mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
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:
commit
babf5bc3b4
353 changed files with 9916 additions and 3993 deletions
|
|
@ -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
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
8
.github/merge-smoke-tests.json
vendored
8
.github/merge-smoke-tests.json
vendored
|
|
@ -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"
|
||||
}
|
||||
}
|
||||
|
|
|
|||
12
.github/workflows/test-redis-compat.yml
vendored
12
.github/workflows/test-redis-compat.yml
vendored
|
|
@ -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 \
|
||||
|
|
|
|||
6
.github/workflows/test-rust.yml
vendored
6
.github/workflows/test-rust.yml
vendored
|
|
@ -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
|
||||
|
|
|
|||
10
.github/workflows/test-unit.yml
vendored
10
.github/workflows/test-unit.yml
vendored
|
|
@ -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
|
||||
|
|
|
|||
6
Makefile
6
Makefile
|
|
@ -301,7 +301,7 @@ test-rust-extension:
|
|||
UV_PROJECT_ENVIRONMENT="$$temporary/venv" $(UV) sync --python 3.12 --frozen --no-install-project --all-groups --all-extras && \
|
||||
$(UV) pip install --python "$$temporary/venv/bin/python" --no-deps "$$1" && \
|
||||
"$$temporary/venv/bin/python" -I -m mypy.stubtest \
|
||||
--mypy-config-file tests/test_litellm/rust_bridge/stubtest.ini \
|
||||
--mypy-config-file tests/unit/rust_bridge/stubtest.ini \
|
||||
litellm.rust_bridge._native && \
|
||||
LITELLM_RUST=1 LITELLM_LOCAL_MODEL_COST_MAP=True \
|
||||
"$$temporary/venv/bin/python" -I -m pytest --import-mode=importlib -m requires_rust_extension tests/test_litellm_rust
|
||||
|
|
@ -329,10 +329,10 @@ test-unit-integrations: install-test-deps
|
|||
$(UV_RUN) pytest tests/unit/integrations --tb=short -vv -n 4 --durations=20
|
||||
|
||||
test-unit-core-utils: install-test-deps
|
||||
$(UV_RUN) pytest tests/test_litellm/litellm_core_utils --tb=short -vv -n 2 --durations=20
|
||||
$(UV_RUN) pytest tests/unit/litellm_core_utils --tb=short -vv -n 2 --durations=20
|
||||
|
||||
test-unit-other: install-test-deps
|
||||
$(UV_RUN) pytest tests/test_litellm/caching tests/test_litellm/responses tests/unit/secret_managers tests/unit/vector_stores tests/unit/a2a_protocol tests/test_litellm/anthropic_interface tests/unit/completion_extras tests/unit/containers tests/unit/enterprise tests/unit/experimental_mcp_client tests/unit/google_genai tests/unit/images tests/unit/interactions tests/test_litellm/interactions tests/test_litellm/passthrough tests/test_litellm/router_strategy tests/test_litellm/router_utils tests/unit/types --tb=short -vv -n 4 --durations=20
|
||||
$(UV_RUN) pytest tests/unit/caching tests/unit/responses tests/unit/secret_managers tests/unit/vector_stores tests/unit/a2a_protocol tests/test_litellm/anthropic_interface tests/unit/completion_extras tests/unit/containers tests/unit/enterprise tests/unit/experimental_mcp_client tests/unit/google_genai tests/unit/images tests/unit/interactions tests/test_litellm/interactions tests/test_litellm/passthrough tests/unit/router_strategy tests/unit/router_utils tests/unit/types --tb=short -vv -n 4 --durations=20
|
||||
|
||||
test-unit-root: install-test-deps
|
||||
$(UV_RUN) pytest tests/unit/test_*.py tests/test_litellm/test_*.py --tb=short -vv -n 4 --durations=20
|
||||
|
|
|
|||
|
|
@ -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");
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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 = (
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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="",
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
5
litellm/integrations/zerobus/__init__.py
Normal file
5
litellm/integrations/zerobus/__init__.py
Normal file
|
|
@ -0,0 +1,5 @@
|
|||
"""Databricks Zerobus logging integration for LiteLLM."""
|
||||
|
||||
from litellm.integrations.zerobus.logger import ZerobusLogger
|
||||
|
||||
__all__ = ("ZerobusLogger",)
|
||||
161
litellm/integrations/zerobus/client.py
Normal file
161
litellm/integrations/zerobus/client.py
Normal 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)
|
||||
230
litellm/integrations/zerobus/logger.py
Normal file
230
litellm/integrations/zerobus/logger.py
Normal 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)
|
||||
156
litellm/integrations/zerobus/row.py
Normal file
156
litellm/integrations/zerobus/row.py
Normal 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"),
|
||||
}
|
||||
)
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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":
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
152
litellm/litellm_core_utils/sentry_scrubbing.py
Normal file
152
litellm/litellm_core_utils/sentry_scrubbing.py
Normal 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"),
|
||||
)
|
||||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
195
litellm/llms/xai/batches/handler.py
Normal file
195
litellm/llms/xai/batches/handler.py
Normal 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))
|
||||
278
litellm/llms/xai/batches/transformation.py
Normal file
278
litellm/llms/xai/batches/transformation.py
Normal 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()
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
247
litellm/llms/xai/files/transformation.py
Normal file
247
litellm/llms/xai/files/transformation.py
Normal 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
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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}'."
|
||||
),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
13
litellm/proxy/common_utils/validation_error_body.py
Normal file
13
litellm/proxy/common_utils/validation_error_body.py
Normal 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)
|
||||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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}"
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
53
litellm/types/integrations/zerobus.py
Normal file
53
litellm/types/integrations/zerobus.py
Normal 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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
]
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
|
|
|
|||
|
|
@ -86,6 +86,7 @@ LlmCapability = Literal[
|
|||
"tool_search",
|
||||
"tool_search_history",
|
||||
"tool_use",
|
||||
"upstream_stream_failure",
|
||||
"vision",
|
||||
"web_search",
|
||||
"web_search_server_tool",
|
||||
|
|
|
|||
|
|
@ -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}'
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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 ----------
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 = (
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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"):
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
52
tests/test_litellm/proxy/test_zerobus_dashboard_config.py
Normal file
52
tests/test_litellm/proxy/test_zerobus_dashboard_config.py
Normal 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"))
|
||||
|
|
@ -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": []}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"},
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
Loading…
Add table
Reference in a new issue