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

This commit is contained in:
yucheng 2026-09-26 00:27:58 +00:00
commit d12afe1778
279 changed files with 3498 additions and 3288 deletions

View file

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

View file

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

View file

@ -7,9 +7,9 @@
"MODEL-DENY": "tests/test_litellm/proxy/auth/test_auth_checks.py::test_can_object_call_model_denials_return_forbidden[key-key_model_access_denied]",
"COST-EXPLICIT": "tests/unit/test_cost_calculator.py::test_completion_cost_charges_explicit_per_token_rates_over_registered_ones",
"COST-ZERO": "tests/unit/test_cost_calculator.py::test_completion_cost_is_zero_when_explicit_rates_are_zero",
"LOG-CONTENT-ON": "tests/test_litellm/litellm_core_utils/test_litellm_logging.py::test_standard_logging_payload_keeps_message_content_when_message_logging_is_on",
"LOG-CONTENT-OFF": "tests/test_litellm/litellm_core_utils/test_litellm_logging.py::test_standard_logging_payload_redacts_message_content_when_message_logging_is_off",
"CALLBACK-SUCCESS": "tests/test_litellm/litellm_core_utils/test_litellm_logging.py::test_async_success_handler_delivers_standard_logging_payload_to_custom_logger",
"CALLBACK-FAILURE": "tests/test_litellm/litellm_core_utils/test_litellm_logging.py::test_async_failure_handler_delivers_failure_payload_to_custom_logger"
"LOG-CONTENT-ON": "tests/unit/litellm_core_utils/test_litellm_logging.py::test_standard_logging_payload_keeps_message_content_when_message_logging_is_on",
"LOG-CONTENT-OFF": "tests/unit/litellm_core_utils/test_litellm_logging.py::test_standard_logging_payload_redacts_message_content_when_message_logging_is_off",
"CALLBACK-SUCCESS": "tests/unit/litellm_core_utils/test_litellm_logging.py::test_async_success_handler_delivers_standard_logging_payload_to_custom_logger",
"CALLBACK-FAILURE": "tests/unit/litellm_core_utils/test_litellm_logging.py::test_async_failure_handler_delivers_failure_payload_to_custom_logger"
}
}

View file

@ -12,9 +12,9 @@ on:
- "litellm/caching/evicted_client_closer.py"
- "tests/unit/test_redis.py"
- "tests/local_testing/test_caching.py"
- "tests/test_litellm/caching/test_redis_connection_pool.py"
- "tests/test_litellm/caching/test_redis_cluster_cache.py"
- "tests/test_litellm/caching/test_evicted_client_closer.py"
- "tests/unit/caching/test_redis_connection_pool.py"
- "tests/unit/caching/test_redis_cluster_cache.py"
- "tests/unit/caching/test_evicted_client_closer.py"
- ".github/workflows/test-redis-compat.yml"
- "pyproject.toml"
- "uv.lock"
@ -85,9 +85,9 @@ jobs:
redis-server --version
uv run --no-sync pytest \
tests/unit/test_redis.py \
tests/test_litellm/caching/test_redis_connection_pool.py \
tests/test_litellm/caching/test_redis_cluster_cache.py \
tests/test_litellm/caching/test_evicted_client_closer.py \
tests/unit/caching/test_redis_connection_pool.py \
tests/unit/caching/test_redis_cluster_cache.py \
tests/unit/caching/test_evicted_client_closer.py \
tests/local_testing/test_caching.py::test_sync_cluster_authenticates_with_azure_credentials \
tests/local_testing/test_caching.py::test_sync_cluster_authenticates_with_gcp_credentials \
--tb=short -vv \

View file

