Merge remote-tracking branch 'origin/main' into litellm_stream_served_service_tier
Some checks failed
LiteLLM Rust / rust-lint (push) Waiting to run
LiteLLM Rust / rust-test (push) Waiting to run
LiteLLM Rust / rust-wheel (push) Waiting to run
Terraform Provider / gofmt, vet, build, test (push) Has been cancelled
Terraform Provider / Provider endpoints vs proxy OpenAPI schema (push) Has been cancelled

This commit is contained in:
kerry 2026-09-26 00:12:06 +00:00
commit cc446c60ce
564 changed files with 10547 additions and 3874 deletions

View file

@ -5,8 +5,10 @@ flag="${1:?usage: unit_selection.sh <codecov flag>}"
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
}

View file

@ -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

View file

@ -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"
}
}

View file

@ -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 \

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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");

View file

@ -118,7 +118,7 @@ pub(super) struct ParityCase {
pub(super) fn parity_cases() -> Vec<ParityCase> {
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()
}

View file

@ -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 |
| --- | --- |

View file

@ -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

View file

@ -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

View file

@ -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",
]

View file

@ -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),
)

View file

@ -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

View file

@ -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 = (

View file

@ -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",

View file

@ -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

View file

@ -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="",

View file

@ -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://<workspace-id>.zerobus.<region>.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",

View file

@ -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

View file

@ -0,0 +1,5 @@
"""Databricks Zerobus logging integration for LiteLLM."""
from litellm.integrations.zerobus.logger import ZerobusLogger
__all__ = ("ZerobusLogger",)

View file

@ -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)

View file

@ -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://<workspace_id>.zerobus.<region>.<cloud>``, 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://<workspace_id>.zerobus.<region>.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)

View file

@ -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"),
}
)

View file

@ -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,

View file

@ -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":

View file

@ -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):

View file

@ -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"(?<![0-9A-Za-z])[0-9a-f]{64}(?![0-9A-Za-z])")
QUOTED_VALUE: Final = r"'(?:[^'\\]|\\.)*'|\"(?:[^\"\\]|\\.)*\""
BRACKET_ATOM: Final = rf"(?:{QUOTED_VALUE})|[^\[\]{{}}()'\"]"
NESTED_BRACKET_LEVELS: Final = 3
BRACKETED_VALUE: Final = reduce(
lambda inner, _: rf"[\[{{(](?:{BRACKET_ATOM}|{inner})*[\]}})]",
range(NESTED_BRACKET_LEVELS),
rf"[\[{{(](?:{BRACKET_ATOM})*[\]}})]",
)
BARE_VALUE: Final = r"(?!None(?![0-9A-Za-z_]))[^,)\]}\s]+"
class SentryInitOptions(TypedDict):
dsn: ReadOnly[str | None]
traces_sample_rate: ReadOnly[float]
sample_rate: ReadOnly[float]
send_default_pii: ReadOnly[bool]
event_scrubber: ReadOnly[EventScrubber]
before_send: ReadOnly[EventScrubFn]
before_send_transaction: ReadOnly[EventScrubFn]
environment: ReadOnly[str]
def build_repr_field_pattern(field_names: Sequence[str]) -> re.Pattern[str]:
names: Final = "|".join(re.escape(name) for name in field_names)
return re.compile(
rf"(?P<field>(?<![0-9A-Za-z_])(?:{names})=|['\"](?:{names})['\"]:\s*)(?P<value>{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"),
)

View file

@ -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")

View file

@ -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:

View file

@ -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(

View file

@ -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()

View file

@ -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,

View file

@ -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):

View file

@ -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")

View file

@ -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)

View file

@ -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))

View file

@ -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()

View file

@ -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)

View file

@ -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)

View file

@ -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

View file

@ -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)

View file

@ -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})

View file

@ -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))

View file

@ -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,

View file

@ -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(

View file

@ -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)

View file

@ -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

View file

@ -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": <message>}` 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(

View file

@ -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}'."
),
)

View file

@ -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/<id>/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/<id>/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/<id>/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

View file

@ -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,

View file

@ -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)

View file

@ -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(

View file

@ -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:

View file

@ -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

View file

@ -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

View file

@ -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}"

View file

@ -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))

View file

@ -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 = "<system-reminder>"
_REMINDER_CLOSE: Final = "</system-reminder>"
_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
)

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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,

View file

@ -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

View file

@ -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

View file

@ -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))

View file

@ -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

View file

@ -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

View file

@ -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,

View file

@ -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",

View file

@ -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.
]

View file

@ -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"}

View file

@ -86,6 +86,7 @@ LlmCapability = Literal[
"tool_search",
"tool_search_history",
"tool_use",
"upstream_stream_failure",
"vision",
"web_search",
"web_search_server_tool",

View file

@ -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`

View file

@ -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)

View file

@ -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}'
)

View file

@ -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 ----------

View file

@ -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(

View file

@ -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

View file

@ -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())

View file

@ -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

View file

@ -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 = (

View file

@ -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

View file

@ -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.
"""

View file

@ -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]

View file

@ -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."""

View file

@ -1 +0,0 @@
# Levo integration tests

View file

@ -1 +0,0 @@
# This file makes the tests/litellm/litellm_core_utils directory a Python package

File diff suppressed because it is too large Load diff

View file

@ -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"<Encoding {name!r}>"
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="<pad>")
expected.pad(6, direction="left", pad_id=7, pad_type_id=1, pad_token="<pad>")
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)

View file

@ -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")

View file

@ -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",)

View file

@ -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)

View file

@ -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

View file

@ -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

Some files were not shown because too many files have changed in this diff Show more