Merge remote-tracking branch 'origin/main' into litellm_lit8140_docs_tmp

# Conflicts:
#	tests/test_litellm/litellm_core_utils/test_token_counter.py
This commit is contained in:
Shreshth Kharbanda 2026-09-26 00:49:39 +00:00
commit babf5bc3b4
No known key found for this signature in database
353 changed files with 9916 additions and 3993 deletions

View file

@ -5,6 +5,7 @@ flag="${1:?usage: unit_selection.sh <codecov flag>}"
legacy_flags=(
caching-local
core-utils
enterprise-package
enterprise-routing
integrations
@ -32,6 +33,7 @@ legacy_flags=(
legacy_paths() {
case "$1" in
caching-local) echo tests/unit/caching ;;
core-utils) echo tests/unit/litellm_core_utils ;;
enterprise-package)
echo tests/unit/enterprise/integrations
echo tests/unit/enterprise/proxy/auth
@ -42,6 +44,8 @@ legacy_paths() {
echo tests/unit/enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py ;;
enterprise-routing)
echo tests/unit/google_genai
echo tests/unit/router_strategy
echo tests/unit/router_utils
echo tests/unit/enterprise/enterprise_callbacks/send_emails
echo tests/unit/enterprise/proxy/test_afile_retrieve_returns_unified_id.py
echo tests/unit/enterprise/proxy/test_batch_retrieve_input_file_id.py
@ -77,6 +81,7 @@ legacy_paths() {
echo tests/unit/messages
echo tests/unit/rag
echo tests/unit/rerank_api
echo tests/unit/rust_bridge
echo tests/unit/secret_managers
echo tests/unit/vector_stores
echo tests/unit/videos ;;
@ -142,7 +147,9 @@ legacy_paths() {
proxy-db-proxy-utils) echo tests/unit/proxy/test_proxy_utils.py ;;
proxy-extras) echo tests/unit/litellm_proxy_extras ;;
proxy-infra) echo tests/unit/gateway ;;
responses-caching-types) echo tests/unit/types ;;
responses-caching-types)
find tests/unit/responses -name 'test_*.py' -not -path 'tests/unit/responses/mcp/*'
echo tests/unit/types ;;
*) echo "unit_selection.sh: unknown flag $1" >&2; exit 1 ;;
esac
}

View file

@ -369,6 +369,13 @@ workflows:
reruns: 2
base_ref: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.base.ref or "" >>
pull_request_url: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.url or "" >>
- unit:
name: unit-core-utils
flag: core-utils
shards: 2
reruns: 1
base_ref: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.base.ref or "" >>
pull_request_url: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.url or "" >>
- unit:
name: unit-integrations
flag: integrations

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
@ -111,7 +110,6 @@ jobs:
tests/test_litellm/interactions
tests/test_litellm/ocr
tests/test_litellm/passthrough
tests/test_litellm/rust_bridge
tests/test_litellm/test_*.py
unit-flag: misc
workers: 2
@ -228,9 +226,7 @@ jobs:
- shard: responses-caching-types
artifact-name: responses-caching-types
test-path: >-
tests/test_litellm/responses
tests/test_litellm/caching
test-path: ""
unit-flag: responses-caching-types
workers: 2
reruns: 2

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
@ -329,10 +329,10 @@ test-unit-integrations: install-test-deps
$(UV_RUN) pytest tests/unit/integrations --tb=short -vv -n 4 --durations=20
test-unit-core-utils: install-test-deps
$(UV_RUN) pytest tests/test_litellm/litellm_core_utils --tb=short -vv -n 2 --durations=20
$(UV_RUN) pytest tests/unit/litellm_core_utils --tb=short -vv -n 2 --durations=20
test-unit-other: install-test-deps
$(UV_RUN) pytest tests/test_litellm/caching tests/test_litellm/responses tests/unit/secret_managers tests/unit/vector_stores tests/unit/a2a_protocol tests/test_litellm/anthropic_interface tests/unit/completion_extras tests/unit/containers tests/unit/enterprise tests/unit/experimental_mcp_client tests/unit/google_genai tests/unit/images tests/unit/interactions tests/test_litellm/interactions tests/test_litellm/passthrough tests/test_litellm/router_strategy tests/test_litellm/router_utils tests/unit/types --tb=short -vv -n 4 --durations=20
$(UV_RUN) pytest tests/unit/caching tests/unit/responses tests/unit/secret_managers tests/unit/vector_stores tests/unit/a2a_protocol tests/test_litellm/anthropic_interface tests/unit/completion_extras tests/unit/containers tests/unit/enterprise tests/unit/experimental_mcp_client tests/unit/google_genai tests/unit/images tests/unit/interactions tests/test_litellm/interactions tests/test_litellm/passthrough tests/unit/router_strategy tests/unit/router_utils tests/unit/types --tb=short -vv -n 4 --durations=20
test-unit-root: install-test-deps
$(UV_RUN) pytest tests/unit/test_*.py tests/test_litellm/test_*.py --tb=short -vv -n 4 --durations=20

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

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

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

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

@ -68,15 +68,11 @@ def _error_status_and_message(exc: Exception) -> tuple[int, str]:
def _mid_stream_error_sse_event(exc: Exception) -> bytes:
from litellm.anthropic_interface.exceptions.exception_mapping_utils import (
AnthropicExceptionMapping,
anthropic_error_sse_frame,
)
status_code, message = _error_status_and_message(exc)
error_response = AnthropicExceptionMapping.transform_to_anthropic_error(
status_code=status_code,
raw_message=message,
)
return f"event: error\ndata: {json.dumps(error_response)}\n\n".encode()
return anthropic_error_sse_frame(status_code=status_code, raw_message=message).encode()
def _delta_payload_field(delta_type: StreamingContentBlockDeltaType) -> str:

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)

File diff suppressed because it is too large Load diff

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

