mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-02 02:11:58 +00:00
Merge remote-tracking branch 'origin/main' into litellm_agent365_fail_open_default
This commit is contained in:
commit
47e6df7b76
732 changed files with 11540 additions and 11194 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
10
.github/merge-smoke-tests.json
vendored
10
.github/merge-smoke-tests.json
vendored
|
|
@ -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",
|
||||
|
|
|
|||
18
.github/workflows/_test-unit-base.yml
vendored
18
.github/workflows/_test-unit-base.yml
vendored
|
|
@ -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=()
|
||||
|
|
|
|||
4
.github/workflows/test-redis-compat.yml
vendored
4
.github/workflows/test-redis-compat.yml
vendored
|
|
@ -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 \
|
||||
|
|
|
|||
33
.github/workflows/test-unit-proxy-db.yml
vendored
33
.github/workflows/test-unit-proxy-db.yml
vendored
|
|
@ -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 }}
|
||||
|
|
|
|||
42
.github/workflows/test-unit.yml
vendored
42
.github/workflows/test-unit.yml
vendored
|
|
@ -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 }}
|
||||
|
|
|
|||
6
Makefile
6
Makefile
|
|
@ -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
|
||||
|
|
|
|||
16
litellm-rust/Cargo.lock
generated
16
litellm-rust/Cargo.lock
generated
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
},
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)?)
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
}
|
||||
|
|
|
|||
21
litellm-rust/crates/framer/src/framed.rs
Normal file
21
litellm-rust/crates/framer/src/framed.rs
Normal 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()
|
||||
}
|
||||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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(())
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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());
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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());
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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());
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
})
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
||||
|
|
|
|||
|
|
@ -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__":
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
@ -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
|
||||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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]}]
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"},
|
||||
]
|
||||
|
|
@ -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
|
||||
|
|
@ -1 +0,0 @@
|
|||
"""Tests for Gemini files functionality"""
|
||||
|
|
@ -1 +0,0 @@
|
|||
# Gemini Video Generation Tests
|
||||
|
|
@ -1 +0,0 @@
|
|||
# Manus provider tests
|
||||
|
|
@ -1 +0,0 @@
|
|||
# Manus Responses API tests
|
||||
|
|
@ -1 +0,0 @@
|
|||
# MiniMax tests
|
||||
|
|
@ -1 +0,0 @@
|
|||
# MiniMax chat tests
|
||||
|
|
@ -1 +0,0 @@
|
|||
# MiniMax messages tests
|
||||
|
|
@ -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 == ""
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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"
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -1 +0,0 @@
|
|||
|
||||
|
|
@ -1 +0,0 @@
|
|||
# S3 Vectors tests
|
||||
|
|
@ -1 +0,0 @@
|
|||
# S3 Vectors vector store tests
|
||||
|
|
@ -1 +0,0 @@
|
|||
"""Soniox provider tests."""
|
||||
File diff suppressed because it is too large
Load diff
|
|
@ -1 +0,0 @@
|
|||
# Vertex AI Image Edit Tests
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -1 +0,0 @@
|
|||
"""Tests for Vertex AI Gemma-AI models"""
|
||||
|
|
@ -1,3 +0,0 @@
|
|||
"""
|
||||
Tests for Vertex AI video generation.
|
||||
"""
|
||||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"}]}}
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
),
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:")
|
||||
|
|
|
|||
|
|
@ -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
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
Loading…
Add table
Reference in a new issue