Merge remote-tracking branch 'origin/main' into litellm_fips_b6_small_primitives
Some checks failed
LiteLLM Rust / rust-lint (push) Has been cancelled
LiteLLM Rust / rust-test (push) Has been cancelled
LiteLLM Rust / rust-wheel (push) Has been cancelled
Terraform Provider / gofmt, vet, build, test (push) Has been cancelled
Terraform Provider / Provider endpoints vs proxy OpenAPI schema (push) Has been cancelled

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

# Conflicts:
#	tests/unit/proxy/_experimental/mcp_server/test_mcp_env_vars.py
This commit is contained in:
yucheng 2026-10-05 18:05:52 +00:00
commit 3a23cbb2d4
456 changed files with 19083 additions and 15192 deletions

View file

@ -38,7 +38,7 @@ commands:
parameters:
category:
type: enum
enum: ["backend", "client", "provider-harness"]
enum: ["backend", "client", "provider-harness", "redis-compat"]
default: "backend"
steps:
- run:
@ -229,6 +229,19 @@ commands:
- wait_for_service:
url: tcp://localhost:6379
timeout: "60"
install_codecov_cli:
steps:
- run:
name: Install Codecov CLI (pinned v11.3.1)
when: always
command: |
curl -sSLf -o /tmp/codecov https://cli.codecov.io/v11.3.1/linux/codecov
curl -sSLf -o /tmp/codecov.SHA256SUM https://cli.codecov.io/v11.3.1/linux/codecov.SHA256SUM
[ "$(cat /tmp/codecov.SHA256SUM)" = "ca1d64196d2d34771084afe76ea657d581bf628e31d993ff8e52ea09cc88a56d codecov" ]
(cd /tmp && sha256sum -c codecov.SHA256SUM)
chmod +x /tmp/codecov
mkdir -p "$HOME/.local/bin"
mv /tmp/codecov "$HOME/.local/bin/codecov"
start_openai_record_replay_proxy:
description: "Start the record/replay proxy (tests/_openai_record_replay_proxy.py) on host port 8090 and wait until healthy. Models whose api_base points here replay recorded provider responses, so the E2E run neither pays for nor depends on the live provider. The default upstream is OpenAI; a non-OpenAI model must point its api_base at /__recorder_upstream/<host>/ so the recorder forwards there instead of defaulting to OpenAI. Run after uv deps are synced."
steps:
@ -3279,30 +3292,164 @@ jobs:
- store_artifacts:
path: test-results
unit:
postgres_suite:
parameters:
test_path:
type: string
seed:
type: boolean
coverage_flag:
type: string
default: ""
timeout_minutes:
type: integer
machine:
image: ubuntu-2204:2024.04.1
resource_class: large
working_directory: ~/project
parallelism: 4
environment:
DATABASE_URL: postgresql://postgres:postgres@localhost:5432/litellm_test
steps:
- checkout
- skip_if_unrelated_changes
- setup_litellm_test_deps
- start_postgres:
db_name: litellm_test
image: postgres:16@sha256:e17e86066e5ef83e0952a9347f5c792b7ece00972e2aa787a6986f471b3dd3d5
- run:
name: Generate Prisma client
command: uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma
- when:
condition: << parameters.seed >>
steps:
- run:
name: Seed database schema
command: uv run --no-sync prisma db push --schema litellm/proxy/schema.prisma --accept-data-loss
- run:
name: Run << parameters.test_path >>
command: |
mkdir -p test-results
coverage_args=()
if [ -n "<< parameters.coverage_flag >>" ]; then
coverage_args=(--cov=./litellm --cov-report=xml:coverage.xml)
fi
timeout --signal=TERM << parameters.timeout_minutes >>m uv run --no-sync pytest \
<< parameters.test_path >> -vv --tb=short --durations=10 -o junit_family=xunit1 \
--junitxml=test-results/junit.xml "${coverage_args[@]}"
- when:
condition: << parameters.coverage_flag >>
steps:
- install_codecov_cli
- run:
name: Upload coverage
when: always
command: |
[ -f coverage.xml ] || { echo "no coverage.xml produced"; exit 1; }
codecov upload-process --disable-search --fail-on-error -f coverage.xml \
-F << parameters.coverage_flag >> -C "$CIRCLE_SHA1" \
-n "<< parameters.coverage_flag >>-${CIRCLE_BUILD_NUM}" --git-service github
- store_test_results:
path: test-results
- store_artifacts:
path: test-results
mcp_integration:
machine:
image: ubuntu-2204:2024.04.1
resource_class: large
working_directory: ~/project
environment:
COVERAGE_CORE: sysmon
LITELLM_LOCAL_MODEL_COST_MAP: "True"
steps:
- checkout
- skip_if_unrelated_changes
- setup_litellm_test_deps
- run:
name: Generate Prisma client
command: uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma
- run:
name: Run unit tests
name: Install MCP SDK1 peer
command: |
mkdir -p test-results/unit
shard="$(find tests/unit -name 'test_*.py' | sort | circleci tests split --split-by=timings --timings-type=filename)"
if [ -z "${shard}" ]; then echo "shard ${CIRCLE_NODE_INDEX} received no tests/unit files; nothing to run"; exit 0; fi
mapfile -t files < <(printf '%s\n' "${shard}")
set +e
LITELLM_LOCAL_MODEL_COST_MAP=True uv run --no-sync pytest "${files[@]}" -p no:rerunfailures -p no:pytest-retry --timeout=90 -n 4 --dist=loadscope --tb=short -o junit_family=xunit1 --junitxml=test-results/unit/junit.xml
status=$?
set -e
if [ "$status" -eq 5 ]; then echo "pytest collected no tests from tests/unit; passing"; exit 0; fi
exit "$status"
uv venv --python 3.12 .venv-mcp-peer
uv pip install --python .venv-mcp-peer 'mcp==1.28.1' 'langchain-mcp-adapters==0.2.1'
echo "export MCP_TEST_PEER_PYTHON=$PWD/.venv-mcp-peer/bin/python" >> "$BASH_ENV"
- run:
name: Run MCP integration tests
command: |
mkdir -p test-results
env -u OPENAI_API_KEY -u ANTHROPIC_API_KEY \
timeout --signal=TERM 20m uv run --no-sync pytest \
tests/mcp_tests tests/unit/experimental_mcp_client tests/unit/proxy/_experimental/mcp_server \
tests/unit/responses/mcp --tb=short -vv --maxfail=10 -n 2 --dist=loadscope --reruns 0 \
--reruns-delay 1 --timeout=120 --rerun-except "from pytest-timeout" --durations=20 \
--cov=./litellm --cov=./enterprise/litellm_enterprise --cov-report=xml:coverage.xml \
--cov-config=pyproject.toml -o junit_family=xunit1 --junitxml=test-results/junit.xml
- install_codecov_cli
- run:
name: Upload MCP integration coverage
when: always
command: |
[ -f coverage.xml ] || { echo "no coverage.xml produced"; exit 1; }
codecov upload-process --disable-search --fail-on-error -f coverage.xml \
-F mcp-integration -C "$CIRCLE_SHA1" \
-n "mcp-integration-${CIRCLE_BUILD_NUM}" --git-service github
- store_test_results:
path: test-results
- store_artifacts:
path: test-results
redis_compat:
parameters:
redis_py:
type: string
machine:
image: ubuntu-2204:2024.04.1
resource_class: large
working_directory: ~/project
steps:
- checkout
- skip_if_unrelated_changes:
category: redis-compat
- setup_litellm_test_deps
- run:
name: Pin redis-py version
command: |
uv pip install "redis==<< parameters.redis_py >>"
uv run --no-sync python -c "import redis; assert redis.__version__ == '<< parameters.redis_py >>', redis.__version__; print('redis-py', redis.__version__)"
- run:
name: Install redis-server 7.2.16
command: |
mkdir -p "$HOME/.local/bin"
cid="$(docker create redis:7.2.16@sha256:0637954999d01b7c9ce9167db2da50656e2590d3b884f1c600c5f63bb6e6773c)"
docker cp "${cid}":/usr/local/bin/redis-server "$HOME/.local/bin/redis-server"
docker rm "$cid"
redis-server --version | grep -q 'v=7.2.16'
- run:
name: Run Redis compatibility tests
command: |
mkdir -p test-results
env -u CASSETTE_REDIS_URL -u AZURE_CLIENT_ID -u AZURE_CLIENT_SECRET -u AZURE_TENANT_ID \
timeout --signal=TERM 15m uv run --no-sync pytest \
tests/unit/test_redis.py tests/unit/caching/test_redis_connection_pool.py \
tests/unit/caching/test_redis_cluster_cache.py tests/unit/caching/test_evicted_client_closer.py \
tests/local_testing/test_caching.py::test_sync_cluster_authenticates_with_azure_credentials \
tests/local_testing/test_caching.py::test_sync_cluster_authenticates_with_gcp_credentials \
--tb=short -vv --reruns 2 --reruns-delay 1 --durations=20 --cov=./litellm \
--cov-report=xml:coverage.xml -o junit_family=xunit1 --junitxml=test-results/junit.xml
- when:
condition:
equal: ["5.3.1", << parameters.redis_py >>]
steps:
- install_codecov_cli
- run:
name: Upload Redis compatibility coverage
when: always
command: |
[ -f coverage.xml ] || { echo "no coverage.xml produced"; exit 1; }
codecov upload-process --disable-search --fail-on-error -f coverage.xml \
-F redis-compat -C "$CIRCLE_SHA1" \
-n "redis-compat-${CIRCLE_BUILD_NUM}" --git-service github
- store_test_results:
path: test-results
- store_artifacts:
@ -3376,6 +3523,39 @@ workflows:
parameters:
suite: [management, database]
mode: [replica]
- postgres_suite:
name: proxy-behavior
test_path: tests/proxy_behavior
seed: true
coverage_flag: lens-postgres
timeout_minutes: 25
- postgres_suite:
name: proxy-security
test_path: tests/proxy_security_tests
seed: true
timeout_minutes: 15
- postgres_suite:
name: schema-migration
test_path: tests/proxy_migration_tests
seed: false
timeout_minutes: 20
- postgres_suite:
name: roi-database
test_path: tests/integration/database/test_roi_observed.py
seed: false
coverage_flag: roi-postgres
timeout_minutes: 10
- mcp_integration:
name: mcp-integration
- redis_compat:
name: redis-compat-<< matrix.redis_py >>
matrix:
parameters:
redis_py:
- "5.3.1"
- "6.4.0"
- "7.4.1"
- "8.0.1"
build_and_test:
unless:
or:
@ -3385,7 +3565,6 @@ workflows:
jobs:
- using_litellm_on_windows
- windows_release_wheel
- unit
- provider_replay_harness
- base_sdk_install
- local_testing_part1

View file

@ -1,7 +1,7 @@
#!/usr/bin/env bash
set -uo pipefail
category="${1:?usage: classify_changes.sh <backend|client|ui|provider-harness|cost-map-only|mcp-dependencies|windows-release>}"
category="${1:?usage: classify_changes.sh <backend|client|ui|provider-harness|cost-map-only|mcp-dependencies|windows-release|redis-compat>}"
has_client=false
has_backend=false
@ -10,9 +10,14 @@ has_provider_harness=false
has_cost_map=false
has_mcp_dependencies=false
has_windows_release=false
has_redis_compat=false
outside_cost_map_set=false
while IFS= read -r file || [ -n "$file" ]; do
[ -n "$file" ] || continue
case "$file" in
litellm/_redis.py | litellm/_redis_credential_provider.py | litellm/caching/redis_cache.py | litellm/caching/evicted_client_closer.py | tests/unit/test_redis.py | tests/local_testing/test_caching.py | tests/unit/caching/test_redis_connection_pool.py | tests/unit/caching/test_redis_cluster_cache.py | tests/unit/caching/test_evicted_client_closer.py | .circleci/config.yml | .circleci/scripts/classify_changes.sh | .circleci/scripts/path_filter.sh | pyproject.toml | uv.lock)
has_redis_compat=true ;;
esac
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/unit/test_circleci_path_filter.py | tests/unit/test_detect_changes.py)
@ -54,6 +59,9 @@ case "$category" in
windows-release)
[ "$has_windows_release" = true ] && echo run || echo skip
;;
redis-compat)
[ "$has_redis_compat" = true ] && echo run || echo skip
;;
backend)
[ "$has_backend" = true ] && echo run || echo skip
;;

View file

@ -1,7 +1,7 @@
#!/usr/bin/env bash
set -uo pipefail
category="${1:?usage: path_filter.sh <backend|client|provider-harness>}"
category="${1:?usage: path_filter.sh <backend|client|provider-harness|redis-compat>}"
here="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)"
run_full() {

View file

@ -1,187 +0,0 @@
#!/usr/bin/env bash
set -euo pipefail
flag="${1:?usage: unit_selection.sh <codecov flag>}"
legacy_flags=(
caching-local
core-utils
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
proxy-db-db-and-spend
proxy-db-endpoints-and-responses
proxy-db-guardrails-hooks
proxy-db-jwt-and-keys
proxy-db-key-generation
proxy-db-logging-misc
proxy-db-proxy-runtime
proxy-db-proxy-server-core
proxy-db-proxy-utils
proxy-extras
proxy-infra
responses-caching-types
)
legacy_paths() {
case "$1" in
caching-local) echo tests/unit/caching ;;
core-utils) echo tests/unit/litellm_core_utils ;;
enterprise-package)
echo tests/unit/enterprise/integrations
echo tests/unit/enterprise/proxy/auth
echo tests/unit/enterprise/proxy/guardrails
echo tests/unit/enterprise/proxy/hooks
echo tests/unit/enterprise/proxy/management_endpoints
echo tests/unit/enterprise/proxy/test_audit_logging_endpoints.py
echo tests/unit/enterprise/proxy/test_liteadmin.py
echo tests/unit/enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py ;;
enterprise-routing)
echo tests/unit/google_genai
echo tests/unit/router_strategy
echo tests/unit/router_utils
echo tests/unit/proxy/common_utils/test_cache_aware_routing.py
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
echo tests/unit/enterprise/proxy/test_batch_retrieve_registers_missing_output_file_id.py
echo tests/unit/enterprise/proxy/test_batch_retrieve_returns_unified_input_file_id.py
echo tests/unit/enterprise/proxy/test_batch_update_db_managed_output_file_id.py
echo tests/unit/enterprise/proxy/test_deleted_file_returns_403_not_404.py
echo tests/unit/enterprise/proxy/test_enterprise_routes.py
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/harness
echo tests/unit/images
echo tests/unit/interactions
echo tests/unit/messages
echo tests/unit/rag
echo tests/unit/rerank_api
echo tests/unit/rust_bridge
echo tests/unit/secret_managers
echo tests/unit/vector_stores
echo tests/unit/videos ;;
proxy-db-auth-checks)
echo tests/unit/proxy/auth/test_auth_checks.py
echo tests/unit/proxy/auth/test_user_api_key_auth.py
echo tests/unit/proxy/test_credential_slot_registry.py
echo tests/unit/proxy/test_deprecated_key_grace_period.py ;;
proxy-db-budgets)
echo tests/unit/proxy/auth/test_default_end_user_budget_simple.py
echo tests/unit/proxy/hooks/test_unit_test_max_model_budget_limiter.py
echo tests/unit/proxy/test_zero_cost_model_budget_bypass.py ;;
proxy-db-custom-logging)
echo tests/unit/proxy/test_custom_callback_input.py
echo tests/unit/proxy/test_custom_logger_s3_gcs.py ;;
proxy-db-db-and-spend)
echo tests/unit/proxy/common_utils/test_proxy_encrypt_decrypt.py
echo tests/unit/proxy/db/db_transaction_queue/test_e2e_pod_lock_manager.py
echo tests/unit/proxy/db/test_update_daily_tag_spend.py
echo tests/unit/proxy/test_db_schema_changes.py
echo tests/unit/proxy/test_prisma_client_backoff_retry.py
echo tests/unit/proxy/test_update_spend.py
echo tests/unit/skills/test_skills_db.py ;;
proxy-db-endpoints-and-responses)
echo tests/unit/proxy/lens
echo tests/unit/proxy/auth/test_models_fallback_endpoint.py
echo tests/unit/proxy/common_utils/test_check_batch_cost.py
echo tests/unit/proxy/common_utils/test_check_responses_cost.py
echo tests/unit/proxy/common_utils/test_realtime_cache.py
echo tests/unit/proxy/google_endpoints/test_gemini_agents_endpoints.py
echo tests/unit/proxy/google_endpoints/test_google_endpoint_routing.py
echo tests/unit/proxy/google_endpoints/test_google_gemini_proxy_request.py
echo tests/unit/proxy/public_endpoints/test_blog_posts_endpoint.py
echo tests/unit/proxy/response_polling
echo tests/unit/proxy/test_custom_tokenizer_bug.py
echo tests/unit/proxy/test_get_favicon.py
echo tests/unit/proxy/test_get_image.py
echo tests/unit/proxy/test_prompt_test_endpoint.py
echo tests/unit/proxy/test_reducto_ocr_route.py
echo tests/unit/proxy/test_response_polling_pre_call_checks.py
echo tests/unit/proxy/test_ui_path_detection.py ;;
proxy-db-guardrails-hooks)
echo tests/unit/proxy/hooks/test_banned_keyword_list.py
echo tests/unit/proxy/test_proxy_setting_guardrails.py
echo tests/unit/proxy/test_unit_test_proxy_hooks.py ;;
proxy-db-jwt-and-keys)
echo tests/unit/proxy/auth/test_jwt.py
echo tests/unit/proxy/management_endpoints/test_jwt_key_mapping.py
echo tests/unit/proxy/test_proxy_custom_auth.py ;;
proxy-db-key-generation) echo tests/unit/proxy/management_endpoints/test_key_generate_prisma.py ;;
proxy-db-logging-misc)
echo tests/unit/proxy/management_helpers/test_audit_logs_proxy.py
echo tests/unit/proxy/spend_tracking/test_search_api_logging.py
echo tests/unit/proxy/test_proxy_reject_logging.py ;;
proxy-db-proxy-runtime)
echo tests/unit/proxy/auth/test_multipart_bypass_repro.py
echo tests/unit/proxy/auth/test_proxy_routes.py
echo tests/unit/proxy/middleware/test_request_size_limit_middleware.py
echo tests/unit/proxy/test_proxy_config_unit_test.py
echo tests/unit/proxy/test_proxy_token_counter.py
echo tests/unit/proxy/test_server_root_path.py ;;
proxy-db-proxy-server-core)
echo tests/unit/proxy/test__lazy_features.py
echo tests/unit/proxy/test_aproxy_startup.py
echo tests/unit/proxy/test_proxy_server.py ;;
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
echo tests/unit/proxy/management
echo tests/unit/proxy/management_endpoints/test_roi_calculator_endpoints.py
echo tests/unit/proxy/roi_calculator ;;
responses-caching-types)
find tests/unit/responses -name 'test_*.py' -not -path 'tests/unit/responses/mcp/*'
echo tests/unit/types ;;
*) echo "unit_selection.sh: unknown flag $1" >&2; exit 1 ;;
esac
}
expand() {
while read -r path; do
if [ -d "$path" ]; then
find "$path" -name 'test_*.py'
elif [ -f "$path" ]; then
echo "$path"
else
echo "unit_selection.sh: $path does not exist" >&2
exit 1
fi
done
}
if [ "$flag" = unit ]; then
comm -23 \
<(find tests/unit -name 'test_*.py' | sort) \
<(for legacy in "${legacy_flags[@]}"; do legacy_paths "$legacy"; done | expand | sort)
exit 0
fi
legacy_paths "$flag" | expand | sort

View file

@ -1,416 +0,0 @@
version: 2.1
commands:
wait_for_service:
parameters:
url:
type: string
timeout:
type: string
default: "60"
steps:
- run:
name: "Wait for << parameters.url >>"
command: |
TIMEOUT=<< parameters.timeout >>
URL="<< parameters.url >>"
ELAPSED=0
echo "Waiting up to ${TIMEOUT}s for ${URL} ..."
if echo "$URL" | grep -q '^tcp://'; then
HOST=$(echo "$URL" | sed 's|tcp://||' | cut -d: -f1)
PORT=$(echo "$URL" | sed 's|tcp://||' | cut -d: -f2)
while ! bash -c "echo > /dev/tcp/$HOST/$PORT" 2>/dev/null; do
sleep 2; ELAPSED=$((ELAPSED+2))
if [ "$ELAPSED" -ge "$TIMEOUT" ]; then echo "Timed out"; exit 1; fi
done
else
while ! curl -sf --max-time 5 "$URL" > /dev/null 2>&1; do
sleep 2; ELAPSED=$((ELAPSED+2))
if [ "$ELAPSED" -ge "$TIMEOUT" ]; then echo "Timed out"; exit 1; fi
done
fi
echo "Service ready after ${ELAPSED}s"
install_uv:
steps:
- run:
name: Install uv (pinned 0.10.9)
command: |
curl -LsSf -o /tmp/uv-install.sh https://astral.sh/uv/0.10.9/install.sh
echo "7fc46e39cb97290b57169c0c813a17970585ac519139f19006453c99b5f2f45f /tmp/uv-install.sh" | sha256sum -c -
env UV_NO_MODIFY_PATH=1 sh /tmp/uv-install.sh
rm -f /tmp/uv-install.sh
echo 'export PATH="$HOME/.local/bin:$PATH"' >> "$BASH_ENV"
export PATH="$HOME/.local/bin:$PATH"
install_rust:
steps:
- run:
name: Install Rust (rustup 1.28.2, toolchain 1.98.0)
command: |
case "$(uname -m)" in
x86_64)
RUSTUP_TRIPLE=x86_64-unknown-linux-gnu
RUSTUP_SHA256=20a06e644b0d9bd2fbdbfd52d42540bdde820ea7df86e92e533c073da0cdd43c
;;
aarch64)
RUSTUP_TRIPLE=aarch64-unknown-linux-gnu
RUSTUP_SHA256=e3853c5a252fca15252d07cb23a1bdd9377a8c6f3efa01531109281ae47f841c
;;
*)
echo "install_rust: unsupported architecture $(uname -m)" >&2
exit 1
;;
esac
curl -sSLf -o /tmp/rustup-init \
"https://static.rust-lang.org/rustup/archive/1.28.2/${RUSTUP_TRIPLE}/rustup-init"
echo "${RUSTUP_SHA256} /tmp/rustup-init" | sha256sum -c -
chmod +x /tmp/rustup-init
/tmp/rustup-init -y --no-modify-path --profile minimal --default-toolchain 1.98.0
rm -f /tmp/rustup-init
echo 'export PATH="$HOME/.cargo/bin:$PATH"' >> "$BASH_ENV"
export PATH="$HOME/.cargo/bin:$PATH"
rustc --version
cargo --version
install_codecov_cli:
steps:
- run:
name: Install Codecov CLI (pinned v11.3.1)
when: always
command: |
curl -sSLf -o /tmp/codecov https://cli.codecov.io/v11.3.1/linux/codecov
curl -sSLf -o /tmp/codecov.SHA256SUM https://cli.codecov.io/v11.3.1/linux/codecov.SHA256SUM
[ "$(cat /tmp/codecov.SHA256SUM)" = "ca1d64196d2d34771084afe76ea657d581bf628e31d993ff8e52ea09cc88a56d codecov" ]
(cd /tmp && sha256sum -c codecov.SHA256SUM)
chmod +x /tmp/codecov
mkdir -p "$HOME/.local/bin"
mv /tmp/codecov "$HOME/.local/bin/codecov"
setup_litellm_enterprise_pip:
steps:
- run:
name: "Install local version of litellm-enterprise"
command: |
uv run --no-sync python -c "import litellm_enterprise; print('litellm-enterprise OK:', litellm_enterprise.__file__)"
setup_test_deps:
steps:
- install_uv
- install_rust
- restore_cache:
keys:
- v3-integration-uv-cache-{{ checksum "uv.lock" }}
- run:
name: Install Dependencies
command: |
uv sync --frozen --all-groups --all-extras --python 3.12
- setup_litellm_enterprise_pip
- save_cache:
paths:
- ~/.cache/uv
key: v3-integration-uv-cache-{{ checksum "uv.lock" }}
- run:
name: Generate Prisma client
command: uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma
skip_unless_relevant:
parameters:
category:
type: string
default: backend
base_ref:
type: string
default: ""
pull_request_url:
type: string
default: ""
steps:
- run:
name: "Skip job when no << parameters.category >>-relevant files changed"
command: |
export CIRCLE_PULL_REQUEST="${CIRCLE_PULL_REQUEST:-<< parameters.pull_request_url >>}"
export PATH_FILTER_BASE_BRANCH="<< parameters.base_ref >>"
[ -n "$PATH_FILTER_BASE_BRANCH" ] || unset PATH_FILTER_BASE_BRANCH
bash .circleci/scripts/path_filter.sh << parameters.category >>
start_postgres:
parameters:
db_name:
type: string
default: circle_test
image:
type: string
default: postgres:14@sha256:6a70deda415ec296f977890e11aba04a0db9f632a362e3fce45e845e3db74f26
steps:
- run:
name: Start PostgreSQL
command: |
docker run -d \
--name postgres-db \
-e POSTGRES_USER=postgres \
-e POSTGRES_PASSWORD=postgres \
-e POSTGRES_DB=<< parameters.db_name >> \
-p 5432:5432 \
<< parameters.image >>
- wait_for_service:
url: tcp://localhost:5432
timeout: "60"
start_redis:
steps:
- run:
name: Start Redis
command: |
docker run -d \
--name redis-cache \
-p 6379:6379 \
redis:7-alpine@sha256:7aec734b2bb298a1d769fd8729f13b8514a41bf90fcdd1f38ec52267fbaa8ee6
- wait_for_service:
url: tcp://localhost:6379
timeout: "60"
jobs:
unit:
parameters:
flag:
type: string
default: unit
shards:
type: integer
default: 6
workers:
type: integer
default: 4
dist:
type: string
default: loadscope
base_ref:
type: string
default: ""
pull_request_url:
type: string
default: ""
legacy_mcp_peer:
type: boolean
default: false
reruns:
type: integer
default: 0
machine:
image: ubuntu-2204:2024.04.1
resource_class: large
working_directory: ~/project
parallelism: << parameters.shards >>
environment:
COVERAGE_CORE: sysmon
LITELLM_LOCAL_MODEL_COST_MAP: "True"
steps:
- checkout
- skip_unless_relevant:
base_ref: << parameters.base_ref >>
pull_request_url: << parameters.pull_request_url >>
- setup_test_deps
- when:
condition: << parameters.legacy_mcp_peer >>
steps:
- run:
name: Install MCP SDK1 peer
command: |
uv venv --python 3.12 .venv-mcp-peer
uv pip install --python .venv-mcp-peer 'mcp==1.28.1' 'langchain-mcp-adapters==0.2.1'
echo "export MCP_TEST_PEER_PYTHON=$PWD/.venv-mcp-peer/bin/python" >> "$BASH_ENV"
- run:
name: "Run << parameters.flag >> shard"
no_output_timeout: 20m
command: |
mkdir -p test-results/<< parameters.flag >>
selection="$(bash .circleci/scripts/unit_selection.sh << parameters.flag >>)" || { echo "unit_selection.sh failed for << parameters.flag >>"; exit 1; }
[ -n "${selection}" ] || { echo "unit_selection.sh produced no files for << parameters.flag >>"; exit 1; }
shard="$(printf '%s\n' "${selection}" | circleci tests split --split-by=timings --timings-type=filename)" || { echo "circleci tests split failed for << parameters.flag >>"; exit 1; }
[ -n "${shard}" ] || { echo "shard ${CIRCLE_NODE_INDEX} received no << parameters.flag >> files; nothing to run"; exit 0; }
mapfile -t files < <(printf '%s\n' "${shard}")
xdist_args=()
if [ "<< parameters.workers >>" -gt 0 ]; then xdist_args=(-n << parameters.workers >> --dist=<< parameters.dist >>); fi
rerun_args=(-p no:rerunfailures)
if [ "<< parameters.reruns >>" -gt 0 ]; then rerun_args=(--reruns << parameters.reruns >> --reruns-delay 1 --rerun-except "from pytest-timeout"); fi
test_env=(PATH="$PATH" HOME="$HOME" CI=true COVERAGE_CORE="$COVERAGE_CORE" LITELLM_LOCAL_MODEL_COST_MAP="$LITELLM_LOCAL_MODEL_COST_MAP")
if [ -n "${MCP_TEST_PEER_PYTHON:-}" ]; then test_env+=(MCP_TEST_PEER_PYTHON="$MCP_TEST_PEER_PYTHON"); fi
set +e
env -i "${test_env[@]}" \
uv run --no-sync pytest "${files[@]}" "${rerun_args[@]}" -p no:pytest-retry --timeout=90 "${xdist_args[@]}" --tb=short --durations=20 -o junit_family=xunit1 --junitxml=test-results/<< parameters.flag >>/junit.xml --cov=./litellm --cov=./enterprise/litellm_enterprise --cov-report=xml:coverage.xml --cov-config=pyproject.toml
status=$?
set -e
if [ "$status" -eq 5 ]; then echo "pytest collected no tests from the shard; passing"; exit 0; fi
exit "$status"
- install_codecov_cli
- run:
name: Upload coverage
when: always
command: |
[ -f coverage.xml ] || { echo "no coverage.xml produced; skipping upload"; exit 0; }
codecov upload-process --disable-search -f coverage.xml -F << parameters.flag >> -C "$CIRCLE_SHA1" -n "<< parameters.flag >>-${CIRCLE_NODE_INDEX}-${CIRCLE_BUILD_NUM}" --git-service github
- store_test_results:
path: test-results
- store_artifacts:
path: test-results
- store_artifacts:
path: coverage.xml
documentation:
machine:
image: ubuntu-2204:2024.04.1
resource_class: large
working_directory: ~/project
steps:
- checkout
- setup_test_deps
- run:
name: Checkout litellm-docs
command: rm -rf docs/my-website && git clone --depth 1 https://github.com/BerriAI/litellm-docs.git docs/my-website
- run:
name: Run documentation validation
command: |
uv run --no-sync python ./tests/documentation_tests/test_env_keys.py
uv run --no-sync python ./tests/documentation_tests/test_router_settings.py
uv run --no-sync python ./tests/documentation_tests/test_api_docs.py
uv run --no-sync python ./tests/documentation_tests/test_circular_imports.py
integration:
parameters:
suite:
type: string
base_ref:
type: string
default: ""
pull_request_url:
type: string
default: ""
machine:
image: ubuntu-2204:2024.04.1
resource_class: large
working_directory: ~/project
steps:
- checkout
- skip_unless_relevant:
base_ref: << parameters.base_ref >>
pull_request_url: << parameters.pull_request_url >>
- setup_test_deps
- start_postgres:
image: postgres:16@sha256:e17e86066e5ef83e0952a9347f5c792b7ece00972e2aa787a6986f471b3dd3d5
- start_redis
- run:
name: Run owned integration contracts
command: env -i PATH="$PATH" HOME="$HOME" CIRCLE_SHA1="$CIRCLE_SHA1" CIRCLE_WORKFLOW_ID="$CIRCLE_WORKFLOW_ID" bash .circleci/scripts/run_integration.sh << parameters.suite >>
no_output_timeout: 15m
- run:
name: Stop owned database and Redis
when: always
command: |
mkdir -p test-results/integration-<< parameters.suite >>
docker logs postgres-db > test-results/integration-<< parameters.suite >>/postgres.log 2>&1 || true
docker logs redis-cache > test-results/integration-<< parameters.suite >>/redis.log 2>&1 || true
docker rm -f postgres-db redis-cache
test -z "$(docker ps -aq --filter name=postgres-db --filter name=redis-cache)"
- store_test_results:
path: test-results
- store_artifacts:
path: test-results
workflows:
tests:
when: (pipeline.event.name == "push" and pipeline.git.branch == "main") or pipeline.event.name == "api" or (pipeline.event.name == "pull_request" and (pipeline.event.github.pull_request.base.ref == "main" or pipeline.event.github.pull_request.base.ref starts-with "litellm_"))
jobs:
- unit:
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-<< matrix.flag >>
shards: 1
workers: 2
reruns: 2
matrix:
parameters:
flag: [caching-local, proxy-extras, enterprise-routing]
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-mcp-integration
flag: mcp-integration
shards: 1
workers: 2
legacy_mcp_peer: true
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-<< matrix.flag >>
shards: 1
reruns: 2
matrix:
parameters:
flag:
- enterprise-package
- proxy-infra
- responses-caching-types
- proxy-db-auth-checks
- proxy-db-jwt-and-keys
- proxy-db-proxy-server-core
- proxy-db-proxy-runtime
- proxy-db-custom-logging
- proxy-db-logging-misc
- proxy-db-db-and-spend
- proxy-db-guardrails-hooks
- proxy-db-budgets
- 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-core-utils
flag: core-utils
shards: 2
reruns: 1
base_ref: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.base.ref or "" >>
pull_request_url: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.url or "" >>
- unit:
name: unit-integrations
flag: integrations
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
shards: 1
reruns: 2
dist: worksteal
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-key-generation
flag: proxy-db-key-generation
shards: 1
workers: 0
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 "" >>
- documentation
- integration:
name: integration-<< matrix.suite >>
matrix:
parameters:
suite: [sdk]
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 "" >>

