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

This commit is contained in:
yucheng 2026-09-25 19:50:30 +00:00
commit 47e6df7b76
732 changed files with 11540 additions and 11194 deletions

View file

@ -14,12 +14,12 @@ while IFS= read -r file || [ -n "$file" ]; do
[ -n "$file" ] || continue
case "$file" in
*.md | *.mdx) : ;;
pyproject.toml | */pyproject.toml | uv.lock | uv.toml | .python-version | rust-toolchain.toml | litellm-rust/* | litellm/__init__.py | litellm/proxy/proxy_server.py | litellm/*mcp* | tests/*mcp* | litellm/integrations/arize/* | tests/base_sdk_tests/* | scripts/check_mcp_sdk_install.py | .github/workflows/test-mcp-dependency-resolution.yml | .github/actions/detect-changes/* | .github/actions/setup-uv-with-retries/* | .github/actions/cache-cargo-build/* | .github/scripts/detect_changes.sh | .github/scripts/uv_sync_with_retries.sh | .circleci/scripts/classify_changes.sh | tests/test_litellm/test_circleci_path_filter.py | tests/test_litellm/test_detect_changes.py)
pyproject.toml | */pyproject.toml | uv.lock | uv.toml | .python-version | rust-toolchain.toml | litellm-rust/* | litellm/__init__.py | litellm/proxy/proxy_server.py | litellm/*mcp* | tests/*mcp* | litellm/integrations/arize/* | tests/base_sdk_tests/* | scripts/check_mcp_sdk_install.py | .github/workflows/test-mcp-dependency-resolution.yml | .github/actions/detect-changes/* | .github/actions/setup-uv-with-retries/* | .github/actions/cache-cargo-build/* | .github/scripts/detect_changes.sh | .github/scripts/uv_sync_with_retries.sh | .circleci/scripts/classify_changes.sh | tests/unit/test_circleci_path_filter.py | tests/unit/test_detect_changes.py)
has_mcp_dependencies=true ;;
esac
case "$file" in
tests/e2e/*/*.py) : ;;
tests/e2e/*.py | tests/code_coverage_tests/test_provider_cache.py | tests/code_coverage_tests/test_provider_replay_harness.py | tests/test_litellm/test_circleci_path_filter.py | .circleci/* | pyproject.toml | uv.lock)
tests/e2e/*.py | tests/code_coverage_tests/test_provider_cache.py | tests/code_coverage_tests/test_provider_replay_harness.py | tests/unit/test_circleci_path_filter.py | .circleci/* | pyproject.toml | uv.lock)
has_provider_harness=true ;;
esac
case "$file" in

View file

@ -7,7 +7,10 @@ legacy_flags=(
caching-local
enterprise-package
enterprise-routing
llm-other-providers
llm-vertex-ai
mcp-integration
misc
proxy-db-auth-checks
proxy-db-budgets
proxy-db-custom-logging
@ -22,6 +25,7 @@ legacy_flags=(
proxy-db-proxy-utils
proxy-extras
proxy-infra
responses-caching-types
)
legacy_paths() {
@ -36,6 +40,7 @@ legacy_paths() {
echo tests/unit/enterprise/proxy/test_audit_logging_endpoints.py
echo tests/unit/enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py ;;
enterprise-routing)
echo tests/unit/google_genai
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
@ -47,10 +52,31 @@ legacy_paths() {
echo tests/unit/enterprise/proxy/test_file_deletion_blocking.py
echo tests/unit/enterprise/proxy/test_managed_files_access_check.py
echo tests/unit/enterprise/proxy/test_managed_files_hook.py ;;
llm-other-providers) find tests/unit/llms -name 'test_*.py' -not -path 'tests/unit/llms/vertex_ai/*' ;;
llm-vertex-ai) echo tests/unit/llms/vertex_ai ;;
mcp-integration)
echo tests/unit/experimental_mcp_client
echo tests/unit/proxy/_experimental/mcp_server
echo tests/unit/responses/mcp
echo tests/mcp_tests/test_proxy_mcp_e2e.py ;;
misc)
find tests/unit -maxdepth 1 -name 'test_*.py'
echo tests/unit/test_router
echo tests/unit/a2a_protocol
echo tests/unit/batches
echo tests/unit/chat_completions
echo tests/unit/completion_extras
echo tests/unit/containers
echo tests/unit/embeddings
echo tests/unit/endpoints
echo tests/unit/files
echo tests/unit/images
echo tests/unit/interactions
echo tests/unit/messages
echo tests/unit/rag
echo tests/unit/rerank_api
echo tests/unit/vector_stores
echo tests/unit/videos ;;
proxy-db-auth-checks)
echo tests/unit/proxy/auth/test_auth_checks.py
echo tests/unit/proxy/auth/test_user_api_key_auth.py
@ -113,6 +139,7 @@ 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 ;;
*) echo "unit_selection.sh: unknown flag $1" >&2; exit 1 ;;
esac
}

View file

@ -341,6 +341,7 @@ workflows:
flag:
- enterprise-package
- proxy-infra
- responses-caching-types
- proxy-db-auth-checks
- proxy-db-jwt-and-keys
- proxy-db-proxy-server-core
@ -353,6 +354,28 @@ workflows:
- proxy-db-endpoints-and-responses
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-llm-vertex-ai
flag: llm-vertex-ai
shards: 2
workers: 1
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-llm-other-providers
flag: llm-other-providers
shards: 3
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-misc
flag: misc
shards: 2
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-proxy-db-proxy-utils
flag: proxy-db-proxy-utils

View file

@ -1,12 +1,12 @@
{
"cases": {
"CHAT-JSON": "tests/test_litellm/llms/openai/test_openai.py::test_acompletion_returns_json_reply_over_injected_transport",
"CHAT-TEXT-STREAM": "tests/test_litellm/llms/openai/test_openai.py::test_acompletion_streams_text_deltas_over_injected_transport",
"CHAT-TOOL-STREAM": "tests/test_litellm/llms/openai/test_openai.py::test_acompletion_streams_tool_call_arguments_over_injected_transport",
"CHAT-JSON": "tests/unit/llms/openai/test_openai.py::test_acompletion_returns_json_reply_over_injected_transport",
"CHAT-TEXT-STREAM": "tests/unit/llms/openai/test_openai.py::test_acompletion_streams_text_deltas_over_injected_transport",
"CHAT-TOOL-STREAM": "tests/unit/llms/openai/test_openai.py::test_acompletion_streams_tool_call_arguments_over_injected_transport",
"MODEL-ALLOW": "tests/test_litellm/proxy/auth/test_auth_checks.py::test_can_object_call_model_allows_listed_model_for_key",
"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/test_litellm/test_cost_calculator.py::test_completion_cost_charges_explicit_per_token_rates_over_registered_ones",
"COST-ZERO": "tests/test_litellm/test_cost_calculator.py::test_completion_cost_is_zero_when_explicit_rates_are_zero",
"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",

View file

@ -13,12 +13,13 @@ on:
have its path existence-checked like any other token.
required: true
type: string
fork-flag:
unit-flag:
description: >-
Codecov flag of the `.circleci/tests.yml` job that now owns part of
this shard. CircleCI does not run on pull requests from forks, so on
those events this shard also runs the files
`.circleci/scripts/unit_selection.sh` lists for the flag.
this shard. The shard also runs the files
`.circleci/scripts/unit_selection.sh` lists for the flag, on every
event, because the CircleCI pipeline is manual-only while the tests
migrate.
required: false
type: string
default: ""
@ -175,8 +176,7 @@ jobs:
timeout-minutes: ${{ inputs.timeout-minutes }}
env:
TEST_PATH: ${{ inputs.test-path }}
FORK_FLAG: ${{ inputs.fork-flag }}
IS_FORK: ${{ github.event_name == 'pull_request' && github.event.pull_request.head.repo.full_name != github.repository }}
UNIT_FLAG: ${{ inputs.unit-flag }}
MAX_FAILURES: ${{ inputs.max-failures }}
WORKERS: ${{ inputs.workers }}
RERUNS: ${{ inputs.reruns }}
@ -186,11 +186,11 @@ jobs:
run: |
echo "has-coverage=false" >> "$GITHUB_OUTPUT"
selection="${TEST_PATH}"
if [ "${IS_FORK}" = "true" ] && [ -n "${FORK_FLAG}" ]; then
selection="${TEST_PATH} $(bash .circleci/scripts/unit_selection.sh "${FORK_FLAG}" | tr '\n' ' ')"
if [ -n "${UNIT_FLAG}" ]; then
selection="${TEST_PATH} $(bash .circleci/scripts/unit_selection.sh "${UNIT_FLAG}" | tr '\n' ' ')"
fi
if [ -z "${selection// /}" ]; then
echo "shard selection is empty on this event (CircleCI flag ${FORK_FLAG:-none} owns it); nothing to run"
echo "shard selection is empty; nothing to run"
exit 0
fi
pytest_args=()

View file

@ -10,7 +10,7 @@ on:
- "litellm/_redis_credential_provider.py"
- "litellm/caching/redis_cache.py"
- "litellm/caching/evicted_client_closer.py"
- "tests/test_litellm/test_redis.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"
@ -84,7 +84,7 @@ jobs:
run: |
redis-server --version
uv run --no-sync pytest \
tests/test_litellm/test_redis.py \
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 \

View file

@ -22,9 +22,10 @@ concurrency:
#
# `.circleci/tests.yml` runs each group's files on same-repo events under the
# `proxy-db-<group>` Codecov flag; `.circleci/scripts/unit_selection.sh` holds
# the file lists. CircleCI does not build pull requests from forks, so `fork-flag`
# makes the shard run that list there. `test-path` keeps the files that still
# reach real providers and never left tests/proxy_unit_tests.
# the file lists. That pipeline is manual-only while the tests migrate, so
# `unit-flag` makes the shard run that list on every event. `test-path` keeps
# the files that still reach real providers and never left
# tests/proxy_unit_tests.
#
# Design targets:
# * Every shard runs in <= 7 minutes of wall-clock on the default runner.
@ -78,7 +79,7 @@ jobs:
# Must run serially — event-loop conflict with the logging worker.
- test-group: key-generation
test-path: ""
fork-flag: proxy-db-key-generation
unit-flag: proxy-db-key-generation
workers: 0
dist: loadscope
timeout: 20
@ -86,13 +87,13 @@ jobs:
# ---- auth: split into 2 shards ----
- test-group: auth-checks
test-path: ""
fork-flag: proxy-db-auth-checks
unit-flag: proxy-db-auth-checks
workers: 4
dist: loadscope
timeout: 15
- test-group: jwt-and-keys
test-path: ""
fork-flag: proxy-db-jwt-and-keys
unit-flag: proxy-db-jwt-and-keys
workers: 4
dist: loadscope
timeout: 15
@ -100,7 +101,7 @@ jobs:
# ---- test_proxy_utils.py, single shard, worksteal distribution ----
- test-group: proxy-utils
test-path: ""
fork-flag: proxy-db-proxy-utils
unit-flag: proxy-db-proxy-utils
workers: 4
dist: worksteal
timeout: 15
@ -108,13 +109,13 @@ jobs:
# ---- proxy server: split into 2 shards ----
- test-group: proxy-server-core
test-path: "tests/proxy_unit_tests/test_proxy_server_gemini_pass_through.py"
fork-flag: proxy-db-proxy-server-core
unit-flag: proxy-db-proxy-server-core
workers: 4
dist: loadscope
timeout: 15
- test-group: proxy-runtime
test-path: ""
fork-flag: proxy-db-proxy-runtime
unit-flag: proxy-db-proxy-runtime
workers: 4
dist: loadscope
timeout: 15
@ -122,20 +123,20 @@ jobs:
# ---- logging: split into 2 shards ----
- test-group: custom-logging
test-path: "tests/proxy_unit_tests/test_proxy_custom_logger.py"
fork-flag: proxy-db-custom-logging
unit-flag: proxy-db-custom-logging
workers: 4
dist: loadscope
timeout: 15
- test-group: logging-misc
test-path: ""
fork-flag: proxy-db-logging-misc
unit-flag: proxy-db-logging-misc
workers: 4
dist: loadscope
timeout: 15
- test-group: db-and-spend
test-path: ""
fork-flag: proxy-db-db-and-spend
unit-flag: proxy-db-db-and-spend
workers: 4
dist: loadscope
timeout: 15
@ -143,27 +144,27 @@ jobs:
# ---- guardrails + budget + hooks: split into 2 ----
- test-group: guardrails-hooks
test-path: ""
fork-flag: proxy-db-guardrails-hooks
unit-flag: proxy-db-guardrails-hooks
workers: 4
dist: loadscope
timeout: 15
- test-group: budgets
test-path: ""
fork-flag: proxy-db-budgets
unit-flag: proxy-db-budgets
workers: 4
dist: loadscope
timeout: 15
- test-group: endpoints-and-responses
test-path: "tests/proxy_unit_tests/test_proxy_exception_mapping.py"
fork-flag: proxy-db-endpoints-and-responses
unit-flag: proxy-db-endpoints-and-responses
workers: 4
dist: loadscope
timeout: 15
uses: ./.github/workflows/_test-unit-base.yml
with:
test-path: ${{ matrix.test-path }}
fork-flag: ${{ matrix.fork-flag }}
unit-flag: ${{ matrix.unit-flag }}
workers: ${{ matrix.workers }}
reruns: 2
timeout-minutes: ${{ matrix.timeout }}

View file

@ -36,9 +36,9 @@ concurrency:
# Folding it in here is a follow-up, together with generalising that guard into
# assert_ci_coverage.py.
#
# `fork-flag` names the `.circleci/tests.yml` job that now runs part of the
# shard under the same Codecov flag. CircleCI does not build pull requests from
# forks, so the shard still runs those files there and skips them elsewhere.
# `unit-flag` names the `.circleci/tests.yml` job that now runs part of the
# shard under the same Codecov flag. That pipeline is manual-only while the
# tests migrate, so the shard also runs those files on every event.
jobs:
unit:
name: ${{ matrix.shard }}
@ -52,8 +52,8 @@ jobs:
include:
- shard: mcp-integration
artifact-name: mcp-integration
test-path: "tests/mcp_tests tests/test_litellm/experimental_mcp_client"
fork-flag: mcp-integration
test-path: "tests/mcp_tests"
unit-flag: mcp-integration
workers: 2
reruns: 0
timeout-minutes: 20
@ -70,10 +70,9 @@ jobs:
- shard: enterprise-routing
artifact-name: enterprise-routing
test-path: >-
tests/test_litellm/google_genai
tests/test_litellm/router_utils
tests/test_litellm/router_strategy
fork-flag: enterprise-routing
unit-flag: enterprise-routing
workers: 2
reruns: 2
timeout-minutes: 20
@ -90,6 +89,7 @@ jobs:
- shard: Vertex AI
artifact-name: llm-vertex-ai
test-path: "tests/test_litellm/llms/vertex_ai"
unit-flag: llm-vertex-ai
workers: 1
reruns: 2
timeout-minutes: 20
@ -98,6 +98,7 @@ jobs:
- shard: All Other Providers
artifact-name: llm-other-providers
test-path: "tests/test_litellm/llms --ignore=tests/test_litellm/llms/vertex_ai"
unit-flag: llm-other-providers
workers: 2
reruns: 2
timeout-minutes: 20
@ -106,26 +107,13 @@ jobs:
- shard: misc
artifact-name: misc
test-path: >-
tests/test_litellm/batches
tests/test_litellm/secret_managers
tests/test_litellm/a2a_protocol
tests/test_litellm/chat_completions
tests/test_litellm/completion_extras
tests/test_litellm/containers
tests/test_litellm/endpoints
tests/test_litellm/files
tests/test_litellm/images
tests/test_litellm/interactions
tests/test_litellm/messages
tests/test_litellm/embeddings
tests/test_litellm/ocr
tests/test_litellm/passthrough
tests/test_litellm/rag
tests/test_litellm/rerank_api
tests/test_litellm/rust_bridge
tests/test_litellm/vector_stores
tests/test_litellm/videos
tests/test_litellm/test_*.py
unit-flag: misc
workers: 2
reruns: 2
timeout-minutes: 20
@ -205,7 +193,7 @@ jobs:
tests/test_litellm/proxy/types_utils
tests/test_litellm/proxy/logging_endpoints
tests/test_litellm/proxy/test_*.py
fork-flag: proxy-infra
unit-flag: proxy-infra
workers: 4
reruns: 2
timeout-minutes: 20
@ -214,7 +202,7 @@ jobs:
- shard: caching-local
artifact-name: caching-local
test-path: ""
fork-flag: caching-local
unit-flag: caching-local
workers: 2
reruns: 2
timeout-minutes: 20
@ -223,7 +211,7 @@ jobs:
- shard: proxy-extras
artifact-name: proxy-extras
test-path: ""
fork-flag: proxy-extras
unit-flag: proxy-extras
workers: 2
reruns: 2
timeout-minutes: 20
@ -232,7 +220,7 @@ jobs:
- shard: enterprise-package
artifact-name: enterprise-package
test-path: ""
fork-flag: enterprise-package
unit-flag: enterprise-package
workers: 4
reruns: 2
timeout-minutes: 20
@ -243,7 +231,7 @@ jobs:
test-path: >-
tests/test_litellm/responses
tests/test_litellm/caching
tests/test_litellm/types
unit-flag: responses-caching-types
workers: 2
reruns: 2
timeout-minutes: 20
@ -251,7 +239,7 @@ jobs:
uses: ./.github/workflows/_test-unit-base.yml
with:
test-path: ${{ matrix.test-path }}
fork-flag: ${{ matrix.fork-flag || '' }}
unit-flag: ${{ matrix.unit-flag || '' }}
workers: ${{ matrix.workers }}
reruns: ${{ matrix.reruns }}
timeout-minutes: ${{ matrix.timeout-minutes }}

View file

@ -314,7 +314,7 @@ test-unit: install-test-deps
# Matrix test targets (matching CI workflow groups)
test-unit-llms: install-test-deps
$(UV_RUN) pytest tests/test_litellm/llms --tb=short -vv -n 4 --durations=20
$(UV_RUN) pytest tests/unit/llms --tb=short -vv -n 4 --durations=20
test-unit-proxy-guardrails: install-test-deps
$(UV_RUN) pytest tests/test_litellm/proxy/guardrails tests/test_litellm/proxy/management_endpoints tests/test_litellm/proxy/management_helpers --tb=short -vv -n 4 --durations=20
@ -332,10 +332,10 @@ test-unit-core-utils: install-test-deps
$(UV_RUN) pytest tests/test_litellm/litellm_core_utils --tb=short -vv -n 2 --durations=20
test-unit-other: install-test-deps
$(UV_RUN) pytest tests/test_litellm/caching tests/test_litellm/responses tests/test_litellm/secret_managers tests/test_litellm/vector_stores tests/test_litellm/a2a_protocol tests/test_litellm/anthropic_interface tests/test_litellm/completion_extras tests/test_litellm/containers tests/unit/enterprise tests/test_litellm/experimental_mcp_client tests/test_litellm/google_genai tests/test_litellm/images tests/test_litellm/interactions tests/test_litellm/passthrough tests/test_litellm/router_strategy tests/test_litellm/router_utils tests/test_litellm/types --tb=short -vv -n 4 --durations=20
$(UV_RUN) pytest tests/test_litellm/caching tests/test_litellm/responses tests/test_litellm/secret_managers tests/unit/vector_stores tests/unit/a2a_protocol tests/test_litellm/anthropic_interface tests/unit/completion_extras tests/unit/containers tests/unit/enterprise tests/unit/experimental_mcp_client tests/unit/google_genai tests/unit/images tests/unit/interactions tests/test_litellm/interactions tests/test_litellm/passthrough tests/test_litellm/router_strategy tests/test_litellm/router_utils tests/unit/types --tb=short -vv -n 4 --durations=20
test-unit-root: install-test-deps
$(UV_RUN) pytest tests/test_litellm/test_*.py --tb=short -vv -n 4 --durations=20
$(UV_RUN) pytest tests/unit/test_*.py tests/test_litellm/test_*.py --tb=short -vv -n 4 --durations=20
# Proxy unit tests (tests/unit/proxy split alphabetically)
test-proxy-unit-a: install-test-deps

View file

@ -3166,10 +3166,11 @@ dependencies = [
"aws-smithy-types",
"bytes",
"futures-util",
"proptest",
"rstest",
"sse-stream",
"thiserror 2.0.19",
"tokio",
"tokio-util",
]
[[package]]
@ -5468,19 +5469,6 @@ dependencies = [
"wasm-bindgen",
]
[[package]]
name = "sse-stream"
version = "0.2.6"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "c25ac7aff0abd1dbc474536e40416e1102c7dd9bfba0b9861c6d357f835dcfb4"
dependencies = [
"bytes",
"futures-util",
"http-body 1.1.0",
"http-body-util",
"pin-project-lite",
]
[[package]]
name = "stable_deref_trait"
version = "1.2.1"

View file

@ -398,7 +398,7 @@ async fn unsupported_params_are_dropped_under_drop_params_and_rejected_without_i
api_key: Some("sk".into()),
api_base: Some(upstream.uri()),
shaping: MessagesShaping {
capabilities: capabilities.clone(),
capabilities,
drop_params,
..MessagesShaping::default()
},

View file

@ -8,16 +8,17 @@ repository.workspace = true
[features]
default = ["aws", "sse"]
aws = ["dep:aws-smithy-eventstream", "dep:aws-smithy-types"]
sse = ["dep:sse-stream"]
sse = []
[dependencies]
aws-smithy-eventstream = { version = "=0.61.4", optional = true }
aws-smithy-types = { version = "1.6.1", optional = true }
bytes = "1"
futures-util.workspace = true
sse-stream = { version = "=0.2.6", optional = true }
thiserror.workspace = true
tokio-util = { version = "0.7", features = ["codec", "io"] }
[dev-dependencies]
proptest.workspace = true
rstest.workspace = true
tokio.workspace = true

View file

@ -1,66 +1,47 @@
use bytes::{Buf, Bytes, BytesMut};
use futures_util::{Stream, StreamExt};
use aws_smithy_eventstream::frame::{read_message_from, write_message_to};
pub use aws_smithy_types::event_stream::{Header, HeaderValue, Message};
use bytes::BytesMut;
use tokio_util::codec::{Decoder, Encoder};
use aws_smithy_eventstream::frame::read_message_from;
use aws_smithy_types::event_stream::Header;
use crate::{Error, Framer};
use crate::EventStreamError;
const MIN_FRAME_BYTES: usize = 16;
const MAX_FRAME_BYTES: usize = 16 * 1024 * 1024;
#[derive(Clone, Debug, PartialEq)]
pub struct AwsEventStreamFrame {
pub headers: Vec<Header>,
pub payload: Bytes,
}
#[derive(Clone, Copy, Debug, Default)]
pub struct AwsEventStreamFramer;
pub struct AwsEventStreamCodec;
impl Framer for AwsEventStreamFramer {
type Frame = AwsEventStreamFrame;
impl Decoder for AwsEventStreamCodec {
type Item = Message;
type Error = EventStreamError;
fn frame<S, B, E>(self, input: S) -> impl Stream<Item = Result<Self::Frame, Error>> + Send
where
S: Stream<Item = Result<B, E>> + Send,
B: Buf + Send,
E: std::error::Error + Send + Sync + 'static,
{
futures_util::stream::try_unfold(
(Box::pin(input), BytesMut::new()),
|(mut input, mut buffer)| async move {
loop {
if buffer.len() >= 4 {
let length = (&buffer[..4]).get_u32() as usize;
if !(16..=MAX_FRAME_BYTES).contains(&length) {
return Err(Error::InvalidLength(length));
}
if buffer.len() >= length {
let raw = buffer.split_to(length).freeze();
let message = read_message_from(raw)?;
let frame = AwsEventStreamFrame {
headers: message.headers().to_vec(),
payload: message.payload().clone(),
};
return Ok(Some((frame, (input, buffer))));
}
}
match input.next().await {
Some(Ok(mut chunk)) => {
while chunk.has_remaining() {
let bytes = chunk.chunk();
buffer.extend_from_slice(bytes);
let length = bytes.len();
chunk.advance(length);
}
}
Some(Err(error)) => return Err(Error::Body(Box::new(error))),
None if buffer.is_empty() => return Ok(None),
None => return Err(Error::Truncated),
}
}
},
)
.fuse()
fn decode(&mut self, src: &mut BytesMut) -> Result<Option<Message>, EventStreamError> {
let Some(prefix) = src.first_chunk::<4>() else {
return Ok(None);
};
let length = u32::from_be_bytes(*prefix) as usize;
if !(MIN_FRAME_BYTES..=MAX_FRAME_BYTES).contains(&length) {
return Err(EventStreamError::InvalidLength(length));
}
if src.len() < length {
return Ok(None);
}
Ok(Some(read_message_from(src.split_to(length).freeze())?))
}
fn decode_eof(&mut self, src: &mut BytesMut) -> Result<Option<Message>, EventStreamError> {
match self.decode(src)? {
Some(message) => Ok(Some(message)),
None if src.is_empty() => Ok(None),
None => Err(EventStreamError::Truncated),
}
}
}
impl Encoder<Message> for AwsEventStreamCodec {
type Error = EventStreamError;
fn encode(&mut self, message: Message, dst: &mut BytesMut) -> Result<(), EventStreamError> {
Ok(write_message_to(&message, dst)?)
}
}

View file

@ -1,17 +1,21 @@
#[cfg(feature = "sse")]
#[derive(Debug, thiserror::Error)]
pub enum Error {
#[cfg(feature = "sse")]
#[error("SSE framing failed: {0}")]
Sse(#[from] sse_stream::Error),
#[cfg(feature = "aws")]
#[error("AWS EventStream framing failed: {0}")]
Aws(#[from] aws_smithy_eventstream::error::Error),
pub enum SseError {
#[error("body stream failed: {0}")]
Body(#[source] Box<dyn std::error::Error + Send + Sync>),
#[cfg(feature = "aws")]
Body(#[from] std::io::Error),
#[error("SSE field is not UTF-8: {0}")]
InvalidUtf8(#[from] std::str::Utf8Error),
}
#[cfg(feature = "aws")]
#[derive(Debug, thiserror::Error)]
pub enum EventStreamError {
#[error("body stream failed: {0}")]
Body(#[from] std::io::Error),
#[error("invalid AWS EventStream frame length: {0}")]
InvalidLength(usize),
#[cfg(feature = "aws")]
#[error("truncated AWS EventStream frame")]
Truncated,
#[error("malformed AWS EventStream frame: {0}")]
Malformed(#[from] aws_smithy_eventstream::error::Error),
}

View file

@ -0,0 +1,21 @@
use std::io;
use bytes::Buf;
use futures_util::{Stream, StreamExt, TryStreamExt};
use tokio_util::{
codec::{Decoder, FramedRead},
io::StreamReader,
};
pub fn frames<S, B, E, D>(
input: S,
codec: D,
) -> impl Stream<Item = Result<D::Item, D::Error>> + Send
where
S: Stream<Item = Result<B, E>> + Send,
B: Buf + Send,
E: std::error::Error + Send + Sync + 'static,
D: Decoder + Send,
{
FramedRead::new(StreamReader::new(input.map_err(io::Error::other)), codec).fuse()
}

View file

@ -1,8 +1,8 @@
mod error;
mod framer;
mod framed;
pub use error::*;
pub use framer::*;
pub use framed::frames;
#[cfg(feature = "aws")]
pub mod aws_event_stream;

View file

@ -1,43 +1,170 @@
use futures_util::{Stream, StreamExt};
use std::str;
use crate::{Error, Framer};
use bytes::{Buf, BufMut, BytesMut};
use tokio_util::codec::{Decoder, Encoder};
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct SseFrame {
use crate::SseError;
const BOM: &[u8] = b"\xEF\xBB\xBF";
#[derive(Clone, Debug, Default, PartialEq, Eq)]
pub struct SseEvent {
pub event: Option<String>,
pub data: Option<String>,
pub data: String,
pub id: Option<String>,
pub retry: Option<u64>,
}
#[derive(Clone, Copy, Debug, Default)]
pub struct SseFramer;
pub struct SseCodec {
past_bom: bool,
}
impl Framer for SseFramer {
type Frame = SseFrame;
impl Decoder for SseCodec {
type Item = SseEvent;
type Error = SseError;
fn frame<S, B, E>(self, input: S) -> impl Stream<Item = Result<SseFrame, Error>> + Send
where
S: Stream<Item = Result<B, E>> + Send,
B: bytes::Buf + Send,
E: std::error::Error + Send + Sync + 'static,
{
let frames = Box::pin(sse_stream::SseStream::from_bytes_stream(input));
futures_util::stream::try_unfold(frames, |mut frames| async move {
let Some(frame) = frames.next().await else {
return Ok(None);
};
let frame = frame?;
Ok(Some((
SseFrame {
event: frame.event,
data: frame.data,
id: frame.id,
retry: frame.retry,
},
frames,
)))
})
.fuse()
fn decode(&mut self, src: &mut BytesMut) -> Result<Option<SseEvent>, SseError> {
if !self.skip_bom(src) {
return Ok(None);
}
while let Some(end) = block_end(src) {
let block = src.split_to(end);
let pending = lines(&block)
.map(|(line, _)| line)
.take_while(|line| !line.is_empty())
.try_fold(Pending::default(), Pending::apply)?;
if let Some(event) = pending.dispatch() {
return Ok(Some(event));
}
}
Ok(None)
}
fn decode_eof(&mut self, _pending: &mut BytesMut) -> Result<Option<SseEvent>, SseError> {
Ok(None)
}
}
impl SseCodec {
fn skip_bom(&mut self, src: &mut BytesMut) -> bool {
if self.past_bom {
return true;
}
if src.starts_with(BOM) {
src.advance(BOM.len());
} else if BOM.starts_with(src) {
return false;
}
self.past_bom = true;
true
}
}
fn block_end(bytes: &[u8]) -> Option<usize> {
lines(bytes)
.find(|(line, _)| line.is_empty())
.map(|(_, end)| end)
}
fn lines(bytes: &[u8]) -> impl Iterator<Item = (&[u8], usize)> {
let mut cursor: usize = 0;
std::iter::from_fn(move || {
let rest = &bytes[cursor..];
let end = rest.iter().position(|byte| matches!(byte, b'\n' | b'\r'))?;
cursor += end + terminator_len(&rest[end..]);
Some((&rest[..end], cursor))
})
}
fn terminator_len(terminated: &[u8]) -> usize {
match terminated {
[b'\r', b'\n', ..] => 2,
_ => 1,
}
}
#[derive(Default)]
struct Pending {
event: Option<String>,
data: Option<String>,
id: Option<String>,
retry: Option<u64>,
}
impl Pending {
fn apply(self, line: &[u8]) -> Result<Self, SseError> {
let (name, value) = split_field(line);
Ok(match name {
b"event" => Self {
event: Some(str::from_utf8(value)?.to_owned()),
..self
},
b"data" => Self {
data: Some(append_data(self.data, str::from_utf8(value)?)),
..self
},
b"id" if !value.contains(&0) => Self {
id: Some(str::from_utf8(value)?.to_owned()),
..self
},
b"retry" => Self {
retry: parse_retry(value).or(self.retry),
..self
},
_ => self,
})
}
fn dispatch(self) -> Option<SseEvent> {
Some(SseEvent {
event: self.event,
data: self.data?,
id: self.id,
retry: self.retry,
})
}
}
fn split_field(line: &[u8]) -> (&[u8], &[u8]) {
let Some(colon) = line.iter().position(|byte| *byte == b':') else {
return (line, &[]);
};
let value = &line[colon + 1..];
(&line[..colon], value.strip_prefix(b" ").unwrap_or(value))
}
fn append_data(buffer: Option<String>, line: &str) -> String {
match buffer {
Some(existing) => format!("{existing}\n{line}"),
None => line.to_owned(),
}
}
fn parse_retry(value: &[u8]) -> Option<u64> {
if !value.iter().all(u8::is_ascii_digit) {
return None;
}
str::from_utf8(value).ok()?.parse().ok()
}
impl Encoder<SseEvent> for SseCodec {
type Error = SseError;
fn encode(&mut self, event: SseEvent, dst: &mut BytesMut) -> Result<(), SseError> {
if let Some(name) = event.event {
dst.put_slice(format!("event: {name}\n").as_bytes());
}
for line in event.data.split('\n') {
dst.put_slice(format!("data: {line}\n").as_bytes());
}
if let Some(id) = event.id {
dst.put_slice(format!("id: {id}\n").as_bytes());
}
if let Some(retry) = event.retry {
dst.put_slice(format!("retry: {retry}\n").as_bytes());
}
dst.put_u8(b'\n');
Ok(())
}
}

View file

@ -4,89 +4,174 @@ mod support;
use std::io;
use futures_util::TryStreamExt;
use litellm_framing::aws_event_stream::{AwsEventStreamFrame, AwsEventStreamFramer};
use litellm_framing::{Error, Framer};
use bytes::Bytes;
use futures_util::{StreamExt, TryStreamExt, stream};
use litellm_framing::{
EventStreamError,
aws_event_stream::{AwsEventStreamCodec, Header, HeaderValue, Message},
frames,
};
use proptest::prelude::*;
use rstest::{fixture, rstest};
use support::{body_cause, cut_at, encode_all, every, input, runtime};
use support::encode;
async fn collect_aws(bytes: &[u8], chunk_size: usize) -> Result<Vec<AwsEventStreamFrame>, Error> {
AwsEventStreamFramer
.frame(futures_util::stream::iter(
bytes.chunks(chunk_size).map(Ok::<_, io::Error>),
))
async fn collect(pieces: Vec<Bytes>) -> Result<Vec<Message>, EventStreamError> {
frames(input(pieces), AwsEventStreamCodec)
.try_collect()
.await
}
#[fixture]
fn two_frames() -> Vec<u8> {
[encode(b"\xff\x00"), encode(b"second")].concat()
fn message(payload: &[u8]) -> Message {
Message::new(Bytes::copy_from_slice(payload))
.add_header(Header::new(
":event-type",
HeaderValue::String("payload".into()),
))
.add_header(Header::new("sequence", HeaderValue::Int32(7)))
}
#[fixture]
fn payload_frame() -> Vec<u8> {
encode(b"payload")
encode_all(AwsEventStreamCodec, [message(b"payload")])
}
fn header_value() -> impl Strategy<Value = HeaderValue> {
prop_oneof![
"[a-z]{0,8}".prop_map(|text| HeaderValue::String(text.into())),
any::<i32>().prop_map(HeaderValue::Int32),
any::<bool>().prop_map(HeaderValue::Bool),
proptest::collection::vec(any::<u8>(), 0..8)
.prop_map(|bytes| HeaderValue::ByteArray(bytes.into())),
]
}
fn arbitrary_message() -> impl Strategy<Value = Message> {
(
proptest::collection::vec(("[a-z:-]{1,12}", header_value()), 0..3),
proptest::collection::vec(any::<u8>(), 0..32),
)
.prop_map(|(headers, payload)| {
headers.into_iter().fold(
Message::new(Bytes::from(payload)),
|message, (name, value)| message.add_header(Header::new(name, value)),
)
})
}
proptest! {
#[test]
fn any_messages_survive_a_round_trip_through_any_cuts(
messages in proptest::collection::vec(arbitrary_message(), 1..4),
cuts in proptest::collection::vec(0_usize..512, 0..4),
) {
let wire = encode_all(AwsEventStreamCodec, messages.clone());
let decoded = runtime().block_on(collect(cut_at(&wire, cuts))).unwrap();
prop_assert_eq!(decoded, messages);
}
}
#[rstest]
#[case(1)]
#[case(3)]
#[case(12)]
#[case(usize::MAX)]
#[case::prelude_crc(8)]
#[case::message_crc(usize::MAX)]
#[tokio::test]
async fn fragmented_and_coalesced_frames_preserve_typed_headers_and_binary_payloads(
two_frames: Vec<u8>,
#[case] chunk_size: usize,
) {
let chunk_size = chunk_size.min(two_frames.len());
let frames = collect_aws(&two_frames, chunk_size).await.unwrap();
assert_eq!(frames.len(), 2);
assert_eq!(frames[0].payload, &b"\xff\x00"[..]);
assert_eq!(frames[1].payload, "second");
assert_eq!(
frames[0].headers[0].value().as_string().unwrap().as_str(),
"payload"
);
assert_eq!(frames[0].headers[1].value().as_int32(), Ok(7));
}
#[rstest]
#[case(8)]
#[case(0)]
#[tokio::test]
async fn rejects_corrupt_crcs(payload_frame: Vec<u8>, #[case] index: usize) {
let corrupt_index = if index == 0 {
payload_frame.len() - 1
} else {
index
};
async fn a_corrupt_crc_is_malformed(payload_frame: Vec<u8>, #[case] index: usize) {
let mut corrupt = payload_frame;
corrupt[corrupt_index] ^= 1;
assert!(matches!(collect_aws(&corrupt, 3).await, Err(Error::Aws(_))));
}
#[rstest]
#[case(0_u32)]
#[case(15)]
#[case(u32::MAX)]
#[tokio::test]
async fn rejects_invalid_lengths(#[case] length: u32) {
let flipped = index.min(corrupt.len() - 1);
corrupt[flipped] ^= 1;
assert!(matches!(
collect_aws(&length.to_be_bytes(), 1).await,
Err(Error::InvalidLength(_))
collect(every(&corrupt, 3)).await,
Err(EventStreamError::Malformed(_))
));
}
#[rstest]
#[case(1)]
#[case(3)]
#[case(5)]
#[case::zero(0)]
#[case::below_minimum(15)]
#[case::above_maximum(16 * 1024 * 1024 + 1)]
#[case::u32_max(u32::MAX)]
#[tokio::test]
async fn rejects_truncation(payload_frame: Vec<u8>, #[case] end: usize) {
async fn a_length_outside_the_frame_bounds_fails_before_buffering(#[case] length: u32) {
assert!(matches!(
collect_aws(&payload_frame[..end], 1).await,
Err(Error::Truncated)
collect(every(&length.to_be_bytes(), 1)).await,
Err(EventStreamError::InvalidLength(seen)) if seen == length as usize
));
}
#[rstest]
#[case::before_the_length(1)]
#[case::inside_the_prelude(5)]
#[case::one_byte_short(usize::MAX)]
#[tokio::test]
async fn eof_inside_a_frame_is_truncation(payload_frame: Vec<u8>, #[case] end: usize) {
let end = end.min(payload_frame.len() - 1);
assert!(matches!(
collect(every(&payload_frame[..end], 1)).await,
Err(EventStreamError::Truncated)
));
}
const FRAME_OVERHEAD_BYTES: usize = 16;
const MAX_FRAME_BYTES: usize = 16 * 1024 * 1024;
#[tokio::test]
async fn a_frame_at_exactly_the_maximum_length_decodes() {
let largest = Message::new(vec![0xAB; MAX_FRAME_BYTES - FRAME_OVERHEAD_BYTES]);
let wire = encode_all(AwsEventStreamCodec, [largest.clone()]);
assert_eq!(wire.len(), MAX_FRAME_BYTES);
assert_eq!(collect(every(&wire, 1 << 20)).await.unwrap(), vec![largest]);
}
#[tokio::test]
async fn a_frame_one_byte_over_the_maximum_length_is_rejected_by_its_prelude() {
let oversized = Message::new(vec![0xAB; MAX_FRAME_BYTES - FRAME_OVERHEAD_BYTES + 1]);
let wire = encode_all(AwsEventStreamCodec, [oversized]);
assert!(matches!(
collect(every(&wire[..4], 1)).await,
Err(EventStreamError::InvalidLength(length)) if length == MAX_FRAME_BYTES + 1
));
}
#[tokio::test]
async fn an_empty_body_yields_nothing() {
assert_eq!(collect(vec![]).await.unwrap(), vec![]);
}
#[tokio::test]
async fn a_complete_frame_precedes_a_truncated_following_frame() {
let wire = encode_all(AwsEventStreamCodec, [message(b"first"), message(b"second")]);
let mut messages = Box::pin(frames(
input(every(&wire[..wire.len() - 1], 3)),
AwsEventStreamCodec,
));
assert_eq!(messages.next().await.unwrap().unwrap(), message(b"first"));
assert!(matches!(
messages.next().await,
Some(Err(EventStreamError::Truncated))
));
assert!(messages.next().await.is_none());
}
#[tokio::test]
async fn a_body_error_after_a_complete_frame_preserves_its_cause() {
let first = encode_all(AwsEventStreamCodec, [message(b"first")]);
let mut messages = Box::pin(frames(
stream::iter([
Ok(cut_at(&first, [5])[0].clone()),
Ok(cut_at(&first, [5])[1].clone()),
Ok(Bytes::from_static(b"\0\0\0")),
Err(io::Error::new(io::ErrorKind::ConnectionReset, "reset")),
]),
AwsEventStreamCodec,
));
assert_eq!(messages.next().await.unwrap().unwrap(), message(b"first"));
let Some(Err(EventStreamError::Body(body))) = messages.next().await else {
panic!("the body error surfaces");
};
assert_eq!(
body_cause::<io::Error>(&body).unwrap().kind(),
io::ErrorKind::ConnectionReset
);
assert!(messages.next().await.is_none());
}

View file

@ -2,28 +2,64 @@
mod support;
use std::io;
use bytes::Bytes;
use futures_util::{StreamExt, TryStreamExt};
use litellm_framing::{
EventStreamError, SseError,
aws_event_stream::{AwsEventStreamCodec, Message},
frames,
sse::{SseCodec, SseEvent},
};
use proptest::prelude::*;
use support::{body_cause, cut_at, encode_all, every, input, runtime};
use futures_util::TryStreamExt;
use litellm_framing::Framer;
use litellm_framing::aws_event_stream::{AwsEventStreamFrame, AwsEventStreamFramer};
use litellm_framing::sse::SseFramer;
fn delta(data: &str) -> SseEvent {
SseEvent {
event: Some("delta".into()),
data: data.into(),
id: Some("7".into()),
retry: None,
}
}
use support::encode;
fn envelopes(payloads: Vec<Bytes>) -> Vec<u8> {
encode_all(AwsEventStreamCodec, payloads.into_iter().map(Message::new))
}
proptest! {
#[test]
fn an_sse_event_cut_anywhere_across_envelopes_is_reassembled(cut in 0_usize..64, chunk in 1_usize..8) {
let sse = encode_all(SseCodec::default(), [delta("hello")]);
let wire = envelopes(cut_at(&sse, [cut.min(sse.len())]));
let events = runtime().block_on(async {
let payloads = frames(input(every(&wire, chunk)), AwsEventStreamCodec)
.map_ok(|message| message.payload().clone());
frames(payloads, SseCodec::default()).try_collect::<Vec<_>>().await
})
.unwrap();
prop_assert_eq!(events, vec![delta("hello")]);
}
}
#[tokio::test]
async fn hosting_payloads_feed_the_same_sse_framer_across_envelope_boundaries() {
let bytes = [encode(b"event: delta\ndata: hel"), encode(b"lo\nid: 7\n\n")].concat();
let envelopes = AwsEventStreamFramer.frame(futures_util::stream::iter(
bytes.chunks(3).map(Ok::<_, io::Error>),
async fn a_truncated_envelope_after_an_sse_event_keeps_the_event_and_its_cause() {
let complete = encode_all(SseCodec::default(), [delta("complete")]);
let incomplete = encode_all(SseCodec::default(), [delta("incomplete")]);
let wire = envelopes(vec![complete.into(), incomplete.into()]);
let payloads = frames(
input(every(&wire[..wire.len() - 1], 3)),
AwsEventStreamCodec,
)
.map_ok(|message| message.payload().clone());
let mut events = Box::pin(frames(payloads, SseCodec::default()));
assert_eq!(events.next().await.unwrap().unwrap(), delta("complete"));
let Some(Err(SseError::Body(body))) = events.next().await else {
panic!("the envelope error surfaces through the SSE layer");
};
assert!(matches!(
body_cause::<EventStreamError>(&body),
Some(EventStreamError::Truncated)
));
let frames = SseFramer
.frame(envelopes.map_ok(|frame: AwsEventStreamFrame| frame.payload))
.try_collect::<Vec<_>>()
.await
.unwrap();
assert_eq!(frames.len(), 1);
assert_eq!(frames[0].event.as_deref(), Some("delta"));
assert_eq!(frames[0].data.as_deref(), Some("hello"));
assert_eq!(frames[0].id.as_deref(), Some("7"));
assert!(events.next().await.is_none());
}

View file

@ -1,67 +1,169 @@
#![cfg(feature = "sse")]
mod support;
use std::io;
use futures_util::{StreamExt, TryStreamExt};
use litellm_framing::sse::{SseFrame, SseFramer};
use litellm_framing::{Error, Framer};
use bytes::Bytes;
use futures_util::{StreamExt, TryStreamExt, stream};
use litellm_framing::{
SseError, frames,
sse::{SseCodec, SseEvent},
};
use proptest::prelude::*;
use rstest::rstest;
use support::{body_cause, cut_at, encode_all, every, input, runtime};
async fn collect_sse(chunks: &[&[u8]]) -> Result<Vec<SseFrame>, Error> {
SseFramer
.frame(futures_util::stream::iter(
chunks.iter().copied().map(Ok::<_, io::Error>),
))
async fn collect(pieces: Vec<Bytes>) -> Result<Vec<SseEvent>, SseError> {
frames(input(pieces), SseCodec::default())
.try_collect()
.await
}
fn event(name: Option<&str>, data: &str) -> SseEvent {
SseEvent {
event: name.map(str::to_owned),
data: data.to_owned(),
id: None,
retry: None,
}
}
fn sse_event() -> impl Strategy<Value = SseEvent> {
(
proptest::option::of("[^\r\n\0]{0,8}"),
"[^\r\0]{0,16}",
proptest::option::of("[^\r\n\0]{0,8}"),
proptest::option::of(any::<u64>()),
)
.prop_map(|(event, data, id, retry)| SseEvent {
event,
data,
id,
retry,
})
}
fn terminators() -> impl Strategy<Value = &'static [u8]> {
prop_oneof![Just(&b"\n"[..]), Just(&b"\r\n"[..]), Just(&b"\r"[..])]
}
proptest! {
#[test]
fn any_events_survive_a_round_trip_through_any_terminator_and_any_cuts(
events in proptest::collection::vec(sse_event(), 1..4),
terminator in terminators(),
cuts in proptest::collection::vec(0_usize..256, 0..4),
bom in any::<bool>(),
) {
let lf_wire = encode_all(SseCodec::default(), events.clone());
let body: Vec<u8> = lf_wire
.iter()
.flat_map(|byte| if *byte == b'\n' { terminator.to_vec() } else { vec![*byte] })
.collect();
let wire = if bom { [&b"\xEF\xBB\xBF"[..], &body].concat() } else { body };
let decoded = runtime().block_on(collect(cut_at(&wire, cuts))).unwrap();
prop_assert_eq!(decoded, events);
}
}
#[rstest]
#[case(
&[&b":ping\r\nevent: delta\r\nid: 7\r\nretry: 10\r\ndata: \xe2"[..], &b"\x82"[..], &b"\xac\r"[..], &b"\ndata: next\r\n\r"[..], &b"\ndata: [DONE]\n\n"[..]],
vec![
SseFrame {
event: Some("delta".into()),
data: Some("€\nnext".into()),
id: Some("7".into()),
retry: Some(10),
},
SseFrame {
event: None,
data: Some("[DONE]".into()),
id: None,
retry: None,
},
]
)]
#[case::comment(b":ping\ndata: x\n\n")]
#[case::unknown_field(b"vendor: 1\ndata: x\n\n")]
#[case::field_without_colon(b"garbage\ndata: x\n\n")]
#[case::retry_with_non_digits(b"retry: soon\ndata: x\n\n")]
#[case::retry_with_a_sign(b"retry: +5\ndata: x\n\n")]
#[case::retry_without_a_value(b"retry:\ndata: x\n\n")]
#[case::id_with_nul(b"id: a\0b\ndata: x\n\n")]
#[tokio::test]
async fn fragmented_utf8_crlf_and_multiline_data_retain_metadata_and_sentinel(
#[case] chunks: &[&[u8]],
#[case] expected: Vec<SseFrame>,
async fn lines_the_spec_ignores_do_not_change_the_event(#[case] wire: &[u8]) {
assert_eq!(
collect(every(wire, 1)).await.unwrap(),
vec![event(None, "x")]
);
}
#[rstest]
#[case::no_data_at_all(b"event: ping\nid: 1\n\ndata: x\n\n", vec![event(None, "x")])]
#[case::empty_data_field(b"data:\n\n", vec![event(None, "")])]
#[case::one_leading_space_stripped(b"data: x\n\n", vec![event(None, " x")])]
#[case::multiline_data(b"data: a\ndata: b\ndata:\n\n", vec![event(None, "a\nb\n")])]
#[case::last_event_name_wins(b"event: a\nevent: b\ndata: x\n\n", vec![event(Some("b"), "x")])]
#[case::last_retry_wins(b"retry: 1\nretry: 2\ndata: x\n\n", vec![SseEvent { retry: Some(2), ..event(None, "x") }])]
#[case::split_utf8_across_lines_is_not_joined(b"data: \xe2\x82\xac\ndata: \xe2\x82\xac\n\n", vec![event(None, "€\n€")])]
#[tokio::test]
async fn dispatch_follows_the_data_buffer(#[case] wire: &[u8], #[case] expected: Vec<SseEvent>) {
assert_eq!(collect(every(wire, 1)).await.unwrap(), expected);
}
#[rstest]
#[case::unterminated_single(b"data: partial\n", vec![])]
#[case::unterminated_tail_after_complete(b"data: complete\n\ndata: unfinished\n", vec![event(None, "complete")])]
#[case::lone_cr_terminates_at_eof(b"data: x\r\r", vec![event(None, "x")])]
#[case::lone_cr_line_then_eof(b"data: x\r", vec![])]
#[tokio::test]
async fn eof_dispatches_only_terminated_events(
#[case] wire: &[u8],
#[case] expected: Vec<SseEvent>,
) {
assert_eq!(collect_sse(chunks).await.unwrap(), expected);
assert_eq!(
collect(vec![Bytes::copy_from_slice(wire)]).await.unwrap(),
expected
);
}
#[rstest]
#[case::inside_the_first_line(vec![&b"data: a\r"[..], &b"\ndata: b\r\n\r\n"[..]])]
#[case::inside_the_blank_line(vec![&b"data: a\r\ndata: b\r\n\r"[..], &b"\n"[..]])]
#[tokio::test]
async fn a_crlf_split_across_chunks_is_one_terminator(#[case] pieces: Vec<&[u8]>) {
let pieces = pieces.into_iter().map(Bytes::copy_from_slice).collect();
assert_eq!(collect(pieces).await.unwrap(), vec![event(None, "a\nb")]);
}
#[tokio::test]
async fn eof_does_not_dispatch_an_unterminated_frame() {
assert!(collect_sse(&[b"data: partial\n"]).await.unwrap().is_empty());
async fn a_bom_is_stripped_only_at_the_start_of_the_stream() {
let wire = b"\xEF\xBB\xBFdata: a\n\n\xEF\xBB\xBFdata: b\ndata: c\n\n";
let decoded = collect(every(wire, 2)).await.unwrap();
assert_eq!(decoded, vec![event(None, "a"), event(None, "c")]);
}
#[tokio::test]
async fn invalid_utf8_in_a_field_fails_after_earlier_events_and_terminates() {
let mut events = Box::pin(frames(
input(every(b"data: ok\n\ndata: \xff\n\n", 3)),
SseCodec::default(),
));
assert_eq!(events.next().await.unwrap().unwrap(), event(None, "ok"));
assert!(matches!(
events.next().await,
Some(Err(SseError::InvalidUtf8(_)))
));
assert!(events.next().await.is_none());
}
#[rstest]
#[case(io::ErrorKind::ConnectionReset)]
#[case(io::ErrorKind::UnexpectedEof)]
#[tokio::test]
async fn framing_errors_terminate_and_preserve_input_error_causes(#[case] kind: io::ErrorKind) {
let mut frames = Box::pin(SseFramer.frame(futures_util::stream::iter([
Err(io::Error::new(kind, "reset")),
Ok(&b"data: later\n\n"[..]),
])));
let error = frames.next().await.unwrap().unwrap_err();
assert!(matches!(
error,
Error::Sse(sse_stream::Error::Body(ref cause))
if cause.downcast_ref::<io::Error>().unwrap().kind() == kind
async fn a_body_error_keeps_earlier_events_and_its_cause_then_terminates(
#[case] kind: io::ErrorKind,
) {
let mut events = Box::pin(frames(
stream::iter([
Ok(&b"data: first\n\ndata: partial"[..]),
Err(io::Error::new(kind, "reset")),
Ok(&b"\n\n"[..]),
]),
SseCodec::default(),
));
assert!(frames.next().await.is_none());
assert!(frames.next().await.is_none());
assert_eq!(events.next().await.unwrap().unwrap(), event(None, "first"));
let Some(Err(SseError::Body(body))) = events.next().await else {
panic!("the body error surfaces");
};
assert_eq!(body_cause::<io::Error>(&body).unwrap().kind(), kind);
assert!(events.next().await.is_none());
assert!(events.next().await.is_none());
}

View file

@ -1,15 +1,57 @@
use aws_smithy_eventstream::frame::write_message_to;
use aws_smithy_types::event_stream::{Header, HeaderValue, Message};
use bytes::Bytes;
#![allow(dead_code)]
pub fn encode(payload: &'static [u8]) -> Vec<u8> {
let message = Message::new(Bytes::from_static(payload))
.add_header(Header::new(
":event-type",
HeaderValue::String("payload".into()),
))
.add_header(Header::new("sequence", HeaderValue::Int32(7)));
let mut bytes = Vec::new();
write_message_to(&message, &mut bytes).unwrap();
bytes
use std::{error::Error, io};
use bytes::{Bytes, BytesMut};
use futures_util::{Stream, stream};
use tokio_util::codec::Encoder;
pub fn encode_all<C, I>(mut codec: C, items: impl IntoIterator<Item = I>) -> Vec<u8>
where
C: Encoder<I>,
C::Error: std::fmt::Debug,
{
let mut wire = BytesMut::new();
for item in items {
codec.encode(item, &mut wire).unwrap();
}
wire.to_vec()
}
pub fn cut_at(bytes: &[u8], offsets: impl IntoIterator<Item = usize>) -> Vec<Bytes> {
let mut sorted: Vec<usize> = offsets
.into_iter()
.filter(|offset| *offset <= bytes.len())
.collect();
sorted.sort_unstable();
sorted.dedup();
let bounds = std::iter::once(0)
.chain(sorted)
.chain(std::iter::once(bytes.len()))
.collect::<Vec<_>>();
bounds
.windows(2)
.map(|pair| Bytes::copy_from_slice(&bytes[pair[0]..pair[1]]))
.collect()
}
pub fn every(bytes: &[u8], size: usize) -> Vec<Bytes> {
bytes
.chunks(size.max(1))
.map(Bytes::copy_from_slice)
.collect()
}
pub fn input(pieces: Vec<Bytes>) -> impl Stream<Item = Result<Bytes, io::Error>> + Send {
stream::iter(pieces.into_iter().map(Ok))
}
pub fn body_cause<T: Error + 'static>(body: &io::Error) -> Option<&T> {
body.get_ref()?.downcast_ref::<T>()
}
pub fn runtime() -> tokio::runtime::Runtime {
tokio::runtime::Builder::new_current_thread()
.build()
.unwrap()
}

View file

@ -2,9 +2,9 @@ use base64::Engine;
use bytes::Buf;
use futures_util::{Stream, StreamExt};
use litellm_framing::{
Framer,
aws_event_stream::{AwsEventStreamFrame, AwsEventStreamFramer},
sse::{SseFrame, SseFramer},
aws_event_stream::{AwsEventStreamCodec, Message},
frames,
sse::{SseCodec, SseEvent},
};
use serde::{Deserialize, Serialize};
use serde_json::{Map, Value};
@ -13,8 +13,6 @@ use serde_json::{Map, Value};
pub enum Error {
#[error("stream framing failed: {0}")]
StreamFraming(String),
#[error("Anthropic SSE frame has no data")]
MissingStreamData,
#[error("Anthropic stream event is invalid: {0}")]
InvalidStreamEvent(String),
#[error("Bedrock event payload is invalid: {0}")]
@ -165,15 +163,14 @@ struct BedrockChunkPayload {
bytes: String,
}
pub fn decode_anthropic_sse_frame(frame: SseFrame) -> Result<AnthropicMessagesStreamEvent, Error> {
let data = frame.data.ok_or(Error::MissingStreamData)?;
serde_json::from_str(&data).map_err(|error| Error::InvalidStreamEvent(error.to_string()))
pub fn decode_anthropic_sse_frame(event: SseEvent) -> Result<AnthropicMessagesStreamEvent, Error> {
serde_json::from_str(&event.data).map_err(|error| Error::InvalidStreamEvent(error.to_string()))
}
pub fn decode_bedrock_anthropic_frame(
frame: AwsEventStreamFrame,
message: Message,
) -> Result<AnthropicMessagesStreamEvent, Error> {
let payload: BedrockChunkPayload = serde_json::from_slice(&frame.payload)
let payload: BedrockChunkPayload = serde_json::from_slice(message.payload())
.map_err(|error| Error::InvalidBedrockPayload(error.to_string()))?;
let event = base64::engine::general_purpose::STANDARD
.decode(payload.bytes)
@ -189,9 +186,8 @@ where
B: Buf + Send,
E: std::error::Error + Send + Sync + 'static,
{
SseFramer.frame(input).map(|frame| {
let frame = frame.map_err(|error| Error::StreamFraming(error.to_string()))?;
decode_anthropic_sse_frame(frame)
frames(input, SseCodec::default()).map(|event| {
decode_anthropic_sse_frame(event.map_err(|error| Error::StreamFraming(error.to_string()))?)
})
}
@ -203,9 +199,10 @@ where
B: Buf + Send,
E: std::error::Error + Send + Sync + 'static,
{
AwsEventStreamFramer.frame(input).map(|frame| {
let frame = frame.map_err(|error| Error::StreamFraming(error.to_string()))?;
decode_bedrock_anthropic_frame(frame)
frames(input, AwsEventStreamCodec).map(|message| {
decode_bedrock_anthropic_frame(
message.map_err(|error| Error::StreamFraming(error.to_string()))?,
)
})
}
@ -247,12 +244,10 @@ mod tests {
#[test]
fn decodes_citations_delta_events() {
let event = decode_anthropic_sse_frame(SseFrame {
let event = decode_anthropic_sse_frame(SseEvent {
event: Some("content_block_delta".into()),
data: Some(
r#"{"type":"content_block_delta","index":0,"delta":{"type":"citations_delta","citation":{"type":"char_location"}}}"#
.into(),
),
data: r#"{"type":"content_block_delta","index":0,"delta":{"type":"citations_delta","citation":{"type":"char_location"}}}"#
.into(),
id: None,
retry: None,
})

View file

@ -16,6 +16,7 @@ from litellm.litellm_core_utils.core_helpers import process_response_headers
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
from litellm.llms.anthropic.common_utils import ANTHROPIC_ERROR_STATUS_CODE_MAP
from litellm.llms.anthropic.experimental_pass_through.messages.utils import INCOMPLETE_STREAM_ERROR_MESSAGE
from litellm.proxy.pass_through_endpoints.success_handler import (
PassThroughEndpointLogging,
)
@ -28,11 +29,6 @@ GLOBAL_PASS_THROUGH_SUCCESS_HANDLER_OBJ: Final = PassThroughEndpointLogging()
_UPSTREAM_PUMP_TASKS: Final[set[asyncio.Task[None]]] = set() # mutable-ok: stdlib strong-ref set for pump tasks
_DETACHED_STREAM_DRAINS: Final[set[asyncio.Task[None]]] = set() # mutable-ok: bounded strong-ref set, detached drains
INCOMPLETE_STREAM_ERROR_MESSAGE: Final = (
"Provider stream ended before emitting a message_stop event; "
"the response is incomplete and any partial content (e.g. tool_use input JSON) may be truncated."
)
def _is_message_stop_chunk(chunk: object) -> bool:
if isinstance(chunk, dict):

View file

@ -15,6 +15,12 @@ if TYPE_CHECKING:
from litellm.exceptions import ContentPolicyViolationError
INCOMPLETE_STREAM_ERROR_MESSAGE: Final = (
"Provider stream ended before emitting a message_stop event; "
"the response is incomplete and any partial content (e.g. tool_use input JSON) may be truncated."
)
def get_safeguard_refusal_stop_details(response: object) -> Mapping[str, Any] | None:
"""
Return the ``stop_details`` of an Anthropic Messages response refused by a

View file

@ -2,20 +2,25 @@
## Translates OpenAI call to Anthropic `/v1/messages` format
import asyncio
import json
import traceback
from collections import deque
from collections.abc import AsyncIterator, Iterator, Mapping
from typing import TYPE_CHECKING, Any, Final
from pydantic import BaseModel, ConfigDict, field_validator
from litellm import verbose_logger
from litellm._logging import redact_internal_details_from_client_message
from litellm._uuid import uuid
from litellm.exceptions import MidStreamFallbackError
from litellm.litellm_core_utils.prompt_templates.common_utils import (
encrypted_reasoning_signature,
)
from litellm.llms.anthropic.experimental_pass_through.messages.utils import (
INCOMPLETE_STREAM_ERROR_MESSAGE,
refusal_stop_details,
responses_output_refusal_text,
)
from litellm.responses.streaming_iterator import stream_error_status_and_message
from litellm.types.llms.anthropic_messages.anthropic_response import AnthropicUsage
from .transformation import (
@ -27,6 +32,72 @@ if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObject
class _UpstreamFailure(BaseModel):
model_config = ConfigDict(frozen=True)
status_code: int | None = None
message: str | None = None
@field_validator("status_code", mode="before")
@classmethod
def http_error_status_or_none(cls, value: object) -> int | None:
candidate: Final = (
value
if isinstance(value, int) and not isinstance(value, bool)
else int(value)
if isinstance(value, str) and value.isdecimal()
else None
)
return candidate if candidate is not None and 400 <= candidate <= 599 else None
@field_validator("message", mode="before")
@classmethod
def str_or_none(cls, value: object) -> str | None:
return value if isinstance(value, str) else None
class _FailedResponse(BaseModel):
model_config = ConfigDict(frozen=True, from_attributes=True)
error: object | None = None
class _FailedResponseEvent(BaseModel):
model_config = ConfigDict(frozen=True, from_attributes=True)
response: _FailedResponse | None = None
def _original_failure(exception: Exception) -> Exception:
failure = exception # rebind-ok: walks the MidStreamFallbackError chain down to the provider failure
while isinstance(failure, MidStreamFallbackError) and failure.original_exception is not None:
failure = failure.original_exception
return failure
def _failure_status_and_message(exception: Exception) -> tuple[int, str]:
original: Final = _original_failure(exception)
failure: Final = _UpstreamFailure.model_validate(
{"status_code": getattr(original, "status_code", None), "message": getattr(original, "message", None)}
)
status_code: Final = failure.status_code if failure.status_code is not None else 500
message: Final = failure.message or str(original) or INCOMPLETE_STREAM_ERROR_MESSAGE
return status_code, message
def _anthropic_error_chunk(status_code: int, message: str) -> dict[str, object]:
from litellm.anthropic_interface.exceptions.exception_mapping_utils import (
AnthropicExceptionMapping,
)
return dict(
AnthropicExceptionMapping.transform_to_anthropic_error(
status_code=status_code,
raw_message=redact_internal_details_from_client_message(message),
)
)
class AnthropicResponsesStreamWrapper:
"""
Wraps a Responses API streaming iterator and re-emits events in Anthropic SSE format.
@ -40,6 +111,7 @@ class AnthropicResponsesStreamWrapper:
response.function_call_arguments.delta -> content_block_delta (input_json_delta)
response.output_item.done -> content_block_delta (signature_delta) + content_block_stop
response.completed -> message_delta + message_stop
response.failed -> error (the stream ends without message_stop)
"""
def __init__(
@ -60,6 +132,7 @@ class AnthropicResponsesStreamWrapper:
self._pending_tool_ids: dict[str, str] = {} # item_id -> call_id / name accumulator
self._sent_message_start = False
self._sent_message_stop = False
self._stream_failed = False
self._chunk_queue: deque[dict[str, object]] = deque()
self._refusal_text: str = ""
self._sync_responses_iterator: Iterator[object] | None = None
@ -293,10 +366,23 @@ class AnthropicResponsesStreamWrapper:
)
return
if event_type == "response.failed":
failed: Final = _FailedResponseEvent.model_validate(event)
status_code, message = stream_error_status_and_message(
failed.response.error if failed.response is not None else None
)
verbose_logger.error(
"AnthropicResponsesStreamWrapper: upstream Responses stream for %s failed (%s): %s",
self.model,
status_code,
message,
)
self._fail_stream(status_code, message)
return
# ---- response completed -> message_delta + message_stop ----
if event_type in (
"response.completed",
"response.failed",
"response.incomplete",
):
response_obj: Final = getattr(event, "response", None) or (
@ -350,21 +436,24 @@ class AnthropicResponsesStreamWrapper:
self._sent_message_stop = True
return
def _fail_stream(self, status_code: int, message: str) -> None:
self._stream_failed = True
self._chunk_queue.append(_anthropic_error_chunk(status_code, message))
def __aiter__(self) -> "AnthropicResponsesStreamWrapper":
return self
async def __anext__(self) -> dict[str, object]:
# Return any queued chunks first
if self._chunk_queue:
return self._chunk_queue.popleft()
if self._stream_failed:
raise StopAsyncIteration
# Emit message_start if not yet done (fallback if response.created wasn't fired)
if not self._sent_message_start:
self._sent_message_start = True
self._chunk_queue.append(self._make_message_start())
return self._chunk_queue.popleft()
# Consume the upstream stream
try:
if hasattr(self.responses_stream, "__aiter__"):
async for event in self.responses_stream:
@ -382,10 +471,19 @@ class AnthropicResponsesStreamWrapper:
return self._chunk_queue.popleft()
except StopAsyncIteration:
pass
except Exception as e:
verbose_logger.error("AnthropicResponsesStreamWrapper error: %s\n%s", e, traceback.format_exc())
except Exception as e: # noqa: BLE001 # every upstream failure becomes a client error event
verbose_logger.exception(
"AnthropicResponsesStreamWrapper: upstream Responses stream for %s failed", self.model
)
self._fail_stream(*_failure_status_and_message(e))
if not self._chunk_queue and not self._sent_message_stop and not self._stream_failed:
verbose_logger.error(
"AnthropicResponsesStreamWrapper: upstream Responses stream for %s ended without a terminal event",
self.model,
)
self._fail_stream(500, INCOMPLETE_STREAM_ERROR_MESSAGE)
# Drain any remaining queued chunks
if self._chunk_queue:
return self._chunk_queue.popleft()

View file

@ -4618,6 +4618,13 @@ class GoogleSSOHandler:
return result or {}
def _raise_if_sso_debug_disabled() -> None:
"""The debug routes run the browser-redirect SSO flow, so they cannot carry a
bearer credential; an explicit opt-in flag is the only way to gate them."""
if get_secret_bool("ENABLE_SSO_DEBUG") is not True:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Not Found")
@router.get("/sso/debug/login", tags=["experimental"], include_in_schema=False)
async def debug_sso_login(request: Request):
"""
@ -4625,6 +4632,8 @@ async def debug_sso_login(request: Request):
PROXY_BASE_URL should be the your deployed proxy endpoint, e.g. PROXY_BASE_URL="https://litellm-production-7002.up.railway.app/"
Example:
"""
_raise_if_sso_debug_disabled()
from litellm.proxy.proxy_server import premium_user
microsoft_client_id: Final = os.getenv("MICROSOFT_CLIENT_ID", None)
@ -4670,6 +4679,8 @@ async def debug_sso_callback(request: Request):
"""
Returns the OpenID object returned by the SSO provider
"""
_raise_if_sso_debug_disabled()
import json
from fastapi.responses import HTMLResponse

View file

@ -230,6 +230,11 @@ def _status_code_for_error_fields(error_type: str | None, error_code: str | None
return next((status for status in map(_status_code_for_error_field, fields) if status is not None), 500)
def stream_error_status_and_message(error_obj: object) -> tuple[int, str]:
message, error_type, error_code = _error_event_fields(error_obj)
return _status_code_for_error_fields(error_type, error_code), message
def _map_stream_error_to_exception(error_obj: object, model: str, custom_llm_provider: str) -> Exception:
from litellm.llms.base_llm.chat.transformation import BaseLLMException

View file

@ -52,7 +52,7 @@ from tests._vcr_redis_persister import (
# network call entirely, so skip tests record nothing (NOOP) and passing tests
# stop carrying a volatile github episode. This matches the established idiom in
# the unit-test suite, which sets the same flag (see e.g.
# tests/test_litellm/test_cost_calculator.py). ``setdefault`` so an explicit
# tests/unit/test_cost_calculator.py). ``setdefault`` so an explicit
# override still wins.
os.environ.setdefault("LITELLM_LOCAL_MODEL_COST_MAP", "True")

View file

@ -13,15 +13,16 @@ def check_for_litellm_module_deletion(base_dir):
del sys.modules[module]
"""
problematic_files = []
test_dir = os.path.join(base_dir, "test_litellm")
candidate_dirs = [os.path.join(base_dir, name) for name in ("test_litellm", "unit")]
test_dirs = [test_dir for test_dir in candidate_dirs if os.path.exists(test_dir)]
if not os.path.exists(test_dir):
print(f"Warning: Directory {test_dir} does not exist.")
if not test_dirs:
print(f"Warning: None of {candidate_dirs} exist.")
return []
print(f"Checking directory: {test_dir}")
print(f"Checking directories: {test_dirs}")
for root, _, files in os.walk(test_dir):
for root, _, files in (entry for test_dir in test_dirs for entry in os.walk(test_dir)):
for file in files:
if file.endswith(".py"):
file_path = os.path.join(root, file)
@ -173,7 +174,7 @@ def main():
f"This can cause import issues and test failures. Files: {problematic_files}"
)
else:
print("✓ No litellm module deletion patterns found in test_litellm directory.")
print("✓ No litellm module deletion patterns found in tests/test_litellm or tests/unit.")
if __name__ == "__main__":

View file

@ -31,7 +31,7 @@ def get_all_functions_called_in_tests(base_dir):
specifically in files containing the word 'router'.
"""
called_functions = set()
test_dirs = ["local_testing", "router_unit_tests", "test_litellm"]
test_dirs = ["local_testing", "router_unit_tests", "test_litellm", "unit"]
for test_dir in test_dirs:
dir_path = os.path.join(base_dir, test_dir)

View file

@ -22,7 +22,7 @@ class TestBedrockGPTOSS(BaseLLMChatTest):
"""Bedrock GPT-OSS intermittently emits truncated toolUse.input deltas on
the live endpoint, which makes the inherited live integration test flaky.
The accumulation side is covered deterministically by
tests/test_litellm/llms/bedrock/chat/test_invoke_handler.py::test_transform_tool_calls_index;
tests/unit/llms/bedrock/chat/test_invoke_handler.py::test_transform_tool_calls_index;
the GPT-OSS-specific request-body transformation is covered by
test_function_calling_request_body_gpt_oss below.
"""

View file

@ -277,4 +277,4 @@ class BaseSkillsAPITest(ABC):
#
# Transformation logic (URL construction, headers, request/response parsing) is
# covered by unit tests in:
# tests/test_litellm/test_anthropic_skills_transformation.py
# tests/unit/test_anthropic_skills_transformation.py

View file

@ -324,7 +324,7 @@ def test_parallel_function_call_anthropic_error_msg(model, messages):
Anthropic (and Bedrock Invoke via ``AnthropicConfig.transform_request``)
inject a dummy tool so CLIs work with ``modify_params`` left off. Bedrock
Converse's no-raise behavior is covered offline in
``tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py``
``tests/unit/llms/bedrock/chat/test_converse_transformation.py``
(see #24158, #27138), which needs no live credentials.
"""
# Force modify_params off as a clean baseline: it exercises the Anthropic

View file

@ -17,7 +17,7 @@ body can still arrive, released once the caller is done with the response.
Nothing here re-tests the shapes ``_handler_may_close_client`` covers -- a
borrowed ``handler.client``, a caller-supplied client, an evicted-but-held
client. Those are pinned in ``tests/test_litellm/llms/custom_httpx/
client. Those are pinned in ``tests/unit/llms/custom_httpx/
test_http_handler.py``. What is uncovered there is the in-flight response, so no
test here may keep the client in a local: that inflates the very refcount under
test, and the test then passes on a broken handler. They hold weak references

View file

@ -4,7 +4,7 @@ Integration tests for SageMaker Nova provider.
These tests require a live SageMaker Nova endpoint and AWS credentials.
They are skipped by default — run manually with:
pytest tests/test_litellm/llms/sagemaker/test_sagemaker_nova_integration.py -v --no-header -rN
pytest tests/local_testing/test_sagemaker_nova_integration.py -v --no-header -rN
Prerequisites:
export AWS_PROFILE=<your-profile> # or set AWS_ACCESS_KEY_ID / AWS_SECRET_ACCESS_KEY
@ -251,7 +251,7 @@ class TestSagemakerNova2LiteIntegration:
Run with:
export SAGEMAKER_NOVA2_LITE_ENDPOINT=<your-nova-2-lite-endpoint>
pytest tests/test_litellm/llms/sagemaker/test_sagemaker_nova_integration.py::TestSagemakerNova2LiteIntegration -v
pytest tests/local_testing/test_sagemaker_nova_integration.py::TestSagemakerNova2LiteIntegration -v
"""
def test_should_accept_reasoning_effort_low(self):

View file

@ -85,7 +85,7 @@ class TestBingGroundingSearch(BaseSearchTest):
class TestBingGroundingSearchTransformation:
"""
Full-stack tests through `litellm.search` / `litellm.asearch` with the HTTP layer mocked.
Transformation details are unit-tested in tests/test_litellm/llms/azure/search/.
Transformation details are unit-tested in tests/unit/llms/azure/search/.
"""
@pytest.fixture(autouse=True)

View file

@ -58,7 +58,7 @@ class TestNimbleSearch(BaseSearchTest):
class TestNimbleSearchTransformation:
"""
Full-stack tests through `litellm.search` / `litellm.asearch` with the HTTP layer mocked.
Transformation details are unit-tested in tests/test_litellm/llms/nimble/search/.
Transformation details are unit-tested in tests/unit/llms/nimble/search/.
"""
@pytest.fixture(autouse=True)

View file

@ -1,387 +0,0 @@
import json
import pytest
import litellm
import litellm.batches.batch_utils as bu
from litellm.types.llms.openai import Batch
GROUNDED_USAGE_METADATA = {
"promptTokenCount": 19,
"candidatesTokenCount": 59,
"thoughtsTokenCount": 406,
"toolUsePromptTokenCount": 73,
"totalTokenCount": 557,
"promptTokensDetails": [{"modality": "TEXT", "tokenCount": 19}],
"candidatesTokensDetails": [{"modality": "TEXT", "tokenCount": 59}],
"toolUsePromptTokensDetails": [{"modality": "TEXT", "tokenCount": 73}],
"trafficType": "ON_DEMAND",
}
PASSTHROUGH_OUTPUT_URI = (
"gs://litellm-bucket/litellm-vertex-files/passthrough/publishers/google/models/gemini-2.5-flash/u/"
"predictions.jsonl"
)
UNGROUNDED_USAGE_METADATA = {
"promptTokenCount": 20,
"candidatesTokenCount": 48,
"thoughtsTokenCount": 195,
"toolUsePromptTokenCount": 73,
"totalTokenCount": 336,
"promptTokensDetails": [{"modality": "TEXT", "tokenCount": 20}],
"trafficType": "ON_DEMAND",
}
def _batch(output_file_id: str) -> Batch:
return Batch(
id="b",
completion_window="24h",
created_at=1,
endpoint="/v1/chat/completions",
input_file_id="f",
object="batch",
status="completed",
output_file_id=output_file_id,
)
def _vertex_jsonl(rows: list[dict]) -> bytes:
return "\n".join(json.dumps(row) for row in rows).encode()
def _vertex_openai_row(custom_id: str, model: str, prompt_tokens: int, completion_tokens: int) -> dict:
return {
"id": f"batch_req_{custom_id}",
"custom_id": custom_id,
"response": {
"status_code": 200,
"request_id": custom_id,
"body": {
"id": f"chatcmpl-{custom_id}",
"object": "chat.completion",
"model": model,
"choices": [{"index": 0, "message": {"role": "assistant", "content": "ok"}, "finish_reason": "stop"}],
"usage": {
"prompt_tokens": prompt_tokens,
"completion_tokens": completion_tokens,
"total_tokens": prompt_tokens + completion_tokens,
},
},
},
"error": None,
}
def _native_vertex_row(usage_metadata: dict, *, grounded: bool, model_version: str | None = "gemini-2.5-flash"):
candidate = {"content": {"role": "model", "parts": [{"text": "ok"}]}, "finishReason": "STOP"}
grounding = {"groundingMetadata": {"webSearchQueries": ["q"]}} if grounded else {}
response = {"candidates": [{**candidate, **grounding}], "usageMetadata": usage_metadata}
return {
"request": {"contents": [{"role": "user", "parts": [{"text": "q"}]}], "tools": [{"googleSearch": {}}]},
"status": "",
"response": {**response, **({"modelVersion": model_version} if model_version else {})},
"processed_time": "2026-09-23T19:02:00.000+00:00",
}
def _capture_cost_calls(monkeypatch, prompt_cost=0.5, completion_cost=0.25) -> list:
import litellm.cost_calculator as cc
calls: list = []
def _calc(**kw):
calls.append(kw)
return (prompt_cost, completion_cost)
monkeypatch.setattr(cc, "batch_cost_calculator", _calc)
return calls
def test_vertex_native_cost_bills_embedding_rows(monkeypatch):
monkeypatch.setitem(litellm.model_cost, "vertex_ai/gemini-embedding-2", {"input_cost_per_token_batches": 1e-7})
rows = [
{
"key": "id_1",
"status": "",
"request": {"content": {"parts": [{"text": "hello world"}]}},
"response": {"embedding": {"values": [0.1, 0.2]}, "usageMetadata": {"promptTokenCount": 2}},
},
{
"key": "id_2",
"status": "",
"request": {"content": {"parts": [{"text": "hello"}]}},
"response": {"embedding": {"values": [0.3]}, "tokenCount": "3"},
},
{"key": "id_3", "status": "INVALID_ARGUMENT", "request": {"content": {"parts": [{"text": ""}]}}},
]
result = bu.calculate_vertex_ai_batch_cost_and_usage(rows, "gemini-embedding-2")
assert (result.successful_requests, result.failed_requests) == (2, 1)
assert (result.usage.prompt_tokens, result.usage.completion_tokens, result.usage.total_tokens) == (5, 0, 5)
assert result.cost == pytest.approx(5 * 1e-7)
assert result.models == ["gemini-embedding-2"]
@pytest.mark.asyncio
async def test_native_vertex_rows_route_to_vertex_cost_path_without_flag(monkeypatch):
monkeypatch.setattr(litellm, "disable_vertex_batch_output_transformation", False, raising=False)
monkeypatch.setattr(
bu, "_aggregate_batch_cost_usage_models", lambda **kw: pytest.fail("generic path should not run")
)
calls = _capture_cost_calls(monkeypatch)
rows = [
_native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True),
_native_vertex_row(UNGROUNDED_USAGE_METADATA, grounded=False),
]
result = await bu.calculate_batch_cost_and_usage(
file_content_dictionary=rows, custom_llm_provider="vertex_ai", model_name="gemini-2.5-flash"
)
assert result.cost == pytest.approx(1.5)
assert (result.successful_requests, result.failed_requests) == (2, 0)
assert result.models == ["gemini-2.5-flash"]
assert {(call["model"], call["custom_llm_provider"]) for call in calls} == {("gemini-2.5-flash", "vertex_ai")}
@pytest.mark.asyncio
async def test_openai_shaped_vertex_rows_keep_the_generic_path_without_flag(monkeypatch):
monkeypatch.setattr(litellm, "disable_vertex_batch_output_transformation", False, raising=False)
monkeypatch.setattr(
bu, "calculate_vertex_ai_batch_cost_and_usage", lambda *a, **kw: pytest.fail("native path should not run")
)
_capture_cost_calls(monkeypatch)
rows = [_vertex_openai_row("request-1", "gemini-2.5-flash", 10, 5)]
result = await bu.calculate_batch_cost_and_usage(
file_content_dictionary=rows, custom_llm_provider="vertex_ai", model_name="gemini-2.5-flash"
)
assert result.successful_requests == 1
@pytest.mark.asyncio
async def test_native_vertex_rows_on_another_provider_keep_the_generic_path(monkeypatch):
monkeypatch.setattr(
bu, "calculate_vertex_ai_batch_cost_and_usage", lambda *a, **kw: pytest.fail("native path should not run")
)
_capture_cost_calls(monkeypatch)
result = await bu.calculate_batch_cost_and_usage(
file_content_dictionary=[_native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True)],
custom_llm_provider="openai",
)
assert result.successful_requests == 0
@pytest.mark.asyncio
async def test_handle_completed_batch_routes_native_rows_without_flag(monkeypatch):
monkeypatch.setattr(litellm, "disable_vertex_batch_output_transformation", False, raising=False)
raw_rows = [_native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True)]
async def fake_fetch(batch, custom_llm_provider, litellm_params=None):
return _vertex_jsonl(raw_rows)
monkeypatch.setattr(bu, "_fetch_batch_output_file_content", fake_fetch)
monkeypatch.setattr(
bu, "_aggregate_batch_cost_usage_models", lambda **kw: pytest.fail("generic path should not run")
)
calls = _capture_cost_calls(monkeypatch, prompt_cost=0.7, completion_cost=0.3)
deployment_model_info = {"input_cost_per_token_batches": 1e-6, "output_cost_per_token_batches": 2e-6}
result = await bu._handle_completed_batch(
_batch(PASSTHROUGH_OUTPUT_URI),
custom_llm_provider="vertex_ai",
model_name="gemini-2.5-flash",
model_info=deployment_model_info,
)
assert result.cost == pytest.approx(1.0)
assert result.usage.total_tokens == 557
assert [call["model_info"] for call in calls] == [deployment_model_info]
def test_native_vertex_usage_is_billed_like_the_online_path(monkeypatch):
calls = _capture_cost_calls(monkeypatch)
grounded = _native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True)
ungrounded = _native_vertex_row(UNGROUNDED_USAGE_METADATA, grounded=False)
result = bu.calculate_vertex_ai_batch_cost_and_usage([grounded, ungrounded], "gemini-2.5-flash")
grounded_usage, ungrounded_usage = (call["usage"] for call in calls)
assert grounded_usage.prompt_tokens == 19
assert grounded_usage.completion_tokens == 59 + 406
assert grounded_usage.completion_tokens_details.reasoning_tokens == 406
assert ungrounded_usage.prompt_tokens == 20 + 73
assert ungrounded_usage.completion_tokens == 48 + 195
assert (result.usage.prompt_tokens, result.usage.completion_tokens, result.usage.total_tokens) == (
19 + 93,
465 + 243,
557 + 336,
)
def test_native_vertex_rows_are_priced_by_model_version_without_a_model_name(monkeypatch):
calls = _capture_cost_calls(monkeypatch)
rows = [
_native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True, model_version="gemini-2.5-flash"),
_native_vertex_row(UNGROUNDED_USAGE_METADATA, grounded=False, model_version="gemini-2.5-pro"),
_native_vertex_row(UNGROUNDED_USAGE_METADATA, grounded=False, model_version=None),
]
result = bu.calculate_vertex_ai_batch_cost_and_usage(rows)
assert [call["model"] for call in calls] == ["gemini-2.5-flash", "gemini-2.5-pro"]
assert result.models == ["gemini-2.5-flash", "gemini-2.5-pro"]
assert result.cost == pytest.approx(1.5)
assert result.successful_requests == 3
assert result.usage.total_tokens == 557 + 336 + 336
def test_native_vertex_rows_without_usage_metadata_count_as_failed(monkeypatch):
_capture_cost_calls(monkeypatch)
rows = [
{"request": {"contents": []}, "status": "Error: bad request", "processed_time": "t"},
{"request": {"contents": []}, "response": {"candidates": []}},
_native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True),
]
result = bu.calculate_vertex_ai_batch_cost_and_usage(rows, "gemini-2.5-flash")
assert (result.successful_requests, result.failed_requests) == (1, 2)
assert result.usage.total_tokens == 557
def test_native_vertex_batch_whose_rows_all_failed_still_names_the_deployment_model(monkeypatch):
calls = _capture_cost_calls(monkeypatch)
rows = [{"request": {"contents": []}, "status": "Error: quota exceeded", "processed_time": "t"}] * 2
result = bu.calculate_vertex_ai_batch_cost_and_usage(rows, "gemini-2.5-flash")
assert result.models == ["gemini-2.5-flash"]
assert (result.successful_requests, result.failed_requests, result.cost) == (0, 2, 0.0)
assert calls == []
def test_native_vertex_rows_are_priced_with_the_deployment_model_info(monkeypatch):
calls = _capture_cost_calls(monkeypatch)
deployment_model_info = {"input_cost_per_token_batches": 1e-6, "output_cost_per_token_batches": 2e-6}
bu.calculate_vertex_ai_batch_cost_and_usage(
[_native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True)],
"gemini-2.5-flash",
model_info=deployment_model_info,
)
assert [call["model_info"] for call in calls] == [deployment_model_info]
@pytest.mark.asyncio
async def test_native_vertex_rows_keep_the_deployment_model_info_through_the_batch_entrypoint(monkeypatch):
calls = _capture_cost_calls(monkeypatch)
deployment_model_info = {"input_cost_per_token_batches": 1e-6}
await bu.calculate_batch_cost_and_usage(
file_content_dictionary=[_native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True)],
custom_llm_provider="vertex_ai",
model_name="gemini-2.5-flash",
model_info=deployment_model_info,
)
assert [call["model_info"] for call in calls] == [deployment_model_info]
def test_native_vertex_rows_are_priced_by_the_deployment_model_over_model_version(monkeypatch):
calls = _capture_cost_calls(monkeypatch)
rows = [_native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True, model_version="gemini-2.5-pro")]
result = bu.calculate_vertex_ai_batch_cost_and_usage(rows, "gemini-2.5-flash")
assert [call["model"] for call in calls] == ["gemini-2.5-flash"]
assert result.models == ["gemini-2.5-flash"]
def test_native_vertex_rows_that_fail_response_validation_count_as_failed(monkeypatch):
calls = _capture_cost_calls(monkeypatch)
rows = [
{"request": {"contents": []}, "response": {"candidates": "nope", "usageMetadata": GROUNDED_USAGE_METADATA}},
_native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True),
]
result = bu.calculate_vertex_ai_batch_cost_and_usage(rows, "gemini-2.5-flash")
assert (result.successful_requests, result.failed_requests) == (1, 1)
assert result.usage.total_tokens == 557
assert len(calls) == 1
@pytest.mark.parametrize("wildcard_model", ["*", "vertex_ai/*"])
def test_native_vertex_rows_under_a_wildcard_deployment_are_priced_by_model_version(monkeypatch, wildcard_model):
calls = _capture_cost_calls(monkeypatch)
rows = [
_native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True, model_version="gemini-2.5-flash"),
_native_vertex_row(UNGROUNDED_USAGE_METADATA, grounded=False, model_version=None),
]
result = bu.calculate_vertex_ai_batch_cost_and_usage(rows, wildcard_model)
assert [call["model"] for call in calls] == ["gemini-2.5-flash", wildcard_model]
assert result.cost == pytest.approx(1.5)
assert (result.successful_requests, result.failed_requests) == (2, 0)
assert result.usage.total_tokens == 557 + 336
def test_native_vertex_row_without_model_version_under_a_wildcard_deployment_bills_its_explicit_prices():
deployment_model_info = {"input_cost_per_token_batches": 1e-6, "output_cost_per_token_batches": 2e-6}
with_version = _native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True, model_version="gemini-2.5-flash")
without_version = _native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True, model_version=None)
twin = bu.calculate_vertex_ai_batch_cost_and_usage([with_version], "vertex_ai/*", model_info=deployment_model_info)
both = bu.calculate_vertex_ai_batch_cost_and_usage(
[with_version, without_version], "vertex_ai/*", model_info=deployment_model_info
)
assert twin.cost > 0
assert both.cost == pytest.approx(2 * twin.cost)
assert (both.successful_requests, both.failed_requests) == (2, 0)
def test_native_vertex_row_the_cost_map_cannot_price_is_billed_at_zero_and_the_rest_still_bills(monkeypatch):
import litellm.cost_calculator as cc
def _calc(**kw):
if kw["model"] == "gemini-unpriced":
raise ValueError("no pricing")
return (0.5, 0.25)
monkeypatch.setattr(cc, "batch_cost_calculator", _calc)
rows = [
_native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True, model_version="gemini-unpriced"),
_native_vertex_row(UNGROUNDED_USAGE_METADATA, grounded=False, model_version="gemini-2.5-flash"),
]
result = bu.calculate_vertex_ai_batch_cost_and_usage(rows)
assert result.cost == pytest.approx(0.75)
assert (result.successful_requests, result.failed_requests) == (2, 0)
assert result.usage.total_tokens == 557 + 336
assert result.models == ["gemini-unpriced", "gemini-2.5-flash"]
@pytest.mark.asyncio
async def test_flag_sends_every_vertex_row_down_the_native_path_when_a_model_is_known(monkeypatch):
monkeypatch.setattr(litellm, "disable_vertex_batch_output_transformation", True, raising=False)
monkeypatch.setattr(
bu, "_aggregate_batch_cost_usage_models", lambda **kw: pytest.fail("generic path should not run")
)
calls = _capture_cost_calls(monkeypatch)
rows = [_vertex_openai_row("request-1", "gemini-2.5-flash", 10, 5)]
result = await bu.calculate_batch_cost_and_usage(
file_content_dictionary=rows, custom_llm_provider="vertex_ai", model_name="gemini-2.5-flash"
)
assert calls == []
assert (result.successful_requests, result.failed_requests) == (0, 1)

View file

@ -1,117 +0,0 @@
from __future__ import annotations
from collections.abc import Mapping
from typing import Final
import pytest
import litellm
from litellm.chat_completions import dispatch
from litellm.rust_bridge.bindings import NativeBinding
from litellm.rust_bridge.catalog import Route, RouteRule, Rules
from litellm.rust_bridge.chat_completions.entrypoints import (
LiteLLMChatCompletionsRequest,
NativeAcompletion,
NativeCompletion,
)
from litellm.rust_bridge.configuration import Rollout
from litellm.types.utils import ModelResponse
MESSAGES: Final = [{"role": "user", "content": "hi"}]
@pytest.mark.asyncio
async def test_public_completion_calls_keep_the_python_result() -> None:
sync_response: Final = litellm.completion(model="openai/test-model", messages=MESSAGES, mock_response="ok")
async_response: Final = await litellm.acompletion(model="openai/test-model", messages=MESSAGES, mock_response="ok")
assert isinstance(sync_response, ModelResponse)
assert isinstance(async_response, ModelResponse)
assert sync_response.choices[0].message.content == "ok"
assert async_response.choices[0].message.content == "ok"
def test_sync_completion_request_projects_public_arguments() -> None:
rules: Final[Rules] = (RouteRule(Route.CHAT_COMPLETIONS, Rollout.RUST_REQUIRED),)
expected: Final = ModelResponse()
def native(
request: LiteLLMChatCompletionsRequest, args: tuple[object, ...], kwargs: Mapping[str, object]
) -> ModelResponse:
assert request.model == "test-model"
assert request.messages == MESSAGES
assert request.custom_llm_provider == "openai"
assert request.stream is True
return expected
binding: Final[NativeBinding[NativeCompletion]] = NativeBinding("completion", validate=lambda _: None)
binding.override(native)
response: Final = dispatch._DISPATCH.run( # pyright: ignore[reportPrivateUsage] # test an explicit route decision
("test-model", MESSAGES),
{"custom_llm_provider": "openai", "stream": True},
python=lambda *args, **kwargs: pytest.fail("required native route must handle this call"),
binding=binding,
native=lambda hook, request, args, kwargs: hook(request, args, kwargs),
rules=rules,
)
assert response is expected
@pytest.mark.asyncio
async def test_async_completion_falls_back_after_native_declines() -> None:
from litellm.rust_bridge.bindings import native_exception_types
native_types: Final = native_exception_types()
if native_types is None:
pytest.skip("native bridge is unavailable")
declined, _ = native_types
expected: Final = ModelResponse()
rules: Final[Rules] = (RouteRule(Route.CHAT_COMPLETIONS, Rollout.RUST_OPT_OUT),)
async def native(
request: LiteLLMChatCompletionsRequest, args: tuple[object, ...], kwargs: Mapping[str, object]
) -> ModelResponse:
raise declined("unsupported")
async def python(*args: object, **kwargs: object) -> ModelResponse:
return expected
binding: Final[NativeBinding[NativeAcompletion]] = NativeBinding("acompletion", validate=lambda _: None)
binding.override(native)
response: Final = await dispatch._ADISPATCH.arun( # pyright: ignore[reportPrivateUsage] # test an explicit route decision
("test-model", MESSAGES),
{},
python=python,
binding=binding,
native=lambda hook, request, args, kwargs: hook(request, args, kwargs),
rules=rules,
)
assert response is expected
def test_internal_acompletion_marker_bypasses_native() -> None:
rules: Final[Rules] = (RouteRule(Route.CHAT_COMPLETIONS, Rollout.RUST_REQUIRED),)
expected: Final = ModelResponse()
def python(*args: object, **kwargs: object) -> ModelResponse:
return expected
def native(
request: LiteLLMChatCompletionsRequest, args: tuple[object, ...], kwargs: Mapping[str, object]
) -> ModelResponse:
pytest.fail("acompletion's inner completion call must stay on Python")
binding: Final[NativeBinding[NativeCompletion]] = NativeBinding("completion", validate=lambda _: None)
binding.override(native)
response: Final = dispatch._DISPATCH.run( # pyright: ignore[reportPrivateUsage] # test an explicit route decision
("test-model", MESSAGES),
{"custom_llm_provider": "openai", "acompletion": True},
python=python,
binding=binding,
native=lambda hook, request, args, kwargs: hook(request, args, kwargs),
rules=rules,
)
assert response is expected

View file

@ -14,6 +14,7 @@ from pathlib import Path
from types import SimpleNamespace
import httpx
import pytest
from pytest_socket import _remove_restrictions
import asyncio
@ -509,6 +510,14 @@ def setup_and_teardown():
print(f"[conftest] Module teardown complete (worker: {worker_id or 'master'})")
def pytest_collectstart():
_remove_restrictions()
def pytest_runtest_setup():
_remove_restrictions()
def pytest_collection_modifyitems(config, items):
"""
Customize test collection order.

View file

@ -13,7 +13,7 @@ def _claude_mapping(messages, response_obj):
def test_claude_mapping_serializes_custom_tool_calls(monkeypatch):
"""
Stub the anthropic module unconditionally: the SDK may be absent (it lives in the
proxy-runtime extra), and the tests/test_litellm/llms/anthropic test package can
proxy-runtime extra), and the tests/unit/llms/anthropic test package can
shadow it on sys.path, so an import probe proves nothing about the real SDK.
"""
stub = types.ModuleType("anthropic")

View file

@ -7,10 +7,6 @@ the litellm_responses bridge provider, which calls litellm.responses() internall
import os
from litellm.interactions.litellm_responses_transformation.transformation import (
LiteLLMResponsesInteractionsConfig,
)
from litellm.types.interactions import Turn
from tests.test_litellm.interactions.base_interactions_test import (
BaseInteractionsTest,
)
@ -30,71 +26,3 @@ class TestLiteLLMResponsesBridge(BaseInteractionsTest):
def get_api_key(self) -> str:
"""Return the OpenAI API key from environment."""
return os.getenv("OPENAI_API_KEY", "")
class TestBridgeInputTransformation:
"""Regression tests for translating Interactions input into Responses API input.
The bridge used to pass Google content parts through raw ({"type": "text"}),
which the Responses API rejects with a 400, and it dropped the role encoded
in step types and in the legacy "model" turn role.
"""
def test_step_input_maps_roles_and_content_types(self):
transformed = LiteLLMResponsesInteractionsConfig._transform_interactions_input_to_responses_input(
[
{"type": "user_input", "content": [{"type": "text", "text": "I like apples."}]},
{"type": "model_output", "content": [{"type": "text", "text": "I like oranges."}]},
{"type": "user_input", "content": [{"type": "text", "text": "What did you say?"}]},
]
)
assert transformed == [
{"role": "user", "content": [{"type": "input_text", "text": "I like apples."}]},
{"role": "assistant", "content": [{"type": "output_text", "text": "I like oranges."}]},
{"role": "user", "content": [{"type": "input_text", "text": "What did you say?"}]},
]
def test_legacy_turn_input_maps_model_role_to_assistant(self):
transformed = LiteLLMResponsesInteractionsConfig._transform_interactions_input_to_responses_input(
[
{"role": "user", "content": [{"type": "text", "text": "I like apples."}]},
{"role": "model", "content": [{"type": "text", "text": "I like oranges."}]},
]
)
assert transformed == [
{"role": "user", "content": [{"type": "input_text", "text": "I like apples."}]},
{"role": "assistant", "content": [{"type": "output_text", "text": "I like oranges."}]},
]
def test_turn_pydantic_model_with_string_content(self):
transformed = LiteLLMResponsesInteractionsConfig._transform_interactions_input_to_responses_input(
[Turn(role="model", content="I like oranges.")]
)
assert transformed == [
{"role": "assistant", "content": [{"type": "output_text", "text": "I like oranges."}]}
]
def test_string_input_passes_through(self):
transformed = LiteLLMResponsesInteractionsConfig._transform_interactions_input_to_responses_input("Hello")
assert transformed == "Hello"
def test_content_list_input_becomes_single_user_message(self):
transformed = LiteLLMResponsesInteractionsConfig._transform_interactions_input_to_responses_input(
[{"type": "text", "text": "Hello"}, "world"]
)
assert transformed == [
{
"role": "user",
"content": [
{"type": "input_text", "text": "Hello"},
{"type": "input_text", "text": "world"},
],
}
]
def test_non_text_content_passes_through_unchanged(self):
image_part = {"type": "image", "data": "base64data", "mime_type": "image/png"}
transformed = LiteLLMResponsesInteractionsConfig._transform_interactions_input_to_responses_input(
[{"type": "user_input", "content": [image_part]}]
)
assert transformed == [{"role": "user", "content": [image_part]}]

View file

@ -9,171 +9,6 @@ import os
import pytest
from litellm.llms.cometapi.chat.transformation import (
CometAPIChatCompletionStreamingHandler,
CometAPIConfig,
)
from litellm.llms.cometapi.common_utils import CometAPIException
class TestCometAPIChatCompletionStreamingHandler:
def test_chunk_parser_successful(self):
handler = CometAPIChatCompletionStreamingHandler(
streaming_response=None, sync_stream=True
)
# Test input chunk
chunk = {
"id": "test_id",
"created": 1234567890,
"model": "gpt-3.5-turbo",
"usage": {"prompt_tokens": 10, "completion_tokens": 20, "total_tokens": 30},
"choices": [
{"delta": {"content": "test content", "reasoning": "test reasoning"}}
],
}
# Parse chunk
result = handler.chunk_parser(chunk)
# Verify response
assert result.id == "test_id"
assert result.object == "chat.completion.chunk"
assert result.created == 1234567890
assert result.model == "gpt-3.5-turbo"
assert result.usage.prompt_tokens == chunk["usage"]["prompt_tokens"]
assert result.usage.completion_tokens == chunk["usage"]["completion_tokens"]
assert result.usage.total_tokens == chunk["usage"]["total_tokens"]
assert len(result.choices) == 1
assert result.choices[0]["delta"]["reasoning_content"] == "test reasoning"
def test_chunk_parser_error_response(self):
handler = CometAPIChatCompletionStreamingHandler(
streaming_response=None, sync_stream=True
)
# Test error chunk
error_chunk = {
"error": {
"message": "test error",
"code": 400,
}
}
# Verify error handling
with pytest.raises(CometAPIException) as exc_info:
handler.chunk_parser(error_chunk)
assert "CometAPI Error: test error" in str(exc_info.value)
assert exc_info.value.status_code == 400
def test_chunk_parser_key_error(self):
handler = CometAPIChatCompletionStreamingHandler(
streaming_response=None, sync_stream=True
)
# Test invalid chunk missing required fields
invalid_chunk = {"incomplete": "data"}
# Verify KeyError handling
with pytest.raises(CometAPIException) as exc_info:
handler.chunk_parser(invalid_chunk)
assert "KeyError" in str(exc_info.value)
assert exc_info.value.status_code == 400
class TestCometAPIConfig:
def test_transform_request_basic(self):
"""Test basic request transformation"""
config = CometAPIConfig()
transformed_request = config.transform_request(
model="cometapi/gpt-3.5-turbo",
messages=[{"role": "user", "content": "Hello, world!"}],
optional_params={},
litellm_params={},
headers={},
)
assert transformed_request["model"] == "cometapi/gpt-3.5-turbo"
assert transformed_request["messages"] == [
{"role": "user", "content": "Hello, world!"}
]
def test_transform_request_with_extra_body(self):
"""Test request transformation with extra_body parameters"""
config = CometAPIConfig()
transformed_request = config.transform_request(
model="cometapi/gpt-4",
messages=[{"role": "user", "content": "Hello, world!"}],
optional_params={"extra_body": {"custom_param": "custom_value"}},
litellm_params={},
headers={},
)
# Validate that extra_body parameters are merged into the request
assert transformed_request["custom_param"] == "custom_value"
assert transformed_request["messages"] == [
{"role": "user", "content": "Hello, world!"}
]
def test_cache_control_flag_removal(self):
"""Test cache control flag removal from messages"""
config = CometAPIConfig()
transformed_request = config.transform_request(
model="cometapi/gpt-3.5-turbo",
messages=[
{
"role": "user",
"content": "Hello, world!",
"cache_control": {"type": "ephemeral"},
}
],
optional_params={},
litellm_params={},
headers={},
)
# CometAPI should remove cache_control flags by default
assert transformed_request["messages"][0].get("cache_control") is None
def test_map_openai_params(self):
"""Test OpenAI parameter mapping"""
config = CometAPIConfig()
non_default_params = {
"temperature": 0.7,
"max_tokens": 100,
"top_p": 0.9,
}
mapped_params = config.map_openai_params(
non_default_params=non_default_params,
optional_params={},
model="cometapi/gpt-3.5-turbo",
drop_params=False,
)
assert mapped_params["temperature"] == 0.7
assert mapped_params["max_tokens"] == 100
assert mapped_params["top_p"] == 0.9
def test_get_error_class(self):
"""Test error class creation"""
config = CometAPIConfig()
error = config.get_error_class(
error_message="Test error",
status_code=400,
headers={"Content-Type": "application/json"},
)
assert isinstance(error, CometAPIException)
assert error.message == "Test error"
assert error.status_code == 400
# Integration test example (requires real API key)

View file

@ -1,79 +0,0 @@
import json
from typing import Final
import httpx
import respx
import litellm
def test_completion_merges_leading_system_and_developer_messages_for_chat_template_models(
respx_mock: respx.MockRouter,
):
upstream: Final = respx_mock.post("https://example.databricks.test/serving-endpoints/chat/completions").mock(
return_value=httpx.Response(
status_code=200,
json={
"id": "chatcmpl-123",
"object": "chat.completion",
"created": 1677652288,
"model": "my-custom-model",
"choices": [{"index": 0, "message": {"role": "assistant", "content": "Answer"}, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 9, "completion_tokens": 1, "total_tokens": 10},
},
)
)
response: Final = litellm.completion(
model="databricks/my-custom-model",
messages=[
{"role": "system", "content": "You are terse."},
{"role": "developer", "content": "Skills: none."},
{"role": "user", "content": "Hello"},
],
api_base="https://example.databricks.test/serving-endpoints",
api_key="fake-databricks-api-key",
num_retries=0,
)
assert upstream.call_count == 1
request_body: Final = json.loads(upstream.calls[0].request.read())
assert request_body["messages"] == [
{"role": "system", "content": "You are terse.\n\nSkills: none."},
{"role": "user", "content": "Hello"},
]
assert response.choices[0].message.content == "Answer"
def test_completion_merges_system_messages_when_one_has_empty_content(respx_mock: respx.MockRouter):
upstream: Final = respx_mock.post("https://example.databricks.test/serving-endpoints/chat/completions").mock(
return_value=httpx.Response(
status_code=200,
json={
"id": "chatcmpl-123",
"object": "chat.completion",
"created": 1677652288,
"model": "my-custom-model",
"choices": [{"index": 0, "message": {"role": "assistant", "content": "Answer"}, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 9, "completion_tokens": 1, "total_tokens": 10},
},
)
)
litellm.completion(
model="databricks/my-custom-model",
messages=[
{"role": "system", "content": "You are terse."},
{"role": "system", "content": ""},
{"role": "user", "content": "Hello"},
],
api_base="https://example.databricks.test/serving-endpoints",
api_key="fake-databricks-api-key",
num_retries=0,
)
request_body: Final = json.loads(upstream.calls[0].request.read())
assert request_body["messages"] == [
{"role": "system", "content": "You are terse."},
{"role": "user", "content": "Hello"},
]

View file

@ -1,433 +0,0 @@
"""
Integration tests for DeepInfra rerank functionality.
Tests the full rerank flow following the repository patterns.
"""
import asyncio
import json
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
import litellm
def assert_response_shape(response, custom_llm_provider):
"""Helper function to validate response structure specific to DeepInfra."""
assert hasattr(response, "id")
assert hasattr(response, "results")
assert hasattr(response, "meta")
assert isinstance(response.results, list)
for result in response.results:
assert "index" in result
assert "relevance_score" in result
assert isinstance(result["index"], int)
assert isinstance(result["relevance_score"], (int, float))
# Check meta structure
assert "tokens" in response.meta
assert "billed_units" in response.meta
assert "input_tokens" in response.meta["tokens"]
assert "total_tokens" in response.meta["billed_units"]
@pytest.mark.parametrize("sync_mode", [True, False])
@patch("litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post")
@patch("litellm.llms.custom_httpx.http_handler.HTTPHandler.post")
def test_basic_rerank_deepinfra(mock_sync_post, mock_async_post, sync_mode):
"""Test basic DeepInfra rerank functionality."""
# Mock response data that matches DeepInfra API format
mock_response_data = {
"scores": [0.9, 0.1],
"input_tokens": 25,
"request_id": "deepinfra-request-123",
"inference_status": {
"status": "success",
"runtime_ms": 150,
"cost": 0.0001,
"tokens_generated": 0,
"tokens_input": 25,
},
}
def return_val():
return mock_response_data
api_key = "test_deepinfra_api_key"
api_base = "https://api.deepinfra.com"
if sync_mode:
# Create mock response object for sync
mock_response = MagicMock()
mock_response.json = return_val
mock_response.status_code = 200
mock_response.headers = {"content-type": "application/json"}
mock_response.text = json.dumps(mock_response_data)
mock_sync_post.return_value = mock_response
response = litellm.rerank(
model="deepinfra/Qwen/Qwen3-Reranker-0.6B",
query="hello",
documents=["hello", "world"],
top_n=2,
custom_llm_provider="deepinfra",
api_key=api_key,
api_base=api_base,
)
mock_sync_post.assert_called_once()
else:
# Create mock response object for async
mock_response = AsyncMock()
def return_val():
return mock_response_data
mock_response.json = return_val
mock_response.status_code = 200
mock_response.headers = {"content-type": "application/json"}
mock_response.text = json.dumps(mock_response_data)
mock_async_post.return_value = mock_response
response = asyncio.run(
litellm.arerank(
model="deepinfra/Qwen/Qwen3-Reranker-0.6B",
query="hello",
documents=["hello", "world"],
top_n=2,
custom_llm_provider="deepinfra",
api_key=api_key,
api_base=api_base,
)
)
mock_async_post.assert_called_once()
# Verify response structure
assert response.id == "deepinfra-request-123"
assert response.results is not None
assert len(response.results) == 2
assert response.results[0]["index"] == 0
assert response.results[0]["relevance_score"] == 0.9
assert response.results[1]["index"] == 1
assert response.results[1]["relevance_score"] == 0.1
# Verify metadata
assert response.meta["tokens"]["input_tokens"] == 25
assert response.meta["billed_units"]["total_tokens"] == 25
# Verify hidden params specific to DeepInfra
assert response._hidden_params["status"] == "success"
assert response._hidden_params["runtime_ms"] == 150
assert response._hidden_params["cost"] == 0.0001
# Note: The model name is processed and the 'deepinfra/' prefix is removed
assert response._hidden_params["model"] == "Qwen/Qwen3-Reranker-0.6B"
assert_response_shape(response, custom_llm_provider="deepinfra")
@pytest.mark.parametrize("sync_mode", [True, False])
@patch("litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post")
@patch("litellm.llms.custom_httpx.http_handler.HTTPHandler.post")
def test_deepinfra_rerank_with_queries_param(
mock_sync_post, mock_async_post, sync_mode
):
"""Test DeepInfra rerank with multiple queries parameter."""
mock_response_data = {
"scores": [0.8, 0.6, 0.2],
"input_tokens": 35,
"request_id": "deepinfra-multi-query-123",
"inference_status": {"status": "success", "runtime_ms": 200},
}
def return_val():
return mock_response_data
if sync_mode:
mock_response = MagicMock()
mock_response.json = return_val
mock_response.status_code = 200
mock_response.headers = {"content-type": "application/json"}
mock_response.text = json.dumps(mock_response_data)
mock_sync_post.return_value = mock_response
response = litellm.rerank(
model="deepinfra/Qwen/Qwen3-Reranker-4B",
query="hello",
documents=["hello", "world", "test"],
queries=["hello", "hi there"], # DeepInfra specific param
custom_llm_provider="deepinfra",
api_key="test_key",
api_base="https://api.deepinfra.com",
)
mock_sync_post.assert_called_once()
# Verify that queries parameter was passed in request
call_data = json.loads(mock_sync_post.call_args.kwargs["data"])
assert "queries" in call_data
assert call_data["queries"] == ["hello", "hi there"]
else:
mock_response = AsyncMock()
mock_response.json = return_val
mock_response.status_code = 200
mock_response.headers = {"content-type": "application/json"}
mock_response.text = json.dumps(mock_response_data)
mock_async_post.return_value = mock_response
response = asyncio.run(
litellm.arerank(
model="deepinfra/Qwen/Qwen3-Reranker-4B",
query="hello",
documents=["hello", "world", "test"],
queries=["hello", "hi there"],
custom_llm_provider="deepinfra",
api_key="test_key",
api_base="https://api.deepinfra.com",
)
)
mock_async_post.assert_called_once()
call_data = json.loads(mock_async_post.call_args.kwargs["data"])
assert "queries" in call_data
assert call_data["queries"] == ["hello", "hi there"]
assert response.results is not None
assert len(response.results) == 3
@patch("litellm.llms.custom_httpx.http_handler.HTTPHandler.post")
def test_deepinfra_rerank_with_service_tier(mock_post):
"""Test DeepInfra rerank with service_tier parameter."""
mock_response_data = {
"scores": [0.95, 0.75],
"input_tokens": 30,
"request_id": "deepinfra-premium-123",
}
def return_val():
return mock_response_data
mock_response = MagicMock()
mock_response.json = return_val
mock_response.status_code = 200
mock_response.headers = {"content-type": "application/json"}
mock_response.text = json.dumps(mock_response_data)
mock_post.return_value = mock_response
response = litellm.rerank(
model="deepinfra/Qwen/Qwen3-Reranker-8B",
query="premium search",
documents=["doc1", "doc2"],
service_tier="premium", # DeepInfra specific param
custom_llm_provider="deepinfra",
api_key="test_key",
api_base="https://api.deepinfra.com",
)
mock_post.assert_called_once()
# Verify URL
call_url = mock_post.call_args.kwargs["url"]
assert "api.deepinfra.com/inference/Qwen/Qwen3-Reranker-8B" in call_url
# Verify request contains service_tier
call_data = json.loads(mock_post.call_args.kwargs["data"])
assert call_data["service_tier"] == "premium"
assert response.results is not None
@patch("litellm.llms.custom_httpx.http_handler.HTTPHandler.post")
def test_deepinfra_rerank_with_env_vars(mock_post, monkeypatch):
"""Test DeepInfra rerank with environment variable configuration."""
monkeypatch.setenv("DEEPINFRA_API_KEY", "env_test_key")
monkeypatch.setenv("DEEPINFRA_API_BASE", "https://custom-deepinfra.com")
mock_response_data = {
"scores": [0.88, 0.22],
"input_tokens": 28,
"request_id": "env-test-123",
}
def return_val():
return mock_response_data
mock_response = MagicMock()
mock_response.json = return_val
mock_response.status_code = 200
mock_response.headers = {"content-type": "application/json"}
mock_response.text = json.dumps(mock_response_data)
mock_post.return_value = mock_response
response = litellm.rerank(
model="deepinfra/Qwen/Qwen3-Reranker-0.6B",
query="hello",
documents=["hello", "world"],
custom_llm_provider="deepinfra",
)
mock_post.assert_called_once()
# Verify headers contain env API key
headers = mock_post.call_args.kwargs.get("headers", {})
assert "Bearer env_test_key" in headers.get("Authorization", "")
assert response.results is not None
@patch("litellm.llms.custom_httpx.http_handler.HTTPHandler.post")
def test_deepinfra_rerank_error_handling(mock_post):
"""Test DeepInfra rerank error handling."""
error_response = {"detail": {"error": "Invalid API key"}}
def return_val():
return error_response
mock_response = MagicMock()
mock_response.status_code = 401
mock_response.json = return_val
mock_response.text = json.dumps(error_response)
mock_response.headers = {"content-type": "application/json"}
mock_post.return_value = mock_response
# The current implementation handles errors gracefully, so we expect a successful response
# with the error information in the hidden params
response = litellm.rerank(
model="deepinfra/Qwen/Qwen3-Reranker-0.6B",
query="hello",
documents=["hello", "world"],
custom_llm_provider="deepinfra",
api_key="invalid_key",
api_base="https://api.deepinfra.com",
)
# Verify that the response contains error information
assert (
response._hidden_params["status"] == "unknown"
) # Default status when error occurs
@patch("litellm.llms.custom_httpx.http_handler.HTTPHandler.post")
def test_deepinfra_rerank_defaults_api_base_when_missing(mock_post, monkeypatch):
"""With no api_base anywhere, the call still goes out against DeepInfra's own base."""
monkeypatch.delenv("DEEPINFRA_API_BASE", raising=False)
mock_response = MagicMock()
mock_response.json = lambda: {"scores": [0.9, 0.1], "input_tokens": 20}
mock_response.status_code = 200
mock_response.headers = {"content-type": "application/json"}
mock_post.return_value = mock_response
response = litellm.rerank(
model="deepinfra/Qwen/Qwen3-Reranker-0.6B",
query="hello",
documents=["hello", "world"],
custom_llm_provider="deepinfra",
api_key="test_key",
# api_base is intentionally missing
)
assert "api.deepinfra.com" in mock_post.call_args.kwargs["url"]
assert [result["relevance_score"] for result in response.results] == [0.9, 0.1]
@patch("litellm.llms.custom_httpx.http_handler.HTTPHandler.post")
def test_deepinfra_rerank_request_format(mock_post):
"""Test that the request is properly formatted for DeepInfra API."""
mock_response_data = {"scores": [0.9, 0.1], "input_tokens": 20}
def return_val():
return mock_response_data
mock_response = MagicMock()
mock_response.json = return_val
mock_response.status_code = 200
mock_response.headers = {"content-type": "application/json"}
mock_response.text = json.dumps(mock_response_data)
mock_post.return_value = mock_response
response = litellm.rerank(
model="deepinfra/Qwen/Qwen3-Reranker-0.6B",
query="test query",
documents=["doc1", "doc2"],
custom_llm_provider="deepinfra",
api_key="test_key",
api_base="https://api.deepinfra.com",
instruction="custom instruction",
webhook="https://webhook.example.com",
)
mock_post.assert_called_once()
# Verify URL format
call_url = mock_post.call_args.kwargs["url"]
assert call_url == "https://api.deepinfra.com/inference/Qwen/Qwen3-Reranker-0.6B"
# Verify headers
headers = mock_post.call_args.kwargs["headers"]
assert headers["Authorization"] == "Bearer test_key"
assert headers["accept"] == "application/json"
assert headers["content-type"] == "application/json"
# Verify request body format
request_data = json.loads(mock_post.call_args.kwargs["data"])
assert request_data["queries"] == [
"test query",
"test query",
] # DeepInfra requires queries to match documents length
assert request_data["documents"] == ["doc1", "doc2"]
assert request_data["instruction"] == "custom instruction"
assert request_data["webhook"] == "https://webhook.example.com"
assert response.results is not None
def test_deepinfra_rerank_models():
"""Test that DeepInfra Qwen rerank models are recognized."""
# These should not raise errors during model validation
models = [
"deepinfra/Qwen/Qwen3-Reranker-0.6B",
"deepinfra/Qwen/Qwen3-Reranker-4B",
"deepinfra/Qwen/Qwen3-Reranker-8B",
]
for model in models:
resolved_model, provider, _, api_base = litellm.get_llm_provider(model=model)
assert provider == "deepinfra"
assert resolved_model == model.removeprefix("deepinfra/")
assert api_base == "https://api.deepinfra.com/v1/openai"
@patch("litellm.llms.custom_httpx.http_handler.HTTPHandler.post")
def test_deepinfra_rerank_minimal_response(mock_post):
"""Test handling of minimal DeepInfra response."""
# Minimal response with just scores
mock_response_data = {"scores": [0.7, 0.3]}
def return_val():
return mock_response_data
mock_response = MagicMock()
mock_response.json = return_val
mock_response.status_code = 200
mock_response.headers = {"content-type": "application/json"}
mock_response.text = json.dumps(mock_response_data)
mock_post.return_value = mock_response
response = litellm.rerank(
model="deepinfra/Qwen/Qwen3-Reranker-0.6B",
query="hello",
documents=["hello", "world"],
custom_llm_provider="deepinfra",
api_key="test_key",
api_base="https://api.deepinfra.com",
)
# Should handle minimal response gracefully
assert response.results is not None
assert len(response.results) == 2
assert response.results[0]["relevance_score"] == 0.7
assert response.results[1]["relevance_score"] == 0.3
# Should have default values for missing fields
assert response.meta["tokens"]["input_tokens"] == 0 # Default when missing
assert response._hidden_params["status"] == "unknown" # Default when missing

View file

@ -1 +0,0 @@
"""Tests for Gemini files functionality"""

View file

@ -1 +0,0 @@
# Gemini Video Generation Tests

View file

@ -1 +0,0 @@
# Manus provider tests

View file

@ -1 +0,0 @@
# Manus Responses API tests

View file

@ -1 +0,0 @@
# MiniMax tests

View file

@ -1 +0,0 @@
# MiniMax chat tests

View file

@ -1 +0,0 @@
# MiniMax messages tests

View file

@ -1,19 +1,9 @@
import os
from typing import Dict
from unittest.mock import MagicMock
import httpx
import litellm
import pytest
from litellm.llms.base_llm.audio_transcription.transformation import (
BaseAudioTranscriptionConfig,
)
from litellm.llms.mistral.audio_transcription.transformation import (
MistralAudioTranscriptionConfig,
)
from litellm.types.utils import TranscriptionResponse
from litellm.utils import ProviderConfigManager
from tests.llm_translation.base_audio_transcription_unit_tests import (
BaseLLMAudioTranscriptionTest,
)
@ -37,184 +27,3 @@ class TestMistralAudioTranscription(BaseLLMAudioTranscriptionTest):
"Async audio transcription test for Mistral is skipped in this suite; "
"async test plugins (e.g. pytest-asyncio/anyio) are not configured here."
)
def test_mistral_audio_transcription_config_installed():
"""Ensure Mistral audio transcription config is registered with ProviderConfigManager."""
config = ProviderConfigManager.get_provider_audio_transcription_config(
model="mistral/voxtral-mini-latest",
provider=litellm.LlmProviders.MISTRAL,
)
assert config is not None
assert isinstance(config, BaseAudioTranscriptionConfig)
assert isinstance(config, MistralAudioTranscriptionConfig)
def test_mistral_audio_transcription_get_complete_url():
config = MistralAudioTranscriptionConfig()
url = config.get_complete_url(
api_base=None,
api_key="fake-key",
model="voxtral-mini-latest",
optional_params={},
litellm_params={},
)
assert url == "https://api.mistral.ai/v1/audio/transcriptions"
def test_mistral_audio_transcription_get_complete_url_custom_base():
config = MistralAudioTranscriptionConfig()
url = config.get_complete_url(
api_base="https://custom.api.example.com/v1/",
api_key="fake-key",
model="voxtral-mini-latest",
optional_params={},
litellm_params={},
)
assert url == "https://custom.api.example.com/v1/audio/transcriptions"
def test_mistral_audio_transcription_validate_environment():
config = MistralAudioTranscriptionConfig()
headers = config.validate_environment(
headers={},
model="voxtral-mini-latest",
messages=[],
optional_params={},
litellm_params={},
api_key="test-key-123",
)
assert headers["Authorization"] == "Bearer test-key-123"
assert headers["accept"] == "application/json"
def test_mistral_audio_transcription_supported_params():
config = MistralAudioTranscriptionConfig()
params = config.get_supported_openai_params("voxtral-mini-latest")
assert "language" in params
assert "temperature" in params
assert "response_format" in params
assert "timestamp_granularities" in params
def test_mistral_audio_transcription_request_transform():
config = MistralAudioTranscriptionConfig()
wav_path = os.path.join(
os.path.dirname(__file__),
"../../../../..",
"tests",
"llm_translation",
"gettysburg.wav",
)
audio_file = open(wav_path, "rb")
result = config.transform_audio_transcription_request(
model="voxtral-mini-latest",
audio_file=audio_file,
optional_params={"language": "en", "temperature": 0.0},
litellm_params={},
)
audio_file.close()
assert isinstance(result.data, dict)
assert result.data["model"] == "voxtral-mini-latest"
assert result.data["language"] == "en"
assert result.data["temperature"] == 0.0
assert result.files is not None
assert "file" in result.files
def test_mistral_audio_transcription_request_with_diarize():
"""Test that Mistral-specific params like diarize are passed through."""
config = MistralAudioTranscriptionConfig()
wav_path = os.path.join(
os.path.dirname(__file__),
"../../../../..",
"tests",
"llm_translation",
"gettysburg.wav",
)
audio_file = open(wav_path, "rb")
result = config.transform_audio_transcription_request(
model="voxtral-mini-latest",
audio_file=audio_file,
optional_params={"diarize": True},
litellm_params={},
)
audio_file.close()
assert isinstance(result.data, dict)
assert result.data["diarize"] == "true"
def test_mistral_audio_transcription_response_transform():
config = MistralAudioTranscriptionConfig()
mock_response = MagicMock(spec=httpx.Response)
mock_response.json.return_value = {"text": "Four score and seven years ago..."}
response = config.transform_audio_transcription_response(mock_response)
assert isinstance(response, TranscriptionResponse)
assert response.text == "Four score and seven years ago..."
def test_mistral_audio_transcription_response_transform_diarized():
"""Test that diarized responses preserve segments and language."""
config = MistralAudioTranscriptionConfig()
mock_response = MagicMock(spec=httpx.Response)
mock_response.json.return_value = {
"model": "voxtral-mini-latest",
"text": "Hello, how are you? I am fine.",
"language": None,
"segments": [
{
"text": "Hello, how are you?",
"start": 0.3,
"end": 2.1,
"speaker_id": "speaker_1",
"type": "transcription_segment",
},
{
"text": "I am fine.",
"start": 2.5,
"end": 3.8,
"speaker_id": "speaker_2",
"type": "transcription_segment",
},
],
"usage": {
"prompt_audio_seconds": 4,
"prompt_tokens": 5,
"total_tokens": 50,
"completion_tokens": 20,
},
}
response = config.transform_audio_transcription_response(mock_response)
assert isinstance(response, TranscriptionResponse)
assert response.text == "Hello, how are you? I am fine."
assert response["segments"] is not None
assert len(response["segments"]) == 2
assert response["segments"][0]["speaker_id"] == "speaker_1"
assert response["segments"][1]["speaker_id"] == "speaker_2"
assert response["language"] is None
def test_mistral_audio_transcription_response_transform_empty():
config = MistralAudioTranscriptionConfig()
mock_response = MagicMock(spec=httpx.Response)
mock_response.json.return_value = {}
response = config.transform_audio_transcription_response(mock_response)
assert isinstance(response, TranscriptionResponse)
assert response.text == ""

View file

@ -3,321 +3,12 @@ Tests for JSON-based provider configuration system.
"""
import os
import sys
from unittest.mock import patch
try:
import pytest
except ImportError:
# pytest not available, will run as standalone script
pytest = None
# Add workspace to path
workspace_path = os.path.abspath(os.path.join(os.path.dirname(__file__), "../../../.."))
sys.path.insert(0, workspace_path)
import pytest
import litellm
class TestJSONProviderLoader:
"""Test JSON provider loading and configuration"""
def test_load_json_providers(self):
"""Test that JSON providers load correctly"""
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
# Verify publicai is loaded
assert JSONProviderRegistry.exists("publicai")
# Get publicai config
publicai = JSONProviderRegistry.get("publicai")
assert publicai is not None
assert publicai.base_url == "https://api.publicai.co/v1"
assert publicai.api_key_env == "PUBLICAI_API_KEY"
assert publicai.api_base_env == "PUBLICAI_API_BASE"
assert publicai.param_mappings.get("max_completion_tokens") == "max_tokens"
def test_dynamic_config_generation(self):
"""Test dynamic config class creation"""
from litellm.llms.openai_like.dynamic_config import create_config_class
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
provider = JSONProviderRegistry.get("publicai")
config_class = create_config_class(provider)
config = config_class()
# Test API info resolution
api_base, api_key = config._get_openai_compatible_provider_info(None, None)
assert api_base == "https://api.publicai.co/v1"
# Test with custom base
api_base, api_key = config._get_openai_compatible_provider_info(
"https://custom.api.com", "test-key"
)
assert api_base == "https://custom.api.com"
assert api_key == "test-key"
def test_parameter_mapping(self):
"""Test parameter mapping works"""
from litellm.llms.openai_like.dynamic_config import create_config_class
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
provider = JSONProviderRegistry.get("publicai")
config_class = create_config_class(provider)
config = config_class()
# Test parameter mapping
optional_params = {}
non_default_params = {"max_completion_tokens": 100, "temperature": 0.7}
result = config.map_openai_params(
non_default_params, optional_params, "gpt-4", False
)
# max_completion_tokens should be mapped to max_tokens
assert "max_tokens" in result
assert result["max_tokens"] == 100
assert "max_completion_tokens" not in result
# temperature should be passed through
assert result["temperature"] == 0.7
def test_supported_params(self):
"""Test that config returns supported params"""
from litellm.llms.openai_like.dynamic_config import create_config_class
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
provider = JSONProviderRegistry.get("publicai")
config_class = create_config_class(provider)
config = config_class()
# Get supported params
supported = config.get_supported_openai_params("gpt-4")
# Should have standard OpenAI params
assert isinstance(supported, list)
assert len(supported) > 0
def test_tool_params_excluded_when_function_calling_not_supported(self):
"""Test that tool-related params are excluded for models that don't support
function calling. Regression test for https://github.com/BerriAI/litellm/issues/21125
"""
from litellm.llms.openai_like.dynamic_config import create_config_class
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
provider = JSONProviderRegistry.get("publicai")
config_class = create_config_class(provider)
config = config_class()
# Mock supports_function_calling to return False
with patch("litellm.utils.supports_function_calling", return_value=False):
supported = config.get_supported_openai_params("some-model-without-fc")
tool_params = [
"tools",
"tool_choice",
"function_call",
"functions",
"parallel_tool_calls",
]
for param in tool_params:
assert (
param not in supported
), f"'{param}' should not be in supported params when function calling is not supported"
# Non-tool params should still be present
assert "temperature" in supported
assert "max_tokens" in supported
assert "stop" in supported
def test_tool_params_included_when_function_calling_supported(self):
"""Test that tool-related params are included for models that support function calling."""
from litellm.llms.openai_like.dynamic_config import create_config_class
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
provider = JSONProviderRegistry.get("publicai")
config_class = create_config_class(provider)
config = config_class()
# Mock supports_function_calling to return True
with patch("litellm.utils.supports_function_calling", return_value=True):
supported = config.get_supported_openai_params("some-model-with-fc")
assert "tools" in supported
assert "tool_choice" in supported
def test_provider_resolution(self):
"""Test that provider resolution finds JSON providers"""
from litellm.litellm_core_utils.get_llm_provider_logic import (
get_llm_provider,
)
model, provider, api_key, api_base = get_llm_provider(
model="publicai/gpt-4",
custom_llm_provider=None,
api_base=None,
api_key=None,
)
assert model == "gpt-4"
assert provider == "publicai"
assert api_base == "https://api.publicai.co/v1"
def test_provider_config_manager(self):
"""Test that ProviderConfigManager returns JSON-based configs"""
from litellm import LlmProviders
from litellm.utils import ProviderConfigManager
config = ProviderConfigManager.get_provider_chat_config(
model="gpt-4", provider=LlmProviders.PUBLICAI
)
assert config is not None
assert config.custom_llm_provider == "publicai"
class TestPinstripes:
"""Tests for Pinstripes JSON-configured provider"""
def test_pinstripes_json_config_exists(self):
"""Test that pinstripes is configured in providers.json"""
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
assert JSONProviderRegistry.exists("pinstripes")
pinstripes = JSONProviderRegistry.get("pinstripes")
assert pinstripes is not None
assert pinstripes.base_url == "https://pinstripes.io/v1"
assert pinstripes.api_key_env == "PINSTRIPES_API_KEY"
assert pinstripes.param_mappings.get("max_completion_tokens") == "max_tokens"
def test_pinstripes_provider_resolution(self):
"""Test that provider resolution finds pinstripes and returns the default base URL"""
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
model, provider, api_key, api_base = get_llm_provider(
model="pinstripes/ps/glm-4.5-air",
custom_llm_provider=None,
api_base=None,
api_key=None,
)
assert model == "ps/glm-4.5-air"
assert provider == "pinstripes"
assert api_base == "https://pinstripes.io/v1"
def test_pinstripes_dynamic_config(self):
"""Test dynamic config class creation for pinstripes"""
from litellm.llms.openai_like.dynamic_config import create_config_class
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
provider = JSONProviderRegistry.get("pinstripes")
config_class = create_config_class(provider)
config = config_class()
api_base, api_key = config._get_openai_compatible_provider_info(None, None)
assert api_base == "https://pinstripes.io/v1"
api_base, api_key = config._get_openai_compatible_provider_info(
"https://custom.pinstripes.io/v1", "test-key"
)
assert api_base == "https://custom.pinstripes.io/v1"
assert api_key == "test-key"
def test_pinstripes_parameter_mapping(self):
"""Test that max_completion_tokens is mapped to max_tokens for pinstripes"""
from litellm.llms.openai_like.dynamic_config import create_config_class
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
provider = JSONProviderRegistry.get("pinstripes")
config_class = create_config_class(provider)
config = config_class()
optional_params = {}
non_default_params = {"max_completion_tokens": 100, "temperature": 0.7}
result = config.map_openai_params(
non_default_params, optional_params, "ps/glm-4.5-air", False
)
assert "max_tokens" in result
assert result["max_tokens"] == 100
assert "max_completion_tokens" not in result
assert result["temperature"] == 0.7
class TestDarkbloom:
def test_darkbloom_json_config_exists(self):
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
darkbloom = JSONProviderRegistry.get("darkbloom")
assert darkbloom is not None
assert darkbloom.base_url == "https://api.darkbloom.dev/v1"
assert darkbloom.api_key_env == "DARKBLOOM_API_KEY"
assert darkbloom.api_base_env == "DARKBLOOM_API_BASE"
assert darkbloom.param_mappings.get("max_completion_tokens") == "max_tokens"
def test_darkbloom_provider_resolution(self):
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
model, provider, api_key, api_base = get_llm_provider(
model="darkbloom/gemma-4-26b",
custom_llm_provider=None,
api_base=None,
api_key=None,
)
assert model == "gemma-4-26b"
assert provider == "darkbloom"
assert api_key is None
assert api_base == "https://api.darkbloom.dev/v1"
def test_darkbloom_dynamic_config(self):
from litellm.llms.openai_like.dynamic_config import create_config_class
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
provider = JSONProviderRegistry.get("darkbloom")
config_class = create_config_class(provider)
config = config_class()
api_base, api_key = config._get_openai_compatible_provider_info(None, None)
assert api_base == "https://api.darkbloom.dev/v1"
api_base, api_key = config._get_openai_compatible_provider_info(
"https://custom.darkbloom.dev/v1", "test-key"
)
assert api_base == "https://custom.darkbloom.dev/v1"
assert api_key == "test-key"
def test_darkbloom_complete_url_appends_endpoint(self):
from litellm.llms.openai_like.dynamic_config import create_config_class
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
provider = JSONProviderRegistry.get("darkbloom")
config_class = create_config_class(provider)
config = config_class()
url = config.get_complete_url(
api_base="https://api.darkbloom.dev/v1",
api_key="test-key",
model="darkbloom/gemma-4-26b",
optional_params={},
litellm_params={},
stream=True,
)
assert url == "https://api.darkbloom.dev/v1/chat/completions"
def test_darkbloom_provider_config_manager(self):
from litellm import LlmProviders
from litellm.utils import ProviderConfigManager
config = ProviderConfigManager.get_provider_chat_config(
model="gemma-4-26b", provider=LlmProviders.DARKBLOOM
)
assert config is not None
assert config.custom_llm_provider == "darkbloom"
class TestPublicAIIntegration:
"""Integration tests for PublicAI provider"""
@ -457,55 +148,3 @@ class TestPublicAIIntegration:
pytest.fail(f"Content list conversion test failed: {str(e)}")
else:
raise
if __name__ == "__main__":
# Run basic tests
print("Testing JSON Provider System...")
test_loader = TestJSONProviderLoader()
print("\n1. Testing JSON provider loading...")
test_loader.test_load_json_providers()
print(" ✓ JSON providers loaded")
print("\n2. Testing dynamic config generation...")
test_loader.test_dynamic_config_generation()
print(" ✓ Dynamic config works")
print("\n3. Testing parameter mapping...")
test_loader.test_parameter_mapping()
print(" ✓ Parameter mapping works")
print("\n4. Testing excluded params...")
test_loader.test_excluded_params()
print(" ✓ Excluded params work")
print("\n5. Testing provider resolution...")
test_loader.test_provider_resolution()
print(" ✓ Provider resolution works")
print("\n6. Testing provider config manager...")
test_loader.test_provider_config_manager()
print(" ✓ Config manager works")
print("\n" + "=" * 50)
print("PublicAI Integration Tests...")
print("=" * 50)
test_integration = TestPublicAIIntegration()
print("\n7. Testing basic completion...")
test_integration.test_publicai_completion_basic()
print("\n8. Testing streaming...")
test_integration.test_publicai_completion_with_streaming()
print("\n9. Testing parameter mapping...")
test_integration.test_publicai_parameter_mapping()
print("\n10. Testing content list conversion...")
test_integration.test_publicai_content_list_conversion()
print("\n" + "=" * 50)
print("✓ All tests passed!")
print("=" * 50)

View file

@ -4,86 +4,12 @@ Related to issue #18794
"""
import os
import sys
from unittest.mock import MagicMock, patch
try:
import pytest
except ImportError:
pytest = None
# Add workspace to path
workspace_path = os.path.abspath(os.path.join(os.path.dirname(__file__), "../../../.."))
sys.path.insert(0, workspace_path)
import pytest
import litellm
class TestXiaomiMiMoProviderConfig:
"""Test Xiaomi MiMo provider configuration"""
def test_xiaomi_mimo_in_provider_list(self):
"""Test that xiaomi_mimo is in the provider list (fixes #18794)"""
from litellm import LlmProviders
# Verify xiaomi_mimo is in the enum
assert hasattr(LlmProviders, "XIAOMI_MIMO")
assert LlmProviders.XIAOMI_MIMO.value == "xiaomi_mimo"
# Verify it's in the provider list
assert "xiaomi_mimo" in litellm.provider_list
def test_xiaomi_mimo_json_config_exists(self):
"""Test that xiaomi_mimo is configured in providers.json"""
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
# Verify xiaomi_mimo is loaded
assert JSONProviderRegistry.exists("xiaomi_mimo")
# Get xiaomi_mimo config
xiaomi_mimo = JSONProviderRegistry.get("xiaomi_mimo")
assert xiaomi_mimo is not None
assert xiaomi_mimo.base_url == "https://api.xiaomimimo.com/v1"
assert xiaomi_mimo.api_key_env == "XIAOMI_MIMO_API_KEY"
assert xiaomi_mimo.param_mappings.get("max_completion_tokens") == "max_tokens"
def test_xiaomi_mimo_provider_resolution(self):
"""Test that provider resolution finds xiaomi_mimo"""
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
model, provider, api_key, api_base = get_llm_provider(
model="xiaomi_mimo/mimo-v2-flash",
custom_llm_provider=None,
api_base=None,
api_key=None,
)
assert model == "mimo-v2-flash"
assert provider == "xiaomi_mimo"
assert api_base == "https://api.xiaomimimo.com/v1"
def test_xiaomi_mimo_router_config(self):
"""Test that xiaomi_mimo can be used in Router configuration (fixes #18794)"""
from litellm import Router
# This should not raise "Unsupported provider - xiaomi_mimo"
router = Router(
model_list=[
{
"model_name": "mimo-v2-flash",
"litellm_params": {
"model": "xiaomi_mimo/mimo-v2-flash",
"api_key": "test-key",
},
}
]
)
# Verify the deployment was created successfully
assert len(router.model_list) == 1
assert router.model_list[0]["model_name"] == "mimo-v2-flash"
class TestXiaomiMiMoIntegration:
"""Integration tests for Xiaomi MiMo provider"""
@ -128,30 +54,3 @@ class TestXiaomiMiMoIntegration:
pytest.fail(f"Xiaomi MiMo completion failed: {str(e)}")
else:
raise
if __name__ == "__main__":
# Run basic tests
print("Testing Xiaomi MiMo Provider...")
test_config = TestXiaomiMiMoProviderConfig()
print("\n1. Testing provider in list...")
test_config.test_xiaomi_mimo_in_provider_list()
print(" ✓ xiaomi_mimo in provider list")
print("\n2. Testing JSON config...")
test_config.test_xiaomi_mimo_json_config_exists()
print(" ✓ xiaomi_mimo JSON config loaded")
print("\n3. Testing provider resolution...")
test_config.test_xiaomi_mimo_provider_resolution()
print(" ✓ Provider resolution works")
print("\n4. Testing router configuration...")
test_config.test_xiaomi_mimo_router_config()
print(" ✓ Router configuration works (issue #18794 fixed)")
print("\n" + "=" * 50)
print("✓ All configuration tests passed!")
print("=" * 50)

View file

@ -54,61 +54,3 @@ def test_ovhcloud_audio_transcription_config_installed():
assert config is not None
assert isinstance(config, BaseAudioTranscriptionConfig)
class TestOVHCloudDurationFieldMigration:
"""Tests for OVHCloud duration -> seconds field migration."""
def test_seconds_field_mapped_to_duration(self):
"""New `seconds` field should be normalized to `duration`."""
from litellm.llms.ovhcloud.audio_transcription.transformation import (
OVHCloudAudioTranscriptionConfig,
)
from unittest.mock import MagicMock
config = OVHCloudAudioTranscriptionConfig()
mock_response = MagicMock()
mock_response.json.return_value = {
"text": "Hello world",
"seconds": 3.14,
}
result = config.transform_audio_transcription_response(mock_response)
assert result.text == "Hello world"
assert result._hidden_params["duration"] == 3.14
def test_legacy_duration_field_still_works(self):
"""Legacy `duration` field should still be accepted."""
from litellm.llms.ovhcloud.audio_transcription.transformation import (
OVHCloudAudioTranscriptionConfig,
)
from unittest.mock import MagicMock
config = OVHCloudAudioTranscriptionConfig()
mock_response = MagicMock()
mock_response.json.return_value = {
"text": "Hello world",
"duration": 2.71,
}
result = config.transform_audio_transcription_response(mock_response)
assert result.text == "Hello world"
assert result._hidden_params["duration"] == 2.71
def test_seconds_zero_mapped_to_duration(self):
"""seconds=0.0 must not be treated as falsy and lost."""
from litellm.llms.ovhcloud.audio_transcription.transformation import (
OVHCloudAudioTranscriptionConfig,
)
from unittest.mock import MagicMock
config = OVHCloudAudioTranscriptionConfig()
mock_response = MagicMock()
mock_response.json.return_value = {"text": "silence", "seconds": 0.0}
result = config.transform_audio_transcription_response(mock_response)
assert result._hidden_params["duration"] == 0.0

View file

@ -6,174 +6,12 @@ import os
import pytest
from litellm.llms.ovhcloud.utils import OVHCloudException
from litellm.utils import get_optional_params
from litellm.llms.ovhcloud.chat.transformation import (
OVHCloudChatCompletionStreamingHandler,
OVHCloudChatConfig,
)
config = OVHCloudChatConfig()
model = "ovhcloud/Mistral-7B-Instruct-v0.3"
class TestOvhCloudChatCompletionStreamingHandler:
def test_chunk_parser_successful(self):
handler = OVHCloudChatCompletionStreamingHandler(
streaming_response=None, sync_stream=True
)
chunk = {
"id": "test_id",
"created": 1234567890,
"model": "gpt-oss-20b",
"usage": {"prompt_tokens": 10, "completion_tokens": 20, "total_tokens": 30},
"choices": [
{"delta": {"content": "test content", "reasoning": "test reasoning"}}
],
}
result = handler.chunk_parser(chunk)
assert result.id == "test_id"
assert result.object == "chat.completion.chunk"
assert result.created == 1234567890
assert result.model == "gpt-oss-20b"
assert result.usage.prompt_tokens == chunk["usage"]["prompt_tokens"]
assert result.usage.completion_tokens == chunk["usage"]["completion_tokens"]
assert result.usage.total_tokens == chunk["usage"]["total_tokens"]
assert len(result.choices) == 1
assert result.choices[0]["delta"]["reasoning_content"] == "test reasoning"
def test_chunk_parser_error_response(self):
handler = OVHCloudChatCompletionStreamingHandler(
streaming_response=None, sync_stream=True
)
error_chunk = {
"error": {
"message": "test error",
"code": 400,
}
}
with pytest.raises(OVHCloudException) as exc_info:
handler.chunk_parser(error_chunk)
assert "OVHCloud Error: test error" in str(exc_info.value)
assert exc_info.value.status_code == 400
def test_chunk_parser_key_error(self):
handler = OVHCloudChatCompletionStreamingHandler(
streaming_response=None, sync_stream=True
)
invalid_chunk = {"incomplete": "data"}
with pytest.raises(OVHCloudException) as exc_info:
handler.chunk_parser(invalid_chunk)
assert "KeyError" in str(exc_info.value)
assert exc_info.value.status_code == 400
class TestOVHCloudConfig:
def test_transform_request_basic(self):
"""Test basic request transformation"""
transformed_request = config.transform_request(
model,
messages=[{"role": "user", "content": "Hello, world!"}],
optional_params={},
litellm_params={},
headers={},
)
assert transformed_request["model"] == model
assert transformed_request["messages"] == [
{"role": "user", "content": "Hello, world!"}
]
def test_transform_request_with_extra_body(self):
"""Test request transformation with extra_body parameters"""
transformed_request = config.transform_request(
model,
messages=[{"role": "user", "content": "Hello, world!"}],
optional_params={"extra_body": {"custom_param": "custom_value"}},
litellm_params={},
headers={},
)
assert transformed_request["custom_param"] == "custom_value"
assert transformed_request["messages"] == [
{"role": "user", "content": "Hello, world!"}
]
def test_map_openai_params(self):
"""Test OpenAI parameter mapping"""
non_default_params = {
"temperature": 0.7,
"max_tokens": 100,
"top_p": 0.9,
}
mapped_params = config.map_openai_params(
non_default_params=non_default_params,
optional_params={},
model=model,
drop_params=False,
)
assert mapped_params["temperature"] == 0.7
assert mapped_params["max_tokens"] == 100
assert mapped_params["top_p"] == 0.9
def test_get_error_class(self):
"""Test error class creation"""
error = config.get_error_class(
error_message="Test error",
status_code=400,
headers={"Content-Type": "application/json"},
)
assert isinstance(error, OVHCloudException)
assert error.message == "Test error"
assert error.status_code == 400
@pytest.mark.parametrize(
"model",
[
"Meta-Llama-3_3-70B-Instruct",
"Meta-Llama-3_1-70B-Instruct",
"Mixtral-8x7B-Instruct-v0.1",
"gpt-oss-120b",
"some-model-not-in-the-cost-map",
],
)
def test_tools_not_filtered_by_static_model_map(self, model):
"""
OVHCloud AI Endpoints are OpenAI-compatible; tools/tool_choice must pass
through for any model. The server is responsible for rejecting unsupported
tool calls — LiteLLM must not strip them based on a stale static catalog.
"""
params = get_optional_params(
model=model,
custom_llm_provider="ovhcloud",
tools=[
{
"type": "function",
"function": {"name": "x", "parameters": {}},
}
],
tool_choice="auto",
)
assert "tools" in params
assert "tool_choice" in params
def test_ovhcloud_integration():
from litellm import completion
@ -285,78 +123,3 @@ def test_ovhcloud_with_custom_base_url():
if __name__ == "__main__":
pytest.main([__file__, "-v"])
class TestOVHCloudReasoningFieldMigration:
"""Tests for OVHCloud reasoning_content -> reasoning field migration."""
def test_streaming_new_reasoning_field(self):
"""New `reasoning` field should be mapped to `reasoning_content`."""
handler = OVHCloudChatCompletionStreamingHandler(
streaming_response=iter([]),
sync_stream=True,
)
chunk = {
"id": "test-id",
"created": 1234567890,
"model": "test-model",
"choices": [
{
"delta": {
"role": "assistant",
"reasoning": "Let me think...",
},
"index": 0,
}
],
}
result = handler.chunk_parser(chunk)
assert result.choices[0]["delta"]["reasoning_content"] == "Let me think..."
def test_streaming_legacy_reasoning_content_unchanged(self):
"""Legacy `reasoning_content` field should pass through untouched."""
handler = OVHCloudChatCompletionStreamingHandler(
streaming_response=iter([]),
sync_stream=True,
)
chunk = {
"id": "test-id",
"created": 1234567890,
"model": "test-model",
"choices": [
{
"delta": {
"role": "assistant",
"reasoning_content": "Already correct field.",
},
"index": 0,
}
],
}
result = handler.chunk_parser(chunk)
assert result.choices[0]["delta"]["reasoning_content"] == "Already correct field."
def test_streaming_both_fields_legacy_wins(self):
"""When both fields present, existing `reasoning_content` is not overwritten."""
handler = OVHCloudChatCompletionStreamingHandler(
streaming_response=iter([]),
sync_stream=True,
)
chunk = {
"id": "test-id",
"created": 1234567890,
"model": "test-model",
"choices": [
{
"delta": {
"reasoning": "new field",
"reasoning_content": "legacy field",
},
"index": 0,
}
],
}
result = handler.chunk_parser(chunk)
assert result.choices[0]["delta"]["reasoning_content"] == "legacy field"

View file

@ -1 +0,0 @@
# S3 Vectors tests

View file

@ -1 +0,0 @@
# S3 Vectors vector store tests

View file

@ -1 +0,0 @@
"""Soniox provider tests."""

View file

@ -1 +0,0 @@
# Vertex AI Image Edit Tests

View file

@ -1,13 +1,9 @@
import os
from unittest.mock import MagicMock, patch
from unittest.mock import patch
import httpx
import pytest
from litellm.llms.vertex_ai.image_generation import (
get_vertex_ai_image_generation_config,
)
from litellm.llms.vertex_ai.image_generation.vertex_gemini_transformation import (
VertexAIGeminiImageGenerationConfig,
)
@ -16,588 +12,6 @@ from litellm.llms.vertex_ai.image_generation.vertex_imagen_transformation import
)
class TestVertexAIGeminiImageGenerationConfig:
def setup_method(self):
"""Set up test fixtures"""
self.config = VertexAIGeminiImageGenerationConfig()
def test_get_supported_openai_params(self):
"""Test get_supported_openai_params returns correct params"""
supported = self.config.get_supported_openai_params("gemini-2.5-flash-image")
assert "n" in supported
assert "size" in supported
def test_map_openai_params_n(self):
"""Test mapping n parameter to candidate_count"""
non_default_params = {"n": 3}
optional_params = {}
result = self.config.map_openai_params(non_default_params, optional_params, "gemini-2.5-flash-image", False)
assert result.get("candidate_count") == 3
def test_map_openai_params_size(self):
"""Test mapping size parameter to aspectRatio"""
non_default_params = {"size": "1024x1024"}
optional_params = {}
result = self.config.map_openai_params(non_default_params, optional_params, "gemini-2.5-flash-image", False)
assert result.get("aspectRatio") == "1:1"
def test_map_openai_params_size_16_9(self):
"""Test mapping 16:9 size"""
non_default_params = {"size": "1792x1024"}
optional_params = {}
result = self.config.map_openai_params(non_default_params, optional_params, "gemini-2.5-flash-image", False)
assert result.get("aspectRatio") == "16:9"
def test_map_size_to_aspect_ratio(self):
"""Test size to aspect ratio mapping"""
assert self.config._map_size_to_aspect_ratio("1024x1024") == "1:1"
assert self.config._map_size_to_aspect_ratio("1792x1024") == "16:9"
assert self.config._map_size_to_aspect_ratio("1024x1792") == "9:16"
assert self.config._map_size_to_aspect_ratio("1280x896") == "4:3"
assert self.config._map_size_to_aspect_ratio("896x1280") == "3:4"
assert self.config._map_size_to_aspect_ratio("unknown") == "1:1" # default
def test_get_supported_openai_params_includes_native_gemini_params(self):
"""Test that native Gemini imageConfig params are supported"""
supported = self.config.get_supported_openai_params("gemini-3-pro-image-preview")
assert "aspectRatio" in supported
assert "aspect_ratio" in supported
assert "imageSize" in supported
assert "image_size" in supported
assert "imageConfig" in supported
def test_map_openai_params_aspect_ratio_camel_case(self):
"""Test mapping native aspectRatio parameter"""
result = self.config.map_openai_params({"aspectRatio": "9:16"}, {}, "gemini-3-pro-image-preview", False)
assert result["aspectRatio"] == "9:16"
def test_map_openai_params_aspect_ratio_snake_case(self):
"""Test mapping native aspect_ratio parameter"""
result = self.config.map_openai_params({"aspect_ratio": "16:9"}, {}, "gemini-3-pro-image-preview", False)
assert result["aspectRatio"] == "16:9"
def test_map_openai_params_image_size_camel_case(self):
"""Test mapping native imageSize parameter"""
result = self.config.map_openai_params({"imageSize": "4K"}, {}, "gemini-3-pro-image-preview", False)
assert result["imageSize"] == "4K"
def test_map_openai_params_image_size_snake_case(self):
"""Test mapping native image_size parameter"""
result = self.config.map_openai_params({"image_size": "2K"}, {}, "gemini-3-pro-image-preview", False)
assert result["imageSize"] == "2K"
def test_map_openai_params_image_config_dict_stored_whole(self):
"""imageConfig dict is stored as-is so all fields survive"""
result = self.config.map_openai_params(
{"imageConfig": {"aspectRatio": "16:9", "imageSize": "2K"}},
{},
"gemini-3.1-flash-image",
False,
)
assert result["imageConfig"] == {"aspectRatio": "16:9", "imageSize": "2K"}
def test_map_openai_params_image_config_all_fields(self):
"""All ImageConfig fields (personGeneration, imageOutputOptions) pass through"""
payload = {
"imageConfig": {
"aspectRatio": "9:16",
"imageSize": "4K",
"personGeneration": "DONT_ALLOW",
"imageOutputOptions": {
"mimeType": "image/jpeg",
"compressionQuality": 80,
},
}
}
result = self.config.map_openai_params(payload, {}, "gemini-3.1-flash-image", False)
assert result["imageConfig"] == payload["imageConfig"]
def test_map_openai_params_image_config_non_dict_warns_and_drops(self):
"""Non-dict imageConfig is dropped with a warning, not silently discarded"""
with patch("litellm.llms.vertex_ai.image_generation.vertex_gemini_transformation.verbose_logger") as mock_log:
result = self.config.map_openai_params(
{"imageConfig": "bad-string-value"}, {}, "gemini-3.1-flash-image", False
)
assert "imageConfig" not in result
mock_log.warning.assert_called_once()
def test_transform_image_generation_request_from_image_config(self):
"""Full imageConfig dict is forwarded verbatim into generationConfig"""
full_config = {
"aspectRatio": "16:9",
"imageSize": "2K",
"personGeneration": "DONT_ALLOW",
"imageOutputOptions": {"mimeType": "image/jpeg", "compressionQuality": 85},
}
mapped = self.config.map_openai_params(
{"imageConfig": full_config},
{},
"gemini-3.1-flash-image",
False,
)
request = self.config.transform_image_generation_request(
model="gemini-3.1-flash-image",
prompt="A nano banana on a desk",
optional_params=mapped,
litellm_params={},
headers={},
)
assert request["generationConfig"]["imageConfig"] == full_config
def test_transform_image_generation_flat_params_override_image_config(self):
"""Explicit flat params win over the same key inside imageConfig"""
request = self.config.transform_image_generation_request(
model="gemini-3.1-flash-image",
prompt="A nano banana",
optional_params={
"imageConfig": {"aspectRatio": "1:1", "personGeneration": "DONT_ALLOW"},
"aspectRatio": "16:9", # should win
},
litellm_params={},
headers={},
)
assert request["generationConfig"]["imageConfig"]["aspectRatio"] == "16:9"
assert request["generationConfig"]["imageConfig"]["personGeneration"] == "DONT_ALLOW"
def test_transform_image_generation_request_basic(self):
"""Test basic request transformation"""
request = self.config.transform_image_generation_request(
model="gemini-2.5-flash-image",
prompt="A nano banana",
optional_params={},
litellm_params={},
headers={},
)
assert "contents" in request
assert "generationConfig" in request
assert request["generationConfig"]["responseModalities"] == ["IMAGE"]
assert request["contents"][0]["parts"][0]["text"] == "A nano banana"
def test_transform_image_generation_request_with_aspect_ratio(self):
"""Test request transformation with aspectRatio"""
request = self.config.transform_image_generation_request(
model="gemini-2.5-flash-image",
prompt="A nano banana",
optional_params={"aspectRatio": "16:9"},
litellm_params={},
headers={},
)
assert request["generationConfig"]["imageConfig"]["aspectRatio"] == "16:9"
def test_transform_image_generation_request_with_image_size(self):
"""Test request transformation with imageSize (Gemini 3 Pro)"""
request = self.config.transform_image_generation_request(
model="gemini-3-pro-image-preview",
prompt="A nano banana",
optional_params={"imageSize": "4K"},
litellm_params={},
headers={},
)
assert request["generationConfig"]["imageConfig"]["imageSize"] == "4K"
def test_map_openai_params_web_search_options(self):
"""Test web_search_options maps to googleSearch tool"""
result = self.config.map_openai_params({"web_search_options": {}}, {}, "gemini-3.1-flash-image-preview", False)
assert result["tools"] == [{"googleSearch": {}}]
def test_transform_image_generation_request_with_web_search_tools(self):
"""Test request transformation includes googleSearch tools"""
request = self.config.transform_image_generation_request(
model="gemini-3.1-flash-image-preview",
prompt="Generate an image of the latest iPhone",
optional_params={"tools": [{"googleSearch": {}}]},
litellm_params={},
headers={},
)
assert request["tools"] == [{"googleSearch": {}}]
def test_transform_image_generation_request_forwards_tool_config(self):
"""Test request transformation forwards toolConfig side-effects from tool mapping"""
mapped = self.config.map_openai_params(
{"tools": [{"googleMaps": {"latitude": 37.7, "longitude": -122.4}}]},
{},
"gemini-3.1-flash-image-preview",
False,
)
request = self.config.transform_image_generation_request(
model="gemini-3.1-flash-image-preview",
prompt="Generate an image of a coffee shop nearby",
optional_params=mapped,
litellm_params={},
headers={},
)
assert request["tools"] == [{"googleMaps": {}}]
assert request["toolConfig"] == {"retrievalConfig": {"latLng": {"latitude": 37.7, "longitude": -122.4}}}
def test_transform_image_generation_request_with_candidate_count(self):
"""Test request transformation with candidate_count"""
request = self.config.transform_image_generation_request(
model="gemini-2.5-flash-image",
prompt="A nano banana",
optional_params={"candidate_count": 2},
litellm_params={},
headers={},
)
assert request["generationConfig"]["candidateCount"] == 2
def test_transform_image_generation_request_with_n(self):
"""Test request transformation with n parameter"""
request = self.config.transform_image_generation_request(
model="gemini-2.5-flash-image",
prompt="A nano banana",
optional_params={"n": 2},
litellm_params={},
headers={},
)
assert request["generationConfig"]["candidateCount"] == 2
def test_transform_image_generation_response(self):
"""Test response transformation"""
mock_response = MagicMock(spec=httpx.Response)
mock_response.status_code = 200
mock_response.json.return_value = {
"candidates": [
{
"content": {
"parts": [
{
"inlineData": {
"mimeType": "image/png",
"data": "base64_encoded_image_data",
}
}
]
}
}
],
"usageMetadata": {
"promptTokenCount": 93,
"promptTokensDetails": [
{
"modality": "TEXT",
"tokenCount": 54,
},
{
"modality": "IMAGE",
"tokenCount": 39,
},
],
"candidatesTokenCount": 17,
"totalTokenCount": 110,
},
}
mock_response.headers = {}
from litellm.types.utils import ImageResponse
model_response = ImageResponse()
result = self.config.transform_image_generation_response(
model="gemini-2.5-flash-image",
raw_response=mock_response,
model_response=model_response,
logging_obj=MagicMock(),
request_data={},
optional_params={},
litellm_params={},
encoding=None,
)
assert len(result.data) == 1
assert result.data[0].b64_json == "base64_encoded_image_data"
assert result.data[0].url is None
assert result.usage.input_tokens == 93
assert result.usage.input_tokens_details.text_tokens == 54
assert result.usage.input_tokens_details.image_tokens == 39
assert result.usage.output_tokens == 17
assert result.usage.total_tokens == 110
def test_transform_image_generation_response_multiple_images(self):
"""Test response transformation with multiple images"""
mock_response = MagicMock(spec=httpx.Response)
mock_response.status_code = 200
mock_response.json.return_value = {
"candidates": [
{
"content": {
"parts": [
{
"inlineData": {
"mimeType": "image/png",
"data": "image1",
}
},
{
"inlineData": {
"mimeType": "image/png",
"data": "image2",
}
},
]
}
}
]
}
mock_response.headers = {}
from litellm.types.utils import ImageResponse
model_response = ImageResponse()
result = self.config.transform_image_generation_response(
model="gemini-2.5-flash-image",
raw_response=mock_response,
model_response=model_response,
logging_obj=MagicMock(),
request_data={},
optional_params={},
litellm_params={},
encoding=None,
)
assert len(result.data) == 2
assert result.data[0].b64_json == "image1"
assert result.data[1].b64_json == "image2"
def test_transform_image_generation_response_signature(self):
"""Test response transformation includes thoughtSignature for Gemini 3 Pro"""
mock_response = MagicMock(spec=httpx.Response)
mock_response.status_code = 200
mock_response.json.return_value = {
"candidates": [
{
"content": {
"parts": [
{
"inlineData": {
"mimeType": "image/png",
"data": "base64_encoded_image_data",
},
"thoughtSignature": "test_signature_abc123",
}
]
}
}
]
}
mock_response.headers = {}
from litellm.types.utils import ImageResponse
model_response = ImageResponse()
result = self.config.transform_image_generation_response(
model="gemini-3-pro-image-preview",
raw_response=mock_response,
model_response=model_response,
logging_obj=MagicMock(),
request_data={},
optional_params={},
litellm_params={},
encoding=None,
)
assert len(result.data) == 1
assert result.data[0].b64_json == "base64_encoded_image_data"
assert result.data[0].provider_specific_fields["thought_signature"] == "test_signature_abc123"
def test_transform_image_generation_response_tracks_web_search_requests(self):
"""Grounding queries are carried onto usage so search spend can be billed"""
mock_response = MagicMock(spec=httpx.Response)
mock_response.status_code = 200
mock_response.json.return_value = {
"candidates": [
{
"content": {
"parts": [
{
"inlineData": {
"mimeType": "image/png",
"data": "base64_encoded_image_data",
}
}
]
},
"groundingMetadata": {"webSearchQueries": ["eiffel tower", "paris skyline"]},
}
],
"usageMetadata": {
"promptTokenCount": 93,
"candidatesTokenCount": 17,
"totalTokenCount": 110,
},
}
mock_response.headers = {}
from litellm.types.utils import ImageResponse
result = self.config.transform_image_generation_response(
model="gemini-2.5-flash-image",
raw_response=mock_response,
model_response=ImageResponse(),
logging_obj=MagicMock(),
request_data={},
optional_params={},
litellm_params={},
encoding=None,
)
assert result.usage.web_search_requests == 2
class TestVertexAIImagenImageGenerationConfig:
def setup_method(self):
"""Set up test fixtures"""
self.config = VertexAIImagenImageGenerationConfig()
def test_get_supported_openai_params(self):
"""Test get_supported_openai_params returns correct params"""
supported = self.config.get_supported_openai_params("imagegeneration@006")
assert "n" in supported
assert "size" in supported
def test_map_openai_params_n(self):
"""Test mapping n parameter to sampleCount"""
non_default_params = {"n": 3}
optional_params = {}
result = self.config.map_openai_params(non_default_params, optional_params, "imagegeneration@006", False)
assert result.get("sampleCount") == 3
def test_map_openai_params_size(self):
"""Test mapping size parameter to aspectRatio"""
non_default_params = {"size": "1024x1024"}
optional_params = {}
result = self.config.map_openai_params(non_default_params, optional_params, "imagegeneration@006", False)
assert result.get("aspectRatio") == "1:1"
def test_map_size_to_aspect_ratio(self):
"""Test size to aspect ratio mapping"""
assert self.config._map_size_to_aspect_ratio("1024x1024") == "1:1"
assert self.config._map_size_to_aspect_ratio("1792x1024") == "16:9"
assert self.config._map_size_to_aspect_ratio("unknown") == "1:1" # default
def test_transform_image_generation_request_basic(self):
"""Test basic request transformation"""
request = self.config.transform_image_generation_request(
model="imagegeneration@006",
prompt="A cat",
optional_params={},
litellm_params={},
headers={},
)
assert "instances" in request
assert "parameters" in request
assert request["instances"][0]["prompt"] == "A cat"
assert request["parameters"]["sampleCount"] == 1
def test_transform_image_generation_request_with_params(self):
"""Test request transformation with parameters"""
request = self.config.transform_image_generation_request(
model="imagegeneration@006",
prompt="A cat",
optional_params={"sampleCount": 2, "aspectRatio": "16:9"},
litellm_params={},
headers={},
)
assert request["parameters"]["sampleCount"] == 2
assert request["parameters"]["aspectRatio"] == "16:9"
def test_transform_image_generation_request_labels_from_metadata(self):
"""Billing labels from litellm_params.metadata.requester_metadata on predict body."""
request = self.config.transform_image_generation_request(
model="imagegeneration@006",
prompt="A cat",
optional_params={},
litellm_params={"metadata": {"requester_metadata": {"team": "platform", "env": "prod"}}},
headers={},
)
assert request["labels"] == {"team": "platform", "env": "prod"}
assert "labels" not in request["parameters"]
def test_transform_image_generation_response(self):
"""Test response transformation"""
mock_response = MagicMock(spec=httpx.Response)
mock_response.status_code = 200
mock_response.json.return_value = {"predictions": [{"bytesBase64Encoded": "base64_encoded_image_data"}]}
mock_response.headers = {}
from litellm.types.utils import ImageResponse
model_response = ImageResponse()
result = self.config.transform_image_generation_response(
model="imagegeneration@006",
raw_response=mock_response,
model_response=model_response,
logging_obj=MagicMock(),
request_data={},
optional_params={},
litellm_params={},
encoding=None,
)
assert len(result.data) == 1
assert result.data[0].b64_json == "base64_encoded_image_data"
assert result.data[0].url is None
def test_transform_image_generation_response_multiple_images(self):
"""Test response transformation with multiple images"""
mock_response = MagicMock(spec=httpx.Response)
mock_response.status_code = 200
mock_response.json.return_value = {
"predictions": [
{"bytesBase64Encoded": "image1"},
{"bytesBase64Encoded": "image2"},
]
}
mock_response.headers = {}
from litellm.types.utils import ImageResponse
model_response = ImageResponse()
result = self.config.transform_image_generation_response(
model="imagegeneration@006",
raw_response=mock_response,
model_response=model_response,
logging_obj=MagicMock(),
request_data={},
optional_params={},
litellm_params={},
encoding=None,
)
assert len(result.data) == 2
assert result.data[0].b64_json == "image1"
assert result.data[1].b64_json == "image2"
class TestGetVertexAIImageGenerationConfig:
"""Test the router function that selects the correct config"""
def test_get_gemini_model_config(self):
"""Test that Gemini models return Gemini config"""
config = get_vertex_ai_image_generation_config("gemini-2.5-flash-image")
assert isinstance(config, VertexAIGeminiImageGenerationConfig)
config = get_vertex_ai_image_generation_config("gemini-3-pro-image-preview")
assert isinstance(config, VertexAIGeminiImageGenerationConfig)
config = get_vertex_ai_image_generation_config("vertex_ai/gemini-2.5-flash-image")
assert isinstance(config, VertexAIGeminiImageGenerationConfig)
def test_get_imagen_model_config(self):
"""Test that Imagen models return Imagen config"""
config = get_vertex_ai_image_generation_config("imagegeneration@006")
assert isinstance(config, VertexAIImagenImageGenerationConfig)
config = get_vertex_ai_image_generation_config("imagen-4.0-generate-001")
assert isinstance(config, VertexAIImagenImageGenerationConfig)
config = get_vertex_ai_image_generation_config("vertex_ai/imagegeneration@006")
assert isinstance(config, VertexAIImagenImageGenerationConfig)
def test_get_non_gemini_model_config(self):
"""Test that non-Gemini models default to Imagen config"""
config = get_vertex_ai_image_generation_config("some-other-model")
assert isinstance(config, VertexAIImagenImageGenerationConfig)
class TestVertexAIImageGenerationIntegration:
"""Integration tests for Vertex AI image generation"""
@ -642,39 +56,3 @@ class TestVertexAIImageGenerationIntegration:
litellm_params={},
)
assert "Authorization" in headers
def test_gemini_get_complete_url(self):
"""Test Gemini config URL generation"""
config = VertexAIGeminiImageGenerationConfig()
url = config.get_complete_url(
api_base=None,
api_key=None,
model="gemini-2.5-flash-image",
optional_params={},
litellm_params={
"vertex_project": "test-project",
"vertex_location": "us-central1",
},
)
assert "test-project" in url
assert "us-central1" in url
assert "gemini-2.5-flash-image" in url
assert "generateContent" in url
def test_imagen_get_complete_url(self):
"""Test Imagen config URL generation"""
config = VertexAIImagenImageGenerationConfig()
url = config.get_complete_url(
api_base=None,
api_key=None,
model="imagegeneration@006",
optional_params={},
litellm_params={
"vertex_project": "test-project",
"vertex_location": "us-central1",
},
)
assert "test-project" in url
assert "us-central1" in url
assert "imagegeneration@006" in url
assert "predict" in url

View file

@ -1 +0,0 @@
"""Tests for Vertex AI Gemma-AI models"""

View file

@ -1,3 +0,0 @@
"""
Tests for Vertex AI video generation.
"""

View file

@ -1,155 +0,0 @@
from __future__ import annotations
from collections.abc import Mapping
from typing import Final
import pytest
from pydantic import TypeAdapter
import litellm
from litellm.messages import dispatch
from litellm.rust_bridge.bindings import NativeBinding
from litellm.rust_bridge.catalog import Route, RouteRule, Rules
from litellm.rust_bridge.configuration import Rollout
from litellm.rust_bridge.messages.entrypoints import (
LiteLLMMessagesRequest,
NativeAmessages,
NativeMessages,
)
from litellm.types.llms.anthropic_messages.anthropic_response import AnthropicMessagesResponse
MESSAGES: Final = [{"role": "user", "content": "hi"}]
@pytest.mark.asyncio
async def test_public_anthropic_messages_keeps_the_python_result() -> None:
response: Final = await litellm.anthropic_messages(
model="anthropic/claude-sonnet-4-5", messages=MESSAGES, max_tokens=10, mock_response="ok"
)
assert isinstance(response, dict)
content: Final = TypeAdapter(list[dict[str, object]]).validate_python(response.get("content", []))
assert content[0]["text"] == "ok"
def test_sync_messages_request_projects_public_arguments() -> None:
rules: Final[Rules] = (RouteRule(Route.MESSAGES, Rollout.RUST_REQUIRED),)
expected: Final = AnthropicMessagesResponse(model="claude-test")
def native(
request: LiteLLMMessagesRequest, args: tuple[object, ...], kwargs: Mapping[str, object]
) -> AnthropicMessagesResponse:
assert request.model == "claude-test"
assert request.messages == MESSAGES
assert request.max_tokens == 10
assert request.custom_llm_provider == "anthropic"
return expected
binding: Final[NativeBinding[NativeMessages]] = NativeBinding("messages", validate=lambda _: None)
binding.override(native)
response: Final = dispatch._DISPATCH.run( # pyright: ignore[reportPrivateUsage] # test an explicit route decision
(),
{
"model": "claude-test",
"messages": MESSAGES,
"max_tokens": 10,
"custom_llm_provider": "anthropic",
},
python=lambda *args, **kwargs: pytest.fail("required native route must handle this call"),
binding=binding,
native=lambda hook, request, args, kwargs: hook(request, args, kwargs),
rules=rules,
)
assert response is expected
def test_messages_binding_error_delegates_unchanged_to_python() -> None:
rules: Final[Rules] = (RouteRule(Route.MESSAGES, Rollout.RUST_REQUIRED),)
expected: Final = AnthropicMessagesResponse(model="claude-test")
def python(*args: object, **kwargs: object) -> AnthropicMessagesResponse:
return expected
def native(
request: LiteLLMMessagesRequest, args: tuple[object, ...], kwargs: Mapping[str, object]
) -> AnthropicMessagesResponse:
pytest.fail("a call without max_tokens cannot project a request and must stay on Python")
binding: Final[NativeBinding[NativeMessages]] = NativeBinding("messages", validate=lambda _: None)
binding.override(native)
response: Final = dispatch._DISPATCH.run( # pyright: ignore[reportPrivateUsage] # test an explicit route decision
(),
{"model": "claude-test", "messages": MESSAGES, "custom_llm_provider": "anthropic"},
python=python,
binding=binding,
native=lambda hook, request, args, kwargs: hook(request, args, kwargs),
rules=rules,
)
assert response is expected
@pytest.mark.asyncio
async def test_async_messages_falls_back_after_native_declines() -> None:
from litellm.rust_bridge.bindings import native_exception_types
native_types: Final = native_exception_types()
if native_types is None:
pytest.skip("native bridge is unavailable")
declined, _ = native_types
expected: Final = AnthropicMessagesResponse(model="claude-test")
rules: Final[Rules] = (RouteRule(Route.MESSAGES, Rollout.RUST_OPT_OUT),)
async def native(
request: LiteLLMMessagesRequest, args: tuple[object, ...], kwargs: Mapping[str, object]
) -> AnthropicMessagesResponse:
raise declined("unsupported")
async def python(*args: object, **kwargs: object) -> AnthropicMessagesResponse:
return expected
binding: Final[NativeBinding[NativeAmessages]] = NativeBinding("amessages", validate=lambda _: None)
binding.override(native)
response: Final = await dispatch._ADISPATCH.arun( # pyright: ignore[reportPrivateUsage] # test an explicit route decision
(),
{"model": "claude-test", "messages": MESSAGES, "max_tokens": 10},
python=python,
binding=binding,
native=lambda hook, request, args, kwargs: hook(request, args, kwargs),
rules=rules,
)
assert response is expected
def test_internal_is_async_marker_bypasses_native() -> None:
rules: Final[Rules] = (RouteRule(Route.MESSAGES, Rollout.RUST_REQUIRED),)
expected: Final = AnthropicMessagesResponse(model="claude-test")
def python(*args: object, **kwargs: object) -> AnthropicMessagesResponse:
return expected
def native(
request: LiteLLMMessagesRequest, args: tuple[object, ...], kwargs: Mapping[str, object]
) -> AnthropicMessagesResponse:
pytest.fail("anthropic_messages' inner handler call must stay on Python")
binding: Final[NativeBinding[NativeMessages]] = NativeBinding("messages", validate=lambda _: None)
binding.override(native)
response: Final = dispatch._DISPATCH.run( # pyright: ignore[reportPrivateUsage] # test an explicit route decision
(),
{
"model": "claude-test",
"messages": MESSAGES,
"max_tokens": 10,
"custom_llm_provider": "anthropic",
"is_async": True,
},
python=python,
binding=binding,
native=lambda hook, request, args, kwargs: hook(request, args, kwargs),
rules=rules,
)
assert response is expected

View file

@ -30,7 +30,7 @@ from litellm.types.proxy.guardrails.guardrail_hooks.bedrock_guardrails import (
BedrockTextContent,
)
from litellm.types.utils import CallTypes, ModelResponse
from tests.test_litellm.llms.bedrock.event_loop_probe import EventLoopProbe
from tests.unit.llms.bedrock.event_loop_probe import EventLoopProbe
@pytest.mark.asyncio

View file

@ -23,7 +23,7 @@ from litellm.types.proxy.guardrails.guardrail_hooks.bedrock_guardrails import (
BedrockGuardrailResponse,
)
from litellm.types.utils import Choices, Message, ModelResponse
from tests.test_litellm.llms.bedrock.event_loop_probe import EventLoopProbe
from tests.unit.llms.bedrock.event_loop_probe import EventLoopProbe
CONTENT_FILTER_CHECKS = {"contentFilter": {"categories": [{"category": "VIOLENCE"}]}}

View file

@ -8029,6 +8029,37 @@ class TestPKCEStateCookieBinding:
assert result is not None
@pytest.mark.asyncio
@pytest.mark.parametrize("enable_sso_debug_value", [None, "false", "0"])
async def test_sso_debug_routes_return_404_unless_explicitly_enabled(enable_sso_debug_value):
"""
/sso/debug/login and /sso/debug/callback must 404 unless ENABLE_SSO_DEBUG is
explicitly set to a truthy value.
"""
from litellm.proxy.management_endpoints.ui_sso import debug_sso_callback, debug_sso_login
mock_request = MagicMock(spec=Request)
mock_request.base_url = "http://proxy.example.com/"
mock_request.cookies = {}
mock_request.query_params = {}
env = {"GENERIC_CLIENT_ID": "test_client_id"}
if enable_sso_debug_value is not None:
env["ENABLE_SSO_DEBUG"] = enable_sso_debug_value
with patch.dict(os.environ, env, clear=False):
if enable_sso_debug_value is None:
os.environ.pop("ENABLE_SSO_DEBUG", None)
with pytest.raises(HTTPException) as login_exc:
await debug_sso_login(mock_request)
with pytest.raises(HTTPException) as callback_exc:
await debug_sso_callback(mock_request)
assert login_exc.value.status_code == 404
assert callback_exc.value.status_code == 404
@pytest.mark.asyncio
async def test_debug_sso_callback_renders_full_jwt_claims():
"""
@ -8080,7 +8111,7 @@ async def test_debug_sso_callback_renders_full_jwt_claims():
with (
patch.dict(
os.environ,
{"GENERIC_CLIENT_ID": "test_client_id"},
{"GENERIC_CLIENT_ID": "test_client_id", "ENABLE_SSO_DEBUG": "true"},
clear=False,
),
patch(
@ -8165,7 +8196,7 @@ async def test_debug_sso_callback_handles_missing_raw_response():
with (
patch.dict(
os.environ,
{"MICROSOFT_CLIENT_ID": "test_microsoft_id"},
{"MICROSOFT_CLIENT_ID": "test_microsoft_id", "ENABLE_SSO_DEBUG": "true"},
clear=False,
),
patch.object(
@ -8213,7 +8244,7 @@ async def _render_debug_page(provider_env, id_jag_registered, force_inert=False)
return parsed
stack = [
patch.dict(os.environ, provider_env, clear=False),
patch.dict(os.environ, {**provider_env, "ENABLE_SSO_DEBUG": "true"}, clear=False),
patch( # test-quality-ok: endpoint test stubs the upstream generic IdP boundary
"litellm.proxy.management_endpoints.ui_sso.get_generic_sso_response", side_effect=fake_generic
),

View file

@ -24,7 +24,7 @@ from starlette.datastructures import FormData
import litellm
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
from tests.test_litellm.llms.bedrock.event_loop_probe import EventLoopProbe
from tests.unit.llms.bedrock.event_loop_probe import EventLoopProbe
from litellm.constants import LITELLM_PROXY_MASTER_KEY_ALIAS
from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
BaseOpenAIPassThroughHandler,

View file

@ -2823,7 +2823,7 @@ def test_update_router_config_schema_includes_tag_routing_prefix():
# UpdateRouterConfig before calling update_settings; a field missing here
# causes model_dump(exclude_none=True) to silently drop it before
# update_settings is ever called -- the same bug shape LIT-3152 fixed for
# retry_policy (see tests/test_litellm/test_router_retry_policy_update.py).
# retry_policy (see tests/unit/test_router_retry_policy_update.py).
from litellm.types.router import UpdateRouterConfig
config = UpdateRouterConfig(tag_routing_prefix="route:")

View file

@ -3,20 +3,13 @@ Unit tests for litellm.compress().
"""
import os
import importlib
import pytest
import litellm
from litellm.compression.scoring.bm25 import bm25_score_messages
from litellm.compression.scoring.embedding_scorer import embedding_score_messages
from litellm.compression.content_detection import detect_content_type
from litellm.compression.message_stubbing import extract_key, stub_message
from litellm.compression.retrieval_tool import build_retrieval_tool
from litellm.types.utils import CallTypes
CALL_TYPE = CallTypes.completion
ANTHROPIC_CALL_TYPE = CallTypes.anthropic_messages
# ---------------------------------------------------------------------------
@ -24,420 +17,26 @@ ANTHROPIC_CALL_TYPE = CallTypes.anthropic_messages
# ---------------------------------------------------------------------------
def test_bm25_relevance_ranking():
query = "Fix the authentication bug in the login handler"
messages = [
{
"role": "user",
"content": "def login_handler(): authentication check bug fix",
},
{"role": "user", "content": "def render_template(name): css styling layout"},
{"role": "user", "content": "def verify(): authentication token bug handler"},
]
scores = bm25_score_messages(query, messages)
# Messages sharing query terms should score higher than unrelated ones
assert scores[0] > scores[1]
assert scores[2] > scores[1]
def test_bm25_empty_query():
scores = bm25_score_messages("", [{"role": "user", "content": "hello"}])
assert scores == [0.0]
def test_bm25_empty_messages():
scores = bm25_score_messages("query", [])
assert scores == []
def test_bm25_empty_content():
scores = bm25_score_messages("query", [{"role": "user", "content": ""}])
assert scores == [0.0]
# ---------------------------------------------------------------------------
# Content detection
# ---------------------------------------------------------------------------
def test_detect_code():
code = """
import os
from pathlib import Path
def main():
class Foo:
pass
return Foo()
"""
assert detect_content_type(code) == "code"
def test_detect_json():
assert detect_content_type('{"key": "value", "num": 42}') == "json"
assert detect_content_type("[1, 2, 3]") == "json"
def test_detect_text():
assert detect_content_type("This is a plain text paragraph about dogs.") == "text"
def test_detect_empty():
assert detect_content_type("") == "text"
# ---------------------------------------------------------------------------
# Message stubbing
# ---------------------------------------------------------------------------
def test_extract_key_with_filename():
msg = {"role": "user", "content": "# auth.py\ndef authenticate():\n pass"}
used: set = set()
key = extract_key(msg, fallback_index=0, used_keys=used)
assert key == "auth.py"
def test_extract_key_fallback():
msg = {"role": "user", "content": "Some random content without a filename"}
used: set = set()
key = extract_key(msg, fallback_index=5, used_keys=used)
assert key == "message_5"
def test_extract_key_duplicates():
used: set = set()
msg = {"role": "user", "content": "# auth.py\ncode here"}
k1 = extract_key(msg, fallback_index=0, used_keys=used)
k2 = extract_key(msg, fallback_index=1, used_keys=used)
assert k1 == "auth.py"
assert k2 == "auth.py_2"
def test_stub_message():
msg = {"role": "user", "content": "line1\nline2\nline3"}
stubbed = stub_message(msg, "test_key")
assert stubbed["role"] == "user"
assert "test_key" in stubbed["content"]
assert "litellm_content_retrieve" in stubbed["content"]
assert "3 lines" in stubbed["content"]
# ---------------------------------------------------------------------------
# Retrieval tool
# ---------------------------------------------------------------------------
def test_retrieval_tool_schema():
tool = build_retrieval_tool(["auth.py", "utils.py"])
assert tool["type"] == "function"
assert tool["function"]["name"] == "litellm_content_retrieve"
assert "key" in tool["function"]["parameters"]["properties"]
assert tool["function"]["parameters"]["properties"]["key"]["enum"] == [
"auth.py",
"utils.py",
]
assert tool["function"]["parameters"]["required"] == ["key"]
def test_retrieval_tool_description_lists_keys():
tool = build_retrieval_tool(["foo.py", "bar.js"])
desc = tool["function"]["description"]
assert "foo.py" in desc
assert "bar.js" in desc
# ---------------------------------------------------------------------------
# compress() — end-to-end
# ---------------------------------------------------------------------------
def test_compress_below_trigger_passthrough():
messages = [{"role": "user", "content": "hello"}]
result = litellm.compress(messages, model="gpt-4o", call_type=CALL_TYPE)
assert result["messages"] == messages
assert result["cache"] == {}
assert result["tools"] == []
assert result["compression_ratio"] == 0.0
assert result["original_tokens"] == result["compressed_tokens"]
def test_compress_above_trigger():
big_messages = [
{"role": "system", "content": "You are a coding assistant."},
{
"role": "user",
"content": "# auth.py\n" + "def authenticate():\n pass\n" * 2000,
},
{
"role": "user",
"content": "# utils.py\n" + "def helper():\n pass\n" * 2000,
},
{
"role": "user",
"content": "# readme.md\n" + "This is documentation. " * 2000,
},
{"role": "user", "content": "Fix the bug in auth.py"},
]
result = litellm.compress(
big_messages,
model="gpt-4o",
call_type=CALL_TYPE,
compression_trigger=1000,
compression_target=500,
)
assert result["compressed_tokens"] < result["original_tokens"]
assert result["compression_ratio"] > 0
assert len(result["cache"]) > 0
assert len(result["tools"]) == 1
assert result["tools"][0]["function"]["name"] == "litellm_content_retrieve"
def test_compress_anthropic_list_content_is_boundary_stable():
messages = [
{"role": "system", "content": [{"type": "text", "text": "System prompt"}]},
{
"role": "user",
"content": [
{"type": "text", "text": "# a.py\n" + "alpha " * 2000},
{
"type": "image_url",
"image_url": {"url": "https://example.com/a.png"},
},
],
},
{
"role": "user",
"content": [
{"type": "text", "text": "# b.py\n" + "beta " * 2000},
{
"type": "image_url",
"image_url": {"url": "https://example.com/b.png"},
},
],
},
{
"role": "user",
"content": [{"type": "text", "text": "Fix alpha bug in a.py"}],
},
]
result = litellm.compress(
messages=messages,
model="claude-sonnet-4-20250514",
call_type=ANTHROPIC_CALL_TYPE,
compression_trigger=1000,
compression_target=500,
)
assert result["compressed_tokens"] < result["original_tokens"]
assert len(result["messages"]) == len(messages)
assert [m["role"] for m in result["messages"]] == [m["role"] for m in messages]
assert len(result["cache"]) > 0
assert len(result["tools"]) == 1
assert result["tools"][0]["type"] == "custom"
assert result["tools"][0]["name"] == "litellm_content_retrieve"
assert "input_schema" in result["tools"][0]
def test_compress_preserves_system_message():
messages = [
{"role": "system", "content": "System prompt. " * 500},
{"role": "user", "content": "Large file content. " * 5000},
{"role": "user", "content": "Fix the bug"},
]
result = litellm.compress(
messages, model="gpt-4o", call_type=CALL_TYPE, compression_trigger=1000
)
assert result["messages"][0]["role"] == "system"
assert "System prompt" in result["messages"][0]["content"]
def test_compress_preserves_last_user_message():
messages = [
{"role": "user", "content": "Big context " * 5000},
{"role": "user", "content": "Fix the bug in auth.py"},
]
result = litellm.compress(
messages, model="gpt-4o", call_type=CALL_TYPE, compression_trigger=1000
)
last_user = [m for m in result["messages"] if m["role"] == "user"][-1]
assert "Fix the bug in auth.py" in last_user["content"]
def test_compress_preserves_last_assistant_message():
messages = [
{"role": "user", "content": "Big context " * 5000},
{"role": "assistant", "content": "I'll help with that. " * 2000},
{"role": "user", "content": "Now fix the bug"},
]
result = litellm.compress(
messages, model="gpt-4o", call_type=CALL_TYPE, compression_trigger=1000
)
assistant_msgs = [m for m in result["messages"] if m["role"] == "assistant"]
assert len(assistant_msgs) >= 1
# The last assistant message should be preserved (not stubbed)
last_assistant = assistant_msgs[-1]
assert "I'll help with that" in last_assistant["content"]
def test_cache_keys_match_stubs():
messages = [
{"role": "user", "content": "# auth.py\n" + "code " * 5000},
{"role": "user", "content": "Fix it"},
]
result = litellm.compress(
messages, model="gpt-4o", call_type=CALL_TYPE, compression_trigger=1000
)
if result["tools"]:
tool_desc = result["tools"][0]["function"]["description"]
for key in result["cache"]:
assert key in tool_desc
def test_compress_default_target():
"""compression_target defaults to compression_trigger // 2."""
messages = [
{"role": "user", "content": "content " * 5000},
{"role": "user", "content": "query"},
]
result = litellm.compress(
messages, model="gpt-4o", call_type=CALL_TYPE, compression_trigger=2000
)
# Should have compressed — target = 1000
assert result["compressed_tokens"] <= result["original_tokens"]
def test_compress_nested_tool_result_extracts_text_only():
messages = [
{"role": "system", "content": [{"type": "text", "text": "System rules"}]},
{
"role": "user",
"content": [
{"type": "text", "text": "prefix"},
{
"type": "tool_result",
"tool_use_id": "toolu_1",
"content": [
{"type": "text", "text": "nested text fragment"},
{
"type": "image_url",
"image_url": {
"url": "https://example.com/secret-tool.png",
},
},
],
},
{
"type": "image_url",
"image_url": {"url": "https://example.com/top.png"},
},
{"type": "text", "text": " " + ("irrelevant " * 3000)},
],
},
{
"role": "user",
"content": [{"type": "text", "text": "final query that must remain"}],
},
]
result = litellm.compress(
messages=messages,
model="claude-sonnet-4-20250514",
call_type=ANTHROPIC_CALL_TYPE,
compression_trigger=500,
compression_target=100,
)
cached_text = " ".join(result["cache"].values())
assert "nested text fragment" in cached_text
assert "https://example.com/secret-tool.png" not in cached_text
assert "https://example.com/top.png" not in cached_text
def test_compress_default_call_type_is_completion():
result = litellm.compress(
messages=[
{"role": "user", "content": "Large context " * 4000},
{"role": "user", "content": "query"},
],
model="gpt-4o",
compression_trigger=1000,
compression_target=500,
)
assert result["compressed_tokens"] <= result["original_tokens"]
assert isinstance(result["tools"], list)
def test_compress_forwards_embedding_model_params(monkeypatch):
captured = {}
def fake_embedding_score_messages(
query, messages, model, cache=None, embedding_model_params=None
):
captured["query"] = query
captured["model"] = model
captured["embedding_model_params"] = embedding_model_params
return [0.0] * len(messages)
monkeypatch.setattr(
"litellm.compression.scoring.embedding_scorer.embedding_score_messages",
fake_embedding_score_messages,
)
result = litellm.compress(
messages=[
{"role": "user", "content": "Authentication code " * 2000},
{"role": "user", "content": "Fix auth"},
],
model="gpt-4o",
call_type=CALL_TYPE,
compression_trigger=1000,
embedding_model="text-embedding-3-small",
embedding_model_params={"api_base": "https://example-embeddings.test"},
)
assert result["compressed_tokens"] <= result["original_tokens"]
assert captured["model"] == "text-embedding-3-small"
assert captured["embedding_model_params"] == {
"api_base": "https://example-embeddings.test"
}
def test_embedding_scorer_forwards_embedding_model_params(monkeypatch):
captured = {}
class _MockResponse:
data = [
{"embedding": [1.0, 0.0]},
{"embedding": [1.0, 0.0]},
{"embedding": [0.0, 1.0]},
]
def fake_embedding(**kwargs):
captured.update(kwargs)
return _MockResponse()
monkeypatch.setattr(litellm, "embedding", fake_embedding)
scores = embedding_score_messages(
query="auth",
messages=[
{"role": "user", "content": "auth code"},
{"role": "user", "content": "cooking recipe"},
],
model="text-embedding-3-small",
embedding_model_params={"api_base": "https://example-embeddings.test"},
)
assert len(scores) == 2
assert captured["model"] == "text-embedding-3-small"
assert captured["api_base"] == "https://example-embeddings.test"
# ---------------------------------------------------------------------------
# Embedding scorer — integration test (skipped without API key)
# ---------------------------------------------------------------------------
@ -458,210 +57,3 @@ def test_embedding_scorer():
)
assert result["compression_ratio"] > 0
assert len(result["cache"]) > 0
@pytest.mark.parametrize(
"final_user_message, expected_content",
[
("How to cook?", "Unrelated cooking recipes "),
("Fix auth", "Authentication code "),
],
)
def test_simple_compression(final_user_message, expected_content):
messages = [
{"role": "user", "content": "Authentication code " * 2000},
{"role": "user", "content": "Unrelated cooking recipes " * 2000},
{"role": "user", "content": final_user_message},
]
result = litellm.compress(
messages, model="gpt-4o", call_type=CALL_TYPE, compression_trigger=1000
)
if expected_content == "Unrelated cooking recipes ":
assert "Unrelated cooking recipes " in result["messages"][1]["content"]
assert "Authentication code " not in result["messages"][0]["content"]
elif expected_content == "Authentication code ":
assert "Authentication code " in result["messages"][0]["content"]
assert "Unrelated cooking recipes " not in result["messages"][1]["content"]
else:
raise ValueError(f"Unexpected expected_content: {expected_content}")
def test_compress_anthropic_drops_irrelevant_tool_exchange_span(monkeypatch):
compress_module = importlib.import_module("litellm.compression.compress")
def fake_bm25_score_messages(query, messages):
assert "final query" in query
assert len(messages) == 5
# Prefer idx=0 and de-prioritize the tool exchange span (idx=1,2)
return [0.95, 0.01, 0.02, 0.8, 1.0]
def fake_token_counter(model, messages=None, text=None):
if messages is not None:
return 1000
if text is None:
return 0
if "final query" in text:
return 50
if "assistant_tail" in text:
return 20
if "other_blob" in text:
return 220
if "tool_payload_relevant" in text:
return 200
if text == "":
return 1
return 10
monkeypatch.setattr(
compress_module, "bm25_score_messages", fake_bm25_score_messages
)
monkeypatch.setattr(compress_module, "token_counter", fake_token_counter)
messages = [
{"role": "user", "content": "other_blob " * 300},
{
"role": "assistant",
"content": [
{
"type": "tool_use",
"id": "toolu_drop",
"name": "litellm_content_retrieve",
"input": {"key": "message_1"},
}
],
},
{
"role": "user",
"content": [
{
"type": "tool_result",
"tool_use_id": "toolu_drop",
"content": [{"type": "text", "text": "tool_payload_relevant"}],
}
],
},
{"role": "assistant", "content": "assistant_tail"},
{"role": "user", "content": "final query"},
]
result = litellm.compress(
messages=messages,
model="claude-sonnet-4-20250514",
call_type=ANTHROPIC_CALL_TYPE,
compression_trigger=100,
compression_target=280,
)
# idx=1,2 should be dropped atomically (no orphan tool blocks left behind)
assert len(result["messages"]) == 3
assert result["messages"][0]["role"] == "user"
assert "other_blob" in result["messages"][0]["content"]
assert result["messages"][1]["content"] == "assistant_tail"
assert result["messages"][2]["content"] == "final query"
assert result["cache"] == {}
def test_compress_anthropic_keeps_relevant_tool_exchange_span(monkeypatch):
compress_module = importlib.import_module("litellm.compression.compress")
def fake_bm25_score_messages(query, messages):
assert "final query" in query
assert len(messages) == 5
# Prefer the tool exchange span over idx=0
return [0.05, 0.01, 0.92, 0.8, 1.0]
def fake_token_counter(model, messages=None, text=None):
if messages is not None:
return 1000
if text is None:
return 0
if "final query" in text:
return 50
if "assistant_tail" in text:
return 20
if "other_blob" in text:
return 220
if "tool_payload_relevant" in text:
return 200
if text == "":
return 1
return 10
monkeypatch.setattr(
compress_module, "bm25_score_messages", fake_bm25_score_messages
)
monkeypatch.setattr(compress_module, "token_counter", fake_token_counter)
messages = [
{"role": "user", "content": "other_blob " * 300},
{
"role": "assistant",
"content": [
{
"type": "tool_use",
"id": "toolu_keep",
"name": "litellm_content_retrieve",
"input": {"key": "message_1"},
}
],
},
{
"role": "user",
"content": [
{
"type": "tool_result",
"tool_use_id": "toolu_keep",
"content": [{"type": "text", "text": "tool_payload_relevant"}],
}
],
},
{"role": "assistant", "content": "assistant_tail"},
{"role": "user", "content": "final query"},
]
result = litellm.compress(
messages=messages,
model="claude-sonnet-4-20250514",
call_type=ANTHROPIC_CALL_TYPE,
compression_trigger=100,
compression_target=280,
)
assert len(result["messages"]) == 5
assert result["messages"][1]["role"] == "assistant"
assert result["messages"][2]["role"] == "user"
# idx=0 should be compressed instead
assert "litellm_content_retrieve" in result["messages"][0]["content"]
assert len(result["cache"]) == 1
def test_compress_anthropic_malformed_tool_sequence_passes_through():
messages = [
{"role": "user", "content": "other_blob " * 300},
{
"role": "assistant",
"content": [
{
"type": "tool_use",
"id": "toolu_broken",
"name": "litellm_content_retrieve",
"input": {"key": "message_1"},
}
],
},
{"role": "user", "content": [{"type": "text", "text": "missing tool_result"}]},
{"role": "user", "content": "final query"},
]
result = litellm.compress(
messages=messages,
model="claude-sonnet-4-20250514",
call_type=ANTHROPIC_CALL_TYPE,
compression_trigger=100,
compression_target=280,
)
assert result["messages"] == messages
assert result["cache"] == {}
assert result["tools"] == []
assert result["compression_skipped_reason"] == "invalid_anthropic_tool_sequence"

File diff suppressed because it is too large Load diff

View file

@ -2072,3 +2072,348 @@ def test_chat_rows_from_mistral_still_use_token_pricing(monkeypatch):
)
assert result.cost == pytest.approx((10 * 0.001 + 5 * 0.002) / 2)
assert result.usage.total_tokens == 15
GROUNDED_USAGE_METADATA = {
"promptTokenCount": 19,
"candidatesTokenCount": 59,
"thoughtsTokenCount": 406,
"toolUsePromptTokenCount": 73,
"totalTokenCount": 557,
"promptTokensDetails": [{"modality": "TEXT", "tokenCount": 19}],
"candidatesTokensDetails": [{"modality": "TEXT", "tokenCount": 59}],
"toolUsePromptTokensDetails": [{"modality": "TEXT", "tokenCount": 73}],
"trafficType": "ON_DEMAND",
}
PASSTHROUGH_OUTPUT_URI = (
"gs://litellm-bucket/litellm-vertex-files/passthrough/publishers/google/models/gemini-2.5-flash/u/"
"predictions.jsonl"
)
UNGROUNDED_USAGE_METADATA = {
"promptTokenCount": 20,
"candidatesTokenCount": 48,
"thoughtsTokenCount": 195,
"toolUsePromptTokenCount": 73,
"totalTokenCount": 336,
"promptTokensDetails": [{"modality": "TEXT", "tokenCount": 20}],
"trafficType": "ON_DEMAND",
}
def _native_vertex_row(usage_metadata: dict, *, grounded: bool, model_version: str | None = "gemini-2.5-flash"):
candidate = {"content": {"role": "model", "parts": [{"text": "ok"}]}, "finishReason": "STOP"}
grounding = {"groundingMetadata": {"webSearchQueries": ["q"]}} if grounded else {}
response = {"candidates": [{**candidate, **grounding}], "usageMetadata": usage_metadata}
return {
"request": {"contents": [{"role": "user", "parts": [{"text": "q"}]}], "tools": [{"googleSearch": {}}]},
"status": "",
"response": {**response, **({"modelVersion": model_version} if model_version else {})},
"processed_time": "2026-09-23T19:02:00.000+00:00",
}
def _capture_cost_calls(monkeypatch, prompt_cost=0.5, completion_cost=0.25) -> list:
import litellm.cost_calculator as cc
calls: list = []
def _calc(**kw):
calls.append(kw)
return (prompt_cost, completion_cost)
monkeypatch.setattr(cc, "batch_cost_calculator", _calc)
return calls
def test_vertex_native_cost_bills_embedding_rows(monkeypatch):
monkeypatch.setitem(litellm.model_cost, "vertex_ai/gemini-embedding-2", {"input_cost_per_token_batches": 1e-7})
rows = [
{
"key": "id_1",
"status": "",
"request": {"content": {"parts": [{"text": "hello world"}]}},
"response": {"embedding": {"values": [0.1, 0.2]}, "usageMetadata": {"promptTokenCount": 2}},
},
{
"key": "id_2",
"status": "",
"request": {"content": {"parts": [{"text": "hello"}]}},
"response": {"embedding": {"values": [0.3]}, "tokenCount": "3"},
},
{"key": "id_3", "status": "INVALID_ARGUMENT", "request": {"content": {"parts": [{"text": ""}]}}},
]
result = bu.calculate_vertex_ai_batch_cost_and_usage(rows, "gemini-embedding-2")
assert (result.successful_requests, result.failed_requests) == (2, 1)
assert (result.usage.prompt_tokens, result.usage.completion_tokens, result.usage.total_tokens) == (5, 0, 5)
assert result.cost == pytest.approx(5 * 1e-7)
assert result.models == ["gemini-embedding-2"]
@pytest.mark.asyncio
async def test_native_vertex_rows_route_to_vertex_cost_path_without_flag(monkeypatch):
monkeypatch.setattr(litellm, "disable_vertex_batch_output_transformation", False, raising=False)
monkeypatch.setattr(
bu, "_aggregate_batch_cost_usage_models", lambda **kw: pytest.fail("generic path should not run")
)
calls = _capture_cost_calls(monkeypatch)
rows = [
_native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True),
_native_vertex_row(UNGROUNDED_USAGE_METADATA, grounded=False),
]
result = await bu.calculate_batch_cost_and_usage(
file_content_dictionary=rows, custom_llm_provider="vertex_ai", model_name="gemini-2.5-flash"
)
assert result.cost == pytest.approx(1.5)
assert (result.successful_requests, result.failed_requests) == (2, 0)
assert result.models == ["gemini-2.5-flash"]
assert {(call["model"], call["custom_llm_provider"]) for call in calls} == {("gemini-2.5-flash", "vertex_ai")}
@pytest.mark.asyncio
async def test_openai_shaped_vertex_rows_keep_the_generic_path_without_flag(monkeypatch):
monkeypatch.setattr(litellm, "disable_vertex_batch_output_transformation", False, raising=False)
monkeypatch.setattr(
bu, "calculate_vertex_ai_batch_cost_and_usage", lambda *a, **kw: pytest.fail("native path should not run")
)
_capture_cost_calls(monkeypatch)
rows = [_vertex_openai_row("request-1", "gemini-2.5-flash", 10, 5)]
result = await bu.calculate_batch_cost_and_usage(
file_content_dictionary=rows, custom_llm_provider="vertex_ai", model_name="gemini-2.5-flash"
)
assert result.successful_requests == 1
@pytest.mark.asyncio
async def test_native_vertex_rows_on_another_provider_keep_the_generic_path(monkeypatch):
monkeypatch.setattr(
bu, "calculate_vertex_ai_batch_cost_and_usage", lambda *a, **kw: pytest.fail("native path should not run")
)
_capture_cost_calls(monkeypatch)
result = await bu.calculate_batch_cost_and_usage(
file_content_dictionary=[_native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True)],
custom_llm_provider="openai",
)
assert result.successful_requests == 0
@pytest.mark.asyncio
async def test_handle_completed_batch_routes_native_rows_without_flag(monkeypatch):
monkeypatch.setattr(litellm, "disable_vertex_batch_output_transformation", False, raising=False)
raw_rows = [_native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True)]
async def fake_fetch(batch, custom_llm_provider, litellm_params=None):
return _vertex_jsonl(raw_rows)
monkeypatch.setattr(bu, "_fetch_batch_output_file_content", fake_fetch)
monkeypatch.setattr(
bu, "_aggregate_batch_cost_usage_models", lambda **kw: pytest.fail("generic path should not run")
)
calls = _capture_cost_calls(monkeypatch, prompt_cost=0.7, completion_cost=0.3)
deployment_model_info = {"input_cost_per_token_batches": 1e-6, "output_cost_per_token_batches": 2e-6}
result = await bu._handle_completed_batch(
_batch(PASSTHROUGH_OUTPUT_URI),
custom_llm_provider="vertex_ai",
model_name="gemini-2.5-flash",
model_info=deployment_model_info,
)
assert result.cost == pytest.approx(1.0)
assert result.usage.total_tokens == 557
assert [call["model_info"] for call in calls] == [deployment_model_info]
def test_native_vertex_usage_is_billed_like_the_online_path(monkeypatch):
calls = _capture_cost_calls(monkeypatch)
grounded = _native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True)
ungrounded = _native_vertex_row(UNGROUNDED_USAGE_METADATA, grounded=False)
result = bu.calculate_vertex_ai_batch_cost_and_usage([grounded, ungrounded], "gemini-2.5-flash")
grounded_usage, ungrounded_usage = (call["usage"] for call in calls)
assert grounded_usage.prompt_tokens == 19
assert grounded_usage.completion_tokens == 59 + 406
assert grounded_usage.completion_tokens_details.reasoning_tokens == 406
assert ungrounded_usage.prompt_tokens == 20 + 73
assert ungrounded_usage.completion_tokens == 48 + 195
assert (result.usage.prompt_tokens, result.usage.completion_tokens, result.usage.total_tokens) == (
19 + 93,
465 + 243,
557 + 336,
)
def test_native_vertex_rows_are_priced_by_model_version_without_a_model_name(monkeypatch):
calls = _capture_cost_calls(monkeypatch)
rows = [
_native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True, model_version="gemini-2.5-flash"),
_native_vertex_row(UNGROUNDED_USAGE_METADATA, grounded=False, model_version="gemini-2.5-pro"),
_native_vertex_row(UNGROUNDED_USAGE_METADATA, grounded=False, model_version=None),
]
result = bu.calculate_vertex_ai_batch_cost_and_usage(rows)
assert [call["model"] for call in calls] == ["gemini-2.5-flash", "gemini-2.5-pro"]
assert result.models == ["gemini-2.5-flash", "gemini-2.5-pro"]
assert result.cost == pytest.approx(1.5)
assert result.successful_requests == 3
assert result.usage.total_tokens == 557 + 336 + 336
def test_native_vertex_rows_without_usage_metadata_count_as_failed(monkeypatch):
_capture_cost_calls(monkeypatch)
rows = [
{"request": {"contents": []}, "status": "Error: bad request", "processed_time": "t"},
{"request": {"contents": []}, "response": {"candidates": []}},
_native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True),
]
result = bu.calculate_vertex_ai_batch_cost_and_usage(rows, "gemini-2.5-flash")
assert (result.successful_requests, result.failed_requests) == (1, 2)
assert result.usage.total_tokens == 557
def test_native_vertex_batch_whose_rows_all_failed_still_names_the_deployment_model(monkeypatch):
calls = _capture_cost_calls(monkeypatch)
rows = [{"request": {"contents": []}, "status": "Error: quota exceeded", "processed_time": "t"}] * 2
result = bu.calculate_vertex_ai_batch_cost_and_usage(rows, "gemini-2.5-flash")
assert result.models == ["gemini-2.5-flash"]
assert (result.successful_requests, result.failed_requests, result.cost) == (0, 2, 0.0)
assert calls == []
def test_native_vertex_rows_are_priced_with_the_deployment_model_info(monkeypatch):
calls = _capture_cost_calls(monkeypatch)
deployment_model_info = {"input_cost_per_token_batches": 1e-6, "output_cost_per_token_batches": 2e-6}
bu.calculate_vertex_ai_batch_cost_and_usage(
[_native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True)],
"gemini-2.5-flash",
model_info=deployment_model_info,
)
assert [call["model_info"] for call in calls] == [deployment_model_info]
@pytest.mark.asyncio
async def test_native_vertex_rows_keep_the_deployment_model_info_through_the_batch_entrypoint(monkeypatch):
calls = _capture_cost_calls(monkeypatch)
deployment_model_info = {"input_cost_per_token_batches": 1e-6}
await bu.calculate_batch_cost_and_usage(
file_content_dictionary=[_native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True)],
custom_llm_provider="vertex_ai",
model_name="gemini-2.5-flash",
model_info=deployment_model_info,
)
assert [call["model_info"] for call in calls] == [deployment_model_info]
def test_native_vertex_rows_are_priced_by_the_deployment_model_over_model_version(monkeypatch):
calls = _capture_cost_calls(monkeypatch)
rows = [_native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True, model_version="gemini-2.5-pro")]
result = bu.calculate_vertex_ai_batch_cost_and_usage(rows, "gemini-2.5-flash")
assert [call["model"] for call in calls] == ["gemini-2.5-flash"]
assert result.models == ["gemini-2.5-flash"]
def test_native_vertex_rows_that_fail_response_validation_count_as_failed(monkeypatch):
calls = _capture_cost_calls(monkeypatch)
rows = [
{"request": {"contents": []}, "response": {"candidates": "nope", "usageMetadata": GROUNDED_USAGE_METADATA}},
_native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True),
]
result = bu.calculate_vertex_ai_batch_cost_and_usage(rows, "gemini-2.5-flash")
assert (result.successful_requests, result.failed_requests) == (1, 1)
assert result.usage.total_tokens == 557
assert len(calls) == 1
@pytest.mark.parametrize("wildcard_model", ["*", "vertex_ai/*"])
def test_native_vertex_rows_under_a_wildcard_deployment_are_priced_by_model_version(monkeypatch, wildcard_model):
calls = _capture_cost_calls(monkeypatch)
rows = [
_native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True, model_version="gemini-2.5-flash"),
_native_vertex_row(UNGROUNDED_USAGE_METADATA, grounded=False, model_version=None),
]
result = bu.calculate_vertex_ai_batch_cost_and_usage(rows, wildcard_model)
assert [call["model"] for call in calls] == ["gemini-2.5-flash", wildcard_model]
assert result.cost == pytest.approx(1.5)
assert (result.successful_requests, result.failed_requests) == (2, 0)
assert result.usage.total_tokens == 557 + 336
def test_native_vertex_row_without_model_version_under_a_wildcard_deployment_bills_its_explicit_prices():
deployment_model_info = {"input_cost_per_token_batches": 1e-6, "output_cost_per_token_batches": 2e-6}
with_version = _native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True, model_version="gemini-2.5-flash")
without_version = _native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True, model_version=None)
twin = bu.calculate_vertex_ai_batch_cost_and_usage([with_version], "vertex_ai/*", model_info=deployment_model_info)
both = bu.calculate_vertex_ai_batch_cost_and_usage(
[with_version, without_version], "vertex_ai/*", model_info=deployment_model_info
)
assert twin.cost > 0
assert both.cost == pytest.approx(2 * twin.cost)
assert (both.successful_requests, both.failed_requests) == (2, 0)
def test_native_vertex_row_the_cost_map_cannot_price_is_billed_at_zero_and_the_rest_still_bills(monkeypatch):
import litellm.cost_calculator as cc
def _calc(**kw):
if kw["model"] == "gemini-unpriced":
raise ValueError("no pricing")
return (0.5, 0.25)
monkeypatch.setattr(cc, "batch_cost_calculator", _calc)
rows = [
_native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True, model_version="gemini-unpriced"),
_native_vertex_row(UNGROUNDED_USAGE_METADATA, grounded=False, model_version="gemini-2.5-flash"),
]
result = bu.calculate_vertex_ai_batch_cost_and_usage(rows)
assert result.cost == pytest.approx(0.75)
assert (result.successful_requests, result.failed_requests) == (2, 0)
assert result.usage.total_tokens == 557 + 336
assert result.models == ["gemini-unpriced", "gemini-2.5-flash"]
@pytest.mark.asyncio
async def test_flag_sends_every_vertex_row_down_the_native_path_when_a_model_is_known(monkeypatch):
monkeypatch.setattr(litellm, "disable_vertex_batch_output_transformation", True, raising=False)
monkeypatch.setattr(
bu, "_aggregate_batch_cost_usage_models", lambda **kw: pytest.fail("generic path should not run")
)
calls = _capture_cost_calls(monkeypatch)
rows = [_vertex_openai_row("request-1", "gemini-2.5-flash", 10, 5)]
result = await bu.calculate_batch_cost_and_usage(
file_content_dictionary=rows, custom_llm_provider="vertex_ai", model_name="gemini-2.5-flash"
)
assert calls == []
assert (result.successful_requests, result.failed_requests) == (0, 1)

View file

@ -20,6 +20,8 @@ from litellm.rust_bridge.chat_completions.entrypoints import (
)
from litellm.rust_bridge.configuration import Rollout
from litellm.types.utils import ModelResponse
from litellm.chat_completions import dispatch
from litellm.rust_bridge.catalog import Rules
MESSAGES: Final = [{"role": "user", "content": "hi"}]
PYTHON_RULES: Final = ()
@ -256,3 +258,100 @@ async def test_public_acompletion_routes_through_dispatch(monkeypatch: pytest.Mo
NATIVE_ACOMPLETION.reset()
assert result is expected
assert [request.model for request in captured] == ["gpt-4o"]
@pytest.mark.asyncio
async def test_public_completion_calls_keep_the_python_result() -> None:
sync_response: Final = litellm.completion(model="openai/test-model", messages=MESSAGES, mock_response="ok")
async_response: Final = await litellm.acompletion(model="openai/test-model", messages=MESSAGES, mock_response="ok")
assert isinstance(sync_response, ModelResponse)
assert isinstance(async_response, ModelResponse)
assert sync_response.choices[0].message.content == "ok"
assert async_response.choices[0].message.content == "ok"
def test_sync_completion_request_projects_public_arguments() -> None:
rules: Final[Rules] = (RouteRule(Route.CHAT_COMPLETIONS, Rollout.RUST_REQUIRED),)
expected: Final = ModelResponse()
def native(
request: LiteLLMChatCompletionsRequest, args: tuple[object, ...], kwargs: Mapping[str, object]
) -> ModelResponse:
assert request.model == "test-model"
assert request.messages == MESSAGES
assert request.custom_llm_provider == "openai"
assert request.stream is True
return expected
binding: Final[NativeBinding[NativeCompletion]] = NativeBinding("completion", validate=lambda _: None)
binding.override(native)
response: Final = dispatch._DISPATCH.run( # pyright: ignore[reportPrivateUsage] # test an explicit route decision
("test-model", MESSAGES),
{"custom_llm_provider": "openai", "stream": True},
python=lambda *args, **kwargs: pytest.fail("required native route must handle this call"),
binding=binding,
native=lambda hook, request, args, kwargs: hook(request, args, kwargs),
rules=rules,
)
assert response is expected
@pytest.mark.asyncio
async def test_async_completion_falls_back_after_native_declines() -> None:
from litellm.rust_bridge.bindings import native_exception_types
native_types: Final = native_exception_types()
if native_types is None:
pytest.skip("native bridge is unavailable")
declined, _ = native_types
expected: Final = ModelResponse()
rules: Final[Rules] = (RouteRule(Route.CHAT_COMPLETIONS, Rollout.RUST_OPT_OUT),)
async def native(
request: LiteLLMChatCompletionsRequest, args: tuple[object, ...], kwargs: Mapping[str, object]
) -> ModelResponse:
raise declined("unsupported")
async def python(*args: object, **kwargs: object) -> ModelResponse:
return expected
binding: Final[NativeBinding[NativeAcompletion]] = NativeBinding("acompletion", validate=lambda _: None)
binding.override(native)
response: Final = await dispatch._ADISPATCH.arun( # pyright: ignore[reportPrivateUsage] # test an explicit route decision
("test-model", MESSAGES),
{},
python=python,
binding=binding,
native=lambda hook, request, args, kwargs: hook(request, args, kwargs),
rules=rules,
)
assert response is expected
def test_internal_acompletion_marker_bypasses_native() -> None:
rules: Final[Rules] = (RouteRule(Route.CHAT_COMPLETIONS, Rollout.RUST_REQUIRED),)
expected: Final = ModelResponse()
def python(*args: object, **kwargs: object) -> ModelResponse:
return expected
def native(
request: LiteLLMChatCompletionsRequest, args: tuple[object, ...], kwargs: Mapping[str, object]
) -> ModelResponse:
pytest.fail("acompletion's inner completion call must stay on Python")
binding: Final[NativeBinding[NativeCompletion]] = NativeBinding("completion", validate=lambda _: None)
binding.override(native)
response: Final = dispatch._DISPATCH.run( # pyright: ignore[reportPrivateUsage] # test an explicit route decision
("test-model", MESSAGES),
{"custom_llm_provider": "openai", "acompletion": True},
python=python,
binding=binding,
native=lambda hook, request, args, kwargs: hook(request, args, kwargs),
rules=rules,
)
assert response is expected

View file

@ -1,7 +1,14 @@
import asyncio
import base64
import importlib
import os
from collections.abc import Iterator
from collections.abc import Coroutine, Iterator
from dataclasses import dataclass, field
from pathlib import Path
from typing import Final
import boto3
import httpx
import pytest
from pytest_socket import enable_socket, socket_allow_hosts
@ -10,6 +17,17 @@ os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
import litellm # noqa: E402 # litellm reads LITELLM_LOCAL_MODEL_COST_MAP at import
import litellm.router as litellm_router_module # noqa: E402 # same import-time dependency
import litellm.utils as litellm_utils_module # noqa: E402 # same import-time dependency
from litellm._logging import ALL_LOGGERS # noqa: E402 # same import-time dependency
from litellm.anthropic_beta_headers_manager import reload_beta_headers_config # noqa: E402 # same import-time dependency
from litellm.litellm_core_utils.prompt_templates import factory as prompt_factory_module # noqa: E402 # same import-time dependency
from litellm.litellm_core_utils.prompt_templates import ( # noqa: E402 # same import-time dependency
image_handling as image_handling_module,
)
from litellm.llms.gemini.chat import transformation as gemini_chat_transformation_module # noqa: E402 # same import-time dependency
from litellm.llms.custom_httpx.async_client_cleanup import ( # noqa: E402 # same import-time dependency
close_litellm_async_clients,
)
from litellm.proxy.db import tool_registry_writer as tool_registry_writer_module # noqa: E402 # same import-time dependency
LOOPBACK_HOSTS: Final = ["127.0.0.1", "::1", "localhost"]
AMBIENT_AZURE_CREDENTIAL_ENV_VARS: Final = (
@ -20,6 +38,66 @@ AMBIENT_AZURE_CREDENTIAL_ENV_VARS: Final = (
"AZURE_USERNAME",
"AZURE_PASSWORD",
)
AMBIENT_AWS_ENV_VARS: Final = (
"AWS_PROFILE",
"AWS_DEFAULT_PROFILE",
"AWS_CONTAINER_CREDENTIALS_FULL_URI",
"AWS_CONTAINER_CREDENTIALS_RELATIVE_URI",
"AWS_SESSION_TOKEN",
"AWS_ROLE_ARN",
"AWS_WEB_IDENTITY_TOKEN_FILE",
"AWS_BEARER_TOKEN_BEDROCK",
"AWS_REGION_NAME",
"AWS_DEFAULT_REGION",
)
MODULES_WITH_AWS_AUTH_HANDLERS: Final = (
"litellm.main",
"litellm.files.main",
"litellm.rerank_api.main",
"litellm.realtime_api.main",
)
CALLBACK_LISTS: Final = (
"callbacks",
"success_callback",
"failure_callback",
"input_callback",
"_async_success_callback",
"_async_failure_callback",
"_async_input_callback",
)
RESET_TO_NONE_GLOBALS: Final = ("model_fallbacks", "cache")
RESTORED_GLOBALS: Final = (
"disable_aiohttp_transport",
"force_ipv4",
"drop_params",
"secret_manager_client",
"_key_management_system",
"_key_management_settings",
"api_base",
"num_retries",
"modify_params",
"ssl_verify",
"credential_list",
"model_group_settings",
"default_internal_user_params",
"default_team_params",
"prometheus_emit_stream_label",
"vector_store_registry",
"model_cost",
"cost_margin_config",
"cost_discount_config",
"disable_hf_tokenizer_download",
"disable_copilot_system_to_assistant",
"cohere_models",
"anthropic_models",
"token_counter",
"initialized_langfuse_clients",
)
MODULE_LEVEL_CLIENTS: Final = ("module_level_client", "module_level_aclient")
SESSION_CLIENTS: Final = ("base_llm_aiohttp_handler", "httpx_client", "aclient", "client")
ONE_PIXEL_PNG: Final = base64.b64decode(
"iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mNkYPhfDwAChwGA60e6kgAAAABJRU5ErkJggg=="
)
def _allow_loopback_only() -> None:
@ -29,11 +107,116 @@ def _allow_loopback_only() -> None:
_allow_loopback_only()
def pytest_collectstart() -> None:
_allow_loopback_only()
@pytest.hookimpl(trylast=True)
def pytest_runtest_setup() -> None:
_allow_loopback_only()
def _run_coroutine_if_needed(result: object) -> None:
if not asyncio.iscoroutine(result):
return
coroutine: Final[Coroutine[object, object, object]] = result
try:
asyncio.run(coroutine)
except RuntimeError:
try:
loop: Final = asyncio.get_running_loop()
except RuntimeError:
coroutine.close()
return
loop.create_task(coroutine)
def _close_handler_if_needed(handler: object) -> None:
close: Final = getattr(handler, "close", None)
if not callable(close):
return
_run_coroutine_if_needed(close())
def _reset_aws_auth_caches() -> None:
modules: Final = tuple(importlib.import_module(name) for name in MODULES_WITH_AWS_AUTH_HANDLERS)
flushes: Final = (
getattr(getattr(getattr(module, attr_name), "iam_cache", None), "flush_cache", None)
for module in modules
for attr_name in dir(module)
)
for flush in filter(callable, flushes):
flush()
boto3.DEFAULT_SESSION = None
def _flush_client_caches() -> None:
litellm.in_memory_llm_clients_cache.flush_cache()
image_handling_module.in_memory_cache.flush_cache()
_reset_aws_auth_caches()
@pytest.fixture(scope="session")
def isolated_aws_config_files(tmp_path_factory: pytest.TempPathFactory) -> tuple[Path, Path]:
aws_dir: Final = tmp_path_factory.mktemp("aws-config")
credentials: Final = aws_dir / "credentials"
config: Final = aws_dir / "config"
credentials.write_text("", encoding="utf-8")
config.write_text("", encoding="utf-8")
return credentials, config
@pytest.fixture(autouse=True)
def isolate_host_environment(isolated_aws_config_files: tuple[Path, Path]) -> Iterator[None]:
credentials, config = isolated_aws_config_files
with pytest.MonkeyPatch.context() as environment:
environment.setenv("AWS_SHARED_CREDENTIALS_FILE", str(credentials))
environment.setenv("AWS_CONFIG_FILE", str(config))
environment.setenv("AWS_EC2_METADATA_DISABLED", "true")
for name in AMBIENT_AWS_ENV_VARS:
environment.delenv(name, raising=False)
environment.delenv("PROXY_BASE_URL", raising=False)
environment.setenv("LITELLM_CLI_DISABLE_KEYRING", "1")
yield
@pytest.fixture(autouse=True)
def isolate_litellm_globals() -> Iterator[None]:
original_callbacks: Final = {name: list(getattr(litellm, name) or []) for name in CALLBACK_LISTS}
original_reset: Final = {name: getattr(litellm, name) for name in RESET_TO_NONE_GLOBALS}
original_restored: Final = {name: getattr(litellm, name) for name in RESTORED_GLOBALS if hasattr(litellm, name)}
original_clients: Final = {name: litellm.__dict__[name] for name in MODULE_LEVEL_CLIENTS if name in litellm.__dict__}
original_loggers: Final = {
logger: (logger.level, logger.disabled, logger.propagate, list(logger.handlers), list(logger.filters))
for logger in ALL_LOGGERS
}
original_tool_policy_registry: Final = tool_registry_writer_module._tool_policy_registry
_flush_client_caches()
for name in CALLBACK_LISTS:
setattr(litellm, name, [])
for name in RESET_TO_NONE_GLOBALS:
setattr(litellm, name, None)
for name in MODULE_LEVEL_CLIENTS:
litellm.__dict__.pop(name, None)
tool_registry_writer_module._tool_policy_registry = None
yield
_flush_client_caches()
leaked_clients: Final = tuple(litellm.__dict__.pop(name, None) for name in MODULE_LEVEL_CLIENTS)
for name, client in zip(MODULE_LEVEL_CLIENTS, leaked_clients):
if client is not original_clients.get(name):
_close_handler_if_needed(client)
litellm.__dict__.update(original_clients)
for name, value in (original_callbacks | original_reset | original_restored).items():
setattr(litellm, name, value)
for logger, (level, disabled, propagate, handlers, filters) in original_loggers.items():
logger.setLevel(level)
logger.disabled = disabled
logger.propagate = propagate
logger.handlers = handlers
logger.filters = filters
tool_registry_writer_module._tool_policy_registry = original_tool_policy_registry
@pytest.fixture(autouse=True)
def isolate_router_model_cost_state() -> Iterator[None]:
original_live_routers: Final = frozenset(litellm_router_module._live_routers)
@ -41,6 +224,7 @@ def isolate_router_model_cost_state() -> Iterator[None]:
model_key: dict(model_value)
for model_key, model_value in litellm_utils_module._runtime_registered_model_cost.items()
}
litellm_utils_module._invalidate_model_cost_lowercase_map()
yield
for router in tuple(litellm_router_module._live_routers):
litellm_router_module._live_routers.discard(router)
@ -61,6 +245,47 @@ def local_model_cost_map(monkeypatch: pytest.MonkeyPatch) -> Iterator[None]:
litellm.get_model_info.cache_clear()
@pytest.fixture
def local_beta_headers_config(monkeypatch: pytest.MonkeyPatch) -> Iterator[None]:
monkeypatch.setenv("LITELLM_LOCAL_ANTHROPIC_BETA_HEADERS", "True")
reload_beta_headers_config()
yield
monkeypatch.delenv("LITELLM_LOCAL_ANTHROPIC_BETA_HEADERS", raising=False)
reload_beta_headers_config()
@dataclass(slots=True)
class AsyncOnlyImageFetch:
fetched: list[str] = field(default_factory=list) # mutable-ok: tests assert on the URLs fetched, in order
base64_png: str = base64.b64encode(ONE_PIXEL_PNG).decode()
data_url: str = "data:image/png;base64," + base64.b64encode(ONE_PIXEL_PNG).decode()
@pytest.fixture
def async_only_image_fetch(monkeypatch: pytest.MonkeyPatch) -> AsyncOnlyImageFetch:
fetch: Final = AsyncOnlyImageFetch()
def forbid_sync_fetch(client: object, url: str, **kwargs: object) -> httpx.Response:
raise litellm.ImageFetchError(f"sync image fetch ran on the event loop: {url}")
async def serve_png(client: object, url: str, **kwargs: object) -> httpx.Response:
fetch.fetched.append(url)
return httpx.Response(
200, content=ONE_PIXEL_PNG, headers={"content-type": "image/png"}, request=httpx.Request("GET", url)
)
def forbid_sync_convert(url: str, *args: object, **kwargs: object) -> str:
if url.startswith(("http://", "https://")):
raise litellm.ImageFetchError(f"sync convert_url_to_base64 ran on the request path: {url}")
return url
monkeypatch.setattr(image_handling_module, "safe_get", forbid_sync_fetch)
monkeypatch.setattr(image_handling_module, "async_safe_get", serve_png)
for module in (image_handling_module, prompt_factory_module, gemini_chat_transformation_module):
monkeypatch.setattr(module, "convert_url_to_base64", forbid_sync_convert)
return fetch
@pytest.fixture
def no_ambient_azure_credentials(monkeypatch: pytest.MonkeyPatch) -> None:
for name in AMBIENT_AZURE_CREDENTIAL_ENV_VARS:
@ -68,4 +293,9 @@ def no_ambient_azure_credentials(monkeypatch: pytest.MonkeyPatch) -> None:
def pytest_sessionfinish() -> None:
for name in MODULE_LEVEL_CLIENTS:
_close_handler_if_needed(litellm.__dict__.pop(name, None))
for name in SESSION_CLIENTS:
_close_handler_if_needed(getattr(litellm, name, None))
_run_coroutine_if_needed(close_litellm_async_clients())
enable_socket()

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