@ -24,7 +24,7 @@ on:
- ".github/actions/setup-uv-with-retries/**"
- ".github/scripts/smoke_test_native_wheel.py"
- ".github/scripts/verify_linux_native_wheel.py"
- "tests/test_litellm/rust_bridge/native_route_wheel_test.py"
- "tests/unit/rust_bridge/native_route_wheel_test.py"
- ".github/workflows/test-rust.yml"
pull_request:
branches:
@ -52,7 +52,7 @@ on:
- ".github/actions/setup-uv-with-retries/**"
- ".github/scripts/smoke_test_native_wheel.py"
- ".github/scripts/verify_linux_native_wheel.py"
- "tests/test_litellm/rust_bridge/native_route_wheel_test.py"
- "tests/unit/rust_bridge/native_route_wheel_test.py"
- ".github/workflows/test-rust.yml"
permissions:
@ -171,7 +171,7 @@ jobs:
env:
RELEASE_WHEEL_COMMIT_SHA: ${{ github.event.pull_request.head.sha || github.sha }}
- run: python tests/test_litellm/rust_bridge/native_route_wheel_test.py dist/*.whl
- run: python tests/unit/rust_bridge/native_route_wheel_test.py dist/*.whl
- name: Run pytest tests/test_litellm_rust with the compiled extension
run: make test-rust-extension

View file

@ -62,6 +62,7 @@ jobs:
- shard: core-utils
artifact-name: core-utils
test-path: "tests/test_litellm/litellm_core_utils"
unit-flag: core-utils
workers: 2
reruns: 1
timeout-minutes: 20
@ -69,9 +70,7 @@ jobs:
- shard: enterprise-routing
artifact-name: enterprise-routing
test-path: >-
tests/test_litellm/router_utils
tests/test_litellm/router_strategy
test-path: ""
unit-flag: enterprise-routing
workers: 2
reruns: 2
@ -111,7 +110,6 @@ jobs:
tests/test_litellm/interactions
tests/test_litellm/ocr
tests/test_litellm/passthrough
tests/test_litellm/rust_bridge
tests/test_litellm/test_*.py
unit-flag: misc
workers: 2
@ -228,9 +226,7 @@ jobs:
- shard: responses-caching-types
artifact-name: responses-caching-types
test-path: >-
tests/test_litellm/responses
tests/test_litellm/caching
test-path: ""
unit-flag: responses-caching-types
workers: 2
reruns: 2

View file

@ -301,7 +301,7 @@ test-rust-extension:
UV_PROJECT_ENVIRONMENT="$$temporary/venv" $(UV) sync --python 3.12 --frozen --no-install-project --all-groups --all-extras && \
$(UV) pip install --python "$$temporary/venv/bin/python" --no-deps "$$1" && \
"$$temporary/venv/bin/python" -I -m mypy.stubtest \
--mypy-config-file tests/test_litellm/rust_bridge/stubtest.ini \
--mypy-config-file tests/unit/rust_bridge/stubtest.ini \
litellm.rust_bridge._native && \
LITELLM_RUST=1 LITELLM_LOCAL_MODEL_COST_MAP=True \
"$$temporary/venv/bin/python" -I -m pytest --import-mode=importlib -m requires_rust_extension tests/test_litellm_rust
@ -329,10 +329,10 @@ test-unit-integrations: install-test-deps
$(UV_RUN) pytest tests/unit/integrations --tb=short -vv -n 4 --durations=20
test-unit-core-utils: install-test-deps
$(UV_RUN) pytest tests/test_litellm/litellm_core_utils --tb=short -vv -n 2 --durations=20
$(UV_RUN) pytest tests/unit/litellm_core_utils --tb=short -vv -n 2 --durations=20
test-unit-other: install-test-deps
$(UV_RUN) pytest tests/test_litellm/caching tests/test_litellm/responses tests/unit/secret_managers tests/unit/vector_stores tests/unit/a2a_protocol tests/test_litellm/anthropic_interface tests/unit/completion_extras tests/unit/containers tests/unit/enterprise tests/unit/experimental_mcp_client tests/unit/google_genai tests/unit/images tests/unit/interactions tests/test_litellm/interactions tests/test_litellm/passthrough tests/test_litellm/router_strategy tests/test_litellm/router_utils tests/unit/types --tb=short -vv -n 4 --durations=20
$(UV_RUN) pytest tests/unit/caching tests/unit/responses tests/unit/secret_managers tests/unit/vector_stores tests/unit/a2a_protocol tests/test_litellm/anthropic_interface tests/unit/completion_extras tests/unit/containers tests/unit/enterprise tests/unit/experimental_mcp_client tests/unit/google_genai tests/unit/images tests/unit/interactions tests/test_litellm/interactions tests/test_litellm/passthrough tests/unit/router_strategy tests/unit/router_utils tests/unit/types --tb=short -vv -n 4 --durations=20
test-unit-root: install-test-deps
$(UV_RUN) pytest tests/unit/test_*.py tests/test_litellm/test_*.py --tb=short -vv -n 4 --durations=20

View file

@ -33,7 +33,7 @@ mod test_support {
use crate::{LegacyLogging, LegacySurface, PublicCall};
/// The parameters of every `callbacks_legacy_python` function, as the real module declares them.
/// `tests/test_litellm/rust_bridge/test_callbacks_legacy_python.py` pins this file to the Python
/// `tests/unit/rust_bridge/test_callbacks_legacy_python.py` pins this file to the Python
/// signatures, and [`namespace`] binds every fake call against it.
pub(crate) const PYTHON_CONTRACT: &str = include_str!("../python_contract.json");

View file

@ -30,7 +30,7 @@ The HashiCorp Vault backend is enabled with the `hashicorp` feature and reads KV
Native backends consistently distinguish absence from failure instead of swallowing provider errors. Python-compatible resolution maps these results back to the Python handler contract before applying fallback
`hosted_keys` excludes a name for every backend. Python's handler recognizes Azure `SecretClient` and Google `KeyManagementServiceClient` instances before the `local` branch, allowing excluded names to reach those providers. Rust treats that as a routing bug. `test_rust_hosted_keys_exclude_azure_sdk_clients_too` in `tests/test_litellm/rust_bridge/ocr/test_secrets.py` pins this behavior
`hosted_keys` excludes a name for every backend. Python's handler recognizes Azure `SecretClient` and Google `KeyManagementServiceClient` instances before the `local` branch, allowing excluded names to reach those providers. Rust treats that as a routing bug. `test_rust_hosted_keys_exclude_azure_sdk_clients_too` in `tests/unit/rust_bridge/ocr/test_secrets.py` pins this behavior
Google rejects malformed base64 and mismatched CRC32C values instead of accepting corrupted payloads. Python currently ignores the checksum and uses permissive base64 decoding. Rust follows [RFC 4648](https://www.rfc-editor.org/rfc/rfc4648#section-3.3) and [Google's integrity guidance](https://docs.cloud.google.com/secret-manager/docs/data-integrity); `failed_or_missing_reads_are_not_cached` covers rejection and recovery

View file

@ -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,6 +39,7 @@ from litellm.constants import (
DEFAULT_MOCK_RESPONSE_PROMPT_TOKEN_COUNT,
EMPTY_MAPPING,
PROVIDER_REQUEST_ID_HEADERS,
REDACTED_BY_LITELLM,
)
from litellm.cost_calculator import (
RealtimeAPITokenUsageProcessor,
@ -1358,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", ""),
@ -1369,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),
@ -1476,7 +1469,7 @@ class Logging(LiteLLMLoggingBaseClass):
def _print_llm_call_debugging_log(
self,
api_base: str,
headers: dict,
headers: dict | None,
additional_args: dict,
):
"""
@ -1485,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
@ -1563,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(

View file

@ -77,13 +77,13 @@ class _CallerHeadersView(TypedDict):
headers: ReadOnly[dict[str, str]]
# Globally-routable IPs that are cloud-internal. Everything else
# non-public is caught by ``not ip.is_global`` (RFC 6890, as implemented by
# Python's ``ipaddress`` module). This list only holds IPs that are
# publicly routable *and* point to cloud-fabric services reachable from
# inside a VM via special in-fabric routing.
# Cloud-internal IPs that ``ip.is_global`` can report as public. Everything
# else non-public is caught by ``not ip.is_global`` (RFC 6890, as implemented
# by Python's ``ipaddress`` module). Older Python patch releases (3.12.2, for
# one) treat most of 192.0.0.0/24 as global, so it is listed to block it everywhere.
_CLOUD_METADATA_EXCEPTIONS: Final = [
ip_network("168.63.129.16/32"), # Azure Wire Server
ip_network("192.0.0.0/24"),
]
_ALLOWED_SCHEMES: Final = ("http", "https")

View file

@ -60264,13 +60264,16 @@
},
"fireworks_ai/accounts/fireworks/routers/deepseek-v4p1-flash-us": {
"cache_read_input_token_cost": 9e-09,
"cache_read_input_token_cost_priority": 1.125e-08,
"input_cost_per_token": 4.5e-07,
"input_cost_per_token_priority": 5.625e-07,
"litellm_provider": "fireworks_ai",
"max_input_tokens": 1048576,
"max_output_tokens": 393216,
"max_tokens": 393216,
"mode": "chat",
"output_cost_per_token": 1.8e-06,
"output_cost_per_token_priority": 2.25e-06,
"source": "https://docs.fireworks.ai/serverless/pricing",
"supports_function_calling": true,
"supports_prompt_caching": true,
@ -60360,13 +60363,16 @@
},
"fireworks_ai/deepseek-v4p1-flash-us": {
"cache_read_input_token_cost": 9e-09,
"cache_read_input_token_cost_priority": 1.125e-08,
"input_cost_per_token": 4.5e-07,
"input_cost_per_token_priority": 5.625e-07,
"litellm_provider": "fireworks_ai",
"max_input_tokens": 1048576,
"max_output_tokens": 393216,
"max_tokens": 393216,
"mode": "chat",
"output_cost_per_token": 1.8e-06,
"output_cost_per_token_priority": 2.25e-06,
"source": "https://docs.fireworks.ai/serverless/pricing",
"supports_function_calling": true,
"supports_prompt_caching": true,

View file

@ -0,0 +1,13 @@
from collections.abc import Sequence
from typing_extensions import ReadOnly, TypedDict
class ValidationErrorDetail(TypedDict):
type: ReadOnly[str]
loc: ReadOnly[tuple[int | str, ...]]
msg: ReadOnly[str]
def public_validation_errors(errors: Sequence[ValidationErrorDetail]) -> tuple[ValidationErrorDetail, ...]:
return tuple(ValidationErrorDetail(type=error["type"], loc=error["loc"], msg=error["msg"]) for error in errors)

View file

@ -8,8 +8,8 @@ from fastapi import Request
from fastapi.dependencies.utils import get_flat_params
from fastapi.params import ParamTypes
from fastapi.responses import JSONResponse
from typing_extensions import ReadOnly, TypedDict
from litellm.proxy.common_utils.validation_error_body import ValidationErrorDetail
from litellm.types.proxy.management_endpoints.management_v1 import (
ListLinks,
PageLinks,
@ -58,14 +58,6 @@ def escape_like(value: str) -> str:
return value.replace("\\", "\\\\").replace("%", "\\%").replace("_", "\\_")
class ValidationErrorDetail(TypedDict):
"""The keys of a pydantic/FastAPI validation error a problem document needs."""
type: ReadOnly[str]
loc: ReadOnly[tuple[int | str, ...]]
msg: ReadOnly[str]
def _is_length_error_of_rejected_items(error: ValidationErrorDetail, errors: Sequence[ValidationErrorDetail]) -> bool:
"""pydantic counts only items that validated, so a bad item also trips the parent's min_length."""
return error["type"] == "too_short" and any(

View file

@ -476,6 +476,7 @@ from litellm.proxy.common_utils.user_api_key_cache import (
project_spend_counter_key,
tag_cache_key,
)
from litellm.proxy.common_utils.validation_error_body import public_validation_errors
from litellm.proxy.config_resolvers import (
FieldSource,
SettingsStore,
@ -551,7 +552,6 @@ from litellm.proxy.hooks.proxy_track_cost_callback import _ProxyDBLogger, run_sp
from litellm.proxy.image_endpoints.endpoints import router as image_router
from litellm.proxy.list_api.common import (
ManagementProblem,
ValidationErrorDetail,
problem_response,
request_validation_problem,
)
@ -1983,16 +1983,14 @@ class _ExceptionRow(TypedDict, total=False):
@app.exception_handler(RequestValidationError)
async def otel_request_validation_exception_handler(request: Request, exc: RequestValidationError):
public_errors: Final = public_validation_errors(exc.errors())
public_exc: Final = RequestValidationError(public_errors).with_traceback(exc.__traceback__)
if request.url.path.startswith(MANAGEMENT_V1_PREFIX):
validation_errors: Final[Sequence[ValidationErrorDetail]] = exc.errors()
problem: Final = request_validation_problem(validation_errors)
_close_dangling_otel_server_span(request, problem.status, exc=exc)
problem: Final = request_validation_problem(public_errors)
_close_dangling_otel_server_span(request, problem.status, exc=public_exc)
return problem_response(problem)
_close_dangling_otel_server_span(request, 422, exc=exc)
return JSONResponse(
status_code=422,
content={"detail": jsonable_encoder(exc.errors())},
)
_close_dangling_otel_server_span(request, 422, exc=public_exc)
return JSONResponse(status_code=422, content={"detail": public_errors})
@app.exception_handler(Exception)

View file

@ -25,6 +25,7 @@ from litellm.litellm_core_utils.sensitive_data_masker import mask_sensitive_keys
from litellm.proxy._experimental.mcp_server.tool_search import MCP_TOOL_SEARCH_SETTINGS_KEY
from litellm.proxy._types import *
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_utils.validation_error_body import public_validation_errors
from litellm.proxy.config_resolvers import FieldSource, SettingsStore, source_for
from litellm.proxy.config_resolvers.settings_store import ConfigOwnedKeyError
from litellm.proxy.config_resolvers.sso import (
@ -1875,7 +1876,7 @@ async def update_ui_settings(
try:
settings: Final = effective_cls.model_validate(settings_body)
except ValidationError as e:
raise HTTPException(status_code=422, detail=e.errors())
raise HTTPException(status_code=422, detail=public_validation_errors(e.errors()))
unsupported_team_fields: Final = sorted(
frozenset(settings.team_admin_editable_team_fields) - SUPPORTED_TEAM_ADMIN_PERMISSIONS

View file

@ -202,15 +202,26 @@ def capability_classifier_system_prompt(mode: Literal["json_schema", "json_objec
)
def unwrap_classifier_json(content: str) -> str:
"""Remove the optional Markdown fence without repairing or weakening verdict JSON."""
text: Final = content.strip()
if not text.startswith("```"):
return text
unfenced: Final = text.removeprefix("```").removeprefix("json").lstrip("\n\r")
return unfenced.removesuffix("```").strip()
_JSON_DECODER: Final = json.JSONDecoder()
def _complete_json_object_at(content: str, start: int) -> str | None:
try:
_, end = _JSON_DECODER.raw_decode(content, start)
except (ValueError, RecursionError):
return None
return content[start:end]
def extract_classifier_json(content: str) -> str:
"""Return the first complete JSON object in the reply, whatever prose or fence surrounds it.
A reply with no complete object comes back stripped so the caller's validation names the defect."""
object_starts: Final = (index for index, char in enumerate(content) if char == "{")
candidates: Final = (_complete_json_object_at(content, start) for start in object_starts)
return next((candidate for candidate in candidates if candidate is not None), content.strip())
def parse_capability_classifier_verdict(content: str) -> CapabilityClassifierVerdict:
"""Parse raw JSON or the fenced JSON shape tolerated by Switchyard."""
return CapabilityClassifierVerdict.model_validate_json(unwrap_classifier_json(content))
"""Parse the verdict object out of a bare, fenced, or prose-wrapped reply."""
return CapabilityClassifierVerdict.model_validate_json(extract_classifier_json(content))

View file

@ -29,6 +29,7 @@ from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Literal, NamedTuple, cast
from pydantic import BaseModel, TypeAdapter, ValidationError, create_model
from pydantic_core import ErrorDetails
from litellm._logging import verbose_router_logger
from litellm.caching.affinity_cache import claim_affinity_pin
@ -85,8 +86,8 @@ from .capability_classifier import (
CapabilityClassifierForecast,
capability_classifier_response_format,
capability_classifier_system_prompt,
extract_classifier_json,
parse_capability_classifier_verdict,
unwrap_classifier_json,
)
from .classification_rubrics import BUSINESS_TIER_CRITERIA, calibration_examples_section
from .config import (
@ -427,6 +428,41 @@ def _effective_turn_off_message_logging(request_kwargs: Mapping[str, object] | N
)
def _classifier_reply_is_private(request_kwargs: Mapping[str, object] | None) -> bool:
from litellm.litellm_core_utils.initialize_dynamic_callback_params import (
initialize_standard_callback_dynamic_params,
)
from litellm.litellm_core_utils.redact_messages import should_redact_message_logging
kwargs: Final = dict(request_kwargs) if request_kwargs else {}
try:
return should_redact_message_logging(
{
"litellm_params": kwargs,
"standard_callback_dynamic_params": initialize_standard_callback_dynamic_params(kwargs),
}
)
except AttributeError:
return True
def _validation_problem(detail: ErrorDetails) -> str:
location: Final = ".".join(str(part) for part in detail["loc"])
return f"{location}: {detail['msg']}" if location else detail["msg"]
def _log_rejected_classifier_verdict(
error: ValidationError, content: str, request_kwargs: Mapping[str, object] | None
) -> None:
problems: Final = "; ".join(_validation_problem(detail) for detail in error.errors())
reply: Final = (
"raw reply withheld (message logging is off)"
if _classifier_reply_is_private(request_kwargs)
else f"raw reply: {content!r}"
)
verbose_router_logger.warning("ComplexityRouter: classifier verdict rejected (%s); %s", problems, reply)
_REMINDER_OPEN: Final = "<system-reminder>"
_REMINDER_CLOSE: Final = "</system-reminder>"
_DEFAULT_REMINDER_MARKERS: Final = ((_REMINDER_OPEN, _REMINDER_CLOSE),)
@ -2040,7 +2076,7 @@ class ComplexityRouter(CustomLogger):
except Exception as e: # noqa: BLE001 -- every unavailable or invalid judge verdict must fail closed
if breaker is not None and permit is not None:
breaker.record_failure(permit, is_timeout=_is_classifier_timeout(e))
return self._capability_classifier_failure_outcome(f"capability classifier failed ({e})")
return self._capability_classifier_failure_outcome(f"capability classifier failed ({type(e).__name__})")
def _capability_classifier_failure_outcome(self, reason: str, signal: str | None = None) -> ClassificationOutcome:
"""Fail closed to the configured capable tier without consulting another taxonomy."""
@ -2449,7 +2485,11 @@ class ComplexityRouter(CustomLogger):
content, classifier_cost = await self._call_classifier_model(
messages_for_call, request_kwargs, encrypted_task=encrypted_task
)
raw_tier: Final = _LabeledTierClassification.model_validate_json(content).tier
try:
raw_tier: Final = _LabeledTierClassification.model_validate_json(extract_classifier_json(content)).tier
except ValidationError as error:
_log_rejected_classifier_verdict(error, content, request_kwargs)
raise
tier: Final = self.config.resolve_classified_tier(raw_tier)
if tier is None:
raise ValueError(f"LLM classifier returned an unrecognized tier: {raw_tier!r}")
@ -2508,7 +2548,11 @@ class ComplexityRouter(CustomLogger):
max_output_tokens=capability.max_output_tokens,
encrypted_task=encrypted_task,
)
verdict: Final = parse_capability_classifier_verdict(content)
try:
verdict: Final = parse_capability_classifier_verdict(content)
except ValidationError as error:
_log_rejected_classifier_verdict(error, content, request_kwargs)
raise
threshold: Final = verdict.routing_threshold(capability.base_threshold, capability.threshold_step)
calibration: Final = capability.calibration
forecast: Final = CapabilityClassifierForecast(
@ -2563,8 +2607,9 @@ class ComplexityRouter(CustomLogger):
messages_for_call, request_kwargs, encrypted_task=encrypted, max_output_tokens=v2.max_output_tokens
)
try:
verdict: Final = LLMV2Verdict.model_validate_json(unwrap_classifier_json(content))
except ValidationError:
verdict: Final = LLMV2Verdict.model_validate_json(extract_classifier_json(content))
except ValidationError as error:
_log_rejected_classifier_verdict(error, content, request_kwargs)
return self._classifier_failure_outcome("Invalid LLM V2 forecast", prompt, system_prompt)._replace(
classifier_cost=classifier_cost
)

View file

@ -16,6 +16,7 @@ from litellm.llms.base_llm.base_utils import (
from litellm.router_strategy.complexity_router.fuse_presets import ProfileText, resolve_fuse_profile
ShortText: TypeAlias = Annotated[str, StringConstraints(strip_whitespace=True, min_length=1, max_length=512)]
VerdictText: TypeAlias = Annotated[str, StringConstraints(strip_whitespace=True, min_length=1)]
class _SolverProfile(TypedDict):
@ -90,7 +91,7 @@ class LLMV2Demands(BaseModel):
class LLMV2SolverForecast(BaseModel):
model_config = ConfigDict(extra="forbid", frozen=True)
likely_failure: ShortText
likely_failure: VerdictText
p_solve: StrictFloat = Field(ge=0.0, le=1.0)
@ -104,7 +105,7 @@ class LLMV2SolverForecasts(BaseModel):
class LLMV2Verdict(BaseModel):
model_config = ConfigDict(extra="forbid", frozen=True)
crux: ShortText
crux: VerdictText
demands: LLMV2Demands
verification: Literal["relevant", "partial", "unavailable", "unknown"]
forecasts: LLMV2SolverForecasts

View file

@ -60264,13 +60264,16 @@
},
"fireworks_ai/accounts/fireworks/routers/deepseek-v4p1-flash-us": {
"cache_read_input_token_cost": 9e-09,
"cache_read_input_token_cost_priority": 1.125e-08,
"input_cost_per_token": 4.5e-07,
"input_cost_per_token_priority": 5.625e-07,
"litellm_provider": "fireworks_ai",
"max_input_tokens": 1048576,
"max_output_tokens": 393216,
"max_tokens": 393216,
"mode": "chat",
"output_cost_per_token": 1.8e-06,
"output_cost_per_token_priority": 2.25e-06,
"source": "https://docs.fireworks.ai/serverless/pricing",
"supports_function_calling": true,
"supports_prompt_caching": true,
@ -60360,13 +60363,16 @@
},
"fireworks_ai/deepseek-v4p1-flash-us": {
"cache_read_input_token_cost": 9e-09,
"cache_read_input_token_cost_priority": 1.125e-08,
"input_cost_per_token": 4.5e-07,
"input_cost_per_token_priority": 5.625e-07,
"litellm_provider": "fireworks_ai",
"max_input_tokens": 1048576,
"max_output_tokens": 393216,
"max_tokens": 393216,
"mode": "chat",
"output_cost_per_token": 1.8e-06,
"output_cost_per_token_priority": 2.25e-06,
"source": "https://docs.fireworks.ai/serverless/pricing",
"supports_function_calling": true,
"supports_prompt_caching": true,

View file

@ -742,7 +742,7 @@ class BaseResponsesAPITest(ABC):
Passes tools=[{"type": "shell", "environment": {"type": "container_auto"}}];
validates that the request is accepted and returns a valid response.
Only runs for OpenAI; offline coverage for the Azure route lives in
tests/test_litellm/responses/test_responses_api_request_body.py.
tests/unit/responses/test_responses_api_request_body.py.
"""
base_completion_call_args = self.get_base_completion_call_args()
model = (

View file

@ -1800,7 +1800,7 @@ def test_gemini_image_size_limit_exceeded(monkeypatch):
that could cause memory issues and pod crashes.
The image fetch is mocked (mirroring the LargeImageClient pattern in
tests/test_litellm/litellm_core_utils/test_image_handling.py) so the test
tests/unit/litellm_core_utils/test_image_handling.py) so the test
deterministically exercises the size-limit rejection path without any
external network dependency.
"""

View file

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

View file

@ -22,19 +22,6 @@ import litellm
from litellm import router as litellm_router_module
from litellm import utils as litellm_utils_module
from litellm._logging import ALL_LOGGERS
from litellm.litellm_core_utils.cli_keyring import (
KeyringDiscardsWrites,
KeyringUnreachable,
KeyringUnusable,
SecretErase,
SecretErased,
SecretFound,
SecretMissing,
SecretRead,
SecretStored,
SecretStranded,
SecretWrite,
)
from litellm.litellm_core_utils.prompt_templates import (
image_handling as image_handling_module,
)
@ -42,6 +29,7 @@ from litellm.llms.custom_httpx.async_client_cleanup import (
close_litellm_async_clients,
)
from litellm.proxy.db import tool_registry_writer as tool_registry_writer_module
from tests.unit.litellm_core_utils.fake_secret_vault import FakeSecretVault
def _reset_module_level_aws_auth_caches():
@ -128,60 +116,6 @@ def isolate_host_os_keychain(monkeypatch):
monkeypatch.setenv("LITELLM_CLI_DISABLE_KEYRING", "1")
class FakeSecretVault:
"""In-memory stand-in for the OS keychain, injected wherever CLI credential storage is exercised.
`available=False` models a keychain that is locked or has no backend, `writable=False` one that
refuses to store, `erasable=False` one that will not release what it already holds, and `failure`
picks which unusable state those report. `discards=True` is keyring's null backend, which answers
reads and erases like any other yet keeps nothing it is given, so only writes report it.
"""
def __init__(
self,
blob: str | None = None,
*,
available: bool = True,
writable: bool = True,
erasable: bool = True,
discards: bool = False,
failure: KeyringUnusable = KeyringUnreachable(),
) -> None:
self.blob: str | None = blob
self.available: bool = available
self.writable: bool = writable
self.erasable: bool = erasable
self.discards: bool = discards
self.failure: KeyringUnusable = failure
self.reads: int = 0
self.writes: list[str] = []
self.erases: int = 0
def read(self) -> SecretRead:
self.reads += 1
if not self.available:
return self.failure
return SecretMissing() if self.blob is None else SecretFound(self.blob)
def write(self, blob: str) -> SecretWrite:
self.writes.append(blob)
if not (self.available and self.writable):
return self.failure
if self.discards:
return KeyringDiscardsWrites()
self.blob = blob
return SecretStored()
def erase(self) -> SecretErase:
self.erases += 1
if not self.available:
return self.failure
if not self.erasable:
return SecretStranded() if self.blob is not None else SecretErased()
self.blob = None
return SecretErased()
@pytest.fixture
def secret_vault_factory():
"""Build FakeSecretVault instances; see its docstring for the failure modes it can model."""

View file

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

File diff suppressed because it is too large Load diff

View file

@ -1,403 +1,20 @@
import copy
import os
import pickle
import subprocess
import sys
from pathlib import Path
from typing import Final, Literal
import pytest
import tiktoken
from tokenizers import Tokenizer as ReferenceTokenizer
import litellm
from litellm.caching._embedding_router import truncate_embedding_input
from litellm.litellm_core_utils.tokenizer import HuggingFaceTokenizer, OpenAIEncoding
from litellm.utils import claude_json_str
from tests.test_litellm.litellm_core_utils.test_decode_special_tokens import TOKENIZER_JSON
@pytest.mark.parametrize(
"name", ("cl100k_base", "o200k_base", "p50k_base", "p50k_edit", "r50k_base", "gpt2", "o200k_harmony")
)
@pytest.mark.parametrize(
"text", ("hello world", "café 漢字 🙂", "", "a\ud800b", "\ud83d\ude42", "🙂\ud83d\ude42\udfff", " " * 64)
from tests.unit.litellm_core_utils.test_tokenizer import (
UNICODE_TEXTS,
assert_openai_encoding_exposes_the_tiktoken_vocabulary_surface,
assert_openai_encoding_matches_python,
)
NETWORK_ENCODINGS = ("r50k_base", "gpt2")
@pytest.mark.parametrize("name", NETWORK_ENCODINGS)
@pytest.mark.parametrize("text", UNICODE_TEXTS)
def test_openai_encoding_matches_python_unicode_and_batches(name: str, text: str) -> None:
reference: Final = tiktoken.get_encoding(name)
encoding: Final = OpenAIEncoding.from_tiktoken(name)
expected: Final = reference.encode(text)
assert encoding.encode(text) == expected
assert encoding.count(text) == len(expected)
assert encoding.encode_batch([text], num_threads=2) == reference.encode_batch([text], num_threads=2)
assert encoding.encode_ordinary_batch([text]) == reference.encode_ordinary_batch([text])
assert encoding.decode_batch([expected]) == reference.decode_batch([expected])
assert encoding.decode_bytes_batch([expected]) == reference.decode_bytes_batch([expected])
assert_openai_encoding_matches_python(name, text)
@pytest.mark.parametrize("allowed", (frozenset(), frozenset({"<|endoftext|>"}), "all"))
@pytest.mark.parametrize("disallowed", (frozenset(), frozenset({"<|fim_prefix|>"}), "all"))
def test_openai_special_token_options_match_python(
allowed: frozenset[str] | Literal["all"], disallowed: frozenset[str] | Literal["all"]
) -> None:
reference: Final = tiktoken.get_encoding("cl100k_base")
encoding: Final = OpenAIEncoding.from_tiktoken(reference.name)
text: Final = "hello<|endoftext|><|fim_prefix|>world"
allowed_set: Final = reference.special_tokens_set if allowed == "all" else allowed
disallowed_set: Final = reference.special_tokens_set - allowed_set if disallowed == "all" else disallowed
if any(token in text for token in disallowed_set):
with pytest.raises(ValueError, match="disallowed special token"):
encoding.encode(text, allowed_special=allowed, disallowed_special=disallowed)
return
assert encoding.encode(text, allowed_special=allowed, disallowed_special=disallowed) == reference.encode(
text, allowed_special=allowed, disallowed_special=disallowed
)
assert encoding.special_tokens_set == reference.special_tokens_set
assert encoding.eot_token == reference.eot_token
@pytest.mark.parametrize("errors", ("replace", "ignore", "backslashreplace", "strict"))
def test_openai_partial_token_decoding_preserves_error_policy(errors: str) -> None:
reference: Final = tiktoken.get_encoding("cl100k_base")
encoding: Final = OpenAIEncoding.from_tiktoken(reference.name)
tokens: Final = reference.encode("🙂")[:1]
assert encoding.decode_bytes(tokens) == reference.decode_bytes(tokens)
if errors == "strict":
with pytest.raises(UnicodeDecodeError):
encoding.decode(tokens, errors=errors)
return
assert encoding.decode(tokens, errors=errors) == reference.decode(tokens, errors=errors)
assert encoding.decode_tokens_bytes(tokens) == reference.decode_tokens_bytes(tokens)
def test_public_encoding_and_semantic_cache_preserve_truncated_unicode() -> None:
reference: Final = tiktoken.get_encoding(litellm.encoding.name)
text: Final = "🙂"
tokens: Final = reference.encode(text)
assert litellm.encoding.encode(text, disallowed_special=()) == tokens
assert litellm.encoding.encode_batch([text]) == [tokens]
assert litellm.decode(tokens=tokens[:1]) == reference.decode(tokens[:1])
assert truncate_embedding_input(text, "", 1) == reference.decode(tokens[:1])
@pytest.mark.parametrize("add_special_tokens", (True, False))
def test_huggingface_encoding_preserves_result_fields_and_serialization(add_special_tokens: bool) -> None:
reference: Final = ReferenceTokenizer.from_str(TOKENIZER_JSON)
tokenizer: Final = HuggingFaceTokenizer.from_str(TOKENIZER_JSON)
expected: Final = reference.encode("Hello World", add_special_tokens=add_special_tokens)
actual: Final = tokenizer.encode("Hello World", add_special_tokens=add_special_tokens)
assert (actual.ids, actual.tokens, actual.type_ids, actual.offsets, actual.word_ids, actual.sequence_ids) == (
expected.ids,
expected.tokens,
expected.type_ids,
expected.offsets,
expected.word_ids,
expected.sequence_ids,
)
assert (actual.attention_mask, actual.special_tokens_mask, actual.n_sequences, len(actual)) == (
expected.attention_mask,
expected.special_tokens_mask,
expected.n_sequences,
len(expected),
)
assert copy.deepcopy(actual).ids == expected.ids
assert pickle.loads(pickle.dumps(actual)).offsets == expected.offsets
assert tokenizer.decode(actual.ids, skip_special_tokens=False) == reference.decode(
expected.ids, skip_special_tokens=False
)
def test_huggingface_character_offsets_and_pretokenized_pairs_match_python() -> None:
reference: Final = ReferenceTokenizer.from_str(claude_json_str)
tokenizer: Final = HuggingFaceTokenizer.from_str(claude_json_str)
text: Final = "café 漢字 🙂"
actual: Final = tokenizer.encode(text)
expected: Final = reference.encode(text)
assert actual.offsets == expected.offsets
assert actual.ids == expected.ids
assert (
tokenizer.encode(["hello", "world"], ["again"], is_pretokenized=True).ids
== reference.encode(["hello", "world"], ["again"], is_pretokenized=True).ids
)
def test_huggingface_batches_apply_padding_across_inputs() -> None:
reference: Final = ReferenceTokenizer.from_str(TOKENIZER_JSON)
reference.enable_padding(pad_id=0, pad_token="[UNK]")
tokenizer: Final = HuggingFaceTokenizer.from_str(reference.to_str())
inputs: Final = ["Hello", ("Hello World", "World")]
expected: Final = reference.encode_batch(inputs)
actual: Final = tokenizer.encode_batch(inputs)
fast: Final = tokenizer.encode_batch_fast(inputs)
assert [(item.ids, item.attention_mask, item.offsets) for item in actual] == [
(item.ids, item.attention_mask, item.offsets) for item in expected
]
assert [item.ids for item in fast] == [item.ids for item in expected]
assert tokenizer.decode_batch([item.ids for item in actual]) == reference.decode_batch(
[item.ids for item in expected]
)
def test_caller_supplied_huggingface_tokenizer_preserves_public_encode_and_count() -> None:
tokenizer: Final = ReferenceTokenizer.from_str(TOKENIZER_JSON)
custom: Final = {"type": "huggingface_tokenizer", "tokenizer": tokenizer}
expected: Final = tokenizer.encode("Hello World").ids
assert litellm.encode(text="Hello World", custom_tokenizer=custom) == expected
assert litellm.token_counter(text="Hello World", custom_tokenizer=custom) == len(expected)
assert litellm.decode(tokens=expected, custom_tokenizer=custom) == "Hello World"
def test_caller_supplied_tiktoken_treats_special_spellings_as_text() -> None:
tokenizer: Final = tiktoken.get_encoding("cl100k_base")
custom: Final = {"type": "openai_tokenizer", "tokenizer": tokenizer}
text: Final = "<|endoftext|>"
assert litellm.encode(text=text, custom_tokenizer=custom) == tokenizer.encode(text, disallowed_special=())
def test_public_tokenizer_objects_survive_pickle_and_deepcopy(tmp_path: Path) -> None:
custom: Final = litellm.create_tokenizer(TOKENIZER_JSON)
tokenizer: Final = custom["tokenizer"]
path: Final = tmp_path / "tokenizer.json"
tokenizer.save(str(path))
assert copy.deepcopy(custom)["tokenizer"].encode("Hello World").ids == tokenizer.encode("Hello World").ids
assert (
pickle.loads(pickle.dumps(custom))["tokenizer"].encode("Hello World").ids == tokenizer.encode("Hello World").ids
)
assert HuggingFaceTokenizer.from_file(str(path)).encode("Hello World").ids == tokenizer.encode("Hello World").ids
assert copy.deepcopy(litellm.encoding).encode("hello") == litellm.encoding.encode("hello")
assert pickle.loads(pickle.dumps(litellm.encoding)).encode("hello") == litellm.encoding.encode("hello")
@pytest.mark.parametrize("offline", ("0", "1"))
def test_hub_loader_preserves_environment_auth_cache_and_offline(tmp_path: Path, offline: str) -> None:
script: Final = """
import json
import sys
from pathlib import Path
sys.path.insert(0, sys.argv[1])
import httpx
import huggingface_hub
from huggingface_hub.errors import LocalEntryNotFoundError
import litellm
payload = sys.argv[2].encode()
offline = sys.argv[3] == "1"
observed = []
def handle(request):
assert not offline, "offline loading issued a request"
if request.url.path.endswith("/tokenizer.json"):
observed.append(request.headers.get("authorization"))
if request.headers.get("authorization") != "Bearer audit-fixture-token":
return httpx.Response(401)
return httpx.Response(200, headers={"content-length": str(len(payload)), "etag": '"fixture"', "x-repo-commit": "a" * 40}, content=payload if request.method == "GET" else b"")
if not offline:
huggingface_hub.set_client_factory(lambda: httpx.Client(transport=httpx.MockTransport(handle)))
try:
tokenizer = litellm.create_pretrained_tokenizer("test-fixture/tokenizer")["tokenizer"]
except LocalEntryNotFoundError:
assert offline
assert observed == []
else:
assert not offline
assert "Bearer audit-fixture-token" in observed
assert tokenizer.decode(tokenizer.encode("Hello World").ids) == "Hello World"
assert tuple(Path(sys.argv[4]).rglob("tokenizer.json"))
print("compatible")
"""
result: Final = subprocess.run(
[
sys.executable,
"-I",
"-c",
script,
str(Path(litellm.__file__).parent.parent),
TOKENIZER_JSON,
offline,
str(tmp_path / "cache"),
],
capture_output=True,
text=True,
timeout=30,
env={
**os.environ,
"HF_HOME": str(tmp_path / "home"),
"HF_HUB_CACHE": str(tmp_path / "cache"),
"HF_ENDPOINT": "http://127.0.0.1:9",
"HF_TOKEN": "audit-fixture-token",
"HF_HUB_OFFLINE": offline,
"HF_HUB_DISABLE_IMPLICIT_TOKEN": "0",
"LITELLM_LOCAL_MODEL_COST_MAP": "True",
},
)
assert result.returncode == 0, result.stdout + result.stderr
assert result.stdout.strip() == "compatible"
@pytest.mark.parametrize("rust", (None, "0", "1"))
def test_tokenization_without_native_extension_stays_offline(tmp_path: Path, rust: str | None) -> None:
script: Final = """
import importlib.abc
import sys
sys.path.insert(0, sys.argv[1])
def reject_network(event, args):
if event == "socket.connect":
raise AssertionError("tokenizer attempted a network connection")
sys.addaudithook(reject_network)
class Block(importlib.abc.MetaPathFinder):
def find_spec(self, fullname, path=None, target=None):
if fullname == "litellm.rust_bridge._native":
raise ImportError("native extension is unavailable")
sys.meta_path.insert(0, Block())
import litellm
from litellm.rust_bridge.tokenizer import get_encoding
import tiktoken
from tokenizers import Tokenizer
assert isinstance(litellm.encoding, tiktoken.Encoding)
for name in ("cl100k_base", "o200k_base", "o200k_harmony", "p50k_base", "p50k_edit"):
encoding = get_encoding(name)
text = "offline café 漢字 🙂" + " " * 64
assert encoding.decode(encoding.encode(text)) == text
ids = litellm.encode(text="hello world")
assert litellm.decode(tokens=ids) == "hello world"
assert litellm.token_counter(model=None, text="hello world") == len(ids)
custom = litellm.create_tokenizer(sys.argv[2])
assert isinstance(custom["tokenizer"], Tokenizer)
custom["tokenizer"].enable_padding(pad_id=0, pad_token="[UNK]")
assert litellm.decode(tokens=litellm.encode(text="Hello World", custom_tokenizer=custom), custom_tokenizer=custom) == "Hello World"
print("compatible")
"""
result: Final = subprocess.run(
[sys.executable, "-I", "-c", script, str(Path(litellm.__file__).parent.parent), TOKENIZER_JSON],
capture_output=True,
text=True,
timeout=30,
cwd=tmp_path,
env={
**{key: value for key, value in os.environ.items() if key != "LITELLM_RUST"},
**({"LITELLM_RUST": rust} if rust is not None else {}),
"LITELLM_LOCAL_MODEL_COST_MAP": "True",
"TIKTOKEN_CACHE_DIR": str(tmp_path / "unused-tokenizer-cache"),
},
)
assert result.returncode == 0, result.stdout + result.stderr
assert result.stdout.strip() == "compatible"
assert not (tmp_path / "unused-tokenizer-cache").exists()
@pytest.mark.parametrize("is_pretokenized", (False, True))
def test_huggingface_batch_sequence_containers_match_python(is_pretokenized: bool) -> None:
reference: Final = ReferenceTokenizer.from_str(TOKENIZER_JSON)
tokenizer: Final = HuggingFaceTokenizer.from_str(TOKENIZER_JSON)
inputs: Final = [["Hello", "World"], ("Hello", "World")]
actual: Final = tokenizer.encode_batch(inputs, is_pretokenized=is_pretokenized)
expected: Final = reference.encode_batch(inputs, is_pretokenized=is_pretokenized)
assert [(item.ids, item.type_ids, item.sequence_ids) for item in actual] == [
(item.ids, item.type_ids, item.sequence_ids) for item in expected
]
@pytest.mark.parametrize("name", ("cl100k_base", "o200k_base", "p50k_edit", "gpt2"))
@pytest.mark.parametrize("name", ("gpt2",))
def test_openai_encoding_exposes_the_tiktoken_vocabulary_surface(name: str) -> None:
reference: Final = tiktoken.get_encoding(name)
encoding: Final = OpenAIEncoding.from_tiktoken(name)
text: Final = "hello fanta"
assert repr(encoding) == repr(reference) == f"<Encoding {name!r}>"
assert (encoding.name, encoding.n_vocab, encoding.max_token_value) == (
reference.name,
reference.n_vocab,
reference.max_token_value,
)
assert encoding.token_byte_values() == reference.token_byte_values()
assert encoding.encode_single_token("hello") == reference.encode_single_token("hello")
assert encoding.encode_single_token(b"<|endoftext|>") == reference.eot_token
assert [encoding.is_special_token(token) for token in (0, reference.eot_token)] == [False, True]
assert encoding.decode_with_offsets(reference.encode(text)) == reference.decode_with_offsets(reference.encode(text))
assert encoding.encode_to_numpy(text).tolist() == reference.encode_to_numpy(text).tolist()
stable, completions = encoding.encode_with_unstable(text)
expected_stable, expected_completions = reference.encode_with_unstable(text)
assert (stable, sorted(completions)) == (expected_stable, sorted(expected_completions))
with pytest.raises(KeyError):
encoding.encode_single_token("<|not-a-token|>")
def test_huggingface_tokenizer_exposes_the_tokenizers_vocabulary_surface() -> None:
reference: Final = ReferenceTokenizer.from_str(TOKENIZER_JSON)
reference.enable_padding(pad_id=0, pad_token="[UNK]", length=4)
reference.enable_truncation(max_length=3, stride=1, strategy="only_first", direction="left")
tokenizer: Final = HuggingFaceTokenizer.from_str(reference.to_str())
assert tokenizer.token_to_id("Hello") == reference.token_to_id("Hello") == 1
assert tokenizer.id_to_token(3) == reference.id_to_token(3) == "[BOS]"
assert tokenizer.id_to_token(99) is None
assert tokenizer.get_vocab() == reference.get_vocab()
assert tokenizer.get_vocab(with_added_tokens=False) == reference.get_vocab(with_added_tokens=False)
assert tokenizer.get_vocab_size() == reference.get_vocab_size() == 4
assert tokenizer.get_vocab_size(with_added_tokens=False) == reference.get_vocab_size(with_added_tokens=False)
added: Final = tokenizer.get_added_tokens_decoder()
expected_added: Final = reference.get_added_tokens_decoder()
assert {token_id: str(token) for token_id, token in added.items()} == {
token_id: str(token) for token_id, token in expected_added.items()
}
assert added[3].special == expected_added[3].special
assert tokenizer.num_special_tokens_to_add(False) == reference.num_special_tokens_to_add(False) == 1
assert tokenizer.num_special_tokens_to_add(True) == reference.num_special_tokens_to_add(True) == 0
assert tokenizer.padding == reference.padding
assert tokenizer.truncation == reference.truncation
assert tokenizer.encode_special_tokens == reference.encode_special_tokens is False
assert HuggingFaceTokenizer.from_buffer(TOKENIZER_JSON.encode()).encode("Hello").ids == [3, 1]
assert HuggingFaceTokenizer.from_str(TOKENIZER_JSON).padding is None
assert HuggingFaceTokenizer.from_str(TOKENIZER_JSON).truncation is None
def test_huggingface_encoding_exposes_the_tokenizers_lookup_and_mutation_surface() -> None:
reference: Final = ReferenceTokenizer.from_str(claude_json_str)
tokenizer: Final = HuggingFaceTokenizer.from_str(claude_json_str)
text: Final = "hello wide world"
actual: Final = tokenizer.encode(text, "again")
expected: Final = reference.encode(text, "again")
lookups: Final = (
lambda encoding: [encoding.token_to_chars(index) for index in range(len(encoding))],
lambda encoding: [encoding.token_to_word(index) for index in range(len(encoding))],
lambda encoding: [encoding.token_to_sequence(index) for index in range(len(encoding))],
lambda encoding: [encoding.char_to_token(position) for position in range(len(text))],
lambda encoding: [encoding.char_to_word(position) for position in range(len(text))],
lambda encoding: [encoding.char_to_token(position, 1) for position in range(5)],
lambda encoding: [encoding.word_to_tokens(word) for word in range(3)],
lambda encoding: [encoding.word_to_chars(word) for word in range(3)],
lambda encoding: [encoding.word_to_tokens(0, 1), encoding.word_to_chars(0, 1)],
)
for lookup in lookups:
assert lookup(actual) == lookup(expected)
assert repr(actual) == repr(expected)
actual.truncate(4, stride=1, direction="left")
expected.truncate(4, stride=1, direction="left")
assert (actual.ids, [item.ids for item in actual.overflowing]) == (
expected.ids,
[item.ids for item in expected.overflowing],
)
actual.pad(6, direction="left", pad_id=7, pad_type_id=1, pad_token="<pad>")
expected.pad(6, direction="left", pad_id=7, pad_type_id=1, pad_token="<pad>")
assert (actual.ids, actual.attention_mask, actual.type_ids, actual.tokens) == (
expected.ids,
expected.attention_mask,
expected.type_ids,
expected.tokens,
)
actual.set_sequence_id(3)
expected.set_sequence_id(3)
assert actual.sequence_ids == expected.sequence_ids
merged: Final = type(actual).merge([actual, tokenizer.encode("more")])
assert merged.ids == type(expected).merge([expected, reference.encode("more")]).ids
assert merged.offsets == type(expected).merge([expected, reference.encode("more")]).offsets
with pytest.raises(ValueError, match="direction"):
actual.pad(8, direction="sideways")
assert_openai_encoding_exposes_the_tiktoken_vocabulary_surface(name)

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -1033,7 +1033,7 @@ def test_qdrant_semantic_cache_defaults_embedding_timeout():
@pytest.mark.asyncio
async def test_qdrant_async_embedding_truncates_off_the_event_loop(monkeypatch):
from tests.large_text import text
from tests.test_litellm.litellm_core_utils.event_loop_lag import (
from tests.unit.litellm_core_utils.event_loop_lag import (
assert_loop_stayed_free,
timed_with_loop_lags,
warm_tokenizer,

View file

@ -1392,7 +1392,7 @@ def test_redis_semantic_cache_defaults_embedding_timeout():
@pytest.mark.asyncio
async def test_redis_async_embedding_truncates_off_the_event_loop(monkeypatch):
from tests.large_text import text
from tests.test_litellm.litellm_core_utils.event_loop_lag import (
from tests.unit.litellm_core_utils.event_loop_lag import (
assert_loop_stayed_free,
timed_with_loop_lags,
warm_tokenizer,

View file

@ -528,7 +528,7 @@ async def test_pre_call_hook_no_compression_records_no_savings(monkeypatch):
@pytest.mark.asyncio
async def test_pre_call_hook_counts_tokens_off_the_event_loop():
from tests.large_text import text
from tests.test_litellm.litellm_core_utils.event_loop_lag import (
from tests.unit.litellm_core_utils.event_loop_lag import (
assert_loop_stayed_free,
timed_with_loop_lags,
warm_tokenizer,

View file

@ -0,0 +1,15 @@
import importlib
import pytest
from tests.unit.litellm_core_utils.fake_secret_vault import FakeSecretVault
@pytest.fixture(autouse=True, scope="session")
def bundled_tiktoken_cache() -> None:
importlib.import_module("litellm.litellm_core_utils.default_encoding")
@pytest.fixture
def secret_vault_factory() -> type[FakeSecretVault]:
return FakeSecretVault

View file

@ -0,0 +1,67 @@
from litellm.litellm_core_utils.cli_keyring import (
KeyringDiscardsWrites,
KeyringUnreachable,
KeyringUnusable,
SecretErase,
SecretErased,
SecretFound,
SecretMissing,
SecretRead,
SecretStored,
SecretStranded,
SecretWrite,
)
class FakeSecretVault:
"""In-memory stand-in for the OS keychain, injected wherever CLI credential storage is exercised.
`available=False` models a keychain that is locked or has no backend, `writable=False` one that
refuses to store, `erasable=False` one that will not release what it already holds, and `failure`
picks which unusable state those report. `discards=True` is keyring's null backend, which answers
reads and erases like any other yet keeps nothing it is given, so only writes report it.
"""
def __init__(
self,
blob: str | None = None,
*,
available: bool = True,
writable: bool = True,
erasable: bool = True,
discards: bool = False,
failure: KeyringUnusable = KeyringUnreachable(),
) -> None:
self.blob: str | None = blob
self.available: bool = available
self.writable: bool = writable
self.erasable: bool = erasable
self.discards: bool = discards
self.failure: KeyringUnusable = failure
self.reads: int = 0
self.writes: list[str] = []
self.erases: int = 0
def read(self) -> SecretRead:
self.reads += 1
if not self.available:
return self.failure
return SecretMissing() if self.blob is None else SecretFound(self.blob)
def write(self, blob: str) -> SecretWrite:
self.writes.append(blob)
if not (self.available and self.writable):
return self.failure
if self.discards:
return KeyringDiscardsWrites()
self.blob = blob
return SecretStored()
def erase(self) -> SecretErase:
self.erases += 1
if not self.available:
return self.failure
if not self.erasable:
return SecretStranded() if self.blob is not None else SecretErased()
self.blob = None
return SecretErased()

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