mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
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
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:
commit
3a23cbb2d4
456 changed files with 19083 additions and 15192 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
;;
|
||||
|
|
|
|||
|
|
@ -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() {
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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 "" >>
|
||||
2
.github/ci-coverage-allowlist.yml
vendored
2
.github/ci-coverage-allowlist.yml
vendored
|
|
@ -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
|
||||
|
|
|
|||
17
.github/e2e-stack/down.sh
vendored
17
.github/e2e-stack/down.sh
vendored
|
|
@ -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
|
||||
83
.github/e2e-stack/redact_output.py
vendored
83
.github/e2e-stack/redact_output.py
vendored
|
|
@ -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, {'"': """}))
|
||||
)
|
||||
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())
|
||||
55
.github/e2e-stack/secrets_to_env.py
vendored
55
.github/e2e-stack/secrets_to_env.py
vendored
|
|
@ -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())
|
||||
52
.github/e2e-stack/select_tests.py
vendored
52
.github/e2e-stack/select_tests.py
vendored
|
|
@ -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())
|
||||
231
.github/e2e-stack/up.sh
vendored
231
.github/e2e-stack/up.sh
vendored
|
|
@ -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"
|
||||
89
.github/scripts/assert_ci_coverage.py
vendored
89
.github/scripts/assert_ci_coverage.py
vendored
|
|
@ -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))
|
||||
|
|
|
|||
2
.github/scripts/verify_linux_native_wheel.py
vendored
2
.github/scripts/verify_linux_native_wheel.py
vendored
|
|
@ -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),
|
||||
|
|
|
|||
29
.github/workflows/_test-unit-base.yml
vendored
29
.github/workflows/_test-unit-base.yml
vendored
|
|
@ -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
|
||||
|
|
|
|||
265
.github/workflows/test-e2e-changed.yml
vendored
265
.github/workflows/test-e2e-changed.yml
vendored
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
185
.github/workflows/test-postgres.yml
vendored
185
.github/workflows/test-postgres.yml
vendored
|
|
@ -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
|
||||
106
.github/workflows/test-redis-compat.yml
vendored
106
.github/workflows/test-redis-compat.yml
vendored
|
|
@ -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
|
||||
100
.github/workflows/test-unit-proxy-db.yml
vendored
100
.github/workflows/test-unit-proxy-db.yml
vendored
|
|
@ -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 }}
|
||||
|
|
|
|||
124
.github/workflows/test-unit.yml
vendored
124
.github/workflows/test-unit.yml
vendored
|
|
@ -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' }}
|
||||
|
|
|
|||
|
|
@ -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"},
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
"""
|
||||
|
|
|
|||
26
litellm-rust/Cargo.lock
generated
26
litellm-rust/Cargo.lock
generated
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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" }
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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 },
|
||||
}
|
||||
|
|
@ -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 { .. })));
|
||||
}
|
||||
}
|
||||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
@ -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,
|
||||
}
|
||||
|
|
@ -1 +0,0 @@
|
|||
SELECT 10;
|
||||
|
|
@ -1 +0,0 @@
|
|||
SELECT 1;
|
||||
|
|
@ -1 +0,0 @@
|
|||
SELECT 2;
|
||||
|
|
@ -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);
|
||||
}
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
#[derive(Debug, thiserror::Error)]
|
||||
#[derive(Clone, Debug, thiserror::Error)]
|
||||
pub enum Error {
|
||||
#[error("invalid ClickHouse insert row")]
|
||||
InvalidRow,
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
||||
|
|
|
|||
308
litellm-rust/crates/storage-clickhouse/src/migrate.rs
Normal file
308
litellm-rust/crates/storage-clickhouse/src/migrate.rs
Normal 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,
|
||||
}
|
||||
}
|
||||
476
litellm-rust/crates/storage-clickhouse/tests/migrations.rs
Normal file
476
litellm-rust/crates/storage-clickhouse/tests/migrations.rs
Normal 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
|
||||
);
|
||||
}
|
||||
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"] }
|
||||
|
|
|
|||
|
|
@ -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>),
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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, "ed_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)]
|
||||
|
|
|
|||
|
|
@ -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(())
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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());
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
56
litellm-rust/crates/traces/src/request.rs
Normal file
56
litellm-rust/crates/traces/src/request.rs
Normal 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,
|
||||
}
|
||||
|
|
@ -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>(),
|
||||
),
|
||||
])
|
||||
}
|
||||
|
|
|
|||
93
litellm-rust/crates/traces/tests/request_schema.rs
Normal file
93
litellm-rust/crates/traces/tests/request_schema.rs
Normal 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());
|
||||
}
|
||||
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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]]:
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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 "
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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 {}
|
||||
|
||||
|
|
|
|||
|
|
@ -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 = ""
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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, ...]:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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", "")
|
||||
|
|
|
|||
|
|
@ -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
Loading…
Add table
Reference in a new issue