View file

@ -21,7 +21,7 @@ test_paths:
Live-provider caching cases in tests/local_testing that remain outside CI. Jobs that
glob that directory either deselect them (local_testing_part1 and part2 carry `-k "... and
not caching and not cache"`) or keep only another keyword (langfuse, router, assistants).
Separately, test-redis-compat.yml selects two IAM cluster authentication tests in
Separately, the CircleCI redis-compat jobs select two IAM cluster authentication tests in
test_caching.py by node ID. It does not run that file's other tests.
The gap was eight files and 118 tests when measured 2026-08-20; the five keyless files now
run in the caching-local shard, leaving live cases in these three. Measured 2026-08-21 with no provider

View file

@ -1,17 +0,0 @@
#!/usr/bin/env bash
set -uo pipefail
STACK_DIR="${E2E_STACK_DIR:-${RUNNER_TEMP:-/tmp}/litellm-e2e-stack}"
for pid_file in "${STACK_DIR}"/pids/*.pid; do
[[ -f "${pid_file}" ]] || continue
pkill -TERM -P "$(cat "${pid_file}")" 2>/dev/null
kill -TERM "$(cat "${pid_file}")" 2>/dev/null
rm -f "${pid_file}"
done
for container in e2e-nginx e2e-keycloak e2e-valkey e2e-jaeger e2e-postgres; do
docker rm -f "${container}" >/dev/null 2>&1
done
exit 0

View file

@ -1,83 +0,0 @@
import argparse
import os
import sys
from functools import reduce
from pathlib import Path
from typing import Final
from xml.sax.saxutils import escape
from pydantic import JsonValue, TypeAdapter, ValidationError
from secrets_to_env import MIN_MASKED_LENGTH
REDACTED: Final = "***"
json_adapter: Final[TypeAdapter[JsonValue]] = TypeAdapter(JsonValue)
def string_leaves(node: JsonValue) -> tuple[str, ...]:
match node:
case str():
return (node,)
case list():
return tuple(leaf for child in node for leaf in string_leaves(child))
case dict():
return tuple(leaf for child in node.values() for leaf in string_leaves(child))
return ()
def field_lines(value: str) -> tuple[str, ...]:
try:
return tuple(line for leaf in string_leaves(json_adapter.validate_json(value)) for line in leaf.splitlines())
except ValidationError:
return ()
def masked_values(values_files: tuple[Path, ...]) -> tuple[str, ...]:
values: Final = frozenset(
line.split("=", 1)[1].strip().strip("'")
for path in values_files
for line in path.read_text().splitlines()
if "=" in line
)
texts: Final = frozenset(text for value in values for text in (value, *field_lines(value)))
renderings: Final = frozenset(
rendering
for text in texts
if len(text) >= MIN_MASKED_LENGTH
for rendering in (text, escape(text), escape(text, {'"': "&quot;"}))
)
return tuple(sorted(renderings, key=lambda rendering: (-len(rendering), rendering)))
def redact(text: str, values: tuple[str, ...]) -> str:
return reduce(lambda redacted, value: redacted.replace(value, REDACTED), values, text)
def write_redacted(source: Path, out_dir: Path, values: tuple[str, ...]) -> None:
target: Final = out_dir / source.name
with os.fdopen(os.open(target, os.O_WRONLY | os.O_CREAT | os.O_EXCL | os.O_NOFOLLOW, 0o600), "w") as handle:
_ = handle.write(redact(source.read_text(errors="replace"), values))
def main() -> int:
parser: Final = argparse.ArgumentParser()
_ = parser.add_argument("--values", action="append", type=Path, required=True)
_ = parser.add_argument("--out", type=Path, required=True)
_ = parser.add_argument("files", nargs="*", type=Path)
args: Final = parser.parse_args()
values_files: Final = tuple(args.values)
out_dir: Final[Path] = args.out
sources: Final = tuple(args.files)
try:
values: Final = masked_values(values_files)
out_dir.mkdir(mode=0o700, exist_ok=True)
for source in sources:
write_redacted(source, out_dir, values)
except OSError as error:
_ = sys.stderr.write(f"could not redact {error.filename}\n")
return 1
_ = sys.stdout.write(f"redacted {len(sources)} file(s) into {out_dir}\n")
return 0
if __name__ == "__main__":
sys.exit(main())

View file

@ -1,55 +0,0 @@
import os
import re
import sys
from pathlib import Path
from typing import Final
from pydantic import TypeAdapter, ValidationError
secrets_adapter: Final[TypeAdapter[dict[str, str]]] = TypeAdapter(dict[str, str])
ENV_NAME: Final = re.compile(r"[A-Za-z_][A-Za-z0-9_]*")
MIN_MASKED_LENGTH: Final = 8
ACTIONS_RUNNER_FLAG: Final = "GITHUB_ACTIONS"
def main() -> int:
env_path: Final = Path(sys.argv[1])
try:
secrets: Final = {
key: value.rstrip("\r\n") for key, value in secrets_adapter.validate_json(sys.stdin.read()).items()
}
except (ValidationError, UnicodeError):
_ = sys.stderr.write("expected a JSON object containing string environment values\n")
return 1
unusable: Final = tuple(
key
for key, value in secrets.items()
if ENV_NAME.fullmatch(key) is None or any(char in value for char in "'\n\r\0")
)
if unusable:
_ = sys.stderr.write(
f"these names or values cannot be represented in both bash and dotenv: {' '.join(sorted(unusable))}\n"
)
return 1
if os.environ.get(ACTIONS_RUNNER_FLAG) == "true":
_ = sys.stdout.write(
"".join(
f"::add-mask::{value.replace('%', '%25')}\n"
for value in secrets.values()
if len(value) >= MIN_MASKED_LENGTH
)
)
sys.stdout.flush()
lines: Final = tuple(f"{key}='{value}'" for key, value in secrets.items() if value)
try:
with os.fdopen(os.open(env_path, os.O_WRONLY | os.O_APPEND | os.O_CREAT | os.O_NOFOLLOW, 0o600), "w") as handle:
os.fchmod(handle.fileno(), 0o600)
_ = handle.write("\n".join(lines) + "\n")
except OSError:
_ = sys.stderr.write("could not write the environment file\n")
return 1
return 0
if __name__ == "__main__":
sys.exit(main())

View file

@ -1,52 +0,0 @@
import re
import sys
from typing import Final
SELECTABLE: Final = re.compile(r"^tests/e2e/([A-Za-z0-9_.-]+/)*test_[A-Za-z0-9_.-]+\.py$")
UNSUPPORTED: Final = re.compile(
r"^tests/e2e/(ui|claude_code|load|migrations)/"
r"|^tests/e2e/mcp/test_mcp_oauth_happy_path_e2e\.py$"
r"|^tests/e2e/llm_translation/realtime/test_realtime_pipecat_audio_e2e\.py$"
r"|^tests/e2e/batches/test_managed_files_enforcement_e2e\.py$"
r"|^tests/e2e/guardrails/test_presidio_masking_e2e\.py$"
r"|^tests/e2e/logging/test_otel_v2_langfuse_generation_output_e2e\.py$"
r"|^tests/e2e/logging/test_langsmith_batch_serialization_e2e\.py$"
r"|^tests/e2e/logging/test_s3_log_e2e\.py$"
r"|^tests/e2e/secret_manager/"
)
HARNESS: Final = re.compile(
r"^tests/e2e/[A-Za-z0-9_.-]+\.(py|ini)$"
r"|^tests/e2e/idp_realm\.json$"
r"|^tests/e2e/management/(management_client|jwt_actors|conftest)\.py$"
r"|^tests/e2e/coverage_registry/management_cases\.py$"
r"|^tests/e2e/gateway/"
r"|^\.github/e2e-stack/"
r"|^\.github/workflows/test-e2e-changed\.yml$"
)
UNEXPANDED: Final = re.compile(r"[*?\[]")
def is_selectable(path: str) -> bool:
return SELECTABLE.match(path) is not None and UNSUPPORTED.match(path) is None
def select(changed: tuple[str, ...], canary: tuple[str, ...]) -> tuple[str, ...]:
direct: Final = frozenset(path for path in changed if is_selectable(path))
harness_changed: Final = any(HARNESS.match(path) for path in changed)
canary_tests: Final = frozenset(path for path in canary if harness_changed and is_selectable(path))
return tuple(sorted(direct | canary_tests))
def main() -> int:
canary: Final = tuple(sys.argv[1:])
unexpanded: Final = tuple(path for path in canary if UNEXPANDED.search(path))
if unexpanded:
_ = sys.stderr.write(f"the canary paths reached the selector unexpanded: {' '.join(unexpanded)}\n")
return 1
changed: Final = tuple(line.strip() for line in sys.stdin if line.strip())
_ = sys.stdout.write(" ".join(select(changed, canary)) + "\n")
return 0
if __name__ == "__main__":
sys.exit(main())

View file

@ -1,231 +0,0 @@
#!/usr/bin/env bash
set -euo pipefail
umask 077
REPO_ROOT="$(cd "$(dirname "${BASH_SOURCE[0]}")/../.." && pwd)"
STACK_DIR="${E2E_STACK_DIR:-${RUNNER_TEMP:-/tmp}/litellm-e2e-stack}"
CERTS_DIR="${STACK_DIR}/certs"
LOGS_DIR="${STACK_DIR}/logs"
PIDS_DIR="${STACK_DIR}/pids"
POSTGRES_IMAGE="${E2E_POSTGRES_IMAGE:-postgres:16.6}"
VALKEY_IMAGE="${E2E_VALKEY_IMAGE:-valkey/valkey:8.1.4@sha256:81db6d39e1bba3b3ff32bd3a1b19a6d69690f94a3954ec131277b9a26b95b3aa}"
JAEGER_IMAGE="${E2E_JAEGER_IMAGE:-jaegertracing/jaeger:2.10.0}"
NGINX_IMAGE="${E2E_NGINX_IMAGE:-nginx:1.29.1-alpine@sha256:42a516af16b852e33b7682d5ef8acbd5d13fe08fecadc7ed98605ba5e3b26ab8}"
LB_PORT="${E2E_LB_PORT:-4000}"
GATEWAY_PORT_1="${E2E_GATEWAY_PORT_1:-4010}"
GATEWAY_PORT_2="${E2E_GATEWAY_PORT_2:-4011}"
BACKEND_PORT="${E2E_BACKEND_PORT:-4001}"
REDIS_PORT="${E2E_REDIS_PORT:-6379}"
DATABASE_HOST="${E2E_DATABASE_HOST:-127.0.0.1}"
DATABASE_PORT="${E2E_DATABASE_PORT:-5432}"
DATABASE_USER="${E2E_DATABASE_USER:-litellm}"
DATABASE_PASSWORD="${E2E_DATABASE_PASSWORD:-dbpassword9090}"
DATABASE_NAME="${E2E_DATABASE_NAME:-litellm}"
JAEGER_OTLP_PORT="${E2E_JAEGER_OTLP_PORT:-4318}"
JAEGER_OTLP_TLS_PORT="${E2E_JAEGER_OTLP_TLS_PORT:-4319}"
JAEGER_QUERY_PORT="${E2E_JAEGER_QUERY_PORT:-16686}"
KEYCLOAK_PORT="${E2E_KEYCLOAK_PORT:-8081}"
MASTER_KEY="${LITELLM_MASTER_KEY:-sk-e2e-$(openssl rand -hex 16)}"
mkdir -p "${CERTS_DIR}" "${LOGS_DIR}" "${PIDS_DIR}"
chmod 700 "${STACK_DIR}" "${LOGS_DIR}" "${PIDS_DIR}"
chmod 755 "${CERTS_DIR}"
log() { printf 'e2e-stack: %s\n' "$*"; }
port_open() { (exec 3<>"/dev/tcp/127.0.0.1/$1") 2>/dev/null; }
wait_for() {
local label="$1" check="$2" deadline=$((SECONDS + ${3:-120}))
until eval "${check}"; do
if ((SECONDS >= deadline)); then
log "timed out waiting for ${label}"
exit 1
fi
sleep 2
done
log "${label} is up"
}
if [[ -f "${REPO_ROOT}/tests/e2e/.env" ]]; then
set -a
source "${REPO_ROOT}/tests/e2e/.env"
set +a
fi
if [[ -z "${DD_API_KEY:-}" ]]; then
log "DD_API_KEY is empty; the gateway config enables the datadog callback, so put a Datadog API key in tests/e2e/.env"
exit 1
fi
export DD_SITE="${DD_SITE:-datadoghq.com}"
if ! port_open "${DATABASE_PORT}"; then
docker run -d --name e2e-postgres -p "${DATABASE_PORT}:5432" \
-e "POSTGRES_USER=${DATABASE_USER}" -e "POSTGRES_PASSWORD=${DATABASE_PASSWORD}" -e "POSTGRES_DB=${DATABASE_NAME}" \
"${POSTGRES_IMAGE}" >/dev/null
fi
wait_for "postgres" "port_open ${DATABASE_PORT}"
if ! port_open "${JAEGER_QUERY_PORT}"; then
docker run -d --name e2e-jaeger -p "${JAEGER_OTLP_PORT}:4318" -p "${JAEGER_QUERY_PORT}:16686" \
"${JAEGER_IMAGE}" >/dev/null
fi
wait_for "jaeger" "curl -fs http://127.0.0.1:${JAEGER_QUERY_PORT}/api/services >/dev/null"
openssl genrsa -out "${CERTS_DIR}/ca.key" 2048 2>/dev/null
openssl req -x509 -new -nodes -key "${CERTS_DIR}/ca.key" -sha256 -days 7 \
-subj "/CN=litellm-e2e-ca" \
-addext "basicConstraints=critical,CA:TRUE" -addext "keyUsage=critical,keyCertSign,cRLSign" \
-out "${CERTS_DIR}/ca.crt" 2>/dev/null
openssl genrsa -out "${CERTS_DIR}/server.key" 2048 2>/dev/null
openssl req -new -key "${CERTS_DIR}/server.key" -subj "/CN=localhost" -out "${CERTS_DIR}/server.csr" 2>/dev/null
openssl x509 -req -in "${CERTS_DIR}/server.csr" -CA "${CERTS_DIR}/ca.crt" -CAkey "${CERTS_DIR}/ca.key" \
-CAcreateserial -days 7 -sha256 \
-extfile <(printf 'basicConstraints=CA:FALSE\nkeyUsage=critical,digitalSignature,keyEncipherment\nextendedKeyUsage=serverAuth\nsubjectAltName=DNS:localhost,IP:127.0.0.1\n') \
-out "${CERTS_DIR}/server.crt" 2>/dev/null
chmod 644 "${CERTS_DIR}"/*.key "${CERTS_DIR}"/*.crt
CERTIFI_BUNDLE="$(cd "${REPO_ROOT}" && uv run --no-sync python -c 'import certifi; print(certifi.where())')"
cat "${CERTIFI_BUNDLE}" "${CERTS_DIR}/ca.crt" > "${CERTS_DIR}/ca-bundle.pem"
docker rm -f e2e-valkey >/dev/null 2>&1 || true
docker run -d --name e2e-valkey -p "${REDIS_PORT}:${REDIS_PORT}" -v "${CERTS_DIR}:/certs:ro" \
"${VALKEY_IMAGE}" valkey-server \
--cluster-enabled yes --port 0 --tls-port "${REDIS_PORT}" \
--tls-cert-file /certs/server.crt --tls-key-file /certs/server.key --tls-ca-cert-file /certs/ca.crt \
--tls-auth-clients no --cluster-announce-ip 127.0.0.1 >/dev/null
VALKEY_CLI="docker exec e2e-valkey valkey-cli --tls --cacert /certs/ca.crt -h 127.0.0.1 -p ${REDIS_PORT}"
wait_for "valkey" "${VALKEY_CLI} ping 2>/dev/null | grep -q PONG"
${VALKEY_CLI} cluster addslotsrange 0 16383 >/dev/null
wait_for "valkey cluster" "${VALKEY_CLI} cluster info 2>/dev/null | grep -q cluster_state:ok"
CONFIG_SOURCE="${REPO_ROOT}/tests/e2e/gateway/stage_mirror_ci_config.yml"
CONFIG_PATH="${CONFIG_SOURCE}"
if [[ "${REDIS_PORT}" != "6379" ]]; then
CONFIG_PATH="${STACK_DIR}/litellm-config.yml"
sed "s/port: 6379/port: ${REDIS_PORT}/" "${CONFIG_SOURCE}" > "${CONFIG_PATH}"
fi
SERVER_ENV=(
"LITELLM_MASTER_KEY=${MASTER_KEY}"
"DATABASE_HOST=${DATABASE_HOST}"
"DATABASE_PORT=${DATABASE_PORT}"
"DATABASE_USER=${DATABASE_USER}"
"DATABASE_PASSWORD=${DATABASE_PASSWORD}"
"DATABASE_NAME=${DATABASE_NAME}"
"DISABLE_SCHEMA_UPDATE=true"
"REDIS_HOST=127.0.0.1"
"REDIS_PORT=${REDIS_PORT}"
"REDIS_CLUSTER_NODES=[{\"host\":\"127.0.0.1\",\"port\":${REDIS_PORT}}]"
"CONFIG_FILE_PATH=${CONFIG_PATH}"
"STORE_MODEL_IN_DB=True"
"OTEL_EXPORTER_OTLP_PROTOCOL=http/protobuf"
"OTEL_EXPORTER_OTLP_ENDPOINT=https://127.0.0.1:${JAEGER_OTLP_TLS_PORT}"
"SSL_CERT_FILE=${CERTS_DIR}/ca-bundle.pem"
"PYTHONPATH=${REPO_ROOT}"
"JWT_PUBLIC_KEY_URL=http://127.0.0.1:${KEYCLOAK_PORT}/realms/litellm-e2e/protocol/openid-connect/certs"
"JWT_ISSUER=http://127.0.0.1:${KEYCLOAK_PORT}/realms/litellm-e2e"
"JWT_AUDIENCE=litellm-e2e"
)
if [[ -n "${VERTEXAI_CREDENTIALS:-}" ]]; then
printf '%s' "${VERTEXAI_CREDENTIALS}" > "${STACK_DIR}/vertex-adc.json"
SERVER_ENV+=("GOOGLE_APPLICATION_CREDENTIALS=${STACK_DIR}/vertex-adc.json")
fi
cd "${REPO_ROOT}"
env "${SERVER_ENV[@]}" "E2E_KEYCLOAK_PORT=${KEYCLOAK_PORT}" bash .github/e2e-stack/start-idp.sh
log "running migrations"
env "${SERVER_ENV[@]}" uv run --no-sync python migrations/run.py >"${LOGS_DIR}/migrations.log" 2>&1
start_server() {
local name="$1"; shift
env -u AWS_ROLE_NAME "${SERVER_ENV[@]}" "$@" >"${LOGS_DIR}/${name}.log" 2>&1 &
echo $! > "${PIDS_DIR}/${name}.pid"
}
if [[ "$(uname)" == "Linux" ]]; then
NGINX_UPSTREAM_HOST=127.0.0.1
NGINX_DOCKER_ARGS=(--network host)
else
NGINX_UPSTREAM_HOST=host.docker.internal
NGINX_DOCKER_ARGS=(-p "${LB_PORT}:${LB_PORT}" -p "${JAEGER_OTLP_TLS_PORT}:${JAEGER_OTLP_TLS_PORT}")
fi
cat > "${STACK_DIR}/nginx.conf" <<EOF
events {}
http {
map \$http_upgrade \$connection_upgrade {
default upgrade;
'' close;
}
upstream litellm_gateways {
server ${NGINX_UPSTREAM_HOST}:${GATEWAY_PORT_1};
server ${NGINX_UPSTREAM_HOST}:${GATEWAY_PORT_2};
}
server {
listen ${LB_PORT};
client_max_body_size 100m;
location / {
proxy_pass http://litellm_gateways;
proxy_http_version 1.1;
proxy_set_header Host \$host;
proxy_set_header X-Forwarded-For \$proxy_add_x_forwarded_for;
proxy_set_header X-Forwarded-Proto \$scheme;
proxy_set_header Upgrade \$http_upgrade;
proxy_set_header Connection \$connection_upgrade;
proxy_buffering off;
proxy_read_timeout 600s;
proxy_send_timeout 600s;
}
}
server {
listen ${JAEGER_OTLP_TLS_PORT} ssl;
ssl_certificate /certs/server.crt;
ssl_certificate_key /certs/server.key;
client_max_body_size 100m;
location / {
proxy_pass http://${NGINX_UPSTREAM_HOST}:${JAEGER_OTLP_PORT};
}
}
}
EOF
docker rm -f e2e-nginx >/dev/null 2>&1 || true
docker run -d --name e2e-nginx "${NGINX_DOCKER_ARGS[@]}" \
-v "${STACK_DIR}/nginx.conf:/etc/nginx/nginx.conf:ro" \
-v "${CERTS_DIR}:/certs:ro" "${NGINX_IMAGE}" >/dev/null
wait_for "Jaeger OTLP TLS listener" \
"curl -sS --cacert ${CERTS_DIR}/ca.crt https://127.0.0.1:${JAEGER_OTLP_TLS_PORT}/ -o /dev/null -w '%{http_code}' | grep -qE '^[2345]'"
start_server backend uv run --no-sync uvicorn backend.main:app --host 0.0.0.0 --port "${BACKEND_PORT}"
start_server gateway-1 uv run --no-sync uvicorn gateway.main:app --workers 1 --host 0.0.0.0 --port "${GATEWAY_PORT_1}"
start_server gateway-2 uv run --no-sync uvicorn gateway.main:app --workers 1 --host 0.0.0.0 --port "${GATEWAY_PORT_2}"
wait_for "backend" "curl -fs http://127.0.0.1:${BACKEND_PORT}/health/liveliness >/dev/null" 300
wait_for "gateway-1" "curl -fs http://127.0.0.1:${GATEWAY_PORT_1}/health/liveliness >/dev/null" 300
wait_for "gateway-2" "curl -fs http://127.0.0.1:${GATEWAY_PORT_2}/health/liveliness >/dev/null" 300
wait_for "load balancer" "curl -fs http://127.0.0.1:${LB_PORT}/health/liveliness >/dev/null" 60
cat > "${STACK_DIR}/stack.env" <<EOF
LITELLM_PROXY_URL=http://127.0.0.1:${LB_PORT}
LITELLM_CONTROL_PLANE_URL=http://127.0.0.1:${BACKEND_PORT}
LITELLM_PROXY_REPLICA_URLS=http://127.0.0.1:${GATEWAY_PORT_1},http://127.0.0.1:${GATEWAY_PORT_2}
LITELLM_MASTER_KEY=${MASTER_KEY}
REDIS_HOST=127.0.0.1
REDIS_PORT=${REDIS_PORT}
E2E_OTEL_QUERY_URL=http://127.0.0.1:${JAEGER_QUERY_PORT}
E2E_OTEL_EXPORTER_ENDPOINT=https://127.0.0.1:${JAEGER_OTLP_TLS_PORT}
E2E_KEYCLOAK_URL=http://127.0.0.1:${KEYCLOAK_PORT}
E2E_KEYCLOAK_ADMIN_USER=admin
E2E_KEYCLOAK_ADMIN_PASSWORD=e2e-ephemeral-idp-not-a-secret
SSL_CERT_FILE=${CERTS_DIR}/ca-bundle.pem
DATABASE_URL=postgresql://${DATABASE_USER}:${DATABASE_PASSWORD}@${DATABASE_HOST}:${DATABASE_PORT}/${DATABASE_NAME}
EOF
log "stack is up; pytest env written to ${STACK_DIR}/stack.env"

View file

@ -9,7 +9,6 @@ import sys
import warnings
from collections.abc import Callable, Iterable, Mapping, Sequence
from dataclasses import dataclass
from types import MappingProxyType
from typing import Final
import yaml
@ -22,11 +21,12 @@ TESTS_ROOT = REPO_ROOT / "tests"
ALLOWLIST_KEYS = frozenset({"description", "test_paths", "dockerfiles"})
PATH_FILTER_KEYS = frozenset({"paths", "paths-ignore"})
TEST_PATH_KEYS = frozenset({"test-path", "test-paths"})
TEST_PATH_KEYS = frozenset({"test-path", "test-paths", "test_path"})
DOCKERFILE_INPUT_KEYS = frozenset({"file", "dockerfile"})
TEST_RUNNER_RE = re.compile(r"\bpytest\b|\bcircleci tests\b|\bhelm unittest\b|\bplaywright test\b|\bpython[0-9.]*\s")
IMAGE_BUILD_RE = re.compile(r"\bdocker\s+(?:buildx\s+)?build\b")
TEST_TOKEN_RE = re.compile(r"tests/[A-Za-z0-9_./*?-]+")
IGNORE_ARG_RE: Final = re.compile(r"--ignore(?:-glob)?[= ](\S+)")
DOCKERFILE_TOKEN_RE = re.compile(r"[A-Za-z0-9_./-]*Dockerfile[A-Za-z0-9_.-]*")
COMMENT_RE = re.compile(r"^\s*#.*$", re.MULTILINE)
GLOB_CHARS = frozenset("*?")
@ -72,6 +72,17 @@ class Scalar:
value: str
@dataclass(frozen=True, slots=True)
class Selection:
included: frozenset[str]
ignored: frozenset[str]
def covers(self, relative_path: str) -> bool:
return any(_token_covers(token, relative_path) for token in self.included) and not any(
_token_covers(token, relative_path) for token in self.ignored
)
@dataclass(frozen=True, slots=True)
class Finding:
subject: str
@ -111,44 +122,32 @@ def _uncommented(value: str) -> str:
return COMMENT_RE.sub("", value)
def _invoked_test_tokens(scalars: Iterable[Scalar]) -> frozenset[str]:
return frozenset(
match.group(0).rstrip("/")
def _selection_for_scalar(scalar: Scalar) -> Selection:
text: Final = _uncommented(scalar.value)
ignored: Final = frozenset().union(
*(
frozenset(token.rstrip("/") for token in TEST_TOKEN_RE.findall(ignored_argument))
for ignored_argument in IGNORE_ARG_RE.findall(text)
)
)
included: Final = frozenset(
match.group(0).rstrip("/") for match in TEST_TOKEN_RE.finditer(IGNORE_ARG_RE.sub("", text))
)
return Selection(included=included, ignored=ignored)
def _invoked_selections(scalars: Iterable[Scalar]) -> tuple[Selection, ...]:
selected_scalars: Final = tuple(
scalar
for scalar in scalars
if scalar.key in TEST_PATH_KEYS or TEST_RUNNER_RE.search(scalar.value)
for match in TEST_TOKEN_RE.finditer(_uncommented(scalar.value))
)
selections: Final = tuple(_selection_for_scalar(scalar) for scalar in selected_scalars)
return tuple(selection for selection in selections if selection.included)
SELECTION_ARM_RE = re.compile(r"(?ms)^\s*([A-Za-z0-9_|*-]+)\)\s*(.*?);;")
def _unit_selection_arms(repo_root: pathlib.Path = REPO_ROOT) -> Mapping[str, frozenset[str]]:
script: Final = repo_root / ".circleci/scripts/unit_selection.sh"
if not script.is_file():
return MappingProxyType({})
text: Final = _uncommented(script.read_text())
return MappingProxyType(
{
label: frozenset(match.group(0).rstrip("/") for match in TEST_TOKEN_RE.finditer(body))
for label, body in SELECTION_ARM_RE.findall(text)
}
)
def _unit_selection_tokens(repo_root: pathlib.Path = REPO_ROOT) -> frozenset[str]:
return frozenset(token for tokens in _unit_selection_arms(repo_root).values() for token in tokens)
def _wired_unit_flags(scalars: Iterable[Scalar]) -> frozenset[str]:
return frozenset(scalar.value for scalar in scalars if scalar.key == "unit-flag" and "${{" not in scalar.value)
def _shard_tokens(scalars: Iterable[Scalar], arms: Mapping[str, frozenset[str]]) -> frozenset[str]:
wired: Final = _wired_unit_flags(scalars)
return _invoked_test_tokens(scalars) | frozenset(
token for label, tokens in arms.items() if label in wired for token in tokens
)
def _invoked_test_tokens(scalars: Iterable[Scalar]) -> frozenset[str]:
return frozenset().union(*(selection.included for selection in _invoked_selections(scalars)))
def _built_dockerfile_tokens(scalars: Iterable[Scalar]) -> frozenset[str]:
@ -211,11 +210,12 @@ def _dockerfiles() -> tuple[str, ...]:
)
def _uncovered_tests(allowlist: Allowlist, tokens: frozenset[str]) -> tuple[Finding, ...]:
def _uncovered_tests(allowlist: Allowlist, selections: tuple[Selection, ...]) -> tuple[Finding, ...]:
uncovered = tuple(
relative_path
for relative_path in _test_files()
if not any(_token_covers(token, relative_path) for token in tokens) and not allowlist.covers_test(relative_path)
if not any(selection.covers(relative_path) for selection in selections)
and not allowlist.covers_test(relative_path)
)
directories = tuple(dict.fromkeys(path.rsplit("/", 1)[0] for path in uncovered))
return tuple(
@ -505,7 +505,7 @@ def _check_slices() -> int:
def _check_shards() -> int:
findings = _unassigned_shard_children(_shard_tokens(_all_scalars(), _unit_selection_arms()))
findings = _unassigned_shard_children(_invoked_test_tokens(_all_scalars()))
if findings:
_report(
"test directories and files that no shard claims",
@ -568,6 +568,9 @@ def _integration_ownership(repo_root: pathlib.Path = REPO_ROOT) -> tuple[frozens
browser_paths: Final = frozenset(node.split("::", 1)[0] for node in browser_nodes)
circle_path: Final = repo_root / ".circleci/config.yml"
circle: Final = yaml.safe_load(circle_path.read_text()) if circle_path.exists() else {}
circle_test_path_tokens: Final = _invoked_test_tokens(
scalar for scalar in _scalars(circle, "config.yml") if scalar.key == "test_path"
)
steps: Final = circle.get("jobs", {}).get("integration_contracts", {}).get("steps", ())
invoked: Final = any(
".circleci/scripts/run_integration.sh" in scalar.value
@ -609,9 +612,9 @@ def _integration_ownership(repo_root: pathlib.Path = REPO_ROOT) -> tuple[frozens
if any(_token_covers(token, path) for token in gha_tokens)
)
+ tuple(
Finding(path, "GitHub-owned integration contract has no invoking workflow")
Finding(path, "GitHub-owned integration contract has no invoking job")
for path in sorted(github_files)
if not any(_token_covers(token, path) for token in gha_tokens)
if not any(_token_covers(token, path) for token in gha_tokens | circle_test_path_tokens)
)
+ tuple(
Finding(path, "GitHub-owned integration file is missing")
@ -675,7 +678,11 @@ def main() -> int:
integration_paths, ownership_findings = _integration_ownership()
test_findings = (
_uncovered_tests(allowlist, _invoked_test_tokens(scalars) | _unit_selection_tokens() | integration_paths)
_uncovered_tests(
allowlist,
_invoked_selections(scalars)
+ (Selection(included=integration_paths, ignored=frozenset()),),
)
+ ownership_findings
)
dockerfile_findings = _uncovered_dockerfiles(allowlist, _built_dockerfile_tokens(scalars))

View file

@ -214,7 +214,7 @@ def main(
native_module: Final = load_native_module(native_path)
native_module_loads: Final = native_module is not None
panic_test_hook_absent: Final = native_module is not None and not hasattr(native_module, "_panic_for_test")
native_size_limit: Final = 45_000_000
native_size_limit: Final = 48_000_000
native_size_within_limit: Final = native_member.file_size <= native_size_limit
validations: Final = (
(f"Python tag is {EXPECTED_PYTHON_TAG}", python_tag == EXPECTED_PYTHON_TAG),

View file

@ -13,16 +13,6 @@ on:
have its path existence-checked like any other token.
required: true
type: string
unit-flag:
description: >-
Codecov flag of the `.circleci/tests.yml` job that now owns part of
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: ""
workers:
description: "Number of pytest-xdist workers"
required: false
@ -79,12 +69,6 @@ on:
description: "Unique name for the coverage artifact (must be unique per run)"
required: true
type: string
legacy-mcp-peer:
description: "Install the isolated SDK1 peer for MCP compatibility tests"
required: false
type: boolean
default: false
permissions:
contents: read
@ -142,17 +126,10 @@ jobs:
- name: Install dependencies
if: steps.changes.outputs.decision != 'skip'
timeout-minutes: 8
env:
LEGACY_MCP_PEER: ${{ inputs.legacy-mcp-peer }}
run: |
diff -u model_prices_and_context_window.json litellm/model_prices_and_context_window_backup.json
.github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router --extra saml
.github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router --extra saml --extra caching --extra extra_proxy --extra proxy-runtime
uv run --no-sync python -c 'import os, sys; print(sys.version); assert f"{sys.version_info.major}.{sys.version_info.minor}" == os.environ["UV_PYTHON"]'
if [ "$LEGACY_MCP_PEER" = "true" ]; then
uv venv --python "${UV_PYTHON}" .venv-mcp-peer
uv pip install --python .venv-mcp-peer 'mcp==1.28.1' 'langchain-mcp-adapters==0.2.1'
echo "MCP_TEST_PEER_PYTHON=$GITHUB_WORKSPACE/.venv-mcp-peer/bin/python" >> "$GITHUB_ENV"
fi
- name: Cache Prisma binaries
if: steps.changes.outputs.decision != 'skip'
@ -171,7 +148,6 @@ jobs:
timeout-minutes: ${{ inputs.timeout-minutes }}
env:
TEST_PATH: ${{ inputs.test-path }}
UNIT_FLAG: ${{ inputs.unit-flag }}
MAX_FAILURES: ${{ inputs.max-failures }}
WORKERS: ${{ inputs.workers }}
RERUNS: ${{ inputs.reruns }}
@ -181,9 +157,6 @@ jobs:
run: |
echo "has-coverage=false" >> "$GITHUB_OUTPUT"
selection="${TEST_PATH}"
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; nothing to run"
exit 0

View file

@ -1,265 +0,0 @@
name: e2e-changed-tests
on:
pull_request:
concurrency:
group: e2e-changed-${{ github.event.pull_request.number }}
cancel-in-progress: true
permissions: {}
jobs:
detect:
name: Detect changed e2e tests
runs-on: ubuntu-latest
timeout-minutes: 5
permissions:
contents: read
pull-requests: read
outputs:
tests: ${{ steps.changed.outputs.tests }}
any: ${{ steps.changed.outputs.any }}
steps:
- name: Checkout the selector and the canary suite
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
sparse-checkout: |
.github/e2e-stack
tests/e2e/access_control
tests/e2e/management/test_jwt_management_e2e.py
tests/e2e/other/test_jwt_auth_e2e.py
persist-credentials: false
ref: ${{ github.sha }}
- name: List the e2e test files this PR added or modified
id: changed
env:
GH_TOKEN: ${{ github.token }}
REPO: ${{ github.repository }}
PR_NUMBER: ${{ github.event.pull_request.number }}
HEAD_SHA: ${{ github.event.pull_request.head.sha }}
run: |
gh api "repos/${REPO}/pulls/${PR_NUMBER}" \
--jq 'select(.head.sha == env.HEAD_SHA and .changed_files < 3000) | .head.sha' \
| grep -Fxq "${HEAD_SHA}"
files="$(gh api "repos/${REPO}/pulls/${PR_NUMBER}/files" --paginate \
--jq '.[] | select(.status != "removed") | .filename')"
gh api "repos/${REPO}/pulls/${PR_NUMBER}" --jq '.head.sha' | grep -Fxq "${HEAD_SHA}"
tests="$(printf '%s\n' "${files}" \
| python3 .github/e2e-stack/select_tests.py tests/e2e/access_control/test_*.py \
tests/e2e/management/test_jwt_management_e2e.py tests/e2e/other/test_jwt_auth_e2e.py)"
echo "tests=${tests}" >> "${GITHUB_OUTPUT}"
if [ -n "${tests}" ]; then
echo "any=true" >> "${GITHUB_OUTPUT}"
echo "selected e2e tests: ${tests}"
else
echo "any=false" >> "${GITHUB_OUTPUT}"
echo "no changed e2e test files supported by this stack; nothing to run"
fi
run:
name: Run changed e2e tests against the stage-mirror stack
needs: detect
if: needs.detect.outputs.any == 'true' && github.event.pull_request.head.repo.full_name == github.repository
runs-on: ubuntu-latest
timeout-minutes: 90
environment: e2e-changed
permissions:
contents: read
id-token: write
services:
postgres:
image: postgres:16.6
env:
POSTGRES_USER: litellm
POSTGRES_PASSWORD: dbpassword9090
POSTGRES_DB: litellm
ports:
- 5432:5432
options: >-
--health-cmd "pg_isready -U litellm"
--health-interval 5s
--health-timeout 5s
--health-retries 10
jaeger:
image: jaegertracing/jaeger:2.10.0
ports:
- 4318:4318
- 16686:16686
steps:
- name: Validate configuration
env:
ROLE: ${{ vars.E2E_AWS_ROLE_TO_ASSUME }}
run: test -n "${ROLE}" || { echo "::error::Set repo variable E2E_AWS_ROLE_TO_ASSUME to an OIDC role with read access to the e2e secrets"; exit 1; }
- name: Checkout
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
persist-credentials: false
ref: ${{ github.sha }}
- name: Set up Python
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
with:
python-version: "3.13"
- name: Set up uv
uses: ./.github/actions/setup-uv-with-retries
with:
version: "0.10.9"
- name: Cache the Rust build
uses: ./.github/actions/cache-cargo-build
- name: Install dependencies
run: |
.github/scripts/uv_sync_with_retries.sh --frozen \
--extra proxy --extra proxy-runtime --extra extra_proxy \
--extra semantic-router --extra bedrock-realtime \
--group ci --group proxy-dev --group e2e-dev
uv pip install "pipecat-ai[openai]==1.4.0"
- name: Cache Prisma binaries
uses: ./.github/actions/cache-prisma-binaries
- name: Generate Prisma client
run: uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma
- name: Install Playwright chromium
run: uv run --no-sync playwright install --with-deps chromium
- name: Configure AWS credentials
id: aws
uses: aws-actions/configure-aws-credentials@e7f100cf4c008499ea8adda475de1042d6975c7b # v6.2.0
with:
role-to-assume: ${{ vars.E2E_AWS_ROLE_TO_ASSUME }}
aws-region: us-east-1
role-session-name: litellm-e2e-changed-${{ github.run_id }}
role-duration-seconds: 900
output-env-credentials: false
output-credentials: true
- name: Fetch provider credentials from AWS Secrets Manager
env:
AWS_ACCESS_KEY_ID: ${{ steps.aws.outputs.aws-access-key-id }}
AWS_SECRET_ACCESS_KEY: ${{ steps.aws.outputs.aws-secret-access-key }}
AWS_SESSION_TOKEN: ${{ steps.aws.outputs.aws-session-token }}
AWS_DEFAULT_REGION: us-east-1
run: |
umask 077
aws secretsmanager get-secret-value --secret-id litellm-e2e-changed-provider-keys \
--query SecretString --output text \
| uv run --no-sync python .github/e2e-stack/secrets_to_env.py tests/e2e/.env
aws secretsmanager get-secret-value --secret-id litellm-e2e-changed-license \
--query SecretString --output text \
| jq -R -s '{"LITELLM_LICENSE": .}' \
| uv run --no-sync python .github/e2e-stack/secrets_to_env.py tests/e2e/.env
- name: Boot the stage-mirror stack
id: boot
run: |
umask 077
if ! bash .github/e2e-stack/up.sh > "${RUNNER_TEMP}/e2e-boot.log" 2>&1; then
echo "::error::stage-mirror stack failed to boot; raw logs are not published"
exit 1
fi
- name: Export stack environment
run: |
master_key="$(grep '^LITELLM_MASTER_KEY=' "${RUNNER_TEMP}/litellm-e2e-stack/stack.env" | cut -d= -f2-)"
echo "::add-mask::${master_key}"
cat "${RUNNER_TEMP}/litellm-e2e-stack/stack.env" >> "${GITHUB_ENV}"
- name: Run the selected tests three times
env:
TESTS: ${{ needs.detect.outputs.tests }}
E2E_FIXTURE_MODE: live
E2E_PROVIDER_EDGE_HOST_REACHABLE: '1'
E2E_OWNED_GATEWAY: '1'
COLUMNS: '400'
run: |
umask 077
read -r -a test_files <<< "${TESTS}"
for pass in 1 2 3; do
report="${RUNNER_TEMP}/e2e-pass-${pass}.xml"
log="${RUNNER_TEMP}/e2e-pass-${pass}.log"
echo "::group::pass ${pass} of 3"
set +e
uv run --no-sync pytest "${test_files[@]}" --rootdir=. -v --reruns 0 -p no:cacheprovider \
-o junit_family=xunit1 --junitxml="${report}" > "${log}" 2>&1
status=$?
uv run --no-sync python .github/e2e-stack/assert_tests_ran.py "${report}" "${test_files[@]}"
verified=$?
set -e
grep -E '^(FAILED|ERROR) ' "${log}" || true
grep -E '^=+ .* in [0-9.]+s( \([0-9:]+\))? =+$' "${log}" | tail -n 1
echo "::endgroup::"
if [ "${status}" = "5" ]; then
echo "::error::the selected files collected no runnable tests, so nothing was verified"
exit 1
fi
if [ "${status}" != "0" ]; then
echo "::error::pass ${pass} of 3 failed with exit code ${status}"
exit "${status}"
fi
if [ "${verified}" != "0" ]; then
echo "::error::pass ${pass} of 3 did not verify every selected file"
exit 1
fi
echo "pass ${pass} of 3 passed"
done
- name: Redact the pytest output
if: always() && steps.boot.outcome == 'success'
run: |
umask 077
shopt -s nullglob
uv run --no-sync python .github/e2e-stack/redact_output.py \
--values tests/e2e/.env --values "${RUNNER_TEMP}/litellm-e2e-stack/stack.env" \
--out "${RUNNER_TEMP}/e2e-redacted" "${RUNNER_TEMP}"/e2e-pass-*.log "${RUNNER_TEMP}"/e2e-pass-*.xml
- name: Keep the redacted pytest output
if: always() && steps.boot.outcome == 'success'
uses: actions/upload-artifact@4cec3d8aa04e39d1a68397de0c4cd6fb9dce8ec1 # v4.6.1
with:
name: e2e-changed-pytest-output-${{ github.run_attempt }}
path: ${{ runner.temp }}/e2e-redacted
retention-days: 14
if-no-files-found: ignore
- name: Stop the stack
if: always() && steps.boot.outcome != 'skipped'
run: bash .github/e2e-stack/down.sh
- name: Remove credentials and raw output
if: always()
run: |
rm -f tests/e2e/.env "${RUNNER_TEMP}/e2e-boot.log" "${RUNNER_TEMP}"/e2e-pass-*.log "${RUNNER_TEMP}"/e2e-pass-*.xml
rm -rf "${RUNNER_TEMP}/litellm-e2e-stack" "${RUNNER_TEMP}/e2e-redacted"
gate:
name: e2e-changed-tests
needs: [detect, run]
if: always()
runs-on: ubuntu-latest
timeout-minutes: 5
steps:
- name: Require three successful passes when tests changed
env:
DETECT_RESULT: ${{ needs.detect.result }}
ANY_TESTS: ${{ needs.detect.outputs.any }}
RUN_RESULT: ${{ needs.run.result }}
run: |
if [ "${DETECT_RESULT}" != "success" ]; then
echo "::error::changed-test detection did not succeed"
exit 1
fi
if [ "${ANY_TESTS}" = "false" ]; then
echo "no changed e2e test files supported by this stack; nothing to run"
exit 0
fi
if [ "${ANY_TESTS}" != "true" ] || [ "${RUN_RESULT}" != "success" ]; then
echo "::error::selected e2e tests require an approved, successful run; fork PRs must run from a reviewed same-repository branch"
exit 1
fi

View file

@ -5,6 +5,27 @@ on:
branches:
- main
- "litellm_**"
paths:
- "**/pyproject.toml"
- "uv.lock"
- "uv.toml"
- ".python-version"
- "rust-toolchain.toml"
- "litellm-rust/**"
- "litellm/__init__.py"
- "litellm/proxy/proxy_server.py"
- "litellm/**/*mcp*"
- "litellm/**/*mcp*/**"
- "litellm/integrations/arize/**"
- "scripts/check_mcp_sdk_install.py"
- "tests/base_sdk_tests/**"
- ".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"
permissions:
contents: read

