diff --git a/.circleci/scripts/unit_selection.sh b/.circleci/scripts/unit_selection.sh index e9e5dd3d66b..3f4f5620176 100755 --- a/.circleci/scripts/unit_selection.sh +++ b/.circleci/scripts/unit_selection.sh @@ -5,8 +5,10 @@ flag="${1:?usage: unit_selection.sh }" legacy_flags=( caching-local + core-utils enterprise-package enterprise-routing + integrations llm-other-providers llm-vertex-ai mcp-integration @@ -31,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 @@ -41,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 @@ -52,6 +57,7 @@ legacy_paths() { echo tests/unit/enterprise/proxy/test_file_deletion_blocking.py echo tests/unit/enterprise/proxy/test_managed_files_access_check.py echo tests/unit/enterprise/proxy/test_managed_files_hook.py ;; + integrations) echo tests/unit/integrations ;; llm-other-providers) find tests/unit/llms -name 'test_*.py' -not -path 'tests/unit/llms/vertex_ai/*' ;; llm-vertex-ai) echo tests/unit/llms/vertex_ai ;; mcp-integration) @@ -75,6 +81,8 @@ 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 ;; proxy-db-auth-checks) @@ -139,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 994d67da64d..a9cd21bad5e 100644 --- a/.circleci/tests.yml +++ b/.circleci/tests.yml @@ -369,6 +369,20 @@ 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 + shards: 2 + reruns: 3 + 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-misc flag: misc 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 2fa05879350..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 @@ -80,7 +79,8 @@ jobs: - shard: integrations artifact-name: integrations - test-path: "tests/test_litellm/integrations" + test-path: "" + unit-flag: integrations workers: 2 reruns: 3 timeout-minutes: 20 @@ -107,11 +107,9 @@ jobs: - shard: misc artifact-name: misc test-path: >- - tests/test_litellm/secret_managers 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 e86047b1987..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 @@ -326,13 +326,13 @@ test-unit-proxy-misc: install-test-deps $(UV_RUN) pytest tests/test_litellm/proxy/_experimental tests/test_litellm/proxy/agent_endpoints tests/test_litellm/proxy/anthropic_endpoints tests/test_litellm/proxy/common_utils tests/test_litellm/proxy/discovery_endpoints tests/test_litellm/proxy/experimental tests/test_litellm/proxy/google_endpoints tests/test_litellm/proxy/health_endpoints tests/test_litellm/proxy/image_endpoints tests/test_litellm/proxy/middleware tests/test_litellm/proxy/openai_files_endpoint tests/test_litellm/proxy/pass_through_endpoints tests/test_litellm/proxy/prompts tests/test_litellm/proxy/public_endpoints tests/test_litellm/proxy/response_api_endpoints tests/test_litellm/proxy/shutdown tests/test_litellm/proxy/spend_tracking tests/test_litellm/proxy/ui_crud_endpoints tests/test_litellm/proxy/vector_store_endpoints tests/test_litellm/proxy/test_*.py --tb=short -vv -n 4 --durations=20 test-unit-integrations: install-test-deps - $(UV_RUN) pytest tests/test_litellm/integrations --tb=short -vv -n 4 --durations=20 + $(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/test_litellm/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/cookbook/litellm_proxy_server/grafana_dashboard/dashboard_all_metrics/readme.md b/cookbook/litellm_proxy_server/grafana_dashboard/dashboard_all_metrics/readme.md index 70562b18aa6..a3869213be2 100644 --- a/cookbook/litellm_proxy_server/grafana_dashboard/dashboard_all_metrics/readme.md +++ b/cookbook/litellm_proxy_server/grafana_dashboard/dashboard_all_metrics/readme.md @@ -2,7 +2,7 @@ Every `litellm_*` metric family the proxy can expose on `/metrics` (136 families across 97 panels), grouped into rows: proxy traffic, latency, spend and tokens, cache, LLM API deployments, key and team rate limits, budgets, guardrails, MCP, managed files and batches, users and teams, the Redis circuit breaker, the spend log cleanup job, and the `prometheus_system` service callback metrics (per-service latency, request and failure rates, spend update queue sizes). Panel titles are the metric names so you can grep the JSON for the metric you care about -Import `grafana_dashboard.json` from **Dashboards > New > Import** and pick your Prometheus data source when prompted (the `DS_PROMETHEUS` variable). Counters are plotted as `rate()` over `$__rate_interval`, histograms as p50 / p95 / p99, gauges as the raw value grouped by the most useful label. Every query names the metric exactly as the proxy emits it (counters carry the `_total` suffix the Prometheus client adds), and `tests/test_litellm/integrations/test_prometheus_metric_name_consistency.py` fails if a metric is renamed without updating this dashboard +Import `grafana_dashboard.json` from **Dashboards > New > Import** and pick your Prometheus data source when prompted (the `DS_PROMETHEUS` variable). Counters are plotted as `rate()` over `$__rate_interval`, histograms as p50 / p95 / p99, gauges as the raw value grouped by the most useful label. Every query names the metric exactly as the proxy emits it (counters carry the `_total` suffix the Prometheus client adds), and `tests/unit/integrations/test_prometheus_metric_name_consistency.py` fails if a metric is renamed without updating this dashboard The first eleven rows need only `callbacks: ["prometheus"]`. The last three rows and the `litellm_admission_*` panels are emitted by other subsystems and stay empty until those are on: the service callback row needs `service_callback: ["prometheus_system"]` in `litellm_settings`, the circuit breaker row needs a Redis cache, the cleanup row needs spend log retention, and admission control needs its middleware enabled. Within the base rows, many panels only fill in once the matching feature is in use: budgets need keys, teams, users or orgs with `max_budget` set, cache panels need caching on, guardrail and MCP panels need those features configured, deployment health needs the router with more than one deployment or a failure to record, and `litellm_in_flight_requests` needs traffic at scrape time. An empty panel for a feature you do not use is expected 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-hashicorp/tests/secret_manager/support.rs b/litellm-rust/crates/secrets-hashicorp/tests/secret_manager/support.rs index 30e46f93248..e1797784b84 100644 --- a/litellm-rust/crates/secrets-hashicorp/tests/secret_manager/support.rs +++ b/litellm-rust/crates/secrets-hashicorp/tests/secret_manager/support.rs @@ -118,7 +118,7 @@ pub(super) struct ParityCase { pub(super) fn parity_cases() -> Vec { serde_json::from_str(include_str!(concat!( env!("CARGO_MANIFEST_DIR"), - "/../../../tests/test_litellm/secret_managers/hashicorp_vault_parity.json" + "/../../../tests/unit/secret_managers/hashicorp_vault_parity.json" ))) .unwrap() } diff --git a/litellm-rust/crates/secrets/PARITY.md b/litellm-rust/crates/secrets/PARITY.md index aeed4ba4b83..9439da93d65 100644 --- a/litellm-rust/crates/secrets/PARITY.md +++ b/litellm-rust/crates/secrets/PARITY.md @@ -32,7 +32,7 @@ The rollout decision comes from [catalog.py](../../../litellm/rust_bridge/catalo `_SecretManagerRuntime` is a private implementation detail, not a replacement SDK class. Its async methods return Futures; public `async def` methods retain lazy coroutine creation and `asyncio.create_task` support. Passing the same names and arguments is insufficient to claim parity until the remaining return-value, error, cache and configuration differences above are closed -## [tests/test_litellm/secret_managers/test_aws_secret_manager_replication.py](../../../tests/test_litellm/secret_managers/test_aws_secret_manager_replication.py) +## [tests/unit/secret_managers/test_aws_secret_manager_replication.py](../../../tests/unit/secret_managers/test_aws_secret_manager_replication.py) | Python test | Rust coverage or boundary | | --- | --- | @@ -48,7 +48,7 @@ The rollout decision comes from [catalog.py](../../../litellm/rust_bridge/catalo | `test_replicate_secret_http_error_raises` | [direct_replication_returns_response_or_service_error](../secrets-aws/tests/secret_manager/writes.rs) | | `test_replicate_secret_timeout_raises` | [write_and_replication_timeouts_remain_errors](../secrets-aws/tests/secret_manager/writes.rs) | -## [tests/test_litellm/secret_managers/test_aws_secret_manager_rotation.py](../../../tests/test_litellm/secret_managers/test_aws_secret_manager_rotation.py) +## [tests/unit/secret_managers/test_aws_secret_manager_rotation.py](../../../tests/unit/secret_managers/test_aws_secret_manager_rotation.py) | Python test | Rust coverage or boundary | | --- | --- | @@ -59,7 +59,7 @@ The rollout decision comes from [catalog.py](../../../litellm/rust_bridge/catalo | `test_write_secret_to_name_inside_recovery_window_restores_and_stores_new_value` | [recovery_window_alias_is_restored_updated_and_tagged](../secrets-aws/tests/secret_manager/writes.rs) | | `test_write_secret_to_live_existing_name_still_fails_without_overwriting` | [create_failure_does_not_overwrite_an_alias_without_a_deletion_date](../secrets-aws/tests/secret_manager/writes.rs) | -## [tests/test_litellm/secret_managers/test_aws_secret_manager_v2.py](../../../tests/test_litellm/secret_managers/test_aws_secret_manager_v2.py) +## [tests/unit/secret_managers/test_aws_secret_manager_v2.py](../../../tests/unit/secret_managers/test_aws_secret_manager_v2.py) | Python test | Rust coverage or boundary | | --- | --- | @@ -70,14 +70,14 @@ The rollout decision comes from [catalog.py](../../../litellm/rust_bridge/catalo | `test_prepare_request_explicit_bedrock_runtime_endpoint_param_still_wins` | [endpoint_overrides_replace_the_service_and_override_the_region](../secrets-aws/tests/secret_manager/configuration.rs) | | `test_prepare_request_env_bedrock_runtime_endpoint_still_wins` | [endpoint_overrides_replace_the_service_and_override_the_region](../secrets-aws/tests/secret_manager/configuration.rs) | -## [tests/test_litellm/secret_managers/test_base_secret_manager.py](../../../tests/test_litellm/secret_managers/test_base_secret_manager.py) +## [tests/unit/secret_managers/test_base_secret_manager.py](../../../tests/unit/secret_managers/test_base_secret_manager.py) | Python test | Rust coverage or boundary | | --- | --- | | `test_raise_if_unsafe_secret_name_rejects_traversal_and_line_breaks` | [names_reject_path_traversal_and_control_characters](../secrets-types/tests/rotation.rs) | | `test_raise_if_unsafe_secret_name_allows_legitimate_aliases` | [names_allow_safe_values](../secrets-types/tests/rotation.rs) | -## [tests/test_litellm/secret_managers/test_custom_secret_manager.py](../../../tests/test_litellm/secret_managers/test_custom_secret_manager.py) +## [tests/unit/secret_managers/test_custom_secret_manager.py](../../../tests/unit/secret_managers/test_custom_secret_manager.py) | Python test | Rust coverage or boundary | | --- | --- | @@ -89,7 +89,7 @@ The rollout decision comes from [catalog.py](../../../litellm/rust_bridge/catalo | `test_custom_secret_manager_integration_with_litellm` | [manager_strings_are_coerced_like_literal_eval](../secrets/tests/resolution.rs) | | `test_minimal_custom_secret_manager` | Exercises the Python example subclass itself or Python default methods. Caller-authored Python implementations remain Python callbacks; resolver integration is covered by `manager_strings_are_coerced_like_literal_eval` | -## [tests/test_litellm/secret_managers/test_cyberark_secret_manager.py](../../../tests/test_litellm/secret_managers/test_cyberark_secret_manager.py) +## [tests/unit/secret_managers/test_cyberark_secret_manager.py](../../../tests/unit/secret_managers/test_cyberark_secret_manager.py) | Python test | Rust coverage or boundary | | --- | --- | @@ -97,7 +97,7 @@ The rollout decision comes from [catalog.py](../../../litellm/rust_bridge/catalo | `test_async_write_matches_parity_fixture` | [writes_match_python_parity_fixture](../secrets-cyberark/tests/secret_manager/writes.rs) | | `test_missing_credentials_raise_value_error` | [new_validates_credentials_before_license_and_configuration](../secrets-cyberark/tests/secret_manager/configuration.rs) | -## [tests/test_litellm/secret_managers/test_get_azure_ad_token_provider.py](../../../tests/test_litellm/secret_managers/test_get_azure_ad_token_provider.py) +## [tests/unit/secret_managers/test_get_azure_ad_token_provider.py](../../../tests/unit/secret_managers/test_get_azure_ad_token_provider.py) | Python test | Rust coverage or boundary | | --- | --- | @@ -115,7 +115,7 @@ The rollout decision comes from [catalog.py](../../../litellm/rust_bridge/catalo | `test_get_azure_ad_token_provider_prefers_workload_identity_over_managed_identity` | Credential-selection contract belongs to `litellm-auth-azure`, not a secrets crate | | `test_get_azure_ad_token_provider_defaults_to_default_azure_credential` | Credential-selection contract belongs to `litellm-auth-azure`, not a secrets crate | -## [tests/test_litellm/secret_managers/test_hashicorp_secret_manager.py](../../../tests/test_litellm/secret_managers/test_hashicorp_secret_manager.py) +## [tests/unit/secret_managers/test_hashicorp_secret_manager.py](../../../tests/unit/secret_managers/test_hashicorp_secret_manager.py) | Python test | Rust coverage or boundary | | --- | --- | @@ -130,13 +130,13 @@ The rollout decision comes from [catalog.py](../../../litellm/rust_bridge/catalo | `test_tls_login_uses_login_namespace` | [tls_login_posts_the_role_and_uses_the_client_identity](../secrets-hashicorp/tests/secret_manager/configuration.rs) | | `test_configuration_matches_native_parity_fixture` | [configuration_matches_python_parity_fixture](../secrets-hashicorp/tests/secret_manager/configuration.rs) | -## [tests/test_litellm/secret_managers/test_secret_manager_handler.py](../../../tests/test_litellm/secret_managers/test_secret_manager_handler.py) +## [tests/unit/secret_managers/test_secret_manager_handler.py](../../../tests/unit/secret_managers/test_secret_manager_handler.py) | Python test | Rust coverage or boundary | | --- | --- | | `test_azure_key_vault_matches_rust_parity_fixture` | [parity_fixture_matches_python_backend_contract](../secrets-azure/tests/key_vault.rs) | -## [tests/test_litellm/secret_managers/test_secret_managers_main.py](../../../tests/test_litellm/secret_managers/test_secret_managers_main.py) +## [tests/unit/secret_managers/test_secret_managers_main.py](../../../tests/unit/secret_managers/test_secret_managers_main.py) | Python test | Rust coverage or boundary | | --- | --- | 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/experimental_mcp_client/client.py b/litellm/experimental_mcp_client/client.py index 01670be74c8..4e3b92edc89 100644 --- a/litellm/experimental_mcp_client/client.py +++ b/litellm/experimental_mcp_client/client.py @@ -10,6 +10,7 @@ import os from collections.abc import AsyncIterator, Awaitable, Callable, Generator, Sequence from contextlib import AbstractAsyncContextManager from functools import partial +from importlib.metadata import version from types import MappingProxyType from typing import Final, TypeAlias, TypeVar, cast @@ -34,8 +35,16 @@ _TransportContext: TypeAlias = AbstractAsyncContextManager[_TransportStreams] from mcp.types import ( METHOD_NOT_FOUND, REQUEST_TIMEOUT, + ClientCapabilities, + ElicitationCapability, + FormElicitationCapability, GetPromptRequestParams, GetPromptResult, + Implementation, + InitializedNotification, + InitializeRequest, + InitializeRequestParams, + InitializeResult, InputRequiredResult, ListPromptsResult, ListResourcesResult, @@ -44,12 +53,14 @@ from mcp.types import ( PaginatedResult, Prompt, ResourceTemplate, + SamplingCapability, ServerNotification, + UrlElicitationCapability, ) from mcp.types import CallToolRequestParams as MCPCallToolRequestParams from mcp.types import CallToolResult as MCPCallToolResult from mcp.types import Tool as MCPTool -from pydantic import AnyUrl +from pydantic import AnyUrl, TypeAdapter from litellm._logging import verbose_logger from litellm.constants import ( @@ -64,11 +75,13 @@ from litellm.proxy._experimental.mcp_server.mcp_debug import capture_upstream_er from litellm.proxy._experimental.mcp_server.result_conversion import error_text_result from litellm.types.llms.custom_http import VerifyTypes from litellm.types.mcp import ( + MCP_LEGACY_VERSIONS, MCPAuth, MCPAuthType, MCPStdioConfig, MCPTransport, MCPTransportType, + MCPUpstreamProtocol, credential_redirect_hook, has_header, without_header, @@ -386,7 +399,9 @@ class MCPClient: sampling_callback: Callable | None = None, elicitation_callback: Callable | None = None, logging_callback: Callable | None = None, + protocol_version: MCPUpstreamProtocol = "auto", ): + self.protocol_version: MCPUpstreamProtocol = TypeAdapter(MCPUpstreamProtocol).validate_python(protocol_version) self.server_url: str = server_url self.transport_type: MCPTransport = transport_type self.auth_type: MCPAuthType = auth_type @@ -525,6 +540,35 @@ class MCPClient: return safe_env + async def _initialize_session(self, session: ClientSession) -> InitializeResult: + if self.protocol_version == "auto": + automatic: Final = await session.initialize() + if automatic.protocol_version not in MCP_LEGACY_VERSIONS: + raise MCPError(code=-32022, message="Upstream selected an unsupported MCP protocol version") + return automatic + result: Final = await session.send_request( + InitializeRequest( + params=InitializeRequestParams( + protocol_version=self.protocol_version, + client_info=Implementation(name="litellm", version=version("litellm")), + capabilities=ClientCapabilities( + sampling=SamplingCapability() if self._sampling_callback is not None else None, + elicitation=ElicitationCapability( + form=FormElicitationCapability(), url=UrlElicitationCapability() + ) + if self._elicitation_callback is not None + else None, + ), + ) + ), + InitializeResult, + ) + if result.protocol_version != self.protocol_version: + raise MCPError(code=-32022, message="Upstream did not accept the configured MCP protocol version") + session.adopt(result) + await session.send_notification(InitializedNotification()) + return result + async def _execute_session_operation( self, transport_ctx: _TransportContext, @@ -579,7 +623,7 @@ class MCPClient: ) session: Final = await session_ctx.__aenter__() try: - init_result: Final = await session.initialize() + init_result: Final = await self._initialize_session(session) instructions: Final = getattr(init_result, "instructions", None) self._last_initialize_instructions = ( instructions.strip() or None if isinstance(instructions, str) else None 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/levo/README.md b/litellm/integrations/levo/README.md index 5296acb7ff4..1fbd202d9a5 100644 --- a/litellm/integrations/levo/README.md +++ b/litellm/integrations/levo/README.md @@ -92,7 +92,7 @@ litellm/integrations/levo/ ## Testing -See the test files in `tests/test_litellm/integrations/levo/`: +See the test files in `tests/unit/integrations/levo/`: - `test_levo.py`: Unit tests for configuration - `test_levo_integration.py`: Integration tests for callback registration 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 4e5643ad02c..b1536382753 100644 --- a/litellm/llms/anthropic/experimental_pass_through/adapters/streaming_iterator.py +++ b/litellm/llms/anthropic/experimental_pass_through/adapters/streaming_iterator.py @@ -70,15 +70,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/integrations/code_interpreter_interception/__init__.py b/litellm/llms/xai/batches/__init__.py similarity index 100% rename from tests/test_litellm/integrations/code_interpreter_interception/__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/integrations/gitlab/__init__.py b/litellm/llms/xai/files/__init__.py similarity index 100% rename from tests/test_litellm/integrations/gitlab/__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 5a207dc4c02..a670200c132 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -51436,13 +51436,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, @@ -51450,8 +51453,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 @@ -51480,9 +51486,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, @@ -51490,6 +51500,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, @@ -51502,9 +51514,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, @@ -51512,6 +51528,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, @@ -59483,13 +59501,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, @@ -59497,20 +59518,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, @@ -59519,8 +59546,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": [ @@ -62787,13 +62817,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, @@ -62801,21 +62834,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, @@ -62823,21 +62862,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, @@ -62845,8 +62890,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 @@ -63067,13 +63115,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, @@ -63081,20 +63132,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, @@ -63102,20 +63159,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" @@ -63127,20 +63190,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" @@ -63152,8 +63221,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 }, @@ -75459,13 +75531,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, @@ -75473,8 +75548,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 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/_experimental/mcp_server/capabilities.py b/litellm/proxy/_experimental/mcp_server/capabilities.py new file mode 100644 index 00000000000..bfd00327eb4 --- /dev/null +++ b/litellm/proxy/_experimental/mcp_server/capabilities.py @@ -0,0 +1,145 @@ +from collections.abc import Callable, Mapping +from dataclasses import dataclass +from itertools import product +from types import MappingProxyType +from typing import Final, Literal + +from mcp.server.context import CallNext, HandlerResult, ServerRequestContext +from mcp.shared.exceptions import MCPError +from mcp.types import DiscoverResult, InitializeRequestParams, InitializeResult, ServerCapabilities +from mcp_types.methods import CLIENT_REQUESTS +from mcp_types.version import HANDSHAKE_PROTOCOL_VERSIONS, LATEST_HANDSHAKE_VERSION +from pydantic import TypeAdapter + +from litellm.types.mcp import MCP_LEGACY_VERSIONS, MCPAdvertisedVersions, MCPLegacyVersion, MCPSpecVersion, MCPTransport + +GATEWAY_OPERATIONS: Final = frozenset( + { + "tools/list", + "tools/call", + "prompts/list", + "prompts/get", + "resources/list", + "resources/read", + "resources/templates/list", + } +) + + +@dataclass(frozen=True, slots=True) +class RevisionSupport: + transports: frozenset[MCPTransport] + operations: frozenset[str] + results: frozenset[Literal["complete", "input_required"]] + extensions: frozenset[str] + completed: bool + + +REVISION_SUPPORT: Final[Mapping[str, RevisionSupport]] = MappingProxyType( + { + version.value: RevisionSupport( + transports=frozenset(MCPTransport) + if version.value in HANDSHAKE_PROTOCOL_VERSIONS + else frozenset({MCPTransport.http, MCPTransport.stdio}), + operations=frozenset(method for method in GATEWAY_OPERATIONS if (method, version.value) in CLIENT_REQUESTS), + results=frozenset({"complete"}) + if version.value in HANDSHAKE_PROTOCOL_VERSIONS + else frozenset({"complete", "input_required"}), + extensions=frozenset(), + completed=version.value in HANDSHAKE_PROTOCOL_VERSIONS, + ) + for version in MCPSpecVersion + } +) +_COMPLETED_REVISIONS: Final = tuple(version for version, support in REVISION_SUPPORT.items() if support.completed) +TRANSLATION_PAIRS: Final = frozenset(product(_COMPLETED_REVISIONS, repeat=2)) +_ADVERTISED_VERSIONS: Final[TypeAdapter[tuple[MCPLegacyVersion, ...]]] = TypeAdapter(MCPAdvertisedVersions) + + +def configured_versions() -> tuple[str, ...]: + from litellm.proxy.proxy_server import general_settings_view + + configured: Final = general_settings_view().get("mcp_advertised_versions") + return _ADVERTISED_VERSIONS.validate_python(MCP_LEGACY_VERSIONS if configured is None else configured) + + +def build_discovery( + *, + configured: tuple[str, ...], + revision: str, + transport: MCPTransport, + authorized_operations: frozenset[str], + upstream_versions: frozenset[str], + capabilities: ServerCapabilities, + client_extensions: frozenset[str] = frozenset(), + upstream_extensions: frozenset[str] = frozenset(), + instructions: str | None = None, +) -> DiscoverResult: + supported: Final = tuple( + version + for version, support in REVISION_SUPPORT.items() + if version in configured and support.completed and transport in support.transports + ) + revision_support: Final = REVISION_SUPPORT.get(revision) + operations: Final[frozenset[str]] = ( + authorized_operations & revision_support.operations + if revision in supported + and revision_support is not None + and any((revision, upstream) in TRANSLATION_PAIRS for upstream in upstream_versions) + else frozenset() + ) + extensions: Final[frozenset[str]] = ( + revision_support.extensions & client_extensions & upstream_extensions + if operations and revision_support is not None + else frozenset() + ) + caller_capabilities: Final = capabilities.model_copy(deep=True) + return DiscoverResult( + supported_versions=list(supported), + capabilities=ServerCapabilities( + tools=caller_capabilities.tools if {"tools/list", "tools/call"} <= operations else None, + prompts=caller_capabilities.prompts if {"prompts/list", "prompts/get"} <= operations else None, + resources=caller_capabilities.resources if {"resources/list", "resources/read"} <= operations else None, + extensions={ + key: value for key, value in (caller_capabilities.extensions or {}).items() if key in extensions + } + or None, + ), + instructions=instructions, + cache_scope="private", + ttl_ms=0, + ) + + +class GatewayVersionPolicy: + def __init__(self, versions: Callable[[], tuple[str, ...]] = configured_versions) -> None: + self._versions = versions + + async def __call__(self, ctx: ServerRequestContext[object, object], call_next: CallNext) -> HandlerResult: + versions: Final = self._versions() + requested: Final = ( + InitializeRequestParams.model_validate(ctx.params or {}).protocol_version + if ctx.method == "initialize" + else ctx.protocol_version + ) + negotiated: Final = ( + (requested if requested in HANDSHAKE_PROTOCOL_VERSIONS else LATEST_HANDSHAKE_VERSION) + if ctx.method == "initialize" + else requested + ) + if negotiated not in versions: + raise MCPError(code=-32022, message="Unsupported MCP protocol version", data={"supported": list(versions)}) + result: Final = await call_next(ctx) + if ctx.method != "initialize": + return result + initialized: Final = InitializeResult.model_validate(result) + discovery: Final = build_discovery( + configured=versions, + revision=initialized.protocol_version, + transport=MCPTransport.http, + authorized_operations=GATEWAY_OPERATIONS, + upstream_versions=frozenset(HANDSHAKE_PROTOCOL_VERSIONS), + capabilities=initialized.capabilities, + instructions=initialized.instructions, + ) + return initialized.model_copy(update={"capabilities": discovery.capabilities}) diff --git a/litellm/proxy/_experimental/mcp_server/contracts.py b/litellm/proxy/_experimental/mcp_server/contracts.py index a88d400282c..a0a08dc08ce 100644 --- a/litellm/proxy/_experimental/mcp_server/contracts.py +++ b/litellm/proxy/_experimental/mcp_server/contracts.py @@ -28,6 +28,7 @@ class OperationContext: client_ip: str | None = None mcp_proxy_mode: bool = False wire_compat: WireCompat = WireCompat.LEGACY + protocol_version: str | None = None def __post_init__(self) -> None: object.__setattr__(self, "_caller", copy_caller(self._caller)) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 24cae976174..6baa695433c 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -193,6 +193,7 @@ from litellm.types.mcp import ( MCPAuth, MCPStdioConfig, MCPTokenEndpointAuthMethod, + MCPUpstreamProtocol, has_header, without_header, ) @@ -340,6 +341,7 @@ class MCPServerConfig(TypedDict, total=False): whatever the admin wrote, and each read applies its own default.""" server_id: ReadOnly[str] + protocol_version: ReadOnly[MCPUpstreamProtocol] alias: str description: str mcp_info: MCPInfo @@ -2549,6 +2551,9 @@ class MCPServerManager: new_server = MCPServer( server_id=server_id, name=name_for_prefix, + protocol_version=TypeAdapter(MCPUpstreamProtocol).validate_python( + server_config.get("protocol_version", mcp_info.get("protocol_version", "auto")) + ), alias=alias, server_name=server_name, spec_path=server_config.get("spec_path", None), @@ -3109,6 +3114,9 @@ class MCPServerManager: new_server: Final = MCPServer( server_id=mcp_server.server_id, name=name_for_prefix, + protocol_version=TypeAdapter(MCPUpstreamProtocol).validate_python( + _mcp_info.get("protocol_version", "auto") + ), alias=getattr(mcp_server, "alias", None), server_name=getattr(mcp_server, "server_name", None), url=mcp_server.url, @@ -4145,6 +4153,7 @@ class MCPServerManager: cred_provider: UpstreamCredentialProvider | None = None, raw_headers: Mapping[str, str] | None = None, client_ip: str | None = None, + protocol_version_override: MCPUpstreamProtocol | None = None, ) -> MCPClient: """ Create an MCPClient instance for the given server. @@ -4168,6 +4177,9 @@ class MCPServerManager: """ record_auth_resolution(server.server_id, AuthResolution.unresolved) resolved_server: Final = await self.ensure_oauth_metadata_discovered(server) + protocol_version: Final = ( + protocol_version_override if protocol_version_override is not None else resolved_server.protocol_version + ) transport: Final = resolved_server.transport or MCPTransport.sse spec = None if transport == MCPTransport.stdio else _to_server_spec_fail_closed(resolved_server) provider: Final = cred_provider or self._cred_provider @@ -4249,6 +4261,7 @@ class MCPServerManager: return MCPClient( server_url="", # Not used for stdio transport_type=transport, + protocol_version=protocol_version, auth_type=resolved_server.auth_type, auth_value=auth_value, timeout=(resolved_server.timeout if resolved_server.timeout is not None else MCP_CLIENT_TIMEOUT), @@ -4281,6 +4294,7 @@ class MCPServerManager: MCPClient( server_url=server_url, transport_type=transport, + protocol_version=protocol_version, auth_type=resolved_server.auth_type, timeout=( resolved_server.timeout if resolved_server.timeout is not None else MCP_CLIENT_TIMEOUT @@ -4324,6 +4338,7 @@ class MCPServerManager: MCPClient( server_url=server_url, transport_type=transport, + protocol_version=protocol_version, auth_type=resolved_server.auth_type, auth_value=auth_value, auth_header_name=auth_header_name, diff --git a/litellm/proxy/_experimental/mcp_server/operations.py b/litellm/proxy/_experimental/mcp_server/operations.py index ebd26e4bf87..a19246b6e90 100644 --- a/litellm/proxy/_experimental/mcp_server/operations.py +++ b/litellm/proxy/_experimental/mcp_server/operations.py @@ -14,6 +14,8 @@ from mcp.types import ( CallToolRequest, CallToolRequestParams, CallToolResult, + DiscoverRequest, + DiscoverResult, GetPromptRequest, GetPromptRequestParams, GetPromptResult, @@ -28,10 +30,14 @@ from mcp.types import ( ListToolsResult, PaginatedRequestParams, Prompt, + PromptsCapability, ReadResourceRequest, ReadResourceRequestParams, + ResourcesCapability, ResourceTemplate, + ServerCapabilities, TextContent, + ToolsCapability, ) from mcp.types import Tool as MCPTool from pydantic import AnyUrl, ConfigDict, Field, TypeAdapter @@ -51,6 +57,11 @@ from litellm.proxy._experimental.mcp_server.byok_credential_cache import ( cache_byok_credential, get_cached_byok_credential, ) +from litellm.proxy._experimental.mcp_server.capabilities import ( + GATEWAY_OPERATIONS, + build_discovery, + configured_versions, +) from litellm.proxy._experimental.mcp_server.contracts import ( AuthorizedToolCall, OperationContext, @@ -122,7 +133,9 @@ from litellm.proxy.litellm_pre_call_utils import ( ) from litellm.types.mcp import ( DEFAULT_CREDENTIAL_HEADER, + MCP_LEGACY_VERSIONS, MCPAuth, + MCPTransport, without_header, ) from litellm.types.mcp_server.mcp_server_manager import MCPInfo, MCPServer @@ -2657,7 +2670,11 @@ class _McpDeniedDetail(TypedDict): async def _execute_handle_list_tools( - context: OperationContext, params: PaginatedRequestParams, host_progress_callback: ProgressCallback | None = None + context: OperationContext, + params: PaginatedRequestParams, + host_progress_callback: ProgressCallback | None = None, + *, + log_list_tools_to_spendlogs: bool = True, ) -> ListToolsResult: try: ( @@ -2700,7 +2717,7 @@ async def _execute_handle_list_tools( mcp_server_auth_headers=mcp_server_auth_headers, oauth2_headers=oauth2_headers, raw_headers=raw_headers, - log_list_tools_to_spendlogs=True, + log_list_tools_to_spendlogs=log_list_tools_to_spendlogs, list_tools_log_source="mcp_protocol", client_ip=_client_ip, ) @@ -3065,6 +3082,7 @@ def prepare_context( client_ip: str | None = None, mcp_proxy_mode: bool = False, wire_compat: WireCompat = WireCompat.LEGACY, + protocol_version: str | None = None, ) -> OperationContext: return OperationContext( _caller=user_api_key_auth, @@ -3076,11 +3094,13 @@ def prepare_context( client_ip=client_ip, mcp_proxy_mode=mcp_proxy_mode, wire_compat=wire_compat, + protocol_version=protocol_version, ) GatewayOperation: TypeAlias = ( AuthorizedToolCall + | DiscoverRequest | ListToolsRequest | CallToolRequest | ListPromptsRequest @@ -3090,7 +3110,8 @@ GatewayOperation: TypeAlias = ( | ReadResourceRequest ) GatewayResult: TypeAlias = ( - ListToolsResult + DiscoverResult + | ListToolsResult | CallToolResult | InputRequiredResult | ListPromptsResult @@ -3105,6 +3126,9 @@ class GatewayOperations: def __init__(self, host_progress_callback: ProgressCallback | None = None) -> None: self._host_progress_callback = host_progress_callback + @overload + async def execute(self, operation: DiscoverRequest, context: OperationContext) -> DiscoverResult: ... + @overload async def execute( self, operation: AuthorizedToolCall, context: OperationContext @@ -3137,6 +3161,51 @@ class GatewayOperations: async def execute(self, operation: GatewayOperation, context: OperationContext) -> GatewayResult: match operation: + case DiscoverRequest(): + listings: Final = ( + () + if context.mcp_proxy_mode + else (ListPromptsRequest(), ListResourcesRequest(), ListResourceTemplatesRequest()) + ) + tasks: Final = ( + asyncio.create_task( + _execute_handle_list_tools( + context, + PaginatedRequestParams(), + self._host_progress_callback, + log_list_tools_to_spendlogs=False, + ) + ), + *(asyncio.create_task(self.execute(listing, context)) for listing in listings), + ) + try: + results: Final = await asyncio.gather(*tasks) + finally: + for task in tasks: + task.cancel() + await asyncio.gather(*tasks, return_exceptions=True) + return build_discovery( + configured=configured_versions(), + revision=context.protocol_version or "2025-11-25", + transport=MCPTransport.http, + authorized_operations=GATEWAY_OPERATIONS, + upstream_versions=frozenset(MCP_LEGACY_VERSIONS), + capabilities=ServerCapabilities( + tools=ToolsCapability() + if any(isinstance(result, ListToolsResult) and result.tools for result in results) + else None, + prompts=PromptsCapability() + if any(isinstance(result, ListPromptsResult) and result.prompts for result in results) + else None, + resources=ResourcesCapability() + if any( + (isinstance(result, ListResourcesResult) and result.resources) + or (isinstance(result, ListResourceTemplatesResult) and result.resource_templates) + for result in results + ) + else None, + ), + ) case AuthorizedToolCall(): auth, token, _servers, server_headers, oauth_headers, headers, _client_ip = context.legacy_auth() return await _execute_mcp_tool( diff --git a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py index 7f519e2c0d9..02694f110b1 100644 --- a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py @@ -1375,7 +1375,16 @@ if MCP_AVAILABLE: and headers.get(MCPRequestHandler.LITELLM_API_KEY_HEADER_NAME_PRIMARY) else None ) - return _StagedServerTest(request=request, mcp_auth_header=mcp_auth_header, oauth2_headers=oauth2_headers) + preview_request: Final = ( + request.model_copy( + update={"mcp_info": {**(request.mcp_info or {}), "protocol_version": saved_server.protocol_version}} + ) + if saved_server is not None and "protocol_version" not in (request.mcp_info or {}) + else request + ) + return _StagedServerTest( + request=preview_request, mcp_auth_header=mcp_auth_header, oauth2_headers=oauth2_headers + ) async def _list_tools_within(client: MCPClient, deadline: float) -> list[MCPTool] | None: with anyio.move_on_after(deadline): @@ -1512,6 +1521,7 @@ if MCP_AVAILABLE: extra_headers=merged_headers, stdio_env=stdio_env, cred_provider=preview_cred_provider, + protocol_version_override=server_model.protocol_version, ) return await operation(client) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 1bd31d971b0..555aebc7434 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -125,12 +125,14 @@ def unsupported_protocol_version(scope: Scope) -> str | None: ``HANDSHAKE_PROTOCOL_VERSIONS`` to the modern single-exchange path, which bypasses litellm's session/auth model, so the ASGI entry rejects it. """ + from litellm.proxy._experimental.mcp_server.capabilities import configured_versions + headers: Final[Iterable[tuple[bytes, bytes]]] = scope.get("headers") or () values: Final = tuple( raw.decode("latin-1").strip() for key, raw in headers if key.lower() == _MCP_PROTOCOL_VERSION_HEADER ) for value in values: - if value and value not in HANDSHAKE_PROTOCOL_VERSIONS: + if value and value not in configured_versions(): return value return None @@ -149,7 +151,10 @@ try: from mcp.server.session import ServerSession as _McpServerSession from mcp.types import ( BlobResourceContents, + DiscoverRequest, + DiscoverResult, GetPromptResult, + RequestParams, ResourceTemplate, TextResourceContents, ) @@ -526,11 +531,11 @@ if MCP_AVAILABLE: PaginatedRequestParams, ReadResourceRequestParams, ) - from mcp_types.version import HANDSHAKE_PROTOCOL_VERSIONS from litellm.proxy._experimental.mcp_server.auth.litellm_auth_handler import ( MCPAuthenticatedUser, ) + from litellm.proxy._experimental.mcp_server.capabilities import GatewayVersionPolicy, configured_versions from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( MCPServerManager, global_mcp_server_manager, @@ -585,6 +590,7 @@ if MCP_AVAILABLE: name=LITELLM_MCP_SERVER_NAME, version=LITELLM_MCP_SERVER_VERSION, ) + server.middleware.append(GatewayVersionPolicy()) server.create_initialization_options = types.MethodType(_gateway_create_initialization_options, server) sse: Final[SseServerTransport] = SseServerTransport("/sse/messages") @@ -830,6 +836,7 @@ if MCP_AVAILABLE: client_ip, _mcp_proxy_mode.get(), wire_compat_for(ctx.protocol_version), + ctx.protocol_version, ) async def handle_list_tools(ctx: ServerRequestContext, params: PaginatedRequestParams) -> ListToolsResult: @@ -948,6 +955,11 @@ if MCP_AVAILABLE: ReadResourceRequest(params=params), context ) + async def discover(ctx: ServerRequestContext, params: RequestParams) -> DiscoverResult: + async with _legacy_operation_context(ctx, trace=False) as context: + return await operations.GatewayOperations().execute(DiscoverRequest(params=params), context) + + server.add_request_handler("server/discover", RequestParams, discover) server.add_request_handler("tools/list", PaginatedRequestParams, handle_list_tools) server.add_request_handler("tools/call", CallToolRequestParams, mcp_server_tool_call) server.add_request_handler("prompts/list", PaginatedRequestParams, list_prompts) @@ -1954,7 +1966,7 @@ if MCP_AVAILABLE: reject_disallowed_mcp_origin(StarletteRequest(scope)) bad_version: Final = unsupported_protocol_version(scope) if bad_version is not None: - supported: Final = ", ".join(sorted(HANDSHAKE_PROTOCOL_VERSIONS)) + supported: Final = ", ".join(configured_versions()) await JSONResponse( status_code=400, content={ # mutable-ok: JSON-RPC error payload @@ -2299,7 +2311,7 @@ if MCP_AVAILABLE: reject_disallowed_mcp_origin(StarletteRequest(scope)) bad_version: Final = unsupported_protocol_version(scope) if bad_version is not None: - supported: Final = ", ".join(sorted(HANDSHAKE_PROTOCOL_VERSIONS)) + supported: Final = ", ".join(configured_versions()) await JSONResponse( status_code=400, content={ # mutable-ok: JSON-RPC error payload diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 3da30070b9e..4597872d84e 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -37,6 +37,7 @@ from litellm.types.llms.openai import ( ResponsesAPIResponse, ) from litellm.types.mcp import ( + MCPAdvertisedVersions, MCPAllowedClient, MCPAuth, MCPAuthType, @@ -2998,6 +2999,11 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase): None, description="Custom CIDR ranges that define internal/private networks for MCP access control. When set, only these ranges are treated as internal. Defaults to RFC 1918 private ranges (10.0.0.0/8, 172.16.0.0/12, 192.168.0.0/16, 127.0.0.0/8).", ) + mcp_advertised_versions: MCPAdvertisedVersions | None = Field( + None, + description="MCP revisions enabled by the gateway. Defaults to all completed legacy revisions. " + "Modern protocol serving and Apps/Tasks remain disabled.", + ) mcp_allowed_clients: list[MCPAllowedClient] | None = Field( None, description="MCP client applications admitted by the gateway, each an {alias, value} pair where alias is the name shown in the dashboard and logs and value is the identity that must match exactly. When set, every MCP request must carry a client identity equal to one of the values: a JWT caller is identified by the claim named in litellm_jwtauth.mcp_client_id_jwt_field, any other caller by the header named in mcp_client_id_header. A request with no resolvable identity, or an unlisted one, is rejected with 403. Unset means every client is admitted.", @@ -4013,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`.""" @@ -5349,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 7e1989b0aab..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) @@ -2781,15 +2796,6 @@ class ProxyBaseLLMRequestProcessing: request=request, ) if route_type == "aresponses": - # Streaming /v1/responses returns here without - # reaching the non-streaming ownership tail below. - # Wrap the SSE generator so container ownership is - # written once the upstream iterator finishes - # assembling ``completed_response`` — otherwise - # code-interpreter containers created during the - # stream stay unregistered and follow-up file API - # calls 403. Covers the background-polling path - # too, which loops ``body_iterator`` end-to-end. selected_data_generator = ( ProxyBaseLLMRequestProcessing._wrap_responses_stream_for_container_ownership( original_stream_response=response, @@ -3011,50 +3017,50 @@ class ProxyBaseLLMRequestProcessing: wrapped_generator: Any, user_api_key_dict: UserAPIKeyAuth, ): - """Forward SSE chunks, then record container ownership at stream end. + """Forward SSE chunks and record container ownership before the terminal chunk goes out. Streaming ``/v1/responses`` short-circuits out of ``base_process_llm_request`` before the non-streaming ownership - tail runs, so without this wrap the - ``LiteLLM_ManagedObjectTable`` row for any container created - during the stream is never written and follow-up file API calls - return 403. + tail runs. The OpenAI SDK closes the connection at ``data: [DONE]`` + and starlette cancels the body task on disconnect, so a write that + waits for the generator to finish never lands. The iterator sets + ``completed_response`` before it hands over its terminal chunk, so + the ``LiteLLM_ManagedObjectTable`` row is written the moment it + appears, ahead of the chunk carrying ``response.completed``. """ - try: - async for chunk in wrapped_generator: + async for chunk in wrapped_generator: + completed_obj = ProxyBaseLLMRequestProcessing._extract_completed_responses_response( + original_stream_response + ) + if completed_obj is None: yield chunk - finally: - try: - completed_obj: Final = ProxyBaseLLMRequestProcessing._extract_completed_responses_response( - original_stream_response - ) - if completed_obj is not None: - await ProxyBaseLLMRequestProcessing._record_container_owners_from_responses_if_needed( - response=completed_obj, - user_api_key_dict=user_api_key_dict, - ) - else: - # Silent skip caused #30210: the proxy's Router wrapper - # of the responses streaming iterator wasn't propagating - # ``completed_response``, so this hook recorded nothing - # and follow-up /v1/containers//files calls 403'd - # for non-admin keys with no proxy-side hint. Log a - # warning so future regressions of the same shape - # surface in operator logs. - verbose_proxy_logger.warning( - "Container ownership recording skipped on streaming " - "/v1/responses: no completed_response on stream " - "iterator %s. If this stream created any tool " - "container (e.g. code_interpreter), follow-up " - "/v1/containers//files calls will 403 for " - "non-admin keys.", - type(original_stream_response).__name__, - ) - except Exception as e: - verbose_proxy_logger.exception( - "Container ownership recording failed after streaming responses call: %s", - e, - ) + continue + await ProxyBaseLLMRequestProcessing._record_container_owners_from_responses_if_needed( + response=completed_obj, + user_api_key_dict=user_api_key_dict, + ) + yield chunk + async for remaining_chunk in wrapped_generator: + yield remaining_chunk + return + late_completed_obj: Final = ProxyBaseLLMRequestProcessing._extract_completed_responses_response( + original_stream_response + ) + if late_completed_obj is not None: + await ProxyBaseLLMRequestProcessing._record_container_owners_from_responses_if_needed( + response=late_completed_obj, + user_api_key_dict=user_api_key_dict, + ) + return + verbose_proxy_logger.warning( + "Container ownership recording skipped on streaming " + "/v1/responses: no completed_response on stream " + "iterator %s. If this stream created any tool " + "container (e.g. code_interpreter), follow-up " + "/v1/containers//files calls will 403 for " + "non-admin keys.", + type(original_stream_response).__name__, + ) async def base_passthrough_process_llm_request( self, @@ -3861,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, @@ -3870,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 @@ -3886,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( @@ -3931,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) @@ -3979,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, @@ -4001,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 @@ -4019,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 f8d9c9c7be9..11e2107a02b 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) @@ -6342,6 +6340,11 @@ class ProxyConfig: if general_settings is None: general_settings = {} + if general_settings.get("mcp_advertised_versions") is not None: + from litellm.types.mcp import MCPAdvertisedVersions + + TypeAdapter(MCPAdvertisedVersions).validate_python(general_settings["mcp_advertised_versions"]) + if os.getenv("NUM_WORKERS", "1") != "1" and redis_usage_cache is None: warn_login_counters_are_per_worker(os.getenv("NUM_WORKERS", "1")) if declared_proxy_ranges(general_settings) is None: 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/mcp.py b/litellm/types/mcp.py index 83e719810d5..fec5e84c8df 100644 --- a/litellm/types/mcp.py +++ b/litellm/types/mcp.py @@ -4,7 +4,7 @@ import enum import re from collections.abc import Awaitable, Callable, Mapping from types import MappingProxyType -from typing import TYPE_CHECKING, Any, Final, Literal +from typing import TYPE_CHECKING, Annotated, Any, Final, Literal from urllib.parse import urlsplit import httpx @@ -34,6 +34,8 @@ class MCPSpecVersion(str, enum.Enum): nov_2024 = "2024-11-05" mar_2025 = "2025-03-26" jun_2025 = "2025-06-18" + nov_2025 = "2025-11-25" + jul_2026 = "2026-07-28" class MCPAuth(str, enum.Enum): @@ -59,7 +61,17 @@ DEFAULT_SUBJECT_TOKEN_TYPE: Final = "urn:ietf:params:oauth:token-type:access_tok # MCP Literals MCPTransportType = Literal[MCPTransport.sse, MCPTransport.http, MCPTransport.stdio] -MCPSpecVersionType = Literal[MCPSpecVersion.nov_2024, MCPSpecVersion.mar_2025, MCPSpecVersion.jun_2025] +MCPLegacyVersion = Literal["2024-11-05", "2025-03-26", "2025-06-18", "2025-11-25"] +MCP_LEGACY_VERSIONS: Final[tuple[MCPLegacyVersion, ...]] = ("2024-11-05", "2025-03-26", "2025-06-18", "2025-11-25") +MCPUpstreamProtocol = MCPLegacyVersion | Literal["auto"] +MCPAdvertisedVersions = Annotated[tuple[MCPLegacyVersion, ...], Field(min_length=1)] +MCPSpecVersionType = Literal[ + MCPSpecVersion.nov_2024, + MCPSpecVersion.mar_2025, + MCPSpecVersion.jun_2025, + MCPSpecVersion.nov_2025, + MCPSpecVersion.jul_2026, +] MCPAuthType = ( Literal[ MCPAuth.none, diff --git a/litellm/types/mcp_server/mcp_server_manager.py b/litellm/types/mcp_server/mcp_server_manager.py index cb32299b143..c3b106c11d5 100644 --- a/litellm/types/mcp_server/mcp_server_manager.py +++ b/litellm/types/mcp_server/mcp_server_manager.py @@ -1,7 +1,7 @@ from datetime import datetime -from typing import Any, Final, Literal +from typing import Annotated, Any, Final, Literal -from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator +from pydantic import AfterValidator, BaseModel, ConfigDict, Field, TypeAdapter, field_validator, model_validator from typing_extensions import Self from litellm.types.mcp import ( @@ -10,11 +10,19 @@ from litellm.types.mcp import ( MCPAuthType, MCPTokenEndpointAuthMethod, MCPTransportType, + MCPUpstreamProtocol, normalize_upstream_header_name, ) + # MCPInfo now allows arbitrary additional fields for custom metadata -MCPInfo = dict[str, Any] +def _validate_mcp_protocol_metadata(value: dict[str, object]) -> dict[str, object]: + if "protocol_version" in value: + TypeAdapter(MCPUpstreamProtocol).validate_python(value["protocol_version"]) + return value + + +MCPInfo = Annotated[dict[str, Any], AfterValidator(_validate_mcp_protocol_metadata)] class MCPOAuthMetadata(BaseModel): @@ -66,6 +74,7 @@ class MCPServer(BaseModel): server_name: str | None = None url: str | None = None transport: MCPTransportType + protocol_version: MCPUpstreamProtocol = "auto" spec_path: str | None = None auth_type: MCPAuthType | None = None authentication_token: str | None = None @@ -246,6 +255,14 @@ class MCPServer(BaseModel): """ return self.oauth2_flow == "client_credentials" + @model_validator(mode="after") + def resolve_protocol_version(self) -> Self: + if "protocol_version" not in self.model_fields_set and self.mcp_info is not None: + self.protocol_version = TypeAdapter(MCPUpstreamProtocol).validate_python( + self.mcp_info.get("protocol_version", "auto") + ) + return self + @model_validator(mode="after") def validate_identity_binding_mode(self) -> Self: binding: Final = self.oauth_identity_binding diff --git a/litellm/types/router.py b/litellm/types/router.py index b72809f625f..d545f7ae639 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 3e306b48887..cd336c9b989 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 @@ -3729,6 +3732,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 @@ -3742,6 +3746,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 @@ -3766,6 +3771,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 @@ -4140,7 +4146,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 4ea0769ea11..be4388802f9 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" ), @@ -9357,6 +9366,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 5a207dc4c02..a670200c132 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -51436,13 +51436,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, @@ -51450,8 +51453,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 @@ -51480,9 +51486,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, @@ -51490,6 +51500,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, @@ -51502,9 +51514,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, @@ -51512,6 +51528,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, @@ -59483,13 +59501,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, @@ -59497,20 +59518,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, @@ -59519,8 +59546,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": [ @@ -62787,13 +62817,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, @@ -62801,21 +62834,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, @@ -62823,21 +62862,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, @@ -62845,8 +62890,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 @@ -63067,13 +63115,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, @@ -63081,20 +63132,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, @@ -63102,20 +63159,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" @@ -63127,20 +63190,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" @@ -63152,8 +63221,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 }, @@ -75459,13 +75531,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, @@ -75473,8 +75548,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 diff --git a/model_prices_and_context_window.schema.json b/model_prices_and_context_window.schema.json index 35624045fdf..fa1828c780a 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, @@ -351,6 +356,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, @@ -708,6 +718,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 7bdf5b9c0e0..3ad6189bd86 100644 --- a/tests/e2e/coverage_registry/llm_conversational.yaml +++ b/tests/e2e/coverage_registry/llm_conversational.yaml @@ -86,6 +86,7 @@ - {id: llm.responses.azure_openai.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: responses, route: azure_openai, capability: basic, streaming: nonstream, assertions: [works], source: "response_api_endpoints/endpoints.py:26", rationale: "Responses w/ Azure OpenAI (smoke)"} - {id: llm.responses.azure_openai.tool_use.nonstream.works, module: llm, tier: P1, subject_endpoint: responses, route: azure_openai, capability: tool_use, streaming: nonstream, assertions: [works], source: "model_prices json", rationale: "Responses tool calls w/ Azure OpenAI"} - {id: llm.responses.azure_openai.code_interpreter.nonstream.works, module: llm, tier: P1, subject_endpoint: responses, route: azure_openai, capability: code_interpreter, streaming: nonstream, assertions: [works], source: "llm_translation/test_containers_e2e.py", rationale: "An implicit code_interpreter container on an Azure deployment that carries its own api_base must serve GET /v1/containers/{id}/files/{fid}/content by its native cntr_ id to a team service-account key, the customer's shape (#27921, #28990)", fail_before_fix: proven} +- {id: llm.responses.azure_openai.code_interpreter.stream.works, module: llm, tier: P1, subject_endpoint: responses, route: azure_openai, capability: code_interpreter, streaming: stream, assertions: [works], source: "llm_translation/test_containers_e2e.py", rationale: "A container created by a streamed /v1/responses code_interpreter call must serve /v1/containers/{id}/files to the same service-account key right after the OpenAI SDK closes at [DONE]; the ownership row used to be written after the stream and the disconnect cancelled it (LIT-8612)", fail_before_fix: proven} - {id: llm.chat_completions.together_ai.thinking.nonstream.works, module: llm, tier: P1, subject_endpoint: chat_completions, route: together_ai, capability: thinking, streaming: nonstream, assertions: [works], source: "llm_translation/test_together_ai_e2e.py", rationale: "Together reasoning surfaces as reasoning_content (LIT-5960)"} - {id: llm.chat_completions.together_ai.thinking.stream.works, module: llm, tier: P1, subject_endpoint: chat_completions, route: together_ai, capability: thinking, streaming: stream, assertions: [works], source: "llm_translation/test_together_ai_e2e.py", rationale: "Together reasoning deltas stream as reasoning_content"} - {id: llm.chat_completions.together_ai.thinking.nonstream.template_kwargs_forwarded, module: llm, tier: P1, subject_endpoint: chat_completions, route: together_ai, capability: thinking, streaming: nonstream, assertions: [template_kwargs_forwarded], source: "llm_translation/test_together_ai_e2e.py", rationale: "chat_template_kwargs reaches Together and turns thinking off"} @@ -103,6 +104,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/LLM_TRANSLATION_COVERAGE_MATRIX.md b/tests/e2e/llm_translation/LLM_TRANSLATION_COVERAGE_MATRIX.md index a18c81fa01d..92330db0530 100644 --- a/tests/e2e/llm_translation/LLM_TRANSLATION_COVERAGE_MATRIX.md +++ b/tests/e2e/llm_translation/LLM_TRANSLATION_COVERAGE_MATRIX.md @@ -48,7 +48,7 @@ most likely to silently break and the one a mock can't prove works. |----------|---------------|-----------|------------|-------------|--------| | Chat | live (spend suite) | live (spend suite) | gap | live | partial | | Embeddings | live (spend suite) | n/a | n/a | live | covered | -| Responses (Azure code_interpreter container files) | live | gap | live | gap | partial | +| Responses (Azure code_interpreter container files) | live | live | live | gap | partial | | Image / audio / rerank / realtime | - | - | - | - | gap | ## This suite's files @@ -63,6 +63,7 @@ most likely to silently break and the one a mock can't prove works. | `test_anthropic_passthrough_tool_call_logs_cost` | anthropic native, tool call, cost | | `test_vertex_passthrough_via_managed_model_logs_cost` | vertex_ai native, non-stream, cost | | `test_service_account_key_reads_container_file_by_native_id` | azure responses code_interpreter, non-stream, native container id, service-account key | +| `test_service_account_key_reads_container_file_created_by_a_streamed_response` | azure responses code_interpreter, stream, native container id, service-account key, upload right after `[DONE]` | Vertex keeps the credential on the proxy like gemini/anthropic, but the deployment is added at runtime instead of declared in the gateway config: the test POSTs `/model/new` diff --git a/tests/e2e/llm_translation/test_containers_e2e.py b/tests/e2e/llm_translation/test_containers_e2e.py index 887aecb8df1..3048a830810 100644 --- a/tests/e2e/llm_translation/test_containers_e2e.py +++ b/tests/e2e/llm_translation/test_containers_e2e.py @@ -31,10 +31,11 @@ A proxy whose env carries ``AZURE_API_BASE`` for the same resource masks the second regression, since the global-credential fallback then reaches the container anyway. -The streaming variant is not here: a streamed ``/v1/responses`` writes the -container ownership row only after the ``[DONE]`` frame, and the OpenAI SDK -closes the connection at ``[DONE]``, so the write is cancelled and every -follow-up container call 403s (LIT-8612). That cell comes with its fix. +The streaming cell repeats the flow with ``stream=True`` and uploads right +after the last event. The OpenAI SDK closes the connection at ``[DONE]``, so an +ownership row written after the stream is cancelled with the body task and every +follow-up container call 403s (LIT-8612); the row has to land before the +``response.completed`` frame goes out. """ from __future__ import annotations @@ -52,7 +53,7 @@ from lifecycle import ResourceManager from management.management_client import ManagementClient, build_client from models import KeyGenerateBody, KeyGenerateResponse, LiteLLMParamsBody, TeamNewBody, UserNewBody from openai import OpenAI -from openai.types.responses import Response, ResponseCodeInterpreterToolCall +from openai.types.responses import Response, ResponseCodeInterpreterToolCall, ResponseCompletedEvent from openai.types.responses.tool_param import CodeInterpreter from proxy_client import ProxyClient from sdk_clients import NO_PROXY_CACHE, SdkClients @@ -120,6 +121,24 @@ def _response_with_code_interpreter(client: OpenAI, model: str) -> Response: ) +def _streamed_response_with_code_interpreter(client: OpenAI, model: str) -> Response: + events: Final = tuple( + client.with_options(timeout=CODE_INTERPRETER_TIMEOUT).responses.create( + model=model, + input=PROMPT, + tools=[CODE_INTERPRETER], + tool_choice="required", + stream=True, + extra_body=NO_PROXY_CACHE, + ) + ) + assert events, "responses stream returned no events" + assert isinstance(events[-1], ResponseCompletedEvent), ( + f"responses stream did not terminate with response.completed: {events[-1].type}" + ) + return events[-1].response + + def _container_id(response: Response) -> str: calls: Final = tuple(item for item in response.output if isinstance(item, ResponseCodeInterpreterToolCall)) assert calls, f"no code_interpreter_call in the responses output: {response.output!r}" @@ -165,3 +184,17 @@ class TestAzureContainerFiles: f"container id is not the provider's own id: {native_id}" ) _assert_file_round_trip(client, native_id, marker) + + @pytest.mark.covers("llm.responses.azure_openai.code_interpreter.stream.works") + def test_service_account_key_reads_container_file_created_by_a_streamed_response( + self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients + ) -> None: + marker: Final = unique_marker() + model: Final = _register_two_azure_deployments(proxy, resources, marker) + key: Final = _service_account_key(proxy, resources, build_client(proxy), marker, model) + client: Final = sdk.openai(key) + native_id: Final = _native_container_id( + _container_id(_streamed_response_with_code_interpreter(client, model)) + ) + resources.defer(lambda: client.containers.delete(native_id, extra_query=AZURE_PROVIDER_QUERY)) + _assert_file_round_trip(client, native_id, marker) 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 7f27fd28207..dac863552e6 100644 --- a/tests/e2e/models.py +++ b/tests/e2e/models.py @@ -609,6 +609,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/integration/mcp/test_mcp_protocol_errors.py b/tests/integration/mcp/test_mcp_protocol_errors.py index bb06d8c6068..257e90e9d0e 100644 --- a/tests/integration/mcp/test_mcp_protocol_errors.py +++ b/tests/integration/mcp/test_mcp_protocol_errors.py @@ -84,3 +84,42 @@ def test_jsonrpc_error_and_malformed_tool_result_remain_errors(gateway: Gateway) control: Final = call_tool(gateway, key, identity, names["add"], {"a": 3, "b": 5}) assert control.status_code == 200 and control.json()["isError"] is False, control.text assert control.json()["content"][0]["text"] == "8" + + +@pytest.mark.parametrize("ingress", ("http", "sse")) +def test_configured_revision_blocks_unadvertised_handshake_and_keeps_allowed_control( + gateway: Gateway, tmp_path, ingress: str +) -> None: + import asyncio + from pathlib import Path + + import yaml + from integration._support.mcp import mcp_peer + from integration._support.process import owned_proxy + from litellm.experimental_mcp_client.client import MCPClient + from litellm.types.mcp import MCPTransport + from mcp import MCPError + from mcp.types import CallToolRequestParams + + with mcp_peer() as upstream, gateway.scenario() as scenario: + alias: Final = "restricted" + uuid.uuid4().hex[:8] + identity: Final = register_mcp(scenario, upstream, alias) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["general_settings"]["mcp_advertised_versions"] = ["2024-11-05"] + config_path: Final = tmp_path / "restricted.yaml" + config_path.write_text(yaml.safe_dump(config)) + with owned_proxy(gateway, tmp_path, {"DISABLE_SCHEMA_UPDATE": "true"}, config=config_path) as restricted: + endpoint: Final = str(restricted.client.base_url).rstrip("/") + ("/mcp/sse" if ingress == "sse" else "/mcp") + headers: Final = {"Authorization": f"Bearer {key}", "x-mcp-servers": identity} + denied: Final = MCPClient(server_url=endpoint, transport_type=MCPTransport(ingress), protocol_version="2025-11-25", extra_headers=headers) + allowed: Final = MCPClient(server_url=endpoint, transport_type=MCPTransport(ingress), protocol_version="2024-11-05", extra_headers=headers) + + async def exercise() -> None: + with pytest.raises(MCPError, match="Unsupported MCP protocol version"): + await denied.list_tools(raise_on_error=True) + assert f"{alias}-add" in tuple(tool.name for tool in await allowed.list_tools(raise_on_error=True)) + result: Final = await allowed.call_tool(CallToolRequestParams(name=f"{alias}-add", arguments={"a": 2, "b": 5})) + assert result.is_error is False and result.content[0].text == "7" + + asyncio.run(exercise()) diff --git a/tests/integration/mcp/test_mcp_transports.py b/tests/integration/mcp/test_mcp_transports.py index 1862f11d07e..16004cdd501 100644 --- a/tests/integration/mcp/test_mcp_transports.py +++ b/tests/integration/mcp/test_mcp_transports.py @@ -155,3 +155,43 @@ def test_server_initiated_sampling_and_elicitation_surface_as_errors_not_success assert outcome.error is not None, outcome.raw assert outcome.text is None or not outcome.text.startswith(("sampled:", "elicited:")), outcome.raw assert len(tool_calls(peer.drain())) == 1 + + +@pytest.mark.parametrize("downstream", ("2024-11-05", "2025-03-26", "2025-06-18", "2025-11-25")) +@pytest.mark.parametrize("upstream", ("2024-11-05", "2025-03-26", "2025-06-18", "2025-11-25")) +@pytest.mark.parametrize("peer_kind", ("http", "sse", "stdio")) +@pytest.mark.parametrize("ingress", ("http", "sse")) +def test_pinned_revision_pairs_list_and_call_through_gateway( + gateway: Gateway, downstream: str, upstream: str, peer_kind: PeerKind, ingress: str +) -> None: + import asyncio + + from mcp.types import CallToolRequestParams + + from litellm.experimental_mcp_client.client import MCPClient + from litellm.types.mcp import MCPTransport + + with peer_of(peer_kind) as peer, gateway.scenario() as scenario: + alias: Final = "versions" + uuid.uuid4().hex[:8] + identity: Final = register_mcp(scenario, peer, alias, mcp_info={"protocol_version": upstream}) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + endpoint: Final = str(gateway.client.base_url).rstrip("/") + ("/mcp/sse" if ingress == "sse" else "/mcp") + client: Final = MCPClient( + server_url=endpoint, transport_type=MCPTransport(ingress), protocol_version=downstream, + extra_headers={"Authorization": f"Bearer {key}", "x-mcp-servers": identity}, timeout=15, + ) + + async def exercise() -> None: + tools: Final = await client.list_tools(raise_on_error=True) + assert f"{alias}-add" in tuple(tool.name for tool in tools) + result: Final = await client.call_tool(CallToolRequestParams(name=f"{alias}-add", arguments={"a": 3, "b": 4})) + assert result.is_error is False + assert result.content[0].text == "7" + + peer.drain() + asyncio.run(exercise()) + observed: Final = peer.drain() + negotiations: Final = tuple(item["body"] for item in observed if item["body"].get("method") == "initialize") + assert negotiations, "The operation must reach the upstream negotiation" + assert all(request["params"]["protocolVersion"] == upstream for request in negotiations), negotiations + assert len(tool_calls(observed)) == 1 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/integrations/levo/__init__.py b/tests/test_litellm/integrations/levo/__init__.py deleted file mode 100644 index 1560e78b7b9..00000000000 --- a/tests/test_litellm/integrations/levo/__init__.py +++ /dev/null @@ -1 +0,0 @@ -# Levo integration tests 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 eccf44a1bda..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,442 +1,6 @@ -#### What this tests #### -# This tests litellm.token_counter.token_counter() function -import asyncio -import base64 -import importlib -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, - high_detail_image_token_upper_bound, - 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(): @@ -449,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}" @@ -485,1117 +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() 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/_experimental/mcp_server/test_capabilities.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_capabilities.py new file mode 100644 index 00000000000..f104e655704 --- /dev/null +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_capabilities.py @@ -0,0 +1,106 @@ +from typing import Final + +import pytest +from mcp import Client +from mcp.server import Server +from mcp.shared.exceptions import MCPError +from mcp.types import PromptsCapability, ResourcesCapability, ServerCapabilities, ToolsCapability +from mcp_types.version import HANDSHAKE_PROTOCOL_VERSIONS + +from litellm.proxy._experimental.mcp_server.capabilities import ( + GATEWAY_OPERATIONS, + REVISION_SUPPORT, + TRANSLATION_PAIRS, + GatewayVersionPolicy, + build_discovery, +) +from litellm.types.mcp import MCPTransport + + +@pytest.mark.parametrize("revision", HANDSHAKE_PROTOCOL_VERSIONS) +@pytest.mark.parametrize("transport", tuple(MCPTransport)) +def test_discovery_only_exposes_authorized_completed_support(revision, transport): + result = build_discovery( + configured=(revision, "2026-07-28", "unknown"), + revision=revision, + transport=transport, + authorized_operations=frozenset({"tools/list", "tools/call"}), + upstream_versions=frozenset(HANDSHAKE_PROTOCOL_VERSIONS), + capabilities=ServerCapabilities( + tools=ToolsCapability(), prompts=PromptsCapability(), resources=ResourcesCapability(), + extensions={"io.modelcontextprotocol/ui": {}}, + ), + client_extensions=frozenset({"io.modelcontextprotocol/ui"}), + upstream_extensions=frozenset({"io.modelcontextprotocol/ui"}), + ) + assert result.supported_versions == [revision] + assert result.capabilities.tools is not None + assert result.capabilities.prompts is None + assert result.capabilities.resources is None + assert result.capabilities.extensions is None + assert result.capabilities.tasks is None + assert result.cache_scope == "private" + assert result.ttl_ms == 0 + + +@pytest.mark.parametrize("upstream", [frozenset(), frozenset({"unknown"}), frozenset({"2026-07-28"})]) +def test_unproven_translation_never_advertises_operations(upstream): + result = build_discovery( + configured=HANDSHAKE_PROTOCOL_VERSIONS, + revision="2025-11-25", + transport=MCPTransport.http, + authorized_operations=GATEWAY_OPERATIONS, + upstream_versions=upstream, + capabilities=ServerCapabilities(tools=ToolsCapability()), + ) + assert result.capabilities.tools is None + + +@pytest.mark.parametrize("revision", ["2026-07-28", "unknown", "2024-11-05"]) +def test_unadvertised_revision_never_gains_capabilities(revision): + result = build_discovery( + configured=("2025-11-25",), revision=revision, transport=MCPTransport.http, + authorized_operations=GATEWAY_OPERATIONS, upstream_versions=frozenset(HANDSHAKE_PROTOCOL_VERSIONS), + capabilities=ServerCapabilities(tools=ToolsCapability()), + ) + assert result.capabilities.tools is None + + +def test_discovery_results_do_not_share_mutable_capabilities(): + capabilities = ServerCapabilities(tools=ToolsCapability(), prompts=PromptsCapability(), resources=ResourcesCapability()) + args = dict( + configured=HANDSHAKE_PROTOCOL_VERSIONS, revision="2025-11-25", transport=MCPTransport.http, + upstream_versions=frozenset(HANDSHAKE_PROTOCOL_VERSIONS), capabilities=capabilities, + ) + allowed = build_discovery(**args, authorized_operations=GATEWAY_OPERATIONS) + denied = build_discovery(**args, authorized_operations=frozenset()) + assert allowed.capabilities.prompts is not None + assert allowed.capabilities.resources is not None + assert denied.capabilities.model_dump(exclude_none=True) == {} + assert allowed.capabilities.tools is not None + allowed.capabilities.tools.list_changed = True + assert capabilities.tools.list_changed is not True + + +def test_modern_candidates_do_not_enable_public_serving(): + modern = REVISION_SUPPORT["2026-07-28"] + assert modern.completed is False + assert "input_required" in modern.results + assert MCPTransport.sse not in modern.transports + assert not any("2026-07-28" in pair for pair in TRANSLATION_PAIRS) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("versions,accepted", [(("2025-11-25",), True), (("2025-06-18",), False)]) +async def test_version_policy_gates_the_actual_sdk_handshake(versions, accepted): + server: Final = Server("test-gateway", version="1") + server.middleware.append(GatewayVersionPolicy(lambda: versions)) + if accepted: + async with Client(server, mode="legacy") as client: + assert client.protocol_version == "2025-11-25" + result = await client.session.send_ping() + assert result is not None + else: + with pytest.RaisesGroup(pytest.RaisesExc(MCPError, match="Unsupported MCP protocol version"), flatten_subgroups=True): + async with Client(server, mode="legacy"): + pytest.fail("The excluded revision must not initialize") diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index 97b242831a2..ab00ec4da1e 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -10444,11 +10444,12 @@ async def test_active_request_ctx_var_feeds_auth_resolution_recording(_mcp_reque ) @pytest.mark.parametrize("handler", ("handle_streamable_http_mcp", "handle_sse_mcp")) async def test_streamable_http_rejects_modern_protocol_version( - header_value: str, expected_rejected: bool, handler: str + header_value: str, expected_rejected: bool, handler: str, monkeypatch: pytest.MonkeyPatch ) -> None: from litellm.proxy._experimental.mcp_server import server as mcp_module from litellm.proxy._experimental.mcp_server.server import unsupported_protocol_version + monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {}) scope: Scope = { "type": "http", "method": "POST", @@ -10659,3 +10660,32 @@ async def test_legacy_sse_mount_emits_message_endpoint( await incoming.put({"type": "http.disconnect"}) await asyncio.wait_for(task, 2) assert await post(initialization) == 404 + + +@pytest.mark.parametrize("revision,rejected", [("2024-11-05", False), ("2025-11-25", True), ("2026-07-28", True)]) +def test_protocol_header_respects_configured_advertisement(revision, rejected): + from litellm.proxy import proxy_server + from litellm.proxy._experimental.mcp_server.server import unsupported_protocol_version + + with patch.dict(proxy_server.general_settings, {"mcp_advertised_versions": ["2024-11-05"]}): + result = unsupported_protocol_version({"headers": [(b"mcp-protocol-version", revision.encode())]}) + assert result == (revision if rejected else None) + + +@pytest.mark.asyncio +async def test_discovery_adapter_preserves_authenticated_context(_mcp_request_ctx): + from mcp.types import DiscoverResult, RequestParams, ServerCapabilities + from litellm.proxy._experimental.mcp_server import server + + expected = DiscoverResult(supported_versions=["2025-11-25"], capabilities=ServerCapabilities()) + dispatched = AsyncMock(return_value=expected) + auth = UserAPIKeyAuth(user_id="discover-caller") + with ( + patch.object(server, "get_or_extract_auth_context", AsyncMock(return_value=(auth, None, ["allowed"], None, None, None, None))), + patch.object(server.operations.GatewayOperations, "execute", dispatched), + ): + result = await server.discover(_mcp_request_ctx(), RequestParams()) + assert result is expected + context = dispatched.await_args.args[1] + assert context.user_api_key_auth.user_id == "discover-caller" + assert context.mcp_servers == ("allowed",) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 64d94065674..16bffa1a356 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -63,7 +63,7 @@ from litellm.proxy._types import ( UserAPIKeyAuth, ) from litellm.types.llms.custom_http import httpxSpecialProvider -from litellm.types.mcp import MCPAuth, MCPAuthType +from litellm.types.mcp import MCPAuth, MCPAuthType, MCPUpstreamProtocol from litellm.types.mcp_server.mcp_server_manager import MCPOAuthMetadata, MCPServer from litellm.caching.caching import DualCache from litellm.caching.llm_caching_handler import LLMClientCache @@ -14637,3 +14637,27 @@ class TestSharedIdentifierPrefixWarning: assert "srv-b" in shared_warnings[0] assert "srv-c" not in shared_warnings[0] assert "'shared'" in shared_warnings[0] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("revision", ["auto", "2024-11-05", "2025-03-26", "2025-06-18", "2025-11-25"]) +async def test_configured_protocol_reaches_the_upstream_client(config_only_mcp_manager_factory, revision): + manager = config_only_mcp_manager_factory() + await manager.load_servers_from_config({"versions": {"url": "http://127.0.0.1:9/mcp", "transport": "http", "protocol_version": revision}}) + server = next(iter(manager.config_mcp_servers.values())) + client = await manager._create_mcp_client(server) + assert server.protocol_version == revision + assert client.protocol_version == revision + + +@pytest.mark.parametrize("revision", ("auto", "2024-11-05", "2025-06-18")) +@pytest.mark.parametrize("explicit", (None, "auto", "2025-11-25")) +def test_runtime_protocol_metadata_preserves_explicit_precedence( + revision: MCPUpstreamProtocol, explicit: MCPUpstreamProtocol | None +) -> None: + server: Final = MCPServer.model_validate({ + "server_id": "preview", "name": "preview", "transport": "http", + "mcp_info": {"protocol_version": revision}, + **({"protocol_version": explicit} if explicit is not None else {}), + }) + assert server.protocol_version == (explicit if explicit is not None else revision) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_operations.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_operations.py index 81f81045740..bb900de4f98 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_operations.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_operations.py @@ -542,3 +542,126 @@ async def test_local_tool_json_array_is_converted_once_for_the_caller_revision(c assert [block.text for block in result.content] == [body] assert result.structured_content == (["a", "b"] if compat == "modern" else None) + + +@pytest.mark.asyncio +async def test_discovery_preserves_caller_scope_and_proxy_restrictions(): + from mcp.types import DiscoverRequest, ListToolsResult, Tool + + listed = AsyncMock(return_value=ListToolsResult(tools=[Tool(name="allowed", input_schema={"type": "object"})])) + context = prepare_context(UserAPIKeyAuth(user_id="scoped"), mcp_servers=["only-this"], mcp_proxy_mode=True, protocol_version="2025-06-18") + with patch("litellm.proxy._experimental.mcp_server.operations._execute_handle_list_tools", listed): + result = await GatewayOperations().execute(DiscoverRequest(), context) + assert result.capabilities.tools is not None + assert result.capabilities.resources is None + assert result.capabilities.prompts is None + assert listed.await_args.args[0] is context + assert listed.await_args.args[0].user_api_key_auth.user_id == "scoped" + assert listed.await_args.args[0].mcp_servers == ("only-this",) + + +@pytest.mark.asyncio +async def test_discovery_denial_cannot_advertise_tools(): + from mcp.types import DiscoverRequest + from fastapi import HTTPException + + denied = AsyncMock(side_effect=HTTPException(status_code=403, detail="Forbidden")) + with patch("litellm.proxy._experimental.mcp_server.operations._execute_handle_list_tools", denied): + with pytest.raises(HTTPException) as error: + await GatewayOperations().execute(DiscoverRequest(), prepare_context(UserAPIKeyAuth(user_id="denied"))) + assert error.value.status_code == 403 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("available", ["none", "resources", "templates", "prompts"]) +async def test_discovery_lists_each_capability_with_the_same_caller(available): + from mcp.types import ( + DiscoverRequest, ListToolsResult, ListPromptsResult, ListResourcesResult, + ListResourceTemplatesResult, Prompt, Resource, ResourceTemplate, + ) + from litellm.proxy._experimental.mcp_server import operations + + context = prepare_context(UserAPIKeyAuth(user_id="scoped"), mcp_servers=["authorized"]) + tools = AsyncMock(return_value=ListToolsResult(tools=[])) + prompts = AsyncMock(return_value=ListPromptsResult(prompts=[Prompt(name="allowed")] if available == "prompts" else [])) + resources = AsyncMock(return_value=ListResourcesResult(resources=[Resource(name="allowed", uri="test://allowed")] if available == "resources" else [])) + templates = AsyncMock(return_value=ListResourceTemplatesResult(resource_templates=[ResourceTemplate(name="allowed", uri_template="test://{id}")] if available == "templates" else [])) + with ( + patch.object(operations, "_execute_handle_list_tools", tools), + patch.object(operations, "_execute_list_prompts", prompts), + patch.object(operations, "_execute_list_resources", resources), + patch.object(operations, "_execute_list_resource_templates", templates), + ): + result = await GatewayOperations().execute(DiscoverRequest(), context) + assert result.capabilities.tools is None + assert (result.capabilities.prompts is not None) == (available == "prompts") + assert (result.capabilities.resources is not None) == (available in {"resources", "templates"}) + for listing in (tools, prompts, resources, templates): + assert listing.await_args.args[0] is context + + +@pytest.mark.asyncio +@pytest.mark.parametrize("outcome", ["success", "failure", "cancel"]) +async def test_discovery_concurrent_listings_drain_on_failure_and_cancellation(outcome): + from mcp.types import DiscoverRequest, ListToolsResult, ListPromptsResult, ListResourcesResult, ListResourceTemplatesResult + from litellm.proxy._experimental.mcp_server import operations + + ready = [asyncio.Event() for _ in range(4)] + closed = [asyncio.Event() for _ in range(4)] + release = asyncio.Event() + responses = (ListToolsResult(tools=[]), ListPromptsResult(prompts=[]), ListResourcesResult(resources=[]), ListResourceTemplatesResult(resource_templates=[])) + + def listing(index): + async def run(*args, **kwargs): + ready[index].set() + try: + await release.wait() + if index == 0 and outcome == "failure": + raise ValueError("discovery failed") + if outcome != "success": + await asyncio.Event().wait() + return responses[index] + finally: + closed[index].set() + return run + + with ( + patch.object(operations, "_execute_handle_list_tools", side_effect=listing(0)) as tools, + patch.object(operations, "_execute_list_prompts", side_effect=listing(1)), + patch.object(operations, "_execute_list_resources", side_effect=listing(2)), + patch.object(operations, "_execute_list_resource_templates", side_effect=listing(3)), + ): + task = asyncio.create_task(GatewayOperations().execute(DiscoverRequest(), prepare_context(UserAPIKeyAuth(user_id="scoped")))) + try: + await asyncio.wait_for(asyncio.gather(*(event.wait() for event in ready)), 1) + if outcome == "cancel": + task.cancel() + else: + release.set() + if outcome == "success": + result = await asyncio.wait_for(task, 1) + assert result.capabilities.model_dump(exclude_none=True) == {} + else: + with pytest.raises(asyncio.CancelledError if outcome == "cancel" else ValueError): + await asyncio.wait_for(task, 1) + assert all(event.is_set() for event in closed) + assert tools.call_args.kwargs["log_list_tools_to_spendlogs"] is False + finally: + task.cancel() + await asyncio.gather(task, return_exceptions=True) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("log_enabled", [False, True]) +async def test_tools_listing_preserves_explicit_spend_log_policy(log_enabled): + from mcp.types import PaginatedRequestParams + from litellm.proxy._experimental.mcp_server import operations + + listing = AsyncMock(return_value=operations.AggregateToolListing(tools=[], outcomes={})) + with patch.object(operations, "_list_mcp_tools", listing): + result = await operations._execute_handle_list_tools( + prepare_context(UserAPIKeyAuth(user_id="caller")), PaginatedRequestParams(), + log_list_tools_to_spendlogs=log_enabled, + ) + assert result.tools == [] + assert listing.await_args.kwargs["log_list_tools_to_spendlogs"] is log_enabled diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py index 4da120cb26f..e82ab28bb4c 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py @@ -28,7 +28,7 @@ from litellm.proxy._types import ( UserAPIKeyAuth, ) from litellm.proxy.auth.user_api_key_auth import user_api_key_auth -from litellm.types.mcp import MCPAuth, MCPTransport +from litellm.types.mcp import MCPAuth, MCPTransport, MCPUpstreamProtocol from litellm.types.mcp_server.mcp_server_manager import MCPServer _OK_TOOL_RESULT: Final = CallToolResult(content=[TextContent(type="text", text='{"result": "ok"}')], is_error=False) @@ -1476,7 +1476,7 @@ class TestListToolsRestAPI: monkeypatch, ): """The REST tools/list path should include tools beyond the upstream first page.""" - from mcp.types import ListToolsResult, PaginatedRequestParams + from mcp.types import Implementation, InitializeResult, ListToolsResult, PaginatedRequestParams, ServerCapabilities from mcp.types import Tool as MCPTool import litellm.experimental_mcp_client.client as mcp_client_module @@ -1512,7 +1512,11 @@ class TestListToolsRestAPI: mock_session_ctx = AsyncMock() mock_session_instance = AsyncMock() - mock_session_instance.initialize = AsyncMock(return_value=None) + mock_session_instance.initialize = AsyncMock(return_value=InitializeResult( + protocol_version="2025-11-25", + capabilities=ServerCapabilities(), + server_info=Implementation(name="stub", version="1"), + )) mock_session_instance.list_tools.side_effect = [ ListToolsResult( tools=[ @@ -4628,3 +4632,74 @@ class TestClientAllowlistOnRestRoutes: assert denied.value.detail["error"] == "Forbidden" assert "'claude-code'" in denied.value.detail["details"] acting.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("revision", ("auto", "2024-11-05", "2025-06-18")) +async def test_preview_client_honors_protocol_metadata(revision: MCPUpstreamProtocol) -> None: + from litellm.experimental_mcp_client.client import MCPClient + + payload: Final = NewMCPServerRequest( + server_name="preview", url="http://127.0.0.1:9/mcp", transport="http", + auth_type=MCPAuth.none, mcp_info={"protocol_version": revision}, + ) + + async def inspect_client(client: MCPClient) -> dict[str, str]: + return {"protocol_version": client.protocol_version} + + result: Final = await rest_endpoints._execute_with_mcp_client(payload, inspect_client) + assert result == {"protocol_version": revision} + + +@pytest.mark.asyncio +@pytest.mark.parametrize("auth_type", (MCPAuth.none, MCPAuth.bearer_token, MCPAuth.oauth2)) +@pytest.mark.parametrize( + ("metadata", "expected"), + ( + (None, "2025-11-25"), + ({}, "2025-11-25"), + ({"description": "edited"}, "2025-11-25"), + ({"protocol_version": "auto"}, "auto"), + ({"protocol_version": "2024-11-05"}, "2024-11-05"), + ({"protocol_version": "2025-06-18"}, "2025-06-18"), + ), +) +async def test_saved_preview_protocol_omission_and_explicit_edits( + monkeypatch: pytest.MonkeyPatch, auth_type: MCPAuth, + metadata: dict[str, str] | None, expected: MCPUpstreamProtocol, +) -> None: + from starlette.datastructures import Headers + + from litellm.experimental_mcp_client.client import MCPClient + from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager + from litellm.proxy.management_endpoints import mcp_management_endpoints + + saved: Final = MCPServer( + server_id="saved-preview", name="preview", url="https://example.com/mcp", + transport="http", auth_type=auth_type, protocol_version="2025-11-25", + authentication_token="stored-token", + authorization_url="https://example.com/authorize", token_url="https://example.com/token", + ) + manager: Final = MCPServerManager() + manager.registry = {saved.server_id: saved} + monkeypatch.setattr(rest_endpoints, "global_mcp_server_manager", manager) + monkeypatch.setattr(mcp_management_endpoints, "global_mcp_server_manager", manager) + payload: Final = NewMCPServerRequest( + server_id=saved.server_id, server_name=saved.name, url=saved.url, transport="http", + auth_type=auth_type, mcp_info=metadata, + authorization_url=saved.authorization_url, token_url=saved.token_url, + ) + staged: Final = rest_endpoints._stage_server_test( + payload, Headers({"x-litellm-api-key": "sk-admin", "authorization": "Bearer preview-token"}) + ) + + async def inspect_client(client: MCPClient) -> dict[str, str]: + return {"protocol_version": client.protocol_version} + + result: Final = await rest_endpoints._execute_with_mcp_client( + staged.request, inspect_client, + mcp_auth_header=staged.mcp_auth_header, oauth2_headers=staged.oauth2_headers, + ) + assert result == {"protocol_version": expected} + assert saved.protocol_version == "2025-11-25" + assert payload.mcp_info == metadata 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_proxy_config.py b/tests/test_litellm/proxy/proxy_server/test_proxy_config.py index 7e198bc9131..b2ef327f50e 100644 --- a/tests/test_litellm/proxy/proxy_server/test_proxy_config.py +++ b/tests/test_litellm/proxy/proxy_server/test_proxy_config.py @@ -4903,3 +4903,19 @@ async def test_model_refresh_updates_availability_catalog_and_retains_it_on_db_f assert await pc._get_models_from_db(client) == [] assert pc.auto_router_db_catalog == () assert find_many.await_count == 3 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("versions", [None, ["2024-11-05"], [], ["2026-07-28"], ["unknown"]]) +async def test_proxy_config_validates_advertised_mcp_versions_at_load(tmp_path, monkeypatch, versions): + config = tmp_path / "mcp-versions.yaml" + config.write_text(json.dumps({"model_list": [], "general_settings": {"mcp_advertised_versions": versions}})) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None) + monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", False) + monkeypatch.delenv("LITELLM_CONFIG_BUCKET_NAME", raising=False) + if versions is None or versions == ["2024-11-05"]: + _, _, settings = await ProxyConfig().load_config(router=None, config_file_path=str(config)) + assert settings["mcp_advertised_versions"] == versions + return + with pytest.raises(ValidationError): + await ProxyConfig().load_config(router=None, config_file_path=str(config)) 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__types.py b/tests/test_litellm/proxy/test__types.py index f5abe0561db..b43a75d3323 100644 --- a/tests/test_litellm/proxy/test__types.py +++ b/tests/test_litellm/proxy/test__types.py @@ -377,3 +377,21 @@ def test_change_password_request_passwords_hidden_from_repr(): for rendered in (repr(request), str(request)): assert "hunter2hunter2" not in rendered assert "NewP@ssw0rd-2026" not in rendered +@pytest.mark.parametrize("versions", [[], ["2099-01-01"], ["2026-07-28"]]) +def test_mcp_advertised_versions_reject_unavailable_revisions(versions): + from pydantic import ValidationError + + from litellm.proxy._types import ConfigGeneralSettings + + with pytest.raises(ValidationError): + ConfigGeneralSettings(mcp_advertised_versions=versions) + + +@pytest.mark.parametrize("revision", ["2026-07-28", "unknown", None]) +def test_mcp_metadata_rejects_unavailable_upstream_protocol(revision): + from litellm.proxy._types import NewMCPServerRequest, UpdateMCPServerRequest + + payload = {"server_id": "test", "transport": "http", "url": "https://example.com/mcp", "mcp_info": {"protocol_version": revision}} + for model in (NewMCPServerRequest, UpdateMCPServerRequest): + with pytest.raises(ValidationError): + model.model_validate(payload) diff --git a/tests/test_litellm/proxy/test_common_request_processing.py b/tests/test_litellm/proxy/test_common_request_processing.py index f2dab6389db..4b54df98aeb 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 @@ -2544,6 +2547,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 @@ -9955,3 +10015,270 @@ class TestErrorLogCarriesCallId: record: Final = caplog.records[-1] assert record.litellm_call_id == call_id 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 + written after the SSE generator is exhausted never lands. The row must be + written before the chunk carrying ``response.completed`` is handed to the + client.""" + + CHUNKS: Final = ( + 'data: {"type":"response.created"}\n\n', + 'data: {"type":"response.output_text.delta"}\n\n', + 'data: {"type":"response.completed"}\n\n', + "data: [DONE]\n\n", + ) + TERMINAL_INDEX: Final = 2 + + @staticmethod + def _completed_event() -> SimpleNamespace: + return SimpleNamespace( + type="response.completed", + response=SimpleNamespace( + id="resp_lit8612", + output=[SimpleNamespace(type="code_interpreter_call", container_id="cntr_lit8612")], + ), + ) + + async def _sse(self, stream: SimpleNamespace, populate_at: int) -> AsyncGenerator[str, None]: + for index, chunk in enumerate(self.CHUNKS): + if index == populate_at: + stream.completed_response = self._completed_event() + yield chunk + if populate_at == len(self.CHUNKS): + stream.completed_response = self._completed_event() + + async def _await_counts_per_chunk(self, populate_at: int) -> tuple[tuple[tuple[str, int], ...], AsyncMock]: + stream: Final = SimpleNamespace(completed_response=None, _hidden_params={"custom_llm_provider": "azure"}) + recorder: Final = AsyncMock(return_value=None) + with patch( + "litellm.proxy.container_endpoints.ownership.record_container_owners_from_responses_response", recorder + ): + wrapped: Final = ProxyBaseLLMRequestProcessing._wrap_responses_stream_for_container_ownership( + original_stream_response=stream, + wrapped_generator=self._sse(stream, populate_at), + user_api_key_dict=ProxyUserAPIKeyAuth(api_key="sk-test", team_id="team-1"), + ) + observed: Final = tuple([(chunk, recorder.await_count) async for chunk in wrapped]) + return observed, recorder + + async def test_row_is_written_before_the_terminal_chunk_reaches_the_client(self) -> None: + observed, recorder = await self._await_counts_per_chunk(populate_at=self.TERMINAL_INDEX) + + assert tuple(chunk for chunk, _ in observed) == self.CHUNKS + assert tuple(count for _, count in observed) == (0, 0, 1, 1) + recorder.assert_awaited_once() + assert recorder.await_args.kwargs["response"].output[0].container_id == "cntr_lit8612" + assert recorder.await_args.kwargs["user_api_key_dict"].team_id == "team-1" + + async def test_row_is_still_written_when_the_iterator_completes_only_at_exhaustion(self) -> None: + observed, recorder = await self._await_counts_per_chunk(populate_at=len(self.CHUNKS)) + + assert tuple(chunk for chunk, _ in observed) == self.CHUNKS + assert tuple(count for _, count in observed) == (0, 0, 0, 0) + recorder.assert_awaited_once() 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/integrations/open_telemetry/__init__.py b/tests/unit/expected_responses_api_request/__init__.py similarity index 100% rename from tests/test_litellm/integrations/open_telemetry/__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/experimental_mcp_client/test_mcp_client.py b/tests/unit/experimental_mcp_client/test_mcp_client.py index 368e34c455d..1a56227b008 100644 --- a/tests/unit/experimental_mcp_client/test_mcp_client.py +++ b/tests/unit/experimental_mcp_client/test_mcp_client.py @@ -58,6 +58,15 @@ from litellm.types.mcp_server.mcp_server_manager import MCPServer _JSONRPC_MESSAGE_ADAPTER: Final = TypeAdapter(JSONRPCMessage) +def _initialized(instructions: str | None = None) -> InitializeResult: + return InitializeResult( + protocol_version=LATEST_HANDSHAKE_VERSION, + capabilities=ServerCapabilities(), + server_info=Implementation(name="test", version="1"), + instructions=instructions, + ) + + class _MockTransportClient(MCPClient): """An MCPClient whose streamable-HTTP transport runs on an httpx2 MockTransport.""" @@ -125,7 +134,7 @@ class TestMCPClient: mock_stdio_client.return_value = mock_stdio_ctx mock_session_instance = AsyncMock() - mock_session_instance.initialize = AsyncMock() + mock_session_instance.initialize = AsyncMock(return_value=_initialized()) mock_session_ctx = AsyncMock() mock_session_ctx.__aenter__.return_value = mock_session_instance mock_session_ctx.__aexit__.return_value = None @@ -168,7 +177,7 @@ class TestMCPClient: # Mock the session with patch("litellm.experimental_mcp_client.client.ClientSession") as mock_session: mock_session_instance = AsyncMock() - mock_session_instance.initialize = AsyncMock() + mock_session_instance.initialize = AsyncMock(return_value=_initialized()) mock_session_ctx = AsyncMock() mock_session_ctx.__aenter__.return_value = mock_session_instance mock_session_ctx.__aexit__.return_value = None @@ -214,7 +223,7 @@ class TestMCPClient: # Mock the session with patch("litellm.experimental_mcp_client.client.ClientSession") as mock_session: mock_session_instance = AsyncMock() - mock_session_instance.initialize = AsyncMock() + mock_session_instance.initialize = AsyncMock(return_value=_initialized()) mock_session_ctx = AsyncMock() mock_session_ctx.__aenter__.return_value = mock_session_instance mock_session_ctx.__aexit__.return_value = None @@ -266,7 +275,7 @@ class TestMCPClient: # Mock the session with patch("litellm.experimental_mcp_client.client.ClientSession") as mock_session: mock_session_instance = AsyncMock() - mock_session_instance.initialize = AsyncMock() + mock_session_instance.initialize = AsyncMock(return_value=_initialized()) mock_session_ctx = AsyncMock() mock_session_ctx.__aenter__.return_value = mock_session_instance mock_session_ctx.__aexit__.return_value = None @@ -413,8 +422,7 @@ class TestMCPClientInstructionsCapture: ) mock_session = AsyncMock() - init_result = MagicMock() - init_result.instructions = " upstream says hello " + init_result = _initialized(" upstream says hello ") mock_session.initialize = AsyncMock(return_value=init_result) session_ctx = MagicMock() @@ -442,8 +450,7 @@ class TestMCPClientInstructionsCapture: ) mock_session = AsyncMock() - init_result = MagicMock() - init_result.instructions = None + init_result = _initialized() mock_session.initialize = AsyncMock(return_value=init_result) session_ctx = MagicMock() @@ -600,8 +607,7 @@ class TestExecuteSessionOperationSurfacesTransportError: @patch("litellm.experimental_mcp_client.client.ClientSession") async def test_cleanup_error_after_success_is_swallowed(self, mock_session_cls): client = MCPClient(server_url="http://example.com/mcp", transport_type="http") - init_result = MagicMock() - init_result.instructions = None + init_result = _initialized() self._make_session(mock_session_cls, AsyncMock(return_value=init_result)) transport_ctx = self._make_transport(_FakeExceptionGroup("late", [httpx2.ConnectError("late cleanup error")])) @@ -634,7 +640,7 @@ class TestExecuteSessionOperationSurfacesTransportError: @pytest.mark.parametrize("original_error", (False, True)) @patch("litellm.experimental_mcp_client.client.ClientSession") async def test_session_exit_cancellation_preserves_original_failure(self, session_class, original_error): - self._make_session(session_class, AsyncMock(return_value=None)) + self._make_session(session_class, AsyncMock(return_value=_initialized())) cancelled: Final = asyncio.CancelledError("cancelled while closing session") session_class.return_value.__aexit__ = AsyncMock(side_effect=cancelled) original: Final = RuntimeError("operation failed") @@ -656,7 +662,7 @@ class TestExecuteSessionOperationSurfacesTransportError: @pytest.mark.parametrize("phase", ("session", "transport")) @patch("litellm.experimental_mcp_client.client.ClientSession") async def test_cleanup_preserves_process_exit(self, session_class, phase, signal_type): - self._make_session(session_class, AsyncMock(return_value=None)) + self._make_session(session_class, AsyncMock(return_value=_initialized())) signal: Final = signal_type("process stopping") if phase == "session": session_class.return_value.__aexit__ = AsyncMock(side_effect=signal) @@ -670,7 +676,7 @@ class TestExecuteSessionOperationSurfacesTransportError: @pytest.mark.asyncio @patch("litellm.experimental_mcp_client.client.ClientSession") async def test_session_and_termination_share_one_cleanup_deadline(self, session_class): - self._make_session(session_class, AsyncMock(return_value=None)) + self._make_session(session_class, AsyncMock(return_value=_initialized())) deleting: Final = asyncio.Event() async def close_session(*args): @@ -1883,16 +1889,17 @@ async def test_sse_read_failure_is_preserved() -> None: @pytest.mark.asyncio +@pytest.mark.parametrize("protocol_version", ["auto", "2025-06-18"]) @pytest.mark.parametrize("transport", [MCPTransport.sse, MCPTransport.stdio]) @pytest.mark.parametrize("mode", ["ok", "closed", "silent"]) -async def test_transport_completion_and_normal_messages(transport: MCPTransport, mode: str) -> None: +async def test_transport_completion_and_normal_messages(transport: MCPTransport, mode: str, protocol_version: str) -> None: from mcp import ClientSession from litellm.proxy._experimental.mcp_server.rest_endpoints import _connection_error_message logging_callback: Final = AsyncMock() read_timeout: Final = 0.2 if mode == "silent" else 30 client: Final = MCPClient( - server_url="https://example.com/sse", transport_type=transport, timeout=read_timeout, logging_callback=logging_callback + server_url="https://example.com/sse", transport_type=transport, timeout=read_timeout, logging_callback=logging_callback, protocol_version=protocol_version ) async def operation(session: ClientSession) -> CallToolResult: @@ -2754,12 +2761,13 @@ async def test_http_close_cancellation_cannot_turn_into_success(original_error: @pytest.mark.asyncio +@pytest.mark.parametrize("protocol_version", ("auto", "2025-06-18")) @pytest.mark.parametrize("cancel_mode", ("scope", "task", "wait_for", "read_timeout")) @pytest.mark.parametrize("concurrency", (1, 5)) @pytest.mark.parametrize("termination", ("ok", "hang", "hang_body")) @pytest.mark.parametrize("raise_on_error", (False, True)) async def test_cancellation_delivers_termination_over_tcp( - cancel_mode: str, concurrency: int, termination: str, raise_on_error: bool + cancel_mode: str, concurrency: int, termination: str, raise_on_error: bool, protocol_version: str ) -> None: started: Final = asyncio.Event() scope_ready: Final[asyncio.Future[anyio.CancelScope]] = asyncio.get_running_loop().create_future() @@ -2813,6 +2821,8 @@ async def test_cancellation_delivers_termination_over_tcp( await stop.wait() return if payload["method"] == "initialize": + if cancel_mode != "read_timeout": + await asyncio.sleep(0.75) response: Final = json.dumps( { "jsonrpc": "2.0", @@ -2839,7 +2849,7 @@ async def test_cancellation_delivers_termination_over_tcp( listener: Final = await asyncio.start_server(handle_connection, "127.0.0.1", 0) port: Final = listener.sockets[0].getsockname()[1] client: Final = MCPClient( - server_url=f"http://127.0.0.1:{port}/mcp", timeout=2 if cancel_mode == "read_timeout" else 0.5 if termination != "ok" else 30 + server_url=f"http://127.0.0.1:{port}/mcp", protocol_version=protocol_version, timeout=2 if cancel_mode == "read_timeout" else 30 ) async def calls(): @@ -2866,7 +2876,7 @@ async def test_cancellation_delivers_termination_over_tcp( try: task: Final = asyncio.create_task(invoke()) - await asyncio.wait_for(started.wait(), 3) + await asyncio.wait_for(started.wait(), 30) if cancel_mode == "scope": (await scope_ready).deadline = anyio.current_time() + 0.2 if cancel_mode == "task": @@ -2901,3 +2911,52 @@ async def test_cancellation_delivers_termination_over_tcp( closed: Final = await asyncio.wait_for(asyncio.gather(*connections, return_exceptions=True), 2) assert all(result is None or isinstance(result, asyncio.CancelledError) for result in closed), closed await asyncio.wait_for(listener.wait_closed(), 2) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("revision", ["2024-11-05", "2025-03-26", "2025-06-18", "2025-11-25", "auto"]) +@pytest.mark.parametrize("accepted", [True, False]) +@pytest.mark.parametrize("callbacks", [False, True]) +async def test_configured_upstream_revision_is_offered_and_checked(revision, accepted, callbacks): + from mcp.types import JSONRPCRequest + from mcp_types.version import LATEST_HANDSHAKE_VERSION + + offered = LATEST_HANDSHAKE_VERSION if revision == "auto" else revision + + def respond(request): + if request.method == "DELETE": + return httpx2.Response(200) + payload = _JSONRPC_MESSAGE_ADAPTER.validate_json(request.content) + if not isinstance(payload, JSONRPCRequest): + return httpx2.Response(202) + if payload.method == "initialize": + assert payload.params["protocolVersion"] == offered + assert ("sampling" in payload.params["capabilities"]) == callbacks + assert ("elicitation" in payload.params["capabilities"]) == callbacks + return httpx2.Response(200, json={ + "jsonrpc": "2.0", "id": payload.id, + "result": {"protocolVersion": offered if accepted else "unsupported", + "capabilities": {"tools": {}}, "serverInfo": {"name": "upstream", "version": "1"}}, + }) + assert accepted, "No operation may execute after failed version negotiation" + return httpx2.Response(200, json={"jsonrpc": "2.0", "id": payload.id, "result": {"tools": [{"name": "echo", "inputSchema": {"type": "object"}}]}}) + + client = _MockTransportClient( + respond, server_url="https://example.com/mcp", protocol_version=revision, + sampling_callback=AsyncMock() if callbacks else None, + elicitation_callback=AsyncMock() if callbacks else None, + ) + if accepted: + result = await client.list_tools(raise_on_error=True) + assert [tool.name for tool in result] == ["echo"] + else: + with pytest.raises((MCPError, RuntimeError), match="protocol version"): + await client.list_tools(raise_on_error=True) + + +@pytest.mark.parametrize("revision", ["2026-07-28", "unknown", "", None]) +def test_upstream_protocol_configuration_rejects_unavailable_modes(revision): + from pydantic import ValidationError + + with pytest.raises(ValidationError): + MCPClient(protocol_version=revision) diff --git a/tests/test_litellm/litellm_core_utils/audio_utils/__init__.py b/tests/unit/integrations/SlackAlerting/__init__.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/audio_utils/__init__.py rename to tests/unit/integrations/SlackAlerting/__init__.py diff --git a/tests/test_litellm/integrations/SlackAlerting/test_budget_alert_types.py b/tests/unit/integrations/SlackAlerting/test_budget_alert_types.py similarity index 100% rename from tests/test_litellm/integrations/SlackAlerting/test_budget_alert_types.py rename to tests/unit/integrations/SlackAlerting/test_budget_alert_types.py diff --git a/tests/test_litellm/integrations/SlackAlerting/test_hanging_request_check.py b/tests/unit/integrations/SlackAlerting/test_hanging_request_check.py similarity index 100% rename from tests/test_litellm/integrations/SlackAlerting/test_hanging_request_check.py rename to tests/unit/integrations/SlackAlerting/test_hanging_request_check.py diff --git a/tests/test_litellm/integrations/SlackAlerting/test_model_deprecation_alert.py b/tests/unit/integrations/SlackAlerting/test_model_deprecation_alert.py similarity index 100% rename from tests/test_litellm/integrations/SlackAlerting/test_model_deprecation_alert.py rename to tests/unit/integrations/SlackAlerting/test_model_deprecation_alert.py diff --git a/tests/test_litellm/integrations/SlackAlerting/test_ms_teams.py b/tests/unit/integrations/SlackAlerting/test_ms_teams.py similarity index 100% rename from tests/test_litellm/integrations/SlackAlerting/test_ms_teams.py rename to tests/unit/integrations/SlackAlerting/test_ms_teams.py diff --git a/tests/test_litellm/integrations/SlackAlerting/test_slack_alerting.py b/tests/unit/integrations/SlackAlerting/test_slack_alerting.py similarity index 100% rename from tests/test_litellm/integrations/SlackAlerting/test_slack_alerting.py rename to tests/unit/integrations/SlackAlerting/test_slack_alerting.py diff --git a/tests/test_litellm/integrations/SlackAlerting/test_slack_alerting_digest.py b/tests/unit/integrations/SlackAlerting/test_slack_alerting_digest.py similarity index 100% rename from tests/test_litellm/integrations/SlackAlerting/test_slack_alerting_digest.py rename to tests/unit/integrations/SlackAlerting/test_slack_alerting_digest.py diff --git a/tests/test_litellm/integrations/SlackAlerting/test_slack_alerting_utils.py b/tests/unit/integrations/SlackAlerting/test_slack_alerting_utils.py similarity index 100% rename from tests/test_litellm/integrations/SlackAlerting/test_slack_alerting_utils.py rename to tests/unit/integrations/SlackAlerting/test_slack_alerting_utils.py diff --git a/tests/test_litellm/integrations/SlackAlerting/test_user_spend_alerts.py b/tests/unit/integrations/SlackAlerting/test_user_spend_alerts.py similarity index 100% rename from tests/test_litellm/integrations/SlackAlerting/test_user_spend_alerts.py rename to tests/unit/integrations/SlackAlerting/test_user_spend_alerts.py diff --git a/tests/test_litellm/litellm_core_utils/llm_response_utils/__init__.py b/tests/unit/integrations/arize/__init__.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/llm_response_utils/__init__.py rename to tests/unit/integrations/arize/__init__.py diff --git a/tests/test_litellm/integrations/arize/test_arize.py b/tests/unit/integrations/arize/test_arize.py similarity index 100% rename from tests/test_litellm/integrations/arize/test_arize.py rename to tests/unit/integrations/arize/test_arize.py diff --git a/tests/test_litellm/integrations/arize/test_arize_health_check.py b/tests/unit/integrations/arize/test_arize_health_check.py similarity index 100% rename from tests/test_litellm/integrations/arize/test_arize_health_check.py rename to tests/unit/integrations/arize/test_arize_health_check.py diff --git a/tests/test_litellm/integrations/arize/test_arize_otel_coexistence.py b/tests/unit/integrations/arize/test_arize_otel_coexistence.py similarity index 100% rename from tests/test_litellm/integrations/arize/test_arize_otel_coexistence.py rename to tests/unit/integrations/arize/test_arize_otel_coexistence.py diff --git a/tests/test_litellm/integrations/arize/test_arize_phoenix.py b/tests/unit/integrations/arize/test_arize_phoenix.py similarity index 100% rename from tests/test_litellm/integrations/arize/test_arize_phoenix.py rename to tests/unit/integrations/arize/test_arize_phoenix.py diff --git a/tests/test_litellm/integrations/arize/test_arize_utils.py b/tests/unit/integrations/arize/test_arize_utils.py similarity index 100% rename from tests/test_litellm/integrations/arize/test_arize_utils.py rename to tests/unit/integrations/arize/test_arize_utils.py diff --git a/tests/test_litellm/router_strategy/adaptive_router/__init__.py b/tests/unit/integrations/azure_storage/__init__.py similarity index 100% rename from tests/test_litellm/router_strategy/adaptive_router/__init__.py rename to tests/unit/integrations/azure_storage/__init__.py diff --git a/tests/test_litellm/integrations/azure_storage/test_azure_storage.py b/tests/unit/integrations/azure_storage/test_azure_storage.py similarity index 100% rename from tests/test_litellm/integrations/azure_storage/test_azure_storage.py rename to tests/unit/integrations/azure_storage/test_azure_storage.py diff --git a/tests/test_litellm/rust_bridge/__init__.py b/tests/unit/integrations/bitbucket/__init__.py similarity index 100% rename from tests/test_litellm/rust_bridge/__init__.py rename to tests/unit/integrations/bitbucket/__init__.py diff --git a/tests/test_litellm/integrations/bitbucket/test_bitbucket_integration.py b/tests/unit/integrations/bitbucket/test_bitbucket_integration.py similarity index 100% rename from tests/test_litellm/integrations/bitbucket/test_bitbucket_integration.py rename to tests/unit/integrations/bitbucket/test_bitbucket_integration.py diff --git a/tests/test_litellm/integrations/bitbucket/test_bitbucket_prompt_manager.py b/tests/unit/integrations/bitbucket/test_bitbucket_prompt_manager.py similarity index 93% rename from tests/test_litellm/integrations/bitbucket/test_bitbucket_prompt_manager.py rename to tests/unit/integrations/bitbucket/test_bitbucket_prompt_manager.py index d6668bf9ad8..a1a88653ee6 100644 --- a/tests/test_litellm/integrations/bitbucket/test_bitbucket_prompt_manager.py +++ b/tests/unit/integrations/bitbucket/test_bitbucket_prompt_manager.py @@ -307,37 +307,6 @@ def test_bitbucket_prompt_manager_render_template_not_found(): manager.prompt_manager.render_template("nonexistent", {"some": "variable"}) -@patch("litellm.integrations.bitbucket.bitbucket_prompt_manager.BitBucketClient") -def test_bitbucket_prompt_manager_integration(mock_client_class): - """Test BitBucketPromptManager integration with BitBucketClient.""" - # Mock the BitBucket client - mock_client = MagicMock() - mock_client.get_file_content.return_value = """--- -model: gpt-4 -temperature: 0.7 ---- -Hello {{name}}!""" - mock_client_class.return_value = mock_client - - config = { - "workspace": "test-workspace", - "repository": "test-repo", - "access_token": "test-token", - } - - manager = BitBucketPromptManager(config, prompt_id="test_prompt") - - # Should have loaded the prompt - assert "test_prompt" in manager.prompt_manager.prompts - template = manager.prompt_manager.prompts["test_prompt"] - assert template.model == "gpt-4" - assert template.temperature == 0.7 - - # Test rendering - rendered = manager.prompt_manager.render_template("test_prompt", {"name": "World"}) - assert rendered == "Hello World!" - - def test_bitbucket_prompt_manager_parse_prompt_to_messages(): """Test parsing prompt content into messages.""" config = { diff --git a/tests/test_litellm/rust_bridge/chat_completions/__init__.py b/tests/unit/integrations/cloudzero/__init__.py similarity index 100% rename from tests/test_litellm/rust_bridge/chat_completions/__init__.py rename to tests/unit/integrations/cloudzero/__init__.py diff --git a/tests/test_litellm/integrations/cloudzero/test_cloudzero.py b/tests/unit/integrations/cloudzero/test_cloudzero.py similarity index 100% rename from tests/test_litellm/integrations/cloudzero/test_cloudzero.py rename to tests/unit/integrations/cloudzero/test_cloudzero.py diff --git a/tests/test_litellm/integrations/cloudzero/test_cloudzero_database.py b/tests/unit/integrations/cloudzero/test_cloudzero_database.py similarity index 100% rename from tests/test_litellm/integrations/cloudzero/test_cloudzero_database.py rename to tests/unit/integrations/cloudzero/test_cloudzero_database.py diff --git a/tests/test_litellm/integrations/cloudzero/test_cz_stream_api.py b/tests/unit/integrations/cloudzero/test_cz_stream_api.py similarity index 100% rename from tests/test_litellm/integrations/cloudzero/test_cz_stream_api.py rename to tests/unit/integrations/cloudzero/test_cz_stream_api.py diff --git a/tests/test_litellm/integrations/cloudzero/test_dry_run_endpoint.py b/tests/unit/integrations/cloudzero/test_dry_run_endpoint.py similarity index 100% rename from tests/test_litellm/integrations/cloudzero/test_dry_run_endpoint.py rename to tests/unit/integrations/cloudzero/test_dry_run_endpoint.py diff --git a/tests/test_litellm/integrations/cloudzero/test_transform.py b/tests/unit/integrations/cloudzero/test_transform.py similarity index 100% rename from tests/test_litellm/integrations/cloudzero/test_transform.py rename to tests/unit/integrations/cloudzero/test_transform.py diff --git a/tests/test_litellm/rust_bridge/messages/__init__.py b/tests/unit/integrations/code_interpreter_interception/__init__.py similarity index 100% rename from tests/test_litellm/rust_bridge/messages/__init__.py rename to tests/unit/integrations/code_interpreter_interception/__init__.py diff --git a/tests/test_litellm/integrations/code_interpreter_interception/test_handler.py b/tests/unit/integrations/code_interpreter_interception/test_handler.py similarity index 100% rename from tests/test_litellm/integrations/code_interpreter_interception/test_handler.py rename to tests/unit/integrations/code_interpreter_interception/test_handler.py 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/integrations/conftest.py b/tests/unit/integrations/conftest.py similarity index 94% rename from tests/test_litellm/integrations/conftest.py rename to tests/unit/integrations/conftest.py index adc8e36e0af..48a01ed8d48 100644 --- a/tests/test_litellm/integrations/conftest.py +++ b/tests/unit/integrations/conftest.py @@ -1,6 +1,7 @@ import functools import http.server import ipaddress +import os import queue import ssl import threading @@ -74,6 +75,14 @@ def write_self_signed_cert(directory: Path, stem: str) -> tuple[Path, Path]: return certificate_path, key_path +@pytest.fixture(autouse=True) +def restore_process_environment() -> Iterator[None]: + original: Final = dict(os.environ) + yield + os.environ.clear() + os.environ.update(original) + + @pytest.fixture def tls_sink(tmp_path: Path) -> Iterator[TlsSink]: certificate_path, key_path = write_self_signed_cert(tmp_path, "sink") diff --git a/tests/test_litellm/rust_bridge/ocr/__init__.py b/tests/unit/integrations/datadog/__init__.py similarity index 100% rename from tests/test_litellm/rust_bridge/ocr/__init__.py rename to tests/unit/integrations/datadog/__init__.py diff --git a/tests/test_litellm/integrations/datadog/test_datadog_cost_management.py b/tests/unit/integrations/datadog/test_datadog_cost_management.py similarity index 100% rename from tests/test_litellm/integrations/datadog/test_datadog_cost_management.py rename to tests/unit/integrations/datadog/test_datadog_cost_management.py diff --git a/tests/test_litellm/integrations/datadog/test_datadog_llm_obs.py b/tests/unit/integrations/datadog/test_datadog_llm_obs.py similarity index 100% rename from tests/test_litellm/integrations/datadog/test_datadog_llm_obs.py rename to tests/unit/integrations/datadog/test_datadog_llm_obs.py diff --git a/tests/test_litellm/integrations/datadog/test_datadog_llm_obs_agent.py b/tests/unit/integrations/datadog/test_datadog_llm_obs_agent.py similarity index 100% rename from tests/test_litellm/integrations/datadog/test_datadog_llm_obs_agent.py rename to tests/unit/integrations/datadog/test_datadog_llm_obs_agent.py diff --git a/tests/test_litellm/integrations/datadog/test_datadog_logger_batching.py b/tests/unit/integrations/datadog/test_datadog_logger_batching.py similarity index 100% rename from tests/test_litellm/integrations/datadog/test_datadog_logger_batching.py rename to tests/unit/integrations/datadog/test_datadog_logger_batching.py diff --git a/tests/test_litellm/integrations/datadog/test_datadog_metrics.py b/tests/unit/integrations/datadog/test_datadog_metrics.py similarity index 100% rename from tests/test_litellm/integrations/datadog/test_datadog_metrics.py rename to tests/unit/integrations/datadog/test_datadog_metrics.py diff --git a/tests/test_litellm/integrations/datadog/test_datadog_tags_regression.py b/tests/unit/integrations/datadog/test_datadog_tags_regression.py similarity index 100% rename from tests/test_litellm/integrations/datadog/test_datadog_tags_regression.py rename to tests/unit/integrations/datadog/test_datadog_tags_regression.py diff --git a/tests/test_litellm/integrations/datadog/test_datadog_team_handler.py b/tests/unit/integrations/datadog/test_datadog_team_handler.py similarity index 100% rename from tests/test_litellm/integrations/datadog/test_datadog_team_handler.py rename to tests/unit/integrations/datadog/test_datadog_team_handler.py diff --git a/tests/test_litellm/rust_bridge/responses/__init__.py b/tests/unit/integrations/dotprompt/__init__.py similarity index 100% rename from tests/test_litellm/rust_bridge/responses/__init__.py rename to tests/unit/integrations/dotprompt/__init__.py diff --git a/tests/test_litellm/integrations/dotprompt/chat_prompt.prompt b/tests/unit/integrations/dotprompt/chat_prompt.prompt similarity index 100% rename from tests/test_litellm/integrations/dotprompt/chat_prompt.prompt rename to tests/unit/integrations/dotprompt/chat_prompt.prompt diff --git a/tests/test_litellm/integrations/dotprompt/chat_prompt.v1.prompt b/tests/unit/integrations/dotprompt/chat_prompt.v1.prompt similarity index 100% rename from tests/test_litellm/integrations/dotprompt/chat_prompt.v1.prompt rename to tests/unit/integrations/dotprompt/chat_prompt.v1.prompt diff --git a/tests/test_litellm/integrations/dotprompt/chat_prompt.v2.prompt b/tests/unit/integrations/dotprompt/chat_prompt.v2.prompt similarity index 100% rename from tests/test_litellm/integrations/dotprompt/chat_prompt.v2.prompt rename to tests/unit/integrations/dotprompt/chat_prompt.v2.prompt diff --git a/tests/test_litellm/integrations/dotprompt/coding_assistant.prompt b/tests/unit/integrations/dotprompt/coding_assistant.prompt similarity index 100% rename from tests/test_litellm/integrations/dotprompt/coding_assistant.prompt rename to tests/unit/integrations/dotprompt/coding_assistant.prompt diff --git a/tests/test_litellm/integrations/dotprompt/sample_prompt.prompt b/tests/unit/integrations/dotprompt/sample_prompt.prompt similarity index 100% rename from tests/test_litellm/integrations/dotprompt/sample_prompt.prompt rename to tests/unit/integrations/dotprompt/sample_prompt.prompt diff --git a/tests/test_litellm/integrations/dotprompt/test_prompt_manager.py b/tests/unit/integrations/dotprompt/test_prompt_manager.py similarity index 100% rename from tests/test_litellm/integrations/dotprompt/test_prompt_manager.py rename to tests/unit/integrations/dotprompt/test_prompt_manager.py diff --git a/tests/unit/integrations/focus/__init__.py b/tests/unit/integrations/focus/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/integrations/focus/test_csv_serializer.py b/tests/unit/integrations/focus/test_csv_serializer.py similarity index 100% rename from tests/test_litellm/integrations/focus/test_csv_serializer.py rename to tests/unit/integrations/focus/test_csv_serializer.py diff --git a/tests/test_litellm/integrations/focus/test_destination_factory.py b/tests/unit/integrations/focus/test_destination_factory.py similarity index 100% rename from tests/test_litellm/integrations/focus/test_destination_factory.py rename to tests/unit/integrations/focus/test_destination_factory.py diff --git a/tests/test_litellm/integrations/focus/test_focus_database.py b/tests/unit/integrations/focus/test_focus_database.py similarity index 100% rename from tests/test_litellm/integrations/focus/test_focus_database.py rename to tests/unit/integrations/focus/test_focus_database.py diff --git a/tests/test_litellm/integrations/focus/test_focus_gcs_destination.py b/tests/unit/integrations/focus/test_focus_gcs_destination.py similarity index 100% rename from tests/test_litellm/integrations/focus/test_focus_gcs_destination.py rename to tests/unit/integrations/focus/test_focus_gcs_destination.py diff --git a/tests/test_litellm/integrations/focus/test_focus_transformer.py b/tests/unit/integrations/focus/test_focus_transformer.py similarity index 100% rename from tests/test_litellm/integrations/focus/test_focus_transformer.py rename to tests/unit/integrations/focus/test_focus_transformer.py diff --git a/tests/test_litellm/integrations/focus/test_mavvrik_destination.py b/tests/unit/integrations/focus/test_mavvrik_destination.py similarity index 100% rename from tests/test_litellm/integrations/focus/test_mavvrik_destination.py rename to tests/unit/integrations/focus/test_mavvrik_destination.py diff --git a/tests/test_litellm/integrations/focus/test_s3_destination.py b/tests/unit/integrations/focus/test_s3_destination.py similarity index 100% rename from tests/test_litellm/integrations/focus/test_s3_destination.py rename to tests/unit/integrations/focus/test_s3_destination.py diff --git a/tests/test_litellm/integrations/focus/test_transformer.py b/tests/unit/integrations/focus/test_transformer.py similarity index 100% rename from tests/test_litellm/integrations/focus/test_transformer.py rename to tests/unit/integrations/focus/test_transformer.py diff --git a/tests/test_litellm/integrations/focus/test_vantage_destination.py b/tests/unit/integrations/focus/test_vantage_destination.py similarity index 100% rename from tests/test_litellm/integrations/focus/test_vantage_destination.py rename to tests/unit/integrations/focus/test_vantage_destination.py diff --git a/tests/unit/integrations/gitlab/__init__.py b/tests/unit/integrations/gitlab/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/integrations/gitlab/test_gitlab_client.py b/tests/unit/integrations/gitlab/test_gitlab_client.py similarity index 100% rename from tests/test_litellm/integrations/gitlab/test_gitlab_client.py rename to tests/unit/integrations/gitlab/test_gitlab_client.py diff --git a/tests/test_litellm/integrations/gitlab/test_gitlab_integration.py b/tests/unit/integrations/gitlab/test_gitlab_integration.py similarity index 100% rename from tests/test_litellm/integrations/gitlab/test_gitlab_integration.py rename to tests/unit/integrations/gitlab/test_gitlab_integration.py diff --git a/tests/test_litellm/integrations/gitlab/test_gitlab_prompt_manager.py b/tests/unit/integrations/gitlab/test_gitlab_prompt_manager.py similarity index 100% rename from tests/test_litellm/integrations/gitlab/test_gitlab_prompt_manager.py rename to tests/unit/integrations/gitlab/test_gitlab_prompt_manager.py diff --git a/tests/unit/integrations/langfuse/__init__.py b/tests/unit/integrations/langfuse/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/integrations/langfuse/test_gemini_cached_tokens.py b/tests/unit/integrations/langfuse/test_gemini_cached_tokens.py similarity index 100% rename from tests/test_litellm/integrations/langfuse/test_gemini_cached_tokens.py rename to tests/unit/integrations/langfuse/test_gemini_cached_tokens.py diff --git a/tests/test_litellm/integrations/langfuse/test_langfuse_prompt_management.py b/tests/unit/integrations/langfuse/test_langfuse_prompt_management.py similarity index 100% rename from tests/test_litellm/integrations/langfuse/test_langfuse_prompt_management.py rename to tests/unit/integrations/langfuse/test_langfuse_prompt_management.py diff --git a/tests/test_litellm/integrations/langfuse/test_langfuse_sdk.py b/tests/unit/integrations/langfuse/test_langfuse_sdk.py similarity index 100% rename from tests/test_litellm/integrations/langfuse/test_langfuse_sdk.py rename to tests/unit/integrations/langfuse/test_langfuse_sdk.py diff --git a/tests/unit/integrations/newrelic/__init__.py b/tests/unit/integrations/newrelic/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/integrations/newrelic/test_newrelic.py b/tests/unit/integrations/newrelic/test_newrelic.py similarity index 100% rename from tests/test_litellm/integrations/newrelic/test_newrelic.py rename to tests/unit/integrations/newrelic/test_newrelic.py diff --git a/tests/test_litellm/integrations/newrelic/test_newrelic_metrics.py b/tests/unit/integrations/newrelic/test_newrelic_metrics.py similarity index 100% rename from tests/test_litellm/integrations/newrelic/test_newrelic_metrics.py rename to tests/unit/integrations/newrelic/test_newrelic_metrics.py diff --git a/tests/test_litellm/integrations/newrelic/test_newrelic_team_handler.py b/tests/unit/integrations/newrelic/test_newrelic_team_handler.py similarity index 100% rename from tests/test_litellm/integrations/newrelic/test_newrelic_team_handler.py rename to tests/unit/integrations/newrelic/test_newrelic_team_handler.py diff --git a/tests/unit/integrations/open_telemetry/__init__.py b/tests/unit/integrations/open_telemetry/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/integrations/open_telemetry/_helpers.py b/tests/unit/integrations/open_telemetry/_helpers.py similarity index 100% rename from tests/test_litellm/integrations/open_telemetry/_helpers.py rename to tests/unit/integrations/open_telemetry/_helpers.py diff --git a/tests/test_litellm/integrations/open_telemetry/conftest.py b/tests/unit/integrations/open_telemetry/conftest.py similarity index 100% rename from tests/test_litellm/integrations/open_telemetry/conftest.py rename to tests/unit/integrations/open_telemetry/conftest.py diff --git a/tests/unit/integrations/open_telemetry/data/__init__.py b/tests/unit/integrations/open_telemetry/data/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/integrations/open_telemetry/data/captured_kwargs.json b/tests/unit/integrations/open_telemetry/data/captured_kwargs.json similarity index 100% rename from tests/test_litellm/integrations/open_telemetry/data/captured_kwargs.json rename to tests/unit/integrations/open_telemetry/data/captured_kwargs.json diff --git a/tests/test_litellm/integrations/open_telemetry/data/captured_response.json b/tests/unit/integrations/open_telemetry/data/captured_response.json similarity index 100% rename from tests/test_litellm/integrations/open_telemetry/data/captured_response.json rename to tests/unit/integrations/open_telemetry/data/captured_response.json diff --git a/tests/test_litellm/integrations/open_telemetry/test_otel_admin_endpoints.py b/tests/unit/integrations/open_telemetry/test_otel_admin_endpoints.py similarity index 100% rename from tests/test_litellm/integrations/open_telemetry/test_otel_admin_endpoints.py rename to tests/unit/integrations/open_telemetry/test_otel_admin_endpoints.py diff --git a/tests/test_litellm/integrations/open_telemetry/test_otel_exception_handler.py b/tests/unit/integrations/open_telemetry/test_otel_exception_handler.py similarity index 100% rename from tests/test_litellm/integrations/open_telemetry/test_otel_exception_handler.py rename to tests/unit/integrations/open_telemetry/test_otel_exception_handler.py diff --git a/tests/test_litellm/integrations/open_telemetry/test_otel_passthrough_endpoints.py b/tests/unit/integrations/open_telemetry/test_otel_passthrough_endpoints.py similarity index 100% rename from tests/test_litellm/integrations/open_telemetry/test_otel_passthrough_endpoints.py rename to tests/unit/integrations/open_telemetry/test_otel_passthrough_endpoints.py diff --git a/tests/test_litellm/integrations/open_telemetry/test_otel_unified_endpoints.py b/tests/unit/integrations/open_telemetry/test_otel_unified_endpoints.py similarity index 100% rename from tests/test_litellm/integrations/open_telemetry/test_otel_unified_endpoints.py rename to tests/unit/integrations/open_telemetry/test_otel_unified_endpoints.py diff --git a/tests/test_litellm/integrations/open_telemetry/test_passthrough_parent_span.py b/tests/unit/integrations/open_telemetry/test_passthrough_parent_span.py similarity index 100% rename from tests/test_litellm/integrations/open_telemetry/test_passthrough_parent_span.py rename to tests/unit/integrations/open_telemetry/test_passthrough_parent_span.py diff --git a/tests/unit/integrations/otel/__init__.py b/tests/unit/integrations/otel/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/integrations/otel/test_db_endpoint.py b/tests/unit/integrations/otel/test_db_endpoint.py similarity index 100% rename from tests/test_litellm/integrations/otel/test_db_endpoint.py rename to tests/unit/integrations/otel/test_db_endpoint.py diff --git a/tests/test_litellm/integrations/otel/test_langfuse_logger.py b/tests/unit/integrations/otel/test_langfuse_logger.py similarity index 100% rename from tests/test_litellm/integrations/otel/test_langfuse_logger.py rename to tests/unit/integrations/otel/test_langfuse_logger.py diff --git a/tests/test_litellm/integrations/otel/test_otel_v2_baggage.py b/tests/unit/integrations/otel/test_otel_v2_baggage.py similarity index 100% rename from tests/test_litellm/integrations/otel/test_otel_v2_baggage.py rename to tests/unit/integrations/otel/test_otel_v2_baggage.py diff --git a/tests/test_litellm/integrations/otel/test_otel_v2_components.py b/tests/unit/integrations/otel/test_otel_v2_components.py similarity index 99% rename from tests/test_litellm/integrations/otel/test_otel_v2_components.py rename to tests/unit/integrations/otel/test_otel_v2_components.py index 07705e17d9a..fd10210c5ba 100644 --- a/tests/test_litellm/integrations/otel/test_otel_v2_components.py +++ b/tests/unit/integrations/otel/test_otel_v2_components.py @@ -42,7 +42,7 @@ from opentelemetry.trace.propagation.tracecontext import ( # noqa: E402 ) import litellm # noqa: E402 -from conftest import TlsSink # noqa: E402 +from tests.unit.integrations.conftest import TlsSink # noqa: E402 from litellm.integrations.otel.plumbing import context as ctx_mod # noqa: E402 from litellm.integrations.otel.plumbing import providers # noqa: E402 from litellm.integrations.otel.model.config import OpenTelemetryV2Config # noqa: E402 diff --git a/tests/test_litellm/integrations/otel/test_otel_v2_config_baggage_parenting_guardrails.py b/tests/unit/integrations/otel/test_otel_v2_config_baggage_parenting_guardrails.py similarity index 100% rename from tests/test_litellm/integrations/otel/test_otel_v2_config_baggage_parenting_guardrails.py rename to tests/unit/integrations/otel/test_otel_v2_config_baggage_parenting_guardrails.py diff --git a/tests/test_litellm/integrations/otel/test_otel_v2_destinations.py b/tests/unit/integrations/otel/test_otel_v2_destinations.py similarity index 100% rename from tests/test_litellm/integrations/otel/test_otel_v2_destinations.py rename to tests/unit/integrations/otel/test_otel_v2_destinations.py diff --git a/tests/test_litellm/integrations/otel/test_otel_v2_dynamic.py b/tests/unit/integrations/otel/test_otel_v2_dynamic.py similarity index 100% rename from tests/test_litellm/integrations/otel/test_otel_v2_dynamic.py rename to tests/unit/integrations/otel/test_otel_v2_dynamic.py diff --git a/tests/test_litellm/integrations/otel/test_otel_v2_emitter.py b/tests/unit/integrations/otel/test_otel_v2_emitter.py similarity index 100% rename from tests/test_litellm/integrations/otel/test_otel_v2_emitter.py rename to tests/unit/integrations/otel/test_otel_v2_emitter.py diff --git a/tests/test_litellm/integrations/otel/test_otel_v2_logger.py b/tests/unit/integrations/otel/test_otel_v2_logger.py similarity index 99% rename from tests/test_litellm/integrations/otel/test_otel_v2_logger.py rename to tests/unit/integrations/otel/test_otel_v2_logger.py index 00c1343f72e..d478c670e58 100644 --- a/tests/test_litellm/integrations/otel/test_otel_v2_logger.py +++ b/tests/unit/integrations/otel/test_otel_v2_logger.py @@ -2945,14 +2945,6 @@ def test_success_without_pre_call_emits_deferred_span(): assert spans[0].end_time == 101_500_000_000 -def test_no_carrier_and_no_payload_is_noop(): - logger, exporter = _logger() - asyncio.run( - logger.async_log_success_event({"litellm_params": {}}, None, None, None) - ) - assert exporter.get_finished_spans() == () - - def test_second_close_for_same_call_does_not_duplicate_span(): """Success and failure can both fire on one logging object for the same call id. The first close pops the carrier and finishes the boundary span; the diff --git a/tests/test_litellm/integrations/otel/test_otel_v2_metrics.py b/tests/unit/integrations/otel/test_otel_v2_metrics.py similarity index 100% rename from tests/test_litellm/integrations/otel/test_otel_v2_metrics.py rename to tests/unit/integrations/otel/test_otel_v2_metrics.py diff --git a/tests/test_litellm/integrations/otel/test_otel_v2_mount.py b/tests/unit/integrations/otel/test_otel_v2_mount.py similarity index 100% rename from tests/test_litellm/integrations/otel/test_otel_v2_mount.py rename to tests/unit/integrations/otel/test_otel_v2_mount.py diff --git a/tests/test_litellm/integrations/otel/test_otel_v2_multibackend.py b/tests/unit/integrations/otel/test_otel_v2_multibackend.py similarity index 100% rename from tests/test_litellm/integrations/otel/test_otel_v2_multibackend.py rename to tests/unit/integrations/otel/test_otel_v2_multibackend.py diff --git a/tests/test_litellm/integrations/otel/test_otel_v2_presets.py b/tests/unit/integrations/otel/test_otel_v2_presets.py similarity index 100% rename from tests/test_litellm/integrations/otel/test_otel_v2_presets.py rename to tests/unit/integrations/otel/test_otel_v2_presets.py diff --git a/tests/test_litellm/integrations/otel/test_otel_v2_sources_of_truth.py b/tests/unit/integrations/otel/test_otel_v2_sources_of_truth.py similarity index 100% rename from tests/test_litellm/integrations/otel/test_otel_v2_sources_of_truth.py rename to tests/unit/integrations/otel/test_otel_v2_sources_of_truth.py diff --git a/tests/test_litellm/integrations/otel/test_otel_v2_vendor_mappers.py b/tests/unit/integrations/otel/test_otel_v2_vendor_mappers.py similarity index 100% rename from tests/test_litellm/integrations/otel/test_otel_v2_vendor_mappers.py rename to tests/unit/integrations/otel/test_otel_v2_vendor_mappers.py diff --git a/tests/test_litellm/integrations/otel/test_runtime.py b/tests/unit/integrations/otel/test_runtime.py similarity index 100% rename from tests/test_litellm/integrations/otel/test_runtime.py rename to tests/unit/integrations/otel/test_runtime.py diff --git a/tests/test_litellm/integrations/rubrik_test_helpers.py b/tests/unit/integrations/rubrik_test_helpers.py similarity index 100% rename from tests/test_litellm/integrations/rubrik_test_helpers.py rename to tests/unit/integrations/rubrik_test_helpers.py diff --git a/tests/test_litellm/integrations/test_agentops.py b/tests/unit/integrations/test_agentops.py similarity index 100% rename from tests/test_litellm/integrations/test_agentops.py rename to tests/unit/integrations/test_agentops.py diff --git a/tests/test_litellm/integrations/test_anthropic_cache_control_hook.py b/tests/unit/integrations/test_anthropic_cache_control_hook.py similarity index 100% rename from tests/test_litellm/integrations/test_anthropic_cache_control_hook.py rename to tests/unit/integrations/test_anthropic_cache_control_hook.py diff --git a/tests/test_litellm/integrations/test_athina.py b/tests/unit/integrations/test_athina.py similarity index 100% rename from tests/test_litellm/integrations/test_athina.py rename to tests/unit/integrations/test_athina.py diff --git a/tests/test_litellm/integrations/test_azure_sentinel.py b/tests/unit/integrations/test_azure_sentinel.py similarity index 100% rename from tests/test_litellm/integrations/test_azure_sentinel.py rename to tests/unit/integrations/test_azure_sentinel.py diff --git a/tests/test_litellm/integrations/test_braintrust_logging.py b/tests/unit/integrations/test_braintrust_logging.py similarity index 100% rename from tests/test_litellm/integrations/test_braintrust_logging.py rename to tests/unit/integrations/test_braintrust_logging.py diff --git a/tests/test_litellm/integrations/test_braintrust_span_name.py b/tests/unit/integrations/test_braintrust_span_name.py similarity index 100% rename from tests/test_litellm/integrations/test_braintrust_span_name.py rename to tests/unit/integrations/test_braintrust_span_name.py diff --git a/tests/test_litellm/integrations/test_custom_guardrail.py b/tests/unit/integrations/test_custom_guardrail.py similarity index 100% rename from tests/test_litellm/integrations/test_custom_guardrail.py rename to tests/unit/integrations/test_custom_guardrail.py diff --git a/tests/test_litellm/integrations/test_custom_guardrail_recursion.py b/tests/unit/integrations/test_custom_guardrail_recursion.py similarity index 100% rename from tests/test_litellm/integrations/test_custom_guardrail_recursion.py rename to tests/unit/integrations/test_custom_guardrail_recursion.py diff --git a/tests/test_litellm/integrations/test_custom_prompt_management.py b/tests/unit/integrations/test_custom_prompt_management.py similarity index 100% rename from tests/test_litellm/integrations/test_custom_prompt_management.py rename to tests/unit/integrations/test_custom_prompt_management.py diff --git a/tests/test_litellm/integrations/test_deepeval.py b/tests/unit/integrations/test_deepeval.py similarity index 100% rename from tests/test_litellm/integrations/test_deepeval.py rename to tests/unit/integrations/test_deepeval.py diff --git a/tests/test_litellm/integrations/test_galileo.py b/tests/unit/integrations/test_galileo.py similarity index 100% rename from tests/test_litellm/integrations/test_galileo.py rename to tests/unit/integrations/test_galileo.py diff --git a/tests/test_litellm/integrations/test_guardrail_logging_sync.py b/tests/unit/integrations/test_guardrail_logging_sync.py similarity index 100% rename from tests/test_litellm/integrations/test_guardrail_logging_sync.py rename to tests/unit/integrations/test_guardrail_logging_sync.py diff --git a/tests/test_litellm/integrations/test_helicone.py b/tests/unit/integrations/test_helicone.py similarity index 100% rename from tests/test_litellm/integrations/test_helicone.py rename to tests/unit/integrations/test_helicone.py diff --git a/tests/test_litellm/integrations/test_langfuse.py b/tests/unit/integrations/test_langfuse.py similarity index 100% rename from tests/test_litellm/integrations/test_langfuse.py rename to tests/unit/integrations/test_langfuse.py diff --git a/tests/test_litellm/integrations/test_langfuse_otel.py b/tests/unit/integrations/test_langfuse_otel.py similarity index 100% rename from tests/test_litellm/integrations/test_langfuse_otel.py rename to tests/unit/integrations/test_langfuse_otel.py diff --git a/tests/test_litellm/integrations/test_langsmith_init.py b/tests/unit/integrations/test_langsmith_init.py similarity index 100% rename from tests/test_litellm/integrations/test_langsmith_init.py rename to tests/unit/integrations/test_langsmith_init.py diff --git a/tests/test_litellm/integrations/test_lunary.py b/tests/unit/integrations/test_lunary.py similarity index 100% rename from tests/test_litellm/integrations/test_lunary.py rename to tests/unit/integrations/test_lunary.py diff --git a/tests/test_litellm/integrations/test_mlflow.py b/tests/unit/integrations/test_mlflow.py similarity index 100% rename from tests/test_litellm/integrations/test_mlflow.py rename to tests/unit/integrations/test_mlflow.py diff --git a/tests/test_litellm/integrations/test_openmeter.py b/tests/unit/integrations/test_openmeter.py similarity index 100% rename from tests/test_litellm/integrations/test_openmeter.py rename to tests/unit/integrations/test_openmeter.py diff --git a/tests/test_litellm/integrations/test_opentelemetry.py b/tests/unit/integrations/test_opentelemetry.py similarity index 99% rename from tests/test_litellm/integrations/test_opentelemetry.py rename to tests/unit/integrations/test_opentelemetry.py index 974961f2eb5..52eeec31e71 100644 --- a/tests/test_litellm/integrations/test_opentelemetry.py +++ b/tests/unit/integrations/test_opentelemetry.py @@ -33,7 +33,7 @@ from parameterized import parameterized import requests -from conftest import TlsSink, write_self_signed_cert +from tests.unit.integrations.conftest import TlsSink, write_self_signed_cert import litellm from litellm.integrations import opentelemetry as otel_module from litellm.integrations.opentelemetry import ( @@ -1244,64 +1244,6 @@ class TestOpenTelemetry(unittest.TestCase): time.sleep(self.POLL_INTERVAL) return [] - @patch("litellm.integrations.opentelemetry.datetime") - def test_create_guardrail_span_with_valid_info(self, mock_datetime): - # Setup - otel = OpenTelemetry() - otel.tracer = MagicMock() - mock_span = MagicMock() - otel.tracer.start_span.return_value = mock_span - - # Create guardrail information - guardrail_info = { - "guardrail_name": "test_guardrail", - "guardrail_mode": "input", - "masked_entity_count": {"CREDIT_CARD": 2}, - "guardrail_response": "filtered_content", - "start_time": 1609459200.0, - "end_time": 1609459201.0, - } - - # Create a kwargs dict with standard_logging_object containing guardrail information - kwargs = { - "standard_logging_object": {"guardrail_information": [guardrail_info]} - } - - # Call the method - otel._create_guardrail_span(kwargs=kwargs, context=None) - - # Assertions - otel.tracer.start_span.assert_called_once() - - # print all calls to mock_span.set_attribute - print("Calls to mock_span.set_attribute:") - for call in mock_span.set_attribute.call_args_list: - print(call) - - # Check that the span has the correct attributes set - mock_span.set_attribute.assert_any_call("guardrail_name", "test_guardrail") - mock_span.set_attribute.assert_any_call("guardrail_mode", "input") - mock_span.set_attribute.assert_any_call( - "guardrail_response", safe_dumps("filtered_content") - ) - mock_span.set_attribute.assert_any_call( - "masked_entity_count", safe_dumps({"CREDIT_CARD": 2}) - ) - - # Verify that the span was ended - mock_span.end.assert_called_once() - - def test_create_guardrail_span_with_no_info(self): - # Setup - otel = OpenTelemetry() - otel.tracer = MagicMock() - - # Test with no guardrail information - kwargs = {"standard_logging_object": {}} - otel._create_guardrail_span(kwargs=kwargs, context=None) - - # Verify that start_span was never called - otel.tracer.start_span.assert_not_called() def test_get_tracer_to_use_for_request_with_dynamic_headers(self): """Test that get_tracer_to_use_for_request returns a dynamic tracer when dynamic headers are present.""" @@ -5461,10 +5403,6 @@ class TestOpenTelemetryPreprocessingDuration(unittest.TestCase): ) assert "litellm.preprocessing.duration_ms" not in self._attr(span, exp) - def test_none_span_is_noop(self): - OpenTelemetry().set_preprocessing_duration_attribute( - None, {"first_api_call_start_time": datetime(2026, 1, 1)} - ) def test_non_dict_container_is_noop(self): otel = OpenTelemetry() diff --git a/tests/test_litellm/integrations/test_opentelemetry_dynamic_imports.py b/tests/unit/integrations/test_opentelemetry_dynamic_imports.py similarity index 100% rename from tests/test_litellm/integrations/test_opentelemetry_dynamic_imports.py rename to tests/unit/integrations/test_opentelemetry_dynamic_imports.py diff --git a/tests/test_litellm/integrations/test_opik_utils.py b/tests/unit/integrations/test_opik_utils.py similarity index 100% rename from tests/test_litellm/integrations/test_opik_utils.py rename to tests/unit/integrations/test_opik_utils.py diff --git a/tests/test_litellm/integrations/test_otel_guardrail_violation_spans.py b/tests/unit/integrations/test_otel_guardrail_violation_spans.py similarity index 100% rename from tests/test_litellm/integrations/test_otel_guardrail_violation_spans.py rename to tests/unit/integrations/test_otel_guardrail_violation_spans.py diff --git a/tests/test_litellm/integrations/test_otel_team_attributes_matrix.py b/tests/unit/integrations/test_otel_team_attributes_matrix.py similarity index 100% rename from tests/test_litellm/integrations/test_otel_team_attributes_matrix.py rename to tests/unit/integrations/test_otel_team_attributes_matrix.py diff --git a/tests/test_litellm/integrations/test_prometheus_api_promql_escape.py b/tests/unit/integrations/test_prometheus_api_promql_escape.py similarity index 100% rename from tests/test_litellm/integrations/test_prometheus_api_promql_escape.py rename to tests/unit/integrations/test_prometheus_api_promql_escape.py diff --git a/tests/test_litellm/integrations/test_prometheus_budget_metric_guard.py b/tests/unit/integrations/test_prometheus_budget_metric_guard.py similarity index 100% rename from tests/test_litellm/integrations/test_prometheus_budget_metric_guard.py rename to tests/unit/integrations/test_prometheus_budget_metric_guard.py diff --git a/tests/test_litellm/integrations/test_prometheus_budget_metrics_db_lookups.py b/tests/unit/integrations/test_prometheus_budget_metrics_db_lookups.py similarity index 100% rename from tests/test_litellm/integrations/test_prometheus_budget_metrics_db_lookups.py rename to tests/unit/integrations/test_prometheus_budget_metrics_db_lookups.py diff --git a/tests/test_litellm/integrations/test_prometheus_budget_metrics_timeout.py b/tests/unit/integrations/test_prometheus_budget_metrics_timeout.py similarity index 100% rename from tests/test_litellm/integrations/test_prometheus_budget_metrics_timeout.py rename to tests/unit/integrations/test_prometheus_budget_metrics_timeout.py diff --git a/tests/test_litellm/integrations/test_prometheus_cache_metrics.py b/tests/unit/integrations/test_prometheus_cache_metrics.py similarity index 99% rename from tests/test_litellm/integrations/test_prometheus_cache_metrics.py rename to tests/unit/integrations/test_prometheus_cache_metrics.py index aa031bb813b..21f13f0ec5c 100644 --- a/tests/test_litellm/integrations/test_prometheus_cache_metrics.py +++ b/tests/unit/integrations/test_prometheus_cache_metrics.py @@ -1,7 +1,7 @@ """ Unit tests for cache Prometheus metrics. -Run with: uv run pytest tests/test_litellm/integrations/test_prometheus_cache_metrics.py -v +Run with: uv run pytest tests/unit/integrations/test_prometheus_cache_metrics.py -v """ import pytest diff --git a/tests/test_litellm/integrations/test_prometheus_caller_identity.py b/tests/unit/integrations/test_prometheus_caller_identity.py similarity index 100% rename from tests/test_litellm/integrations/test_prometheus_caller_identity.py rename to tests/unit/integrations/test_prometheus_caller_identity.py diff --git a/tests/test_litellm/integrations/test_prometheus_carried_budget_state.py b/tests/unit/integrations/test_prometheus_carried_budget_state.py similarity index 100% rename from tests/test_litellm/integrations/test_prometheus_carried_budget_state.py rename to tests/unit/integrations/test_prometheus_carried_budget_state.py diff --git a/tests/test_litellm/integrations/test_prometheus_client_ip_user_agent.py b/tests/unit/integrations/test_prometheus_client_ip_user_agent.py similarity index 100% rename from tests/test_litellm/integrations/test_prometheus_client_ip_user_agent.py rename to tests/unit/integrations/test_prometheus_client_ip_user_agent.py diff --git a/tests/test_litellm/integrations/test_prometheus_custom_metadata_label_counts.py b/tests/unit/integrations/test_prometheus_custom_metadata_label_counts.py similarity index 100% rename from tests/test_litellm/integrations/test_prometheus_custom_metadata_label_counts.py rename to tests/unit/integrations/test_prometheus_custom_metadata_label_counts.py diff --git a/tests/test_litellm/integrations/test_prometheus_deployment_state_proxy_rejects.py b/tests/unit/integrations/test_prometheus_deployment_state_proxy_rejects.py similarity index 100% rename from tests/test_litellm/integrations/test_prometheus_deployment_state_proxy_rejects.py rename to tests/unit/integrations/test_prometheus_deployment_state_proxy_rejects.py diff --git a/tests/test_litellm/integrations/test_prometheus_end_user_cardinality.py b/tests/unit/integrations/test_prometheus_end_user_cardinality.py similarity index 100% rename from tests/test_litellm/integrations/test_prometheus_end_user_cardinality.py rename to tests/unit/integrations/test_prometheus_end_user_cardinality.py diff --git a/tests/test_litellm/integrations/test_prometheus_input_sequence_length_label.py b/tests/unit/integrations/test_prometheus_input_sequence_length_label.py similarity index 100% rename from tests/test_litellm/integrations/test_prometheus_input_sequence_length_label.py rename to tests/unit/integrations/test_prometheus_input_sequence_length_label.py diff --git a/tests/test_litellm/integrations/test_prometheus_invalid_key_filtering.py b/tests/unit/integrations/test_prometheus_invalid_key_filtering.py similarity index 100% rename from tests/test_litellm/integrations/test_prometheus_invalid_key_filtering.py rename to tests/unit/integrations/test_prometheus_invalid_key_filtering.py diff --git a/tests/test_litellm/integrations/test_prometheus_labels.py b/tests/unit/integrations/test_prometheus_labels.py similarity index 100% rename from tests/test_litellm/integrations/test_prometheus_labels.py rename to tests/unit/integrations/test_prometheus_labels.py diff --git a/tests/test_litellm/integrations/test_prometheus_mcp_tool_metrics.py b/tests/unit/integrations/test_prometheus_mcp_tool_metrics.py similarity index 99% rename from tests/test_litellm/integrations/test_prometheus_mcp_tool_metrics.py rename to tests/unit/integrations/test_prometheus_mcp_tool_metrics.py index 22c36f00ca9..da5a0b35e9d 100644 --- a/tests/test_litellm/integrations/test_prometheus_mcp_tool_metrics.py +++ b/tests/unit/integrations/test_prometheus_mcp_tool_metrics.py @@ -5,7 +5,7 @@ These metrics expose ``mcp_tool_call_metadata`` in Prometheus so Grafana dashboards can break down MCP usage by server and tool name. Run with: - uv run pytest tests/test_litellm/integrations/test_prometheus_mcp_tool_metrics.py -v + uv run pytest tests/unit/integrations/test_prometheus_mcp_tool_metrics.py -v """ from typing import get_args diff --git a/tests/test_litellm/integrations/test_prometheus_media_generation_metrics.py b/tests/unit/integrations/test_prometheus_media_generation_metrics.py similarity index 100% rename from tests/test_litellm/integrations/test_prometheus_media_generation_metrics.py rename to tests/unit/integrations/test_prometheus_media_generation_metrics.py diff --git a/tests/test_litellm/integrations/test_prometheus_metric_name_consistency.py b/tests/unit/integrations/test_prometheus_metric_name_consistency.py similarity index 100% rename from tests/test_litellm/integrations/test_prometheus_metric_name_consistency.py rename to tests/unit/integrations/test_prometheus_metric_name_consistency.py diff --git a/tests/test_litellm/integrations/test_prometheus_metrics_endpoint.py b/tests/unit/integrations/test_prometheus_metrics_endpoint.py similarity index 100% rename from tests/test_litellm/integrations/test_prometheus_metrics_endpoint.py rename to tests/unit/integrations/test_prometheus_metrics_endpoint.py diff --git a/tests/test_litellm/integrations/test_prometheus_missing_metrics.py b/tests/unit/integrations/test_prometheus_missing_metrics.py similarity index 100% rename from tests/test_litellm/integrations/test_prometheus_missing_metrics.py rename to tests/unit/integrations/test_prometheus_missing_metrics.py diff --git a/tests/test_litellm/integrations/test_prometheus_none_metadata.py b/tests/unit/integrations/test_prometheus_none_metadata.py similarity index 100% rename from tests/test_litellm/integrations/test_prometheus_none_metadata.py rename to tests/unit/integrations/test_prometheus_none_metadata.py diff --git a/tests/test_litellm/integrations/test_prometheus_overhead_with_guardrails.py b/tests/unit/integrations/test_prometheus_overhead_with_guardrails.py similarity index 100% rename from tests/test_litellm/integrations/test_prometheus_overhead_with_guardrails.py rename to tests/unit/integrations/test_prometheus_overhead_with_guardrails.py diff --git a/tests/test_litellm/integrations/test_prometheus_queue_guardrail_metrics.py b/tests/unit/integrations/test_prometheus_queue_guardrail_metrics.py similarity index 100% rename from tests/test_litellm/integrations/test_prometheus_queue_guardrail_metrics.py rename to tests/unit/integrations/test_prometheus_queue_guardrail_metrics.py diff --git a/tests/test_litellm/integrations/test_prometheus_rate_limit_labels.py b/tests/unit/integrations/test_prometheus_rate_limit_labels.py similarity index 100% rename from tests/test_litellm/integrations/test_prometheus_rate_limit_labels.py rename to tests/unit/integrations/test_prometheus_rate_limit_labels.py diff --git a/tests/test_litellm/integrations/test_prometheus_remaining_tokens_router_fallback.py b/tests/unit/integrations/test_prometheus_remaining_tokens_router_fallback.py similarity index 100% rename from tests/test_litellm/integrations/test_prometheus_remaining_tokens_router_fallback.py rename to tests/unit/integrations/test_prometheus_remaining_tokens_router_fallback.py diff --git a/tests/test_litellm/integrations/test_prometheus_requested_model_cardinality.py b/tests/unit/integrations/test_prometheus_requested_model_cardinality.py similarity index 100% rename from tests/test_litellm/integrations/test_prometheus_requested_model_cardinality.py rename to tests/unit/integrations/test_prometheus_requested_model_cardinality.py diff --git a/tests/test_litellm/integrations/test_prometheus_service_tier_label.py b/tests/unit/integrations/test_prometheus_service_tier_label.py similarity index 98% rename from tests/test_litellm/integrations/test_prometheus_service_tier_label.py rename to tests/unit/integrations/test_prometheus_service_tier_label.py index b2212c4ff41..8b8131b5af2 100644 --- a/tests/test_litellm/integrations/test_prometheus_service_tier_label.py +++ b/tests/unit/integrations/test_prometheus_service_tier_label.py @@ -6,7 +6,7 @@ between the tier a provider served and the tier a caller requested, and the end-to-end emit wiring through async_log_success_event. Run with: - uv run pytest tests/test_litellm/integrations/test_prometheus_service_tier_label.py -v + uv run pytest tests/unit/integrations/test_prometheus_service_tier_label.py -v """ import datetime diff --git a/tests/test_litellm/integrations/test_prometheus_services.py b/tests/unit/integrations/test_prometheus_services.py similarity index 100% rename from tests/test_litellm/integrations/test_prometheus_services.py rename to tests/unit/integrations/test_prometheus_services.py diff --git a/tests/test_litellm/integrations/test_prometheus_spend_capture_rate.py b/tests/unit/integrations/test_prometheus_spend_capture_rate.py similarity index 100% rename from tests/test_litellm/integrations/test_prometheus_spend_capture_rate.py rename to tests/unit/integrations/test_prometheus_spend_capture_rate.py diff --git a/tests/test_litellm/integrations/test_prometheus_spend_logs_metadata.py b/tests/unit/integrations/test_prometheus_spend_logs_metadata.py similarity index 100% rename from tests/test_litellm/integrations/test_prometheus_spend_logs_metadata.py rename to tests/unit/integrations/test_prometheus_spend_logs_metadata.py diff --git a/tests/test_litellm/integrations/test_prometheus_stream_label.py b/tests/unit/integrations/test_prometheus_stream_label.py similarity index 100% rename from tests/test_litellm/integrations/test_prometheus_stream_label.py rename to tests/unit/integrations/test_prometheus_stream_label.py diff --git a/tests/test_litellm/integrations/test_prometheus_token_detail_metrics.py b/tests/unit/integrations/test_prometheus_token_detail_metrics.py similarity index 99% rename from tests/test_litellm/integrations/test_prometheus_token_detail_metrics.py rename to tests/unit/integrations/test_prometheus_token_detail_metrics.py index 5e3846d6fa2..03d67316080 100644 --- a/tests/test_litellm/integrations/test_prometheus_token_detail_metrics.py +++ b/tests/unit/integrations/test_prometheus_token_detail_metrics.py @@ -6,7 +6,7 @@ from the Usage object that providers report. They are sparse — only incremented when the underlying detail is populated and > 0. Run with: - uv run pytest tests/test_litellm/integrations/test_prometheus_token_detail_metrics.py -v + uv run pytest tests/unit/integrations/test_prometheus_token_detail_metrics.py -v """ from typing import get_args diff --git a/tests/test_litellm/integrations/test_prometheus_user_team_metrics.py b/tests/unit/integrations/test_prometheus_user_team_metrics.py similarity index 98% rename from tests/test_litellm/integrations/test_prometheus_user_team_metrics.py rename to tests/unit/integrations/test_prometheus_user_team_metrics.py index 0fc91748af2..ab0fb67b52f 100644 --- a/tests/test_litellm/integrations/test_prometheus_user_team_metrics.py +++ b/tests/unit/integrations/test_prometheus_user_team_metrics.py @@ -102,27 +102,6 @@ class TestPrometheusUserTeamCountMetrics: f"litellm_teams_count_metric should accept value {value}: {e}" ) - def test_user_count_metric_with_zero(self, prometheus_logger): - """Test that user count metric handles zero users""" - metric = prometheus_logger.litellm_total_users_metric - - # Should handle zero gracefully - try: - metric.set(0) - assert True - except Exception as e: - pytest.fail(f"litellm_total_users_metric should handle zero: {e}") - - def test_team_count_metric_with_zero(self, prometheus_logger): - """Test that team count metric handles zero teams""" - metric = prometheus_logger.litellm_teams_count_metric - - # Should handle zero gracefully - try: - metric.set(0) - assert True - except Exception as e: - pytest.fail(f"litellm_teams_count_metric should handle zero: {e}") def test_metrics_can_be_updated_multiple_times(self, prometheus_logger): """Test that metrics can be updated multiple times (simulating refresh cycle)""" diff --git a/tests/test_litellm/integrations/test_prometheus_zero_cost_metric.py b/tests/unit/integrations/test_prometheus_zero_cost_metric.py similarity index 100% rename from tests/test_litellm/integrations/test_prometheus_zero_cost_metric.py rename to tests/unit/integrations/test_prometheus_zero_cost_metric.py diff --git a/tests/test_litellm/integrations/test_prompt_manager_ssti.py b/tests/unit/integrations/test_prompt_manager_ssti.py similarity index 100% rename from tests/test_litellm/integrations/test_prompt_manager_ssti.py rename to tests/unit/integrations/test_prompt_manager_ssti.py diff --git a/tests/test_litellm/integrations/test_responses_background_cost.py b/tests/unit/integrations/test_responses_background_cost.py similarity index 100% rename from tests/test_litellm/integrations/test_responses_background_cost.py rename to tests/unit/integrations/test_responses_background_cost.py diff --git a/tests/test_litellm/integrations/test_rubrik.py b/tests/unit/integrations/test_rubrik.py similarity index 99% rename from tests/test_litellm/integrations/test_rubrik.py rename to tests/unit/integrations/test_rubrik.py index 4a2ee487c65..f3fea292bde 100644 --- a/tests/test_litellm/integrations/test_rubrik.py +++ b/tests/unit/integrations/test_rubrik.py @@ -19,7 +19,7 @@ from litellm.integrations.rubrik import ( ) from litellm.proxy._types import UserAPIKeyAuth -from tests.test_litellm.integrations.rubrik_test_helpers import ( +from tests.unit.integrations.rubrik_test_helpers import ( make_inputs_with_tools, make_tool_call_dict, ) diff --git a/tests/test_litellm/integrations/test_s3.py b/tests/unit/integrations/test_s3.py similarity index 100% rename from tests/test_litellm/integrations/test_s3.py rename to tests/unit/integrations/test_s3.py diff --git a/tests/test_litellm/integrations/test_s3_v2.py b/tests/unit/integrations/test_s3_v2.py similarity index 100% rename from tests/test_litellm/integrations/test_s3_v2.py rename to tests/unit/integrations/test_s3_v2.py diff --git a/tests/test_litellm/integrations/test_shadow_eval_logger.py b/tests/unit/integrations/test_shadow_eval_logger.py similarity index 100% rename from tests/test_litellm/integrations/test_shadow_eval_logger.py rename to tests/unit/integrations/test_shadow_eval_logger.py diff --git a/tests/test_litellm/integrations/test_weave_otel.py b/tests/unit/integrations/test_weave_otel.py similarity index 100% rename from tests/test_litellm/integrations/test_weave_otel.py rename to tests/unit/integrations/test_weave_otel.py diff --git a/tests/unit/integrations/websearch_interception/__init__.py b/tests/unit/integrations/websearch_interception/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/integrations/websearch_interception/test_websearch_agentic_loop_cap.py b/tests/unit/integrations/websearch_interception/test_websearch_agentic_loop_cap.py similarity index 100% rename from tests/test_litellm/integrations/websearch_interception/test_websearch_agentic_loop_cap.py rename to tests/unit/integrations/websearch_interception/test_websearch_agentic_loop_cap.py diff --git a/tests/test_litellm/integrations/websearch_interception/test_websearch_chat_completion.py b/tests/unit/integrations/websearch_interception/test_websearch_chat_completion.py similarity index 100% rename from tests/test_litellm/integrations/websearch_interception/test_websearch_chat_completion.py rename to tests/unit/integrations/websearch_interception/test_websearch_chat_completion.py diff --git a/tests/test_litellm/integrations/websearch_interception/test_websearch_interception_handler.py b/tests/unit/integrations/websearch_interception/test_websearch_interception_handler.py similarity index 100% rename from tests/test_litellm/integrations/websearch_interception/test_websearch_interception_handler.py rename to tests/unit/integrations/websearch_interception/test_websearch_interception_handler.py diff --git a/tests/test_litellm/integrations/websearch_interception/test_websearch_interception_thinking.py b/tests/unit/integrations/websearch_interception/test_websearch_interception_thinking.py similarity index 100% rename from tests/test_litellm/integrations/websearch_interception/test_websearch_interception_thinking.py rename to tests/unit/integrations/websearch_interception/test_websearch_interception_thinking.py diff --git a/tests/test_litellm/integrations/websearch_interception/test_websearch_native_blocks.py b/tests/unit/integrations/websearch_interception/test_websearch_native_blocks.py similarity index 100% rename from tests/test_litellm/integrations/websearch_interception/test_websearch_native_blocks.py rename to tests/unit/integrations/websearch_interception/test_websearch_native_blocks.py diff --git a/tests/test_litellm/integrations/websearch_interception/test_websearch_responses.py b/tests/unit/integrations/websearch_interception/test_websearch_responses.py similarity index 100% rename from tests/test_litellm/integrations/websearch_interception/test_websearch_responses.py rename to tests/unit/integrations/websearch_interception/test_websearch_responses.py diff --git a/tests/test_litellm/integrations/websearch_interception/test_websearch_rich_query_shape.py b/tests/unit/integrations/websearch_interception/test_websearch_rich_query_shape.py similarity index 100% rename from tests/test_litellm/integrations/websearch_interception/test_websearch_rich_query_shape.py rename to tests/unit/integrations/websearch_interception/test_websearch_rich_query_shape.py diff --git a/tests/test_litellm/integrations/websearch_interception/test_websearch_short_circuit.py b/tests/unit/integrations/websearch_interception/test_websearch_short_circuit.py similarity index 100% rename from tests/test_litellm/integrations/websearch_interception/test_websearch_short_circuit.py rename to tests/unit/integrations/websearch_interception/test_websearch_short_circuit.py diff --git a/tests/test_litellm/integrations/websearch_interception/test_websearch_streaming_wrap.py b/tests/unit/integrations/websearch_interception/test_websearch_streaming_wrap.py similarity index 100% rename from tests/test_litellm/integrations/websearch_interception/test_websearch_streaming_wrap.py rename to tests/unit/integrations/websearch_interception/test_websearch_streaming_wrap.py diff --git a/tests/test_litellm/integrations/websearch_interception/test_websearch_thinking_constraint.py b/tests/unit/integrations/websearch_interception/test_websearch_thinking_constraint.py similarity index 100% rename from tests/test_litellm/integrations/websearch_interception/test_websearch_thinking_constraint.py rename to tests/unit/integrations/websearch_interception/test_websearch_thinking_constraint.py diff --git a/tests/unit/integrations/zerobus/__init__.py b/tests/unit/integrations/zerobus/__init__.py new file mode 100644 index 00000000000..e69de29bb2d 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/unit/litellm_core_utils/llm_cost_calc/__init__.py b/tests/unit/litellm_core_utils/llm_cost_calc/__init__.py new file mode 100644 index 00000000000..e69de29bb2d 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/unit/litellm_core_utils/prompt_templates/__init__.py b/tests/unit/litellm_core_utils/prompt_templates/__init__.py new file mode 100644 index 00000000000..e69de29bb2d 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/unit/litellm_core_utils/specialty_caches/__init__.py b/tests/unit/litellm_core_utils/specialty_caches/__init__.py new file mode 100644 index 00000000000..e69de29bb2d 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 102038cad30..e8947dd6bf6 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 5f91c204687..2d31f2bd619 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..b1a14e61b96 --- /dev/null +++ b/tests/unit/litellm_core_utils/test_token_counter.py @@ -0,0 +1,1441 @@ +#### What this tests #### +# This tests litellm.token_counter.token_counter() function +import asyncio +import base64 +import importlib +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, + high_detail_image_token_upper_bound, + 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() 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/unit/llms/vertex_ai/rag_engine/__init__.py b/tests/unit/llms/vertex_ai/rag_engine/__init__.py new file mode 100644 index 00000000000..e69de29bb2d 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/_experimental/mcp_server/test_mcp_client_unit.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_client_unit.py index 6438525706a..1a592aa1c9a 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_client_unit.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_client_unit.py @@ -12,7 +12,7 @@ import litellm.experimental_mcp_client.client as mcp_client_module from litellm.experimental_mcp_client.client import MCPClient from litellm.types.mcp import MCPAuth, MCPTransport from mcp.types import CallToolResult as MCPCallToolResult -from mcp.types import ListToolsResult, PaginatedRequestParams +from mcp.types import Implementation, InitializeResult, ListToolsResult, PaginatedRequestParams, ServerCapabilities from mcp.types import Tool as MCPTool @@ -128,6 +128,11 @@ class TestMCPClientUnitTests: mock_session_ctx = AsyncMock() mock_session_class.return_value = mock_session_ctx mock_session_instance = AsyncMock() + mock_session_instance.initialize.return_value = InitializeResult( + protocol_version="2025-11-25", + capabilities=ServerCapabilities(), + server_info=Implementation(name="test-peer", version="1"), + ) mock_session_ctx.__aenter__ = AsyncMock(return_value=mock_session_instance) client = MCPClient( @@ -163,6 +168,11 @@ class TestMCPClientUnitTests: mock_session_ctx = AsyncMock() mock_session_class.return_value = mock_session_ctx mock_session_instance = AsyncMock() + mock_session_instance.initialize.return_value = InitializeResult( + protocol_version="2025-11-25", + capabilities=ServerCapabilities(), + server_info=Implementation(name="test-peer", version="1"), + ) mock_session_ctx.__aenter__ = AsyncMock(return_value=mock_session_instance) mock_tools = [ @@ -204,6 +214,11 @@ class TestMCPClientUnitTests: mock_session_ctx = AsyncMock() mock_session_class.return_value = mock_session_ctx mock_session_instance = AsyncMock() + mock_session_instance.initialize.return_value = InitializeResult( + protocol_version="2025-11-25", + capabilities=ServerCapabilities(), + server_info=Implementation(name="test-peer", version="1"), + ) mock_session_ctx.__aenter__ = AsyncMock(return_value=mock_session_instance) first_page_tools = [ @@ -245,6 +260,11 @@ class TestMCPClientUnitTests: mock_session_ctx = AsyncMock() mock_session_class.return_value = mock_session_ctx mock_session_instance = AsyncMock() + mock_session_instance.initialize.return_value = InitializeResult( + protocol_version="2025-11-25", + capabilities=ServerCapabilities(), + server_info=Implementation(name="test-peer", version="1"), + ) mock_session_ctx.__aenter__ = AsyncMock(return_value=mock_session_instance) mock_session_instance.list_tools.side_effect = [ @@ -277,6 +297,11 @@ class TestMCPClientUnitTests: mock_session_ctx = AsyncMock() mock_session_class.return_value = mock_session_ctx mock_session_instance = AsyncMock() + mock_session_instance.initialize.return_value = InitializeResult( + protocol_version="2025-11-25", + capabilities=ServerCapabilities(), + server_info=Implementation(name="test-peer", version="1"), + ) mock_session_ctx.__aenter__ = AsyncMock(return_value=mock_session_instance) mock_result = MCPCallToolResult(content=[]) @@ -289,7 +314,7 @@ class TestMCPClientUnitTests: assert result == mock_result mock_session_instance.initialize.assert_called_once() mock_session_instance.call_tool.assert_called_once_with( - name="test_tool", arguments={"arg1": "value1"}, progress_callback=ANY + name="test_tool", arguments={"arg1": "value1"}, progress_callback=ANY, allow_input_required=False ) diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py index dfd250338ff..26674373b4e 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py @@ -489,7 +489,7 @@ async def test_sse_mcp_handler_mock(): ) with ( - patch("litellm.proxy._experimental.mcp_server.server.server.run", run), + patch("litellm.proxy._experimental.mcp_server.server.serve_loop", run), patch( "litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED", True, @@ -511,7 +511,7 @@ async def test_sse_mcp_handler_mock(): # Call the handler await handle_sse_mcp(mock_scope, mock_receive, mock_send) - assert run.await_args.args[:2] == (read_stream, write_stream) + assert run.await_args.args[1:3] == (read_stream, write_stream) assert mock_sse.connect_sse.call_args.args[0]["path"] == "/mcp/sse" 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/secret_managers/__init__.py b/tests/unit/secret_managers/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/secret_managers/hashicorp_vault_parity.json b/tests/unit/secret_managers/hashicorp_vault_parity.json similarity index 100% rename from tests/test_litellm/secret_managers/hashicorp_vault_parity.json rename to tests/unit/secret_managers/hashicorp_vault_parity.json diff --git a/tests/test_litellm/secret_managers/test_aws_secret_manager_replication.py b/tests/unit/secret_managers/test_aws_secret_manager_replication.py similarity index 100% rename from tests/test_litellm/secret_managers/test_aws_secret_manager_replication.py rename to tests/unit/secret_managers/test_aws_secret_manager_replication.py diff --git a/tests/test_litellm/secret_managers/test_aws_secret_manager_rotation.py b/tests/unit/secret_managers/test_aws_secret_manager_rotation.py similarity index 100% rename from tests/test_litellm/secret_managers/test_aws_secret_manager_rotation.py rename to tests/unit/secret_managers/test_aws_secret_manager_rotation.py diff --git a/tests/test_litellm/secret_managers/test_aws_secret_manager_v2.py b/tests/unit/secret_managers/test_aws_secret_manager_v2.py similarity index 100% rename from tests/test_litellm/secret_managers/test_aws_secret_manager_v2.py rename to tests/unit/secret_managers/test_aws_secret_manager_v2.py diff --git a/tests/test_litellm/secret_managers/test_base_secret_manager.py b/tests/unit/secret_managers/test_base_secret_manager.py similarity index 100% rename from tests/test_litellm/secret_managers/test_base_secret_manager.py rename to tests/unit/secret_managers/test_base_secret_manager.py diff --git a/tests/test_litellm/secret_managers/test_custom_secret_manager.py b/tests/unit/secret_managers/test_custom_secret_manager.py similarity index 100% rename from tests/test_litellm/secret_managers/test_custom_secret_manager.py rename to tests/unit/secret_managers/test_custom_secret_manager.py diff --git a/tests/test_litellm/secret_managers/test_cyberark_secret_manager.py b/tests/unit/secret_managers/test_cyberark_secret_manager.py similarity index 100% rename from tests/test_litellm/secret_managers/test_cyberark_secret_manager.py rename to tests/unit/secret_managers/test_cyberark_secret_manager.py diff --git a/tests/test_litellm/secret_managers/test_get_azure_ad_token_provider.py b/tests/unit/secret_managers/test_get_azure_ad_token_provider.py similarity index 100% rename from tests/test_litellm/secret_managers/test_get_azure_ad_token_provider.py rename to tests/unit/secret_managers/test_get_azure_ad_token_provider.py diff --git a/tests/test_litellm/secret_managers/test_hashicorp_secret_manager.py b/tests/unit/secret_managers/test_hashicorp_secret_manager.py similarity index 100% rename from tests/test_litellm/secret_managers/test_hashicorp_secret_manager.py rename to tests/unit/secret_managers/test_hashicorp_secret_manager.py diff --git a/tests/test_litellm/secret_managers/test_secret_manager_handler.py b/tests/unit/secret_managers/test_secret_manager_handler.py similarity index 100% rename from tests/test_litellm/secret_managers/test_secret_manager_handler.py rename to tests/unit/secret_managers/test_secret_manager_handler.py diff --git a/tests/test_litellm/secret_managers/test_secret_managers_main.py b/tests/unit/secret_managers/test_secret_managers_main.py similarity index 100% rename from tests/test_litellm/secret_managers/test_secret_managers_main.py rename to tests/unit/secret_managers/test_secret_managers_main.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_unit_shard_missing_paths.py b/tests/unit/test_unit_shard_missing_paths.py index e464402c9d8..0360a227142 100644 --- a/tests/unit/test_unit_shard_missing_paths.py +++ b/tests/unit/test_unit_shard_missing_paths.py @@ -38,8 +38,8 @@ def _run_shard(tmp_path: Path, test_path: str, workers: str) -> subprocess.Compl "PATH": f"{shim_dir}{os.pathsep}{os.environ['PATH']}", "GITHUB_OUTPUT": str(tmp_path / "github_output"), "TEST_PATH": test_path, - "WORKERS": workers, "UNIT_FLAG": "", + "WORKERS": workers, }, capture_output=True, text=True, diff --git a/tests/unit/test_utils.py b/tests/unit/test_utils.py index 768d8955b8e..0cdc52c9a93 100644 --- a/tests/unit/test_utils.py +++ b/tests/unit/test_utils.py @@ -774,6 +774,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"}, @@ -797,6 +798,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"}, @@ -897,6 +899,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 bba67bcf6c2..1e0aed46923 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -28350,6 +28350,11 @@ export interface components { * @description Maximum retention period for spend logs (e.g., '7d' for 7 days). Logs older than this will be deleted. */ maximum_spend_logs_retention_period?: string | null; + /** + * Mcp Advertised Versions + * @description MCP revisions enabled by the gateway. Defaults to all completed legacy revisions. Modern protocol serving and Apps/Tasks remain disabled. + */ + mcp_advertised_versions?: ("2024-11-05" | "2025-03-26" | "2025-06-18" | "2025-11-25")[] | null; /** * Mcp Allowed Clients * @description MCP client applications admitted by the gateway, each an {alias, value} pair where alias is the name shown in the dashboard and logs and value is the identity that must match exactly. When set, every MCP request must carry a client identity equal to one of the values: a JWT caller is identified by the claim named in litellm_jwtauth.mcp_client_id_jwt_field, any other caller by the header named in mcp_client_id_header. A request with no resolvable identity, or an unlisted one, is rejected with 403. Unset means every client is admitted. @@ -32940,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 */ @@ -32974,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 */ @@ -33052,6 +33061,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 */ @@ -33175,6 +33186,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 */ @@ -46725,6 +46738,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 */ @@ -46759,6 +46774,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 */ @@ -46837,6 +46854,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 */ @@ -46960,6 +46979,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]]