Merge branch 'main' into litellm_lit8140_azure_flux2_megapixel_billing

This commit is contained in:
Shreshth Kharbanda 2026-09-25 14:06:08 -07:00
commit 9a7f0bafc8
968 changed files with 14815 additions and 11439 deletions

View file

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

View file

@ -7,7 +7,11 @@ legacy_flags=(
caching-local
enterprise-package
enterprise-routing
integrations
llm-other-providers
llm-vertex-ai
mcp-integration
misc
proxy-db-auth-checks
proxy-db-budgets
proxy-db-custom-logging
@ -22,6 +26,7 @@ legacy_flags=(
proxy-db-proxy-utils
proxy-extras
proxy-infra
responses-caching-types
)
legacy_paths() {
@ -36,6 +41,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 +53,33 @@ 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 ;;
integrations) echo tests/unit/integrations ;;
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/secret_managers
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 +142,7 @@ legacy_paths() {
proxy-db-proxy-utils) echo tests/unit/proxy/test_proxy_utils.py ;;
proxy-extras) echo tests/unit/litellm_proxy_extras ;;
proxy-infra) echo tests/unit/gateway ;;
responses-caching-types) echo tests/unit/types ;;
*) echo "unit_selection.sh: unknown flag $1" >&2; exit 1 ;;
esac
}

View file

@ -341,6 +341,7 @@ workflows:
flag:
- enterprise-package
- proxy-infra
- responses-caching-types
- proxy-db-auth-checks
- proxy-db-jwt-and-keys
- proxy-db-proxy-server-core
@ -353,6 +354,35 @@ 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-integrations
flag: integrations
shards: 2
reruns: 3
base_ref: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.base.ref or "" >>
pull_request_url: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.url or "" >>
- unit:
name: unit-misc
flag: misc
shards: 2
reruns: 2
base_ref: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.base.ref or "" >>
pull_request_url: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.url or "" >>
- unit:
name: unit-proxy-db-proxy-utils
flag: proxy-db-proxy-utils

View file