View file

@ -1,185 +0,0 @@
name: "Postgres Tests"
on:
pull_request:
branches:
- main
- "litellm_**"
push:
branches:
- main
workflow_dispatch:
permissions:
contents: read
concurrency:
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.sha }}
cancel-in-progress: ${{ github.event_name == 'pull_request' }}
jobs:
postgres:
name: ${{ matrix.shard }}
runs-on: ubuntu-latest
timeout-minutes: ${{ matrix.job-timeout-minutes }}
permissions:
contents: read
id-token: write
services:
postgres:
image: postgres:16@sha256:e17e86066e5ef83e0952a9347f5c792b7ece00972e2aa787a6986f471b3dd3d5
env:
POSTGRES_USER: postgres
POSTGRES_PASSWORD: postgres
POSTGRES_DB: litellm_test
ports:
- 5432:5432
options: >-
--health-cmd pg_isready
--health-interval 10s
--health-timeout 5s
--health-retries 10
strategy:
fail-fast: false
matrix:
include:
- shard: roi-database
test-path: "tests/integration/database/test_roi_observed.py"
seed: none
workers: 0
timeout-minutes: 10
job-timeout-minutes: 35
- shard: proxy-behavior
test-path: "tests/proxy_behavior"
seed: db-push
workers: 0
timeout-minutes: 25
job-timeout-minutes: 50
- shard: proxy-security
test-path: "tests/proxy_security_tests"
seed: db-push
workers: 0
timeout-minutes: 15
job-timeout-minutes: 40
- shard: schema-migration
test-path: "tests/proxy_migration_tests"
seed: none
workers: 0
timeout-minutes: 20
job-timeout-minutes: 45
env:
DATABASE_URL: "postgresql://postgres:postgres@localhost:5432/litellm_test"
steps:
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
timeout-minutes: 3
with:
persist-credentials: false
- name: Detect relevant changes
id: changes
timeout-minutes: 2
uses: ./.github/actions/detect-changes
- name: Set up Python
if: steps.changes.outputs.decision != 'skip'
timeout-minutes: 3
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
with:
python-version: "3.12"
- name: Set up uv
if: steps.changes.outputs.decision != 'skip'
timeout-minutes: 3
uses: ./.github/actions/setup-uv-with-retries
with:
version: "0.10.9"
- name: Cache uv dependencies
if: steps.changes.outputs.decision != 'skip' && github.ref == 'refs/heads/main'
timeout-minutes: 5
uses: actions/cache@0057852bfaa89a56745cba8c7296529d2fc39830 # v4.3.0
with:
path: |
~/.cache/uv
.venv
key: ${{ runner.os }}-uv-postgres-${{ hashFiles('uv.lock') }}
restore-keys: |
${{ runner.os }}-uv-postgres-
- name: Cache uv dependencies
if: steps.changes.outputs.decision != 'skip' && github.ref != 'refs/heads/main'
timeout-minutes: 5
uses: actions/cache/restore@0057852bfaa89a56745cba8c7296529d2fc39830 # v4.3.0
with:
path: |
~/.cache/uv
.venv
key: ${{ runner.os }}-uv-postgres-${{ hashFiles('uv.lock') }}
restore-keys: |
${{ runner.os }}-uv-postgres-
- name: Install dependencies
if: steps.changes.outputs.decision != 'skip'
timeout-minutes: 12
run: |
.github/scripts/uv_sync_with_retries.sh --frozen --all-groups --all-extras
- name: Cache Prisma binaries
if: steps.changes.outputs.decision != 'skip'
timeout-minutes: 3
uses: ./.github/actions/cache-prisma-binaries
- name: Generate Prisma client
if: steps.changes.outputs.decision != 'skip'
timeout-minutes: 5
run: |
uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma
- name: Seed database schema
if: steps.changes.outputs.decision != 'skip' && matrix.seed != 'none'
timeout-minutes: 10
run: |
uv run --no-sync prisma db push --schema litellm/proxy/schema.prisma --accept-data-loss
- name: Run tests
if: steps.changes.outputs.decision != 'skip'
timeout-minutes: ${{ matrix.timeout-minutes }}
env:
TEST_PATH: ${{ matrix.test-path }}
WORKERS: ${{ matrix.workers }}
PYTEST_ADDOPTS: ${{ matrix.shard == 'proxy-behavior' && '--cov=./litellm --cov-report=xml:coverage-lens-postgres.xml' || matrix.shard == 'roi-database' && '--cov=./litellm --cov-report=xml:coverage-roi-postgres.xml' || '' }}
run: |
if [ "${WORKERS}" = "0" ]; then
uv run --no-sync pytest ${TEST_PATH:?} -vv --tb=short --durations=10
else
uv run --no-sync pytest ${TEST_PATH:?} -vv --tb=short --durations=10 -n "${WORKERS}"
fi
- name: Upload Lens database coverage
if: steps.changes.outputs.decision != 'skip' && matrix.shard == 'proxy-behavior' && !cancelled()
uses: codecov/codecov-action@303a32d7a59b442fa8d48b6a1cc6825c09c847a5 # v7.1.1
with:
use_oidc: true
version: v11.3.1
root_dir: ${{ github.workspace }}
files: coverage-lens-postgres.xml
flags: lens-postgres
fail_ci_if_error: true
- name: Upload ROI database coverage
if: steps.changes.outputs.decision != 'skip' && matrix.shard == 'roi-database' && !cancelled()
uses: codecov/codecov-action@303a32d7a59b442fa8d48b6a1cc6825c09c847a5 # v7.1.1
with:
use_oidc: true
version: v11.3.1
root_dir: ${{ github.workspace }}
files: coverage-roi-postgres.xml
flags: roi-postgres
fail_ci_if_error: true