@ -4019,6 +4019,18 @@ class AllCallbacks(LiteLLMPydanticObjectBase):
],
)
zerobus: CallbackOnUI = CallbackOnUI(
litellm_callback_name="zerobus",
ui_callback_name="Databricks Zerobus",
litellm_callback_params=[ # mutable-ok: the registry field is typed list
"ZEROBUS_WORKSPACE_URL",
"ZEROBUS_SERVER_ENDPOINT",
"ZEROBUS_CLIENT_ID",
"ZEROBUS_CLIENT_SECRET",
"ZEROBUS_TABLE_NAME",
],
)
class HTTPExceptionErrorDetail(TypedDict):
"""The `{"error": <message>}` shape most proxy endpoints raise as `HTTPException.detail`."""
@ -5355,11 +5367,12 @@ class LiteLLM_JWTAuth(LiteLLMPydanticObjectBase):
default=False,
description=(
"When True, users whose JWT contains no team claims are authenticated "
"using their database team memberships instead of receiving HTTP 403. "
"Usage is attributed to the user's first resolvable DB team, or to the "
"team specified via the x-litellm-team-id request header (validated "
"against DB membership). Requires user_id_upsert=True so that user "
"records exist before the fallback runs."
"using their database team memberships instead of receiving HTTP 403, "
"with usage attributed to the user's first resolvable DB team. Whether or "
"not the JWT carries team claims, the x-litellm-team-id request header may "
"select any team the user is a member of in the database (validated against "
"DB membership); without the header the JWT team stays the default. Requires "
"user_id_upsert=True so that user records exist before the fallback runs."
),
)
issuers: list[JWTIssuerConfig] | None = Field(

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)
@ -3852,6 +3867,7 @@ class ProxyBaseLLMRequestProcessing:
serialize_error: StreamErrorSerializer,
request: Request | None = None,
flush_tail: Callable[[], bytes] | None = None,
seal_open_frame: Callable[[bytes], str] | None = None,
) -> AsyncGenerator[str, None]:
"""
Shared streaming data generator: runs proxy iterator hook, per-chunk hook,
@ -3861,6 +3877,12 @@ class ProxyBaseLLMRequestProcessing:
``flush_tail`` runs once after the upstream iterator completes cleanly and
its non-empty result is yielded, so a serializer that buffers bytes across
chunks can emit anything still held at end of stream.
``seal_open_frame`` is given the tail of what has been yielded when the
error frame goes out, and what it returns is written first. A passthrough
relays raw upstream bytes, so an upstream that hangs up mid-frame leaves the
client inside an open frame, where an error frame would be swallowed or
misparsed instead of raised.
"""
verbose_proxy_logger.debug("inside generator")
# Resolve per-stream (not per-chunk) whether the heavy per-chunk path
@ -3877,6 +3899,7 @@ class ProxyBaseLLMRequestProcessing:
stream_completed = False
client_disconnected = False
delivered_chunk = False
recent_tail = SSE_STREAM_START_TAIL # rebind-ok: rolling window over the yielded bytes
try:
str_so_far = ""
async for chunk in proxy_logging_obj.async_post_call_streaming_iterator_hook(
@ -3922,7 +3945,9 @@ class ProxyBaseLLMRequestProcessing:
# False and refunds. A keepalive ping carries no provider output,
# so it must not suppress that refund.
delivered_chunk = delivered_chunk or chunk != STREAM_SSE_KEEPALIVE_PING_BYTES
yield serialize_chunk(chunk)
serialized = serialize_chunk(chunk)
recent_tail = advance_sse_tail(recent_tail, serialized)
yield serialized
held_tail: Final = flush_tail() if flush_tail is not None else b""
if held_tail:
yield serialize_chunk(held_tail)
@ -3970,7 +3995,9 @@ class ProxyBaseLLMRequestProcessing:
code=stream_error_status,
)
stream_completed = True
yield serialize_error(proxy_exception)
error_frame: Final = serialize_error(proxy_exception)
seal: Final = "" if seal_open_frame is None else seal_open_frame(recent_tail)
yield seal + error_frame if seal else error_frame
finally:
await ProxyBaseLLMRequestProcessing._finalize_streaming_generator_cleanup(
request=request,
@ -3992,7 +4019,7 @@ class ProxyBaseLLMRequestProcessing:
restamp_model: str | None = None,
) -> AsyncGenerator[str, None]:
"""
Anthropic /messages and Google /generateContent streaming data generator require SSE events.
Anthropic /messages streaming data generator, which requires SSE events.
Returns the underlying ``async_streaming_data_generator`` configured with
SSE serializers directly (rather than re-wrapping it in another
@ -4010,11 +4037,13 @@ class ProxyBaseLLMRequestProcessing:
request_data=request_data,
proxy_logging_obj=proxy_logging_obj,
serialize_chunk=ProxyBaseLLMRequestProcessing._sse_chunk_serializer(restamper),
serialize_error=lambda proxy_exc: (
f"{STREAM_SSE_DATA_PREFIX}{json.dumps({'error': proxy_exc.to_dict()})}\n\n"
serialize_error=lambda proxy_exc: anthropic_error_sse_frame(
status_code=error_status_code(proxy_exc, status.HTTP_500_INTERNAL_SERVER_ERROR),
raw_message=proxy_exc.message,
),
request=request,
flush_tail=None if restamper is None else restamper.flush,
seal_open_frame=seal_open_sse_frame,
)
@overload

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)

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

@ -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
@ -3731,6 +3734,7 @@ class CustomPricingLiteLLMParams(MirroredPricingParams):
cache_read_input_token_cost_above_272k_tokens_priority: float | None = None
cache_read_input_token_cost_above_272k_tokens_flex: float | None = None
cache_read_input_token_cost_batches: float | None = None
cache_read_input_token_cost_above_200k_tokens_batches: float | None = None
cache_read_input_token_cost_above_272k_tokens_batches: float | None = None
cache_creation_input_token_cost_batches: float | None = None
cache_creation_input_token_cost_above_272k_tokens_batches: float | None = None
@ -3744,6 +3748,7 @@ class CustomPricingLiteLLMParams(MirroredPricingParams):
input_cost_per_token_above_200k_tokens_priority: float | None = None
input_cost_per_token_above_272k_tokens_priority: float | None = None
input_cost_per_token_above_272k_tokens_flex: float | None = None
input_cost_per_token_above_200k_tokens_batches: float | None = None
input_cost_per_token_above_272k_tokens_batches: float | None = None
input_cost_per_query: float | None = None
input_cost_per_image: float | None = None
@ -3768,6 +3773,7 @@ class CustomPricingLiteLLMParams(MirroredPricingParams):
output_cost_per_token_above_200k_tokens_priority: float | None = None
output_cost_per_token_above_272k_tokens_priority: float | None = None
output_cost_per_token_above_272k_tokens_flex: float | None = None
output_cost_per_token_above_200k_tokens_batches: float | None = None
output_cost_per_token_above_272k_tokens_batches: float | None = None
output_cost_per_character_above_128k_tokens: float | None = None
output_cost_per_image: float | None = None
@ -4142,7 +4148,7 @@ FILE_CONTENT_STREAMING_PROVIDERS: Final[frozenset[str]] = frozenset(
LITELLM_EXECUTED_BATCH_PROVIDERS: Final[frozenset[str]] = frozenset({LlmProviders.HOSTED_VLLM.value})
ListBatchesSupportedProvider = Literal["openai", "azure", "hosted_vllm", "litellm_proxy", "vertex_ai"]
ListBatchesSupportedProvider = Literal["openai", "azure", "hosted_vllm", "litellm_proxy", "vertex_ai", "xai"]
LIST_BATCHES_SUPPORTED_PROVIDERS: Final[frozenset[str]] = frozenset(get_args(ListBatchesSupportedProvider))

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"
),
@ -9358,6 +9367,10 @@ class ProviderConfigManager:
from litellm.llms.mistral.files.transformation import MistralFilesConfig
return MistralFilesConfig()
elif LlmProviders.XAI == provider:
from litellm.llms.xai.files.transformation import XAIFilesConfig
return XAIFilesConfig()
return None
@staticmethod

File diff suppressed because it is too large Load diff

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,
@ -355,6 +360,11 @@
"minimum": 0,
"description": "Rate applied once the prompt exceeds the token threshold in the field name."
},
"input_cost_per_token_above_200k_tokens_batches": {
"type": "number",
"minimum": 0,
"description": "Rate applied once the prompt exceeds the token threshold in the field name."
},
"input_cost_per_token_above_200k_tokens_priority": {
"type": "number",
"minimum": 0,
@ -712,6 +722,11 @@
"minimum": 0,
"description": "Rate applied once the prompt exceeds the token threshold in the field name."
},
"output_cost_per_token_above_200k_tokens_batches": {
"type": "number",
"minimum": 0,
"description": "Rate applied once the prompt exceeds the token threshold in the field name."
},
"output_cost_per_token_above_200k_tokens_priority": {
"type": "number",
"minimum": 0,

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

@ -103,6 +103,8 @@
- {id: llm.chat_completions.anthropic.basic.nonstream.cost_logged, module: llm, tier: P0, subject_endpoint: chat_completions, route: anthropic, capability: basic, streaming: nonstream, assertions: [works, cost_logged], source: "llm_translation/test_conversational_matrix_e2e.py", rationale: "Anthropic over /chat/completions: cost header and spend row agree"}
- {id: llm.chat_completions.anthropic.multi_turn.nonstream.works, module: llm, tier: P0, subject_endpoint: chat_completions, route: anthropic, capability: multi_turn, streaming: nonstream, assertions: [works], source: "llm_translation/test_conversational_matrix_e2e.py", rationale: "Anthropic tool result round trip over /chat/completions"}
- {id: llm.messages.anthropic.multi_turn.nonstream.works, module: llm, tier: P0, subject_endpoint: messages, route: anthropic, capability: multi_turn, streaming: nonstream, assertions: [works], source: "llm_translation/test_conversational_matrix_e2e.py", rationale: "Anthropic tool result round trip over /v1/messages"}
- {id: llm.messages.anthropic.upstream_stream_failure.stream.error_event, module: llm, tier: P1, subject_endpoint: messages, route: anthropic, capability: upstream_stream_failure, streaming: stream, assertions: [error_event], source: "customer report", rationale: "An upstream that hangs up mid-stream must reach Anthropic clients as an event: error frame, not an OpenAI-shaped data-only error they silently drop"}
- {id: llm.messages.anthropic.upstream_stream_failure.stream.error_status, module: llm, tier: P1, subject_endpoint: messages, route: anthropic, capability: upstream_stream_failure, streaming: stream, assertions: [error_status], source: "customer report", rationale: "An upstream that hangs up before its first byte must answer as a JSON error carrying its status, so Anthropic clients raise the status-specific error and retry on it instead of reading a 200 stream that only carries an error event"}
- {id: llm.messages.openai.basic.nonstream.works, module: llm, tier: P0, subject_endpoint: messages, route: openai, capability: basic, streaming: nonstream, assertions: [works], source: "llm_translation/test_conversational_matrix_e2e.py", rationale: "OpenAI models served on the Anthropic Messages contract"}
- {id: llm.messages.openai.basic.stream.works, module: llm, tier: P0, subject_endpoint: messages, route: openai, capability: basic, streaming: stream, assertions: [works], source: "llm_translation/test_conversational_matrix_e2e.py", rationale: "OpenAI over /v1/messages streams the Anthropic event grammar"}
- {id: llm.messages.openai.basic.nonstream.cost_logged, module: llm, tier: P0, subject_endpoint: messages, route: openai, capability: basic, streaming: nonstream, assertions: [works, cost_logged], source: "llm_translation/test_conversational_matrix_e2e.py", rationale: "OpenAI over /v1/messages: cost header and spend row agree"}

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

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

@ -598,6 +598,16 @@ class CountTokensResponse(BaseModel):
input_tokens: int
class AnthropicErrorBody(BaseModel):
type: str
message: str
class AnthropicErrorEvent(BaseModel):
type: Literal["error"]
error: AnthropicErrorBody
# ---------- mcp servers ----------

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

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

@ -5324,16 +5324,22 @@ def test_build_decode_kwargs_warns_for_unscoped_global_fallback_in_mixed_deploym
@pytest.mark.asyncio
async def test_resolve_team_from_header_defers_to_db_membership_only_without_jwt_claims():
async def test_resolve_team_from_header_accepts_db_teams_provisionally_under_fallback_even_with_jwt_claims():
"""With fallback_to_db_teams=True, an x-litellm-team-id header naming an existing
team is accepted provisionally only when the JWT carries no team claims (allowed
set empty). When the JWT does carry team claims, the header must still be validated
against them, and the flag-off behavior must keep rejecting unknown teams."""
team is accepted provisionally whether or not the JWT carries team claims; the
union of JWT teams and DB memberships is enforced by auth_builder's later
membership check. Unknown values still 403, and the flag-off behavior keeps
rejecting teams outside the JWT's allowed set."""
known_ids = frozenset({"team-from-db"})
deferred, _, _ = await _resolve_header("team-from-db", set(), True, _teams_by_id(known_ids), _team_alias_lookup_404)
assert deferred == HeaderTeam(header_value="team-from-db", team_id="team-from-db")
deferred_with_claims, _, _ = await _resolve_header(
"team-from-db", {"team-1"}, True, _teams_by_id(known_ids), _team_alias_lookup_404
)
assert deferred_with_claims == HeaderTeam(header_value="team-from-db", team_id="team-from-db")
with pytest.raises(HTTPException) as exc_info:
await _resolve_header("team-x", {"team-1", "team-2"}, True, _teams_by_id(known_ids), _team_alias_lookup_404)
assert exc_info.value.status_code == 403
@ -5849,6 +5855,7 @@ async def _run_auth_builder_with_header_team(
allowed_team_ids: set,
fake_get_team_by_alias=_team_alias_lookup_404,
route: str = "/chat/completions",
send_header: bool = True,
):
jwt_handler = JWTHandler()
jwt_handler.litellm_jwtauth = jwt_auth_config
@ -5909,7 +5916,7 @@ async def _run_auth_builder_with_header_team(
user_api_key_cache=None,
parent_otel_span=None,
proxy_logging_obj=None,
request_headers={"x-litellm-team-id": header_team_id},
request_headers={"x-litellm-team-id": header_team_id} if send_header else {},
)
@ -7283,6 +7290,130 @@ async def test_auth_builder_header_alias_under_db_fallback_keeps_the_team_allowe
assert allowed["team_id"] == "team_member"
@pytest.mark.asyncio
async def test_auth_builder_header_selects_db_membership_team_when_jwt_also_carries_a_team_claim() -> None:
"""Under fallback_to_db_teams, x-litellm-team-id may name a DB-membership
team the JWT does not claim (LIT-8656): the allowed set is the JWT teams
union the user's DB memberships, not the JWT teams alone. The flag-off
path keeps rejecting the same header against the JWT's allowed teams."""
user_object = LiteLLM_UserTable(
user_id="u_mixed",
user_role=LitellmUserRoles.INTERNAL_USER,
teams=["team_member"],
)
config = LiteLLM_JWTAuth(fallback_to_db_teams=True, team_id_jwt_field="appid")
token = {"sub": "u_mixed", "scope": "", "appid": "team_claimed"}
fake_get_team = _teams_by_id(frozenset({"team_claimed", "team_member"}))
by_membership = await _run_auth_builder_with_header_team(
config, token, "team_member", user_object, fake_get_team, {"team_claimed"}
)
assert by_membership["team_id"] == "team_member"
assert by_membership["team_object"].team_id == "team_member"
by_claim = await _run_auth_builder_with_header_team(
config, token, "team_claimed", user_object, fake_get_team, {"team_claimed"}
)
assert by_claim["team_id"] == "team_claimed"
flag_off = LiteLLM_JWTAuth(fallback_to_db_teams=False, team_id_jwt_field="appid")
with pytest.raises(HTTPException) as exc_info:
await _run_auth_builder_with_header_team(
flag_off, token, "team_member", user_object, fake_get_team, {"team_claimed"}
)
assert exc_info.value.status_code == 403
assert "JWT's allowed teams" in exc_info.value.detail
@pytest.mark.asyncio
async def test_auth_builder_header_non_member_team_is_denied_when_jwt_also_carries_a_team_claim() -> None:
"""A header naming a team the user does not belong to stays a membership
denial even when the JWT carries a team claim, and an existing but
non-member team produces the exact same 403 shape as a nonexistent one so
the response is no oracle for which team ids exist."""
user_object = LiteLLM_UserTable(
user_id="u_mixed",
user_role=LitellmUserRoles.INTERNAL_USER,
teams=["team_member"],
)
config = LiteLLM_JWTAuth(fallback_to_db_teams=True, team_id_jwt_field="appid")
token = {"sub": "u_mixed", "scope": "", "appid": "team_claimed"}
fake_get_team = _teams_by_id(frozenset({"team_claimed", "team_member", "team_other"}))
with pytest.raises(HTTPException) as outsider_exc:
await _run_auth_builder_with_header_team(
config, token, "team_other", user_object, fake_get_team, {"team_claimed"}
)
with pytest.raises(HTTPException) as missing_exc:
await _run_auth_builder_with_header_team(
config, token, "team_ghost", user_object, fake_get_team, {"team_claimed"}
)
assert outsider_exc.value.status_code == 403
assert missing_exc.value.status_code == 403
assert outsider_exc.value.detail == (
"x-litellm-team-id 'team_other' does not resolve to a team id or a unique team alias among your "
"team memberships."
)
assert missing_exc.value.detail.replace("team_ghost", "<team>") == outsider_exc.value.detail.replace(
"team_other", "<team>"
)
assert "exist" not in missing_exc.value.detail
@pytest.mark.asyncio
async def test_auth_builder_no_header_keeps_the_jwt_team_when_fallback_to_db_teams_is_on() -> None:
"""With no x-litellm-team-id header, fallback_to_db_teams must not disturb
the claim path: the JWT's own team claim still binds the request."""
user_object = LiteLLM_UserTable(
user_id="u_mixed",
user_role=LitellmUserRoles.INTERNAL_USER,
teams=["team_member"],
)
config = LiteLLM_JWTAuth(fallback_to_db_teams=True, team_id_jwt_field="appid")
token = {"sub": "u_mixed", "scope": "", "appid": "team_claimed"}
result = await _run_auth_builder_with_header_team(
config,
token,
"team_member",
user_object,
_teams_by_id(frozenset({"team_claimed", "team_member"})),
{"team_claimed"},
send_header=False,
)
assert result["team_id"] == "team_claimed"
@pytest.mark.asyncio
async def test_auth_builder_team_id_default_does_not_widen_the_header_allowed_set() -> None:
"""team_id_default fills in a team for claimless tokens but must not widen
the header's allowed set: a header naming the default team is still held
to DB membership under fallback_to_db_teams."""
user_object = LiteLLM_UserTable(
user_id="u_default",
user_role=LitellmUserRoles.INTERNAL_USER,
teams=["team_member"],
)
config = LiteLLM_JWTAuth(fallback_to_db_teams=True, team_id_default="team_default")
token = {"sub": "u_default", "scope": ""}
with pytest.raises(HTTPException) as exc_info:
await _run_auth_builder_with_header_team(
config,
token,
"team_default",
user_object,
_teams_by_id(frozenset({"team_default", "team_member"})),
set(),
)
assert exc_info.value.status_code == 403
assert exc_info.value.detail == (
"x-litellm-team-id 'team_default' does not resolve to a team id or a unique team alias among your "
"team memberships."
)
@pytest.mark.asyncio
async def test_sync_user_role_and_teams_singular_claim_only_recognized_under_flag():
"""Reading the singular team claim during sync is scoped to fallback_to_db_teams.

View file

@ -13,7 +13,7 @@ from litellm.proxy.client.exceptions import UnauthorizedError
def _load_http_mocking_responses():
"""Load the third-party `responses` package even if test collection creates
a top-level `responses` namespace package from `tests/test_litellm/responses`.
a top-level `responses` namespace package from `tests/unit/responses`.
"""
module = importlib.import_module("responses")
if hasattr(module, "activate"):

View file

@ -6,10 +6,13 @@ import pytest
from fastapi.responses import StreamingResponse
from litellm.proxy.common_request_processing import create_response
from litellm.types.utils import ModelResponse
from litellm.proxy.common_utils.sse_keepalive import (
ANTHROPIC_PING_SSE_CHUNK,
SSE_COMMENT_PING_BYTES,
advance_sse_tail,
resolve_ttft_keepalive_interval,
seal_open_sse_frame,
split_complete_sse_frames,
wrap_passthrough_sse_bytes_with_keepalive_pings,
wrap_sse_stream_with_keepalive_pings,
@ -32,6 +35,12 @@ def test_split_complete_sse_frames_holds_bytes_with_no_complete_frame():
assert split_complete_sse_frames(b"data: unterminated") == (b"", b"data: unterminated")
@pytest.mark.parametrize("chunk", [{"content": "hi"}, ModelResponse()])
def test_advance_sse_tail_ignores_a_chunk_that_is_not_sse_text(chunk: object):
assert advance_sse_tail(b"\n\n", chunk) == b"\n\n"
assert seal_open_sse_frame(advance_sse_tail(b"data: {", chunk)) == "\n" + ANTHROPIC_PING_SSE_CHUNK
@pytest.mark.asyncio
async def test_pings_fill_mid_stream_silence_and_preserve_chunk_order():
async def gappy_stream() -> AsyncGenerator[str, None]:

View file

@ -3674,7 +3674,7 @@ async def test_post_call_success_hook_contains_header_merge_failures(
@pytest.mark.asyncio
async def test_the_project_itpm_reservation_counts_the_request_off_the_event_loop(rate_limiter):
from tests.large_text import text
from tests.test_litellm.litellm_core_utils.event_loop_lag import (
from tests.unit.litellm_core_utils.event_loop_lag import (
assert_loop_stayed_free,
timed_with_loop_lags,
warm_tokenizer,

View file

@ -23,6 +23,7 @@ from starlette.datastructures import UploadFile as StarletteUploadFile
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.constants import DEFAULT_REQUEST_TIMEOUT_SECONDS
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.proxy._types import ProxyException, UserAPIKeyAuth
@ -1165,6 +1166,29 @@ def test_resolve_llm_passthrough_timeout_precedence():
assert resolve_llm_passthrough_timeout() == 6.0
def test_resolve_llm_passthrough_timeout_honors_explicit_global_request_timeout(monkeypatch: pytest.MonkeyPatch):
monkeypatch.setattr("litellm.request_timeout", 44.0, raising=False)
monkeypatch.setattr("litellm.request_timeout_explicitly_set", True, raising=False)
with patch("litellm.proxy.proxy_server.general_settings", {"pass_through_request_timeout": 6}):
assert resolve_llm_passthrough_timeout() == 44.0
assert resolve_llm_passthrough_timeout(kwargs={"stream": True}) == 44.0
assert resolve_llm_passthrough_timeout(router_timeout=120) == 120.0
assert resolve_llm_passthrough_timeout(kwargs={"stream": True}, router_stream_timeout=900) == 900.0
assert resolve_llm_passthrough_timeout(litellm_params={"timeout": 90}) == 90.0
assert resolve_llm_passthrough_timeout(kwargs={"timeout": 45}) == 45.0
def test_resolve_llm_passthrough_timeout_skips_unset_global_request_timeout(monkeypatch: pytest.MonkeyPatch):
monkeypatch.setattr("litellm.request_timeout", float(DEFAULT_REQUEST_TIMEOUT_SECONDS), raising=False)
monkeypatch.setattr("litellm.request_timeout_explicitly_set", False, raising=False)
with patch("litellm.proxy.proxy_server.general_settings", {"pass_through_request_timeout": 6}):
assert resolve_llm_passthrough_timeout() == 6.0
with patch("litellm.proxy.proxy_server.general_settings", {}):
assert resolve_llm_passthrough_timeout() == DEFAULT_PASS_THROUGH_REQUEST_TIMEOUT_SECONDS
def test_resolve_llm_passthrough_timeout_stream_timeout_precedence():
assert (
resolve_llm_passthrough_timeout(

View file

@ -162,7 +162,7 @@ async def test_interrupted_anthropic_stream_recovers_output_tokens_off_the_event
from unittest.mock import AsyncMock
from tests.large_text import text
from tests.test_litellm.litellm_core_utils.event_loop_lag import (
from tests.unit.litellm_core_utils.event_loop_lag import (
assert_loop_stayed_free,
timed_with_loop_lags,
warm_tokenizer,
@ -201,7 +201,7 @@ async def test_failed_anthropic_stream_records_partial_usage_off_the_event_loop(
from unittest.mock import AsyncMock
from tests.large_text import text
from tests.test_litellm.litellm_core_utils.event_loop_lag import (
from tests.unit.litellm_core_utils.event_loop_lag import (
assert_loop_stayed_free,
timed_with_loop_lags,
warm_tokenizer,

View file

@ -265,7 +265,7 @@ def test_close_dangling_otel_server_span_logger_raises_state_cleared_error(monke
@pytest.mark.asyncio
async def test_otel_request_validation_exception_handler_returns_422_detail():
errors = [{"loc": ["body", "model"], "msg": "field required", "type": "missing"}]
errors = [{"loc": ["body", "model"], "msg": "field required", "type": "missing", "input": {"messages": []}}]
exc = RequestValidationError(errors)
request = _make_request()
@ -273,7 +273,69 @@ async def test_otel_request_validation_exception_handler_returns_422_detail():
body = json.loads(response.body)
assert response.status_code == 422
assert normalize(body) == {"detail": exc.errors()}
assert body == {"detail": [{"type": "missing", "loc": ["body", "model"], "msg": "field required"}]}
_SUBMITTED_PASSWORD: Final = "hunter2-Sup3rSecret!"
_PASSWORD_LEAKING_ERRORS: Final = (
{
"type": "missing",
"loc": ["body", "new_password"],
"msg": "Field required",
"input": {"current_password": _SUBMITTED_PASSWORD},
},
{
"type": "value_error",
"loc": ["body", "password"],
"msg": "Value error, password cannot be set via /user/new",
"input": _SUBMITTED_PASSWORD,
"ctx": {"error": ValueError(_SUBMITTED_PASSWORD)},
},
)
_PUBLIC_ERRORS: Final = (
{"type": "missing", "loc": ["body", "new_password"], "msg": "Field required"},
{"type": "value_error", "loc": ["body", "password"], "msg": "Value error, password cannot be set via /user/new"},
)
@pytest.mark.asyncio
async def test_otel_request_validation_exception_handler_never_echoes_the_submitted_body():
"""A pydantic error carries the offending value as ``input`` (the whole body for a
``missing`` error) and input-derived values in ``ctx``; a caller who mistyped a
request holding a password must not get that password back."""
exc = RequestValidationError(list(_PASSWORD_LEAKING_ERRORS))
response = await otel_request_validation_exception_handler(request=_make_request(), exc=exc)
assert response.status_code == 422
assert json.loads(response.body) == {"detail": list(_PUBLIC_ERRORS)}
assert _SUBMITTED_PASSWORD.encode() not in response.body
@pytest.mark.asyncio
async def test_otel_request_validation_exception_handler_hands_the_span_only_the_public_errors(monkeypatch):
"""The OTEL SERVER span's error message is ``str(exc)``, which FastAPI builds from
every error dict ``input`` included, so the span gets the same public-only errors
the caller does, and keeps the traceback the original carried."""
import litellm.proxy.proxy_server as ps
fake_logger = MagicMock()
monkeypatch.setattr(ps, "open_telemetry_logger", fake_logger, raising=False)
exc = RequestValidationError(list(_PASSWORD_LEAKING_ERRORS))
try:
raise exc
except RequestValidationError as raised:
original_traceback = raised.__traceback__
request = _make_request(parent_otel_span=MagicMock())
await otel_request_validation_exception_handler(request=request, exc=exc)
(_span, span_exc, status_code) = fake_logger.record_error_attributes_on_span.call_args.args
assert status_code == 422
assert isinstance(span_exc, RequestValidationError)
assert list(span_exc.errors()) == list(_PUBLIC_ERRORS)
assert _SUBMITTED_PASSWORD not in str(span_exc)
assert span_exc.__traceback__ is original_traceback
@pytest.mark.asyncio

View file

@ -317,6 +317,25 @@ def test_claim_onboarding_link_missing_field_422(client, monkeypatch, mock_prism
assert any("password" in str(item) for item in body["detail"])
def test_claim_onboarding_link_422_never_echoes_the_submitted_password(client):
"""A body that fails validation is answered with the field path and message only;
pydantic's ``input`` (the whole submitted body for a missing field, password
included) must never come back to the caller or land in whatever logs the response."""
password = "hunter2-Sup3rSecret!"
response = client.post(
"/onboarding/claim_token",
json={"invitation_link": "abc", "password": password},
)
assert response.status_code == 422
assert password.encode() not in response.content
detail = response.json()["detail"]
assert detail[0]["loc"] == ["body", "user_id"]
assert detail[0]["msg"]
assert set(detail[0]) == {"type", "loc", "msg"}
def test_claim_onboarding_link_bad_onboarding_jwt_401(
client, monkeypatch, mock_prisma
):

View file

@ -7,6 +7,7 @@ from typing import AsyncGenerator, Callable, Final, Iterator, Literal, Optional,
from urllib.parse import unquote_plus
from unittest.mock import AsyncMock, MagicMock, patch
import anthropic
import httpx
import pytest
from fastapi import HTTPException, Request, Response, status
@ -14,6 +15,7 @@ from fastapi.responses import JSONResponse, StreamingResponse
import litellm
from litellm._uuid import uuid
from litellm.anthropic_interface.exceptions import AnthropicErrorSseFrame, anthropic_error_sse_frame
from litellm.litellm_core_utils.bug_report import (
DISABLE_ENV_VAR,
ISSUE_URL_BASE,
@ -55,6 +57,7 @@ from litellm.proxy.common_request_processing import (
sse_error_payload,
)
from litellm.proxy.common_utils.callback_utils import add_guardrail_to_applied_guardrails_header
from litellm.proxy.common_utils.sse_keepalive import ANTHROPIC_PING_SSE_CHUNK
from litellm.proxy.dd_span_tagger import DDSpanTagger
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.proxy._types import ProxyErrorTypes, ProxyException
@ -2543,6 +2546,63 @@ class TestCommonRequestProcessingHelpers:
assert response.headers["x-litellm-call-id"] == "call-8302"
assert json.loads(response.body) == {"error": {"code": 403, "message": "forbidden"}}
async def test_a_stream_that_fails_before_its_first_byte_answers_as_an_anthropic_json_error(self):
"""A /v1/messages stream whose first chunk is already the error frame has nothing
streamed yet, so the failure answers as JSON with the status the upstream gave,
the shape Anthropic clients raise their status-specific errors on"""
async def stream():
yield anthropic_error_sse_frame(status_code=503, raw_message="upstream unavailable")
yield ANTHROPIC_PING_SSE_CHUNK
generator: Final = stream()
response = await create_response(generator, "text/event-stream", {"x-litellm-call-id": "call-8609"})
assert isinstance(response, JSONResponse)
assert response.status_code == 503
assert response.headers["content-type"] == "application/json"
assert response.headers["x-litellm-call-id"] == "call-8609"
assert json.loads(response.body) == {
"type": "error",
"error": {"type": "api_error", "message": "upstream unavailable"},
}
assert generator.ag_frame is None
async def test_a_stream_that_fails_before_its_first_byte_names_the_call_when_opted_in(self):
async def stream():
yield anthropic_error_sse_frame(status_code=429, raw_message="slow down")
response = await create_response(
stream(),
"text/event-stream",
{"x-litellm-call-id": "call-8609"},
general_settings={"include_call_id_in_error_body": True},
)
assert isinstance(response, JSONResponse)
assert response.status_code == 429
assert json.loads(response.body) == {
"type": "error",
"error": {"type": "rate_limit_error", "message": "slow down", "litellm_call_id": "call-8609"},
}
async def test_an_error_event_after_a_keepalive_ping_still_streams(self):
"""Once a keepalive ping went out the headers are committed, so the error frame
streams as an event instead of turning into a JSON answer"""
async def stream():
yield ANTHROPIC_PING_SSE_CHUNK
yield anthropic_error_sse_frame(status_code=503, raw_message="upstream unavailable")
response = await create_response(stream(), "text/event-stream", {})
assert isinstance(response, StreamingResponse)
assert response.status_code == 200
assert "".join(await self.consume_stream(response)) == (
ANTHROPIC_PING_SSE_CHUNK
+ 'event: error\ndata: {"type": "error", "error": {"type": "api_error", "message": "upstream unavailable"}}\n\n'
)
async def test_create_streaming_response_disables_proxy_buffering(self):
"""Regression for #28384: every StreamingResponse create_response returns
must carry the headers that stop nginx/ingress/Envoy from buffering the
@ -9901,6 +9961,209 @@ class TestErrorLogCarriesCallId:
assert call_id in record.getMessage()
class TestAnthropicMessagesStreamErrorFrame:
"""A ``/v1/messages`` stream that fails after the headers are out has to say so with an
``event: error`` frame. Anthropic clients pick events by name, so a bare ``data:`` line is
skipped and the request looks like it ended with nothing in it"""
@staticmethod
def _sse_generator_failing_with(failure: Exception) -> AsyncGenerator[str, None]:
class FailingUpstream:
def __aiter__(self) -> "FailingUpstream":
return self
async def __anext__(self) -> object:
raise failure
ProxyLogging._callback_capabilities_cache.clear()
return ProxyBaseLLMRequestProcessing.async_sse_data_generator(
response=FailingUpstream(),
user_api_key_dict=ProxyUserAPIKeyAuth(api_key="sk-test"),
request_data={"model": "claude-sonnet-4-5"},
proxy_logging_obj=ProxyLogging(user_api_key_cache=MagicMock()),
)
@pytest.mark.parametrize(
"status_code, expected_error_type",
[
(429, "rate_limit_error"),
(529, "overloaded_error"),
(413, "request_too_large"),
(500, "api_error"),
(502, "api_error"),
(400, "invalid_request_error"),
],
)
async def test_mid_stream_failure_arrives_as_an_anthropic_error_event(
self, status_code: int, expected_error_type: str
) -> None:
class UpstreamFailure(Exception):
def __init__(self) -> None:
super().__init__("upstream stopped sending")
self.status_code: Final = status_code
frames: Final = [frame async for frame in self._sse_generator_failing_with(UpstreamFailure())]
assert len(frames) == 1
event_line, data_line, first_blank, second_blank = frames[0].split("\n")
assert isinstance(frames[0], AnthropicErrorSseFrame)
assert frames[0].status_code == status_code
assert event_line == "event: error"
assert (first_blank, second_blank) == ("", "")
payload: Final = json.loads(data_line.removeprefix("data: "))
assert payload["type"] == "error"
assert payload["error"]["type"] == expected_error_type
assert "upstream stopped sending" in payload["error"]["message"]
_CONTENT_DELTA_FRAME: Final = (
b"event: content_block_delta\n"
b'data: {"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"1\\n2\\n3"}}\n\n'
)
_TORN_DATA_LINE: Final = (
b"event: content_block_delta\n"
b'data: {"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"4'
)
_PING: Final = ANTHROPIC_PING_SSE_CHUNK.encode()
@staticmethod
def _upstream_failure(status_code: int) -> Exception:
class UpstreamFailure(Exception):
def __init__(self) -> None:
super().__init__("upstream stopped sending")
self.status_code: Final = status_code
return UpstreamFailure()
@staticmethod
def _sse_generator_cut_after(relayed: Sequence[bytes], failure: Exception) -> AsyncGenerator[str, None]:
class CutUpstream:
def __init__(self) -> None:
self._remaining: Final = iter(relayed)
def __aiter__(self) -> "CutUpstream":
return self
async def __anext__(self) -> object:
chunk: Final = next(self._remaining, None)
if chunk is None:
raise failure
return chunk
ProxyLogging._callback_capabilities_cache.clear()
return ProxyBaseLLMRequestProcessing.async_sse_data_generator(
response=CutUpstream(),
user_api_key_dict=ProxyUserAPIKeyAuth(api_key="sk-test"),
request_data={"model": "claude-sonnet-4-5"},
proxy_logging_obj=ProxyLogging(user_api_key_cache=MagicMock()),
)
@staticmethod
def _as_bytes(chunk: object) -> bytes:
if isinstance(chunk, bytes):
return chunk
assert isinstance(chunk, str)
return chunk.encode()
async def _wire_bytes(self, relayed: Sequence[bytes]) -> bytes:
stream: Final = self._sse_generator_cut_after(relayed, self._upstream_failure(500))
return b"".join([self._as_bytes(chunk) async for chunk in stream])
@staticmethod
def _error_frame_after(wire: bytes, relayed: bytes) -> bytes:
assert wire.startswith(relayed), f"the wire did not open with {relayed!r}: {wire!r}"
return wire.removeprefix(relayed)
@staticmethod
def _assert_error_frame(frame: bytes) -> None:
event_line, data_line, first_blank, second_blank = frame.split(b"\n")
assert event_line == b"event: error"
assert (first_blank, second_blank) == (b"", b"")
payload: Final = json.loads(data_line.removeprefix(b"data: "))
assert payload["type"] == "error"
assert "upstream stopped sending" in payload["error"]["message"]
@pytest.mark.parametrize(
"torn, seal",
[
(_TORN_DATA_LINE, b"\n" + _PING),
(b"event: content_bl", b"\n" + _PING),
(b"event: content_block_delta\n", _PING),
(b'event: content_block_delta\r\ndata: {"type":"content_block_delta"}\r\n', _PING),
],
ids=["mid_data_line", "mid_event_line", "after_a_complete_line", "after_a_crlf_line"],
)
async def test_a_frame_the_upstream_tore_is_closed_as_a_ping_before_the_error_event(
self, torn: bytes, seal: bytes
) -> None:
wire: Final = await self._wire_bytes((self._CONTENT_DELTA_FRAME, torn))
self._assert_error_frame(self._error_frame_after(wire, self._CONTENT_DELTA_FRAME + torn + seal))
async def test_a_cut_at_a_frame_boundary_gets_the_error_event_alone(self) -> None:
wire: Final = await self._wire_bytes((self._CONTENT_DELTA_FRAME,))
self._assert_error_frame(self._error_frame_after(wire, self._CONTENT_DELTA_FRAME))
async def test_a_torn_frame_still_raises_the_error_in_the_anthropic_sdk(self) -> None:
wire: Final = await self._wire_bytes((self._CONTENT_DELTA_FRAME, self._TORN_DATA_LINE))
def serve(request: httpx.Request) -> httpx.Response:
return httpx.Response(200, headers={"content-type": "text/event-stream"}, content=wire)
client: Final = anthropic.Anthropic(
api_key="sk-test",
base_url="http://proxy.test",
http_client=httpx.Client(transport=httpx.MockTransport(serve)),
max_retries=0,
)
with pytest.raises(anthropic.APIStatusError) as raised:
for _ in client.messages.create(
model="claude-sonnet-4-5", max_tokens=16, messages=[{"role": "user", "content": "count"}], stream=True
):
pass
body: Final = raised.value.body
assert isinstance(body, dict)
assert body["type"] == "error"
assert "upstream stopped sending" in body["error"]["message"]
async def test_a_failure_before_the_first_byte_answers_with_its_status_as_json(self) -> None:
response: Final = await create_response(
self._sse_generator_failing_with(self._upstream_failure(502)), "text/event-stream", {}
)
assert isinstance(response, JSONResponse)
assert response.status_code == 502
body: Final = json.loads(response.body)
assert body["type"] == "error"
assert body["error"]["type"] == "api_error"
assert "upstream stopped sending" in body["error"]["message"]
async def test_a_failure_before_the_first_byte_raises_with_its_status_in_the_anthropic_sdk(self) -> None:
response: Final = await create_response(
self._sse_generator_failing_with(self._upstream_failure(502)), "text/event-stream", {}
)
assert isinstance(response, JSONResponse)
def serve(request: httpx.Request) -> httpx.Response:
return httpx.Response(response.status_code, headers=dict(response.headers), content=response.body)
client: Final = anthropic.Anthropic(
api_key="sk-test",
base_url="http://proxy.test",
http_client=httpx.Client(transport=httpx.MockTransport(serve)),
max_retries=0,
)
with pytest.raises(anthropic.APIStatusError) as raised:
client.messages.create(
model="claude-sonnet-4-5", max_tokens=16, messages=[{"role": "user", "content": "count"}], stream=True
)
assert raised.value.status_code == 502
body: Final = raised.value.body
assert isinstance(body, dict)
assert body["type"] == "error"
assert "upstream stopped sending" in body["error"]["message"]
class TestStreamingContainerOwnershipRecordedBeforeDone:
"""Regression for LIT-8612: the OpenAI SDK closes the connection at
``data: [DONE]`` and starlette cancels the body task, so an ownership row

View file

@ -14974,7 +14974,7 @@ def test_settings_store_exposes_dashboard_saved_mcp_client_allowlist_to_the_mcp_
async def test_token_counter_keeps_the_event_loop_free_during_a_huggingface_count(monkeypatch):
from tests.large_text import text
from tests.test_litellm.litellm_core_utils.event_loop_lag import (
from tests.unit.litellm_core_utils.event_loop_lag import (
assert_loop_stayed_free,
timed_with_loop_lags,
warm_tokenizer,
@ -14995,7 +14995,7 @@ async def test_token_counter_loads_a_custom_tokenizer_off_the_event_loop(monkeyp
from litellm.rust_bridge._native import Tokenizer
from litellm import Router
from tests.test_litellm.litellm_core_utils.event_loop_lag import assert_loop_stayed_free, timed_with_loop_lags
from tests.unit.litellm_core_utils.event_loop_lag import assert_loop_stayed_free, timed_with_loop_lags
claude_tokenizer: Final = litellm.utils._select_tokenizer("claude-fable-5")["tokenizer"]

View file

@ -2021,7 +2021,7 @@ async def test_a_dispatched_failure_is_counted_off_the_event_loop():
from unittest.mock import AsyncMock, patch
from tests.large_text import text
from tests.test_litellm.litellm_core_utils.event_loop_lag import (
from tests.unit.litellm_core_utils.event_loop_lag import (
assert_loop_stayed_free,
timed_with_loop_lags,
warm_tokenizer,

View file

@ -0,0 +1,52 @@
from pathlib import Path
from typing import Final
from pydantic import BaseModel, TypeAdapter
import litellm
from litellm.integrations.custom_logger import CustomLogger
class DashboardField(BaseModel):
type: str
required: bool
class DashboardCallbackConfig(BaseModel):
id: str
displayName: str
logo: str
supports_key_team_logging: bool
dynamic_params: dict[str, DashboardField]
def _zerobus_config() -> DashboardCallbackConfig:
path: Final = Path(litellm.__file__).parent / "integrations" / "callback_configs.json"
configs: Final = TypeAdapter(tuple[DashboardCallbackConfig, ...]).validate_json(path.read_text())
return next(config for config in configs if config.id == "zerobus")
def test_zerobus_appears_in_the_dashboard_callback_dropdown():
"""The dropdown is served from callback_configs.json, so an entry only in the dashboard source is invisible."""
entry = _zerobus_config()
assert entry.displayName == "Databricks Zerobus"
assert entry.supports_key_team_logging is False
assert entry.dynamic_params["ZEROBUS_CLIENT_SECRET"].type == "password"
assert all(field.required is True for field in entry.dynamic_params.values())
def test_the_dropdown_logo_asset_exists():
"""A logo the dashboard cannot resolve degrades silently to a letter tile."""
logo = _zerobus_config().logo
repo_root = Path(litellm.__file__).parent.parent
asset = repo_root / "ui" / "litellm-dashboard" / "public" / "assets" / "logos" / logo
assert asset.is_file()
def test_the_dropdown_fields_are_the_env_vars_the_logger_reads():
"""Naming the fields as stored means the edit form prefills saved values instead of showing blanks."""
fields = tuple(_zerobus_config().dynamic_params)
assert fields == tuple(CustomLogger.get_callback_env_vars("zerobus"))

View file

@ -3817,6 +3817,22 @@ class TestTeamAdminEditableTeamFieldsSetting:
assert response.status_code == 422
def test_patch_422_never_echoes_the_submitted_value(self, monkeypatch):
self._as_proxy_admin(monkeypatch)
submitted = "hunter2-Sup3rSecret!"
try:
response = client.patch("/update/ui_settings", json={"team_admin_editable_team_fields": submitted})
finally:
app.dependency_overrides.clear()
assert response.status_code == 422
assert submitted.encode() not in response.content
detail = response.json()["detail"]
assert detail[0]["loc"] == ["team_admin_editable_team_fields"]
assert detail[0]["msg"]
assert set(detail[0]) == {"type", "loc", "msg"}
def test_patch_persists_and_syncs_the_list_to_general_settings(self, monkeypatch):
mock_prisma = self._as_proxy_admin(monkeypatch)
general_settings: dict = {"team_admin_editable_team_fields": []}

View file

@ -1,124 +0,0 @@
from dataclasses import astuple
from typing import Final
import pytest
import litellm
from litellm.rust_bridge.messages import route_host
pytestmark = pytest.mark.usefixtures("local_model_cost_map")
def _flag_model(monkeypatch: pytest.MonkeyPatch, name: str, **flags: bool) -> None:
monkeypatch.setitem(
litellm.model_cost,
name,
{
"litellm_provider": "anthropic",
"mode": "chat",
"input_cost_per_token": 0,
"output_cost_per_token": 0,
**flags,
},
)
def test_capabilities_come_from_the_model_map_under_the_callers_provider(monkeypatch: pytest.MonkeyPatch) -> None:
_flag_model(
monkeypatch,
"claude-test-adaptive",
supports_reasoning=True,
supports_adaptive_thinking=True,
supports_output_config=True,
supports_xhigh_reasoning_effort=True,
supports_sampling_params=False,
)
capabilities: Final = route_host.model_capabilities("anthropic/claude-test-adaptive", None)
assert capabilities.supports_adaptive_thinking
assert capabilities.supports_output_config
assert not capabilities.supports_legacy_thinking
assert not capabilities.supports_sampling_params
assert capabilities.effort_tiers.xhigh
assert not capabilities.effort_tiers.max
def test_unmapped_model_keeps_sampling_params_and_no_reasoning_features() -> None:
capabilities: Final = route_host.model_capabilities("anthropic/not-a-real-model", None)
assert capabilities.supports_sampling_params
assert not capabilities.supports_reasoning
assert not capabilities.supports_adaptive_thinking
assert not any(astuple(capabilities.effort_tiers))
@pytest.mark.parametrize(
("global_flag", "kwargs", "expected"),
[
(False, {}, False),
(True, {}, True),
(False, {"drop_params": "true"}, True),
(False, {"drop_params": "nonsense"}, False),
(False, {"drop_params": False}, False),
],
)
def test_drop_params_merges_the_global_flag_with_the_request(
monkeypatch: pytest.MonkeyPatch, global_flag: bool, kwargs: dict[str, object], expected: bool
) -> None:
monkeypatch.setattr(litellm, "drop_params", global_flag)
assert route_host.shaping("anthropic/not-a-real-model", None, kwargs)["drop_params"] is expected
@pytest.mark.parametrize(
("configured", "expected"),
[
(["tools[*].input_examples", 3, "metadata.user_id"], ("tools[*].input_examples", "metadata.user_id")),
("tools", ()),
(None, ()),
],
)
def test_additional_drop_params_keep_only_string_paths(configured: object, expected: tuple[str, ...]) -> None:
shaping: Final = route_host.shaping("anthropic/not-a-real-model", None, {"additional_drop_params": configured})
assert shaping["additional_drop_params"] == expected
def test_native_request_rejections_map_to_the_public_400() -> None:
from types import MappingProxyType
from litellm.rust_bridge.messages.entrypoints import LiteLLMMessagesRequest
request: Final = LiteLLMMessagesRequest(
model="anthropic/claude-sonnet-5",
messages=(),
max_tokens=8,
stream=None,
api_key=None,
api_base=None,
custom_llm_provider=None,
kwargs=MappingProxyType({}),
)
rejected: Final = ValueError("claude-sonnet-5 does not support top_k=5")
rejected.messages_request_error = True # pyright: ignore[reportAttributeAccessIssue] # marker the native host sets
mapped: Final = route_host.map_failure(rejected, request, "anthropic")
assert isinstance(mapped, litellm.BadRequestError)
assert mapped.status_code == 400
assert "does not support top_k=5" in mapped.message
assert mapped.model == "claude-sonnet-5"
assert not isinstance(route_host.map_failure(ValueError("plain"), request, "anthropic"), litellm.BadRequestError)
def test_stream_hidden_params_projects_upstream_headers_the_way_the_python_handler_does() -> None:
hidden: Final = route_host.stream_hidden_params(
(("request-id", "req_upstream_123"), ("x-ratelimit-remaining-requests", "41"))
)
additional: Final = hidden["additional_headers"]
assert isinstance(additional, dict)
assert additional["llm_provider-request-id"] == "req_upstream_123"
assert additional["x-ratelimit-remaining-requests"] == "41"
assert "request-id" not in additional

View file

@ -7,7 +7,7 @@ from tokenizers import Tokenizer as ReferenceTokenizer
from litellm.rust_bridge import _native
from litellm.utils import claude_json_str
from tests.test_litellm.litellm_core_utils.test_decode_special_tokens import TOKENIZER_JSON
from tests.unit.litellm_core_utils.test_decode_special_tokens import TOKENIZER_JSON
pytestmark = pytest.mark.requires_rust_extension

View file

@ -94,7 +94,7 @@ class _AgentChunk:
@pytest.mark.asyncio
async def test_stream_completion_counts_tokens_off_the_event_loop(monkeypatch):
from tests.large_text import text
from tests.test_litellm.litellm_core_utils.event_loop_lag import (
from tests.unit.litellm_core_utils.event_loop_lag import (
assert_loop_stayed_free,
timed_with_loop_lags,
warm_tokenizer,

View file

@ -469,7 +469,7 @@ class _UsageRecorder(CustomLogger):
@pytest.mark.asyncio
async def test_asend_message_counts_usage_off_the_event_loop(monkeypatch):
from tests.large_text import text
from tests.test_litellm.litellm_core_utils.event_loop_lag import (
from tests.unit.litellm_core_utils.event_loop_lag import (
assert_loop_stayed_free,
timed_with_loop_lags,
warm_tokenizer,

View file

@ -3,8 +3,15 @@ Tests for AnthropicExceptionMapping class in litellm/anthropic_interface/excepti
"""
import json
from typing import Final
from litellm.anthropic_interface.exceptions import AnthropicExceptionMapping
import pytest
from litellm.anthropic_interface.exceptions import (
AnthropicErrorSseFrame,
AnthropicExceptionMapping,
anthropic_error_sse_frame,
)
class TestCreateErrorResponse:
@ -206,3 +213,42 @@ class TestTransformToAnthropicError:
)
assert result["type"] == "error"
assert result["error"]["message"] == '["error1", "error2"]'
class TestAnthropicErrorSseFrame:
@pytest.mark.parametrize(
("status_code", "expected_error_type"),
[(429, "rate_limit_error"), (503, "api_error"), (400, "invalid_request_error")],
)
def test_the_frame_is_one_error_event_carrying_the_anthropic_envelope(
self, status_code: int, expected_error_type: str
) -> None:
frame: Final = anthropic_error_sse_frame(status_code=status_code, raw_message="upstream unavailable")
event_line, data_line, first_blank, second_blank = frame.split("\n")
assert event_line == "event: error"
assert (first_blank, second_blank) == ("", "")
assert json.loads(data_line.removeprefix("data: ")) == {
"type": "error",
"error": {"type": expected_error_type, "message": "upstream unavailable"},
}
def test_the_frame_remembers_the_status_and_body_it_was_built_from(self) -> None:
frame: Final = anthropic_error_sse_frame(status_code=503, raw_message="upstream unavailable")
assert isinstance(frame, AnthropicErrorSseFrame)
assert frame.status_code == 503
data_line: Final = frame.split("\n")[1]
assert data_line == f"data: {json.dumps(frame.json_body(call_id=None))}"
def test_the_json_body_names_the_call_only_when_asked(self) -> None:
frame: Final = anthropic_error_sse_frame(status_code=503, raw_message="upstream unavailable")
assert frame.json_body(call_id="call-1") == {
"type": "error",
"error": {"type": "api_error", "message": "upstream unavailable", "litellm_call_id": "call-1"},
}
assert frame.json_body(call_id=None) == {
"type": "error",
"error": {"type": "api_error", "message": "upstream unavailable"},
}