@ -1,12 +1,12 @@
{
"cases": {
"CHAT-JSON": "tests/test_litellm/llms/openai/test_openai.py::test_acompletion_returns_json_reply_over_injected_transport",
"CHAT-TEXT-STREAM": "tests/test_litellm/llms/openai/test_openai.py::test_acompletion_streams_text_deltas_over_injected_transport",
"CHAT-TOOL-STREAM": "tests/test_litellm/llms/openai/test_openai.py::test_acompletion_streams_tool_call_arguments_over_injected_transport",
"CHAT-JSON": "tests/unit/llms/openai/test_openai.py::test_acompletion_returns_json_reply_over_injected_transport",
"CHAT-TEXT-STREAM": "tests/unit/llms/openai/test_openai.py::test_acompletion_streams_text_deltas_over_injected_transport",
"CHAT-TOOL-STREAM": "tests/unit/llms/openai/test_openai.py::test_acompletion_streams_tool_call_arguments_over_injected_transport",
"MODEL-ALLOW": "tests/test_litellm/proxy/auth/test_auth_checks.py::test_can_object_call_model_allows_listed_model_for_key",
"MODEL-DENY": "tests/test_litellm/proxy/auth/test_auth_checks.py::test_can_object_call_model_denials_return_forbidden[key-key_model_access_denied]",
"COST-EXPLICIT": "tests/test_litellm/test_cost_calculator.py::test_completion_cost_charges_explicit_per_token_rates_over_registered_ones",
"COST-ZERO": "tests/test_litellm/test_cost_calculator.py::test_completion_cost_is_zero_when_explicit_rates_are_zero",
"COST-EXPLICIT": "tests/unit/test_cost_calculator.py::test_completion_cost_charges_explicit_per_token_rates_over_registered_ones",
"COST-ZERO": "tests/unit/test_cost_calculator.py::test_completion_cost_is_zero_when_explicit_rates_are_zero",
"LOG-CONTENT-ON": "tests/test_litellm/litellm_core_utils/test_litellm_logging.py::test_standard_logging_payload_keeps_message_content_when_message_logging_is_on",
"LOG-CONTENT-OFF": "tests/test_litellm/litellm_core_utils/test_litellm_logging.py::test_standard_logging_payload_redacts_message_content_when_message_logging_is_off",
"CALLBACK-SUCCESS": "tests/test_litellm/litellm_core_utils/test_litellm_logging.py::test_async_success_handler_delivers_standard_logging_payload_to_custom_logger",

View file

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

View file

@ -10,7 +10,7 @@ on:
- "litellm/_redis_credential_provider.py"
- "litellm/caching/redis_cache.py"
- "litellm/caching/evicted_client_closer.py"
- "tests/test_litellm/test_redis.py"
- "tests/unit/test_redis.py"
- "tests/local_testing/test_caching.py"
- "tests/test_litellm/caching/test_redis_connection_pool.py"
- "tests/test_litellm/caching/test_redis_cluster_cache.py"
@ -84,7 +84,7 @@ jobs:
run: |
redis-server --version
uv run --no-sync pytest \
tests/test_litellm/test_redis.py \
tests/unit/test_redis.py \
tests/test_litellm/caching/test_redis_connection_pool.py \
tests/test_litellm/caching/test_redis_cluster_cache.py \
tests/test_litellm/caching/test_evicted_client_closer.py \

View file

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

View file

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

View file

@ -314,7 +314,7 @@ test-unit: install-test-deps
# Matrix test targets (matching CI workflow groups)
test-unit-llms: install-test-deps
$(UV_RUN) pytest tests/test_litellm/llms --tb=short -vv -n 4 --durations=20
$(UV_RUN) pytest tests/unit/llms --tb=short -vv -n 4 --durations=20
test-unit-proxy-guardrails: install-test-deps
$(UV_RUN) pytest tests/test_litellm/proxy/guardrails tests/test_litellm/proxy/management_endpoints tests/test_litellm/proxy/management_helpers --tb=short -vv -n 4 --durations=20
@ -326,16 +326,16 @@ test-unit-proxy-misc: install-test-deps
$(UV_RUN) pytest tests/test_litellm/proxy/_experimental tests/test_litellm/proxy/agent_endpoints tests/test_litellm/proxy/anthropic_endpoints tests/test_litellm/proxy/common_utils tests/test_litellm/proxy/discovery_endpoints tests/test_litellm/proxy/experimental tests/test_litellm/proxy/google_endpoints tests/test_litellm/proxy/health_endpoints tests/test_litellm/proxy/image_endpoints tests/test_litellm/proxy/middleware tests/test_litellm/proxy/openai_files_endpoint tests/test_litellm/proxy/pass_through_endpoints tests/test_litellm/proxy/prompts tests/test_litellm/proxy/public_endpoints tests/test_litellm/proxy/response_api_endpoints tests/test_litellm/proxy/shutdown tests/test_litellm/proxy/spend_tracking tests/test_litellm/proxy/ui_crud_endpoints tests/test_litellm/proxy/vector_store_endpoints tests/test_litellm/proxy/test_*.py --tb=short -vv -n 4 --durations=20
test-unit-integrations: install-test-deps
$(UV_RUN) pytest tests/test_litellm/integrations --tb=short -vv -n 4 --durations=20
$(UV_RUN) pytest tests/unit/integrations --tb=short -vv -n 4 --durations=20
test-unit-core-utils: install-test-deps
$(UV_RUN) pytest tests/test_litellm/litellm_core_utils --tb=short -vv -n 2 --durations=20
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/unit/secret_managers tests/unit/vector_stores tests/unit/a2a_protocol tests/test_litellm/anthropic_interface tests/unit/completion_extras tests/unit/containers tests/unit/enterprise tests/unit/experimental_mcp_client tests/unit/google_genai tests/unit/images tests/unit/interactions tests/test_litellm/interactions tests/test_litellm/passthrough tests/test_litellm/router_strategy tests/test_litellm/router_utils tests/unit/types --tb=short -vv -n 4 --durations=20
test-unit-root: install-test-deps
$(UV_RUN) pytest tests/test_litellm/test_*.py --tb=short -vv -n 4 --durations=20
$(UV_RUN) pytest tests/unit/test_*.py tests/test_litellm/test_*.py --tb=short -vv -n 4 --durations=20
# Proxy unit tests (tests/unit/proxy split alphabetically)
test-proxy-unit-a: install-test-deps

View file

@ -2,7 +2,7 @@
Every `litellm_*` metric family the proxy can expose on `/metrics` (136 families across 97 panels), grouped into rows: proxy traffic, latency, spend and tokens, cache, LLM API deployments, key and team rate limits, budgets, guardrails, MCP, managed files and batches, users and teams, the Redis circuit breaker, the spend log cleanup job, and the `prometheus_system` service callback metrics (per-service latency, request and failure rates, spend update queue sizes). Panel titles are the metric names so you can grep the JSON for the metric you care about
Import `grafana_dashboard.json` from **Dashboards > New > Import** and pick your Prometheus data source when prompted (the `DS_PROMETHEUS` variable). Counters are plotted as `rate()` over `$__rate_interval`, histograms as p50 / p95 / p99, gauges as the raw value grouped by the most useful label. Every query names the metric exactly as the proxy emits it (counters carry the `_total` suffix the Prometheus client adds), and `tests/test_litellm/integrations/test_prometheus_metric_name_consistency.py` fails if a metric is renamed without updating this dashboard
Import `grafana_dashboard.json` from **Dashboards > New > Import** and pick your Prometheus data source when prompted (the `DS_PROMETHEUS` variable). Counters are plotted as `rate()` over `$__rate_interval`, histograms as p50 / p95 / p99, gauges as the raw value grouped by the most useful label. Every query names the metric exactly as the proxy emits it (counters carry the `_total` suffix the Prometheus client adds), and `tests/unit/integrations/test_prometheus_metric_name_consistency.py` fails if a metric is renamed without updating this dashboard
The first eleven rows need only `callbacks: ["prometheus"]`. The last three rows and the `litellm_admission_*` panels are emitted by other subsystems and stay empty until those are on: the service callback row needs `service_callback: ["prometheus_system"]` in `litellm_settings`, the circuit breaker row needs a Redis cache, the cleanup row needs spend log retention, and admission control needs its middleware enabled. Within the base rows, many panels only fill in once the matching feature is in use: budgets need keys, teams, users or orgs with `max_budget` set, cache panels need caching on, guardrail and MCP panels need those features configured, deployment health needs the router with more than one deployment or a failure to record, and `litellm_in_flight_requests` needs traffic at scrape time. An empty panel for a feature you do not use is expected

168
litellm-rust/Cargo.lock generated
View file

@ -73,6 +73,15 @@ version = "1.0.104"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "330a5ed07fa54e4702c9d6c4174f74427fc0ef6e214bbd677ae50a5099946470"
[[package]]
name = "arbitrary"
version = "1.4.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "c3d036a3c4ab069c7b410a2ce876bd74808d2d0888a82667669f8e783a898bf1"
dependencies = [
"derive_arbitrary",
]
[[package]]
name = "arc-swap"
version = "1.9.2"
@ -1470,6 +1479,17 @@ dependencies = [
"serde_core",
]
[[package]]
name = "derive_arbitrary"
version = "1.4.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "1e567bd82dcff979e4b03460c307b3cdc9e96fde3d73bed1496d2bc75d9dd62a"
dependencies = [
"proc-macro2",
"quote",
"syn 2.0.119",
]
[[package]]
name = "derive_builder"
version = "0.20.2"
@ -1643,6 +1663,16 @@ version = "2.5.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "da7c62ceae207dd37ea5b845da6a0696c799f85e97da1ab5b7910be3c1c80223"
[[package]]
name = "filetime"
version = "0.2.29"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "5c287a33c7f0a620c38e641e7f60827713987b3c0f26e8ddc9462cc69cf75759"
dependencies = [
"cfg-if",
"libc",
]
[[package]]
name = "find-msvc-tools"
version = "0.1.9"
@ -3136,10 +3166,11 @@ dependencies = [
"aws-smithy-types",
"bytes",
"futures-util",
"proptest",
"rstest",
"sse-stream",
"thiserror 2.0.19",
"tokio",
"tokio-util",
]
[[package]]
@ -3453,6 +3484,27 @@ dependencies = [
"veil",
]
[[package]]
name = "litellm-testkit"
version = "0.1.0"
dependencies = [
"flate2",
"futures-util",
"reqwest 0.12.28",
"rstest",
"semver",
"serde",
"serde_json",
"sha2 0.10.9",
"tar",
"target-lexicon",
"tempfile",
"thiserror 2.0.19",
"tokio",
"toml",
"zip",
]
[[package]]
name = "litellm-token-counter"
version = "0.1.0"
@ -5206,6 +5258,15 @@ dependencies = [
"serde_core",
]
[[package]]
name = "serde_spanned"
version = "1.1.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "6662b5879511e06e8999a8a235d848113e942c9124f211511b16466ee2995f26"
dependencies = [
"serde_core",
]
[[package]]
name = "serde_urlencoded"
version = "0.7.1"
@ -5408,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"
@ -5537,6 +5585,17 @@ version = "0.2.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "7b2093cf4c8eb1e67749a6762251bc9cd836b6fc171623bd0a9d324d37af2417"
[[package]]
name = "tar"
version = "0.4.46"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "3f6221d9a6003c78398e3b239969f352578258df48c8eb051caadae0015bc840"
dependencies = [
"filetime",
"libc",
"xattr",
]
[[package]]
name = "target-lexicon"
version = "0.13.5"
@ -5806,6 +5865,30 @@ dependencies = [
"tokio",
]
[[package]]
name = "toml"
version = "0.9.12+spec-1.1.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "cf92845e79fc2e2def6a5d828f0801e29a2f8acc037becc5ab08595c7d5e9863"
dependencies = [
"indexmap 2.14.0",
"serde_core",
"serde_spanned",
"toml_datetime 0.7.5+spec-1.1.0",
"toml_parser",
"toml_writer",
"winnow 0.7.15",
]
[[package]]
name = "toml_datetime"
version = "0.7.5+spec-1.1.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "92e1cfed4a3038bc5a127e35a2d360f145e1f4b971b551a2ba5fd7aedf7e1347"
dependencies = [
"serde_core",
]
[[package]]
name = "toml_datetime"
version = "1.1.1+spec-1.1.0"
@ -5822,9 +5905,9 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "6975367e4d2ef766d86af01ffad14b622fecc8d4357a998fbc4deb6e9bacaf9b"
dependencies = [
"indexmap 2.14.0",
"toml_datetime",
"toml_datetime 1.1.1+spec-1.1.0",
"toml_parser",
"winnow",
"winnow 1.0.4",
]
[[package]]
@ -5833,9 +5916,15 @@ version = "1.1.3+spec-1.1.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "1d38ac1cf9b95face32296c0a3ede1fdc270627c9d9c02a7274dd6d960dc4d56"
dependencies = [
"winnow",
"winnow 1.0.4",
]
[[package]]
name = "toml_writer"
version = "1.1.2+spec-1.1.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "7d56353a2a665ad0f41a421187180aab746c8c325620617ad883a99a1cbe66d2"
[[package]]
name = "tonic"
version = "0.14.6"
@ -6597,6 +6686,12 @@ version = "0.52.6"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "589f6da84c646204747d1270a2a5661ea66ed1cced2631d546fdfb155959f9ec"
[[package]]
name = "winnow"
version = "0.7.15"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "df79d97927682d2fd8adb29682d1140b343be4ac0f08fd68b7765d9c059d3945"
[[package]]
name = "winnow"
version = "1.0.4"
@ -6659,6 +6754,16 @@ dependencies = [
"time",
]
[[package]]
name = "xattr"
version = "1.6.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "32e45ad4206f6d2479085147f02bc2ef834ac85886624a23575ae137c8aa8156"
dependencies = [
"libc",
"rustix",
]
[[package]]
name = "xmlparser"
version = "0.13.6"
@ -6784,6 +6889,23 @@ dependencies = [
"syn 2.0.119",
]
[[package]]
name = "zip"
version = "2.4.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "fabe6324e908f85a1c52063ce7aa26b68dcb7eb6dbc83a2d148403c9bc3eba50"
dependencies = [
"arbitrary",
"crc32fast",
"crossbeam-utils",
"displaydoc",
"flate2",
"indexmap 2.14.0",
"memchr",
"thiserror 2.0.19",
"zopfli",
]
[[package]]
name = "zlib-rs"
version = "0.6.7"
@ -6795,3 +6917,15 @@ name = "zmij"
version = "1.0.23"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "29666d0abbfad1e3dc4dcf6144730dd3a3ab225bbbdac83319345b1b44ccfc1b"
[[package]]
name = "zopfli"
version = "0.8.3"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "f05cd8797d63865425ff89b5c4a48804f35ba0ce8d125800027ad6017d2b5249"
dependencies = [
"bumpalo",
"crc32fast",
"log",
"simd-adler32",
]

View file

@ -81,6 +81,12 @@ tokio = { version = "1", features = ["rt-multi-thread", "macros", "time", "net"]
tokio-tungstenite = { version = "0.24", default-features = false, features = ["connect", "rustls-tls-native-roots"] }
futures-util = { version = "0.3", default-features = false, features = ["sink", "std"] }
base64 = "0.22"
flate2 = "1"
semver = "1"
tar = "0.4"
target-lexicon = "0.13.5"
tempfile = "3"
zip = { version = "2", default-features = false, features = ["deflate"] }
moka = { version = "0.12.16", features = ["future"] }
strum = { version = "0.28.0", features = ["derive"] }
url = "2.5.8"

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -118,7 +118,7 @@ pub(super) struct ParityCase {
pub(super) fn parity_cases() -> Vec<ParityCase> {
serde_json::from_str(include_str!(concat!(
env!("CARGO_MANIFEST_DIR"),
"/../../../tests/test_litellm/secret_managers/hashicorp_vault_parity.json"
"/../../../tests/unit/secret_managers/hashicorp_vault_parity.json"
)))
.unwrap()
}

View file

@ -32,7 +32,7 @@ The rollout decision comes from [catalog.py](../../../litellm/rust_bridge/catalo
`_SecretManagerRuntime` is a private implementation detail, not a replacement SDK class. Its async methods return Futures; public `async def` methods retain lazy coroutine creation and `asyncio.create_task` support. Passing the same names and arguments is insufficient to claim parity until the remaining return-value, error, cache and configuration differences above are closed
## [tests/test_litellm/secret_managers/test_aws_secret_manager_replication.py](../../../tests/test_litellm/secret_managers/test_aws_secret_manager_replication.py)
## [tests/unit/secret_managers/test_aws_secret_manager_replication.py](../../../tests/unit/secret_managers/test_aws_secret_manager_replication.py)
| Python test | Rust coverage or boundary |
| --- | --- |
@ -48,7 +48,7 @@ The rollout decision comes from [catalog.py](../../../litellm/rust_bridge/catalo
| `test_replicate_secret_http_error_raises` | [direct_replication_returns_response_or_service_error](../secrets-aws/tests/secret_manager/writes.rs) |
| `test_replicate_secret_timeout_raises` | [write_and_replication_timeouts_remain_errors](../secrets-aws/tests/secret_manager/writes.rs) |
## [tests/test_litellm/secret_managers/test_aws_secret_manager_rotation.py](../../../tests/test_litellm/secret_managers/test_aws_secret_manager_rotation.py)
## [tests/unit/secret_managers/test_aws_secret_manager_rotation.py](../../../tests/unit/secret_managers/test_aws_secret_manager_rotation.py)
| Python test | Rust coverage or boundary |
| --- | --- |
@ -59,7 +59,7 @@ The rollout decision comes from [catalog.py](../../../litellm/rust_bridge/catalo
| `test_write_secret_to_name_inside_recovery_window_restores_and_stores_new_value` | [recovery_window_alias_is_restored_updated_and_tagged](../secrets-aws/tests/secret_manager/writes.rs) |
| `test_write_secret_to_live_existing_name_still_fails_without_overwriting` | [create_failure_does_not_overwrite_an_alias_without_a_deletion_date](../secrets-aws/tests/secret_manager/writes.rs) |
## [tests/test_litellm/secret_managers/test_aws_secret_manager_v2.py](../../../tests/test_litellm/secret_managers/test_aws_secret_manager_v2.py)
## [tests/unit/secret_managers/test_aws_secret_manager_v2.py](../../../tests/unit/secret_managers/test_aws_secret_manager_v2.py)
| Python test | Rust coverage or boundary |
| --- | --- |
@ -70,14 +70,14 @@ The rollout decision comes from [catalog.py](../../../litellm/rust_bridge/catalo
| `test_prepare_request_explicit_bedrock_runtime_endpoint_param_still_wins` | [endpoint_overrides_replace_the_service_and_override_the_region](../secrets-aws/tests/secret_manager/configuration.rs) |
| `test_prepare_request_env_bedrock_runtime_endpoint_still_wins` | [endpoint_overrides_replace_the_service_and_override_the_region](../secrets-aws/tests/secret_manager/configuration.rs) |
## [tests/test_litellm/secret_managers/test_base_secret_manager.py](../../../tests/test_litellm/secret_managers/test_base_secret_manager.py)
## [tests/unit/secret_managers/test_base_secret_manager.py](../../../tests/unit/secret_managers/test_base_secret_manager.py)
| Python test | Rust coverage or boundary |
| --- | --- |
| `test_raise_if_unsafe_secret_name_rejects_traversal_and_line_breaks` | [names_reject_path_traversal_and_control_characters](../secrets-types/tests/rotation.rs) |
| `test_raise_if_unsafe_secret_name_allows_legitimate_aliases` | [names_allow_safe_values](../secrets-types/tests/rotation.rs) |
## [tests/test_litellm/secret_managers/test_custom_secret_manager.py](../../../tests/test_litellm/secret_managers/test_custom_secret_manager.py)
## [tests/unit/secret_managers/test_custom_secret_manager.py](../../../tests/unit/secret_managers/test_custom_secret_manager.py)
| Python test | Rust coverage or boundary |
| --- | --- |
@ -89,7 +89,7 @@ The rollout decision comes from [catalog.py](../../../litellm/rust_bridge/catalo
| `test_custom_secret_manager_integration_with_litellm` | [manager_strings_are_coerced_like_literal_eval](../secrets/tests/resolution.rs) |
| `test_minimal_custom_secret_manager` | Exercises the Python example subclass itself or Python default methods. Caller-authored Python implementations remain Python callbacks; resolver integration is covered by `manager_strings_are_coerced_like_literal_eval` |
## [tests/test_litellm/secret_managers/test_cyberark_secret_manager.py](../../../tests/test_litellm/secret_managers/test_cyberark_secret_manager.py)
## [tests/unit/secret_managers/test_cyberark_secret_manager.py](../../../tests/unit/secret_managers/test_cyberark_secret_manager.py)
| Python test | Rust coverage or boundary |
| --- | --- |
@ -97,7 +97,7 @@ The rollout decision comes from [catalog.py](../../../litellm/rust_bridge/catalo
| `test_async_write_matches_parity_fixture` | [writes_match_python_parity_fixture](../secrets-cyberark/tests/secret_manager/writes.rs) |
| `test_missing_credentials_raise_value_error` | [new_validates_credentials_before_license_and_configuration](../secrets-cyberark/tests/secret_manager/configuration.rs) |
## [tests/test_litellm/secret_managers/test_get_azure_ad_token_provider.py](../../../tests/test_litellm/secret_managers/test_get_azure_ad_token_provider.py)
## [tests/unit/secret_managers/test_get_azure_ad_token_provider.py](../../../tests/unit/secret_managers/test_get_azure_ad_token_provider.py)
| Python test | Rust coverage or boundary |
| --- | --- |
@ -115,7 +115,7 @@ The rollout decision comes from [catalog.py](../../../litellm/rust_bridge/catalo
| `test_get_azure_ad_token_provider_prefers_workload_identity_over_managed_identity` | Credential-selection contract belongs to `litellm-auth-azure`, not a secrets crate |
| `test_get_azure_ad_token_provider_defaults_to_default_azure_credential` | Credential-selection contract belongs to `litellm-auth-azure`, not a secrets crate |
## [tests/test_litellm/secret_managers/test_hashicorp_secret_manager.py](../../../tests/test_litellm/secret_managers/test_hashicorp_secret_manager.py)
## [tests/unit/secret_managers/test_hashicorp_secret_manager.py](../../../tests/unit/secret_managers/test_hashicorp_secret_manager.py)
| Python test | Rust coverage or boundary |
| --- | --- |
@ -130,13 +130,13 @@ The rollout decision comes from [catalog.py](../../../litellm/rust_bridge/catalo
| `test_tls_login_uses_login_namespace` | [tls_login_posts_the_role_and_uses_the_client_identity](../secrets-hashicorp/tests/secret_manager/configuration.rs) |
| `test_configuration_matches_native_parity_fixture` | [configuration_matches_python_parity_fixture](../secrets-hashicorp/tests/secret_manager/configuration.rs) |
## [tests/test_litellm/secret_managers/test_secret_manager_handler.py](../../../tests/test_litellm/secret_managers/test_secret_manager_handler.py)
## [tests/unit/secret_managers/test_secret_manager_handler.py](../../../tests/unit/secret_managers/test_secret_manager_handler.py)
| Python test | Rust coverage or boundary |
| --- | --- |
| `test_azure_key_vault_matches_rust_parity_fixture` | [parity_fixture_matches_python_backend_contract](../secrets-azure/tests/key_vault.rs) |
## [tests/test_litellm/secret_managers/test_secret_managers_main.py](../../../tests/test_litellm/secret_managers/test_secret_managers_main.py)
## [tests/unit/secret_managers/test_secret_managers_main.py](../../../tests/unit/secret_managers/test_secret_managers_main.py)
| Python test | Rust coverage or boundary |
| --- | --- |

View file

@ -0,0 +1,32 @@
[package]
name = "litellm-testkit"
version = "0.1.0"
edition.workspace = true
license.workspace = true
repository.workspace = true
publish = false
[dependencies]
flate2.workspace = true
reqwest.workspace = true
serde.workspace = true
semver.workspace = true
serde_json.workspace = true
sha2.workspace = true
tar.workspace = true
target-lexicon.workspace = true
thiserror.workspace = true
tokio = { workspace = true, features = ["fs", "process"] }
zip.workspace = true
[dev-dependencies]
flate2.workspace = true
rstest.workspace = true
sha2.workspace = true
tar.workspace = true
target-lexicon.workspace = true
futures-util.workspace = true
tempfile.workspace = true
tokio.workspace = true
toml = "0.9"
zip.workspace = true

View file

@ -0,0 +1,181 @@
use std::collections::BTreeMap;
use std::path::Path;
use semver::Version;
use serde::Deserialize;
use super::{
Configure, Drive, Install, LaunchSpec, Outcome, Prompt, Settings, Usage, Wire, env, json_lines,
path_string,
};
use crate::install::release::parse;
use crate::install::{Packaging, Release};
use crate::{Error, Fetch, Target};
const RELEASES: &str = "https://downloads.claude.ai/claude-code-releases";
pub struct ClaudeCode;
#[derive(Deserialize)]
struct Manifest {
platforms: BTreeMap<String, Platform>,
}
#[derive(Deserialize)]
struct Platform {
checksum: String,
}
impl Install for ClaudeCode {
fn binary(&self) -> &'static str {
"claude"
}
async fn release(
&self,
fetch: &impl Fetch,
version: &Version,
target: Target,
) -> Result<Release, Error> {
let manifest_url = format!("{RELEASES}/{version}/manifest.json");
let manifest: Manifest = parse(&manifest_url, &fetch.get(&manifest_url).await?)?;
let key = format!(
"{}-{}{}",
target.os_name(),
target.arch_name(),
target.musl_suffix()
);
let platform = manifest
.platforms
.get(&key)
.ok_or_else(|| Error::AssetNotFound(key.clone()))?;
Ok(Release {
url: format!("{RELEASES}/{version}/{key}/claude"),
asset: key,
sha256: platform.checksum.clone(),
packaging: Packaging::Bare,
})
}
}
impl Configure for ClaudeCode {
fn configure(
&self,
_version: &Version,
settings: &Settings,
home: &Path,
) -> Result<LaunchSpec, Error> {
if settings.wire != Wire::Messages {
return Err(Error::UnsupportedWire {
agent: "claude",
wire: settings.wire,
});
}
Ok(LaunchSpec {
env: env([
("HOME", path_string(home)),
("CLAUDE_CONFIG_DIR", path_string(&home.join(".claude"))),
("ANTHROPIC_BASE_URL", settings.base_url.clone()),
("ANTHROPIC_AUTH_TOKEN", settings.api_key.clone()),
("ANTHROPIC_MODEL", settings.model.clone()),
("DISABLE_AUTOUPDATER", "1".to_owned()),
]),
files: BTreeMap::new(),
})
}
}
#[derive(Deserialize)]
#[serde(tag = "type", rename_all = "snake_case")]
enum Event {
Assistant {
message: AssistantMessage,
},
Result(Finished),
#[serde(other)]
Other,
}
#[derive(Deserialize)]
struct AssistantMessage {
content: Vec<Block>,
}
#[derive(Deserialize)]
struct Block {
#[serde(rename = "type")]
kind: String,
name: Option<String>,
}
#[derive(Deserialize)]
struct Finished {
is_error: bool,
result: Option<String>,
usage: Option<TokenUsage>,
}
#[derive(Deserialize)]
struct TokenUsage {
input_tokens: u64,
output_tokens: u64,
}
impl Drive for ClaudeCode {
fn args(&self, _version: &Version, settings: &Settings, prompt: &Prompt) -> Vec<String> {
let base = [
"-p",
&prompt.text,
"--output-format",
"stream-json",
"--verbose",
"--model",
&settings.model,
];
let tools = ["--allowedTools", "Bash,Read,Write,Edit"];
base.into_iter()
.chain(tools.into_iter().filter(|_| prompt.allow_tools))
.map(str::to_owned)
.collect()
}
fn parse(&self, _version: &Version, stdout: &str) -> Outcome {
let events: Vec<Event> = json_lines(stdout).collect();
let tool_calls = events
.iter()
.filter_map(|event| match event {
Event::Assistant { message } => Some(&message.content),
_ => None,
})
.flatten()
.filter(|block| block.kind == "tool_use")
.filter_map(|block| block.name.clone())
.collect();
let finished = events.into_iter().find_map(|event| match event {
Event::Result(finished) => Some(finished),
_ => None,
});
let Some(finished) = finished else {
return Outcome {
tool_calls,
..Outcome::default()
};
};
let result = finished.result.unwrap_or_default();
let (text, errors) = if finished.is_error {
(String::new(), vec![result])
} else {
(result, Vec::new())
};
Outcome {
text,
tool_calls,
usage: finished.usage.map_or_else(Usage::default, |usage| Usage {
input_tokens: usage.input_tokens,
output_tokens: usage.output_tokens,
}),
errors,
exit_code: None,
}
}
}

View file

@ -0,0 +1,174 @@
use std::collections::BTreeMap;
use std::path::{Path, PathBuf};
use semver::Version;
use serde::Deserialize;
use super::{
Configure, Drive, Install, LaunchSpec, Outcome, Prompt, Settings, Usage, Wire, env, json_lines,
path_string, quoted, v1,
};
use crate::install::release::github_release;
use crate::install::{Packaging, Release};
use crate::target::{Arch, Os};
use crate::{Error, Fetch, Target};
const RELEASES: &str = "https://api.github.com/repos/openai/codex/releases/tags";
pub struct Codex;
fn triple(target: Target) -> String {
let arch = match target.arch {
Arch::Aarch64 => "aarch64",
Arch::X86_64 => "x86_64",
};
match target.os {
Os::Macos => format!("{arch}-apple-darwin"),
Os::Linux => format!("{arch}-unknown-linux-musl"),
}
}
impl Install for Codex {
fn binary(&self) -> &'static str {
"codex"
}
async fn release(
&self,
fetch: &impl Fetch,
version: &Version,
target: Target,
) -> Result<Release, Error> {
let triple = triple(target);
github_release(
fetch,
RELEASES,
&format!("rust-v{version}"),
&format!("codex-{triple}.tar.gz"),
Packaging::TarGz {
member: format!("codex-{triple}"),
},
)
.await
}
}
impl Configure for Codex {
fn configure(
&self,
_version: &Version,
settings: &Settings,
home: &Path,
) -> Result<LaunchSpec, Error> {
if settings.wire != Wire::Responses {
return Err(Error::UnsupportedWire {
agent: "codex",
wire: settings.wire,
});
}
let config = format!(
"model = {model}\nmodel_provider = \"litellm\"\n\n[model_providers.litellm]\nname = \"LiteLLM\"\nbase_url = {base_url}\nenv_key = \"LITELLM_API_KEY\"\nwire_api = \"responses\"\n",
model = quoted(&settings.model),
base_url = quoted(&v1(settings)),
);
Ok(LaunchSpec {
env: env([
("HOME", path_string(home)),
("CODEX_HOME", path_string(&home.join(".codex"))),
("LITELLM_API_KEY", settings.api_key.clone()),
]),
files: BTreeMap::from([(PathBuf::from(".codex/config.toml"), config)]),
})
}
}
#[derive(Deserialize)]
enum EventKind {
#[serde(rename = "item.completed")]
ItemCompleted,
#[serde(rename = "turn.completed")]
TurnCompleted,
#[serde(rename = "turn.failed")]
TurnFailed,
#[serde(other)]
Other,
}
#[derive(Deserialize)]
struct Event {
#[serde(rename = "type")]
kind: EventKind,
item: Option<Item>,
usage: Option<TokenUsage>,
error: Option<Failure>,
}
#[derive(Deserialize)]
struct Item {
#[serde(rename = "type")]
kind: String,
text: Option<String>,
}
#[derive(Deserialize)]
struct TokenUsage {
input_tokens: u64,
output_tokens: u64,
}
#[derive(Deserialize)]
struct Failure {
message: String,
}
const NON_TOOL_ITEMS: [&str; 3] = ["agent_message", "reasoning", "error"];
impl Drive for Codex {
fn args(&self, _version: &Version, _settings: &Settings, prompt: &Prompt) -> Vec<String> {
let sandbox = ["--sandbox", "workspace-write"];
["exec", "--json", "--skip-git-repo-check"]
.into_iter()
.chain(sandbox.into_iter().filter(|_| prompt.allow_tools))
.chain([prompt.text.as_str()])
.map(str::to_owned)
.collect()
}
fn parse(&self, _version: &Version, stdout: &str) -> Outcome {
let events: Vec<Event> = json_lines(stdout).collect();
let items: Vec<&Item> = events
.iter()
.filter(|event| matches!(event.kind, EventKind::ItemCompleted))
.filter_map(|event| event.item.as_ref())
.collect();
Outcome {
text: items
.iter()
.rev()
.find(|item| item.kind == "agent_message")
.and_then(|item| item.text.clone())
.unwrap_or_default(),
tool_calls: items
.iter()
.filter(|item| !NON_TOOL_ITEMS.contains(&item.kind.as_str()))
.map(|item| item.kind.clone())
.collect(),
usage: events
.iter()
.filter(|event| matches!(event.kind, EventKind::TurnCompleted))
.filter_map(|event| event.usage.as_ref())
.map(|usage| Usage {
input_tokens: usage.input_tokens,
output_tokens: usage.output_tokens,
})
.fold(Usage::default(), |total, turn| total + turn),
errors: events
.iter()
.filter(|event| matches!(event.kind, EventKind::TurnFailed))
.filter_map(|event| event.error.as_ref())
.map(|failure| failure.message.clone())
.collect(),
exit_code: None,
}
}
}

View file

@ -0,0 +1,69 @@
use std::collections::BTreeMap;
use std::path::{Path, PathBuf};
use semver::Version;
use crate::Error;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum Wire {
ChatCompletions,
Messages,
Responses,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct Settings {
pub base_url: String,
pub api_key: String,
pub model: String,
pub wire: Wire,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct LaunchSpec {
pub env: BTreeMap<String, String>,
pub files: BTreeMap<PathBuf, String>,
}
impl LaunchSpec {
pub fn write_files(&self, home: &Path) -> std::io::Result<()> {
self.files.iter().try_for_each(|(relative, contents)| {
let path = home.join(relative);
if let Some(parent) = path.parent() {
std::fs::create_dir_all(parent)?;
}
std::fs::write(path, contents)
})
}
}
pub trait Configure {
fn configure(
&self,
version: &Version,
settings: &Settings,
home: &Path,
) -> Result<LaunchSpec, Error>;
}
pub(crate) fn env(
pairs: impl IntoIterator<Item = (&'static str, String)>,
) -> BTreeMap<String, String> {
pairs
.into_iter()
.map(|(key, value)| (key.to_owned(), value))
.collect()
}
pub(crate) fn path_string(path: &Path) -> String {
path.to_string_lossy().into_owned()
}
pub(crate) fn quoted(value: &str) -> String {
serde_json::Value::from(value).to_string()
}
pub(crate) fn v1(settings: &Settings) -> String {
format!("{}/v1", settings.base_url.trim_end_matches('/'))
}

View file

@ -0,0 +1,57 @@
use std::ops::Add;
use semver::Version;
use crate::Settings;
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct Prompt {
pub text: String,
pub allow_tools: bool,
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub struct Usage {
pub input_tokens: u64,
pub output_tokens: u64,
}
impl Add for Usage {
type Output = Self;
fn add(self, other: Self) -> Self {
Self {
input_tokens: self.input_tokens + other.input_tokens,
output_tokens: self.output_tokens + other.output_tokens,
}
}
}
#[derive(Clone, Debug, Default, PartialEq, Eq)]
pub struct Outcome {
pub text: String,
pub tool_calls: Vec<String>,
pub usage: Usage,
pub errors: Vec<String>,
pub exit_code: Option<i32>,
}
impl Outcome {
pub fn succeeded(&self) -> bool {
self.exit_code == Some(0) && self.errors.is_empty()
}
}
pub trait Drive {
fn args(&self, version: &Version, settings: &Settings, prompt: &Prompt) -> Vec<String>;
fn parse(&self, version: &Version, stdout: &str) -> Outcome;
}
pub(crate) fn json_lines<'a, T: serde::de::DeserializeOwned + 'a>(
stdout: &'a str,
) -> impl Iterator<Item = T> + 'a {
stdout
.lines()
.filter_map(|line| serde_json::from_str(line).ok())
}

View file

@ -0,0 +1,17 @@
use std::future::Future;
use semver::Version;
use crate::install::Release;
use crate::{Error, Fetch, Target};
pub trait Install: Sync {
fn binary(&self) -> &'static str;
fn release(
&self,
fetch: &impl Fetch,
version: &Version,
target: Target,
) -> impl Future<Output = Result<Release, Error>> + Send;
}

View file

@ -0,0 +1,20 @@
mod claude;
mod codex;
mod configure;
mod drive;
mod install;
mod opencode;
pub use claude::ClaudeCode;
pub use codex::Codex;
pub use configure::{Configure, LaunchSpec, Settings, Wire};
pub use drive::{Drive, Outcome, Prompt, Usage};
pub use install::Install;
pub use opencode::Opencode;
pub(crate) use configure::{env, path_string, quoted, v1};
pub(crate) use drive::json_lines;
pub trait Agent: Install + Configure + Drive {}
impl<T: Install + Configure + Drive> Agent for T {}

View file

@ -0,0 +1,187 @@
use std::collections::BTreeMap;
use std::path::{Path, PathBuf};
use semver::Version;
use serde::Deserialize;
use super::{
Configure, Drive, Install, LaunchSpec, Outcome, Prompt, Settings, Usage, Wire, env, json_lines,
path_string, v1,
};
use crate::install::release::github_release;
use crate::install::{Packaging, Release};
use crate::target::Os;
use crate::{Error, Fetch, Target};
const RELEASES: &str = "https://api.github.com/repos/sst/opencode/releases/tags";
pub struct Opencode;
impl Install for Opencode {
fn binary(&self) -> &'static str {
"opencode"
}
async fn release(
&self,
fetch: &impl Fetch,
version: &Version,
target: Target,
) -> Result<Release, Error> {
let stem = format!(
"opencode-{}-{}{}",
target.os_name(),
target.arch_name(),
target.musl_suffix()
);
let member = "opencode".to_owned();
let (asset, packaging) = match target.os {
Os::Macos => (format!("{stem}.zip"), Packaging::Zip { member }),
Os::Linux => (format!("{stem}.tar.gz"), Packaging::TarGz { member }),
};
github_release(fetch, RELEASES, &format!("v{version}"), &asset, packaging).await
}
}
impl Configure for Opencode {
fn configure(
&self,
_version: &Version,
settings: &Settings,
home: &Path,
) -> Result<LaunchSpec, Error> {
let npm = match settings.wire {
Wire::ChatCompletions => "@ai-sdk/openai-compatible",
Wire::Responses => "@ai-sdk/openai",
Wire::Messages => "@ai-sdk/anthropic",
};
let config = serde_json::json!({
"$schema": "https://opencode.ai/config.json",
"model": format!("litellm/{}", settings.model),
"provider": {
"litellm": {
"npm": npm,
"name": "LiteLLM",
"options": { "baseURL": v1(settings), "apiKey": settings.api_key },
"models": { settings.model.clone(): { "name": settings.model } },
}
},
});
Ok(LaunchSpec {
env: env([
("HOME", path_string(home)),
("XDG_CONFIG_HOME", path_string(&home.join(".config"))),
("XDG_DATA_HOME", path_string(&home.join(".local/share"))),
("OPENCODE_DISABLE_AUTOUPDATE", "true".to_owned()),
]),
files: BTreeMap::from([(
PathBuf::from(".config/opencode/opencode.json"),
config.to_string(),
)]),
})
}
}
#[derive(Deserialize)]
#[serde(tag = "type", rename_all = "snake_case")]
enum Event {
Text {
part: TextPart,
},
ToolUse {
part: ToolPart,
},
StepFinish {
part: StepFinish,
},
Error {
error: Failure,
},
#[serde(other)]
Other,
}
#[derive(Deserialize)]
struct TextPart {
text: String,
}
#[derive(Deserialize)]
struct ToolPart {
tool: String,
}
#[derive(Deserialize)]
struct StepFinish {
tokens: Tokens,
}
#[derive(Deserialize)]
struct Tokens {
input: u64,
output: u64,
}
#[derive(Deserialize)]
struct Failure {
name: String,
data: Option<FailureData>,
}
#[derive(Deserialize)]
struct FailureData {
message: Option<String>,
}
impl Drive for Opencode {
fn args(&self, _version: &Version, _settings: &Settings, prompt: &Prompt) -> Vec<String> {
["run", "--format", "json", &prompt.text]
.map(str::to_owned)
.to_vec()
}
fn parse(&self, _version: &Version, stdout: &str) -> Outcome {
let events: Vec<Event> = json_lines(stdout).collect();
Outcome {
text: events
.iter()
.rev()
.find_map(|event| match event {
Event::Text { part } => Some(part.text.clone()),
_ => None,
})
.unwrap_or_default(),
tool_calls: events
.iter()
.filter_map(|event| match event {
Event::ToolUse { part } => Some(part.tool.clone()),
_ => None,
})
.collect(),
usage: events
.iter()
.filter_map(|event| match event {
Event::StepFinish { part } => Some(Usage {
input_tokens: part.tokens.input,
output_tokens: part.tokens.output,
}),
_ => None,
})
.fold(Usage::default(), |total, step| total + step),
errors: events
.iter()
.filter_map(|event| match event {
Event::Error { error } => Some(
error
.data
.as_ref()
.and_then(|data| data.message.clone())
.unwrap_or_else(|| error.name.clone()),
),
_ => None,
})
.collect(),
exit_code: None,
}
}
}

View file

@ -0,0 +1,56 @@
use std::io;
use std::path::PathBuf;
use thiserror::Error;
use crate::Wire;
#[derive(Debug, Error)]
pub enum Error {
#[error("unsupported target {0}")]
UnsupportedTarget(String),
#[error("{0} is not a plain x.y.z release version")]
InvalidVersion(String),
#[error("request to {url} failed")]
Request {
url: String,
#[source]
source: reqwest::Error,
},
#[error("{url} answered with status {status}")]
Status { url: String, status: u16 },
#[error("release metadata at {url} is malformed")]
Metadata {
url: String,
#[source]
source: serde_json::Error,
},
#[error("release has no asset named {0}")]
AssetNotFound(String),
#[error("release publishes no sha256 for {0}")]
MissingChecksum(String),
#[error("sha256 mismatch for {asset}: expected {expected}, got {actual}")]
ChecksumMismatch {
asset: String,
expected: String,
actual: String,
},
#[error("archive does not contain {0}")]
ArchiveMemberNotFound(String),
#[error("archive is unreadable")]
Archive(#[source] io::Error),
#[error("zip archive is unreadable")]
Zip(#[from] zip::result::ZipError),
#[error("{binary} reports version '{reported}', expected {expected}")]
VersionMismatch {
binary: PathBuf,
expected: String,
reported: String,
},
#[error("{agent} cannot talk to the gateway over {wire:?}")]
UnsupportedWire { agent: &'static str, wire: Wire },
#[error("agent did not finish within {0:?}")]
Timeout(std::time::Duration),
#[error("io failure")]
Io(#[from] io::Error),
}

View file

@ -0,0 +1,52 @@
use std::io::{Cursor, Read};
use flate2::read::GzDecoder;
use sha2::{Digest, Sha256};
use super::release::Packaging;
use crate::Error;
pub(crate) fn verify_sha256(asset: &str, expected: &str, bytes: &[u8]) -> Result<(), Error> {
let actual = format!("{:x}", Sha256::digest(bytes));
if actual.eq_ignore_ascii_case(expected) {
return Ok(());
}
Err(Error::ChecksumMismatch {
asset: asset.to_owned(),
expected: expected.to_owned(),
actual,
})
}
pub(crate) fn extract_binary(packaging: &Packaging, bytes: &[u8]) -> Result<Vec<u8>, Error> {
match packaging {
Packaging::Bare => Ok(bytes.to_vec()),
Packaging::TarGz { member } => extract_tar_gz(member, bytes),
Packaging::Zip { member } => extract_zip(member, bytes),
}
}
fn extract_tar_gz(member: &str, bytes: &[u8]) -> Result<Vec<u8>, Error> {
let mut archive = tar::Archive::new(GzDecoder::new(bytes));
for entry in archive.entries().map_err(Error::Archive)? {
let mut entry = entry.map_err(Error::Archive)?;
let path = entry.path().map_err(Error::Archive)?;
if path.file_name().is_some_and(|name| name == member) {
let mut binary = Vec::new();
entry.read_to_end(&mut binary).map_err(Error::Archive)?;
return Ok(binary);
}
}
Err(Error::ArchiveMemberNotFound(member.to_owned()))
}
fn extract_zip(member: &str, bytes: &[u8]) -> Result<Vec<u8>, Error> {
let mut archive = zip::ZipArchive::new(Cursor::new(bytes))?;
let mut file = archive.by_name(member).map_err(|error| match error {
zip::result::ZipError::FileNotFound => Error::ArchiveMemberNotFound(member.to_owned()),
other => Error::Zip(other),
})?;
let mut binary = Vec::new();
file.read_to_end(&mut binary).map_err(Error::Archive)?;
Ok(binary)
}

View file

@ -0,0 +1,55 @@
use std::future::Future;
use crate::Error;
pub trait Fetch: Sync {
fn get(&self, url: &str) -> impl Future<Output = Result<Vec<u8>, Error>> + Send;
}
pub struct HttpFetch {
client: reqwest::Client,
github_token: Option<String>,
}
impl HttpFetch {
pub fn new(github_token: Option<String>) -> Self {
Self {
client: reqwest::Client::new(),
github_token,
}
}
pub fn from_env() -> Self {
Self::new(std::env::var("GITHUB_TOKEN").ok())
}
}
impl Fetch for HttpFetch {
async fn get(&self, url: &str) -> Result<Vec<u8>, Error> {
let request = self
.client
.get(url)
.header("user-agent", "litellm-testkit")
.header("accept", "application/json, application/octet-stream");
let request = match (
&self.github_token,
url.starts_with("https://api.github.com/"),
) {
(Some(token), true) => request.bearer_auth(token),
_ => request,
};
let request_error = |source| Error::Request {
url: url.to_owned(),
source,
};
let response = request.send().await.map_err(request_error)?;
let status = response.status();
if !status.is_success() {
return Err(Error::Status {
url: url.to_owned(),
status: status.as_u16(),
});
}
Ok(response.bytes().await.map_err(request_error)?.to_vec())
}
}

View file

@ -0,0 +1,118 @@
mod archive;
mod fetch;
pub(crate) mod release;
use std::os::unix::fs::PermissionsExt;
use std::path::{Path, PathBuf};
use std::process::Stdio;
use std::sync::atomic::{AtomicU64, Ordering};
use semver::Version;
use tokio::fs;
use tokio::process::Command;
use crate::{Error, Install, Target};
use archive::{extract_binary, verify_sha256};
static STAGING_COUNTER: AtomicU64 = AtomicU64::new(0);
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct Installed {
pub version: Version,
pub binary: PathBuf,
}
pub struct Installer<F> {
fetch: F,
cache_root: PathBuf,
target: Target,
}
impl<F: Fetch> Installer<F> {
pub fn new(fetch: F, cache_root: impl Into<PathBuf>, target: Target) -> Self {
Self {
fetch,
cache_root: cache_root.into(),
target,
}
}
pub async fn install(
&self,
agent: &impl Install,
version: &Version,
) -> Result<Installed, Error> {
validate_release(version)?;
let dir = self
.cache_root
.join(agent.binary())
.join(version.to_string());
let binary = dir.join(agent.binary());
let installed = Installed {
version: version.clone(),
binary: binary.clone(),
};
if fs::try_exists(&binary).await? && probe_version(&binary, version).await.is_ok() {
return Ok(installed);
}
let release = agent.release(&self.fetch, version, self.target).await?;
let archive = self.fetch.get(&release.url).await?;
verify_sha256(&release.asset, &release.sha256, &archive)?;
let contents = extract_binary(&release.packaging, &archive)?;
fs::create_dir_all(&dir).await?;
let staging = dir.join(format!(
".{}.{}.{}.partial",
agent.binary(),
std::process::id(),
STAGING_COUNTER.fetch_add(1, Ordering::Relaxed)
));
fs::write(&staging, contents).await?;
fs::set_permissions(&staging, std::fs::Permissions::from_mode(0o755)).await?;
fs::rename(&staging, &binary).await?;
match probe_version(&binary, version).await {
Ok(()) => Ok(installed),
Err(error) => {
fs::remove_file(&binary).await?;
Err(error)
}
}
}
}
fn validate_release(version: &Version) -> Result<(), Error> {
if version.pre.is_empty() && version.build.is_empty() {
return Ok(());
}
Err(Error::InvalidVersion(version.to_string()))
}
async fn probe_version(binary: &Path, expected: &Version) -> Result<(), Error> {
let home = std::env::temp_dir();
let output = Command::new(binary)
.arg("--version")
.env_clear()
.env("HOME", home)
.env("DISABLE_AUTOUPDATER", "1")
.stdin(Stdio::null())
.output()
.await?;
let stdout = String::from_utf8_lossy(&output.stdout);
if stdout
.split_whitespace()
.filter_map(|token| Version::parse(token).ok())
.any(|reported| &reported == expected)
{
return Ok(());
}
Err(Error::VersionMismatch {
binary: binary.to_owned(),
expected: expected.to_string(),
reported: stdout.trim().to_owned(),
})
}
pub use fetch::{Fetch, HttpFetch};
pub use release::{Packaging, Release};

View file

@ -0,0 +1,65 @@
use serde::Deserialize;
use crate::{Error, Fetch};
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum Packaging {
Bare,
TarGz { member: String },
Zip { member: String },
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct Release {
pub asset: String,
pub url: String,
pub sha256: String,
pub packaging: Packaging,
}
#[derive(Deserialize)]
struct GithubRelease {
assets: Vec<GithubAsset>,
}
#[derive(Deserialize)]
struct GithubAsset {
name: String,
digest: Option<String>,
browser_download_url: String,
}
pub(crate) async fn github_release(
fetch: &impl Fetch,
releases_url: &str,
tag: &str,
asset_name: &str,
packaging: Packaging,
) -> Result<Release, Error> {
let url = format!("{releases_url}/{tag}");
let release: GithubRelease = parse(&url, &fetch.get(&url).await?)?;
let asset = release
.assets
.into_iter()
.find(|asset| asset.name == asset_name)
.ok_or_else(|| Error::AssetNotFound(asset_name.to_owned()))?;
let sha256 = asset
.digest
.as_deref()
.and_then(|digest| digest.strip_prefix("sha256:"))
.ok_or_else(|| Error::MissingChecksum(asset_name.to_owned()))?
.to_owned();
Ok(Release {
asset: asset.name,
url: asset.browser_download_url,
sha256,
packaging,
})
}
pub(crate) fn parse<T: for<'de> Deserialize<'de>>(url: &str, body: &[u8]) -> Result<T, Error> {
serde_json::from_slice(body).map_err(|source| Error::Metadata {
url: url.to_owned(),
source,
})
}

View file

@ -0,0 +1,15 @@
mod agent;
mod error;
mod install;
mod session;
mod target;
pub use agent::{
Agent, ClaudeCode, Codex, Configure, Drive, Install, LaunchSpec, Opencode, Outcome, Prompt,
Settings, Usage, Wire,
};
pub use error::Error;
pub use install::{Fetch, HttpFetch, Installed, Installer, Packaging, Release};
pub use semver::Version;
pub use session::Session;
pub use target::{Arch, Os, Target};

View file

@ -0,0 +1,76 @@
use std::collections::BTreeMap;
use std::path::PathBuf;
use std::process::Stdio;
use std::time::Duration;
use semver::Version;
use tokio::process::Command;
use tokio::time::timeout;
use crate::{Configure, Drive, Error, Installed, Outcome, Prompt, Settings};
const STDERR_LIMIT_CHARS: usize = 2000;
pub struct Session {
binary: PathBuf,
home: PathBuf,
version: Version,
settings: Settings,
env: BTreeMap<String, String>,
}
impl Session {
pub fn prepare(
agent: &impl Configure,
installed: &Installed,
settings: Settings,
home: impl Into<PathBuf>,
) -> Result<Self, Error> {
let home = home.into();
let spec = agent.configure(&installed.version, &settings, &home)?;
spec.write_files(&home)?;
Ok(Self {
binary: installed.binary.clone(),
home,
version: installed.version.clone(),
settings,
env: spec.env,
})
}
pub async fn run(
&self,
agent: &impl Drive,
prompt: &Prompt,
limit: Duration,
) -> Result<Outcome, Error> {
let child = Command::new(&self.binary)
.args(agent.args(&self.version, &self.settings, prompt))
.env_clear()
.env("PATH", "/usr/bin:/bin")
.envs(&self.env)
.current_dir(&self.home)
.stdin(Stdio::null())
.kill_on_drop(true)
.output();
let output = timeout(limit, child)
.await
.map_err(|_| Error::Timeout(limit))??;
let parsed = agent.parse(&self.version, &String::from_utf8_lossy(&output.stdout));
let failed_silently = !output.status.success() && parsed.errors.is_empty();
Ok(Outcome {
errors: if failed_silently {
vec![
String::from_utf8_lossy(&output.stderr)
.chars()
.take(STDERR_LIMIT_CHARS)
.collect(),
]
} else {
parsed.errors
},
exit_code: output.status.code(),
..parsed
})
}
}

View file

@ -0,0 +1,69 @@
use target_lexicon::{Architecture, Environment, OperatingSystem, Triple};
use crate::Error;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum Os {
Macos,
Linux,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum Arch {
Aarch64,
X86_64,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct Target {
pub os: Os,
pub arch: Arch,
pub musl: bool,
}
impl Target {
pub fn host() -> Result<Self, Error> {
Self::try_from(&Triple::host())
}
pub(crate) const fn os_name(self) -> &'static str {
match self.os {
Os::Macos => "darwin",
Os::Linux => "linux",
}
}
pub(crate) const fn arch_name(self) -> &'static str {
match self.arch {
Arch::Aarch64 => "arm64",
Arch::X86_64 => "x64",
}
}
pub(crate) const fn musl_suffix(self) -> &'static str {
if self.musl { "-musl" } else { "" }
}
}
impl TryFrom<&Triple> for Target {
type Error = Error;
fn try_from(triple: &Triple) -> Result<Self, Error> {
let unsupported = || Error::UnsupportedTarget(triple.to_string());
let os = match triple.operating_system {
OperatingSystem::Darwin(_) | OperatingSystem::MacOSX(_) => Os::Macos,
OperatingSystem::Linux => Os::Linux,
_ => return Err(unsupported()),
};
let arch = match triple.architecture {
Architecture::Aarch64(_) => Arch::Aarch64,
Architecture::X86_64 => Arch::X86_64,
_ => return Err(unsupported()),
};
Ok(Self {
os,
arch,
musl: triple.environment == Environment::Musl,
})
}
}

View file

@ -0,0 +1,133 @@
use std::path::Path;
use litellm_testkit::{ClaudeCode, Codex, Configure, Error, Opencode, Settings, Version, Wire};
use rstest::rstest;
fn settings(wire: Wire) -> Settings {
Settings {
base_url: "http://localhost:4000/".to_owned(),
api_key: "sk-test \"quoted\"".to_owned(),
model: "some-model".to_owned(),
wire,
}
}
fn version() -> Version {
Version::new(1, 2, 3)
}
#[rstest]
#[case(&ClaudeCode, Wire::Messages)]
#[case(&Codex, Wire::Responses)]
#[case(&Opencode, Wire::ChatCompletions)]
fn every_agent_runs_inside_the_given_home(#[case] agent: &impl Configure, #[case] wire: Wire) {
let home = Path::new("/scratch/home");
let spec = agent.configure(&version(), &settings(wire), home).unwrap();
assert_eq!(spec.env["HOME"], "/scratch/home");
assert!(
spec.env
.iter()
.filter(|(key, _)| key.ends_with("_HOME") || key.as_str() == "CLAUDE_CONFIG_DIR")
.all(|(_, value)| value.starts_with("/scratch/home"))
);
assert!(spec.files.keys().all(|path| path.is_relative()));
}
#[rstest]
#[case::claude_code(&ClaudeCode, &[Wire::ChatCompletions, Wire::Responses])]
#[case::codex(&Codex, &[Wire::ChatCompletions, Wire::Messages])]
fn wires_an_agent_cannot_speak_are_refused(
#[case] agent: &impl Configure,
#[case] refused: &[Wire],
) {
refused.iter().for_each(|wire| {
let result = agent.configure(&version(), &settings(*wire), Path::new("/h"));
assert!(matches!(result, Err(Error::UnsupportedWire { wire: got, .. }) if got == *wire));
});
}
#[test]
fn claude_code_points_at_the_gateway_root_with_the_key_and_model() {
let spec = ClaudeCode
.configure(&version(), &settings(Wire::Messages), Path::new("/h"))
.unwrap();
assert_eq!(spec.env["ANTHROPIC_BASE_URL"], "http://localhost:4000/");
assert_eq!(spec.env["ANTHROPIC_AUTH_TOKEN"], "sk-test \"quoted\"");
assert_eq!(spec.env["ANTHROPIC_MODEL"], "some-model");
}
#[test]
fn codex_config_is_valid_toml_routing_the_responses_api_to_the_gateway() {
let dir = tempfile::tempdir().unwrap();
let spec = Codex
.configure(&version(), &settings(Wire::Responses), dir.path())
.unwrap();
spec.write_files(dir.path()).unwrap();
let config: toml::Table =
toml::from_str(&std::fs::read_to_string(dir.path().join(".codex/config.toml")).unwrap())
.unwrap();
let provider = &config["model_providers"]["litellm"];
assert_eq!(config["model"].as_str(), Some("some-model"));
assert_eq!(config["model_provider"].as_str(), Some("litellm"));
assert_eq!(
provider["base_url"].as_str(),
Some("http://localhost:4000/v1")
);
assert_eq!(provider["wire_api"].as_str(), Some("responses"));
let key_var = provider["env_key"].as_str().unwrap();
assert_eq!(spec.env[key_var], "sk-test \"quoted\"");
}
#[rstest]
#[case(Wire::ChatCompletions)]
#[case(Wire::Responses)]
#[case(Wire::Messages)]
fn opencode_config_is_valid_json_registering_the_gateway_model(#[case] wire: Wire) {
let dir = tempfile::tempdir().unwrap();
let spec = Opencode
.configure(&version(), &settings(wire), dir.path())
.unwrap();
spec.write_files(dir.path()).unwrap();
let config: serde_json::Value = serde_json::from_str(
&std::fs::read_to_string(dir.path().join(".config/opencode/opencode.json")).unwrap(),
)
.unwrap();
let provider = &config["provider"]["litellm"];
assert_eq!(config["model"], "litellm/some-model");
assert_eq!(provider["options"]["baseURL"], "http://localhost:4000/v1");
assert_eq!(provider["options"]["apiKey"], "sk-test \"quoted\"");
assert!(provider["models"]["some-model"].is_object());
}
#[test]
fn opencode_uses_a_different_provider_package_for_every_wire() {
let package = |wire| {
let dir = tempfile::tempdir().unwrap();
let spec = Opencode
.configure(&version(), &settings(wire), dir.path())
.unwrap();
let config: serde_json::Value =
serde_json::from_str(spec.files.values().next().unwrap()).unwrap();
config["provider"]["litellm"]["npm"]
.as_str()
.unwrap()
.to_owned()
};
let packages = [Wire::ChatCompletions, Wire::Responses, Wire::Messages].map(package);
assert_eq!(
packages
.iter()
.collect::<std::collections::BTreeSet<_>>()
.len(),
packages.len()
);
}

View file

@ -0,0 +1,262 @@
mod support;
use std::str::FromStr;
use litellm_testkit::{ClaudeCode, Codex, Error, Installer, Opencode, Target, Version};
use rstest::rstest;
use serde_json::json;
use support::{FakeFetch, script_printing, sha256, tar_gz, zip_archive};
use target_lexicon::Triple;
fn target(triple: &str) -> Target {
Target::try_from(&Triple::from_str(triple).unwrap()).unwrap()
}
fn linux() -> Target {
target("x86_64-unknown-linux-gnu")
}
fn version() -> Version {
Version::new(9, 8, 7)
}
fn github_release(asset: &str, download_url: &str, digest: Option<String>) -> Vec<u8> {
json!({
"assets": [
{ "name": "unrelated.txt", "digest": "sha256:00", "browser_download_url": "https://example.test/unrelated" },
{ "name": asset, "digest": digest, "browser_download_url": download_url },
]
})
.to_string()
.into_bytes()
}
fn claude_routes(binary: &[u8], checksum: &str) -> Vec<(String, Vec<u8>)> {
let base = "https://downloads.claude.ai/claude-code-releases/9.8.7";
let manifest = json!({ "platforms": { "linux-x64": { "checksum": checksum } } });
vec![
(
format!("{base}/manifest.json"),
manifest.to_string().into_bytes(),
),
(format!("{base}/linux-x64/claude"), binary.to_vec()),
]
}
fn codex_routes(archive: Vec<u8>, digest: Option<String>) -> Vec<(String, Vec<u8>)> {
let release = github_release(
"codex-x86_64-unknown-linux-musl.tar.gz",
"https://example.test/codex.tar.gz",
digest,
);
vec![
(
"https://api.github.com/repos/openai/codex/releases/tags/rust-v9.8.7".to_owned(),
release,
),
("https://example.test/codex.tar.gz".to_owned(), archive),
]
}
#[tokio::test]
async fn claude_bare_binary_is_installed_and_runnable() {
let binary = script_printing("9.8.7 (Claude Code)");
let fetch = FakeFetch::new(claude_routes(&binary, &sha256(&binary)));
let cache = tempfile::tempdir().unwrap();
let installed = Installer::new(&fetch, cache.path(), linux())
.install(&ClaudeCode, &version())
.await
.unwrap();
assert_eq!(installed.binary, cache.path().join("claude/9.8.7/claude"));
assert_eq!(std::fs::read(&installed.binary).unwrap(), binary);
}
#[tokio::test]
async fn codex_binary_is_extracted_from_the_tarball_under_its_own_name() {
let binary = script_printing("codex-cli 9.8.7");
let archive = tar_gz("codex-x86_64-unknown-linux-musl", &binary);
let fetch = FakeFetch::new(codex_routes(
archive.clone(),
Some(format!("sha256:{}", sha256(&archive))),
));
let cache = tempfile::tempdir().unwrap();
let installed = Installer::new(&fetch, cache.path(), linux())
.install(&Codex, &version())
.await
.unwrap();
assert_eq!(std::fs::read(&installed.binary).unwrap(), binary);
assert_eq!(installed.binary, cache.path().join("codex/9.8.7/codex"));
}
#[tokio::test]
async fn opencode_binary_is_extracted_from_the_darwin_zip() {
let binary = script_printing("9.8.7");
let archive = zip_archive("opencode", &binary);
let release = github_release(
"opencode-darwin-arm64.zip",
"https://example.test/opencode.zip",
Some(format!("sha256:{}", sha256(&archive))),
);
let fetch = FakeFetch::new([
(
"https://api.github.com/repos/sst/opencode/releases/tags/v9.8.7".to_owned(),
release,
),
("https://example.test/opencode.zip".to_owned(), archive),
]);
let cache = tempfile::tempdir().unwrap();
let installed = Installer::new(&fetch, cache.path(), target("aarch64-apple-darwin"))
.install(&Opencode, &version())
.await
.unwrap();
assert_eq!(std::fs::read(&installed.binary).unwrap(), binary);
}
#[tokio::test]
async fn tampered_download_is_rejected_and_nothing_is_left_behind() {
let binary = script_printing("9.8.7 (Claude Code)");
let fetch = FakeFetch::new(claude_routes(&binary, &sha256(b"what the vendor signed")));
let cache = tempfile::tempdir().unwrap();
let result = Installer::new(&fetch, cache.path(), linux())
.install(&ClaudeCode, &version())
.await;
assert!(matches!(result, Err(Error::ChecksumMismatch { .. })));
assert!(!cache.path().join("claude/9.8.7").exists());
}
#[tokio::test]
async fn github_asset_without_a_digest_is_refused() {
let archive = tar_gz(
"codex-x86_64-unknown-linux-musl",
&script_printing("codex-cli 9.8.7"),
);
let fetch = FakeFetch::new(codex_routes(archive, None));
let cache = tempfile::tempdir().unwrap();
let result = Installer::new(&fetch, cache.path(), linux())
.install(&Codex, &version())
.await;
assert!(matches!(result, Err(Error::MissingChecksum(_))));
}
#[tokio::test]
async fn binary_reporting_a_different_version_is_removed() {
let binary = script_printing("1.0.0 (Claude Code)");
let fetch = FakeFetch::new(claude_routes(&binary, &sha256(&binary)));
let cache = tempfile::tempdir().unwrap();
let result = Installer::new(&fetch, cache.path(), linux())
.install(&ClaudeCode, &version())
.await;
assert!(matches!(result, Err(Error::VersionMismatch { .. })));
assert!(!cache.path().join("claude/9.8.7/claude").exists());
}
#[tokio::test]
async fn second_install_reuses_the_cached_binary_without_downloading() {
let binary = script_printing("9.8.7 (Claude Code)");
let fetch = FakeFetch::new(claude_routes(&binary, &sha256(&binary)));
let cache = tempfile::tempdir().unwrap();
let installer = Installer::new(&fetch, cache.path(), linux());
let first = installer.install(&ClaudeCode, &version()).await.unwrap();
let calls_after_first = fetch.calls();
let second = installer.install(&ClaudeCode, &version()).await.unwrap();
assert_eq!(first, second);
assert_eq!(fetch.calls(), calls_after_first);
}
#[tokio::test]
async fn corrupted_cache_entry_is_replaced_by_a_fresh_download() {
let binary = script_printing("9.8.7 (Claude Code)");
let fetch = FakeFetch::new(claude_routes(&binary, &sha256(&binary)));
let cache = tempfile::tempdir().unwrap();
let installer = Installer::new(&fetch, cache.path(), linux());
let installed = installer.install(&ClaudeCode, &version()).await.unwrap();
std::fs::write(&installed.binary, script_printing("0.0.1")).unwrap();
installer.install(&ClaudeCode, &version()).await.unwrap();
assert_eq!(std::fs::read(&installed.binary).unwrap(), binary);
}
#[rstest]
#[case("9.8.7-beta.1")]
#[case("9.8.7+build.5")]
#[tokio::test]
async fn pre_releases_never_reach_the_network_or_the_filesystem(#[case] version: &str) {
let fetch = FakeFetch::new([]);
let cache = tempfile::tempdir().unwrap();
let result = Installer::new(&fetch, cache.path(), linux())
.install(&ClaudeCode, &Version::parse(version).unwrap())
.await;
assert!(matches!(result, Err(Error::InvalidVersion(_))));
assert_eq!(fetch.calls(), 0);
assert_eq!(std::fs::read_dir(cache.path()).unwrap().count(), 0);
}
#[tokio::test]
async fn musl_linux_picks_the_musl_claude_build() {
let binary = script_printing("9.8.7 (Claude Code)");
let base = "https://downloads.claude.ai/claude-code-releases/9.8.7";
let manifest = json!({ "platforms": {
"linux-x64": { "checksum": sha256(b"glibc build") },
"linux-x64-musl": { "checksum": sha256(&binary) },
} });
let fetch = FakeFetch::new([
(
format!("{base}/manifest.json"),
manifest.to_string().into_bytes(),
),
(format!("{base}/linux-x64-musl/claude"), binary.clone()),
]);
let cache = tempfile::tempdir().unwrap();
let installed = Installer::new(&fetch, cache.path(), target("x86_64-unknown-linux-musl"))
.install(&ClaudeCode, &version())
.await
.unwrap();
assert_eq!(std::fs::read(&installed.binary).unwrap(), binary);
}
#[rstest]
#[case("x86_64-pc-windows-msvc")]
#[case("riscv64gc-unknown-linux-gnu")]
#[case("wasm32-unknown-unknown")]
fn targets_no_agent_ships_for_are_rejected(#[case] triple: &str) {
let result = Target::try_from(&Triple::from_str(triple).unwrap());
assert!(matches!(result, Err(Error::UnsupportedTarget(_))));
}
#[tokio::test]
async fn concurrent_installs_of_the_same_version_both_succeed() {
let binary = script_printing("9.8.7 (Claude Code)");
let fetch = FakeFetch::new(claude_routes(&binary, &sha256(&binary)));
let cache = tempfile::tempdir().unwrap();
let installer = Installer::new(&fetch, cache.path(), linux());
let wanted = version();
let installs =
futures_util::future::join_all((0..8).map(|_| installer.install(&ClaudeCode, &wanted)))
.await;
assert!(installs.iter().all(Result::is_ok));
assert_eq!(
std::fs::read(&installs[0].as_ref().unwrap().binary).unwrap(),
binary
);
}

View file

@ -0,0 +1,133 @@
//! Drives the real agents through a real gateway. Run with `cargo test -p litellm-testkit --test live -- --ignored`
//! after exporting `TESTKIT_GATEWAY_URL`, `TESTKIT_GATEWAY_KEY`, one `TESTKIT_MODEL_<WIRE>` per wire
//! (`MESSAGES`, `RESPONSES`, `CHAT_COMPLETIONS`) and one `TESTKIT_<AGENT>_VERSION` per agent
//! (`CLAUDE`, `CODEX`, `OPENCODE`). `TESTKIT_CACHE_DIR` and `GITHUB_TOKEN` are optional.
use std::path::PathBuf;
use std::time::Duration;
use litellm_testkit::{
Agent, ClaudeCode, Codex, HttpFetch, Installer, Opencode, Outcome, Prompt, Session, Settings,
Target, Version, Wire,
};
use rstest::rstest;
const LIMIT: Duration = Duration::from_secs(180);
fn required(name: &str) -> String {
std::env::var(name).unwrap_or_else(|_| panic!("{name} must be set to run the live tests"))
}
fn model_var(wire: Wire) -> &'static str {
match wire {
Wire::Messages => "TESTKIT_MODEL_MESSAGES",
Wire::Responses => "TESTKIT_MODEL_RESPONSES",
Wire::ChatCompletions => "TESTKIT_MODEL_CHAT_COMPLETIONS",
}
}
async fn drive(
agent: &impl Agent,
version_var: &str,
wire: Wire,
model: Option<&str>,
prompt: Prompt,
) -> Outcome {
let cache = std::env::var("TESTKIT_CACHE_DIR")
.map(PathBuf::from)
.unwrap_or_else(|_| std::env::temp_dir().join("litellm-testkit-cache"));
let installer = Installer::new(HttpFetch::from_env(), cache, Target::host().unwrap());
let installed = installer
.install(agent, &Version::parse(&required(version_var)).unwrap())
.await
.unwrap();
let settings = Settings {
base_url: required("TESTKIT_GATEWAY_URL"),
api_key: required("TESTKIT_GATEWAY_KEY"),
model: model.map_or_else(|| required(model_var(wire)), str::to_owned),
wire,
};
let home = tempfile::tempdir().unwrap();
let session = Session::prepare(agent, &installed, settings, home.path()).unwrap();
session.run(agent, &prompt, LIMIT).await.unwrap()
}
fn text_prompt() -> Prompt {
Prompt {
text: "Reply with the single word: pong".to_owned(),
allow_tools: false,
}
}
fn tool_prompt() -> Prompt {
Prompt {
text: "Run the shell command 'echo tool-ok' and reply with exactly its output.".to_owned(),
allow_tools: true,
}
}
#[rstest]
#[case::claude_messages(&ClaudeCode, "TESTKIT_CLAUDE_VERSION", Wire::Messages)]
#[case::codex_responses(&Codex, "TESTKIT_CODEX_VERSION", Wire::Responses)]
#[case::opencode_chat(&Opencode, "TESTKIT_OPENCODE_VERSION", Wire::ChatCompletions)]
#[case::opencode_responses(&Opencode, "TESTKIT_OPENCODE_VERSION", Wire::Responses)]
#[case::opencode_messages(&Opencode, "TESTKIT_OPENCODE_VERSION", Wire::Messages)]
#[ignore = "needs a live gateway, see the module docs"]
#[tokio::test]
async fn plain_prompt_gets_an_answer_and_token_usage(
#[case] agent: &impl Agent,
#[case] version_var: &str,
#[case] wire: Wire,
) {
let outcome = drive(agent, version_var, wire, None, text_prompt()).await;
assert!(outcome.succeeded(), "{outcome:?}");
assert!(outcome.text.to_lowercase().contains("pong"), "{outcome:?}");
assert!(outcome.usage.output_tokens > 0, "{outcome:?}");
}
#[rstest]
#[case::claude_messages(&ClaudeCode, "TESTKIT_CLAUDE_VERSION", Wire::Messages)]
#[case::codex_responses(&Codex, "TESTKIT_CODEX_VERSION", Wire::Responses)]
#[case::opencode_chat(&Opencode, "TESTKIT_OPENCODE_VERSION", Wire::ChatCompletions)]
#[case::opencode_responses(&Opencode, "TESTKIT_OPENCODE_VERSION", Wire::Responses)]
#[case::opencode_messages(&Opencode, "TESTKIT_OPENCODE_VERSION", Wire::Messages)]
#[ignore = "needs a live gateway, see the module docs"]
#[tokio::test]
async fn tool_use_is_reported_and_its_result_reaches_the_answer(
#[case] agent: &impl Agent,
#[case] version_var: &str,
#[case] wire: Wire,
) {
let outcome = drive(agent, version_var, wire, None, tool_prompt()).await;
assert!(outcome.succeeded(), "{outcome:?}");
assert!(!outcome.tool_calls.is_empty(), "{outcome:?}");
assert!(outcome.text.contains("tool-ok"), "{outcome:?}");
}
#[rstest]
#[case::claude_messages(&ClaudeCode, "TESTKIT_CLAUDE_VERSION", Wire::Messages)]
#[case::codex_responses(&Codex, "TESTKIT_CODEX_VERSION", Wire::Responses)]
#[case::opencode_chat(&Opencode, "TESTKIT_OPENCODE_VERSION", Wire::ChatCompletions)]
#[case::opencode_responses(&Opencode, "TESTKIT_OPENCODE_VERSION", Wire::Responses)]
#[case::opencode_messages(&Opencode, "TESTKIT_OPENCODE_VERSION", Wire::Messages)]
#[ignore = "needs a live gateway, see the module docs"]
#[tokio::test]
async fn model_the_gateway_rejects_is_reported_as_an_error(
#[case] agent: &impl Agent,
#[case] version_var: &str,
#[case] wire: Wire,
) {
let outcome = drive(
agent,
version_var,
wire,
Some("testkit-no-such-model"),
text_prompt(),
)
.await;
assert!(!outcome.succeeded(), "{outcome:?}");
assert!(!outcome.errors.is_empty(), "{outcome:?}");
}

View file

@ -0,0 +1,155 @@
use std::os::unix::fs::PermissionsExt;
use std::path::{Path, PathBuf};
use std::time::Duration;
use litellm_testkit::{
Configure, Drive, Error, Installed, LaunchSpec, Outcome, Prompt, Session, Settings, Version,
Wire,
};
struct Scripted;
impl Configure for Scripted {
fn configure(
&self,
version: &Version,
_settings: &Settings,
home: &Path,
) -> Result<LaunchSpec, Error> {
Ok(LaunchSpec {
env: [
("AGENT_HOME".to_owned(), home.to_string_lossy().into_owned()),
("AGENT_SAW_VERSION".to_owned(), version.to_string()),
]
.into(),
files: [(
PathBuf::from("conf/agent.toml"),
"configured = true\n".to_owned(),
)]
.into(),
})
}
}
impl Drive for Scripted {
fn args(&self, _version: &Version, _settings: &Settings, prompt: &Prompt) -> Vec<String> {
vec!["--prompt".to_owned(), prompt.text.clone()]
}
fn parse(&self, _version: &Version, stdout: &str) -> Outcome {
Outcome {
text: stdout.to_owned(),
..Outcome::default()
}
}
}
fn settings() -> Settings {
Settings {
base_url: "http://gateway.test".to_owned(),
api_key: "sk-test".to_owned(),
model: "some-model".to_owned(),
wire: Wire::Messages,
}
}
fn prompt(text: &str) -> Prompt {
Prompt {
text: text.to_owned(),
allow_tools: false,
}
}
fn session(script: &str) -> (Session, tempfile::TempDir) {
let dir = tempfile::tempdir().unwrap();
let binary = dir.path().join("agent");
std::fs::write(&binary, format!("#!/bin/sh\n{script}\n")).unwrap();
std::fs::set_permissions(&binary, std::fs::Permissions::from_mode(0o755)).unwrap();
let home = dir.path().join("home");
std::fs::create_dir(&home).unwrap();
let installed = Installed {
version: Version::new(4, 5, 6),
binary,
};
(
Session::prepare(&Scripted, &installed, settings(), home).unwrap(),
dir,
)
}
const LIMIT: Duration = Duration::from_secs(20);
#[tokio::test]
async fn prepare_writes_the_config_files_under_home() {
let (_session, dir) = session("true");
let written = std::fs::read_to_string(dir.path().join("home/conf/agent.toml")).unwrap();
assert_eq!(written, "configured = true\n");
}
#[tokio::test]
async fn configure_and_drive_are_given_the_installed_version() {
let (session, _dir) = session("echo \"$AGENT_SAW_VERSION\"");
let outcome = session.run(&Scripted, &prompt("hi"), LIMIT).await.unwrap();
assert_eq!(outcome.text.trim(), "4.5.6");
}
#[tokio::test]
async fn agent_runs_in_home_with_only_its_own_environment() {
let (session, dir) = session("pwd -P; env");
let outcome = session.run(&Scripted, &prompt("hi"), LIMIT).await.unwrap();
let home = dir.path().join("home").canonicalize().unwrap();
assert_eq!(outcome.text.lines().next().unwrap(), home.to_string_lossy());
assert!(outcome.text.contains("AGENT_HOME="));
assert!(
!outcome.text.contains("CARGO_"),
"test runner environment leaked into the agent"
);
}
#[tokio::test]
async fn prompt_reaches_the_agent_as_one_untouched_argument() {
let (session, _dir) = session("printf '%s|' \"$@\"");
let text = "two spaces; $(echo injected) 'quoted'";
let outcome = session.run(&Scripted, &prompt(text), LIMIT).await.unwrap();
assert_eq!(outcome.text, format!("--prompt|{text}|"));
}
#[tokio::test]
async fn clean_exit_is_a_success() {
let (session, _dir) = session("echo done");
let outcome = session.run(&Scripted, &prompt("hi"), LIMIT).await.unwrap();
assert_eq!(outcome.exit_code, Some(0));
assert!(outcome.succeeded());
}
#[tokio::test]
async fn failing_exit_without_a_parsed_error_reports_stderr() {
let (session, _dir) = session("echo boom >&2; exit 3");
let outcome = session.run(&Scripted, &prompt("hi"), LIMIT).await.unwrap();
assert_eq!(outcome.exit_code, Some(3));
assert!(!outcome.succeeded());
assert_eq!(outcome.errors, ["boom\n"]);
}
#[tokio::test]
async fn agent_that_outlives_the_limit_is_stopped() {
let (session, _dir) = session("sleep 30");
let result = session
.run(&Scripted, &prompt("hi"), Duration::from_millis(200))
.await;
assert!(matches!(result, Err(Error::Timeout(_))));
}

View file

@ -0,0 +1,70 @@
use std::collections::HashMap;
use std::io::Write;
use std::sync::atomic::{AtomicUsize, Ordering};
use litellm_testkit::{Error, Fetch};
use sha2::{Digest, Sha256};
pub struct FakeFetch {
routes: HashMap<String, Vec<u8>>,
calls: AtomicUsize,
}
impl FakeFetch {
pub fn new(routes: impl IntoIterator<Item = (String, Vec<u8>)>) -> Self {
Self {
routes: routes.into_iter().collect(),
calls: AtomicUsize::new(0),
}
}
pub fn calls(&self) -> usize {
self.calls.load(Ordering::SeqCst)
}
}
impl Fetch for FakeFetch {
async fn get(&self, url: &str) -> Result<Vec<u8>, Error> {
self.calls.fetch_add(1, Ordering::SeqCst);
self.routes.get(url).cloned().ok_or_else(|| Error::Status {
url: url.to_owned(),
status: 404,
})
}
}
impl Fetch for &FakeFetch {
async fn get(&self, url: &str) -> Result<Vec<u8>, Error> {
(*self).get(url).await
}
}
pub fn sha256(bytes: &[u8]) -> String {
format!("{:x}", Sha256::digest(bytes))
}
pub fn script_printing(output: &str) -> Vec<u8> {
format!("#!/bin/sh\necho '{output}'\n").into_bytes()
}
pub fn tar_gz(member: &str, contents: &[u8]) -> Vec<u8> {
let mut builder = tar::Builder::new(Vec::new());
let mut header = tar::Header::new_gnu();
header.set_size(contents.len() as u64);
header.set_mode(0o755);
header.set_cksum();
builder.append_data(&mut header, member, contents).unwrap();
let tarball = builder.into_inner().unwrap();
let mut encoder = flate2::write::GzEncoder::new(Vec::new(), flate2::Compression::default());
encoder.write_all(&tarball).unwrap();
encoder.finish().unwrap()
}
pub fn zip_archive(member: &str, contents: &[u8]) -> Vec<u8> {
let mut writer = zip::ZipWriter::new(std::io::Cursor::new(Vec::new()));
writer
.start_file(member, zip::write::SimpleFileOptions::default())
.unwrap();
writer.write_all(contents).unwrap();
writer.finish().unwrap().into_inner()
}

View file

@ -10,6 +10,7 @@ import os
from collections.abc import AsyncIterator, Awaitable, Callable, Generator, Sequence
from contextlib import AbstractAsyncContextManager
from functools import partial
from importlib.metadata import version
from types import MappingProxyType
from typing import Final, TypeAlias, TypeVar, cast
@ -34,8 +35,16 @@ _TransportContext: TypeAlias = AbstractAsyncContextManager[_TransportStreams]
from mcp.types import (
METHOD_NOT_FOUND,
REQUEST_TIMEOUT,
ClientCapabilities,
ElicitationCapability,
FormElicitationCapability,
GetPromptRequestParams,
GetPromptResult,
Implementation,
InitializedNotification,
InitializeRequest,
InitializeRequestParams,
InitializeResult,
InputRequiredResult,
ListPromptsResult,
ListResourcesResult,
@ -44,12 +53,14 @@ from mcp.types import (
PaginatedResult,
Prompt,
ResourceTemplate,
SamplingCapability,
ServerNotification,
UrlElicitationCapability,
)
from mcp.types import CallToolRequestParams as MCPCallToolRequestParams
from mcp.types import CallToolResult as MCPCallToolResult
from mcp.types import Tool as MCPTool
from pydantic import AnyUrl
from pydantic import AnyUrl, TypeAdapter
from litellm._logging import verbose_logger
from litellm.constants import (
@ -64,11 +75,13 @@ from litellm.proxy._experimental.mcp_server.mcp_debug import capture_upstream_er
from litellm.proxy._experimental.mcp_server.result_conversion import error_text_result
from litellm.types.llms.custom_http import VerifyTypes
from litellm.types.mcp import (
MCP_LEGACY_VERSIONS,
MCPAuth,
MCPAuthType,
MCPStdioConfig,
MCPTransport,
MCPTransportType,
MCPUpstreamProtocol,
credential_redirect_hook,
has_header,
without_header,
@ -386,7 +399,9 @@ class MCPClient:
sampling_callback: Callable | None = None,
elicitation_callback: Callable | None = None,
logging_callback: Callable | None = None,
protocol_version: MCPUpstreamProtocol = "auto",
):
self.protocol_version: MCPUpstreamProtocol = TypeAdapter(MCPUpstreamProtocol).validate_python(protocol_version)
self.server_url: str = server_url
self.transport_type: MCPTransport = transport_type
self.auth_type: MCPAuthType = auth_type
@ -525,6 +540,35 @@ class MCPClient:
return safe_env
async def _initialize_session(self, session: ClientSession) -> InitializeResult:
if self.protocol_version == "auto":
automatic: Final = await session.initialize()
if automatic.protocol_version not in MCP_LEGACY_VERSIONS:
raise MCPError(code=-32022, message="Upstream selected an unsupported MCP protocol version")
return automatic
result: Final = await session.send_request(
InitializeRequest(
params=InitializeRequestParams(
protocol_version=self.protocol_version,
client_info=Implementation(name="litellm", version=version("litellm")),
capabilities=ClientCapabilities(
sampling=SamplingCapability() if self._sampling_callback is not None else None,
elicitation=ElicitationCapability(
form=FormElicitationCapability(), url=UrlElicitationCapability()
)
if self._elicitation_callback is not None
else None,
),
)
),
InitializeResult,
)
if result.protocol_version != self.protocol_version:
raise MCPError(code=-32022, message="Upstream did not accept the configured MCP protocol version")
session.adopt(result)
await session.send_notification(InitializedNotification())
return result
async def _execute_session_operation(
self,
transport_ctx: _TransportContext,
@ -579,7 +623,7 @@ class MCPClient:
)
session: Final = await session_ctx.__aenter__()
try:
init_result: Final = await session.initialize()
init_result: Final = await self._initialize_session(session)
instructions: Final = getattr(init_result, "instructions", None)
self._last_initialize_instructions = (
instructions.strip() or None if isinstance(instructions, str) else None

View file

@ -92,7 +92,7 @@ litellm/integrations/levo/
## Testing
See the test files in `tests/test_litellm/integrations/levo/`:
See the test files in `tests/unit/integrations/levo/`:
- `test_levo.py`: Unit tests for configuration
- `test_levo_integration.py`: Integration tests for callback registration

View file

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

View file

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

View file

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

View file

@ -0,0 +1,145 @@
from collections.abc import Callable, Mapping
from dataclasses import dataclass
from itertools import product
from types import MappingProxyType
from typing import Final, Literal
from mcp.server.context import CallNext, HandlerResult, ServerRequestContext
from mcp.shared.exceptions import MCPError
from mcp.types import DiscoverResult, InitializeRequestParams, InitializeResult, ServerCapabilities
from mcp_types.methods import CLIENT_REQUESTS
from mcp_types.version import HANDSHAKE_PROTOCOL_VERSIONS, LATEST_HANDSHAKE_VERSION
from pydantic import TypeAdapter
from litellm.types.mcp import MCP_LEGACY_VERSIONS, MCPAdvertisedVersions, MCPLegacyVersion, MCPSpecVersion, MCPTransport
GATEWAY_OPERATIONS: Final = frozenset(
{
"tools/list",
"tools/call",
"prompts/list",
"prompts/get",
"resources/list",
"resources/read",
"resources/templates/list",
}
)
@dataclass(frozen=True, slots=True)
class RevisionSupport:
transports: frozenset[MCPTransport]
operations: frozenset[str]
results: frozenset[Literal["complete", "input_required"]]
extensions: frozenset[str]
completed: bool
REVISION_SUPPORT: Final[Mapping[str, RevisionSupport]] = MappingProxyType(
{
version.value: RevisionSupport(
transports=frozenset(MCPTransport)
if version.value in HANDSHAKE_PROTOCOL_VERSIONS
else frozenset({MCPTransport.http, MCPTransport.stdio}),
operations=frozenset(method for method in GATEWAY_OPERATIONS if (method, version.value) in CLIENT_REQUESTS),
results=frozenset({"complete"})
if version.value in HANDSHAKE_PROTOCOL_VERSIONS
else frozenset({"complete", "input_required"}),
extensions=frozenset(),
completed=version.value in HANDSHAKE_PROTOCOL_VERSIONS,
)
for version in MCPSpecVersion
}
)
_COMPLETED_REVISIONS: Final = tuple(version for version, support in REVISION_SUPPORT.items() if support.completed)
TRANSLATION_PAIRS: Final = frozenset(product(_COMPLETED_REVISIONS, repeat=2))
_ADVERTISED_VERSIONS: Final[TypeAdapter[tuple[MCPLegacyVersion, ...]]] = TypeAdapter(MCPAdvertisedVersions)
def configured_versions() -> tuple[str, ...]:
from litellm.proxy.proxy_server import general_settings_view
configured: Final = general_settings_view().get("mcp_advertised_versions")
return _ADVERTISED_VERSIONS.validate_python(MCP_LEGACY_VERSIONS if configured is None else configured)
def build_discovery(
*,
configured: tuple[str, ...],
revision: str,
transport: MCPTransport,
authorized_operations: frozenset[str],
upstream_versions: frozenset[str],
capabilities: ServerCapabilities,
client_extensions: frozenset[str] = frozenset(),
upstream_extensions: frozenset[str] = frozenset(),
instructions: str | None = None,
) -> DiscoverResult:
supported: Final = tuple(
version
for version, support in REVISION_SUPPORT.items()
if version in configured and support.completed and transport in support.transports
)
revision_support: Final = REVISION_SUPPORT.get(revision)
operations: Final[frozenset[str]] = (
authorized_operations & revision_support.operations
if revision in supported
and revision_support is not None
and any((revision, upstream) in TRANSLATION_PAIRS for upstream in upstream_versions)
else frozenset()
)
extensions: Final[frozenset[str]] = (
revision_support.extensions & client_extensions & upstream_extensions
if operations and revision_support is not None
else frozenset()
)
caller_capabilities: Final = capabilities.model_copy(deep=True)
return DiscoverResult(
supported_versions=list(supported),
capabilities=ServerCapabilities(
tools=caller_capabilities.tools if {"tools/list", "tools/call"} <= operations else None,
prompts=caller_capabilities.prompts if {"prompts/list", "prompts/get"} <= operations else None,
resources=caller_capabilities.resources if {"resources/list", "resources/read"} <= operations else None,
extensions={
key: value for key, value in (caller_capabilities.extensions or {}).items() if key in extensions
}
or None,
),
instructions=instructions,
cache_scope="private",
ttl_ms=0,
)
class GatewayVersionPolicy:
def __init__(self, versions: Callable[[], tuple[str, ...]] = configured_versions) -> None:
self._versions = versions
async def __call__(self, ctx: ServerRequestContext[object, object], call_next: CallNext) -> HandlerResult:
versions: Final = self._versions()
requested: Final = (
InitializeRequestParams.model_validate(ctx.params or {}).protocol_version
if ctx.method == "initialize"
else ctx.protocol_version
)
negotiated: Final = (
(requested if requested in HANDSHAKE_PROTOCOL_VERSIONS else LATEST_HANDSHAKE_VERSION)
if ctx.method == "initialize"
else requested
)
if negotiated not in versions:
raise MCPError(code=-32022, message="Unsupported MCP protocol version", data={"supported": list(versions)})
result: Final = await call_next(ctx)
if ctx.method != "initialize":
return result
initialized: Final = InitializeResult.model_validate(result)
discovery: Final = build_discovery(
configured=versions,
revision=initialized.protocol_version,
transport=MCPTransport.http,
authorized_operations=GATEWAY_OPERATIONS,
upstream_versions=frozenset(HANDSHAKE_PROTOCOL_VERSIONS),
capabilities=initialized.capabilities,
instructions=initialized.instructions,
)
return initialized.model_copy(update={"capabilities": discovery.capabilities})

View file

@ -28,6 +28,7 @@ class OperationContext:
client_ip: str | None = None
mcp_proxy_mode: bool = False
wire_compat: WireCompat = WireCompat.LEGACY
protocol_version: str | None = None
def __post_init__(self) -> None:
object.__setattr__(self, "_caller", copy_caller(self._caller))

View file

@ -193,6 +193,7 @@ from litellm.types.mcp import (
MCPAuth,
MCPStdioConfig,
MCPTokenEndpointAuthMethod,
MCPUpstreamProtocol,
has_header,
without_header,
)
@ -340,6 +341,7 @@ class MCPServerConfig(TypedDict, total=False):
whatever the admin wrote, and each read applies its own default."""
server_id: ReadOnly[str]
protocol_version: ReadOnly[MCPUpstreamProtocol]
alias: str
description: str
mcp_info: MCPInfo
@ -2549,6 +2551,9 @@ class MCPServerManager:
new_server = MCPServer(
server_id=server_id,
name=name_for_prefix,
protocol_version=TypeAdapter(MCPUpstreamProtocol).validate_python(
server_config.get("protocol_version", mcp_info.get("protocol_version", "auto"))
),
alias=alias,
server_name=server_name,
spec_path=server_config.get("spec_path", None),
@ -3109,6 +3114,9 @@ class MCPServerManager:
new_server: Final = MCPServer(
server_id=mcp_server.server_id,
name=name_for_prefix,
protocol_version=TypeAdapter(MCPUpstreamProtocol).validate_python(
_mcp_info.get("protocol_version", "auto")
),
alias=getattr(mcp_server, "alias", None),
server_name=getattr(mcp_server, "server_name", None),
url=mcp_server.url,
@ -4145,6 +4153,7 @@ class MCPServerManager:
cred_provider: UpstreamCredentialProvider | None = None,
raw_headers: Mapping[str, str] | None = None,
client_ip: str | None = None,
protocol_version_override: MCPUpstreamProtocol | None = None,
) -> MCPClient:
"""
Create an MCPClient instance for the given server.
@ -4168,6 +4177,9 @@ class MCPServerManager:
"""
record_auth_resolution(server.server_id, AuthResolution.unresolved)
resolved_server: Final = await self.ensure_oauth_metadata_discovered(server)
protocol_version: Final = (
protocol_version_override if protocol_version_override is not None else resolved_server.protocol_version
)
transport: Final = resolved_server.transport or MCPTransport.sse
spec = None if transport == MCPTransport.stdio else _to_server_spec_fail_closed(resolved_server)
provider: Final = cred_provider or self._cred_provider
@ -4249,6 +4261,7 @@ class MCPServerManager:
return MCPClient(
server_url="", # Not used for stdio
transport_type=transport,
protocol_version=protocol_version,
auth_type=resolved_server.auth_type,
auth_value=auth_value,
timeout=(resolved_server.timeout if resolved_server.timeout is not None else MCP_CLIENT_TIMEOUT),
@ -4281,6 +4294,7 @@ class MCPServerManager:
MCPClient(
server_url=server_url,
transport_type=transport,
protocol_version=protocol_version,
auth_type=resolved_server.auth_type,
timeout=(
resolved_server.timeout if resolved_server.timeout is not None else MCP_CLIENT_TIMEOUT
@ -4324,6 +4338,7 @@ class MCPServerManager:
MCPClient(
server_url=server_url,
transport_type=transport,
protocol_version=protocol_version,
auth_type=resolved_server.auth_type,
auth_value=auth_value,
auth_header_name=auth_header_name,

View file

@ -14,6 +14,8 @@ from mcp.types import (
CallToolRequest,
CallToolRequestParams,
CallToolResult,
DiscoverRequest,
DiscoverResult,
GetPromptRequest,
GetPromptRequestParams,
GetPromptResult,
@ -28,10 +30,14 @@ from mcp.types import (
ListToolsResult,
PaginatedRequestParams,
Prompt,
PromptsCapability,
ReadResourceRequest,
ReadResourceRequestParams,
ResourcesCapability,
ResourceTemplate,
ServerCapabilities,
TextContent,
ToolsCapability,
)
from mcp.types import Tool as MCPTool
from pydantic import AnyUrl, ConfigDict, Field, TypeAdapter
@ -51,6 +57,11 @@ from litellm.proxy._experimental.mcp_server.byok_credential_cache import (
cache_byok_credential,
get_cached_byok_credential,
)
from litellm.proxy._experimental.mcp_server.capabilities import (
GATEWAY_OPERATIONS,
build_discovery,
configured_versions,
)
from litellm.proxy._experimental.mcp_server.contracts import (
AuthorizedToolCall,
OperationContext,
@ -122,7 +133,9 @@ from litellm.proxy.litellm_pre_call_utils import (
)
from litellm.types.mcp import (
DEFAULT_CREDENTIAL_HEADER,
MCP_LEGACY_VERSIONS,
MCPAuth,
MCPTransport,
without_header,
)
from litellm.types.mcp_server.mcp_server_manager import MCPInfo, MCPServer
@ -2657,7 +2670,11 @@ class _McpDeniedDetail(TypedDict):
async def _execute_handle_list_tools(
context: OperationContext, params: PaginatedRequestParams, host_progress_callback: ProgressCallback | None = None
context: OperationContext,
params: PaginatedRequestParams,
host_progress_callback: ProgressCallback | None = None,
*,
log_list_tools_to_spendlogs: bool = True,
) -> ListToolsResult:
try:
(
@ -2700,7 +2717,7 @@ async def _execute_handle_list_tools(
mcp_server_auth_headers=mcp_server_auth_headers,
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
log_list_tools_to_spendlogs=True,
log_list_tools_to_spendlogs=log_list_tools_to_spendlogs,
list_tools_log_source="mcp_protocol",
client_ip=_client_ip,
)
@ -3065,6 +3082,7 @@ def prepare_context(
client_ip: str | None = None,
mcp_proxy_mode: bool = False,
wire_compat: WireCompat = WireCompat.LEGACY,
protocol_version: str | None = None,
) -> OperationContext:
return OperationContext(
_caller=user_api_key_auth,
@ -3076,11 +3094,13 @@ def prepare_context(
client_ip=client_ip,
mcp_proxy_mode=mcp_proxy_mode,
wire_compat=wire_compat,
protocol_version=protocol_version,
)
GatewayOperation: TypeAlias = (
AuthorizedToolCall
| DiscoverRequest
| ListToolsRequest
| CallToolRequest
| ListPromptsRequest
@ -3090,7 +3110,8 @@ GatewayOperation: TypeAlias = (
| ReadResourceRequest
)
GatewayResult: TypeAlias = (
ListToolsResult
DiscoverResult
| ListToolsResult
| CallToolResult
| InputRequiredResult
| ListPromptsResult
@ -3105,6 +3126,9 @@ class GatewayOperations:
def __init__(self, host_progress_callback: ProgressCallback | None = None) -> None:
self._host_progress_callback = host_progress_callback
@overload
async def execute(self, operation: DiscoverRequest, context: OperationContext) -> DiscoverResult: ...
@overload
async def execute(
self, operation: AuthorizedToolCall, context: OperationContext
@ -3137,6 +3161,51 @@ class GatewayOperations:
async def execute(self, operation: GatewayOperation, context: OperationContext) -> GatewayResult:
match operation:
case DiscoverRequest():
listings: Final = (
()
if context.mcp_proxy_mode
else (ListPromptsRequest(), ListResourcesRequest(), ListResourceTemplatesRequest())
)
tasks: Final = (
asyncio.create_task(
_execute_handle_list_tools(
context,
PaginatedRequestParams(),
self._host_progress_callback,
log_list_tools_to_spendlogs=False,
)
),
*(asyncio.create_task(self.execute(listing, context)) for listing in listings),
)
try:
results: Final = await asyncio.gather(*tasks)
finally:
for task in tasks:
task.cancel()
await asyncio.gather(*tasks, return_exceptions=True)
return build_discovery(
configured=configured_versions(),
revision=context.protocol_version or "2025-11-25",
transport=MCPTransport.http,
authorized_operations=GATEWAY_OPERATIONS,
upstream_versions=frozenset(MCP_LEGACY_VERSIONS),
capabilities=ServerCapabilities(
tools=ToolsCapability()
if any(isinstance(result, ListToolsResult) and result.tools for result in results)
else None,
prompts=PromptsCapability()
if any(isinstance(result, ListPromptsResult) and result.prompts for result in results)
else None,
resources=ResourcesCapability()
if any(
(isinstance(result, ListResourcesResult) and result.resources)
or (isinstance(result, ListResourceTemplatesResult) and result.resource_templates)
for result in results
)
else None,
),
)
case AuthorizedToolCall():
auth, token, _servers, server_headers, oauth_headers, headers, _client_ip = context.legacy_auth()
return await _execute_mcp_tool(

View file

@ -1375,7 +1375,16 @@ if MCP_AVAILABLE:
and headers.get(MCPRequestHandler.LITELLM_API_KEY_HEADER_NAME_PRIMARY)
else None
)
return _StagedServerTest(request=request, mcp_auth_header=mcp_auth_header, oauth2_headers=oauth2_headers)
preview_request: Final = (
request.model_copy(
update={"mcp_info": {**(request.mcp_info or {}), "protocol_version": saved_server.protocol_version}}
)
if saved_server is not None and "protocol_version" not in (request.mcp_info or {})
else request
)
return _StagedServerTest(
request=preview_request, mcp_auth_header=mcp_auth_header, oauth2_headers=oauth2_headers
)
async def _list_tools_within(client: MCPClient, deadline: float) -> list[MCPTool] | None:
with anyio.move_on_after(deadline):
@ -1512,6 +1521,7 @@ if MCP_AVAILABLE:
extra_headers=merged_headers,
stdio_env=stdio_env,
cred_provider=preview_cred_provider,
protocol_version_override=server_model.protocol_version,
)
return await operation(client)

View file

@ -125,12 +125,14 @@ def unsupported_protocol_version(scope: Scope) -> str | None:
``HANDSHAKE_PROTOCOL_VERSIONS`` to the modern single-exchange path, which
bypasses litellm's session/auth model, so the ASGI entry rejects it.
"""
from litellm.proxy._experimental.mcp_server.capabilities import configured_versions
headers: Final[Iterable[tuple[bytes, bytes]]] = scope.get("headers") or ()
values: Final = tuple(
raw.decode("latin-1").strip() for key, raw in headers if key.lower() == _MCP_PROTOCOL_VERSION_HEADER
)
for value in values:
if value and value not in HANDSHAKE_PROTOCOL_VERSIONS:
if value and value not in configured_versions():
return value
return None
@ -149,7 +151,10 @@ try:
from mcp.server.session import ServerSession as _McpServerSession
from mcp.types import (
BlobResourceContents,
DiscoverRequest,
DiscoverResult,
GetPromptResult,
RequestParams,
ResourceTemplate,
TextResourceContents,
)
@ -526,11 +531,11 @@ if MCP_AVAILABLE:
PaginatedRequestParams,
ReadResourceRequestParams,
)
from mcp_types.version import HANDSHAKE_PROTOCOL_VERSIONS
from litellm.proxy._experimental.mcp_server.auth.litellm_auth_handler import (
MCPAuthenticatedUser,
)
from litellm.proxy._experimental.mcp_server.capabilities import GatewayVersionPolicy, configured_versions
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
MCPServerManager,
global_mcp_server_manager,
@ -585,6 +590,7 @@ if MCP_AVAILABLE:
name=LITELLM_MCP_SERVER_NAME,
version=LITELLM_MCP_SERVER_VERSION,
)
server.middleware.append(GatewayVersionPolicy())
server.create_initialization_options = types.MethodType(_gateway_create_initialization_options, server)
sse: Final[SseServerTransport] = SseServerTransport("/sse/messages")
@ -830,6 +836,7 @@ if MCP_AVAILABLE:
client_ip,
_mcp_proxy_mode.get(),
wire_compat_for(ctx.protocol_version),
ctx.protocol_version,
)
async def handle_list_tools(ctx: ServerRequestContext, params: PaginatedRequestParams) -> ListToolsResult:
@ -948,6 +955,11 @@ if MCP_AVAILABLE:
ReadResourceRequest(params=params), context
)
async def discover(ctx: ServerRequestContext, params: RequestParams) -> DiscoverResult:
async with _legacy_operation_context(ctx, trace=False) as context:
return await operations.GatewayOperations().execute(DiscoverRequest(params=params), context)
server.add_request_handler("server/discover", RequestParams, discover)
server.add_request_handler("tools/list", PaginatedRequestParams, handle_list_tools)
server.add_request_handler("tools/call", CallToolRequestParams, mcp_server_tool_call)
server.add_request_handler("prompts/list", PaginatedRequestParams, list_prompts)
@ -1954,7 +1966,7 @@ if MCP_AVAILABLE:
reject_disallowed_mcp_origin(StarletteRequest(scope))
bad_version: Final = unsupported_protocol_version(scope)
if bad_version is not None:
supported: Final = ", ".join(sorted(HANDSHAKE_PROTOCOL_VERSIONS))
supported: Final = ", ".join(configured_versions())
await JSONResponse(
status_code=400,
content={ # mutable-ok: JSON-RPC error payload
@ -2299,7 +2311,7 @@ if MCP_AVAILABLE:
reject_disallowed_mcp_origin(StarletteRequest(scope))
bad_version: Final = unsupported_protocol_version(scope)
if bad_version is not None:
supported: Final = ", ".join(sorted(HANDSHAKE_PROTOCOL_VERSIONS))
supported: Final = ", ".join(configured_versions())
await JSONResponse(
status_code=400,
content={ # mutable-ok: JSON-RPC error payload

View file

@ -37,6 +37,7 @@ from litellm.types.llms.openai import (
ResponsesAPIResponse,
)
from litellm.types.mcp import (
MCPAdvertisedVersions,
MCPAllowedClient,
MCPAuth,
MCPAuthType,
@ -2998,6 +2999,11 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
None,
description="Custom CIDR ranges that define internal/private networks for MCP access control. When set, only these ranges are treated as internal. Defaults to RFC 1918 private ranges (10.0.0.0/8, 172.16.0.0/12, 192.168.0.0/16, 127.0.0.0/8).",
)
mcp_advertised_versions: MCPAdvertisedVersions | None = Field(
None,
description="MCP revisions enabled by the gateway. Defaults to all completed legacy revisions. "
"Modern protocol serving and Apps/Tasks remain disabled.",
)
mcp_allowed_clients: list[MCPAllowedClient] | None = Field(
None,
description="MCP client applications admitted by the gateway, each an {alias, value} pair where alias is the name shown in the dashboard and logs and value is the identity that must match exactly. When set, every MCP request must carry a client identity equal to one of the values: a JWT caller is identified by the claim named in litellm_jwtauth.mcp_client_id_jwt_field, any other caller by the header named in mcp_client_id_header. A request with no resolvable identity, or an unlisted one, is rejected with 403. Unset means every client is admitted.",

View file

@ -2781,15 +2781,6 @@ class ProxyBaseLLMRequestProcessing:
request=request,
)
if route_type == "aresponses":
# Streaming /v1/responses returns here without
# reaching the non-streaming ownership tail below.
# Wrap the SSE generator so container ownership is
# written once the upstream iterator finishes
# assembling ``completed_response`` — otherwise
# code-interpreter containers created during the
# stream stay unregistered and follow-up file API
# calls 403. Covers the background-polling path
# too, which loops ``body_iterator`` end-to-end.
selected_data_generator = (
ProxyBaseLLMRequestProcessing._wrap_responses_stream_for_container_ownership(
original_stream_response=response,
@ -3011,50 +3002,50 @@ class ProxyBaseLLMRequestProcessing:
wrapped_generator: Any,
user_api_key_dict: UserAPIKeyAuth,
):
"""Forward SSE chunks, then record container ownership at stream end.
"""Forward SSE chunks and record container ownership before the terminal chunk goes out.
Streaming ``/v1/responses`` short-circuits out of
``base_process_llm_request`` before the non-streaming ownership
tail runs, so without this wrap the
``LiteLLM_ManagedObjectTable`` row for any container created
during the stream is never written and follow-up file API calls
return 403.
tail runs. The OpenAI SDK closes the connection at ``data: [DONE]``
and starlette cancels the body task on disconnect, so a write that
waits for the generator to finish never lands. The iterator sets
``completed_response`` before it hands over its terminal chunk, so
the ``LiteLLM_ManagedObjectTable`` row is written the moment it
appears, ahead of the chunk carrying ``response.completed``.
"""
try:
async for chunk in wrapped_generator:
async for chunk in wrapped_generator:
completed_obj = ProxyBaseLLMRequestProcessing._extract_completed_responses_response(
original_stream_response
)
if completed_obj is None:
yield chunk
finally:
try:
completed_obj: Final = ProxyBaseLLMRequestProcessing._extract_completed_responses_response(
original_stream_response
)
if completed_obj is not None:
await ProxyBaseLLMRequestProcessing._record_container_owners_from_responses_if_needed(
response=completed_obj,
user_api_key_dict=user_api_key_dict,
)
else:
# Silent skip caused #30210: the proxy's Router wrapper
# of the responses streaming iterator wasn't propagating
# ``completed_response``, so this hook recorded nothing
# and follow-up /v1/containers/<id>/files calls 403'd
# for non-admin keys with no proxy-side hint. Log a
# warning so future regressions of the same shape
# surface in operator logs.
verbose_proxy_logger.warning(
"Container ownership recording skipped on streaming "
"/v1/responses: no completed_response on stream "
"iterator %s. If this stream created any tool "
"container (e.g. code_interpreter), follow-up "
"/v1/containers/<id>/files calls will 403 for "
"non-admin keys.",
type(original_stream_response).__name__,
)
except Exception as e:
verbose_proxy_logger.exception(
"Container ownership recording failed after streaming responses call: %s",
e,
)
continue
await ProxyBaseLLMRequestProcessing._record_container_owners_from_responses_if_needed(
response=completed_obj,
user_api_key_dict=user_api_key_dict,
)
yield chunk
async for remaining_chunk in wrapped_generator:
yield remaining_chunk
return
late_completed_obj: Final = ProxyBaseLLMRequestProcessing._extract_completed_responses_response(
original_stream_response
)
if late_completed_obj is not None:
await ProxyBaseLLMRequestProcessing._record_container_owners_from_responses_if_needed(
response=late_completed_obj,
user_api_key_dict=user_api_key_dict,
)
return
verbose_proxy_logger.warning(
"Container ownership recording skipped on streaming "
"/v1/responses: no completed_response on stream "
"iterator %s. If this stream created any tool "
"container (e.g. code_interpreter), follow-up "
"/v1/containers/<id>/files calls will 403 for "
"non-admin keys.",
type(original_stream_response).__name__,
)
async def base_passthrough_process_llm_request(
self,

View file

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

View file

@ -6342,6 +6342,11 @@ class ProxyConfig:
if general_settings is None:
general_settings = {}
if general_settings.get("mcp_advertised_versions") is not None:
from litellm.types.mcp import MCPAdvertisedVersions
TypeAdapter(MCPAdvertisedVersions).validate_python(general_settings["mcp_advertised_versions"])
if os.getenv("NUM_WORKERS", "1") != "1" and redis_usage_cache is None:
warn_login_counters_are_per_worker(os.getenv("NUM_WORKERS", "1"))
if declared_proxy_ranges(general_settings) is None:

View file

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

View file

@ -4,7 +4,7 @@ import enum
import re
from collections.abc import Awaitable, Callable, Mapping
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Literal
from typing import TYPE_CHECKING, Annotated, Any, Final, Literal
from urllib.parse import urlsplit
import httpx
@ -34,6 +34,8 @@ class MCPSpecVersion(str, enum.Enum):
nov_2024 = "2024-11-05"
mar_2025 = "2025-03-26"
jun_2025 = "2025-06-18"
nov_2025 = "2025-11-25"
jul_2026 = "2026-07-28"
class MCPAuth(str, enum.Enum):
@ -59,7 +61,17 @@ DEFAULT_SUBJECT_TOKEN_TYPE: Final = "urn:ietf:params:oauth:token-type:access_tok
# MCP Literals
MCPTransportType = Literal[MCPTransport.sse, MCPTransport.http, MCPTransport.stdio]
MCPSpecVersionType = Literal[MCPSpecVersion.nov_2024, MCPSpecVersion.mar_2025, MCPSpecVersion.jun_2025]
MCPLegacyVersion = Literal["2024-11-05", "2025-03-26", "2025-06-18", "2025-11-25"]
MCP_LEGACY_VERSIONS: Final[tuple[MCPLegacyVersion, ...]] = ("2024-11-05", "2025-03-26", "2025-06-18", "2025-11-25")
MCPUpstreamProtocol = MCPLegacyVersion | Literal["auto"]
MCPAdvertisedVersions = Annotated[tuple[MCPLegacyVersion, ...], Field(min_length=1)]
MCPSpecVersionType = Literal[
MCPSpecVersion.nov_2024,
MCPSpecVersion.mar_2025,
MCPSpecVersion.jun_2025,
MCPSpecVersion.nov_2025,
MCPSpecVersion.jul_2026,
]
MCPAuthType = (
Literal[
MCPAuth.none,

View file

@ -1,7 +1,7 @@
from datetime import datetime
from typing import Any, Final, Literal
from typing import Annotated, Any, Final, Literal
from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator
from pydantic import AfterValidator, BaseModel, ConfigDict, Field, TypeAdapter, field_validator, model_validator
from typing_extensions import Self
from litellm.types.mcp import (
@ -10,11 +10,19 @@ from litellm.types.mcp import (
MCPAuthType,
MCPTokenEndpointAuthMethod,
MCPTransportType,
MCPUpstreamProtocol,
normalize_upstream_header_name,
)
# MCPInfo now allows arbitrary additional fields for custom metadata
MCPInfo = dict[str, Any]
def _validate_mcp_protocol_metadata(value: dict[str, object]) -> dict[str, object]:
if "protocol_version" in value:
TypeAdapter(MCPUpstreamProtocol).validate_python(value["protocol_version"])
return value
MCPInfo = Annotated[dict[str, Any], AfterValidator(_validate_mcp_protocol_metadata)]
class MCPOAuthMetadata(BaseModel):
@ -66,6 +74,7 @@ class MCPServer(BaseModel):
server_name: str | None = None
url: str | None = None
transport: MCPTransportType
protocol_version: MCPUpstreamProtocol = "auto"
spec_path: str | None = None
auth_type: MCPAuthType | None = None
authentication_token: str | None = None
@ -246,6 +255,14 @@ class MCPServer(BaseModel):
"""
return self.oauth2_flow == "client_credentials"
@model_validator(mode="after")
def resolve_protocol_version(self) -> Self:
if "protocol_version" not in self.model_fields_set and self.mcp_info is not None:
self.protocol_version = TypeAdapter(MCPUpstreamProtocol).validate_python(
self.mcp_info.get("protocol_version", "auto")
)
return self
@model_validator(mode="after")
def validate_identity_binding_mode(self) -> Self:
binding: Final = self.oauth_identity_binding

View file

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

View file

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

View file

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

View file

@ -85,6 +85,7 @@
- {id: llm.responses.azure_openai.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: responses, route: azure_openai, capability: basic, streaming: nonstream, assertions: [works], source: "response_api_endpoints/endpoints.py:26", rationale: "Responses w/ Azure OpenAI (smoke)"}
- {id: llm.responses.azure_openai.tool_use.nonstream.works, module: llm, tier: P1, subject_endpoint: responses, route: azure_openai, capability: tool_use, streaming: nonstream, assertions: [works], source: "model_prices json", rationale: "Responses tool calls w/ Azure OpenAI"}
- {id: llm.responses.azure_openai.code_interpreter.nonstream.works, module: llm, tier: P1, subject_endpoint: responses, route: azure_openai, capability: code_interpreter, streaming: nonstream, assertions: [works], source: "llm_translation/test_containers_e2e.py", rationale: "An implicit code_interpreter container on an Azure deployment that carries its own api_base must serve GET /v1/containers/{id}/files/{fid}/content by its native cntr_ id to a team service-account key, the customer's shape (#27921, #28990)", fail_before_fix: proven}
- {id: llm.responses.azure_openai.code_interpreter.stream.works, module: llm, tier: P1, subject_endpoint: responses, route: azure_openai, capability: code_interpreter, streaming: stream, assertions: [works], source: "llm_translation/test_containers_e2e.py", rationale: "A container created by a streamed /v1/responses code_interpreter call must serve /v1/containers/{id}/files to the same service-account key right after the OpenAI SDK closes at [DONE]; the ownership row used to be written after the stream and the disconnect cancelled it (LIT-8612)", fail_before_fix: proven}
- {id: llm.chat_completions.together_ai.thinking.nonstream.works, module: llm, tier: P1, subject_endpoint: chat_completions, route: together_ai, capability: thinking, streaming: nonstream, assertions: [works], source: "llm_translation/test_together_ai_e2e.py", rationale: "Together reasoning surfaces as reasoning_content (LIT-5960)"}
- {id: llm.chat_completions.together_ai.thinking.stream.works, module: llm, tier: P1, subject_endpoint: chat_completions, route: together_ai, capability: thinking, streaming: stream, assertions: [works], source: "llm_translation/test_together_ai_e2e.py", rationale: "Together reasoning deltas stream as reasoning_content"}
- {id: llm.chat_completions.together_ai.thinking.nonstream.template_kwargs_forwarded, module: llm, tier: P1, subject_endpoint: chat_completions, route: together_ai, capability: thinking, streaming: nonstream, assertions: [template_kwargs_forwarded], source: "llm_translation/test_together_ai_e2e.py", rationale: "chat_template_kwargs reaches Together and turns thinking off"}

View file

@ -48,7 +48,7 @@ most likely to silently break and the one a mock can't prove works.
|----------|---------------|-----------|------------|-------------|--------|
| Chat | live (spend suite) | live (spend suite) | gap | live | partial |
| Embeddings | live (spend suite) | n/a | n/a | live | covered |
| Responses (Azure code_interpreter container files) | live | gap | live | gap | partial |
| Responses (Azure code_interpreter container files) | live | live | live | gap | partial |
| Image / audio / rerank / realtime | - | - | - | - | gap |
## This suite's files
@ -63,6 +63,7 @@ most likely to silently break and the one a mock can't prove works.
| `test_anthropic_passthrough_tool_call_logs_cost` | anthropic native, tool call, cost |
| `test_vertex_passthrough_via_managed_model_logs_cost` | vertex_ai native, non-stream, cost |
| `test_service_account_key_reads_container_file_by_native_id` | azure responses code_interpreter, non-stream, native container id, service-account key |
| `test_service_account_key_reads_container_file_created_by_a_streamed_response` | azure responses code_interpreter, stream, native container id, service-account key, upload right after `[DONE]` |
Vertex keeps the credential on the proxy like gemini/anthropic, but the deployment is
added at runtime instead of declared in the gateway config: the test POSTs `/model/new`

View file

@ -31,10 +31,11 @@ A proxy whose env carries ``AZURE_API_BASE`` for the same resource masks the
second regression, since the global-credential fallback then reaches the
container anyway.
The streaming variant is not here: a streamed ``/v1/responses`` writes the
container ownership row only after the ``[DONE]`` frame, and the OpenAI SDK
closes the connection at ``[DONE]``, so the write is cancelled and every
follow-up container call 403s (LIT-8612). That cell comes with its fix.
The streaming cell repeats the flow with ``stream=True`` and uploads right
after the last event. The OpenAI SDK closes the connection at ``[DONE]``, so an
ownership row written after the stream is cancelled with the body task and every
follow-up container call 403s (LIT-8612); the row has to land before the
``response.completed`` frame goes out.
"""
from __future__ import annotations
@ -52,7 +53,7 @@ from lifecycle import ResourceManager
from management.management_client import ManagementClient, build_client
from models import KeyGenerateBody, KeyGenerateResponse, LiteLLMParamsBody, TeamNewBody, UserNewBody
from openai import OpenAI
from openai.types.responses import Response, ResponseCodeInterpreterToolCall
from openai.types.responses import Response, ResponseCodeInterpreterToolCall, ResponseCompletedEvent
from openai.types.responses.tool_param import CodeInterpreter
from proxy_client import ProxyClient
from sdk_clients import NO_PROXY_CACHE, SdkClients
@ -120,6 +121,24 @@ def _response_with_code_interpreter(client: OpenAI, model: str) -> Response:
)
def _streamed_response_with_code_interpreter(client: OpenAI, model: str) -> Response:
events: Final = tuple(
client.with_options(timeout=CODE_INTERPRETER_TIMEOUT).responses.create(
model=model,
input=PROMPT,
tools=[CODE_INTERPRETER],
tool_choice="required",
stream=True,
extra_body=NO_PROXY_CACHE,
)
)
assert events, "responses stream returned no events"
assert isinstance(events[-1], ResponseCompletedEvent), (
f"responses stream did not terminate with response.completed: {events[-1].type}"
)
return events[-1].response
def _container_id(response: Response) -> str:
calls: Final = tuple(item for item in response.output if isinstance(item, ResponseCodeInterpreterToolCall))
assert calls, f"no code_interpreter_call in the responses output: {response.output!r}"
@ -165,3 +184,17 @@ class TestAzureContainerFiles:
f"container id is not the provider's own id: {native_id}"
)
_assert_file_round_trip(client, native_id, marker)
@pytest.mark.covers("llm.responses.azure_openai.code_interpreter.stream.works")
def test_service_account_key_reads_container_file_created_by_a_streamed_response(
self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients
) -> None:
marker: Final = unique_marker()
model: Final = _register_two_azure_deployments(proxy, resources, marker)
key: Final = _service_account_key(proxy, resources, build_client(proxy), marker, model)
client: Final = sdk.openai(key)
native_id: Final = _native_container_id(
_container_id(_streamed_response_with_code_interpreter(client, model))
)
resources.defer(lambda: client.containers.delete(native_id, extra_query=AZURE_PROVIDER_QUERY))
_assert_file_round_trip(client, native_id, marker)

View file

@ -84,3 +84,42 @@ def test_jsonrpc_error_and_malformed_tool_result_remain_errors(gateway: Gateway)
control: Final = call_tool(gateway, key, identity, names["add"], {"a": 3, "b": 5})
assert control.status_code == 200 and control.json()["isError"] is False, control.text
assert control.json()["content"][0]["text"] == "8"
@pytest.mark.parametrize("ingress", ("http", "sse"))
def test_configured_revision_blocks_unadvertised_handshake_and_keeps_allowed_control(
gateway: Gateway, tmp_path, ingress: str
) -> None:
import asyncio
from pathlib import Path
import yaml
from integration._support.mcp import mcp_peer
from integration._support.process import owned_proxy
from litellm.experimental_mcp_client.client import MCPClient
from litellm.types.mcp import MCPTransport
from mcp import MCPError
from mcp.types import CallToolRequestParams
with mcp_peer() as upstream, gateway.scenario() as scenario:
alias: Final = "restricted" + uuid.uuid4().hex[:8]
identity: Final = register_mcp(scenario, upstream, alias)
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
config["general_settings"]["mcp_advertised_versions"] = ["2024-11-05"]
config_path: Final = tmp_path / "restricted.yaml"
config_path.write_text(yaml.safe_dump(config))
with owned_proxy(gateway, tmp_path, {"DISABLE_SCHEMA_UPDATE": "true"}, config=config_path) as restricted:
endpoint: Final = str(restricted.client.base_url).rstrip("/") + ("/mcp/sse" if ingress == "sse" else "/mcp")
headers: Final = {"Authorization": f"Bearer {key}", "x-mcp-servers": identity}
denied: Final = MCPClient(server_url=endpoint, transport_type=MCPTransport(ingress), protocol_version="2025-11-25", extra_headers=headers)
allowed: Final = MCPClient(server_url=endpoint, transport_type=MCPTransport(ingress), protocol_version="2024-11-05", extra_headers=headers)
async def exercise() -> None:
with pytest.raises(MCPError, match="Unsupported MCP protocol version"):
await denied.list_tools(raise_on_error=True)
assert f"{alias}-add" in tuple(tool.name for tool in await allowed.list_tools(raise_on_error=True))
result: Final = await allowed.call_tool(CallToolRequestParams(name=f"{alias}-add", arguments={"a": 2, "b": 5}))
assert result.is_error is False and result.content[0].text == "7"
asyncio.run(exercise())

View file

@ -155,3 +155,43 @@ def test_server_initiated_sampling_and_elicitation_surface_as_errors_not_success
assert outcome.error is not None, outcome.raw
assert outcome.text is None or not outcome.text.startswith(("sampled:", "elicited:")), outcome.raw
assert len(tool_calls(peer.drain())) == 1
@pytest.mark.parametrize("downstream", ("2024-11-05", "2025-03-26", "2025-06-18", "2025-11-25"))
@pytest.mark.parametrize("upstream", ("2024-11-05", "2025-03-26", "2025-06-18", "2025-11-25"))
@pytest.mark.parametrize("peer_kind", ("http", "sse", "stdio"))
@pytest.mark.parametrize("ingress", ("http", "sse"))
def test_pinned_revision_pairs_list_and_call_through_gateway(
gateway: Gateway, downstream: str, upstream: str, peer_kind: PeerKind, ingress: str
) -> None:
import asyncio
from mcp.types import CallToolRequestParams
from litellm.experimental_mcp_client.client import MCPClient
from litellm.types.mcp import MCPTransport
with peer_of(peer_kind) as peer, gateway.scenario() as scenario:
alias: Final = "versions" + uuid.uuid4().hex[:8]
identity: Final = register_mcp(scenario, peer, alias, mcp_info={"protocol_version": upstream})
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
endpoint: Final = str(gateway.client.base_url).rstrip("/") + ("/mcp/sse" if ingress == "sse" else "/mcp")
client: Final = MCPClient(
server_url=endpoint, transport_type=MCPTransport(ingress), protocol_version=downstream,
extra_headers={"Authorization": f"Bearer {key}", "x-mcp-servers": identity}, timeout=15,
)
async def exercise() -> None:
tools: Final = await client.list_tools(raise_on_error=True)
assert f"{alias}-add" in tuple(tool.name for tool in tools)
result: Final = await client.call_tool(CallToolRequestParams(name=f"{alias}-add", arguments={"a": 3, "b": 4}))
assert result.is_error is False
assert result.content[0].text == "7"
peer.drain()
asyncio.run(exercise())
observed: Final = peer.drain()
negotiations: Final = tuple(item["body"] for item in observed if item["body"].get("method") == "initialize")
assert negotiations, "The operation must reach the upstream negotiation"
assert all(request["params"]["protocolVersion"] == upstream for request in negotiations), negotiations
assert len(tool_calls(observed)) == 1

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -1 +0,0 @@
# Levo integration tests

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

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