View file

@ -1,106 +0,0 @@
name: "Unit Tests: Redis Client Version Compatibility"
on:
pull_request:
branches:
- main
- "litellm_**"
paths:
- "litellm/_redis.py"
- "litellm/_redis_credential_provider.py"
- "litellm/caching/redis_cache.py"
- "litellm/caching/evicted_client_closer.py"
- "tests/unit/test_redis.py"
- "tests/local_testing/test_caching.py"
- "tests/unit/caching/test_redis_connection_pool.py"
- "tests/unit/caching/test_redis_cluster_cache.py"
- "tests/unit/caching/test_evicted_client_closer.py"
- ".github/workflows/test-redis-compat.yml"
- "pyproject.toml"
- "uv.lock"
permissions:
contents: read
concurrency:
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }}
cancel-in-progress: true
jobs:
redis-compat:
name: "redis-py ${{ matrix.redis-version }}"
runs-on: ubuntu-latest
timeout-minutes: 15
permissions:
contents: read
id-token: write
strategy:
fail-fast: false
matrix:
# 5.3.1 is the version pinned in uv.lock (redisvl caps it below 6); the
# newer legs prove the inspect.signature introspection in litellm/_redis.py
# keeps extracting kwargs on the redis-py releases people actually run now.
# Only the exact release 6.0.0 is skipped: rq (pulled by the proxy extra)
# specifies `redis != 6`, which excludes 6.0.0 alone, so 6.4.0 stands in
# for the 6.x line.
redis-version: ["5.3.1", "6.4.0", "7.4.1", "8.0.1"]
steps:
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
persist-credentials: false
- name: Set up Python
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
with:
python-version: "3.12"
- name: Set up uv
uses: ./.github/actions/setup-uv-with-retries
with:
version: "0.10.9"
- name: Install dependencies
run: |
.github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra extra_proxy --extra semantic-router
- name: Pin redis-py to the matrix version
env:
REDIS_VERSION: ${{ matrix.redis-version }}
run: |
uv pip install "redis==${REDIS_VERSION:?}"
uv run --no-sync python -c "import redis; assert redis.__version__ == '${REDIS_VERSION:?}', redis.__version__; print('redis-py', redis.__version__)"
- name: Build Redis for cluster authentication tests
run: |
curl --fail --location --retry 3 https://download.redis.io/releases/redis-7.2.16.tar.gz -o "$RUNNER_TEMP/redis-7.2.16.tar.gz"
echo "960a8ec15e34ff40e57ff16837b26b33bd81f2da6d24497bb63de532a323a18e $RUNNER_TEMP/redis-7.2.16.tar.gz" | sha256sum --check
tar -xzf "$RUNNER_TEMP/redis-7.2.16.tar.gz" -C "$RUNNER_TEMP"
make -C "$RUNNER_TEMP/redis-7.2.16" -j2 MALLOC=libc OPTIMIZATION=-O1 redis-server
echo "$RUNNER_TEMP/redis-7.2.16/src" >> "$GITHUB_PATH"
- name: Run redis unit tests
run: |
redis-server --version
uv run --no-sync pytest \
tests/unit/test_redis.py \
tests/unit/caching/test_redis_connection_pool.py \
tests/unit/caching/test_redis_cluster_cache.py \
tests/unit/caching/test_evicted_client_closer.py \
tests/local_testing/test_caching.py::test_sync_cluster_authenticates_with_azure_credentials \
tests/local_testing/test_caching.py::test_sync_cluster_authenticates_with_gcp_credentials \
--tb=short -vv \
--reruns 2 \
--reruns-delay 1 \
--durations=20 \
--cov=./litellm --cov-report=xml:coverage-redis.xml
- name: Upload Redis coverage
if: matrix.redis-version == '5.3.1'
uses: codecov/codecov-action@0fb7174895f61a3b6b78fc075e0cd60383518dac # v5.5.5
with:
use_oidc: true
files: coverage-redis.xml
flags: redis-compat
fail_ci_if_error: false

View file

@ -20,13 +20,6 @@ concurrency:
# rather than alphabetical letter ranges. Adding a new test file means adding it
# to whichever group it belongs to, not reshuffling slices.
#
# `.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. 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.
# Most of a shard's time is pytest plugin load + xdist worker imports +
@ -78,93 +71,136 @@ jobs:
include:
# Must run serially — event-loop conflict with the logging worker.
- test-group: key-generation
test-path: ""
unit-flag: proxy-db-key-generation
test-path: >-
tests/unit/proxy/management_endpoints/test_key_generate_prisma.py
workers: 0
dist: loadscope
timeout: 20
# ---- auth: split into 2 shards ----
- test-group: auth-checks
test-path: ""
unit-flag: proxy-db-auth-checks
test-path: >-
tests/unit/proxy/auth/test_auth_checks.py
tests/unit/proxy/auth/test_user_api_key_auth.py
tests/unit/proxy/test_credential_slot_registry.py
tests/unit/proxy/test_deprecated_key_grace_period.py
workers: 4
dist: loadscope
timeout: 15
- test-group: jwt-and-keys
test-path: ""
unit-flag: proxy-db-jwt-and-keys
test-path: >-
tests/unit/proxy/auth/test_jwt.py
tests/unit/proxy/management_endpoints/test_jwt_key_mapping.py
tests/unit/proxy/test_proxy_custom_auth.py
workers: 4
dist: loadscope
timeout: 15
# ---- test_proxy_utils.py, single shard, worksteal distribution ----
- test-group: proxy-utils
test-path: ""
unit-flag: proxy-db-proxy-utils
test-path: >-
tests/unit/proxy/test_proxy_utils.py
workers: 4
dist: worksteal
timeout: 15
# ---- proxy server: split into 2 shards ----
- test-group: proxy-server-core
test-path: "tests/proxy_unit_tests/test_proxy_server_gemini_pass_through.py"
unit-flag: proxy-db-proxy-server-core
test-path: >-
tests/proxy_unit_tests/test_proxy_server_gemini_pass_through.py
tests/unit/proxy/test__lazy_features.py
tests/unit/proxy/test_aproxy_startup.py
tests/unit/proxy/test_proxy_server.py
workers: 4
dist: loadscope
timeout: 15
- test-group: proxy-runtime
test-path: ""
unit-flag: proxy-db-proxy-runtime
test-path: >-
tests/unit/proxy/auth/test_multipart_bypass_repro.py
tests/unit/proxy/auth/test_proxy_routes.py
tests/unit/proxy/middleware/test_request_size_limit_middleware.py
tests/unit/proxy/test_proxy_config_unit_test.py
tests/unit/proxy/test_proxy_token_counter.py
tests/unit/proxy/test_server_root_path.py
workers: 4
dist: loadscope
timeout: 15
# ---- logging: split into 2 shards ----
- test-group: custom-logging
test-path: "tests/proxy_unit_tests/test_proxy_custom_logger.py"
unit-flag: proxy-db-custom-logging
test-path: >-
tests/proxy_unit_tests/test_proxy_custom_logger.py
tests/unit/proxy/test_custom_callback_input.py
tests/unit/proxy/test_custom_logger_s3_gcs.py
workers: 4
dist: loadscope
timeout: 15
- test-group: logging-misc
test-path: ""
unit-flag: proxy-db-logging-misc
test-path: >-
tests/unit/proxy/management_helpers/test_audit_logs_proxy.py
tests/unit/proxy/spend_tracking/test_search_api_logging.py
tests/unit/proxy/test_proxy_reject_logging.py
workers: 4
dist: loadscope
timeout: 15
- test-group: db-and-spend
test-path: ""
unit-flag: proxy-db-db-and-spend
test-path: >-
tests/unit/proxy/common_utils/test_proxy_encrypt_decrypt.py
tests/unit/proxy/db/db_transaction_queue/test_e2e_pod_lock_manager.py
tests/unit/proxy/db/test_update_daily_tag_spend.py
tests/unit/proxy/test_db_schema_changes.py
tests/unit/proxy/test_prisma_client_backoff_retry.py
tests/unit/proxy/test_update_spend.py
tests/unit/skills/test_skills_db.py
workers: 4
dist: loadscope
timeout: 15
# ---- guardrails + budget + hooks: split into 2 ----
- test-group: guardrails-hooks
test-path: ""
unit-flag: proxy-db-guardrails-hooks
test-path: >-
tests/unit/proxy/hooks/test_banned_keyword_list.py
tests/unit/proxy/test_proxy_setting_guardrails.py
tests/unit/proxy/test_unit_test_proxy_hooks.py
workers: 4
dist: loadscope
timeout: 15
- test-group: budgets
test-path: ""
unit-flag: proxy-db-budgets
test-path: >-
tests/unit/proxy/auth/test_default_end_user_budget_simple.py
tests/unit/proxy/hooks/test_unit_test_max_model_budget_limiter.py
tests/unit/proxy/test_zero_cost_model_budget_bypass.py
workers: 4
dist: loadscope
timeout: 15
- test-group: endpoints-and-responses
test-path: "tests/proxy_unit_tests/test_proxy_exception_mapping.py"
unit-flag: proxy-db-endpoints-and-responses
test-path: >-
tests/proxy_unit_tests/test_proxy_exception_mapping.py
tests/unit/proxy/lens
tests/unit/proxy/auth/test_models_fallback_endpoint.py
tests/unit/proxy/common_utils/test_check_batch_cost.py
tests/unit/proxy/common_utils/test_check_responses_cost.py
tests/unit/proxy/common_utils/test_realtime_cache.py
tests/unit/proxy/google_endpoints/test_gemini_agents_endpoints.py
tests/unit/proxy/google_endpoints/test_google_endpoint_routing.py
tests/unit/proxy/google_endpoints/test_google_gemini_proxy_request.py
tests/unit/proxy/public_endpoints/test_blog_posts_endpoint.py
tests/unit/proxy/response_polling
tests/unit/proxy/test_custom_tokenizer_bug.py
tests/unit/proxy/test_get_favicon.py
tests/unit/proxy/test_get_image.py
tests/unit/proxy/test_prompt_test_endpoint.py
tests/unit/proxy/test_reducto_ocr_route.py
tests/unit/proxy/test_response_polling_pre_call_checks.py
tests/unit/proxy/test_ui_path_detection.py
workers: 4
dist: loadscope
timeout: 15
uses: ./.github/workflows/_test-unit-base.yml
with:
test-path: ${{ matrix.test-path }}
unit-flag: ${{ matrix.unit-flag }}
workers: ${{ matrix.workers }}
reruns: 2
timeout-minutes: ${{ matrix.timeout }}

View file

@ -35,10 +35,6 @@ concurrency:
# a matrix and carries a shard-coverage guard that reads that file by name.
# Folding it in here is a follow-up, together with generalising that guard into
# assert_ci_coverage.py.
#
# `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 }}
@ -50,19 +46,11 @@ jobs:
fail-fast: false
matrix:
include:
- shard: mcp-integration
artifact-name: mcp-integration
test-path: "tests/mcp_tests"
unit-flag: mcp-integration
workers: 2
reruns: 0
timeout-minutes: 20
job-timeout-minutes: 60
- shard: core-utils
artifact-name: core-utils
test-path: tests/unit/decisions
unit-flag: core-utils
test-path: >-
tests/unit/decisions
tests/unit/litellm_core_utils
workers: 2
reruns: 1
timeout-minutes: 20
@ -70,8 +58,22 @@ jobs:
- shard: enterprise-routing
artifact-name: enterprise-routing
test-path: ""
unit-flag: enterprise-routing
test-path: >-
tests/unit/google_genai
tests/unit/router_strategy
tests/unit/router_utils
tests/unit/proxy/common_utils/test_cache_aware_routing.py
tests/unit/enterprise/enterprise_callbacks/send_emails
tests/unit/enterprise/proxy/test_afile_retrieve_returns_unified_id.py
tests/unit/enterprise/proxy/test_batch_retrieve_input_file_id.py
tests/unit/enterprise/proxy/test_batch_retrieve_registers_missing_output_file_id.py
tests/unit/enterprise/proxy/test_batch_retrieve_returns_unified_input_file_id.py
tests/unit/enterprise/proxy/test_batch_update_db_managed_output_file_id.py
tests/unit/enterprise/proxy/test_deleted_file_returns_403_not_404.py
tests/unit/enterprise/proxy/test_enterprise_routes.py
tests/unit/enterprise/proxy/test_file_deletion_blocking.py
tests/unit/enterprise/proxy/test_managed_files_access_check.py
tests/unit/enterprise/proxy/test_managed_files_hook.py
workers: 2
reruns: 2
timeout-minutes: 20
@ -82,7 +84,7 @@ jobs:
test-path: >-
tests/test_litellm/integrations
tests/test_litellm/tracing
unit-flag: integrations
tests/unit/integrations
workers: 2
reruns: 3
timeout-minutes: 20
@ -90,8 +92,8 @@ jobs:
- shard: Vertex AI
artifact-name: llm-vertex-ai
test-path: ""
unit-flag: llm-vertex-ai
test-path: >-
tests/unit/llms/vertex_ai
workers: 1
reruns: 2
timeout-minutes: 20
@ -99,8 +101,10 @@ jobs:
- shard: All Other Providers
artifact-name: llm-other-providers
test-path: ""
unit-flag: llm-other-providers
test-path: >-
tests/unit/llms
--ignore=tests/unit/llms/vertex_ai
--ignore=tests/unit/llms/base_llm/batches/base_batches_config_test.py
workers: 2
reruns: 2
timeout-minutes: 20
@ -110,7 +114,27 @@ jobs:
artifact-name: misc
test-path: >-
tests/test_litellm/test_*.py
unit-flag: misc
tests/unit/test_*.py
tests/unit/test_router
tests/unit/a2a_protocol
tests/unit/batches
tests/unit/chat_completions
tests/unit/completion_extras
tests/unit/containers
tests/unit/embeddings
tests/unit/endpoints
tests/unit/files
tests/unit/harness
tests/unit/images
tests/unit/interactions
tests/unit/messages
tests/unit/rag
tests/unit/rerank_api
tests/unit/rust_bridge
tests/unit/secret_managers
tests/unit/vector_stores
tests/unit/videos
--ignore=tests/unit/rust_bridge/native_route_wheel_test.py
workers: 2
reruns: 2
timeout-minutes: 20
@ -218,7 +242,10 @@ jobs:
tests/unit/proxy/enterprise_billing
tests/unit/proxy/types_utils
tests/unit/proxy/logging_endpoints
unit-flag: proxy-infra
tests/unit/gateway
tests/unit/proxy/management
tests/unit/proxy/management_endpoints/test_roi_calculator_endpoints.py
tests/unit/proxy/roi_calculator
workers: 4
reruns: 2
timeout-minutes: 20
@ -260,8 +287,8 @@ jobs:
- shard: caching-local
artifact-name: caching-local
test-path: ""
unit-flag: caching-local
test-path: >-
tests/unit/caching
workers: 2
reruns: 2
timeout-minutes: 20
@ -269,8 +296,8 @@ jobs:
- shard: proxy-extras
artifact-name: proxy-extras
test-path: ""
unit-flag: proxy-extras
test-path: >-
tests/unit/litellm_proxy_extras
workers: 2
reruns: 2
timeout-minutes: 20
@ -278,8 +305,15 @@ jobs:
- shard: enterprise-package
artifact-name: enterprise-package
test-path: ""
unit-flag: enterprise-package
test-path: >-
tests/unit/enterprise/integrations
tests/unit/enterprise/proxy/auth
tests/unit/enterprise/proxy/guardrails
tests/unit/enterprise/proxy/hooks
tests/unit/enterprise/proxy/management_endpoints
tests/unit/enterprise/proxy/test_audit_logging_endpoints.py
tests/unit/enterprise/proxy/test_liteadmin.py
tests/unit/enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py
workers: 4
reruns: 2
timeout-minutes: 20
@ -287,19 +321,41 @@ jobs:
- shard: responses-caching-types
artifact-name: responses-caching-types
test-path: ""
unit-flag: responses-caching-types
test-path: >-
tests/unit/responses
tests/unit/types
--ignore=tests/unit/responses/mcp
workers: 2
reruns: 2
timeout-minutes: 20
job-timeout-minutes: 60
- shard: unit
artifact-name: unit
test-path: >-
tests/unit/anthropic_interface
tests/unit/compression
tests/unit/enterprise/enterprise_callbacks/test_callback_controls.py
tests/unit/enterprise/enterprise_callbacks/test_llm_guard.py
tests/unit/enterprise/enterprise_callbacks/test_secret_detection.py
tests/unit/integration_support
tests/unit/models
tests/unit/ocr
tests/unit/passthrough
tests/unit/realtime_api
tests/unit/repositories
tests/unit/sandbox
tests/unit/skills/test_skills_main.py
tests/unit/tracing
workers: 2
reruns: 0
timeout-minutes: 20
job-timeout-minutes: 60
uses: ./.github/workflows/_test-unit-base.yml
with:
test-path: ${{ matrix.test-path }}
unit-flag: ${{ matrix.unit-flag || '' }}
workers: ${{ matrix.workers }}
reruns: ${{ matrix.reruns }}
timeout-minutes: ${{ matrix.timeout-minutes }}
job-timeout-minutes: ${{ matrix.job-timeout-minutes }}
artifact-name: ${{ matrix.artifact-name }}
legacy-mcp-peer: ${{ matrix.shard == 'mcp-integration' }}

View file