View file

@ -464,6 +464,40 @@ def test_total_cost_applies_the_long_context_batch_tier_per_line():
assert result.cost == pytest.approx((300_000 * 2e-6) + (10 * 6e-6) + (100 * 1e-6) + (10 * 4e-6))
def test_xai_output_lines_bill_reasoning_tokens_as_completion_tokens():
row = _success_row(
model="grok-4.3",
usage={
"prompt_tokens": 615,
"completion_tokens": 3,
"total_tokens": 993,
"completion_tokens_details": {"reasoning_tokens": 375},
},
)
result = bu._aggregate_batch_cost_usage_models(
entries=[row],
custom_llm_provider="xai",
model_info=ModelInfo(
key="xai/grok-4.3",
max_tokens=None,
max_input_tokens=None,
max_output_tokens=None,
input_cost_per_token=1.25e-6,
output_cost_per_token=2.5e-6,
litellm_provider="xai",
mode="chat",
supported_openai_params=None,
input_cost_per_token_batches=1e-6,
output_cost_per_token_batches=2e-6,
),
)
assert result.usage.completion_tokens == 378
assert result.usage.total_tokens == 993
assert result.cost == pytest.approx((615 * 1e-6) + (378 * 2e-6))
def test_total_usage_empty_is_zero():
result = bu._aggregate_batch_cost_usage_models(entries=[], custom_llm_provider="openai")
assert result.cost == 0.0

View file

@ -39,6 +39,11 @@ from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
from litellm._logging import verbose_logger
import logging
import json
import httpx
import respx
from fastapi.testclient import TestClient
from litellm.caching.caching_handler import _PENDING_CACHE_WRITES
def setup_cache():
@ -1062,6 +1067,9 @@ def test_is_chat_completion_cached_dict():
assert _is_chat_completion_cached_dict(
{"id": "other", "object": "chat.completion.chunk", "choices": []}
)
assert _is_chat_completion_cached_dict(
{"id": "no-object", "choices": [{"index": 0}]}
)
assert not _is_chat_completion_cached_dict(
{"id": "resp_abc", "object": "response", "output": []}
)
@ -1432,3 +1440,799 @@ def test_convert_cached_responses_result_parameterized(
assert result is not None
assert result.id == cached_result["id"]
assert result.status == cached_result["status"]
@pytest.mark.asyncio
async def test_process_async_embedding_cached_response():
llm_caching_handler = LLMCachingHandler(
original_function=MagicMock(),
request_kwargs={},
start_time=datetime.now(),
)
args = {
"cached_result": [
{
"embedding": [-0.025122925639152527, -0.019487135112285614],
"index": 0,
"object": "embedding",
}
]
}
mock_logging_obj = MagicMock()
mock_logging_obj.async_success_handler = AsyncMock()
response, cache_hit = llm_caching_handler._process_async_embedding_cached_response(
final_embedding_cached_response=None,
cached_result=args["cached_result"],
kwargs={"model": "text-embedding-ada-002", "input": "test"},
logging_obj=mock_logging_obj,
start_time=datetime.now(),
model="text-embedding-ada-002",
)
assert cache_hit
print(f"response: {response}")
assert len(response.data) == 1
@pytest.mark.asyncio
async def test_embedding_cache_preserves_prompt_tokens_details():
"""Test that prompt_tokens_details (including image_count) survives a full cache hit."""
llm_caching_handler = LLMCachingHandler(
original_function=MagicMock(),
request_kwargs={},
start_time=datetime.now(),
)
cached_result = [
{
"embedding": [-0.025, -0.019],
"index": 0,
"object": "embedding",
"model": "amazon.titan-embed-image-v1",
"prompt_tokens_details": {"image_count": 1},
}
]
mock_logging_obj = MagicMock()
mock_logging_obj.async_success_handler = AsyncMock()
response, cache_hit = llm_caching_handler._process_async_embedding_cached_response(
final_embedding_cached_response=None,
cached_result=cached_result,
kwargs={"model": "amazon.titan-embed-image-v1", "input": "base64imagedata"},
logging_obj=mock_logging_obj,
start_time=datetime.now(),
model="amazon.titan-embed-image-v1",
)
assert cache_hit
assert response.usage is not None
assert response.usage.prompt_tokens_details is not None
assert response.usage.prompt_tokens_details.image_count == 1
@pytest.mark.asyncio
async def test_embedding_cache_backward_compat_no_prompt_tokens_details():
"""Test that old cached items without prompt_tokens_details still work."""
llm_caching_handler = LLMCachingHandler(
original_function=MagicMock(),
request_kwargs={},
start_time=datetime.now(),
)
# Old-format cached item — no prompt_tokens_details field
cached_result = [
{
"embedding": [-0.025, -0.019],
"index": 0,
"object": "embedding",
"model": "text-embedding-ada-002",
}
]
mock_logging_obj = MagicMock()
mock_logging_obj.async_success_handler = AsyncMock()
response, cache_hit = llm_caching_handler._process_async_embedding_cached_response(
final_embedding_cached_response=None,
cached_result=cached_result,
kwargs={"model": "text-embedding-ada-002", "input": "test"},
logging_obj=mock_logging_obj,
start_time=datetime.now(),
model="text-embedding-ada-002",
)
assert cache_hit
assert response.usage is not None
assert response.usage.prompt_tokens_details is None
@pytest.mark.asyncio
async def test_embedding_cache_aggregates_multiple_image_counts():
"""Test that image_count is summed correctly across multiple cached items."""
llm_caching_handler = LLMCachingHandler(
original_function=MagicMock(),
request_kwargs={},
start_time=datetime.now(),
)
cached_result = [
{
"embedding": [-0.025, -0.019],
"index": 0,
"object": "embedding",
"model": "amazon.titan-embed-image-v1",
"prompt_tokens_details": {"image_count": 1},
},
{
"embedding": [0.031, 0.042],
"index": 1,
"object": "embedding",
"model": "amazon.titan-embed-image-v1",
"prompt_tokens_details": {"image_count": 1},
},
]
mock_logging_obj = MagicMock()
mock_logging_obj.async_success_handler = AsyncMock()
response, cache_hit = llm_caching_handler._process_async_embedding_cached_response(
final_embedding_cached_response=None,
cached_result=cached_result,
kwargs={
"model": "amazon.titan-embed-image-v1",
"input": ["img1", "img2"],
},
logging_obj=mock_logging_obj,
start_time=datetime.now(),
model="amazon.titan-embed-image-v1",
)
assert cache_hit
assert response.usage.prompt_tokens_details is not None
assert response.usage.prompt_tokens_details.image_count == 2
def test_combine_usage_merges_prompt_tokens_details():
"""Test that combine_usage merges prompt_tokens_details from both Usage objects."""
from litellm.types.utils import PromptTokensDetailsWrapper, Usage
llm_caching_handler = LLMCachingHandler(
original_function=MagicMock(),
request_kwargs={},
start_time=datetime.now(),
)
usage1 = Usage(
prompt_tokens=10,
completion_tokens=0,
total_tokens=10,
prompt_tokens_details=PromptTokensDetailsWrapper(image_count=1),
)
usage2 = Usage(
prompt_tokens=20,
completion_tokens=0,
total_tokens=20,
prompt_tokens_details=PromptTokensDetailsWrapper(image_count=2),
)
combined = llm_caching_handler.combine_usage(usage1, usage2)
assert combined.prompt_tokens == 30
assert combined.total_tokens == 30
assert combined.prompt_tokens_details is not None
assert combined.prompt_tokens_details.image_count == 3
def test_combine_usage_handles_none_details():
"""Test that combine_usage works when one or both sides have null prompt_tokens_details."""
from litellm.types.utils import PromptTokensDetailsWrapper, Usage
llm_caching_handler = LLMCachingHandler(
original_function=MagicMock(),
request_kwargs={},
start_time=datetime.now(),
)
# Both null
usage_a = Usage(prompt_tokens=10, completion_tokens=0, total_tokens=10)
usage_b = Usage(prompt_tokens=20, completion_tokens=0, total_tokens=20)
combined = llm_caching_handler.combine_usage(usage_a, usage_b)
assert combined.prompt_tokens_details is None
# Only first has details
usage_c = Usage(
prompt_tokens=10,
completion_tokens=0,
total_tokens=10,
prompt_tokens_details=PromptTokensDetailsWrapper(image_count=1),
)
combined = llm_caching_handler.combine_usage(usage_c, usage_b)
assert combined.prompt_tokens_details is not None
assert combined.prompt_tokens_details.image_count == 1
# Only second has details
combined = llm_caching_handler.combine_usage(usage_a, usage_c)
assert combined.prompt_tokens_details is not None
assert combined.prompt_tokens_details.image_count == 1
def _build_logging_obj(call_type: str, stream: bool):
import uuid as _uuid
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
return LiteLLMLogging(
litellm_call_id=str(datetime.now()),
call_type=call_type,
model="gpt-5.4",
messages=[],
function_id=str(_uuid.uuid4()),
stream=stream,
start_time=datetime.now(),
)
def test_convert_cached_responses_bridge_chat_completion_nonstream():
"""openai/responses chat-completions bridge: non-streaming cache hit replays as ModelResponse."""
from litellm import responses
from litellm.types.utils import CallTypes, ModelResponse
caching_handler = LLMCachingHandler(
original_function=responses, request_kwargs={}, start_time=datetime.now()
)
cached_result = {
"id": "chatcmpl-bridge-nonstream",
"object": "chat.completion",
"created": int(time.time()),
"model": "gpt-5.4",
"choices": [
{
"index": 0,
"message": {"role": "assistant", "content": "Hi!"},
"finish_reason": "stop",
}
],
"usage": {"prompt_tokens": 7, "completion_tokens": 11, "total_tokens": 18},
}
result = caching_handler._convert_cached_result_to_model_response(
cached_result=cached_result,
call_type=CallTypes.responses.value,
kwargs={
"model": "gpt-5.4",
"stream": False,
"messages": [{"role": "user", "content": "hi"}],
},
logging_obj=_build_logging_obj(CallTypes.responses.value, stream=False),
model="gpt-5.4",
args=(),
)
assert isinstance(result, ModelResponse)
assert result.choices[0].message.content == "Hi!"
def test_convert_cached_responses_legacy_nonstream_path():
"""Genuine ResponsesAPIResponse dict (no chatcmpl/choices) falls through legacy path."""
from litellm import responses
from litellm.types.llms.openai import ResponsesAPIResponse
from litellm.types.utils import CallTypes
caching_handler = LLMCachingHandler(
original_function=responses, request_kwargs={}, start_time=datetime.now()
)
cached_result = {
"id": "resp_legacy_nonstream",
"created_at": int(time.time()),
"status": "completed",
"model": "gpt-4o",
"object": "response",
"output": [
{
"type": "message",
"id": "msg_legacy",
"status": "completed",
"role": "assistant",
"content": [
{
"type": "output_text",
"text": "legacy response",
"annotations": [],
}
],
}
],
}
result = caching_handler._convert_cached_result_to_model_response(
cached_result=cached_result,
call_type=CallTypes.responses.value,
kwargs={"model": "gpt-4o", "input": "hi", "stream": False},
logging_obj=_build_logging_obj(CallTypes.responses.value, stream=False),
model="gpt-4o",
args=(),
)
assert isinstance(result, ResponsesAPIResponse)
assert result.id == "resp_legacy_nonstream"
def test_convert_cached_responses_legacy_stream_path():
"""Genuine ResponsesAPIResponse dict (no chatcmpl/choices) on stream falls through legacy path."""
from litellm import responses
from litellm.responses.streaming_iterator import (
CachedResponsesAPIStreamingIterator,
)
from litellm.types.utils import CallTypes
caching_handler = LLMCachingHandler(
original_function=responses, request_kwargs={}, start_time=datetime.now()
)
cached_result = {
"id": "resp_legacy_stream",
"created_at": int(time.time()),
"status": "completed",
"model": "gpt-4o",
"object": "response",
"output": [
{
"type": "message",
"id": "msg_legacy_stream",
"status": "completed",
"role": "assistant",
"content": [
{
"type": "output_text",
"text": "legacy stream",
"annotations": [],
}
],
}
],
}
result = caching_handler._convert_cached_result_to_model_response(
cached_result=cached_result,
call_type=CallTypes.responses.value,
kwargs={"model": "gpt-4o", "input": "hi", "stream": True},
logging_obj=_build_logging_obj(CallTypes.responses.value, stream=True),
model="gpt-4o",
args=(),
)
assert isinstance(result, CachedResponsesAPIStreamingIterator)
@pytest.mark.asyncio
async def test_embedding_cache_restores_stored_prompt_tokens_for_image_input():
"""Image-embedding cache hit restores prompt_tokens=0 from the stored value
instead of recomputing a bogus count by tokenizing the base64 input."""
llm_caching_handler = LLMCachingHandler(
original_function=MagicMock(),
request_kwargs={},
start_time=datetime.now(),
)
# base64-like blob — token_counter over this would return a large nonzero count
image_input = "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAQAAAC1HAwCAAAAC0lEQVR42mNk" * 50
cached_result = [
{
"embedding": [-0.025, -0.019],
"index": 0,
"object": "embedding",
"model": "amazon.titan-embed-image-v1",
"prompt_tokens": 0,
"prompt_tokens_details": {"image_count": 1},
}
]
mock_logging_obj = MagicMock()
mock_logging_obj.async_success_handler = AsyncMock()
response, cache_hit = llm_caching_handler._process_async_embedding_cached_response(
final_embedding_cached_response=None,
cached_result=cached_result,
kwargs={"model": "amazon.titan-embed-image-v1", "input": image_input},
logging_obj=mock_logging_obj,
start_time=datetime.now(),
model="amazon.titan-embed-image-v1",
)
assert cache_hit
assert response.usage is not None
assert response.usage.prompt_tokens == 0
assert response.usage.total_tokens == 0
assert response.usage.prompt_tokens_details.image_count == 1
@pytest.mark.asyncio
async def test_embedding_cache_sums_stored_prompt_tokens_across_items():
"""A multi-item cache hit sums the stored per-item prompt_tokens back to the total."""
llm_caching_handler = LLMCachingHandler(
original_function=MagicMock(),
request_kwargs={},
start_time=datetime.now(),
)
cached_result = [
{
"embedding": [-0.01],
"index": 0,
"object": "embedding",
"model": "text-embedding-3-small",
"prompt_tokens": 5,
},
{
"embedding": [-0.02],
"index": 1,
"object": "embedding",
"model": "text-embedding-3-small",
"prompt_tokens": 4,
},
]
mock_logging_obj = MagicMock()
mock_logging_obj.async_success_handler = AsyncMock()
response, cache_hit = llm_caching_handler._process_async_embedding_cached_response(
final_embedding_cached_response=None,
cached_result=cached_result,
kwargs={"model": "text-embedding-3-small", "input": ["hello world", "foo bar"]},
logging_obj=mock_logging_obj,
start_time=datetime.now(),
model="text-embedding-3-small",
)
assert cache_hit
assert response.usage.prompt_tokens == 9
assert response.usage.total_tokens == 9
@pytest.mark.asyncio
async def test_embedding_cache_falls_back_to_token_counter_for_legacy_entries():
"""Legacy cache entries with no stored prompt_tokens still recompute via token_counter
for str inputs (backward compatibility)."""
llm_caching_handler = LLMCachingHandler(
original_function=MagicMock(),
request_kwargs={},
start_time=datetime.now(),
)
# No prompt_tokens key — pre-fix entry
cached_result = [
{
"embedding": [-0.025, -0.019],
"index": 0,
"object": "embedding",
"model": "text-embedding-ada-002",
},
]
mock_logging_obj = MagicMock()
mock_logging_obj.async_success_handler = AsyncMock()
response, cache_hit = llm_caching_handler._process_async_embedding_cached_response(
final_embedding_cached_response=None,
cached_result=cached_result,
kwargs={"model": "text-embedding-ada-002", "input": "hello world"},
logging_obj=mock_logging_obj,
start_time=datetime.now(),
model="text-embedding-ada-002",
)
assert cache_hit
# token_counter over "hello world" yields a nonzero count — fallback path still runs
assert response.usage.prompt_tokens > 0
@pytest.mark.asyncio
async def test_embedding_cache_hit_sets_custom_llm_provider_on_logging_obj():
"""A full embedding cache hit must stamp the resolved provider onto the logging
obj so spend logs record the provider instead of None/unknown."""
from litellm.types.utils import CallTypes
llm_caching_handler = LLMCachingHandler(
original_function=MagicMock(),
request_kwargs={},
start_time=datetime.now(),
)
cached_result = [
{
"embedding": [-0.025, -0.019],
"index": 0,
"object": "embedding",
"model": "text-embedding-3-small",
"prompt_tokens": 5,
}
]
logging_obj = _build_logging_obj(CallTypes.aembedding.value, stream=False)
logging_obj.async_success_handler = AsyncMock()
response, cache_hit = llm_caching_handler._process_async_embedding_cached_response(
final_embedding_cached_response=None,
cached_result=cached_result,
kwargs={"model": "text-embedding-3-small", "input": "hello world"},
logging_obj=logging_obj,
start_time=datetime.now(),
model="text-embedding-3-small",
)
assert cache_hit
assert logging_obj.model_call_details["custom_llm_provider"] == "openai"
def test_sync_stream_responses_cache_hit_sets_custom_llm_provider_on_logging_obj(monkeypatch):
import litellm
from litellm.caching.caching import Cache
from litellm.types.utils import CallTypes
monkeypatch.setattr(litellm, "cache", Cache(type="local"))
kwargs = {"model": "azure/gpt-5.4-mini", "input": "hello", "stream": True}
cached_response = {
"id": "resp_sync_stream",
"created_at": int(time.time()),
"status": "completed",
"model": "gpt-5.4-mini",
"object": "response",
"output": [
{
"type": "message",
"id": "msg_sync_stream",
"status": "completed",
"role": "assistant",
"content": [{"type": "output_text", "text": "hi", "annotations": []}],
}
],
}
litellm.cache.add_cache(json.dumps(cached_response), **kwargs)
handler = LLMCachingHandler(original_function=litellm.responses, request_kwargs=kwargs, start_time=datetime.now())
logging_obj = _build_logging_obj(CallTypes.responses.value, stream=True)
hit = handler._sync_get_cache(
model="azure/gpt-5.4-mini",
original_function=litellm.responses,
logging_obj=logging_obj,
start_time=datetime.now(),
call_type=CallTypes.responses.value,
kwargs=kwargs,
args=(),
)
assert hit.cached_result is not None
assert logging_obj.model_call_details["custom_llm_provider"] == "azure"
assert logging_obj.model_call_details["litellm_params"]["custom_llm_provider"] == "azure"
def test_request_kwargs_does_not_retain_logging_obj():
"""
The caching handler lives on logging_obj._llm_caching_handler, so keeping
litellm_logging_obj inside request_kwargs closes a reference cycle
(Logging -> LLMCachingHandler -> kwargs -> Logging). That cycle keeps the
full request payload alive until a generational GC pass instead of being
freed by refcount when the request finishes; under bursts of large-token
requests this presents as stepwise RSS growth that never returns to
baseline. Other kwargs (messages included) must be preserved.
"""
logging_obj = MagicMock()
kwargs = {
"model": "gpt-4o",
"messages": [{"role": "user", "content": "hello"}],
"litellm_logging_obj": logging_obj,
}
handler = LLMCachingHandler(
original_function=MagicMock(),
request_kwargs=kwargs,
start_time=datetime.now(),
)
assert "litellm_logging_obj" not in handler.request_kwargs
assert handler.request_kwargs["messages"] == kwargs["messages"]
assert handler.request_kwargs["model"] == "gpt-4o"
def test_async_cache_write_completes_when_asyncio_run_closes_the_loop(monkeypatch):
"""
Regression test for the SDK losing async cache writes in short-lived scripts:
async_set_cache dispatched the write as a bare fire-and-forget task, so
asyncio.run cancelled it at loop close before the write landed (LIT-6184,
deterministic with hiredis installed). The write must survive loop shutdown.
"""
import litellm
writes = []
class _SlowWriteCache:
supported_call_types = ["acompletion"]
cache = None
async def async_add_cache(self, result, dynamic_cache_object=None, **kwargs):
await asyncio.sleep(0.2)
writes.append(result)
async def acompletion(**kwargs):
return None
handler = LLMCachingHandler(
original_function=acompletion,
request_kwargs={},
start_time=datetime.now(),
)
monkeypatch.setattr(litellm, "cache", _SlowWriteCache())
async def _short_lived_script():
await handler.async_set_cache(
result=litellm.ModelResponse(),
original_function=acompletion,
kwargs={},
)
asyncio.run(_short_lived_script())
assert len(writes) == 1
@pytest.mark.asyncio
async def test_cache_hit_records_the_looked_up_key_as_the_preset_cache_key(monkeypatch):
"""The spend log for a cache hit must reuse the key the lookup already computed instead of hashing again."""
import litellm
from litellm.caching.caching import Cache
from litellm.types.utils import CallTypes
async def acompletion(**kwargs):
return None
monkeypatch.setattr(litellm, "cache", Cache(type="local"))
kwargs = {"model": "gpt-5.4", "messages": [{"role": "user", "content": "hello"}], "caching": True}
await litellm.cache.async_add_cache(
litellm.ModelResponse(choices=[{"message": {"role": "assistant", "content": "hi"}}]), **kwargs
)
handler = LLMCachingHandler(original_function=acompletion, request_kwargs=kwargs, start_time=datetime.now())
logging_obj = _build_logging_obj(CallTypes.acompletion.value, stream=False)
logging_obj.async_success_handler = AsyncMock()
hit = await handler._async_get_cache(
model="gpt-5.4",
original_function=acompletion,
logging_obj=logging_obj,
start_time=datetime.now(),
call_type=CallTypes.acompletion.value,
kwargs=kwargs,
args=(),
)
assert hit is not None and hit.cached_result is not None
assert handler.preset_cache_key is not None
assert logging_obj.litellm_params["preset_cache_key"] == handler.preset_cache_key
assert hit.cached_result._hidden_params["cache_key"] == handler.preset_cache_key
@pytest.mark.asyncio
async def test_converted_stream_cache_hit_replayed_as_plain_object_logs_at_hit_time(monkeypatch):
import litellm
from litellm.caching.caching import Cache
from litellm.types.utils import CallTypes
async def aanthropic_messages(**kwargs):
return None
monkeypatch.setattr(litellm, "cache", Cache(type="local"))
kwargs = {
"model": "claude-sonnet-5",
"messages": [{"role": "user", "content": "hello"}],
"max_tokens": 16,
"caching": True,
"stream": False,
"_websearch_interception_converted_stream": True,
}
cached_message = {
"id": "msg_1",
"type": "message",
"role": "assistant",
"content": [{"type": "text", "text": "hi"}],
}
await litellm.cache.async_add_cache(cached_message, **kwargs)
handler = LLMCachingHandler(original_function=aanthropic_messages, request_kwargs=kwargs, start_time=datetime.now())
logging_obj = _build_logging_obj(CallTypes.aanthropic_messages.value, stream=False)
logging_obj.async_success_handler = AsyncMock()
logging_obj.handle_sync_success_callbacks_for_async_calls = MagicMock()
hit = await handler._async_get_cache(
model="claude-sonnet-5",
original_function=aanthropic_messages,
logging_obj=logging_obj,
start_time=datetime.now(),
call_type=CallTypes.aanthropic_messages.value,
kwargs=kwargs,
args=(),
)
assert hit is not None and hit.cached_result == cached_message
logging_obj.handle_sync_success_callbacks_for_async_calls.assert_called_once()
assert logging_obj.handle_sync_success_callbacks_for_async_calls.call_args.kwargs["cache_hit"] is True
@pytest.mark.asyncio
async def test_agentic_loop_followup_cache_hit_with_converted_stream_marker_replays_as_plain_object(monkeypatch):
import litellm
from litellm.caching.caching import Cache
from litellm.types.utils import CallTypes
async def acompletion(**kwargs):
return None
monkeypatch.setattr(litellm, "cache", Cache(type="local"))
kwargs = {
"model": "gpt-5.6",
"messages": [{"role": "user", "content": "run the code"}],
"caching": True,
"stream": False,
"_code_interpreter_interception_converted_stream": True,
"_agentic_loop_depth": 1,
}
await litellm.cache.async_add_cache(
litellm.ModelResponse(choices=[{"message": {"role": "assistant", "content": "done"}}]), **kwargs
)
handler = LLMCachingHandler(original_function=acompletion, request_kwargs=kwargs, start_time=datetime.now())
logging_obj = _build_logging_obj(CallTypes.acompletion.value, stream=False)
logging_obj.async_success_handler = AsyncMock()
logging_obj.handle_sync_success_callbacks_for_async_calls = MagicMock()
hit = await handler._async_get_cache(
model="gpt-5.6",
original_function=acompletion,
logging_obj=logging_obj,
start_time=datetime.now(),
call_type=CallTypes.acompletion.value,
kwargs=kwargs,
args=(),
)
assert hit is not None and isinstance(hit.cached_result, litellm.ModelResponse)
assert hit.cached_result.choices[0].message.content == "done"
logging_obj.handle_sync_success_callbacks_for_async_calls.assert_called_once()
assert logging_obj.handle_sync_success_callbacks_for_async_calls.call_args.kwargs["cache_hit"] is True
@pytest.mark.asyncio
async def test_partial_embedding_cache_hit_sends_only_misses_and_keeps_input_order(monkeypatch):
import litellm
from litellm import CustomLLM
from litellm.caching.caching import Cache
from litellm.types.utils import Embedding, EmbeddingResponse
class RecordingEmbedder(CustomLLM):
provider_inputs: tuple[tuple[str, ...], ...] = ()
async def aembedding(self, model, input, model_response, **kwargs) -> EmbeddingResponse:
self.provider_inputs = (*self.provider_inputs, tuple(input))
return EmbeddingResponse(
model=model,
data=[
Embedding(embedding=[float(len(text))], index=idx, object="embedding")
for idx, text in enumerate(input)
],
)
embedder = RecordingEmbedder()
monkeypatch.setattr(litellm, "custom_provider_map", [{"provider": "recording-embedder", "custom_handler": embedder}])
monkeypatch.setattr(litellm, "provider_list", [*litellm.provider_list, "recording-embedder"])
monkeypatch.setattr(litellm, "_custom_providers", [*litellm._custom_providers, "recording-embedder"])
monkeypatch.setattr(litellm, "cache", Cache(type="local"))
await litellm.aembedding(model="recording-embedder/m", input=["aa", "bbbb"])
await asyncio.gather(*_PENDING_CACHE_WRITES)
mixed_input = ["c", "aa", "ddd", "bbbb", "eeeee"]
response = await litellm.aembedding(model="recording-embedder/m", input=mixed_input)
await asyncio.gather(*_PENDING_CACHE_WRITES)
assert embedder.provider_inputs == (("aa", "bbbb"), ("c", "ddd", "eeeee")), embedder.provider_inputs
assert [item["index"] for item in response.data] == [0, 1, 2, 3, 4]
assert [item["embedding"] for item in response.data] == [[float(len(text))] for text in mixed_input]
assert response._hidden_params["cache_hit"] is True, "a partial hit must still be reported as a cache hit"
repeat = await litellm.aembedding(model="recording-embedder/m", input=mixed_input)
assert len(embedder.provider_inputs) == 2, embedder.provider_inputs
assert [item["embedding"] for item in repeat.data] == [[float(len(text))] for text in mixed_input]

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