@ -46,6 +46,8 @@ from litellm.proxy._types import (
UserAPIKeyAuth,
WebhookEvent,
)
from litellm.repositories.table_repositories import InvitationLinkRepository
from litellm.repositories.user_repository import UserRepository
from litellm.secret_managers.main import get_secret_bool
from litellm.types.integrations.slack_alerting import LITELLM_LOGO_URL
@ -844,7 +846,7 @@ class BaseEmailLogger(CustomLogger):
)
return None
user_row = await prisma_client.db.litellm_usertable.find_unique(
user_row = await UserRepository(prisma_client).table.find_unique(
where={"user_id": user_id}
)
@ -929,7 +931,7 @@ class BaseEmailLogger(CustomLogger):
try:
# Try to get existing invitation
existing_invitations = (
await prisma_client.db.litellm_invitationlink.find_many(
await InvitationLinkRepository(prisma_client).table.find_many(
where={"user_id": user_id},
order={"created_at": "desc"},
)

View file

@ -15,6 +15,10 @@ from litellm.constants import (
MANAGED_OBJECT_STALENESS_CUTOFF_DAYS,
MAX_OBJECTS_PER_POLL_CYCLE,
)
from litellm.repositories.table_repositories import ManagedObjectRepository
from litellm.repositories.team_repository import TeamRepository
from litellm.repositories.user_repository import UserRepository
from litellm.repositories.verification_token_repository import VerificationTokenRepository
if TYPE_CHECKING:
from prisma import models as prisma_models
@ -58,25 +62,19 @@ class _ManagedObjectRow(Protocol):
def _managed_object_table(prisma_client: "PrismaClient") -> "TableActions[_ManagedObjectRow]":
table: Final[TableActions[_ManagedObjectRow]] = prisma_client.db.litellm_managedobjecttable
return table
return ManagedObjectRepository(prisma_client).table
def _user_table(prisma_client: "PrismaClient") -> "TableActions[prisma_models.LiteLLM_UserTable]":
table: Final[TableActions[prisma_models.LiteLLM_UserTable]] = prisma_client.db.litellm_usertable
return table
return UserRepository(prisma_client).table
def _token_table(prisma_client: "PrismaClient") -> "TableActions[prisma_models.LiteLLM_VerificationToken]":
table: Final[TableActions[prisma_models.LiteLLM_VerificationToken]] = (
prisma_client.db.litellm_verificationtoken
)
return table
return VerificationTokenRepository(prisma_client).table
def _team_table(prisma_client: "PrismaClient") -> "TableActions[prisma_models.LiteLLM_TeamTable]":
table: Final[TableActions[prisma_models.LiteLLM_TeamTable]] = prisma_client.db.litellm_teamtable
return table
return TeamRepository(prisma_client).table
class CheckBatchCost:

View file

@ -16,6 +16,7 @@ from litellm.constants import (
MAX_OBJECTS_PER_POLL_CYCLE,
STALE_OBJECT_CLEANUP_BATCH_SIZE,
)
from litellm.repositories.table_repositories import ManagedObjectRepository
from litellm.responses.utils import ResponsesAPIRequestUtils
from litellm.types.llms.openai import ResponsesAPIResponse
from litellm.types.utils import BACKGROUND_RESPONSE_COST_POLL_CALL_ORIGIN
@ -43,8 +44,7 @@ class _ManagedObjectRow(Protocol):
def _managed_object_table(prisma_client: "PrismaClient") -> "TableActions[_ManagedObjectRow]":
table: Final[TableActions[_ManagedObjectRow]] = prisma_client.db.litellm_managedobjecttable
return table
return ManagedObjectRepository(prisma_client).table
class CheckResponsesCost:

View file

@ -1647,7 +1647,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
owner_filter: Final = build_owner_filter(user_api_key_dict)
if owner_filter is None:
return FileListPage(**build_list_page([]))
return FileListPage.model_validate(build_list_page([]))
if after:
cursor_row = await _managed_file_table(self.prisma_client).find_first(
@ -1686,7 +1686,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
cursor_id = chunk[-1].unified_file_id
chunk_size = max(chunk_size, FILE_LIST_CONTINUATION_CHUNK_SIZE)
return FileListPage(**build_list_page(matches[:page_size], has_more=len(matches) > page_size))
return FileListPage.model_validate(build_list_page(matches[:page_size], has_more=len(matches) > page_size))
def _is_batch_polling_enabled(self) -> bool:
"""

View file

@ -4085,26 +4085,6 @@ dependencies = [
"strum",
]
[[package]]
name = "litellm-migrate"
version = "0.1.0"
dependencies = [
"litellm-migrate-macros",
"rstest",
]
[[package]]
name = "litellm-migrate-macros"
version = "0.1.0"
dependencies = [
"proc-macro2",
"quote",
"rstest",
"syn 2.0.119",
"tempfile",
"thiserror 2.0.19",
]
[[package]]
name = "litellm-model-catalog"
version = "0.1.0"
@ -4170,6 +4150,7 @@ dependencies = [
"serde_json",
"serde_with",
"sha2 0.10.9",
"sqlx",
"strum",
"thiserror 2.0.19",
"tokio",
@ -4368,6 +4349,9 @@ dependencies = [
"rstest",
"serde",
"serde_json",
"serde_with",
"sqlx",
"testcontainers-modules",
"thiserror 2.0.19",
"tokio",
"url",
@ -4498,7 +4482,6 @@ dependencies = [
"hmac 0.12.1",
"jsonschema",
"litellm-http",
"litellm-migrate",
"litellm-storage-clickhouse",
"litellm-traces",
"litellm-traces-cache",
@ -4509,6 +4492,7 @@ dependencies = [
"serde",
"serde_json",
"sha2 0.10.9",
"sqlx",
"strum",
"testcontainers-modules",
"thiserror 2.0.19",

View file

@ -16,8 +16,6 @@ litellm-traces = { path = "crates/traces" }
litellm-traces-cache = { path = "crates/traces-cache" }
litellm-traces-clickhouse = { path = "crates/traces-clickhouse" }
litellm-storage-clickhouse = { path = "crates/storage-clickhouse" }
litellm-migrate = { path = "crates/migrate" }
litellm-migrate-macros = { path = "crates/migrate-macros" }
litellm-core = { path = "crates/core" }
litellm-gateway-mcp = { path = "crates/gateway-mcp" }
litellm-gateway = { path = "crates/gateway" }

View file

@ -1,19 +0,0 @@
[package]
name = "litellm-migrate-macros"
version = "0.1.0"
edition.workspace = true
license.workspace = true
repository.workspace = true
[lib]
proc-macro = true
[dependencies]
proc-macro2.workspace = true
quote.workspace = true
syn = { workspace = true, features = ["parsing", "printing", "proc-macro"] }
thiserror.workspace = true
[dev-dependencies]
rstest.workspace = true
tempfile.workspace = true

View file

@ -1,21 +0,0 @@
use std::io;
#[derive(Debug, thiserror::Error)]
pub enum Error {
#[error("could not read migrations directory `{path}`")]
ReadDirectory {
path: String,
#[source]
source: io::Error,
},
#[error(
"migration name `{name}` must be `<digits>_<description>.sql` with a `[a-z0-9_]` description"
)]
InvalidName { name: String },
#[error("migration version `{version}` is declared more than once")]
DuplicateVersion { version: u64 },
#[error("migrations directory `{path}` contains no migrations")]
Empty { path: String },
#[error("migration path `{path}` is not valid UTF-8")]
NonUtf8Path { path: String },
}

View file

@ -1,199 +0,0 @@
mod error;
use std::path::{Path, PathBuf};
use error::Error;
use proc_macro::TokenStream;
use quote::quote;
use syn::LitStr;
struct Entry {
version: u64,
description: String,
path: PathBuf,
}
fn resolve(dir: &Path) -> Result<Vec<Entry>, Error> {
let mut entries = Vec::new();
let files = std::fs::read_dir(dir).map_err(|source| Error::ReadDirectory {
path: dir.display().to_string(),
source,
})?;
for file in files {
let file = file.map_err(|source| Error::ReadDirectory {
path: dir.display().to_string(),
source,
})?;
let path = file.path();
let name = path
.file_name()
.and_then(|name| name.to_str())
.ok_or_else(|| Error::NonUtf8Path {
path: path.display().to_string(),
})?
.to_owned();
let invalid = || Error::InvalidName { name: name.clone() };
let stem = name
.strip_suffix(".sql")
.filter(|_| file.file_type().is_ok_and(|kind| kind.is_file()))
.and_then(|stem| stem.split_once('_'))
.filter(|(version, description)| {
!version.is_empty()
&& version.bytes().all(|b| b.is_ascii_digit())
&& !description.is_empty()
&& description
.bytes()
.all(|b| b.is_ascii_lowercase() || b.is_ascii_digit() || b == b'_')
})
.ok_or_else(invalid)?;
let version = stem.0.parse::<u64>().map_err(|_| invalid())?;
entries.push(Entry {
version,
description: stem.1.to_owned(),
path,
});
}
if entries.is_empty() {
return Err(Error::Empty {
path: dir.display().to_string(),
});
}
entries.sort_by_key(|entry| entry.version);
for pair in entries.windows(2) {
if pair[0].version == pair[1].version {
return Err(Error::DuplicateVersion {
version: pair[0].version,
});
}
}
Ok(entries)
}
fn resolve_input(lit: &LitStr) -> Result<Vec<Entry>, Error> {
let root = std::env::var("CARGO_MANIFEST_DIR")
.map(PathBuf::from)
.unwrap_or_default();
let dir = root.join(lit.value());
let dir = dir.canonicalize().map_err(|source| Error::ReadDirectory {
path: dir.display().to_string(),
source,
})?;
if dir.to_str().is_none() {
return Err(Error::NonUtf8Path {
path: dir.display().to_string(),
});
}
resolve(&dir)
}
#[proc_macro]
pub fn migrate(input: TokenStream) -> TokenStream {
let lit = syn::parse_macro_input!(input as LitStr);
match resolve_input(&lit) {
Ok(entries) => {
let migrations = entries.iter().map(|entry| {
let version = entry.version;
let description = &entry.description;
let path = entry
.path
.to_str()
.expect("canonical migration path is UTF-8");
quote! {
::litellm_migrate::Migration {
version: #version,
description: #description,
sql: ::core::include_str!(#path),
}
}
});
quote! { &[#(#migrations),*] }.into()
}
Err(err) => syn::Error::new(lit.span(), err).to_compile_error().into(),
}
}
#[cfg(test)]
mod tests {
use std::fs;
use rstest::rstest;
use tempfile::TempDir;
use super::{Error, resolve};
fn migrations_dir(files: &[&str]) -> TempDir {
let dir = TempDir::new().expect("tempdir");
for file in files {
fs::write(dir.path().join(file), "SELECT 1").expect("write fixture");
}
dir
}
#[rstest]
fn orders_versions_numerically() {
let dir = migrations_dir(&["10_tenth.sql", "2_second.sql", "1_first.sql"]);
let entries = resolve(dir.path()).expect("resolves");
let versions: Vec<u64> = entries.iter().map(|entry| entry.version).collect();
let descriptions: Vec<&str> = entries
.iter()
.map(|entry| entry.description.as_str())
.collect();
assert_eq!(versions, [1, 2, 10]);
assert_eq!(descriptions, ["first", "second", "tenth"]);
}
#[rstest]
#[case::dash_in_version(&["0001-dash.sql"])]
#[case::not_sql(&["notes.txt"])]
#[case::empty_description(&["0001_.sql"])]
#[case::non_digit_version(&["x_name.sql"])]
#[case::uppercase_description(&["0001_Upper.sql"])]
#[case::no_underscore(&["0001.sql"])]
#[case::plus_sign_version(&["+10_add.sql"])]
fn rejects_invalid_names(#[case] files: &[&str]) {
let dir = migrations_dir(files);
assert!(matches!(
resolve(dir.path()),
Err(Error::InvalidName { .. })
));
}
#[rstest]
fn rejects_subdirectories() {
let dir = migrations_dir(&["0001_a.sql"]);
fs::create_dir(dir.path().join("0002_b.sql")).expect("subdir");
assert!(matches!(
resolve(dir.path()),
Err(Error::InvalidName { .. })
));
}
#[cfg(unix)]
#[rstest]
fn rejects_symlinks() {
let dir = migrations_dir(&["0001_a.sql"]);
let target = TempDir::new().expect("tempdir");
let target_file = target.path().join("real.sql");
fs::write(&target_file, "SELECT 2").expect("write fixture");
std::os::unix::fs::symlink(&target_file, dir.path().join("0002_b.sql")).expect("symlink");
assert!(matches!(
resolve(dir.path()),
Err(Error::InvalidName { .. })
));
}
#[rstest]
fn rejects_duplicate_versions() {
let dir = migrations_dir(&["0001_a.sql", "1_b.sql"]);
assert!(matches!(
resolve(dir.path()),
Err(Error::DuplicateVersion { version: 1 })
));
}
#[rstest]
fn rejects_empty_directory() {
let dir = migrations_dir(&[]);
assert!(matches!(resolve(dir.path()), Err(Error::Empty { .. })));
}
}

View file

@ -1,12 +0,0 @@
[package]
name = "litellm-migrate"
version = "0.1.0"
edition.workspace = true
license.workspace = true
repository.workspace = true
[dependencies]
litellm-migrate-macros.workspace = true
[dev-dependencies]
rstest.workspace = true

View file

@ -1,5 +0,0 @@
# Migrations
`litellm-migrate` exports the `Migration` struct and the `migrate!` macro that embeds a directory of `<digits>_<description>.sql` files at compile time, sorted by numeric version
The crate does not apply or track migrations; callers decide how and when the embedded SQL runs

View file

@ -1,8 +0,0 @@
pub use litellm_migrate_macros::migrate;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct Migration {
pub version: u64,
pub description: &'static str,
pub sql: &'static str,
}

View file

@ -1 +0,0 @@
SELECT 10;

View file

@ -1 +0,0 @@
SELECT 1;

View file

@ -1 +0,0 @@
SELECT 2;

View file

@ -1,21 +0,0 @@
use litellm_migrate::Migration;
use rstest::rstest;
const MIGRATIONS: &[Migration] = litellm_migrate::migrate!("tests/fixtures/migrations");
#[rstest]
#[case::first(0, 1, "first", include_str!("fixtures/migrations/1_first.sql"))]
#[case::second(1, 2, "second", include_str!("fixtures/migrations/2_second.sql"))]
#[case::tenth(2, 10, "tenth", include_str!("fixtures/migrations/10_tenth.sql"))]
fn embeds_every_file_sorted_by_numeric_version(
#[case] index: usize,
#[case] version: u64,
#[case] description: &str,
#[case] sql: &str,
) {
assert_eq!(MIGRATIONS.len(), 3);
let migration = &MIGRATIONS[index];
assert_eq!(migration.version, version);
assert_eq!(migration.description, description);
assert_eq!(migration.sql, sql);
}

View file

@ -74,6 +74,7 @@ serde_with.workspace = true
criterion.workspace = true
futures-util.workspace = true
rstest.workspace = true
sqlx = { workspace = true, features = ["migrate"] }
sha2.workspace = true
tokio-tungstenite.workspace = true
wiremock.workspace = true

View file

@ -52,8 +52,7 @@ fn map_error_ref(error: &Error) -> PyErr {
| Error::InvalidParameters
| Error::InvalidScope => PyValueError::new_err(error.to_string()),
Error::Task
| Error::SchemaFailed(_)
| Error::SchemaTransport
| Error::Migration(_)
| Error::MissingSecret
| Error::Busy
| Error::ProvisionFailed(_)
@ -452,7 +451,14 @@ mod tests {
)]
#[case::insert_budget(Error::InsertTooLarge, "OverflowError")]
#[case::scope(Error::InvalidScope, "ValueError")]
#[case::schema(Error::SchemaFailed(503), "RuntimeError")]
#[case::schema(
Error::Storage(litellm_storage_clickhouse::Error::SchemaFailed(503)),
"RuntimeError"
)]
#[case::migration(
Error::Migration(sqlx::migrate::MigrateError::VersionMismatch(1)),
"RuntimeError"
)]
#[case::reader(Error::MissingSecret, "RuntimeError")]
#[case::storage(
Error::Storage(litellm_storage_clickhouse::Error::InvalidUrl),

View file

@ -3,3 +3,7 @@
`litellm-storage-clickhouse` exports `Storage`, a writer and bounded reader derived from one ClickHouse URL and database. It also exports bounded HTTP read and insert execution
The crate has no trace tables, OTLP types, or named trace queries. `litellm-traces-clickhouse` supplies those rules and uses this storage for both trace rows and spend rows
It also applies embedded SQLx migrations through the `_sqlx_migrations` ledger
Migration files are append-only, and changed applied files are rejected by their checksums. Startup migrations must be replay-safe schema changes because the runner records success after execution without dirty states or locks. Backfills belong in coordinated jobs outside proxy startup. The replay policy lives in `ClickHouseMigrate::apply`, `dirty_version`, and `lock`; a Keeper-backed or deploy-time runner changes only those methods

View file

@ -11,11 +11,14 @@ flate2.workspace = true
litellm-http.workspace = true
serde.workspace = true
serde_json.workspace = true
serde_with.workspace = true
sqlx = { workspace = true, features = ["migrate"] }
thiserror.workspace = true
url.workspace = true
[dev-dependencies]
litellm-http = { workspace = true, features = ["test-support"] }
rstest.workspace = true
testcontainers-modules = { version = "0.15.0", features = ["clickhouse"] }
tokio.workspace = true
wiremock.workspace = true

View file

@ -1,4 +1,4 @@
#[derive(Debug, thiserror::Error)]
#[derive(Clone, Debug, thiserror::Error)]
pub enum Error {
#[error("invalid ClickHouse insert row")]
InvalidRow,

View file

@ -1,9 +1,11 @@
mod error;
mod insert;
mod migrate;
mod read;
pub use error::Error;
pub use insert::{insert_compressed_rows, insert_encoded_rows};
pub use migrate::{ClickHouseMigrate, execute_statement, storage_error};
pub use read::{Parameter, Query, READ_LIMITS, ReadLimits, execute_read, fetch, fetch_json};
use url::Url;

View file

@ -0,0 +1,308 @@
use std::{
future::Future,
pin::Pin,
time::{Duration, Instant},
};
use litellm_http::Client;
use serde::Deserialize;
use serde_with::{DisplayFromStr, PickFirst, serde_as};
use sqlx::{
Error as SqlxError,
migrate::{AppliedMigration, Migrate, MigrateError, Migration},
};
use crate::{Connection, Error, READ_LIMITS, valid_identifier};
pub async fn execute_statement(
client: &Client,
connection: &Connection,
sql: &str,
timeout: Duration,
) -> Result<(), Error> {
let body = execute_sql(client, connection, sql, timeout).await?;
if !body.trim().is_empty() {
return Err(Error::InvalidResponse);
}
Ok(())
}
async fn execute_sql(
client: &Client,
connection: &Connection,
sql: &str,
timeout: Duration,
) -> Result<String, Error> {
let mut url = connection.url().clone();
let pairs: Vec<_> = url
.query_pairs()
.filter(|(key, _)| {
!matches!(
key.as_ref(),
"query" | "wait_end_of_query" | "send_progress_in_http_headers" | "async_insert"
)
})
.map(|(key, value)| (key.into_owned(), value.into_owned()))
.collect();
url.query_pairs_mut()
.clear()
.extend_pairs(pairs)
.append_pair("wait_end_of_query", "1")
.append_pair("send_progress_in_http_headers", "0")
.append_pair("async_insert", "0");
let mut response = client
.post(url)
.timeout(timeout)
.body(sql.to_owned())
.send()
.await
.map_err(|_| Error::Transport)?;
if !response.status().is_success() {
return Err(Error::SchemaFailed(response.status().as_u16()));
}
let mut body = Vec::new();
while let Some(chunk) = response.chunk().await.map_err(|_| Error::Transport)? {
if body.len() + chunk.len() > READ_LIMITS.response_bytes {
return Err(Error::ResponseTooLarge);
}
body.extend_from_slice(&chunk);
}
String::from_utf8(body).map_err(|_| Error::InvalidResponse)
}
/// The startup runner records success after execution, never marks migrations dirty, and skips locks
/// A Keeper-backed or deploy-time runner changes only `apply`, `dirty_version`, and `lock`
pub struct ClickHouseMigrate<'a, R> {
client: &'a Client,
connection: &'a Connection,
database: &'a str,
render: R,
timeout: Duration,
}
impl<'a, R> ClickHouseMigrate<'a, R>
where
R: Fn(&str) -> String + Send + Sync,
{
pub fn new(
client: &'a Client,
connection: &'a Connection,
database: &'a str,
render: R,
timeout: Duration,
) -> Result<Self, Error> {
if !valid_identifier(database) {
return Err(Error::InvalidSchema);
}
Ok(Self {
client,
connection,
database,
render,
timeout,
})
}
}
#[serde_as]
#[derive(Deserialize)]
struct Applied {
#[serde_as(as = "PickFirst<(_, DisplayFromStr)>")]
version: i64,
checksum: String,
}
fn migrate_error(error: Error) -> MigrateError {
MigrateError::Execute(SqlxError::AnyDriverError(Box::new(error)))
}
fn migrate_execution_error(error: Error, version: i64) -> MigrateError {
MigrateError::ExecuteMigration(SqlxError::AnyDriverError(Box::new(error)), version)
}
fn encode_hex(bytes: &[u8]) -> String {
bytes
.iter()
.map(|byte| format!("{byte:02x}"))
.collect::<Vec<_>>()
.join("")
}
fn decode_hex(value: &str) -> Result<Vec<u8>, Error> {
let (pairs, remainder) = value.as_bytes().as_chunks::<2>();
if !remainder.is_empty() {
return Err(Error::InvalidResponse);
}
pairs
.iter()
.map(|pair| {
let high = decode_hex_digit(pair[0]).ok_or(Error::InvalidResponse)?;
let low = decode_hex_digit(pair[1]).ok_or(Error::InvalidResponse)?;
Ok((high << 4) | low)
})
.collect()
}
fn decode_hex_digit(value: u8) -> Option<u8> {
match value {
b'0'..=b'9' => Some(value - b'0'),
b'a'..=b'f' => Some(value - b'a' + 10),
b'A'..=b'F' => Some(value - b'A' + 10),
_ => None,
}
}
fn escape_sql_string(value: &str) -> String {
value.replace('\\', "\\\\").replace('\'', "\\'")
}
type MigrateFuture<'e, T> = Pin<Box<dyn Future<Output = T> + Send + 'e>>;
impl<R> Migrate for ClickHouseMigrate<'_, R>
where
R: Fn(&str) -> String + Send + Sync,
{
fn create_schema_if_not_exists<'e>(
&'e mut self,
schema_name: &'e str,
) -> MigrateFuture<'e, Result<(), MigrateError>> {
Box::pin(async move {
if !valid_identifier(schema_name) {
return Err(migrate_error(Error::InvalidSchema));
}
let statement = format!("CREATE DATABASE IF NOT EXISTS `{schema_name}`");
execute_statement(self.client, self.connection, &statement, self.timeout)
.await
.map_err(migrate_error)
})
}
fn ensure_migrations_table<'e>(
&'e mut self,
table_name: &'e str,
) -> MigrateFuture<'e, Result<(), MigrateError>> {
Box::pin(async move {
let database = format!("`{}`", self.database);
execute_statement(
self.client,
self.connection,
&format!("CREATE DATABASE IF NOT EXISTS {database}"),
self.timeout,
)
.await
.map_err(migrate_error)?;
execute_statement(
self.client,
self.connection,
&format!(
"CREATE TABLE IF NOT EXISTS {database}.{table_name} \
(version Int64, description String, installed_on DateTime64(3) DEFAULT now64(3), \
success Bool, checksum String, execution_time Int64) ENGINE = MergeTree ORDER BY version"
),
self.timeout,
)
.await
.map_err(migrate_error)
})
}
fn dirty_version<'e>(
&'e mut self,
_table_name: &'e str,
) -> MigrateFuture<'e, Result<Option<i64>, MigrateError>> {
Box::pin(async { Ok(None) })
}
fn list_applied_migrations<'e>(
&'e mut self,
table_name: &'e str,
) -> MigrateFuture<'e, Result<Vec<AppliedMigration>, MigrateError>> {
Box::pin(async move {
let database = format!("`{}`", self.database);
let statement = format!(
"SELECT DISTINCT version, checksum FROM {database}.{table_name} \
WHERE success ORDER BY version FORMAT JSONEachRow"
);
let response = execute_sql(self.client, self.connection, &statement, self.timeout)
.await
.map_err(migrate_error)?;
response
.lines()
.filter(|line| !line.trim().is_empty())
.map(|line| {
let row = serde_json::from_str::<Applied>(line)
.map_err(|_| migrate_error(Error::InvalidResponse))?;
let checksum = decode_hex(&row.checksum).map_err(migrate_error)?;
Ok(AppliedMigration {
version: row.version,
checksum: checksum.into(),
})
})
.collect()
})
}
fn lock(&mut self) -> MigrateFuture<'_, Result<(), MigrateError>> {
Box::pin(async { Ok(()) })
}
fn unlock(&mut self) -> MigrateFuture<'_, Result<(), MigrateError>> {
Box::pin(async { Ok(()) })
}
fn apply<'e>(
&'e mut self,
table_name: &'e str,
migration: &'e Migration,
) -> MigrateFuture<'e, Result<Duration, MigrateError>> {
Box::pin(async move {
let started_at = Instant::now();
let statement = (self.render)(migration.sql.as_str());
execute_statement(self.client, self.connection, &statement, self.timeout)
.await
.map_err(|error| migrate_execution_error(error, migration.version))?;
let elapsed = started_at.elapsed();
let execution_time = elapsed.as_nanos().min(i64::MAX as u128) as i64;
let description = escape_sql_string(&migration.description);
let checksum = encode_hex(&migration.checksum);
let database = format!("`{}`", self.database);
execute_statement(
self.client,
self.connection,
&format!(
"INSERT INTO {database}.{table_name} \
(version, description, success, checksum, execution_time) \
VALUES ({}, '{}', true, '{}', {execution_time})",
migration.version, description, checksum
),
self.timeout,
)
.await
.map_err(migrate_error)?;
Ok(elapsed)
})
}
fn revert<'e>(
&'e mut self,
_table_name: &'e str,
_migration: &'e Migration,
) -> MigrateFuture<'e, Result<Duration, MigrateError>> {
Box::pin(async {
Err(MigrateError::Execute(SqlxError::AnyDriverError(Box::new(
std::io::Error::other("ClickHouse migrations are forward-only"),
))))
})
}
}
pub fn storage_error(error: &MigrateError) -> Option<&Error> {
let error = match error {
MigrateError::Execute(error) | MigrateError::ExecuteMigration(error, _) => error,
_ => return None,
};
match error {
SqlxError::AnyDriverError(error) => error.downcast_ref(),
_ => None,
}
}

View file

@ -0,0 +1,476 @@
use std::time::Duration;
use litellm_http::Client;
use litellm_storage_clickhouse::{
ClickHouseMigrate, Connection, Error, READ_LIMITS, execute_statement, storage_error,
};
use rstest::{fixture, rstest};
use sqlx::{
SqlStr,
migrate::{Migrate, MigrateError, Migration, MigrationType, Migrator},
};
use testcontainers_modules::{
clickhouse::ClickHouse,
testcontainers::{ContainerAsync, ImageExt, runners::AsyncRunner},
};
use wiremock::{
Mock, MockServer, ResponseTemplate,
matchers::{body_string, method, query_param},
};
const CLICKHOUSE_TAG: &str =
"26.9.6.6@sha256:eb4870e7ca7ed70c259eebfcfbee6cf797017f6b5436c2926bbbfe3d4d28486e";
const DATABASE: &str = "storage_migrate_test";
const REQUEST_TIMEOUT: Duration = Duration::from_secs(10);
const SELECT_APPLIED: &str = "SELECT DISTINCT version, checksum FROM `trace_test`._sqlx_migrations \
WHERE success ORDER BY version FORMAT JSONEachRow";
type TestResult<T = ()> = Result<T, Box<dyn std::error::Error>>;
struct ClickHouseDatabase {
_container: ContainerAsync<ClickHouse>,
url: String,
client: Client,
}
#[fixture]
async fn database() -> TestResult<ClickHouseDatabase> {
let container = ClickHouse::default()
.with_tag(CLICKHOUSE_TAG)
.with_env_var("CLICKHOUSE_SKIP_USER_SETUP", "1")
.start()
.await?;
let url = format!(
"http://{}:{}",
container.get_host().await?,
container.get_host_port_ipv4(8123).await?
);
Ok(ClickHouseDatabase {
_container: container,
url,
client: Client::no_redirect_for_test(),
})
}
#[fixture]
async fn mock_server() -> MockServer {
let server = MockServer::start().await;
Mock::given(method("POST"))
.respond_with(ResponseTemplate::new(200))
.with_priority(10)
.mount(&server)
.await;
server
}
fn migration(version: i64, sql: &'static str) -> Migration {
Migration::new(
version,
format!("migration_{version}").into(),
MigrationType::Simple,
SqlStr::from_static(sql),
false,
)
}
fn migrator(migrations: Vec<Migration>) -> Migrator {
Migrator {
ignore_missing: true,
locking: false,
..Migrator::with_migrations(migrations)
}
}
fn render_database(sql: &str) -> String {
sql.replace("{database}", &format!("`{DATABASE}`"))
}
async fn run_migrations<R>(
database: &ClickHouseDatabase,
migrator: &Migrator,
schema: &str,
render: R,
) -> Result<(), MigrateError>
where
R: Fn(&str) -> String + Send + Sync,
{
let connection = Connection::writer(&database.url).expect("valid ClickHouse URL");
let mut adapter = ClickHouseMigrate::new(
&database.client,
&connection,
schema,
render,
REQUEST_TIMEOUT,
)
.expect("valid schema");
migrator.run_direct(None, &mut adapter, false).await
}
async fn execute_write(database: &ClickHouseDatabase, sql: &str) -> TestResult {
database
.client
.post(&database.url)
.body(sql.to_owned())
.send()
.await?
.error_for_status()?;
Ok(())
}
async fn read_json(database: &ClickHouseDatabase, sql: &str) -> TestResult<serde_json::Value> {
let response = database
.client
.post(&database.url)
.body(sql.to_owned())
.send()
.await?
.error_for_status()?;
Ok(serde_json::from_str(&response.text().await?)?)
}
async fn ledger_versions(database: &ClickHouseDatabase) -> TestResult<Vec<i64>> {
let response = read_json(
database,
&format!(
"SELECT version FROM `{DATABASE}`._sqlx_migrations \
GROUP BY version ORDER BY version FORMAT JSON"
),
)
.await?;
Ok(response["data"]
.as_array()
.expect("ClickHouse returns versions")
.iter()
.map(|row| row["version"].as_i64().expect("version is Int64"))
.collect())
}
fn encode_hex(bytes: &[u8]) -> String {
bytes
.iter()
.map(|byte| format!("{byte:02x}"))
.collect::<Vec<_>>()
.join("")
}
#[rstest]
#[tokio::test]
async fn only_pending_migrations_execute_on_the_second_run(
#[future(awt)] database: TestResult<ClickHouseDatabase>,
) -> TestResult {
let database = database?;
let migrator = migrator(vec![
migration(
1,
"CREATE TABLE IF NOT EXISTS {database}.migration_one (id UInt8) ENGINE = MergeTree ORDER BY id",
),
migration(
2,
"CREATE TABLE IF NOT EXISTS {database}.migration_two (id UInt8) ENGINE = MergeTree ORDER BY id",
),
]);
run_migrations(&database, &migrator, DATABASE, render_database).await?;
run_migrations(&database, &migrator, DATABASE, render_database).await?;
execute_statement(
&database.client,
&Connection::writer(&database.url)?,
"SYSTEM FLUSH LOGS",
REQUEST_TIMEOUT,
)
.await?;
let queries = read_json(
&database,
"SELECT count() AS executions FROM system.query_log \
WHERE type = 'QueryFinish' AND query LIKE \
'CREATE TABLE IF NOT EXISTS `storage_migrate_test`.migration_%' FORMAT JSON",
)
.await?;
assert_eq!(queries["data"][0]["executions"].as_u64(), Some(2));
assert_eq!(ledger_versions(&database).await?, vec![1, 2]);
Ok(())
}
#[rstest]
#[tokio::test]
async fn edited_migration_checksum_returns_version_mismatch(
#[future(awt)] database: TestResult<ClickHouseDatabase>,
) -> TestResult {
let database = database?;
let original = migrator(vec![migration(
1,
"CREATE TABLE IF NOT EXISTS {database}.original (id UInt8) ENGINE = MergeTree ORDER BY id",
)]);
let changed = migrator(vec![migration(
1,
"CREATE TABLE IF NOT EXISTS {database}.changed (id UInt8) ENGINE = MergeTree ORDER BY id",
)]);
run_migrations(&database, &original, DATABASE, render_database).await?;
assert!(matches!(
run_migrations(&database, &changed, DATABASE, render_database).await,
Err(MigrateError::VersionMismatch(1))
));
Ok(())
}
#[rstest]
#[tokio::test]
async fn failed_migration_is_not_recorded_and_retains_storage_error(
#[future(awt)] database: TestResult<ClickHouseDatabase>,
) -> TestResult {
let database = database?;
let migrator = migrator(vec![migration(1, "THIS IS NOT VALID CLICKHOUSE SQL")]);
let error = run_migrations(&database, &migrator, DATABASE, render_database)
.await
.expect_err("invalid SQL must fail");
assert!(matches!(&error, MigrateError::ExecuteMigration(_, 1)));
assert!(matches!(
storage_error(&error),
Some(Error::SchemaFailed(_))
));
let rows = read_json(
&database,
&format!(
"SELECT count() AS rows FROM `{DATABASE}`._sqlx_migrations \
WHERE version = 1 FORMAT JSON"
),
)
.await?;
assert_eq!(rows["data"][0]["rows"].as_u64(), Some(0));
Ok(())
}
#[rstest]
fn invalid_database_identifier_is_rejected() {
let client = Client::no_redirect_for_test();
let connection = Connection::writer("http://127.0.0.1:1").expect("valid URL");
assert!(matches!(
ClickHouseMigrate::new(
&client,
&connection,
"storage_test; DROP DATABASE default",
str::to_owned,
REQUEST_TIMEOUT,
),
Err(Error::InvalidSchema)
));
}
#[rstest]
#[tokio::test]
async fn unknown_source_version_is_tolerated(
#[future(awt)] database: TestResult<ClickHouseDatabase>,
) -> TestResult {
let database = database?;
run_migrations(&database, &migrator(vec![]), DATABASE, render_database).await?;
execute_write(
&database,
&format!(
"INSERT INTO `{DATABASE}`._sqlx_migrations \
(version, description, success, checksum, execution_time) \
VALUES (99, 'unknown', true, '{}', 0)",
"00".repeat(48)
),
)
.await?;
let migrator = migrator(vec![migration(
1,
"CREATE TABLE IF NOT EXISTS {database}.known (id UInt8) ENGINE = MergeTree ORDER BY id",
)]);
run_migrations(&database, &migrator, DATABASE, render_database).await?;
assert_eq!(ledger_versions(&database).await?, vec![1, 99]);
Ok(())
}
#[rstest]
#[tokio::test]
async fn duplicate_ledger_rows_are_tolerated(
#[future(awt)] database: TestResult<ClickHouseDatabase>,
) -> TestResult {
let database = database?;
let applied = migration(
1,
"CREATE TABLE IF NOT EXISTS {database}.duplicate_test (id UInt8) ENGINE = MergeTree ORDER BY id",
);
let checksum = encode_hex(&applied.checksum);
let migrator = migrator(vec![applied]);
run_migrations(&database, &migrator, DATABASE, render_database).await?;
execute_write(
&database,
&format!(
"INSERT INTO `{DATABASE}`._sqlx_migrations \
(version, description, success, checksum, execution_time) \
VALUES (1, 'migration_1', true, '{checksum}', 0)"
),
)
.await?;
run_migrations(&database, &migrator, DATABASE, render_database).await?;
let rows = read_json(
&database,
&format!(
"SELECT count() AS rows, uniqExact(version) AS versions \
FROM `{DATABASE}`._sqlx_migrations FORMAT JSON"
),
)
.await?;
assert_eq!(rows["data"][0]["rows"].as_u64(), Some(2));
assert_eq!(rows["data"][0]["versions"].as_u64(), Some(1));
Ok(())
}
#[rstest]
#[tokio::test]
async fn revert_reports_forward_only_error() {
let client = Client::no_redirect_for_test();
let connection = Connection::writer("http://127.0.0.1:1").expect("valid URL");
let migration = migration(
1,
"CREATE TABLE IF NOT EXISTS {database}.revert_test (id UInt8) ENGINE = MergeTree ORDER BY id",
);
let mut adapter = ClickHouseMigrate::new(
&client,
&connection,
DATABASE,
render_database,
REQUEST_TIMEOUT,
)
.expect("valid schema");
let error = adapter
.revert("_sqlx_migrations", &migration)
.await
.expect_err("ClickHouse migrations cannot be reverted");
assert!(
error
.to_string()
.contains("ClickHouse migrations are forward-only")
);
}
#[rstest]
#[case::numeric("1")]
#[case::quoted("\"1\"")]
#[tokio::test]
async fn applied_int64_versions_accept_numeric_and_quoted_json(
#[future(awt)] mock_server: MockServer,
#[case] version: &str,
) {
Mock::given(method("POST"))
.and(body_string(SELECT_APPLIED))
.respond_with(
ResponseTemplate::new(200)
.set_body_string(format!("{{\"version\":{version},\"checksum\":\"00\"}}\n")),
)
.mount(&mock_server)
.await;
let client = Client::no_redirect_for_test();
let connection = Connection::writer(&mock_server.uri()).expect("valid URL");
let migrator = migrator(vec![]);
let mut adapter = ClickHouseMigrate::new(
&client,
&connection,
"trace_test",
str::to_owned,
REQUEST_TIMEOUT,
)
.expect("valid schema");
migrator
.run_direct(None, &mut adapter, false)
.await
.expect("applied version parses");
}
#[rstest]
#[tokio::test]
async fn schema_requests_override_unsafe_connection_settings(
#[future(awt)] mock_server: MockServer,
) {
Mock::given(method("POST"))
.and(query_param("wait_end_of_query", "1"))
.and(query_param("send_progress_in_http_headers", "0"))
.and(query_param("async_insert", "0"))
.and(query_param("custom_setting", "preserved"))
.respond_with(ResponseTemplate::new(200))
.expect(1)
.mount(&mock_server)
.await;
let url = format!(
"{}/?wait_end_of_query=0&send_progress_in_http_headers=1&async_insert=1&custom_setting=preserved",
mock_server.uri()
);
execute_statement(
&Client::no_redirect_for_test(),
&Connection::writer(&url).expect("valid URL"),
"CREATE DATABASE IF NOT EXISTS trace_test",
REQUEST_TIMEOUT,
)
.await
.expect("schema execution succeeds");
let requests = mock_server
.received_requests()
.await
.expect("requests recorded");
for name in [
"wait_end_of_query",
"send_progress_in_http_headers",
"async_insert",
] {
assert_eq!(
requests[0]
.url
.query_pairs()
.filter(|(key, _)| key == name)
.count(),
1
);
}
}
#[rstest]
#[tokio::test]
async fn oversized_ledger_response_is_rejected_before_migrations(
#[future(awt)] mock_server: MockServer,
) {
Mock::given(method("POST"))
.and(body_string(SELECT_APPLIED))
.respond_with(
ResponseTemplate::new(200).set_body_string(" ".repeat(READ_LIMITS.response_bytes + 1)),
)
.mount(&mock_server)
.await;
let client = Client::no_redirect_for_test();
let connection = Connection::writer(&mock_server.uri()).expect("valid URL");
let migrator = migrator(vec![]);
let mut adapter = ClickHouseMigrate::new(
&client,
&connection,
"trace_test",
str::to_owned,
REQUEST_TIMEOUT,
)
.expect("valid schema");
let error = migrator
.run_direct(None, &mut adapter, false)
.await
.expect_err("oversized result is rejected");
assert!(matches!(
storage_error(&error),
Some(Error::ResponseTooLarge)
));
assert_eq!(
mock_server
.received_requests()
.await
.expect("requests recorded")
.len(),
3
);
}

View file

@ -15,6 +15,7 @@ use litellm_traces::{
ListTracesParams, ReadAccessParams, SpanDetailParams, SpanErrorParams, TraceIdentityParams,
TraceSpansParams,
},
request::{TRACE_PAGE_SIZE_MAX, TRACE_PAGE_SIZE_MIN},
resolve_trace, to_ui_content,
};
@ -133,7 +134,7 @@ impl TraceReader {
cursor: Option<&str>,
page_size: u32,
) -> Result<Option<Trace>, ReadError<S::Error>> {
if !(1..=500).contains(&page_size) {
if !(u32::from(TRACE_PAGE_SIZE_MIN)..=u32::from(TRACE_PAGE_SIZE_MAX)).contains(&page_size) {
return Err(ReadError::InvalidParameters);
}
let Some(trace_ref) = reference(store, access, trace_id, trace_ref).await? else {

View file

@ -1,6 +1,7 @@
- Own trace schema, row encoding, SQL query adapters and reader provisioning; consume domain types from `litellm-traces`
- Keep generic ClickHouse connections and HTTP execution in `litellm-storage-clickhouse`; keep PyO3 conversion in `python-bridge`
- Keep schema definitions only in `migrations/NNNN_description.sql`, embedded by `litellm_migrate::migrate!`
- Keep schema definitions only in `migrations/NNNN_description.sql`, embedded by `sqlx::migrate!`
- Treat retention TTLs as current configuration: change them in `RETENTION` in `src/schema.rs`, which every startup reapplies, never in a new migration
- Require typed query parameters and SELECT-only readers with server-side limits and tenant isolation
- Bound insert time and encoded bytes; preserve shared values and explicit retry deduplication
- Test storage behavior through the public API against ClickHouse

View file

@ -16,7 +16,6 @@ flate2.workspace = true
futures-util.workspace = true
hmac = "0.12.1"
litellm-http.workspace = true
litellm-migrate.workspace = true
litellm-storage-clickhouse.workspace = true
litellm-traces.workspace = true
litellm-traces-cache.workspace = true
@ -24,6 +23,7 @@ moka.workspace = true
serde.workspace = true
serde_json.workspace = true
sha2.workspace = true
sqlx = { workspace = true, features = ["migrate", "macros"] }
strum.workspace = true
thiserror.workspace = true
time = { workspace = true, features = ["formatting"] }

View file

@ -16,10 +16,6 @@ pub enum Error {
InvalidResponse,
#[error("ClickHouse insert exceeds the encoded size limit")]
InsertTooLarge,
#[error("ClickHouse schema setup failed with HTTP status {0}")]
SchemaFailed(u16),
#[error("ClickHouse schema setup transport failed")]
SchemaTransport,
#[error("trace SQL queries require a configured proxy master key")]
MissingSecret,
#[error("invalid trace query scope")]
@ -39,5 +35,7 @@ pub enum Error {
#[error(transparent)]
Storage(#[from] litellm_storage_clickhouse::Error),
#[error(transparent)]
Migration(#[from] sqlx::migrate::MigrateError),
#[error(transparent)]
Cached(#[from] std::sync::Arc<Error>),
}

View file

@ -33,7 +33,8 @@ pub use query::{QueryHelp, execute_read, query_help, query_sql};
pub use query_access::QueryReaders;
pub use reads::ClickHouseTraces;
pub use schema::{
NORMALIZED_FIELD_DEFINITIONS, NormalizedFieldDefinition, ensure_schema, schema_statements,
NORMALIZED_FIELD_DEFINITIONS, NormalizedFieldDefinition, apply_migrations, ensure_schema,
reconcile_retention, schema_statements,
};
pub use span_row::span_rows;
pub use sql::execute_named_read;

View file

@ -1,15 +1,26 @@
use litellm_http::Client;
use litellm_migrate::Migration;
use litellm_storage_clickhouse::{ClickHouseMigrate, execute_statement, storage_error};
use serde::Serialize;
use sqlx::migrate::Migrator;
use std::time::Duration;
use super::{Connection, Error};
const SCHEMA_REQUEST_TIMEOUT: Duration = Duration::from_secs(30);
const MIGRATIONS: &[Migration] = litellm_migrate::migrate!("migrations");
static MIGRATOR: Migrator = Migrator {
ignore_missing: true,
locking: false,
..sqlx::migrate!("./migrations")
};
pub fn schema_statements(database: &str, retention_days: u32) -> Result<Vec<String>, Error> {
const RETENTION: [(&str, &str); 3] = [
("otel_traces", "toDateTime(Timestamp)"),
("agent_traces_by_key", "toDateTime(StartTs)"),
("spend_logs", "toDateTime(start_time)"),
];
fn validate_schema(database: &str, retention_days: u32) -> Result<(), Error> {
if database.is_empty()
|| !database
.bytes()
@ -18,19 +29,109 @@ pub fn schema_statements(database: &str, retention_days: u32) -> Result<Vec<Stri
{
return Err(Error::InvalidSchema);
}
Ok(())
}
fn render(sql: &str, database: &str, retention_days: u32) -> String {
sql.replace("{database}", database)
.replace("{retention_days}", &retention_days.to_string())
}
pub fn schema_statements(database: &str, retention_days: u32) -> Result<Vec<String>, Error> {
validate_schema(database, retention_days)?;
let database = format!("`{database}`");
Ok(
std::iter::once(format!("CREATE DATABASE IF NOT EXISTS {database}"))
.chain(MIGRATIONS.iter().map(|migration| {
migration
.sql
.replace("{database}", &database)
.replace("{retention_days}", &retention_days.to_string())
}))
.chain(
MIGRATOR
.migrations
.iter()
.map(|migration| render(migration.sql.as_str(), &database, retention_days)),
)
.collect(),
)
}
pub async fn apply_migrations(
client: &Client,
connection: &Connection,
database: &str,
retention_days: u32,
) -> Result<(), Error> {
apply_migrations_with_timeout(
client,
connection,
database,
retention_days,
SCHEMA_REQUEST_TIMEOUT,
)
.await
}
async fn apply_migrations_with_timeout(
client: &Client,
connection: &Connection,
database: &str,
retention_days: u32,
request_timeout: Duration,
) -> Result<(), Error> {
validate_schema(database, retention_days)?;
let quoted_database = format!("`{database}`");
let mut adapter = ClickHouseMigrate::new(
client,
connection,
database,
|sql| render(sql, &quoted_database, retention_days),
request_timeout,
)?;
MIGRATOR
.run_direct(None, &mut adapter, false)
.await
.map_err(|error| match storage_error(&error) {
Some(storage_error) => Error::Storage(storage_error.clone()),
None => Error::Migration(error),
})
}
pub async fn reconcile_retention(
client: &Client,
connection: &Connection,
database: &str,
retention_days: u32,
) -> Result<(), Error> {
reconcile_retention_with_timeout(
client,
connection,
database,
retention_days,
SCHEMA_REQUEST_TIMEOUT,
)
.await
}
async fn reconcile_retention_with_timeout(
client: &Client,
connection: &Connection,
database: &str,
retention_days: u32,
request_timeout: Duration,
) -> Result<(), Error> {
validate_schema(database, retention_days)?;
let database = format!("`{database}`");
for (table, expression) in RETENTION {
execute_statement(
client,
connection,
&format!(
"ALTER TABLE {database}.{table} MODIFY TTL {expression} + INTERVAL {retention_days} DAY"
),
request_timeout,
)
.await?;
}
Ok(())
}
pub async fn ensure_schema(
client: &Client,
connection: &Connection,
@ -54,19 +155,22 @@ async fn ensure_schema_with_timeout(
retention_days: u32,
request_timeout: Duration,
) -> Result<(), Error> {
for statement in schema_statements(database, retention_days)? {
let response = client
.post(connection.url().clone())
.timeout(request_timeout)
.body(statement)
.send()
.await
.map_err(|_| Error::SchemaTransport)?;
if !response.status().is_success() {
return Err(Error::SchemaFailed(response.status().as_u16()));
}
}
Ok(())
apply_migrations_with_timeout(
client,
connection,
database,
retention_days,
request_timeout,
)
.await?;
reconcile_retention_with_timeout(
client,
connection,
database,
retention_days,
request_timeout,
)
.await
}
#[derive(Clone, Copy, Debug, Eq, PartialEq, Serialize)]

View file

@ -1,11 +1,14 @@
use std::{collections::BTreeMap, time::Duration};
use litellm_http::Client;
use litellm_storage_clickhouse::Error as StorageError;
use litellm_traces_clickhouse::{
Connection, Error, InsertTable, NORMALIZED_FIELD_DEFINITIONS, Parameter, ReadQuery,
encode_rows, ensure_schema, execute_named_read, execute_read, schema_statements,
apply_migrations, encode_rows, ensure_schema, execute_named_read, execute_read,
reconcile_retention, schema_statements,
};
use rstest::rstest;
use sqlx::migrate::MigrateError;
mod support;
use support::{ClickHouseDatabase, TestResult, database};
@ -71,6 +74,44 @@ async fn mutation_rows(database: &ClickHouseDatabase) -> TestResult<u64> {
.expect("ClickHouse returns mutation counts as unsigned integers"))
}
fn migration_versions() -> Vec<u64> {
let mut versions = std::fs::read_dir(concat!(env!("CARGO_MANIFEST_DIR"), "/migrations"))
.expect("migration directory exists")
.map(|entry| {
entry
.expect("migration directory entry is readable")
.file_name()
.into_string()
.expect("migration file name is UTF-8")
})
.filter_map(|name| {
name.strip_suffix(".sql")
.and_then(|stem| stem.split('_').next())
.and_then(|version| version.parse::<u64>().ok())
})
.collect::<Vec<_>>();
versions.sort_unstable();
versions
}
async fn migration_ledger_versions(database: &ClickHouseDatabase) -> TestResult<Vec<u64>> {
let response = read_json(
database,
"SELECT version FROM trace_test._sqlx_migrations GROUP BY version ORDER BY version",
)
.await?;
Ok(response["data"]
.as_array()
.expect("ClickHouse returns version rows")
.iter()
.map(|row| {
row["version"]
.as_u64()
.expect("ClickHouse returns versions as unsigned integers")
})
.collect())
}
#[rstest]
#[tokio::test]
async fn schema_supports_span_rollups_and_spend_joins(
@ -80,6 +121,27 @@ async fn schema_supports_span_rollups_and_spend_joins(
let writer = Connection::writer(&database.url)?;
ensure_schema(&database.client, &writer, "trace_test", 7).await?;
ensure_schema(&database.client, &writer, "trace_test", 7).await?;
let expected_versions = migration_versions();
let ledger = read_json(
&database,
"SELECT count() AS rows, uniqExact(version) AS versions, \
countIf(NOT match(checksum, '^[0-9a-f]{96}$')) AS invalid_checksums \
FROM trace_test._sqlx_migrations",
)
.await?;
assert_eq!(
ledger["data"][0]["rows"].as_u64(),
Some(expected_versions.len() as u64)
);
assert_eq!(
ledger["data"][0]["versions"].as_u64(),
Some(expected_versions.len() as u64)
);
assert_eq!(ledger["data"][0]["invalid_checksums"].as_u64(), Some(0));
assert_eq!(
migration_ledger_versions(&database).await?,
expected_versions
);
let timestamp = time::OffsetDateTime::now_utc().unix_timestamp_nanos() as i64;
let span = serde_json::from_value(serde_json::json!({
"Timestamp": timestamp, "TraceId": "trace-1", "SpanId": "span-1", "ParentSpanId": "",
@ -202,6 +264,112 @@ async fn schema_supports_span_rollups_and_spend_joins(
Ok(())
}
#[rstest]
#[case::quoted_versions("output_format_json_quote_64bit_integers=1")]
#[case::asynchronous_inserts(
"async_insert=1&wait_for_async_insert=0&async_insert_busy_timeout_ms=20000"
)]
#[tokio::test]
async fn schema_setup_records_migrations_synchronously_with_configured_settings(
#[future(awt)] database: TestResult<ClickHouseDatabase>,
#[case] settings: &str,
) -> TestResult {
let database = database?;
let writer = Connection::writer(&format!("{}?{settings}", database.url))?;
ensure_schema(&database.client, &writer, "trace_test", 7).await?;
assert_eq!(
migration_ledger_versions(&database).await?,
migration_versions()
);
ensure_schema(&database.client, &writer, "trace_test", 7).await?;
assert_eq!(
migration_ledger_versions(&database).await?,
migration_versions()
);
Ok(())
}
#[rstest]
#[tokio::test]
async fn changed_migration_is_rejected(
#[future(awt)] database: TestResult<ClickHouseDatabase>,
) -> TestResult {
let database = database?;
let writer = Connection::writer(&database.url)?;
ensure_schema(&database.client, &writer, "trace_test", 7).await?;
execute_write(
&database,
"ALTER TABLE trace_test._sqlx_migrations UPDATE checksum = '00' \
WHERE version = 1 SETTINGS mutations_sync = 1",
)
.await?;
assert!(matches!(
ensure_schema(&database.client, &writer, "trace_test", 7).await,
Err(Error::Migration(MigrateError::VersionMismatch(1)))
));
Ok(())
}
#[rstest]
#[tokio::test]
async fn concurrent_schema_setup_succeeds(
#[future(awt)] database: TestResult<ClickHouseDatabase>,
) -> TestResult {
let database = database?;
let writer = Connection::writer(&database.url)?;
let (first, second, third, fourth) = tokio::join!(
ensure_schema(&database.client, &writer, "trace_test", 7),
ensure_schema(&database.client, &writer, "trace_test", 7),
ensure_schema(&database.client, &writer, "trace_test", 7),
ensure_schema(&database.client, &writer, "trace_test", 7),
);
for result in [first, second, third, fourth] {
result?;
}
let tables = read_json(
&database,
"SELECT count() AS tables FROM system.tables \
WHERE database = 'trace_test' AND name IN \
('otel_traces', 'agent_traces_by_key', 'spend_logs')",
)
.await?;
assert_eq!(tables["data"][0]["tables"].as_u64(), Some(3));
assert_eq!(
migration_ledger_versions(&database).await?,
migration_versions()
);
Ok(())
}
#[rstest]
#[tokio::test]
async fn existing_schema_without_ledger_is_adopted(
#[future(awt)] database: TestResult<ClickHouseDatabase>,
) -> TestResult {
let database = database?;
let writer = Connection::writer(&database.url)?;
for statement in schema_statements("trace_test", 7)? {
execute_write(&database, &statement).await?;
}
let timestamp = time::OffsetDateTime::now_utc().unix_timestamp_nanos() as i64;
let span = serde_json::from_value(serde_json::json!({
"Timestamp": timestamp, "TraceId": "trace-adopted", "SpanId": "span-adopted",
"ParentSpanId": "", "ServiceName": "proxy", "SpanName": "request",
"Input": "existing row", "ResourceAttributes": {}, "SpanAttributes": {}
}))?;
insert_rows(&database, "otel_traces", vec![span]).await?;
ensure_schema(&database.client, &writer, "trace_test", 7).await?;
assert_eq!(table_rows(&database, "otel_traces").await?, 1);
assert_eq!(
migration_ledger_versions(&database).await?,
migration_versions()
);
Ok(())
}
#[rstest]
#[tokio::test]
async fn normalized_fields_match_clickhouse_catalog(
@ -785,6 +953,60 @@ async fn retention_changes_materialize_existing_rows_and_remain_idempotent(
Ok(())
}
#[rstest]
#[tokio::test]
async fn retention_reconciliation_updates_each_table_ttl(
#[future(awt)] database: TestResult<ClickHouseDatabase>,
) -> TestResult {
let database = database?;
let writer = Connection::writer(&database.url)?;
apply_migrations(&database.client, &writer, "trace_test", 7).await?;
reconcile_retention(&database.client, &writer, "trace_test", 7).await?;
let ttl_queries = read_json(
&database,
"SELECT name, create_table_query FROM system.tables \
WHERE database = 'trace_test' AND name IN \
('otel_traces', 'agent_traces_by_key', 'spend_logs') ORDER BY name",
)
.await?;
let ttl_queries = ttl_queries["data"].as_array().expect("retention tables");
assert_eq!(
ttl_queries
.iter()
.map(|row| row["name"].as_str().expect("table name"))
.collect::<Vec<_>>(),
["agent_traces_by_key", "otel_traces", "spend_logs"]
);
for row in ttl_queries {
let query = row["create_table_query"]
.as_str()
.expect("table creation query");
assert!(
query.contains("toIntervalDay(7)") || query.contains("INTERVAL 7 DAY"),
"{query}"
);
}
reconcile_retention(&database.client, &writer, "trace_test", 3).await?;
let ttl_queries = read_json(
&database,
"SELECT name, create_table_query FROM system.tables \
WHERE database = 'trace_test' AND name IN \
('otel_traces', 'agent_traces_by_key', 'spend_logs') ORDER BY name",
)
.await?;
for row in ttl_queries["data"].as_array().expect("retention tables") {
let query = row["create_table_query"]
.as_str()
.expect("table creation query");
assert!(
query.contains("toIntervalDay(3)") || query.contains("INTERVAL 3 DAY"),
"{query}"
);
}
Ok(())
}
#[rstest]
#[tokio::test]
async fn schema_statement_timeout_maps_to_transport_error() -> TestResult {
@ -804,7 +1026,7 @@ async fn schema_statement_timeout_maps_to_transport_error() -> TestResult {
.await;
server.abort();
assert!(
matches!(result, Ok(Err(Error::SchemaTransport))),
matches!(result, Ok(Err(Error::Storage(StorageError::Transport)))),
"{result:?}"
);
Ok(())

View file

@ -4,3 +4,8 @@
- Keep ClickHouse schema, row encoding and queries in `litellm-traces-clickhouse`; keep PyO3 conversion in `python-bridge`
- Test decoding and normalization through the public API
- Expose one top-level `Error` enum in `src/error.rs` for decoding and normalization failures
- Own the tracing HTTP contracts: request types in `src/request.rs` and response views exported from `src/schema.rs`
- Python and the dashboard consume them only through generated code: `uv run scripts/generate_trace_types.py` writes `litellm/rust_bridge/trace/generated/`, and `npm run gen:api` in `ui/litellm-dashboard` regenerates `schema.d.ts` from the proxy's OpenAPI
- Declare each bound once as a constant and read it from both the schema attribute and the runtime check
- GET request types accept unknown fields because the routes ignore unknown query parameters; body request types use `deny_unknown_fields`
- Changing a request or response shape changes the public API; ship it in its own behavior-change PR

View file

@ -1,6 +1,8 @@
fn main() {
println!(
"{}",
serde_json::to_string_pretty(&litellm_traces::schema::schemas()).unwrap()
);
let schemas = if std::env::args().nth(1).as_deref() == Some("--requests") {
litellm_traces::schema::request_schemas()
} else {
litellm_traces::schema::schemas()
};
println!("{}", serde_json::to_string_pretty(&schemas).unwrap());
}

View file

@ -15,6 +15,7 @@ mod normalize;
mod otlp;
pub mod query;
mod query_access;
pub mod request;
mod resolve;
#[cfg(feature = "schema")]
pub mod schema;

View file

@ -0,0 +1,56 @@
pub const TRACE_PAGE_SIZE_MIN: u16 = 1;
pub const TRACE_PAGE_SIZE_MAX: u16 = 500;
#[macro_rules_attribute::apply(request_type)]
#[derive(Clone, Debug)]
pub struct TraceListRequest {
/// Window start, unix ms. Default: 24h ago
#[serde(default)]
pub start_ms: Option<i64>,
/// Window end, unix ms. Default: now
#[serde(default)]
pub end_ms: Option<i64>,
#[serde(default)]
#[cfg_attr(feature = "schema", schemars(length(max = 512)))]
pub cursor: Option<String>,
}
#[macro_rules_attribute::apply(request_type)]
#[derive(Clone, Debug)]
pub struct TraceDetailRequest {
#[serde(default)]
pub trace_ref: String,
#[serde(default)]
#[cfg_attr(feature = "schema", schemars(length(max = 512)))]
pub cursor: Option<String>,
#[serde(default)]
#[cfg_attr(
feature = "schema",
schemars(range(min = TRACE_PAGE_SIZE_MIN, max = TRACE_PAGE_SIZE_MAX))
)]
pub page_size: Option<u16>,
}
#[macro_rules_attribute::apply(request_type)]
#[derive(Clone, Debug)]
pub struct TraceSpanRequest {
#[serde(default)]
pub trace_ref: String,
}
#[macro_rules_attribute::apply(request_type)]
#[derive(Clone, Debug)]
pub struct TraceErrorPageRequest {
#[serde(default)]
pub trace_ref: String,
#[serde(default)]
#[cfg_attr(feature = "schema", schemars(length(max = 512)))]
pub cursor: Option<String>,
}
#[macro_rules_attribute::apply(request_type)]
#[derive(Clone, Debug)]
#[serde(deny_unknown_fields)]
pub struct TraceQueryRequest {
pub sql: String,
}

View file

@ -36,6 +36,13 @@ fn received<T: JsonSchema>() -> Schema {
.into_root_schema_for::<T>()
}
fn requested<T: JsonSchema>() -> Schema {
SchemaSettings::draft2020_12()
.for_deserialize()
.into_generator()
.into_root_schema_for::<T>()
}
fn emitted<T: JsonSchema>() -> Schema {
SchemaSettings::draft2020_12()
.for_serialize()
@ -58,3 +65,28 @@ pub fn schemas() -> BTreeMap<&'static str, Schema> {
("SpanErrorPage", emitted::<crate::SpanErrorPage>()),
])
}
pub fn request_schemas() -> BTreeMap<&'static str, Schema> {
BTreeMap::from([
(
"TraceListRequest",
requested::<crate::request::TraceListRequest>(),
),
(
"TraceDetailRequest",
requested::<crate::request::TraceDetailRequest>(),
),
(
"TraceSpanRequest",
requested::<crate::request::TraceSpanRequest>(),
),
(
"TraceErrorPageRequest",
requested::<crate::request::TraceErrorPageRequest>(),
),
(
"TraceQueryRequest",
requested::<crate::request::TraceQueryRequest>(),
),
])
}

View file

@ -0,0 +1,93 @@
#![cfg(feature = "schema")]
use litellm_traces::request::{
TraceDetailRequest, TraceErrorPageRequest, TraceListRequest, TraceQueryRequest,
TraceSpanRequest,
};
use litellm_traces::schema::request_schemas;
use rstest::rstest;
use serde_json::json;
#[rstest]
#[case::list("TraceListRequest")]
#[case::detail("TraceDetailRequest")]
#[case::span("TraceSpanRequest")]
#[case::error_page("TraceErrorPageRequest")]
fn get_request_schemas_ignore_unknown_fields(#[case] name: &str) {
let schemas = request_schemas();
let schema = serde_json::to_value(&schemas[name]).unwrap();
assert_ne!(schema["additionalProperties"], false);
}
#[rstest]
fn query_request_schema_rejects_unknown_fields() {
let schemas = request_schemas();
let schema = serde_json::to_value(&schemas["TraceQueryRequest"]).unwrap();
assert_eq!(schema["additionalProperties"], false);
}
#[rstest]
fn request_schemas_preserve_explicit_constraints() {
let schemas = request_schemas();
let detail = serde_json::to_value(&schemas["TraceDetailRequest"]).unwrap();
let list = serde_json::to_value(&schemas["TraceListRequest"]).unwrap();
let span = serde_json::to_value(&schemas["TraceSpanRequest"]).unwrap();
let error_page = serde_json::to_value(&schemas["TraceErrorPageRequest"]).unwrap();
let query = serde_json::to_value(&schemas["TraceQueryRequest"]).unwrap();
assert_eq!(detail["properties"]["page_size"]["minimum"], 1);
assert_eq!(detail["properties"]["page_size"]["maximum"], 500);
assert_eq!(list["properties"]["cursor"]["maxLength"], 512);
assert_eq!(detail["properties"]["cursor"]["maxLength"], 512);
assert_eq!(error_page["properties"]["cursor"]["maxLength"], 512);
assert!(list["properties"]["start_ms"].get("minimum").is_none());
assert!(list["properties"]["start_ms"].get("maximum").is_none());
assert!(list["properties"]["end_ms"].get("minimum").is_none());
assert!(list["properties"]["end_ms"].get("maximum").is_none());
assert_eq!(detail["properties"]["trace_ref"]["default"], "");
assert_eq!(span["properties"]["trace_ref"]["default"], "");
assert_eq!(error_page["properties"]["trace_ref"]["default"], "");
assert_eq!(query["required"], json!(["sql"]));
assert_eq!(
list["properties"]["start_ms"]["description"],
"Window start, unix ms. Default: 24h ago"
);
assert_eq!(
list["properties"]["end_ms"]["description"],
"Window end, unix ms. Default: now"
);
}
#[rstest]
fn request_models_deserialize_defaults_and_null_cursors() {
let list: TraceListRequest = serde_json::from_value(json!({})).unwrap();
let detail: TraceDetailRequest = serde_json::from_value(json!({})).unwrap();
let span: TraceSpanRequest = serde_json::from_value(json!({})).unwrap();
let error_page: TraceErrorPageRequest = serde_json::from_value(json!({})).unwrap();
assert!(list.start_ms.is_none());
assert!(list.end_ms.is_none());
assert!(list.cursor.is_none());
assert_eq!(detail.trace_ref, "");
assert!(detail.cursor.is_none());
assert!(detail.page_size.is_none());
assert_eq!(span.trace_ref, "");
assert_eq!(error_page.trace_ref, "");
assert!(error_page.cursor.is_none());
let null_cursor: TraceListRequest = serde_json::from_value(json!({"cursor": null})).unwrap();
assert!(null_cursor.cursor.is_none());
}
#[rstest]
fn get_request_models_ignore_unknown_fields_and_query_model_rejects_them() {
assert!(serde_json::from_value::<TraceListRequest>(json!({"unknown": true})).is_ok());
assert!(serde_json::from_value::<TraceDetailRequest>(json!({"unknown": true})).is_ok());
assert!(serde_json::from_value::<TraceSpanRequest>(json!({"unknown": true})).is_ok());
assert!(serde_json::from_value::<TraceErrorPageRequest>(json!({"unknown": true})).is_ok());
assert!(
serde_json::from_value::<TraceQueryRequest>(json!({"sql": "SELECT 1", "unknown": true}))
.is_err()
);
assert!(serde_json::from_value::<TraceQueryRequest>(json!({})).is_err());
}

View file

@ -4,7 +4,6 @@ Base configuration for A2A protocol providers.
from abc import ABC, abstractmethod
from collections.abc import AsyncIterator
from typing import Any
class BaseA2AProviderConfig(ABC):
@ -19,10 +18,10 @@ class BaseA2AProviderConfig(ABC):
async def handle_non_streaming(
self,
request_id: str,
params: dict[str, Any],
params: dict[str, object],
api_base: str | None = None,
**kwargs,
) -> dict[str, Any]:
) -> dict[str, object]:
"""
Handle non-streaming A2A request.
@ -40,10 +39,10 @@ class BaseA2AProviderConfig(ABC):
async def handle_streaming(
self,
request_id: str,
params: dict[str, Any],
params: dict[str, object],
api_base: str | None = None,
**kwargs,
) -> AsyncIterator[dict[str, Any]]:
) -> AsyncIterator[dict[str, object]]:
"""
Handle streaming A2A request.

View file

@ -3,7 +3,7 @@ Bedrock AgentCore A2A provider configuration.
"""
from collections.abc import AsyncIterator
from typing import Any, Final
from typing import Final
from litellm.a2a_protocol.providers.base import BaseA2AProviderConfig
from litellm.a2a_protocol.providers.bedrock_agentcore.handler import (
@ -23,10 +23,10 @@ class BedrockAgentCoreA2AConfig(BaseA2AProviderConfig):
async def handle_non_streaming(
self,
request_id: str,
params: dict[str, Any],
params: dict[str, object],
api_base: str | None = None,
**kwargs,
) -> dict[str, Any]:
) -> dict[str, object]:
"""Handle non-streaming request to AgentCore A2A agent."""
litellm_params: Final = kwargs.get("litellm_params")
if not litellm_params:
@ -43,10 +43,10 @@ class BedrockAgentCoreA2AConfig(BaseA2AProviderConfig):
async def handle_streaming(
self,
request_id: str,
params: dict[str, Any],
params: dict[str, object],
api_base: str | None = None,
**kwargs,
) -> AsyncIterator[dict[str, Any]]:
) -> AsyncIterator[dict[str, object]]:
"""Handle streaming request to AgentCore A2A agent."""
litellm_params: Final = kwargs.get("litellm_params")
if not litellm_params:

View file

@ -7,7 +7,7 @@ completion bridge that would otherwise strip the envelope.
import json
from collections.abc import AsyncIterator, Mapping
from typing import Any, Final
from typing import Final
from litellm._logging import verbose_logger
from litellm.a2a_protocol.providers.bedrock_agentcore.transformation import (
@ -30,9 +30,9 @@ class BedrockAgentCoreA2AHandler:
async def handle_non_streaming(
request_id: str,
params: Mapping[str, object],
litellm_params: dict[str, Any],
litellm_params: dict[str, object],
agent_extra_headers: dict[str, str] | None = None,
) -> dict[str, Any]:
) -> dict[str, object]:
"""
Handle non-streaming A2A request to AgentCore.
@ -77,9 +77,9 @@ class BedrockAgentCoreA2AHandler:
async def handle_streaming(
request_id: str,
params: Mapping[str, object],
litellm_params: dict[str, Any],
litellm_params: dict[str, object],
agent_extra_headers: dict[str, str] | None = None,
) -> AsyncIterator[dict[str, Any]]:
) -> AsyncIterator[dict[str, object]]:
"""
Handle streaming A2A request to AgentCore.

View file

@ -219,7 +219,7 @@ class BedrockAgentCoreA2ATransformation:
return url, signed_headers, signed_body
@staticmethod
async def parse_sse_events(response: _SSELineSource) -> AsyncIterator[dict[str, Any]]:
async def parse_sse_events(response: _SSELineSource) -> AsyncIterator[dict[str, object]]:
"""
Parse SSE events from an httpx streaming response.

View file

@ -1,5 +1,4 @@
from collections.abc import AsyncIterator
from typing import Any
from litellm.a2a_protocol.litellm_completion_bridge.handler import (
A2A_USER_API_KEY_HASH_PARAM,
@ -16,10 +15,10 @@ class LangFlowA2AConfig(BaseA2AProviderConfig):
async def handle_non_streaming(
self,
request_id: str,
params: dict[str, Any],
params: dict[str, object],
api_base: str | None = None,
**kwargs,
) -> dict[str, Any]:
) -> dict[str, object]:
litellm_params = kwargs.get("litellm_params")
if not litellm_params:
raise ValueError(
@ -39,10 +38,10 @@ class LangFlowA2AConfig(BaseA2AProviderConfig):
async def handle_streaming(
self,
request_id: str,
params: dict[str, Any],
params: dict[str, object],
api_base: str | None = None,
**kwargs,
) -> AsyncIterator[dict[str, Any]]:
) -> AsyncIterator[dict[str, object]]:
litellm_params = kwargs.get("litellm_params")
if not litellm_params:
raise ValueError(

View file

@ -2,8 +2,7 @@
Pydantic AI provider configuration.
"""
from collections.abc import AsyncIterator
from typing import Any
from collections.abc import AsyncIterator, Mapping
from litellm.a2a_protocol.providers.base import BaseA2AProviderConfig
from litellm.a2a_protocol.providers.pydantic_ai_agents.handler import PydanticAIHandler
@ -20,9 +19,12 @@ class PydanticAIProviderConfig(BaseA2AProviderConfig):
async def handle_non_streaming(
self,
request_id: str,
params: dict[str, Any],
params: dict[str, object],
api_base: str | None = None,
**kwargs: Any,
*,
timeout: float = 60.0,
agent_extra_headers: Mapping[str, str] | None = None,
**kwargs: object,
) -> dict[str, object]:
"""Handle non-streaming request to Pydantic AI agent."""
if api_base is None:
@ -31,14 +33,14 @@ class PydanticAIProviderConfig(BaseA2AProviderConfig):
request_id=request_id,
params=params,
api_base=api_base,
timeout=kwargs.get("timeout", 60.0),
agent_extra_headers=kwargs.get("agent_extra_headers"),
timeout=timeout,
agent_extra_headers=agent_extra_headers,
)
async def handle_streaming(
self,
request_id: str,
params: dict[str, Any],
params: dict[str, object],
api_base: str | None = None,
**kwargs,
) -> AsyncIterator[dict[str, object]]:

View file

@ -5,8 +5,8 @@ Pydantic AI agents follow A2A protocol but don't support streaming natively.
This handler provides fake streaming by converting non-streaming responses into streaming chunks.
"""
from collections.abc import AsyncIterator
from typing import Any, Final
from collections.abc import AsyncIterator, Mapping
from typing import Final
from litellm._logging import verbose_logger
from litellm.a2a_protocol.providers.pydantic_ai_agents.transformation import (
@ -26,11 +26,11 @@ class PydanticAIHandler:
@staticmethod
async def handle_non_streaming(
request_id: str,
params: dict[str, Any],
params: Mapping[str, object],
api_base: str | None = None,
timeout: float = 60.0,
agent_extra_headers: dict[str, str] | None = None,
) -> dict[str, Any]:
agent_extra_headers: Mapping[str, str] | None = None,
) -> dict[str, object]:
"""
Handle non-streaming request to Pydantic AI agent.
@ -63,13 +63,13 @@ class PydanticAIHandler:
@staticmethod
async def handle_streaming(
request_id: str,
params: dict[str, Any],
params: Mapping[str, object],
api_base: str | None = None,
timeout: float = 60.0,
chunk_size: int = 50,
delay_ms: int = 10,
agent_extra_headers: dict[str, str] | None = None,
) -> AsyncIterator[dict[str, Any]]:
agent_extra_headers: Mapping[str, str] | None = None,
) -> AsyncIterator[dict[str, object]]:
"""
Handle streaming request to Pydantic AI agent with fake streaming.

View file

@ -7,7 +7,7 @@ This module provides fake streaming by converting non-streaming responses into s
import asyncio
from collections.abc import AsyncIterator, Mapping, Sequence
from typing import Any, Final, Protocol, cast, runtime_checkable
from typing import Final, Protocol, runtime_checkable
from uuid import uuid4
from pydantic import TypeAdapter
@ -100,7 +100,7 @@ class PydanticAITransformation:
request_id: str,
max_attempts: int = 30,
poll_interval: float = 0.5,
agent_extra_headers: dict[str, str] | None = None,
agent_extra_headers: Mapping[str, str] | None = None,
) -> dict[str, object]:
"""
Poll for task completion using tasks/get method.
@ -156,7 +156,7 @@ class PydanticAITransformation:
request_id: str,
params: "_SupportsModelDump | _SupportsPydanticDict | Mapping[str, object]",
timeout: float = 60.0,
agent_extra_headers: dict[str, str] | None = None,
agent_extra_headers: Mapping[str, str] | None = None,
) -> dict[str, object]:
"""
Send a request to Pydantic AI agent and return the raw task response.
@ -200,7 +200,7 @@ class PydanticAITransformation:
# Send request to Pydantic AI agent using shared async HTTP client
client: Final = get_async_httpx_client(
llm_provider=cast(Any, "pydantic_ai_agent"),
llm_provider="pydantic_ai_agent",
params={"timeout": timeout},
)
response: Final = await client.post(
@ -242,7 +242,7 @@ class PydanticAITransformation:
request_id: str,
params: "_SupportsModelDump | _SupportsPydanticDict | Mapping[str, object]",
timeout: float = 60.0,
agent_extra_headers: dict[str, str] | None = None,
agent_extra_headers: Mapping[str, str] | None = None,
) -> dict[str, object]:
"""
Send a non-streaming A2A request to Pydantic AI agent and wait for completion.
@ -278,7 +278,7 @@ class PydanticAITransformation:
request_id: str,
params: "_SupportsModelDump | _SupportsPydanticDict | Mapping[str, object]",
timeout: float = 60.0,
agent_extra_headers: dict[str, str] | None = None,
agent_extra_headers: Mapping[str, str] | None = None,
) -> dict[str, object]:
"""
Send a request to Pydantic AI agent and return the raw task response.

View file

@ -3,11 +3,11 @@ A2A provider configuration for IBM watsonx Orchestrate (WXO).
"""
from collections.abc import AsyncIterator
from typing import Any, Final
from litellm.a2a_protocol.providers.base import BaseA2AProviderConfig
from litellm.a2a_protocol.providers.watsonx_orchestrate.handler import (
WatsonxOrchestrateHandler,
WXOLitellmParams,
)
@ -17,12 +17,13 @@ class WatsonxOrchestrateA2AConfig(BaseA2AProviderConfig):
async def handle_non_streaming(
self,
request_id: str,
params: dict[str, Any],
params: dict[str, object],
api_base: str | None = None,
**kwargs: Any,
*,
litellm_params: WXOLitellmParams | None = None,
**kwargs: object,
) -> dict[str, object]:
"""Handle a non-streaming A2A request via WXO runs API."""
litellm_params: Final = kwargs.get("litellm_params")
if not litellm_params:
raise ValueError(
"litellm_params is required for WatsonxOrchestrateA2AConfig "
@ -37,12 +38,13 @@ class WatsonxOrchestrateA2AConfig(BaseA2AProviderConfig):
async def handle_streaming(
self,
request_id: str,
params: dict[str, Any],
params: dict[str, object],
api_base: str | None = None,
**kwargs: Any,
*,
litellm_params: WXOLitellmParams | None = None,
**kwargs: object,
) -> AsyncIterator[dict[str, object]]:
"""Handle a streaming A2A request via WXO streaming runs API."""
litellm_params: Final = kwargs.get("litellm_params")
if not litellm_params:
raise ValueError(
"litellm_params is required for WatsonxOrchestrateA2AConfig "

View file

@ -7,7 +7,7 @@ import hashlib
import json
import time
from collections.abc import AsyncIterator
from typing import Any, Final, NamedTuple, Protocol
from typing import Final, NamedTuple, Protocol
import httpx
from typing_extensions import NotRequired, ReadOnly, TypedDict
@ -228,7 +228,7 @@ class WatsonxOrchestrateHandler:
return run_data
@staticmethod
async def _accumulate_wxo_sse_text(response: Any) -> str:
async def _accumulate_wxo_sse_text(response: _SSELineSource) -> str:
source: Final[_WXOView] = {"sse_source": response}
accumulated_text = ""
async for line in source["sse_source"].aiter_lines():

View file

@ -12,7 +12,7 @@ This module is dependency-injected: callers pass the proxy ``llm_router`` and
from __future__ import annotations
from collections.abc import Sequence
from collections.abc import Mapping, Sequence
from typing import TYPE_CHECKING, Any, Final
import litellm
@ -39,10 +39,10 @@ def resolve_embedding_router(
def build_router_embedding_metadata(
request_metadata: dict[str, Any] | None,
) -> dict[str, Any]:
request_metadata: Mapping[str, object] | None,
) -> Mapping[str, object]:
"""Forward the caller's full metadata, flagged as a semantic-cache embedding."""
metadata: Final[dict[str, Any]] = dict(request_metadata or {})
metadata: Final = dict(request_metadata or {})
metadata["semantic-cache-embedding"] = True
return metadata

View file

@ -446,11 +446,19 @@ class Cache:
2. Else if a model_group is set, then return the model_group as the model. This is used for all requests sent through the litellm.Router()
3. Else use the `model` passed in kwargs
"""
metadata: Final[dict] = kwargs.get("metadata", {}) or {}
litellm_params: Final[dict] = kwargs.get("litellm_params", {}) or {}
metadata_in_litellm_params: Final[dict] = litellm_params.get("metadata", {}) or {}
model_group: Final[str | None] = metadata.get("model_group") or metadata_in_litellm_params.get("model_group")
caching_group: Final = self._get_caching_group(metadata, model_group)
metadata_sources: Final[tuple[dict, ...]] = (
kwargs.get("metadata") or {},
kwargs.get("litellm_metadata") or {},
litellm_params.get("metadata") or {},
litellm_params.get("litellm_metadata") or {},
)
model_group: Final[str | None] = next(
(source["model_group"] for source in metadata_sources if source.get("model_group")), None
)
caching_group: Final = next(
(group for source in metadata_sources if (group := self._get_caching_group(source, model_group))), None
)
return caching_group or model_group or kwargs["model"]
def _get_caching_group(self, metadata: dict, model_group: str | None) -> str | None:

View file

@ -331,7 +331,7 @@ async def afile_retrieve(
else:
response = init_response
return OpenAIFileObject(**response.model_dump())
return OpenAIFileObject.model_validate(response.model_dump())
except Exception as e:
raise e

View file

@ -34,7 +34,7 @@ class FileContentStreamingResponse:
self.custom_llm_provider = custom_llm_provider
self.logging_obj = logging_obj
self.standard_logging_object: StandardLoggingPayload | None = None
self._hidden_params: dict[str, Any] = {}
self._hidden_params: dict[str, object] = {}
self._logging_completed = False
self._close_completed = False
self._start_time = (
@ -121,7 +121,7 @@ class FileContentStreamingResponse:
return response
def _sync_hidden_params(self) -> None:
litellm_params: dict[str, Any] = {}
litellm_params: dict[str, object] = {}
if self.logging_obj is not None:
litellm_params = self.logging_obj.model_call_details.get("litellm_params", {}) or {}

View file

@ -4,7 +4,7 @@ from __future__ import annotations
from collections.abc import Awaitable, Callable, Mapping, Sequence
from dataclasses import dataclass, field
from typing import TYPE_CHECKING, Any, TypeAlias
from typing import TYPE_CHECKING, TypeAlias
from pydantic import BaseModel
@ -42,7 +42,7 @@ class SessionContext:
api_base: str | None = None
endpoint: ModelEndpoint | None = None
instructions: str | None = None
tools: Sequence[Callable[..., Any]] = ()
tools: Sequence[Callable[..., object]] = ()
skills: Sequence[str] = ()
disable_tools: Sequence[str] = ()
permissions: PermissionMode = "full"
@ -50,7 +50,7 @@ class SessionContext:
output: type[BaseModel] | None = None
max_turns: int | None = None
timeout: float | None = None
metadata: Mapping[str, Any] = field(default_factory=dict)
metadata: Mapping[str, object] = field(default_factory=dict)
options: HarnessOptions | None = None
# Set by the handler after each turn.
final_text: str = ""

View file

@ -13,7 +13,7 @@ import asyncio
import os
from collections import deque
from collections.abc import AsyncIterator, Sequence
from typing import Any, Final
from typing import Final
from litellm._logging import verbose_logger
from litellm.constants import HARNESS_STDERR_TAIL_LINES, HARNESS_STREAM_READ_CHUNK_BYTES
@ -125,7 +125,7 @@ class CLIHarnessHandler(BaseHarnessHandler):
self._proc = proc
tail: Final[deque[str]] = deque(maxlen=HARNESS_STDERR_TAIL_LINES) # mutable-ok: bounded stderr ring buffer
stderr_task = asyncio.ensure_future(drain_stderr(proc.stderr, tail))
state: Any = self.config.create_stream_state()
state: Final[object] = self.config.create_stream_state()
exit_code: int | None = None
try:
await send_stdin(proc, request.stdin)

View file

@ -9,7 +9,7 @@ from typing import Any, Literal
@dataclass(frozen=True)
class ClaudeCodeOptions:
config: Mapping[str, Any] = field(default_factory=dict)
config: Mapping[str, object] = field(default_factory=dict)
env: Mapping[str, str] = field(default_factory=dict)
@ -17,20 +17,20 @@ class ClaudeCodeOptions:
class CodexOptions:
reasoning_effort: Literal["low", "medium", "high", "xhigh"] | None = None
web_search: bool = False
config: Mapping[str, Any] = field(default_factory=dict)
config: Mapping[str, object] = field(default_factory=dict)
env: Mapping[str, str] = field(default_factory=dict)
@dataclass(frozen=True)
class OpenCodeOptions:
agent: str = "build"
config: Mapping[str, Any] = field(default_factory=dict)
config: Mapping[str, object] = field(default_factory=dict)
env: Mapping[str, str] = field(default_factory=dict)
@dataclass(frozen=True)
class DeepAgentsOptions:
subagents: Sequence[Any] = ()
subagents: Sequence[object] = ()
recursion_limit: int | None = None

View file

@ -1005,24 +1005,24 @@ def aagent(
`await litellm.aagent(...)` returns a Result. With stream=True it returns an async
iterator of events instead: `async for event in litellm.aagent(..., stream=True)`.
"""
kwargs: dict[str, Any] = { # mutable-ok: forwarded as **kwargs to arun_agent/astream_agent
"sandbox": sandbox,
"model": model,
"api_key": api_key,
"api_base": api_base,
"instructions": instructions,
"tools": tools,
"skills": skills,
"disable_tools": disable_tools,
"permissions": permissions,
"on_approval": on_approval,
"output": output,
"max_turns": max_turns,
"timeout": timeout,
"metadata": metadata,
"options": options,
"install": install,
}
if stream:
return astream_agent(harness, prompt, **kwargs)
return arun_agent(harness, prompt, **kwargs)
call: Final = astream_agent if stream else arun_agent
return call(
harness,
prompt,
sandbox=sandbox,
model=model,
api_key=api_key,
api_base=api_base,
instructions=instructions,
tools=tools,
skills=skills,
disable_tools=disable_tools,
permissions=permissions,
on_approval=on_approval,
output=output,
max_turns=max_turns,
timeout=timeout,
metadata=metadata,
options=options,
install=install,
)

View file

@ -8,7 +8,7 @@ import threading
from collections.abc import AsyncIterator, Callable, Coroutine, Mapping, Sequence
from concurrent.futures import Future
from typing import (
Any,
Final,
TypeVar,
)
@ -417,24 +417,24 @@ def agent(
Prefix the model with `litellm_proxy/` to route every model call through your
LiteLLM AI Gateway.
"""
kwargs: dict[str, Any] = { # mutable-ok: forwarded as **kwargs to _run/_stream
"sandbox": sandbox,
"model": model,
"api_key": api_key,
"api_base": api_base,
"instructions": instructions,
"tools": tools,
"skills": skills,
"disable_tools": disable_tools,
"permissions": permissions,
"on_approval": on_approval,
"output": output,
"max_turns": max_turns,
"timeout": timeout,
"metadata": metadata,
"options": options,
"install": install,
}
if stream:
return _stream(harness, prompt, **kwargs)
return _run(harness, prompt, **kwargs)
call: Final = _stream if stream else _run
return call(
harness,
prompt,
sandbox=sandbox,
model=model,
api_key=api_key,
api_base=api_base,
instructions=instructions,
tools=tools,
skills=skills,
disable_tools=disable_tools,
permissions=permissions,
on_approval=on_approval,
output=output,
max_turns=max_turns,
timeout=timeout,
metadata=metadata,
options=options,
install=install,
)

View file

@ -7,7 +7,7 @@ import json
from collections.abc import Mapping
from dataclasses import dataclass, field
from enum import Enum
from typing import Any, Literal
from typing import Literal
from pydantic import BaseModel
@ -71,7 +71,7 @@ class ToolCall:
id: str
name: str
native_name: str
input: Mapping[str, Any]
input: Mapping[str, object]
builtin: bool = True
@ -100,7 +100,7 @@ class Approval:
"""A request to run a tool. The turn waits until allow() or deny() is called."""
tool: str
input: Mapping[str, Any]
input: Mapping[str, object]
_decision: asyncio.Future[tuple[bool, str]] = field(
default_factory=lambda: asyncio.get_event_loop().create_future(),
compare=False,

View file

@ -2,6 +2,7 @@ import json
from collections.abc import Mapping
from typing import TYPE_CHECKING, Any, Final
from pydantic import ConfigDict, TypeAdapter
from typing_extensions import ReadOnly, TypedDict, override
from litellm._logging import verbose_logger
@ -29,6 +30,8 @@ from litellm.integrations._types.open_inference import (
ToolCallAttributes,
)
_JSON_OBJECT: Final = TypeAdapter(dict[str, object], config=ConfigDict(hide_input_in_errors=True))
class ArizeOTELAttributes(BaseLLMObsOTELAttributes):
@staticmethod
@ -609,9 +612,7 @@ def _coerce_response_obj_for_attrs(response_obj):
text: Final = getattr(response_obj, "text", None)
if isinstance(text, str) and text:
try:
parsed: Final = json.loads(text)
if isinstance(parsed, dict):
return parsed
return _JSON_OBJECT.validate_python(json.loads(text))
except Exception:
pass
return response_obj
@ -1062,9 +1063,7 @@ def _parse_passthrough_response(raw_response_obj, coerced_response_obj, kwargs):
inner = candidate.get("response")
if isinstance(inner, str):
try:
parsed = json.loads(inner)
if isinstance(parsed, dict):
return parsed
return _JSON_OBJECT.validate_python(json.loads(inner))
except Exception:
continue
if isinstance(inner, dict):
@ -1078,9 +1077,7 @@ def _parse_passthrough_response(raw_response_obj, coerced_response_obj, kwargs):
return original
if isinstance(original, str):
try:
parsed = json.loads(original)
if isinstance(parsed, dict):
return parsed
return _JSON_OBJECT.validate_python(json.loads(original))
except Exception:
return None
return None

View file

@ -15,6 +15,7 @@ from types import MappingProxyType
from typing import Any, Final, Literal
import httpx
from pydantic import ConfigDict, TypeAdapter, ValidationError
import litellm
from litellm._logging import verbose_logger
@ -57,6 +58,7 @@ from litellm.types.utils import (
_EMPTY_MAPPING: Final[Mapping[str, object]] = MappingProxyType({})
_EMPTY_MESSAGE: Final[Message] = {"role": "", "content": ""}
_MAX_PARSED_TOOL_ARGUMENT_CHARS: Final = 256 * 1024
_JSON_OBJECT: Final = TypeAdapter(dict[str, object], config=ConfigDict(hide_input_in_errors=True))
_SAFE_REDACTED_MESSAGE_ROLES: Final = frozenset(
{"agent", "assistant", "developer", "function", "model", "system", "tool", "user"}
)
@ -282,8 +284,10 @@ def _to_dd_arguments(raw_arguments: object) -> dict[str, object] | str:
return raw_arguments if isinstance(raw_arguments, dict) else str(raw_arguments)
if len(raw_arguments) > _MAX_PARSED_TOOL_ARGUMENT_CHARS:
return raw_arguments
parsed: Final = safe_json_loads(raw_arguments)
return parsed if isinstance(parsed, dict) else raw_arguments
try:
return _JSON_OBJECT.validate_python(safe_json_loads(raw_arguments))
except ValidationError:
return raw_arguments
def _to_dd_tool_calls(message: Mapping[str, object]) -> tuple[ToolCall, ...]:

View file

@ -4,12 +4,12 @@ import json
import os
import re
import uuid
from collections.abc import Mapping, Sequence
from collections.abc import Callable, Mapping, Sequence
from datetime import datetime, timezone, tzinfo
from typing import Any, Final, Protocol, cast
import httpx
from pydantic import BaseModel, Field
from pydantic import BaseModel, ConfigDict, Field, TypeAdapter
from typing_extensions import ReadOnly, TypedDict
import litellm
@ -35,6 +35,11 @@ GALILEO_CLOUD_API_BASE_URL: Final = "https://api.galileo.ai"
# unavailable, invalid credentials) cannot leak memory unboundedly.
GALILEO_MAX_IN_MEMORY_RECORDS: Final = 1000
_UNTYPED_VALUE: Final = TypeAdapter(object)
_ZERO_ARGUMENT_CALLABLE: Final[TypeAdapter[Callable[[], object]]] = TypeAdapter(
Callable[[], object], config=ConfigDict(hide_input_in_errors=True)
)
class _GalileoLoginBody(TypedDict):
"""Decoded body of the Galileo login response."""
@ -473,9 +478,9 @@ class GalileoObserve(CustomLogger):
if isinstance(value, str):
return value
def _json_default(obj: Any) -> object:
def _json_default(obj: object) -> object:
if hasattr(obj, "model_dump"):
return obj.model_dump()
return _ZERO_ARGUMENT_CALLABLE.validate_python(getattr(obj, "model_dump", None))()
return str(obj)
return json.dumps(value, default=_json_default)
@ -496,7 +501,7 @@ class GalileoObserve(CustomLogger):
if hasattr(message, "json"):
message_json: Final[object] = message.json()
if isinstance(message_json, str):
return json.loads(message_json)
return _UNTYPED_VALUE.validate_python(json.loads(message_json))
return message_json
return message
return None

View file

@ -8,6 +8,8 @@ from datetime import datetime
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, TypedDict, cast
from pydantic import ConfigDict, TypeAdapter
import litellm
from litellm._logging import verbose_logger
from litellm.integrations._types.open_inference import (
@ -95,6 +97,8 @@ class _ResponseWithUsageView(TypedDict, total=False):
usage: "_UsageCompletionTokensView | None"
_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True))
# Cap on credential-scoped providers held at once; each one owns an exporter thread.
_MAX_DYNAMIC_TRACER_PROVIDERS: Final = 256
@ -1633,7 +1637,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
attributes = self.config.attributes
if attributes is None and self.callback_name in (None, "otel"):
otel_settings: Final = (litellm.callback_settings or {}).get("otel") or {}
raw: Final = otel_settings.get("attributes") if isinstance(otel_settings, dict) else None
raw: Final[object] = otel_settings.get("attributes") if isinstance(otel_settings, dict) else None
if raw is not None:
attributes = _build_metric_attribute_filter(raw)
(
@ -2919,7 +2923,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
import json
try:
_parsed: Final[Mapping[str, object]] = json.loads(_raw_response)
_parsed: Final = _JSON_OBJECT.validate_python(json.loads(_raw_response))
for param, val in _parsed.items():
self.safe_set_attribute(
span=span,

View file

@ -1,12 +1,17 @@
"""Public API for Opik payload building."""
from collections.abc import Mapping
from datetime import datetime
from typing import Any, Final
from pydantic import ConfigDict, TypeAdapter
from litellm.integrations.opik import utils
from . import extractors, payload_builders, types
_STANDARD_LOGGING_FIELDS: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True))
def build_opik_payload(
kwargs: dict[str, Any],
@ -36,12 +41,14 @@ def build_opik_payload(
- First element is TracePayload if creating a new trace, None if attaching to existing
- Second element is always SpanPayload
"""
standard_logging_object: Final = kwargs["standard_logging_object"]
standard_logging_object: Final = _STANDARD_LOGGING_FIELDS.validate_python(kwargs["standard_logging_object"])
# Extract litellm params and metadata
litellm_params: Final = kwargs.get("litellm_params", {}) or {}
litellm_metadata: Final = litellm_params.get("metadata", {}) or {}
standard_logging_metadata: Final = standard_logging_object.get("metadata", {}) or {}
standard_logging_metadata: Final = _STANDARD_LOGGING_FIELDS.validate_python(
standard_logging_object.get("metadata", {}) or {}
)
# Extract and merge Opik metadata
opik_metadata: Final = extractors.extract_opik_metadata(litellm_metadata, standard_logging_metadata)

View file

@ -4,8 +4,12 @@ import json
from collections.abc import Mapping
from typing import Any, Final
from pydantic import TypeAdapter
from litellm import _logging
_DECODED_JSON: Final = TypeAdapter(object)
def normalize_provider_name(provider: str | None) -> str | None:
"""
@ -149,7 +153,7 @@ def apply_proxy_header_overrides(
thread_id = value
elif param_key == "tags":
try:
parsed_tags: object = json.loads(value)
parsed_tags = _DECODED_JSON.validate_python(json.loads(value))
if isinstance(parsed_tags, list):
tags.extend(parsed_tags)
except (json.JSONDecodeError, TypeError):

View file

@ -1,7 +1,8 @@
"""Type definitions for Opik payload building."""
from collections.abc import Mapping
from dataclasses import dataclass
from typing import Any, Final, Literal
from typing import Final, Literal
@dataclass
@ -13,9 +14,9 @@ class TracePayload:
name: str
start_time: str
end_time: str
input: Any
output: Any
metadata: dict[str, Any]
input: object
output: object
metadata: Mapping[str, object]
tags: list[str]
thread_id: str | None = None
@ -32,9 +33,9 @@ class SpanPayload:
model: str
start_time: str
end_time: str
input: Any
output: Any
metadata: dict[str, Any]
input: object
output: object
metadata: Mapping[str, object]
tags: list[str]
usage: dict[str, int]
parent_span_id: str | None = None

View file

@ -2,7 +2,8 @@ import configparser
import os
import time
import uuid
from typing import Any, Final
from collections.abc import Mapping, Sequence
from typing import Final
CONFIG_FILE_PATH_DEFAULT: Final[str] = "~/.opik.config"
@ -93,14 +94,14 @@ def create_usage_object(usage):
return usage_dict
def _remove_nulls(x: dict[str, Any]) -> dict[str, Any]:
def _remove_nulls(x: Mapping[str, object]) -> dict[str, object]:
"""Remove None values from dict."""
return {k: v for k, v in x.items() if v is not None}
def get_traces_and_spans_from_payload(
payload: list[dict[str, Any]],
) -> tuple[list[dict[str, Any]], list[dict[str, Any]]]:
payload: Sequence[Mapping[str, object]],
) -> tuple[list[dict[str, object]], list[dict[str, object]]]:
"""
Separate traces and spans from payload.

View file

@ -2,7 +2,7 @@
from enum import Enum
from functools import lru_cache
from typing import Annotated, Any, Final
from typing import Annotated, Final
from pydantic import AliasChoices, BaseModel, Field, TypeAdapter, ValidationError, field_validator, model_validator
from pydantic.fields import FieldInfo
@ -315,7 +315,7 @@ class OpenTelemetryV2Config(BaseSettings):
mode="before",
)
@classmethod
def _split_csv(cls, value: Any) -> Any:
def _split_csv(cls, value: object) -> object:
"""Accept a comma-separated string for list fields.
Env vars are strings, but these fields are lists. Pydantic-settings would

View file

@ -307,7 +307,7 @@ class GenAIMetricRecorder:
return common_attrs
def _bounded_attributes(self, kwargs: Mapping[str, Any]) -> MetricAttributes:
def _bounded_attributes(self, kwargs: Mapping[str, object]) -> MetricAttributes:
"""The datapoint attributes, capped at :data:`METRIC_ATTRIBUTE_CEILING`.
The cap runs BEFORE the operator's include/exclude filter so the filter can
@ -322,7 +322,7 @@ class GenAIMetricRecorder:
attributes = None
if self._callback_name in (None, "otel"):
otel_settings: Final = (litellm.callback_settings or {}).get("otel") or {}
raw: Final = otel_settings.get("attributes") if isinstance(otel_settings, dict) else None
raw: Final[object] = otel_settings.get("attributes") if isinstance(otel_settings, dict) else None
if raw is not None:
attributes = _build_metric_attribute_filter(raw)
# A bad filter (include_list + exclude_list both set, an unfilterable name)

View file

@ -595,7 +595,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
and not (self.s3_drop_on_terminal_error and _is_terminal(response))
and attempt < max_retries - 1
):
wait_time = 2**attempt # 1s, 2s
wait_time = 1 << attempt # 1s, 2s
verbose_logger.log(
logging.DEBUG if _in_flush.get() else logging.WARNING,
"S3 upload returned %s, retrying in %ss (attempt %s/%s) key=%s",
@ -897,7 +897,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
and not (self.s3_drop_on_terminal_error and _is_terminal(response))
and attempt < max_retries - 1
):
wait_time = 2**attempt # 1s, 2s
wait_time = 1 << attempt # 1s, 2s
verbose_logger.warning(
"S3 upload returned %s, retrying in %ss (attempt %s/%s) key=%s",
response.status_code,

View file

@ -6,12 +6,12 @@ It searches the vector store for relevant context, runs the request's pre-call g
over that context, and appends it to the messages.
"""
from collections.abc import Awaitable, Callable, Mapping, Sequence
from collections.abc import Awaitable, Callable, Iterable, Mapping, Sequence
from dataclasses import dataclass
from itertools import chain
from typing import TYPE_CHECKING, Any, Final, Protocol, cast, get_args
from pydantic import TypeAdapter, ValidationError
from pydantic import ConfigDict, TypeAdapter, ValidationError
from typing_extensions import assert_never
import litellm
@ -44,10 +44,12 @@ else:
LiteLLMLoggingObj = Any
SEARCH_FAILURES_FIELD: Final = "vector_store_search_failures"
_PROVIDER_FIELDS_ATTRIBUTE: Final = "provider_specific_fields"
_DEFAULT_FAILURE_MODE: Final[VectorStoreSearchFailureMode] = "annotate"
_FAILURE_MODE_ADAPTER: Final = TypeAdapter(VectorStoreSearchFailureMode)
_OBJECT_ADAPTER: Final = TypeAdapter(object)
_STR_KEYED_ADAPTER: Final = TypeAdapter(dict[str, object])
_ITERABLE_ADAPTER: Final = TypeAdapter(Iterable[object], config=ConfigDict(hide_input_in_errors=True))
_GUARDRAIL_KEYS_THE_PROXY_MERGES_INTO_METADATA: Final = frozenset(
{"guardrails", "guardrail_config", "policies", "include_guardrail_response"}
)
@ -475,7 +477,7 @@ class VectorStorePreCallHook(CustomLogger):
async def async_post_call_streaming_deployment_hook(
self,
request_data: dict,
response_chunk: Any,
response_chunk: object,
call_type: CallTypes | None,
) -> object | None:
"""
@ -496,15 +498,17 @@ class VectorStorePreCallHook(CustomLogger):
return response_chunk
# Add search results to streaming chunk
if hasattr(response_chunk, "choices") and response_chunk.choices:
for choice in response_chunk.choices:
if hasattr(choice, "delta") and choice.delta:
provider_fields = getattr(choice.delta, "provider_specific_fields", None) or {}
choices: Final[object] = getattr(response_chunk, "choices", None)
if choices:
for choice in _ITERABLE_ADAPTER.validate_python(choices):
delta: object = getattr(choice, "delta", None)
if delta:
provider_fields = getattr(delta, _PROVIDER_FIELDS_ATTRIBUTE, None) or {}
if search_results:
provider_fields["search_results"] = search_results
if search_failures:
provider_fields[SEARCH_FAILURES_FIELD] = search_failures
choice.delta.provider_specific_fields = provider_fields
setattr(delta, _PROVIDER_FIELDS_ATTRIBUTE, provider_fields)
# Return modified chunk
return response_chunk

View file

@ -3,8 +3,8 @@ A2A Protocol Transformation for LiteLLM
"""
import uuid
from collections.abc import Iterator, Mapping
from typing import TYPE_CHECKING, Any, Final
from collections.abc import AsyncIterator, Iterator, Mapping
from typing import TYPE_CHECKING, Final
import httpx
@ -72,9 +72,9 @@ class A2AConfig(BaseConfig):
agent_name: str,
api_base: str | None,
api_key: str | None,
headers: dict[str, Any] | None,
optional_params: dict[str, Any],
) -> tuple[str | None, str | None, dict[str, Any] | None]:
headers: dict[str, object] | None,
optional_params: dict[str, object],
) -> tuple[str | None, str | None, dict[str, object] | None]:
"""
Resolve agent configuration from the registry for a registered agent.
@ -376,7 +376,7 @@ class A2AConfig(BaseConfig):
def get_model_response_iterator(
self,
streaming_response: Iterator | Any,
streaming_response: Iterator[str] | AsyncIterator[str] | ModelResponse,
sync_stream: bool,
json_mode: bool | None = False,
) -> BaseModelResponseIterator:
@ -397,7 +397,7 @@ class A2AConfig(BaseConfig):
json_mode=json_mode,
)
def _openai_message_to_a2a_message(self, message: dict[str, Any]) -> dict[str, Any]:
def _openai_message_to_a2a_message(self, message: Mapping[str, object]) -> dict[str, object]:
"""
Convert OpenAI message to A2A message format.

View file

@ -7,9 +7,11 @@ https://github.com/BerriAI/litellm/issues/6592
New config to ensure we introduce this without causing breaking changes for users
"""
from collections.abc import Iterable, Mapping
from typing import TYPE_CHECKING, Any, Final
from aiohttp import ClientResponse
from pydantic import ConfigDict, TypeAdapter
from litellm.llms.openai_like.chat.transformation import OpenAILikeChatConfig
from litellm.types.llms.openai import AllMessageValues
@ -23,6 +25,10 @@ if TYPE_CHECKING:
else:
LiteLLMLoggingObj = Any
_JSON_OBJECTS: Final = TypeAdapter(
Iterable[Mapping[str, object]], config=ConfigDict(strict=True, hide_input_in_errors=True)
)
class AiohttpOpenAIChatConfig(OpenAILikeChatConfig):
def get_complete_url(
@ -73,7 +79,9 @@ class AiohttpOpenAIChatConfig(OpenAILikeChatConfig):
) -> ModelResponse:
_json_response: Final = await raw_response.json()
model_response.id = _json_response.get("id")
model_response.choices = [Choices(**choice) for choice in _json_response.get("choices")]
model_response.choices = [
Choices.model_validate(choice) for choice in _JSON_OBJECTS.validate_python(_json_response.get("choices"))
]
model_response.created = _json_response.get("created")
model_response.model = _json_response.get("model")
model_response.object = _json_response.get("object")

View file

@ -224,6 +224,11 @@ def _write_back_message_text(message: _WritableMessage, target: MessageTextTarge
_TOOL_USE_INPUT_ADAPTER: Final = TypeAdapter(dict[str, object])
_RELEASED_TOOL_USE_STOP: Final = (
b"event: message_delta\n"
b'data: {"type": "message_delta", "delta": {"stop_reason": "tool_use", "stop_sequence": null}, '
b'"usage": {"output_tokens": 0}}\n\n'
)
def _rewritten_tool_use_input(arguments: str) -> Mapping[str, object] | None:
@ -1573,6 +1578,12 @@ class AnthropicMessagesHandler(BaseTranslation):
tool_calls_in_flight=bool(tool_use_fingerprints) and not stream_ended,
)
def released_stream_as_ended(self, responses_so_far: Sequence[object]) -> tuple[object, ...]:
released_key: Final = self.get_streaming_scan_key(responses_so_far)
if released_key is None or not released_key.tool_calls_in_flight:
return tuple(responses_so_far)
return (*responses_so_far, _RELEASED_TOOL_USE_STOP)
@classmethod
def _streamed_tool_use_fingerprints(cls, responses_so_far: Sequence[object]) -> tuple[str, ...]:
return tuple(

View file

@ -27,7 +27,7 @@ class AnthropicTokenCounter(BaseTokenCounter):
self,
model_to_use: str,
messages: list[dict[str, Any]] | None,
contents: list[dict[str, Any]] | None,
contents: list[dict[str, object]] | None,
deployment: dict[str, Any] | None = None,
request_model: str = "",
tools: list[dict[str, Any]] | None = None,

View file

@ -19,6 +19,7 @@ from typing import Final, cast
import httpx
from openai.types.file_deleted import FileDeleted
from pydantic import ConfigDict, TypeAdapter
from litellm.litellm_core_utils.prompt_templates.common_utils import extract_file_data
from litellm.litellm_core_utils.url_utils import encode_url_path_segment
@ -47,6 +48,8 @@ ANTHROPIC_FILES_API_BASE: Final = "https://api.anthropic.com"
ANTHROPIC_FILES_BETA_HEADER: Final = "files-api-2025-04-14"
ANTHROPIC_MESSAGE_BATCH_ID_PREFIX: Final = "msgbatch_"
_JSON_OBJECT: Final = TypeAdapter(dict[str, object], config=ConfigDict(hide_input_in_errors=True))
class AnthropicFilesConfig(BaseFilesConfig):
"""
@ -211,7 +214,7 @@ class AnthropicFilesConfig(BaseFilesConfig):
"created_at": "2025-01-01T00:00:00Z"
}
"""
response_json: Final = raw_response.json()
response_json: Final = _JSON_OBJECT.validate_python(raw_response.json())
return self._parse_anthropic_file(response_json)
def transform_retrieve_file_request(
@ -230,7 +233,7 @@ class AnthropicFilesConfig(BaseFilesConfig):
logging_obj: LiteLLMLoggingObj,
litellm_params: dict,
) -> OpenAIFileObject:
response_json: Final = raw_response.json()
response_json: Final = _JSON_OBJECT.validate_python(raw_response.json())
return self._parse_anthropic_file(response_json)
def transform_delete_file_request(
@ -249,13 +252,9 @@ class AnthropicFilesConfig(BaseFilesConfig):
logging_obj: LiteLLMLoggingObj,
litellm_params: dict,
) -> FileDeleted:
response_json: Final = raw_response.json()
response_json: Final = _JSON_OBJECT.validate_python(raw_response.json())
file_id: Final = response_json.get("id", "")
return FileDeleted(
id=file_id,
deleted=True,
object="file",
)
return FileDeleted.model_validate({"id": file_id, "deleted": True, "object": "file"})
def transform_list_files_request(
self,

View file

@ -5,10 +5,12 @@ Maps OpenAI-compatible audio transcription calls to Azure Speech REST
recognition for short audio.
"""
from typing import Any, Final
from collections.abc import Mapping
from typing import Final
from urllib.parse import urlencode, urlparse
import httpx
from pydantic import ConfigDict, TypeAdapter
from litellm.litellm_core_utils.audio_utils.utils import process_audio_file
from litellm.llms.base_llm.audio_transcription.transformation import (
@ -23,6 +25,8 @@ from litellm.types.llms.openai import (
)
from litellm.types.utils import FileTypes, TranscriptionResponse
_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True))
class AzureSpeechAudioTranscriptionException(BaseLLMException):
pass
@ -127,7 +131,8 @@ class AzureSpeechAudioTranscriptionConfig(BaseAudioTranscriptionConfig):
raw_response: httpx.Response,
) -> TranscriptionResponse:
response_json: Final = raw_response.json()
recognition_status: Final = response_json.get("RecognitionStatus")
payload: Final = _JSON_OBJECT.validate_python(response_json)
recognition_status: Final = payload.get("RecognitionStatus")
if recognition_status is not None and recognition_status != "Success":
raise AzureSpeechAudioTranscriptionException(
message=(f"Azure AI Speech transcription failed with RecognitionStatus={recognition_status}."),
@ -135,7 +140,7 @@ class AzureSpeechAudioTranscriptionConfig(BaseAudioTranscriptionConfig):
headers=raw_response.headers,
)
text: Final = self._extract_text(response_json)
text: Final = self._extract_text(payload)
response: Final = TranscriptionResponse(text=text)
response._hidden_params = response_json
return response
@ -194,9 +199,10 @@ class AzureSpeechAudioTranscriptionConfig(BaseAudioTranscriptionConfig):
return "detailed"
return "simple"
def _extract_text(self, response_json: dict[str, Any]) -> str:
if isinstance(response_json.get("DisplayText"), str):
return response_json["DisplayText"]
def _extract_text(self, response_json: Mapping[str, object]) -> str:
display_text: Final = response_json.get("DisplayText")
if isinstance(display_text, str):
return display_text
nbest: Final = response_json.get("NBest")
if isinstance(nbest, list) and nbest:

View file

@ -30,7 +30,7 @@ class AzureAIAnthropicCountTokensHandler(AzureAIAnthropicCountTokensConfig):
messages: list[dict[str, Any]],
api_key: str,
api_base: str,
litellm_params: dict[str, Any] | None = None,
litellm_params: dict[str, object] | None = None,
timeout: float | httpx.Timeout | None = None,
tools: list[dict[str, Any]] | None = None,
system: object = None,

View file

@ -203,6 +203,11 @@ class BaseTranslation(ABC):
def get_streaming_scan_key(self, responses_so_far: Sequence[object]) -> StreamingScanKey | None:
return None
def released_stream_as_ended(self, responses_so_far: Sequence[object]) -> tuple[object, ...]:
"""The chunks a client left the stream with, closed the way this endpoint ends a stream, so the
end-of-stream scan also inspects tool calls the stream never finished"""
return tuple(responses_so_far)
def build_block_sse_chunks(
self,
exc: "ModifyResponseException",

View file

@ -18,7 +18,7 @@ from __future__ import annotations
from abc import ABC, abstractmethod
from collections.abc import Mapping, Sequence
from dataclasses import dataclass, field
from typing import TYPE_CHECKING, Any, ClassVar, Generic, TypeVar
from typing import TYPE_CHECKING, ClassVar, Generic, TypeVar
from litellm.harness.errors import HarnessError, OptionsMismatch
from litellm.harness.types import Capabilities, Event, Harness
@ -133,7 +133,7 @@ class BaseCLIHarnessConfig(BaseHarnessConfig[OptionsT], Generic[OptionsT, Stream
"""Fresh per-turn parser state."""
@abstractmethod
def transform_stream_line(self, line: Mapping[str, Any], state: StreamStateT) -> Sequence[Event]:
def transform_stream_line(self, line: Mapping[str, object], state: StreamStateT) -> Sequence[Event]:
"""One decoded JSON line from stdout to zero or more events. Pure."""
@abstractmethod

View file

@ -1,6 +1,6 @@
from collections.abc import Mapping, Sequence
from datetime import datetime
from typing import TYPE_CHECKING, Any, Final, cast
from typing import TYPE_CHECKING, Final, Literal
from openai.types.batch import BatchRequestCounts
from openai.types.batch import Metadata as OpenAIBatchMetadata
@ -17,7 +17,12 @@ if TYPE_CHECKING:
# AWS Bedrock model-invocation-job statuses → OpenAI Batch statuses.
# Mirrors the mapping used by `BedrockBatchesConfig.transform_create_batch_response`
# so create / retrieve return consistent statuses.
_BEDROCK_MIJ_STATUS_TO_OPENAI: Final = {
_BEDROCK_MIJ_STATUS_TO_OPENAI: Final[
Mapping[
str,
Literal["validating", "failed", "in_progress", "finalizing", "completed", "expired", "cancelling", "cancelled"],
]
] = {
"Submitted": "validating",
"Validating": "validating",
"Scheduled": "validating",
@ -92,7 +97,7 @@ def _record_counts_from_response(response: Mapping[str, object]) -> BatchRequest
)
def _to_epoch(value: Any) -> int | None:
def _to_epoch(value: object) -> int | None:
if value is None:
return None
if isinstance(value, (int, float)):
@ -349,10 +354,7 @@ class BedrockBatchesHandler:
)
bedrock_status: Final = str(response.get("status", ""))
openai_status: Final = cast(
Any,
_BEDROCK_MIJ_STATUS_TO_OPENAI.get(bedrock_status, "in_progress"),
)
openai_status: Final = _BEDROCK_MIJ_STATUS_TO_OPENAI.get(bedrock_status, "in_progress")
input_uri: Final = response.get("inputDataConfig", {}).get("s3InputDataConfig", {}).get("s3Uri", "")
output_prefix: Final = response.get("outputDataConfig", {}).get("s3OutputDataConfig", {}).get("s3Uri", "")

View file

@ -25,6 +25,7 @@ import base64
from typing import TYPE_CHECKING, Any, Final
import httpx
from httpx._types import RequestFiles
from litellm.llms.base_llm.image_edit.transformation import BaseImageEditConfig
from litellm.llms.bedrock.common_utils import BedrockError
@ -165,7 +166,7 @@ class BedrockStabilityImageEditConfig(BaseImageEditConfig):
image_edit_optional_request_params: dict,
litellm_params: GenericLiteLLMParams,
headers: dict,
) -> tuple[dict, Any]:
) -> tuple[dict, RequestFiles]:
"""
Transform OpenAI-style request to Bedrock Stability request format.

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