mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
chore(mcp): merge main into listed-tool metadata branch
Some checks are pending
LiteLLM Rust / rust-lint (push) Waiting to run
LiteLLM Rust / rust-test (push) Waiting to run
LiteLLM Rust / rust-wheel (push) Waiting to run
Terraform Provider / gofmt, vet, build, test (push) Waiting to run
Terraform Provider / Provider endpoints vs proxy OpenAPI schema (push) Waiting to run
Some checks are pending
LiteLLM Rust / rust-lint (push) Waiting to run
LiteLLM Rust / rust-test (push) Waiting to run
LiteLLM Rust / rust-wheel (push) Waiting to run
Terraform Provider / gofmt, vet, build, test (push) Waiting to run
Terraform Provider / Provider endpoints vs proxy OpenAPI schema (push) Waiting to run
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
commit
8ffd7c107a
682 changed files with 31881 additions and 6765 deletions
|
|
@ -408,7 +408,7 @@ jobs:
|
|||
- run:
|
||||
name: Run Windows-specific test
|
||||
command: |
|
||||
uv run --no-sync python -m pytest tests/windows_tests/ -v
|
||||
uv run --no-sync python -m pytest --tb=short tests/windows_tests/ -v
|
||||
|
||||
windows_release_wheel:
|
||||
executor:
|
||||
|
|
@ -551,7 +551,7 @@ jobs:
|
|||
echo "$TEST_FILES" | circleci tests run \
|
||||
--split-by=timings \
|
||||
--verbose \
|
||||
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \
|
||||
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \
|
||||
-vv \
|
||||
--cov=./litellm --cov=./enterprise/litellm_enterprise \
|
||||
--cov-report=xml \
|
||||
|
|
@ -625,7 +625,7 @@ jobs:
|
|||
echo "$TEST_FILES" | circleci tests run \
|
||||
--split-by=timings \
|
||||
--verbose \
|
||||
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \
|
||||
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \
|
||||
-vv \
|
||||
--cov=./litellm --cov=./enterprise/litellm_enterprise \
|
||||
--cov-report=xml \
|
||||
|
|
@ -697,7 +697,7 @@ jobs:
|
|||
TEST_FILES=$(circleci tests glob "tests/local_testing/**/test_*.py")
|
||||
echo "$TEST_FILES" | circleci tests run \
|
||||
--verbose \
|
||||
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \
|
||||
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \
|
||||
-v \
|
||||
--junitxml=test-results/junit.xml \
|
||||
--durations=5 \
|
||||
|
|
@ -752,7 +752,7 @@ jobs:
|
|||
TEST_FILES=$(circleci tests glob "tests/proxy_admin_ui_tests/**/test_*.py")
|
||||
echo "$TEST_FILES" | circleci tests run \
|
||||
--verbose \
|
||||
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \
|
||||
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \
|
||||
-v \
|
||||
--cov=./litellm --cov=./enterprise/litellm_enterprise --cov-report=xml \
|
||||
--junitxml=test-results/junit.xml \
|
||||
|
|
@ -815,7 +815,7 @@ jobs:
|
|||
echo "$TEST_FILES" | circleci tests run \
|
||||
--split-by=timings \
|
||||
--verbose \
|
||||
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \
|
||||
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \
|
||||
-v \
|
||||
-k 'router' \
|
||||
-n 4 \
|
||||
|
|
@ -859,7 +859,7 @@ jobs:
|
|||
TEST_FILES=$(circleci tests glob "tests/router_unit_tests/**/test_*.py")
|
||||
echo "$TEST_FILES" | circleci tests run \
|
||||
--verbose \
|
||||
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \
|
||||
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \
|
||||
-v \
|
||||
--cov=./litellm --cov=./enterprise/litellm_enterprise --cov-report=xml \
|
||||
--junitxml=test-results/junit.xml \
|
||||
|
|
@ -904,7 +904,7 @@ jobs:
|
|||
TEST_FILES=$(circleci tests glob "tests/local_testing/**/test_*.py")
|
||||
echo "$TEST_FILES" | circleci tests run \
|
||||
--verbose \
|
||||
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \
|
||||
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \
|
||||
-v \
|
||||
--junitxml=test-results/junit.xml \
|
||||
--durations=5 \
|
||||
|
|
@ -948,7 +948,7 @@ jobs:
|
|||
TEST_FILES=$(circleci tests glob "tests/llm_translation/**/test_*.py" | grep -v "^tests/llm_translation/realtime/")
|
||||
echo "$TEST_FILES" | circleci tests run \
|
||||
--verbose \
|
||||
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \
|
||||
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \
|
||||
-v \
|
||||
--junitxml=test-results/junit.xml \
|
||||
--durations=20 \
|
||||
|
|
@ -986,7 +986,7 @@ jobs:
|
|||
TEST_FILES=$(circleci tests glob "tests/llm_translation/realtime/**/test_*.py")
|
||||
echo "$TEST_FILES" | circleci tests run \
|
||||
--verbose \
|
||||
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \
|
||||
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \
|
||||
-vv \
|
||||
--cov=./litellm --cov=./enterprise/litellm_enterprise --cov-report=xml \
|
||||
--junitxml=test-results/junit.xml \
|
||||
|
|
@ -1031,7 +1031,7 @@ jobs:
|
|||
TEST_FILES=$(circleci tests glob "tests/agent_tests/test_*.py")
|
||||
echo "$TEST_FILES" | circleci tests run \
|
||||
--verbose \
|
||||
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \
|
||||
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \
|
||||
-vv -s \
|
||||
--cov=./litellm --cov=./enterprise/litellm_enterprise --cov-report=xml \
|
||||
--junitxml=test-results/junit.xml \
|
||||
|
|
@ -1075,7 +1075,7 @@ jobs:
|
|||
TEST_FILES=$(circleci tests glob "tests/guardrails_tests/**/test_*.py")
|
||||
echo "$TEST_FILES" | circleci tests run \
|
||||
--verbose \
|
||||
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \
|
||||
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \
|
||||
-vv \
|
||||
--cov=./litellm --cov=./enterprise/litellm_enterprise --cov-report=xml \
|
||||
--junitxml=test-results/junit.xml \
|
||||
|
|
@ -1121,7 +1121,7 @@ jobs:
|
|||
TEST_FILES=$(circleci tests glob "tests/unified_google_tests/**/test_*.py")
|
||||
echo "$TEST_FILES" | circleci tests run \
|
||||
--verbose \
|
||||
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \
|
||||
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \
|
||||
-vv -s \
|
||||
--cov=./litellm --cov=./enterprise/litellm_enterprise --cov-report=xml \
|
||||
--junitxml=test-results/junit.xml \
|
||||
|
|
@ -1176,7 +1176,7 @@ jobs:
|
|||
TEST_FILES=$(circleci tests glob "tests/llm_responses_api_testing/**/test_*.py")
|
||||
echo "$TEST_FILES" | circleci tests run \
|
||||
--verbose \
|
||||
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \
|
||||
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \
|
||||
-v \
|
||||
--junitxml=test-results/junit.xml \
|
||||
--durations=5 \
|
||||
|
|
@ -1210,7 +1210,7 @@ jobs:
|
|||
TEST_FILES=$(circleci tests glob "tests/ocr_tests/**/test_*.py")
|
||||
echo "$TEST_FILES" | circleci tests run \
|
||||
--verbose \
|
||||
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \
|
||||
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \
|
||||
-vv \
|
||||
--cov=./litellm --cov=./enterprise/litellm_enterprise --cov-report=xml \
|
||||
--junitxml=test-results/junit.xml \
|
||||
|
|
@ -1254,7 +1254,7 @@ jobs:
|
|||
TEST_FILES=$(circleci tests glob "tests/search_tests/**/test_*.py")
|
||||
echo "$TEST_FILES" | circleci tests run \
|
||||
--verbose \
|
||||
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \
|
||||
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \
|
||||
-vv \
|
||||
--cov=./litellm --cov=./enterprise/litellm_enterprise --cov-report=xml \
|
||||
--junitxml=test-results/junit.xml \
|
||||
|
|
@ -1298,7 +1298,7 @@ jobs:
|
|||
TEST_FILES=$(circleci tests glob "tests/batches_tests/**/test_*.py")
|
||||
echo "$TEST_FILES" | circleci tests run \
|
||||
--verbose \
|
||||
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \
|
||||
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \
|
||||
-vv -s \
|
||||
--cov=./litellm --cov=./enterprise/litellm_enterprise --cov-report=xml \
|
||||
--junitxml=test-results/junit.xml \
|
||||
|
|
@ -1342,7 +1342,7 @@ jobs:
|
|||
TEST_FILES=$(circleci tests glob "tests/litellm_utils_tests/**/test_*.py")
|
||||
echo "$TEST_FILES" | circleci tests run \
|
||||
--verbose \
|
||||
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \
|
||||
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \
|
||||
-vv -s \
|
||||
--cov=./litellm --cov=./enterprise/litellm_enterprise --cov-report=xml \
|
||||
--junitxml=test-results/junit.xml \
|
||||
|
|
@ -1387,7 +1387,7 @@ jobs:
|
|||
TEST_FILES=$(circleci tests glob "tests/pass_through_unit_tests/**/test_*.py")
|
||||
echo "$TEST_FILES" | circleci tests run \
|
||||
--verbose \
|
||||
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \
|
||||
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \
|
||||
-vv \
|
||||
--cov=./litellm --cov=./enterprise/litellm_enterprise --cov-report=xml \
|
||||
--junitxml=test-results/junit.xml \
|
||||
|
|
@ -1432,7 +1432,7 @@ jobs:
|
|||
TEST_FILES=$(circleci tests glob "tests/image_gen_tests/**/test_*.py")
|
||||
echo "$TEST_FILES" | circleci tests run \
|
||||
--verbose \
|
||||
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \
|
||||
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \
|
||||
-v \
|
||||
--junitxml=test-results/junit.xml \
|
||||
--durations=5 \
|
||||
|
|
@ -1466,7 +1466,7 @@ jobs:
|
|||
TEST_FILES=$(circleci tests glob "tests/logging_callback_tests/**/test_*.py")
|
||||
echo "$TEST_FILES" | circleci tests run \
|
||||
--verbose \
|
||||
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \
|
||||
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \
|
||||
-vv \
|
||||
--cov=./litellm --cov=./enterprise/litellm_enterprise --cov-report=xml \
|
||||
-n 4 \
|
||||
|
|
@ -1511,7 +1511,7 @@ jobs:
|
|||
TEST_FILES=$(circleci tests glob "tests/audio_tests/**/test_*.py")
|
||||
echo "$TEST_FILES" | circleci tests run \
|
||||
--verbose \
|
||||
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \
|
||||
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \
|
||||
-vv -s \
|
||||
--cov=./litellm --cov=./enterprise/litellm_enterprise --cov-report=xml \
|
||||
--junitxml=test-results/junit.xml \
|
||||
|
|
@ -1531,61 +1531,6 @@ jobs:
|
|||
paths:
|
||||
- audio_coverage.xml
|
||||
- audio_coverage
|
||||
redis_caching_unit_tests:
|
||||
docker:
|
||||
- *python312_image
|
||||
working_directory: ~/project
|
||||
|
||||
steps:
|
||||
- checkout
|
||||
- skip_if_unrelated_changes
|
||||
- setup_google_dns
|
||||
- restore_cache:
|
||||
keys:
|
||||
- v1-uv-cache-{{ checksum "uv.lock" }}
|
||||
- install_uv
|
||||
- install_rust
|
||||
- run:
|
||||
name: Install Dependencies
|
||||
command: |
|
||||
uv sync --frozen --all-groups --all-extras --python 3.12
|
||||
- save_cache:
|
||||
paths:
|
||||
- ~/.cache/uv
|
||||
key: v1-uv-cache-{{ checksum "uv.lock" }}
|
||||
# Run pytest and generate JUnit XML report
|
||||
- run:
|
||||
name: Run tests
|
||||
command: |
|
||||
mkdir -p test-results
|
||||
TEST_FILES=$(printf "%s\n" \
|
||||
tests/local_testing/test_dual_cache.py \
|
||||
tests/local_testing/test_redis_batch_optimizations.py \
|
||||
tests/local_testing/test_redis_increment_with_floor.py \
|
||||
tests/local_testing/test_router_utils.py)
|
||||
echo "$TEST_FILES" | circleci tests run \
|
||||
--verbose \
|
||||
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \
|
||||
-vv -s \
|
||||
--cov=./litellm --cov=./enterprise/litellm_enterprise --cov-report=xml \
|
||||
--junitxml=test-results/junit.xml \
|
||||
--durations=5 -n 2 \
|
||||
--reruns 2 --reruns-delay 1"
|
||||
no_output_timeout: 20m
|
||||
- run:
|
||||
name: Rename the coverage files
|
||||
command: |
|
||||
mv coverage.xml redis_caching_coverage.xml
|
||||
mv .coverage redis_caching_coverage
|
||||
|
||||
# Store test results
|
||||
- store_test_results:
|
||||
path: test-results
|
||||
- persist_to_workspace:
|
||||
root: .
|
||||
paths:
|
||||
- redis_caching_coverage.xml
|
||||
- redis_caching_coverage
|
||||
installing_litellm_on_python:
|
||||
docker:
|
||||
- *python312_image
|
||||
|
|
@ -1605,7 +1550,7 @@ jobs:
|
|||
- run:
|
||||
name: Run tests
|
||||
command: |
|
||||
uv run --no-sync python -m pytest -vv tests/local_testing/test_basic_python_version.py -k "not legacy_resolver"
|
||||
uv run --no-sync python -m pytest --tb=short -vv tests/local_testing/test_basic_python_version.py -k "not legacy_resolver"
|
||||
|
||||
installing_litellm_on_python_3_13:
|
||||
docker:
|
||||
|
|
@ -1629,7 +1574,7 @@ jobs:
|
|||
- run:
|
||||
name: Run tests
|
||||
command: |
|
||||
uv run --no-sync python -m pytest -v tests/local_testing/test_basic_python_version.py -k "not legacy_resolver"
|
||||
uv run --no-sync python -m pytest --tb=short -v tests/local_testing/test_basic_python_version.py -k "not legacy_resolver"
|
||||
|
||||
installing_litellm_on_python_v2_migration_resolver:
|
||||
docker:
|
||||
|
|
@ -1660,7 +1605,7 @@ jobs:
|
|||
- run:
|
||||
name: Run both migration resolvers against Postgres
|
||||
command: |
|
||||
uv run --no-sync python -m pytest -vv \
|
||||
uv run --no-sync python -m pytest --tb=short -vv \
|
||||
tests/local_testing/test_basic_python_version.py::test_litellm_proxy_server_config_no_general_settings \
|
||||
tests/local_testing/test_basic_python_version.py::test_litellm_proxy_server_config_no_general_settings_legacy_resolver
|
||||
|
||||
|
|
@ -1829,7 +1774,7 @@ jobs:
|
|||
TEST_FILES=$(circleci tests glob "tests/basic_proxy_startup_tests/**/test_*.py")
|
||||
echo "$TEST_FILES" | circleci tests run \
|
||||
--verbose \
|
||||
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \
|
||||
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \
|
||||
-v \
|
||||
--junitxml=test-results/junit-2.xml \
|
||||
--durations=5"
|
||||
|
|
@ -1926,7 +1871,7 @@ jobs:
|
|||
TEST_FILES=$(circleci tests glob "tests/test_*.py")
|
||||
echo "$TEST_FILES" | circleci tests run \
|
||||
--verbose \
|
||||
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \
|
||||
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \
|
||||
-s -v \
|
||||
--junitxml=test-results/junit.xml \
|
||||
-n 4 \
|
||||
|
|
@ -2013,7 +1958,7 @@ jobs:
|
|||
TEST_FILES=$(circleci tests glob "tests/openai_endpoints_tests/**/test_*.py")
|
||||
echo "$TEST_FILES" | circleci tests run \
|
||||
--verbose \
|
||||
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \
|
||||
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \
|
||||
-s -vv \
|
||||
--junitxml=test-results/junit.xml \
|
||||
--durations=5"
|
||||
|
|
@ -2096,7 +2041,7 @@ jobs:
|
|||
TEST_FILES=$(circleci tests glob "tests/otel_tests/**/test_*.py")
|
||||
echo "$TEST_FILES" | circleci tests run \
|
||||
--verbose \
|
||||
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \
|
||||
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \
|
||||
-v \
|
||||
--junitxml=test-results/junit.xml \
|
||||
--durations=5"
|
||||
|
|
@ -2148,7 +2093,7 @@ jobs:
|
|||
TEST_FILES=$(circleci tests glob "tests/basic_proxy_startup_tests/**/test_*.py")
|
||||
echo "$TEST_FILES" | circleci tests run \
|
||||
--verbose \
|
||||
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \
|
||||
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \
|
||||
-v \
|
||||
--junitxml=test-results/junit-2.xml \
|
||||
--durations=5"
|
||||
|
|
@ -2229,7 +2174,7 @@ jobs:
|
|||
TEST_FILES=$(circleci tests glob "tests/spend_tracking_tests/**/test_*.py")
|
||||
echo "$TEST_FILES" | circleci tests run \
|
||||
--verbose \
|
||||
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \
|
||||
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \
|
||||
-vv \
|
||||
--junitxml=test-results/junit.xml \
|
||||
--durations=5"
|
||||
|
|
@ -2334,7 +2279,7 @@ jobs:
|
|||
TEST_FILES=$(circleci tests glob "tests/multi_instance_e2e_tests/**/test_*.py")
|
||||
echo "$TEST_FILES" | circleci tests run \
|
||||
--verbose \
|
||||
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \
|
||||
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \
|
||||
-vv \
|
||||
--junitxml=test-results/junit.xml \
|
||||
--durations=5"
|
||||
|
|
@ -2406,7 +2351,7 @@ jobs:
|
|||
TEST_FILES=$(circleci tests glob "tests/store_model_in_db_tests/**/test_*.py")
|
||||
echo "$TEST_FILES" | circleci tests run \
|
||||
--verbose \
|
||||
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \
|
||||
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \
|
||||
-vv \
|
||||
--junitxml=test-results/junit.xml \
|
||||
--durations=5"
|
||||
|
|
@ -2491,7 +2436,7 @@ jobs:
|
|||
TEST_FILES=$(circleci tests glob "tests/basic_proxy_startup_tests/**/test_*.py")
|
||||
echo "$TEST_FILES" | circleci tests run \
|
||||
--verbose \
|
||||
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \
|
||||
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \
|
||||
-vv \
|
||||
--junitxml=test-results/junit-2.xml \
|
||||
--durations=5"
|
||||
|
|
@ -2588,7 +2533,7 @@ jobs:
|
|||
TEST_FILES=$(circleci tests glob "tests/pass_through_tests/**/test_*.py")
|
||||
echo "$TEST_FILES" | circleci tests run \
|
||||
--verbose \
|
||||
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \
|
||||
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \
|
||||
-v \
|
||||
--junitxml=test-results/junit.xml \
|
||||
--durations=5"
|
||||
|
|
@ -2659,7 +2604,7 @@ jobs:
|
|||
TEST_FILES=$(circleci tests glob "tests/proxy_e2e_anthropic_messages_tests/**/test_*.py")
|
||||
echo "$TEST_FILES" | circleci tests run \
|
||||
--verbose \
|
||||
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest \
|
||||
--command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \
|
||||
-vv -s \
|
||||
--junitxml=test-results/junit.xml \
|
||||
--durations=5"
|
||||
|
|
@ -2689,7 +2634,7 @@ jobs:
|
|||
- run:
|
||||
name: Combine Coverage
|
||||
command: |
|
||||
uv tool run --from 'coverage[toml]==7.10.6' coverage combine realtime_translation_coverage ocr_coverage search_coverage logging_coverage audio_coverage local_testing_part1_coverage local_testing_part2_coverage pass_through_unit_tests_coverage batches_coverage guardrails_coverage redis_caching_coverage agent_coverage google_generate_content_endpoint_coverage litellm_utils_coverage router_unit_tests_coverage auth_ui_unit_tests_coverage
|
||||
uv tool run --from 'coverage[toml]==7.10.6' coverage combine realtime_translation_coverage ocr_coverage search_coverage logging_coverage audio_coverage local_testing_part1_coverage local_testing_part2_coverage pass_through_unit_tests_coverage batches_coverage guardrails_coverage agent_coverage google_generate_content_endpoint_coverage litellm_utils_coverage router_unit_tests_coverage auth_ui_unit_tests_coverage
|
||||
uv tool run --from 'coverage[toml]==7.10.6' coverage xml
|
||||
- codecov/upload:
|
||||
file: ./coverage.xml
|
||||
|
|
@ -3189,7 +3134,7 @@ jobs:
|
|||
name: Test provider capture and replay harness
|
||||
command: |
|
||||
mkdir -p test-results/provider-replay-harness
|
||||
uv run --no-sync pytest -q --noconftest -o addopts= -o pythonpath=tests/e2e -p no:rerunfailures \
|
||||
uv run --no-sync pytest --tb=short -q --noconftest -o addopts= -o pythonpath=tests/e2e -p no:rerunfailures \
|
||||
--junitxml=test-results/provider-replay-harness/junit.xml \
|
||||
tests/e2e/test_provider_edge.py tests/e2e/test_fixture_bundle.py \
|
||||
tests/e2e/test_fixture_canonical.py tests/e2e/test_fixture_mode.py \
|
||||
|
|
@ -3492,7 +3437,6 @@ workflows:
|
|||
- image_gen_testing
|
||||
- logging_testing
|
||||
- audio_testing
|
||||
- redis_caching_unit_tests
|
||||
- upload-coverage:
|
||||
requires:
|
||||
- realtime_translation_testing
|
||||
|
|
@ -3507,7 +3451,6 @@ workflows:
|
|||
- image_gen_testing
|
||||
- logging_testing
|
||||
- audio_testing
|
||||
- redis_caching_unit_tests
|
||||
- langfuse_logging_unit_tests
|
||||
- local_testing_part1
|
||||
- local_testing_part2
|
||||
|
|
|
|||
|
|
@ -191,7 +191,7 @@ if [ "$suite" = management ] || [ "$suite" = mcp ]; then
|
|||
fi
|
||||
|
||||
if [ "$suite" = providers ]; then
|
||||
INTEGRATION_RUN_ID="$integration_identity" .venv/bin/python -m pytest --noconftest -o addopts= \
|
||||
INTEGRATION_RUN_ID="$integration_identity" .venv/bin/python -m pytest --tb=short --noconftest -o addopts= \
|
||||
--strict-markers --strict-config -p no:pytest-retry -p no:rerunfailures --timeout=30 \
|
||||
tests/e2e/test_provider_edge.py::TestReplayMode::test_content_drift_returns_the_miss_status_naming_both_keys \
|
||||
tests/e2e/test_provider_edge.py::TestReplayMode::test_exhausted_key_returns_the_miss_status \
|
||||
|
|
|
|||
|
|
@ -77,6 +77,7 @@ legacy_paths() {
|
|||
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
|
||||
|
|
@ -116,7 +117,7 @@ legacy_paths() {
|
|||
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/test_response_polling_handler.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
|
||||
|
|
@ -145,12 +146,14 @@ legacy_paths() {
|
|||
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)
|
||||
|
|
|
|||
8
.github/ci-coverage-allowlist.yml
vendored
8
.github/ci-coverage-allowlist.yml
vendored
|
|
@ -4,6 +4,14 @@ description: >-
|
|||
by a job nor listed here, so every entry below is a decision on the record.
|
||||
|
||||
test_paths:
|
||||
- reason: >-
|
||||
litellm.agent() end-to-end suite. It drives the real claude, codex and opencode CLIs and
|
||||
deepagents against a live LiteLLM AI Gateway, so it needs those binaries on PATH plus
|
||||
LITELLM_PROXY_API_BASE / LITELLM_PROXY_API_KEY, and skips without them. Run manually
|
||||
before changing litellm/harness; the mocked coverage runs in tests/unit/harness and
|
||||
tests/unit/llms/*/harness
|
||||
paths:
|
||||
- tests/harness_e2e
|
||||
- reason: >-
|
||||
The Rust/Python parity harness is run manually through its local CLI. Recorded replay,
|
||||
fixture generation, and harness checks are intentionally outside pull request CI
|
||||
|
|
|
|||
48
.github/scripts/assert_ci_coverage.py
vendored
48
.github/scripts/assert_ci_coverage.py
vendored
|
|
@ -9,6 +9,7 @@ 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
|
||||
|
|
@ -35,7 +36,7 @@ GLOB_CHARS = frozenset("*?")
|
|||
# itself decomposed one level deeper and is checked through its own entry.
|
||||
SHARDED_ROOTS: tuple[str, ...] = (
|
||||
"tests/test_litellm",
|
||||
"tests/test_litellm/proxy",
|
||||
"tests/unit/proxy",
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -119,11 +120,48 @@ def _invoked_test_tokens(scalars: Iterable[Scalar]) -> frozenset[str]:
|
|||
)
|
||||
|
||||
|
||||
def _unit_selection_tokens(repo_root: pathlib.Path = REPO_ROOT) -> frozenset[str]:
|
||||
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 frozenset()
|
||||
return frozenset(match.group(0).rstrip("/") for match in TEST_TOKEN_RE.finditer(_uncommented(script.read_text())))
|
||||
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 _built_dockerfile_tokens(scalars: Iterable[Scalar]) -> frozenset[str]:
|
||||
|
|
@ -480,7 +518,7 @@ def _check_slices() -> int:
|
|||
|
||||
|
||||
def _check_shards() -> int:
|
||||
findings = _unassigned_shard_children(_invoked_test_tokens(_all_scalars()))
|
||||
findings = _unassigned_shard_children(_shard_tokens(_all_scalars(), _unit_selection_arms()))
|
||||
if findings:
|
||||
_report(
|
||||
"test directories and files that no shard claims",
|
||||
|
|
|
|||
116
.github/workflows/test-unit.yml
vendored
116
.github/workflows/test-unit.yml
vendored
|
|
@ -204,7 +204,6 @@ jobs:
|
|||
tests/unit/proxy/spend_tracking
|
||||
--ignore=tests/unit/proxy/spend_tracking/test_search_api_logging.py
|
||||
tests/unit/proxy/pass_through_endpoints
|
||||
tests/test_litellm/proxy/pass_through_endpoints
|
||||
tests/unit/proxy/_experimental
|
||||
--ignore=tests/unit/proxy/_experimental/mcp_server
|
||||
tests/unit/proxy/experimental
|
||||
|
|
@ -217,93 +216,46 @@ jobs:
|
|||
tests/unit/proxy/enterprise_billing
|
||||
tests/unit/proxy/types_utils
|
||||
tests/unit/proxy/logging_endpoints
|
||||
tests/unit/proxy/test__types.py
|
||||
tests/unit/proxy/test_aiohttp_cleanup_closed.py
|
||||
tests/unit/proxy/test_aiohttp_session_recovery.py
|
||||
tests/unit/proxy/test_api_key_masking_in_errors.py
|
||||
tests/unit/proxy/test_audio_speech_prometheus_hooks.py
|
||||
tests/unit/proxy/test_batch_expiry.py
|
||||
tests/unit/proxy/test_batch_metadata_none_fix.py
|
||||
tests/unit/proxy/test_batch_retrieve_bedrock.py
|
||||
tests/unit/proxy/test_batch_x_litellm_model_encoding.py
|
||||
tests/unit/proxy/test_blocked_response_usage.py
|
||||
tests/unit/proxy/test_body_snapshot_callback_params.py
|
||||
tests/unit/proxy/test_budget_reservation.py
|
||||
tests/unit/proxy/test_bug_report_config.py
|
||||
tests/unit/proxy/test_caching_routes.py
|
||||
tests/unit/proxy/test_chat_completion_metadata.py
|
||||
tests/unit/proxy/test_claude_code_marketplace.py
|
||||
tests/unit/proxy/test_collector.py
|
||||
tests/unit/proxy/test_common_request_processing.py
|
||||
tests/unit/proxy/test_component_allowlists.py
|
||||
tests/unit/proxy/test_conftest.py
|
||||
tests/unit/proxy/test_cors_config.py
|
||||
tests/unit/proxy/test_custom_proxy.py
|
||||
tests/unit/proxy/test_dynamic_mcp_route.py
|
||||
tests/unit/proxy/test_empty_model_list.py
|
||||
tests/unit/proxy/test_enforce_user_param.py
|
||||
tests/unit/proxy/test_fallback_management_endpoints.py
|
||||
tests/unit/proxy/test_fastapi_offline_routes.py
|
||||
tests/unit/proxy/test_filter_models_by_team_access_group.py
|
||||
tests/unit/proxy/test_health_check_functions.py
|
||||
tests/unit/proxy/test_health_check_max_tokens.py
|
||||
tests/unit/proxy/test_init_litellm_callbacks.py
|
||||
tests/unit/proxy/test_langfuse_passthrough_security.py
|
||||
tests/unit/proxy/test_lazy_openapi_snapshot.py
|
||||
tests/unit/proxy/test_litellm_pre_call_utils.py
|
||||
tests/unit/proxy/test_max_budget_env_var.py
|
||||
tests/unit/proxy/test_mcp_asgi_response.py
|
||||
tests/unit/proxy/test_model_based_routing_files_batches.py
|
||||
tests/unit/proxy/test_model_deprecations_endpoint.py
|
||||
tests/unit/proxy/test_model_dump_with_preserved_fields.py
|
||||
tests/unit/proxy/test_model_id_header_propagation.py
|
||||
tests/unit/proxy/test_model_info_default_limits.py
|
||||
tests/unit/proxy/test_model_level_guardrails.py
|
||||
tests/unit/proxy/test_model_list_aliases.py
|
||||
tests/unit/proxy/test_model_list_callback_filter.py
|
||||
tests/unit/proxy/test_model_list_discoverable.py
|
||||
tests/unit/proxy/test_model_list_healthy_only.py
|
||||
tests/unit/proxy/test_modify_response_streaming_passthrough.py
|
||||
tests/unit/proxy/test_native_compaction.py
|
||||
tests/unit/proxy/test_openai_ws_passthrough_routes.py
|
||||
tests/unit/proxy/test_openapi_schema_validation.py
|
||||
tests/unit/proxy/test_plugin_routes.py
|
||||
tests/unit/proxy/test_pointfive_dashboard_config.py
|
||||
tests/unit/proxy/test_pointfive_ui_callback.py
|
||||
tests/unit/proxy/test_pricing_field_strip.py
|
||||
tests/unit/proxy/test_prisma_engine_watchdog.py
|
||||
tests/unit/proxy/test_prisma_migration.py
|
||||
tests/unit/proxy/test_prometheus_cleanup.py
|
||||
tests/unit/proxy/test_prometheus_metrics_server.py
|
||||
tests/unit/proxy/test_provider_url_destination_guard.py
|
||||
tests/unit/proxy/test_proxy_cli.py
|
||||
tests/unit/proxy/test_proxy_logging_hook_detection.py
|
||||
tests/unit/proxy/test_proxy_types.py
|
||||
tests/unit/proxy/test_pyroscope.py
|
||||
tests/unit/proxy/test_read_model_list.py
|
||||
tests/unit/proxy/test_redis_auth_cache_flag.py
|
||||
tests/unit/proxy/test_response_model_sanitization.py
|
||||
tests/unit/proxy/test_route_a2a_models.py
|
||||
tests/unit/proxy/test_route_llm_request.py
|
||||
tests/unit/proxy/test_route_priority.py
|
||||
tests/unit/proxy/test_sensitive_route_auth.py
|
||||
tests/unit/proxy/test_shared_health_check.py
|
||||
tests/unit/proxy/test_spend_log_cleanup.py
|
||||
tests/unit/proxy/test_swagger_chat_completions.py
|
||||
tests/unit/proxy/test_team_member_update.py
|
||||
tests/unit/proxy/test_team_org_move.py
|
||||
tests/unit/proxy/test_tools_allowlist_enforcement.py
|
||||
tests/unit/proxy/test_tracing_endpoints.py
|
||||
tests/unit/proxy/test_update_llm_router_resilience.py
|
||||
tests/unit/proxy/test_zerobus_dashboard_config.py
|
||||
tests/unit/proxy/test_proxy_server_endpoints_and_startup.py
|
||||
tests/unit/proxy/test_proxy_utils_model_creation_and_error_logging.py
|
||||
unit-flag: proxy-infra
|
||||
workers: 4
|
||||
reruns: 2
|
||||
timeout-minutes: 20
|
||||
job-timeout-minutes: 60
|
||||
|
||||
- shard: proxy-infra-root
|
||||
artifact-name: proxy-infra-root
|
||||
test-path: >-
|
||||
tests/unit/proxy/test_*.py
|
||||
--ignore=tests/unit/proxy/test_aproxy_startup.py
|
||||
--ignore=tests/unit/proxy/test_credential_slot_registry.py
|
||||
--ignore=tests/unit/proxy/test_custom_callback_input.py
|
||||
--ignore=tests/unit/proxy/test_custom_logger_s3_gcs.py
|
||||
--ignore=tests/unit/proxy/test_custom_tokenizer_bug.py
|
||||
--ignore=tests/unit/proxy/test_db_schema_changes.py
|
||||
--ignore=tests/unit/proxy/test_deprecated_key_grace_period.py
|
||||
--ignore=tests/unit/proxy/test_get_favicon.py
|
||||
--ignore=tests/unit/proxy/test_get_image.py
|
||||
--ignore=tests/unit/proxy/test_prisma_client_backoff_retry.py
|
||||
--ignore=tests/unit/proxy/test_prompt_test_endpoint.py
|
||||
--ignore=tests/unit/proxy/test_proxy_config_unit_test.py
|
||||
--ignore=tests/unit/proxy/test_proxy_custom_auth.py
|
||||
--ignore=tests/unit/proxy/test_proxy_reject_logging.py
|
||||
--ignore=tests/unit/proxy/test_proxy_server.py
|
||||
--ignore=tests/unit/proxy/test_proxy_setting_guardrails.py
|
||||
--ignore=tests/unit/proxy/test_proxy_token_counter.py
|
||||
--ignore=tests/unit/proxy/test_proxy_utils.py
|
||||
--ignore=tests/unit/proxy/test_reducto_ocr_route.py
|
||||
--ignore=tests/unit/proxy/test_response_polling_pre_call_checks.py
|
||||
--ignore=tests/unit/proxy/test_server_root_path.py
|
||||
--ignore=tests/unit/proxy/test_ui_path_detection.py
|
||||
--ignore=tests/unit/proxy/test_unit_test_proxy_hooks.py
|
||||
--ignore=tests/unit/proxy/test_update_spend.py
|
||||
--ignore=tests/unit/proxy/test_zero_cost_model_budget_bypass.py
|
||||
workers: 4
|
||||
reruns: 2
|
||||
timeout-minutes: 20
|
||||
job-timeout-minutes: 60
|
||||
|
||||
- shard: caching-local
|
||||
artifact-name: caching-local
|
||||
test-path: ""
|
||||
|
|
|
|||
|
|
@ -62,7 +62,7 @@ Never edit or commit `ruff-strict-budget.json`, `type-discipline-budget.json`, `
|
|||
|
||||
If you're trying to create a new function that relies on untyped stuff, instead of adding more Any's and pushing `reportAny` / `reportExplicitAny` closer to their basedpyright ceilings, just validate it in the caller with Pydantic (a model or `TypeAdapter` that returns the typed thing or raises will do) and then pass the now typed variable in
|
||||
|
||||
If you get an LIT001 or LIT002 fail, refactor the code to follow functional programming best practices rather than introducing mutable data structures. For example, build values in one shot with comprehensions or generators wrapped in `tuple()` / `MappingProxyType()` / `frozenset()` instead of seeding an empty `list`/`dict`/`set` and mutating it over time. Ideally, `# mutable-ok` is never used; reach for it only as a genuine last resort when an immutable rewrite is truly impossible, and always pair it with a real reason
|
||||
If you get an LIT001 fail, refactor the code to follow functional programming best practices rather than introducing mutable data structures. For example, build values in one shot with comprehensions or generators wrapped in `tuple()` / `MappingProxyType()` / `frozenset()` instead of seeding an empty `list`/`dict`/`set` and mutating it over time. Ideally, `# mutable-ok` is never used; reach for it only as a genuine last resort when an immutable rewrite is truly impossible, and always pair it with a real reason
|
||||
|
||||
Every lint or type suppression must name the exact rule inside brackets and carry a reason comment, e.g. `# pyright: ignore[reportArgumentType] # stubs lack async overload` or `# noqa: TID251 # <reason>`. `# type: ignore` is banned (LIT009): pyrightconfig.json sets `enableTypeIgnoreComments` to false, so it silently does nothing
|
||||
|
||||
|
|
|
|||
|
|
@ -136,7 +136,7 @@ If you're running broader test suites, proxy tests, or anything that touches Pos
|
|||
make install-test-deps
|
||||
```
|
||||
|
||||
This syncs the locked test environment used across the repo, including `psycopg` v3 plus `psycopg-binary` (used by `pytest-postgresql`), `psycopg2-binary` (used by some proxy E2E tests), and a generated Prisma client for DB-backed proxy tests, so pytest startup matches CI without manual package installs.
|
||||
This syncs the locked test environment used across the repo, including `psycopg` v3 plus `psycopg-binary`, `psycopg2-binary` (used by some proxy E2E tests), and a generated Prisma client for DB-backed proxy tests, so pytest startup matches CI without manual package installs.
|
||||
|
||||
### Running Linting and Formatting Checks
|
||||
|
||||
|
|
|
|||
89
Makefile
89
Makefile
|
|
@ -1,7 +1,7 @@
|
|||
# LiteLLM Makefile
|
||||
# Simple Makefile for running tests and basic development tasks
|
||||
|
||||
.PHONY: help test test-unit test-unit-llms test-unit-proxy-guardrails test-unit-proxy-core test-unit-proxy-misc \
|
||||
.PHONY: help test test-unit test-unit-llms test-unit-proxy-guardrails test-unit-proxy-core test-unit-proxy-misc test-unit-proxy-root \
|
||||
test-unit-integrations test-unit-core-utils test-unit-other test-unit-root \
|
||||
test-proxy-unit-a test-proxy-unit-b test-integration test-unit-helm \
|
||||
test-rust-extension rust-sqlx-prepare \
|
||||
|
|
@ -47,6 +47,7 @@ help:
|
|||
@echo " make test-unit-proxy-guardrails - Run proxy guardrails+mgmt tests (~51 files)"
|
||||
@echo " make test-unit-proxy-core - Run proxy auth+client+db+hooks tests (~52 files)"
|
||||
@echo " make test-unit-proxy-misc - Run proxy misc tests (~77 files)"
|
||||
@echo " make test-unit-proxy-root - Run proxy root-file tests (tests/unit/proxy/test_*.py)"
|
||||
@echo " make test-unit-integrations - Run integration tests (~60 files)"
|
||||
@echo " make test-unit-core-utils - Run core utils tests (~32 files)"
|
||||
@echo " make test-unit-other - Run other tests (caching, responses, etc., ~69 files)"
|
||||
|
|
@ -326,89 +327,11 @@ test-unit-proxy-guardrails: install-test-deps
|
|||
test-unit-proxy-core: install-test-deps
|
||||
$(UV_RUN) pytest tests/unit/proxy/auth tests/unit/proxy/client tests/unit/proxy/db tests/unit/proxy/hooks tests/unit/proxy/policy_engine --ignore=tests/unit/proxy/db/db_transaction_queue/test_e2e_pod_lock_manager.py --ignore=tests/unit/proxy/db/test_update_daily_tag_spend.py --tb=short -vv -n 4 --durations=20
|
||||
|
||||
PROXY_INFRA_ROOT_TESTS := \
|
||||
tests/unit/proxy/test__types.py \
|
||||
tests/unit/proxy/test_aiohttp_cleanup_closed.py \
|
||||
tests/unit/proxy/test_aiohttp_session_recovery.py \
|
||||
tests/unit/proxy/test_api_key_masking_in_errors.py \
|
||||
tests/unit/proxy/test_audio_speech_prometheus_hooks.py \
|
||||
tests/unit/proxy/test_batch_expiry.py \
|
||||
tests/unit/proxy/test_batch_metadata_none_fix.py \
|
||||
tests/unit/proxy/test_batch_retrieve_bedrock.py \
|
||||
tests/unit/proxy/test_batch_x_litellm_model_encoding.py \
|
||||
tests/unit/proxy/test_blocked_response_usage.py \
|
||||
tests/unit/proxy/test_body_snapshot_callback_params.py \
|
||||
tests/unit/proxy/test_budget_reservation.py \
|
||||
tests/unit/proxy/test_bug_report_config.py \
|
||||
tests/unit/proxy/test_caching_routes.py \
|
||||
tests/unit/proxy/test_chat_completion_metadata.py \
|
||||
tests/unit/proxy/test_claude_code_marketplace.py \
|
||||
tests/unit/proxy/test_collector.py \
|
||||
tests/unit/proxy/test_common_request_processing.py \
|
||||
tests/unit/proxy/test_component_allowlists.py \
|
||||
tests/unit/proxy/test_conftest.py \
|
||||
tests/unit/proxy/test_cors_config.py \
|
||||
tests/unit/proxy/test_custom_proxy.py \
|
||||
tests/unit/proxy/test_dynamic_mcp_route.py \
|
||||
tests/unit/proxy/test_empty_model_list.py \
|
||||
tests/unit/proxy/test_enforce_user_param.py \
|
||||
tests/unit/proxy/test_fallback_management_endpoints.py \
|
||||
tests/unit/proxy/test_fastapi_offline_routes.py \
|
||||
tests/unit/proxy/test_filter_models_by_team_access_group.py \
|
||||
tests/unit/proxy/test_health_check_functions.py \
|
||||
tests/unit/proxy/test_health_check_max_tokens.py \
|
||||
tests/unit/proxy/test_init_litellm_callbacks.py \
|
||||
tests/unit/proxy/test_langfuse_passthrough_security.py \
|
||||
tests/unit/proxy/test_lazy_openapi_snapshot.py \
|
||||
tests/unit/proxy/test_litellm_pre_call_utils.py \
|
||||
tests/unit/proxy/test_max_budget_env_var.py \
|
||||
tests/unit/proxy/test_mcp_asgi_response.py \
|
||||
tests/unit/proxy/test_model_based_routing_files_batches.py \
|
||||
tests/unit/proxy/test_model_deprecations_endpoint.py \
|
||||
tests/unit/proxy/test_model_dump_with_preserved_fields.py \
|
||||
tests/unit/proxy/test_model_id_header_propagation.py \
|
||||
tests/unit/proxy/test_model_info_default_limits.py \
|
||||
tests/unit/proxy/test_model_level_guardrails.py \
|
||||
tests/unit/proxy/test_model_list_aliases.py \
|
||||
tests/unit/proxy/test_model_list_callback_filter.py \
|
||||
tests/unit/proxy/test_model_list_discoverable.py \
|
||||
tests/unit/proxy/test_model_list_healthy_only.py \
|
||||
tests/unit/proxy/test_modify_response_streaming_passthrough.py \
|
||||
tests/unit/proxy/test_native_compaction.py \
|
||||
tests/unit/proxy/test_openai_ws_passthrough_routes.py \
|
||||
tests/unit/proxy/test_openapi_schema_validation.py \
|
||||
tests/unit/proxy/test_plugin_routes.py \
|
||||
tests/unit/proxy/test_pointfive_dashboard_config.py \
|
||||
tests/unit/proxy/test_pointfive_ui_callback.py \
|
||||
tests/unit/proxy/test_pricing_field_strip.py \
|
||||
tests/unit/proxy/test_prisma_engine_watchdog.py \
|
||||
tests/unit/proxy/test_prisma_migration.py \
|
||||
tests/unit/proxy/test_prometheus_cleanup.py \
|
||||
tests/unit/proxy/test_prometheus_metrics_server.py \
|
||||
tests/unit/proxy/test_provider_url_destination_guard.py \
|
||||
tests/unit/proxy/test_proxy_cli.py \
|
||||
tests/unit/proxy/test_proxy_logging_hook_detection.py \
|
||||
tests/unit/proxy/test_proxy_types.py \
|
||||
tests/unit/proxy/test_pyroscope.py \
|
||||
tests/unit/proxy/test_read_model_list.py \
|
||||
tests/unit/proxy/test_redis_auth_cache_flag.py \
|
||||
tests/unit/proxy/test_response_model_sanitization.py \
|
||||
tests/unit/proxy/test_route_a2a_models.py \
|
||||
tests/unit/proxy/test_route_llm_request.py \
|
||||
tests/unit/proxy/test_route_priority.py \
|
||||
tests/unit/proxy/test_sensitive_route_auth.py \
|
||||
tests/unit/proxy/test_shared_health_check.py \
|
||||
tests/unit/proxy/test_spend_log_cleanup.py \
|
||||
tests/unit/proxy/test_swagger_chat_completions.py \
|
||||
tests/unit/proxy/test_team_member_update.py \
|
||||
tests/unit/proxy/test_team_org_move.py \
|
||||
tests/unit/proxy/test_tools_allowlist_enforcement.py \
|
||||
tests/unit/proxy/test_tracing_endpoints.py \
|
||||
tests/unit/proxy/test_update_llm_router_resilience.py \
|
||||
tests/unit/proxy/test_zerobus_dashboard_config.py
|
||||
|
||||
test-unit-proxy-misc: install-test-deps
|
||||
$(UV_RUN) pytest tests/unit/proxy/agent_endpoints tests/unit/proxy/anthropic_endpoints tests/unit/proxy/common_utils --ignore=tests/unit/proxy/common_utils/test_cache_aware_routing.py --ignore=tests/unit/proxy/common_utils/test_check_batch_cost.py --ignore=tests/unit/proxy/common_utils/test_check_responses_cost.py --ignore=tests/unit/proxy/common_utils/test_proxy_encrypt_decrypt.py --ignore=tests/unit/proxy/common_utils/test_realtime_cache.py tests/unit/proxy/discovery_endpoints tests/unit/proxy/experimental tests/unit/proxy/google_endpoints tests/unit/proxy/health_endpoints tests/unit/proxy/image_endpoints tests/unit/proxy/middleware --ignore=tests/unit/proxy/middleware/test_request_size_limit_middleware.py tests/unit/proxy/openai_files_endpoint tests/unit/proxy/pass_through_endpoints tests/test_litellm/proxy/pass_through_endpoints tests/unit/proxy/prompts tests/unit/proxy/public_endpoints tests/unit/proxy/response_api_endpoints tests/unit/proxy/shutdown tests/unit/proxy/spend_tracking --ignore=tests/unit/proxy/spend_tracking/test_search_api_logging.py tests/unit/proxy/ui_crud_endpoints tests/unit/proxy/vector_store_endpoints $(PROXY_INFRA_ROOT_TESTS) tests/unit/proxy/test_proxy_server_endpoints_and_startup.py tests/unit/proxy/test_proxy_utils_model_creation_and_error_logging.py tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py --ignore=tests/unit/proxy/google_endpoints/test_gemini_agents_endpoints.py --ignore=tests/unit/proxy/google_endpoints/test_google_endpoint_routing.py --ignore=tests/unit/proxy/google_endpoints/test_google_gemini_proxy_request.py --ignore=tests/unit/proxy/public_endpoints/test_blog_posts_endpoint.py --tb=short -vv -n 4 --durations=20
|
||||
$(UV_RUN) pytest tests/unit/proxy/agent_endpoints tests/unit/proxy/anthropic_endpoints tests/unit/proxy/common_utils --ignore=tests/unit/proxy/common_utils/test_cache_aware_routing.py --ignore=tests/unit/proxy/common_utils/test_check_batch_cost.py --ignore=tests/unit/proxy/common_utils/test_check_responses_cost.py --ignore=tests/unit/proxy/common_utils/test_proxy_encrypt_decrypt.py --ignore=tests/unit/proxy/common_utils/test_realtime_cache.py tests/unit/proxy/discovery_endpoints tests/unit/proxy/experimental tests/unit/proxy/google_endpoints tests/unit/proxy/health_endpoints tests/unit/proxy/image_endpoints tests/unit/proxy/middleware --ignore=tests/unit/proxy/middleware/test_request_size_limit_middleware.py tests/unit/proxy/openai_files_endpoint tests/unit/proxy/pass_through_endpoints tests/unit/proxy/prompts tests/unit/proxy/public_endpoints tests/unit/proxy/response_api_endpoints tests/unit/proxy/shutdown tests/unit/proxy/spend_tracking --ignore=tests/unit/proxy/spend_tracking/test_search_api_logging.py tests/unit/proxy/ui_crud_endpoints tests/unit/proxy/vector_store_endpoints tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py --ignore=tests/unit/proxy/google_endpoints/test_gemini_agents_endpoints.py --ignore=tests/unit/proxy/google_endpoints/test_google_endpoint_routing.py --ignore=tests/unit/proxy/google_endpoints/test_google_gemini_proxy_request.py --ignore=tests/unit/proxy/public_endpoints/test_blog_posts_endpoint.py --tb=short -vv -n 4 --durations=20
|
||||
|
||||
test-unit-proxy-root: install-test-deps
|
||||
$(UV_RUN) pytest tests/unit/proxy/test_*.py --ignore=tests/unit/proxy/test_aproxy_startup.py --ignore=tests/unit/proxy/test_credential_slot_registry.py --ignore=tests/unit/proxy/test_custom_callback_input.py --ignore=tests/unit/proxy/test_custom_logger_s3_gcs.py --ignore=tests/unit/proxy/test_custom_tokenizer_bug.py --ignore=tests/unit/proxy/test_db_schema_changes.py --ignore=tests/unit/proxy/test_deprecated_key_grace_period.py --ignore=tests/unit/proxy/test_get_favicon.py --ignore=tests/unit/proxy/test_get_image.py --ignore=tests/unit/proxy/test_prisma_client_backoff_retry.py --ignore=tests/unit/proxy/test_prompt_test_endpoint.py --ignore=tests/unit/proxy/test_proxy_config_unit_test.py --ignore=tests/unit/proxy/test_proxy_custom_auth.py --ignore=tests/unit/proxy/test_proxy_reject_logging.py --ignore=tests/unit/proxy/test_proxy_server.py --ignore=tests/unit/proxy/test_proxy_setting_guardrails.py --ignore=tests/unit/proxy/test_proxy_token_counter.py --ignore=tests/unit/proxy/test_proxy_utils.py --ignore=tests/unit/proxy/test_reducto_ocr_route.py --ignore=tests/unit/proxy/test_response_polling_pre_call_checks.py --ignore=tests/unit/proxy/test_server_root_path.py --ignore=tests/unit/proxy/test_ui_path_detection.py --ignore=tests/unit/proxy/test_unit_test_proxy_hooks.py --ignore=tests/unit/proxy/test_update_spend.py --ignore=tests/unit/proxy/test_zero_cost_model_budget_bypass.py --tb=short -vv -n 4 --durations=20
|
||||
|
||||
test-unit-integrations: install-test-deps
|
||||
$(UV_RUN) pytest tests/unit/integrations --tb=short -vv -n 4 --durations=20
|
||||
|
|
|
|||
25
README.md
25
README.md
|
|
@ -268,6 +268,31 @@ For MCP OAuth, an upstream may advertise dynamic client registration but refuse
|
|||
|
||||
</details>
|
||||
|
||||
<details>
|
||||
<summary><b>Agents</b> - Run Claude Code, Codex, OpenCode or Deep Agents on any model (Python SDK)</summary>
|
||||
|
||||
### Python SDK - Agents
|
||||
|
||||
```python
|
||||
import litellm
|
||||
from litellm import Harness, sandbox
|
||||
|
||||
result = litellm.agent(
|
||||
Harness.CLAUDE_CODE, # or Harness.CODEX, Harness.OPENCODE, Harness.DEEPAGENTS
|
||||
"Find why tests/test_router.py is flaky and fix it.",
|
||||
sandbox=sandbox.local("./repo"),
|
||||
model="litellm_proxy/claude-sonnet-4-5", # a model group on your AI Gateway
|
||||
)
|
||||
|
||||
print(result.text, result.cost, [f.path for f in result.files])
|
||||
```
|
||||
|
||||
Set `LITELLM_PROXY_API_BASE` and `LITELLM_PROXY_API_KEY` and every model call the agent makes goes through your AI Gateway, tagged `harness,claude_code`. Drop the `litellm_proxy/` prefix to call a provider directly. Install `starlette uvicorn` plus the agent's CLI (`claude`, `codex` or `opencode`), or `deepagents langchain-litellm` for Deep Agents.
|
||||
|
||||
[**Docs: Agent Harnesses**](https://docs.litellm.ai/docs/harness)
|
||||
|
||||
</details>
|
||||
|
||||
### Supported Providers ([Website Supported Models](https://models.litellm.ai/) | [Docs](https://docs.litellm.ai/docs/providers))
|
||||
|
||||
| Provider | `/chat/completions` | `/messages` | `/responses` | `/embeddings` | `/image/generations` | `/audio/transcriptions` | `/audio/speech` | `/moderations` | `/batches` | `/rerank` |
|
||||
|
|
|
|||
|
|
@ -616,7 +616,7 @@ class _ENTERPRISE_SecretDetection(CustomGuardrail):
|
|||
data["prompt"] = self.redact_text(prompt, source="prompt")
|
||||
return 1
|
||||
if isinstance(prompt, list):
|
||||
data["prompt"] = [ # mutable-ok: data["prompt"] is a list on the wire
|
||||
data["prompt"] = [
|
||||
self.redact_text(item, source="prompt")
|
||||
if isinstance(item, str) and item
|
||||
else item
|
||||
|
|
|
|||
|
|
@ -87,3 +87,21 @@ def migration_lock(database_url: str) -> Generator[MigrationCoordinator, None, N
|
|||
f"Timed out waiting for another v2 migration resolver after {wait_seconds}s. "
|
||||
f"Check the running migration or increase {MIGRATION_LOCK_TIMEOUT_ENV_VAR}."
|
||||
)
|
||||
|
||||
|
||||
@contextmanager
|
||||
def held_migration_lock(connection: "psycopg.Connection[tuple[object, ...]]") -> Generator[bool, None, None]:
|
||||
"""A session-level, non-blocking hold of the migration coordinator lock on an autocommit
|
||||
connection, for DDL that cannot run inside a transaction (`CREATE INDEX CONCURRENTLY`).
|
||||
Yields whether the lock was acquired; a v2 resolver or another migration job's index build
|
||||
holding it yields False. Released on exit."""
|
||||
from psycopg.rows import class_row
|
||||
|
||||
with connection.cursor(row_factory=class_row(_LockResult)) as cursor:
|
||||
row: Final = cursor.execute("SELECT pg_try_advisory_lock(%s) AS acquired", (MIGRATION_LOCK_KEY,)).fetchone()
|
||||
acquired: Final = row is not None and row.acquired
|
||||
try:
|
||||
yield acquired
|
||||
finally:
|
||||
if acquired:
|
||||
connection.execute("SELECT pg_advisory_unlock(%s)", (MIGRATION_LOCK_KEY,))
|
||||
|
|
|
|||
|
|
@ -1,4 +1,5 @@
|
|||
import hashlib
|
||||
import re
|
||||
import subprocess
|
||||
from collections.abc import Mapping
|
||||
from dataclasses import dataclass
|
||||
|
|
@ -156,3 +157,48 @@ def baseline_current_schema(
|
|||
"review any feature-specific backfill requirements.",
|
||||
len(migrations),
|
||||
)
|
||||
|
||||
|
||||
_LINE_COMMENT_RE: Final = re.compile(r"--[^\n]*")
|
||||
_BLOCK_COMMENT_RE: Final = re.compile(r"/\*.*?\*/", re.DOTALL)
|
||||
_NO_OP_STATEMENT_RE: Final = re.compile(r"^\s*SELECT\s+1\s*$", re.IGNORECASE)
|
||||
|
||||
|
||||
def is_inert_migration(script: str) -> bool:
|
||||
"""Whether a migration file changes nothing: only comments and `SELECT 1`, so
|
||||
applying it can neither repeat nor skip a database change."""
|
||||
stripped: Final = _LINE_COMMENT_RE.sub("", _BLOCK_COMMENT_RE.sub("", script))
|
||||
return all(not part.strip() or _NO_OP_STATEMENT_RE.match(part) for part in stripped.split(";"))
|
||||
|
||||
|
||||
def roll_back_failed_inert_migration(coordinator: MigrationCoordinator, schema: str, migration: Path) -> bool:
|
||||
"""Roll back the failed ledger row of a migration whose file in this build is inert,
|
||||
so `migrate deploy` applies the inert file on its next pass. The row records an
|
||||
earlier build's attempt at SQL this build no longer ships (an index now built by the
|
||||
migration job), so no database change can be repeated or skipped by replaying
|
||||
the empty file. The caller commits this checkpoint before the next Prisma command.
|
||||
"""
|
||||
from psycopg import sql
|
||||
|
||||
if not is_inert_migration(migration.read_text(encoding="utf-8")):
|
||||
return False
|
||||
coordinator.acquire_prisma_lock()
|
||||
records: Final = _migration_records(coordinator.connection, schema, migration)
|
||||
unfinished: Final = tuple(record for record in records if not record.finished)
|
||||
if len(unfinished) != 1:
|
||||
return False
|
||||
result: Final = coordinator.connection.execute(
|
||||
sql.SQL(
|
||||
"UPDATE {} SET rolled_back_at = current_timestamp "
|
||||
"WHERE id = %s AND finished_at IS NULL AND rolled_back_at IS NULL"
|
||||
).format(sql.Identifier(schema, "_prisma_migrations")),
|
||||
(unfinished[0].id,),
|
||||
)
|
||||
if result.rowcount != 1:
|
||||
raise RuntimeError("Could not roll back the failed inert migration history row; rerun the database setup.")
|
||||
logger.info(
|
||||
"Rolled back the failed history row of %s: this build ships it as an inert migration, "
|
||||
"its index is built by the migration job",
|
||||
migration.parent.name,
|
||||
)
|
||||
return True
|
||||
|
|
|
|||
|
|
@ -1,2 +1,6 @@
|
|||
-- CreateIndex
|
||||
CREATE INDEX IF NOT EXISTS "LiteLLM_SpendLogs_api_key_startTime_idx" ON "LiteLLM_SpendLogs"("api_key", "startTime");
|
||||
-- The (api_key, startTime) index on LiteLLM_SpendLogs is built after migrate deploy,
|
||||
-- through litellm_proxy_extras/request_log_indexes.py: concurrently on a plain table and
|
||||
-- per partition on a partitioned one. The migration job builds it; a serving proxy that
|
||||
-- ran the migrations itself builds it in the background once it serves. A migration
|
||||
-- cannot do either without blocking spend-log writes or failing on a partitioned table.
|
||||
SELECT 1;
|
||||
|
|
|
|||
|
|
@ -1,12 +1,6 @@
|
|||
-- CreateIndex (CONCURRENTLY)
|
||||
--
|
||||
-- Disclaimer:
|
||||
-- - CREATE INDEX CONCURRENTLY cannot run inside a transaction. This migration must stay a
|
||||
-- single statement so Prisma Migrate on PostgreSQL can apply it outside a transaction.
|
||||
-- - Builds are slower and use more I/O than a blocking CREATE INDEX; if the build is
|
||||
-- interrupted, Postgres may leave an INVALID index that must be dropped and recreated.
|
||||
-- - Do not edit this file after it has been applied to any database: Prisma checksums
|
||||
-- migrations; add a new migration instead.
|
||||
-- - Requires PostgreSQL that supports CONCURRENTLY with IF NOT EXISTS (use a new migration
|
||||
-- without IF NOT EXISTS if you must support older versions).
|
||||
CREATE INDEX CONCURRENTLY IF NOT EXISTS "LiteLLM_SpendLogs_litellm_call_id_idx" ON "LiteLLM_SpendLogs"("litellm_call_id");
|
||||
-- The litellm_call_id index on LiteLLM_SpendLogs is built after migrate deploy, through
|
||||
-- litellm_proxy_extras/request_log_indexes.py: concurrently on a plain table and per
|
||||
-- partition on a partitioned one. The migration job builds it; a serving proxy that ran
|
||||
-- the migrations itself builds it in the background once it serves. Postgres refuses
|
||||
-- CREATE INDEX CONCURRENTLY on a partitioned parent, so this migration no longer runs it.
|
||||
SELECT 1;
|
||||
|
|
|
|||
448
litellm-proxy-extras/litellm_proxy_extras/request_log_indexes.py
Normal file
448
litellm-proxy-extras/litellm_proxy_extras/request_log_indexes.py
Normal file
|
|
@ -0,0 +1,448 @@
|
|||
"""The request-log indexes built after `prisma migrate deploy` instead of by a migration:
|
||||
by the migration job, or by a serving proxy that ran the migrations itself (in the
|
||||
background, once it serves).
|
||||
|
||||
A migration cannot build them: a plain `CREATE INDEX` blocks spend-log inserts for the
|
||||
whole build, and `CREATE INDEX CONCURRENTLY` is refused on a partitioned parent
|
||||
(db_scripts/partition_spend_logs.sql). `REQUEST_LOG_INDEXES` is the one list to extend;
|
||||
names match what Prisma derives from the `@@index` declarations in schema.prisma, so an
|
||||
index a database already has is recognized and never rebuilt.
|
||||
"""
|
||||
|
||||
import hashlib
|
||||
import random
|
||||
import re
|
||||
import time
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING, Final
|
||||
|
||||
from litellm_proxy_extras._logging import logger
|
||||
from litellm_proxy_extras.migration_lock import held_migration_lock
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import psycopg
|
||||
from psycopg import sql
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RequestLogIndex:
|
||||
"""One index the migration job owns: the table, the exact Prisma index name and the
|
||||
column list as it would be written after `ON <table>`."""
|
||||
|
||||
table: str
|
||||
name: str
|
||||
definition: str
|
||||
|
||||
@property
|
||||
def columns(self) -> tuple[str, ...]:
|
||||
return tuple(re.findall(r'"([^"]+)"', self.definition))
|
||||
|
||||
def partition_index_name(self, partition: str) -> str:
|
||||
"""The child index name for one partition, built the way Postgres names the
|
||||
children of a partitioned index, and kept within the 63 byte identifier limit."""
|
||||
name: Final = f"{partition}_{self.name.removeprefix(f'{self.table}_')}"
|
||||
if len(name.encode()) <= _IDENTIFIER_MAX_BYTES:
|
||||
return name
|
||||
digest: Final = hashlib.sha256(name.encode()).hexdigest()[:_DIGEST_LENGTH]
|
||||
budget: Final = _IDENTIFIER_MAX_BYTES - _DIGEST_LENGTH - 1
|
||||
kept: Final = next(name[:length] for length in range(len(name), 0, -1) if len(name[:length].encode()) <= budget)
|
||||
return f"{kept}_{digest}"
|
||||
|
||||
|
||||
REQUEST_LOG_INDEXES: Final = (
|
||||
RequestLogIndex("LiteLLM_SpendLogs", "LiteLLM_SpendLogs_api_key_startTime_idx", '("api_key", "startTime")'),
|
||||
RequestLogIndex("LiteLLM_SpendLogs", "LiteLLM_SpendLogs_litellm_call_id_idx", '("litellm_call_id")'),
|
||||
)
|
||||
|
||||
_IDENTIFIER_MAX_BYTES: Final = 63
|
||||
_PARENT_LOCK_TIMEOUT: Final = "2s"
|
||||
_PARENT_LOCK_ATTEMPTS: Final = 30
|
||||
_LOCK_HANDOVER_SECONDS: Final = 2.0
|
||||
_DIGEST_LENGTH: Final = 8
|
||||
_CREATE_INDEX_STATEMENT: Final = re.compile(
|
||||
r'^\s*CREATE\s+(?:UNIQUE\s+)?INDEX\s+(?:CONCURRENTLY\s+)?(?:IF\s+NOT\s+EXISTS\s+)?"(?P<index>[^"]+)"\s+ON\b',
|
||||
re.IGNORECASE,
|
||||
)
|
||||
_TABLE_KIND_SQL: Final = "SELECT c.relkind = 'p' AS partitioned FROM pg_class c WHERE c.oid = to_regclass(%s)"
|
||||
_CHILDREN_WITHOUT_THE_INDEX_SQL: Final = (
|
||||
"SELECT child.relname AS name, n.nspname AS schema, child.relkind = 'p' AS partitioned "
|
||||
"FROM pg_inherits i JOIN pg_class child ON child.oid = i.inhrelid "
|
||||
"JOIN pg_namespace n ON n.oid = child.relnamespace "
|
||||
"WHERE i.inhparent = to_regclass(%s) AND NOT EXISTS ("
|
||||
"SELECT 1 FROM pg_inherits attached JOIN pg_index x ON x.indexrelid = attached.inhrelid "
|
||||
"WHERE attached.inhparent = to_regclass(%s) AND x.indrelid = child.oid) "
|
||||
"ORDER BY child.relname"
|
||||
)
|
||||
_EQUIVALENT_INDEXES_SQL: Final = (
|
||||
"SELECT i.relname AS name, x.indisvalid AS valid "
|
||||
"FROM pg_index x JOIN pg_class i ON i.oid = x.indexrelid JOIN pg_am am ON am.oid = i.relam "
|
||||
"WHERE x.indrelid = to_regclass(%s) AND i.relname <> %s AND am.amname = 'btree' AND NOT x.indisunique "
|
||||
"AND x.indexprs IS NULL AND x.indpred IS NULL AND x.indnkeyatts = x.indnatts "
|
||||
"AND NOT EXISTS (SELECT 1 FROM unnest(x.indoption::int2[]) o WHERE o <> 0) "
|
||||
"AND NOT EXISTS (SELECT 1 FROM unnest(x.indclass::oid[]) c JOIN pg_opclass oc ON oc.oid = c WHERE NOT oc.opcdefault) "
|
||||
"AND NOT EXISTS (SELECT 1 FROM unnest(x.indcollation::oid[]) WITH ORDINALITY c(coll, ord) "
|
||||
"JOIN unnest(x.indkey::int2[]) WITH ORDINALITY k(attnum, ord) ON k.ord = c.ord "
|
||||
"JOIN pg_attribute a ON a.attrelid = x.indrelid AND a.attnum = k.attnum "
|
||||
"WHERE c.coll <> 0 AND c.coll <> a.attcollation) "
|
||||
"AND (SELECT array_agg(a.attname::text ORDER BY k.ord) FROM unnest(x.indkey::int2[]) WITH ORDINALITY k(attnum, ord) "
|
||||
"JOIN pg_attribute a ON a.attrelid = x.indrelid AND a.attnum = k.attnum) = %s::text[] "
|
||||
"AND NOT EXISTS (SELECT 1 FROM pg_inherits WHERE inhrelid = x.indexrelid) "
|
||||
"ORDER BY x.indisvalid DESC, i.relname"
|
||||
)
|
||||
_INDEX_STATE_SQL: Final = (
|
||||
'SELECT x.indisvalid AS valid, t.relname AS "table" '
|
||||
"FROM pg_index x JOIN pg_class t ON t.oid = x.indrelid WHERE x.indexrelid = to_regclass(%s)"
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _Relation:
|
||||
name: str
|
||||
schema: str
|
||||
partitioned: bool
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _IndexState:
|
||||
valid: bool
|
||||
table: str
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _EquivalentIndex:
|
||||
name: str
|
||||
valid: bool
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _TableKind:
|
||||
partitioned: bool
|
||||
|
||||
|
||||
def filter_request_log_index_diff(diff_sql: str, indexes: tuple[RequestLogIndex, ...] = REQUEST_LOG_INDEXES) -> str:
|
||||
"""The `prisma migrate diff` script without the statements that create a migration-job-owned
|
||||
index, which the schema declares and the migrations deliberately do not build."""
|
||||
names: Final = frozenset(index.name for index in indexes)
|
||||
statements: Final = diff_sql.split(";")
|
||||
kept: Final = tuple(statement for statement in statements if not _creates_one_of(statement, names))
|
||||
return ";".join(kept) if any(part.strip() for part in kept) else ""
|
||||
|
||||
|
||||
def _creates_one_of(statement: str, names: frozenset[str]) -> bool:
|
||||
match: Final = _CREATE_INDEX_STATEMENT.match(_without_comments(statement))
|
||||
return match is not None and match["index"] in names
|
||||
|
||||
|
||||
def _without_comments(statement: str) -> str:
|
||||
return "\n".join(line for line in statement.splitlines() if not line.lstrip().startswith("--"))
|
||||
|
||||
|
||||
def _connect(database_url: str) -> "psycopg.Connection[tuple[object, ...]]":
|
||||
import psycopg
|
||||
|
||||
return psycopg.connect(database_url, connect_timeout=10, autocommit=True)
|
||||
|
||||
|
||||
def ensure_request_log_indexes(
|
||||
database_url: str,
|
||||
schema: str,
|
||||
indexes: tuple[RequestLogIndex, ...] = REQUEST_LOG_INDEXES,
|
||||
connect: "Callable[[str], psycopg.Connection[tuple[object, ...]]]" = _connect,
|
||||
) -> bool:
|
||||
"""Build every listed index that is missing or invalid. Each build step runs under
|
||||
the migration coordinator lock, held per statement so a resolver booting on another
|
||||
replica gets in between partitions rather than waiting for the whole table. Any
|
||||
failure is logged and left for the next index build; the result says whether
|
||||
every index ended up valid. Never raises."""
|
||||
import psycopg
|
||||
|
||||
try:
|
||||
with connect(database_url) as connection:
|
||||
connection.execute("SET statement_timeout = 0")
|
||||
results: Final = tuple(_ensure_index(connection, schema, index) for index in indexes)
|
||||
except psycopg.Error as exc:
|
||||
logger.warning("Could not build the request-log indexes, leaving them for the next index build: %s", exc)
|
||||
return False
|
||||
if not all(results):
|
||||
logger.warning("Some request-log indexes are not in place yet, leaving them for the next index build")
|
||||
return False
|
||||
logger.info("Request-log indexes are all in place")
|
||||
return True
|
||||
|
||||
|
||||
def _under_migration_lock(connection: "psycopg.Connection[tuple[object, ...]]", step: Callable[[], bool]) -> bool:
|
||||
with held_migration_lock(connection) as held:
|
||||
if not held:
|
||||
logger.info(
|
||||
"Another process holds the migration lock, leaving the request-log indexes to the next index build"
|
||||
)
|
||||
return False
|
||||
return step()
|
||||
|
||||
|
||||
def _ensure_index(connection: "psycopg.Connection[tuple[object, ...]]", schema: str, index: RequestLogIndex) -> bool:
|
||||
from psycopg.rows import class_row
|
||||
|
||||
with connection.cursor(row_factory=class_row(_TableKind)) as cursor:
|
||||
table: Final = cursor.execute(_TABLE_KIND_SQL, (_regclass_name(connection, schema, index.table),)).fetchone()
|
||||
if table is None:
|
||||
logger.info("Table %s does not exist yet, skipping index %s", index.table, index.name)
|
||||
return True
|
||||
if table.partitioned:
|
||||
return build_index_on_partitioned_table(connection, schema, index)
|
||||
return _build_leaf_index(connection, schema, index.table, index.name, index)
|
||||
|
||||
|
||||
def _regclass_name(connection: "psycopg.Connection[tuple[object, ...]]", schema: str, name: str) -> str:
|
||||
from psycopg import sql
|
||||
|
||||
return sql.Identifier(schema, name).as_string(connection)
|
||||
|
||||
|
||||
def _create_index_statement(
|
||||
connection: "psycopg.Connection[tuple[object, ...]]", prefix: "sql.Composed", definition: str
|
||||
) -> bytes:
|
||||
return (prefix.as_string(connection) + definition).encode()
|
||||
|
||||
|
||||
def _index_state(connection: "psycopg.Connection[tuple[object, ...]]", schema: str, index: str) -> "_IndexState | None":
|
||||
from psycopg.rows import class_row
|
||||
|
||||
with connection.cursor(row_factory=class_row(_IndexState)) as cursor:
|
||||
return cursor.execute(_INDEX_STATE_SQL, (_regclass_name(connection, schema, index),)).fetchone()
|
||||
|
||||
|
||||
def _equivalent_indexes(
|
||||
connection: "psycopg.Connection[tuple[object, ...]]",
|
||||
schema: str,
|
||||
table: str,
|
||||
name: str,
|
||||
index: RequestLogIndex,
|
||||
) -> tuple[_EquivalentIndex, ...]:
|
||||
"""The indexes on `table` other than `name` with the same definition: default btree
|
||||
over the same columns in the same order, no expression, predicate, DESC or custom
|
||||
opclass or collation, and not attached under a partitioned index. Valid ones first."""
|
||||
from psycopg.rows import class_row
|
||||
|
||||
with connection.cursor(row_factory=class_row(_EquivalentIndex)) as cursor:
|
||||
return tuple(
|
||||
cursor.execute(
|
||||
_EQUIVALENT_INDEXES_SQL, (_regclass_name(connection, schema, table), name, list(index.columns))
|
||||
).fetchall()
|
||||
)
|
||||
|
||||
|
||||
def _adopt_equivalent_index(
|
||||
connection: "psycopg.Connection[tuple[object, ...]]",
|
||||
schema: str,
|
||||
table: str,
|
||||
name: str,
|
||||
index: RequestLogIndex,
|
||||
) -> bool:
|
||||
"""Rename a valid index of the same definition under another name (an operator's
|
||||
hand-built copy, say) to the name this code expects, instead of building a second
|
||||
one. RENAME on an index is a catalog change that lets writes through."""
|
||||
from psycopg import sql
|
||||
|
||||
equivalent: Final = next(
|
||||
(found for found in _equivalent_indexes(connection, schema, table, name, index) if found.valid), None
|
||||
)
|
||||
if equivalent is None:
|
||||
return False
|
||||
logger.info(
|
||||
"Renaming the equivalent index %s on %s to %s instead of building a second one", equivalent.name, table, name
|
||||
)
|
||||
connection.execute(
|
||||
sql.SQL("ALTER INDEX {} RENAME TO {}").format(sql.Identifier(schema, equivalent.name), sql.Identifier(name))
|
||||
)
|
||||
return True
|
||||
|
||||
|
||||
def _report_second_copies(
|
||||
connection: "psycopg.Connection[tuple[object, ...]]",
|
||||
schema: str,
|
||||
table: str,
|
||||
name: str,
|
||||
index: RequestLogIndex,
|
||||
concurrently: bool,
|
||||
) -> None:
|
||||
"""Log every other index of the same definition with the statement that removes it.
|
||||
Dropping is the operator's call: a second copy costs writes and disk, never results."""
|
||||
from psycopg import sql
|
||||
|
||||
drop: Final = "DROP INDEX CONCURRENTLY" if concurrently else "DROP INDEX"
|
||||
for copy in _equivalent_indexes(connection, schema, table, name, index):
|
||||
logger.warning(
|
||||
"Index %s on %s is a second copy of %s and only costs writes and disk; remove it with: %s %s",
|
||||
copy.name,
|
||||
table,
|
||||
name,
|
||||
drop,
|
||||
sql.Identifier(schema, copy.name).as_string(connection),
|
||||
)
|
||||
|
||||
|
||||
def _children_without_the_index(
|
||||
connection: "psycopg.Connection[tuple[object, ...]]", schema: str, table: str, index: str
|
||||
) -> tuple[_Relation, ...]:
|
||||
from psycopg.rows import class_row
|
||||
|
||||
with connection.cursor(row_factory=class_row(_Relation)) as cursor:
|
||||
return tuple(
|
||||
cursor.execute(
|
||||
_CHILDREN_WITHOUT_THE_INDEX_SQL,
|
||||
(_regclass_name(connection, schema, table), _regclass_name(connection, schema, index)),
|
||||
).fetchall()
|
||||
)
|
||||
|
||||
|
||||
def _build_leaf_index(
|
||||
connection: "psycopg.Connection[tuple[object, ...]]",
|
||||
schema: str,
|
||||
table: str,
|
||||
name: str,
|
||||
index: RequestLogIndex,
|
||||
) -> bool:
|
||||
"""Build one plain table's or partition's index with CONCURRENTLY so writes keep
|
||||
flowing. The catalog is read under the migration lock, so a replica that saw an
|
||||
invalid index before the lock finds the valid one another replica just built and
|
||||
leaves it. An invalid index left by an interrupted build is dropped and rebuilt; a
|
||||
valid index of the same definition under another name is renamed rather than
|
||||
duplicated; an index of that name on another table is a collision this code will
|
||||
not touch."""
|
||||
from psycopg import sql
|
||||
|
||||
def build() -> bool:
|
||||
existing: Final = _index_state(connection, schema, name)
|
||||
if existing is not None and existing.table != table:
|
||||
logger.warning(
|
||||
"Index %s already exists on %s rather than %s, leaving it alone", name, existing.table, table
|
||||
)
|
||||
return False
|
||||
if existing is not None and existing.valid:
|
||||
return True
|
||||
if existing is not None:
|
||||
logger.info("Dropping the invalid index %s left by an interrupted build on %s", name, table)
|
||||
connection.execute(sql.SQL("DROP INDEX CONCURRENTLY {}").format(sql.Identifier(schema, name)))
|
||||
elif _adopt_equivalent_index(connection, schema, table, name, index):
|
||||
return True
|
||||
logger.info("Building index %s on %s concurrently", name, table)
|
||||
prefix: Final = sql.SQL("CREATE INDEX CONCURRENTLY IF NOT EXISTS {} ON {} ").format(
|
||||
sql.Identifier(name), sql.Identifier(schema, table)
|
||||
)
|
||||
connection.execute(_create_index_statement(connection, prefix, index.definition))
|
||||
built: Final = _index_state(connection, schema, name)
|
||||
return built is not None and built.valid
|
||||
|
||||
current: Final = _index_state(connection, schema, name)
|
||||
if current is None or not current.valid or current.table != table:
|
||||
if not _under_migration_lock(connection, build):
|
||||
return False
|
||||
time.sleep(_LOCK_HANDOVER_SECONDS)
|
||||
_report_second_copies(connection, schema, table, name, index, concurrently=True)
|
||||
return True
|
||||
|
||||
|
||||
def build_index_on_partitioned_table(
|
||||
connection: "psycopg.Connection[tuple[object, ...]]",
|
||||
schema: str,
|
||||
index: RequestLogIndex,
|
||||
table: "str | None" = None,
|
||||
name: "str | None" = None,
|
||||
) -> bool:
|
||||
"""Build the index the way Postgres allows on a partitioned parent: a metadata-only
|
||||
parent index ON ONLY the parent, one CONCURRENTLY build per partition, and ATTACH
|
||||
PARTITION for each child. Partitions that are themselves partitioned get the same
|
||||
treatment one level down. Every step checks the catalog before acting, so an
|
||||
interrupted run resumes where it stopped and a second run finds nothing to do; a
|
||||
parent or child index of the same definition under another name is renamed and
|
||||
used rather than duplicated. The connection must be in autocommit mode. True when
|
||||
the parent index ends up valid."""
|
||||
|
||||
parent_table: Final = index.table if table is None else table
|
||||
parent_index: Final = index.name if name is None else name
|
||||
existing: Final = _index_state(connection, schema, parent_index)
|
||||
if existing is not None and existing.table != parent_table:
|
||||
logger.warning(
|
||||
"Index %s already exists on %s rather than %s, leaving it alone", parent_index, existing.table, parent_table
|
||||
)
|
||||
return False
|
||||
if existing is None and not _under_migration_lock(
|
||||
connection,
|
||||
lambda: (
|
||||
_adopt_equivalent_index(connection, schema, parent_table, parent_index, index)
|
||||
or _create_parent_index(connection, schema, parent_index, parent_table, index)
|
||||
),
|
||||
):
|
||||
return False
|
||||
children: Final = _children_without_the_index(connection, schema, parent_table, parent_index)
|
||||
if not all(_attach_child_index(connection, schema, parent_index, child, index) for child in children):
|
||||
return False
|
||||
final: Final = _index_state(connection, schema, parent_index)
|
||||
if final is None or not final.valid:
|
||||
return False
|
||||
_report_second_copies(connection, schema, parent_table, parent_index, index, concurrently=False)
|
||||
return True
|
||||
|
||||
|
||||
def _create_parent_index(
|
||||
connection: "psycopg.Connection[tuple[object, ...]]",
|
||||
schema: str,
|
||||
name: str,
|
||||
table: str,
|
||||
index: RequestLogIndex,
|
||||
) -> bool:
|
||||
"""Create the metadata-only parent index. Postgres takes a SHARE lock on the
|
||||
parent for that statement, so it waits for in-flight writes and queues new ones
|
||||
behind it; a short lock_timeout with retries keeps every such pause bounded."""
|
||||
import psycopg
|
||||
from psycopg import sql
|
||||
|
||||
prefix: Final = sql.SQL("CREATE INDEX IF NOT EXISTS {} ON ONLY {} ").format(
|
||||
sql.Identifier(name), sql.Identifier(schema, table)
|
||||
)
|
||||
statement: Final = _create_index_statement(connection, prefix, index.definition)
|
||||
connection.execute(sql.SQL("SET lock_timeout = {}").format(sql.Literal(_PARENT_LOCK_TIMEOUT)))
|
||||
try:
|
||||
for _ in range(_PARENT_LOCK_ATTEMPTS):
|
||||
try:
|
||||
connection.execute(statement)
|
||||
return True
|
||||
except psycopg.errors.LockNotAvailable:
|
||||
logger.info("Waiting for in-flight writes to %s before creating the parent index %s", table, name)
|
||||
time.sleep(random.uniform(0.1, 0.5))
|
||||
finally:
|
||||
connection.execute("SET lock_timeout = 0")
|
||||
logger.warning("Could not get the parent lock on %s to create %s, leaving it for the next index build", table, name)
|
||||
return False
|
||||
|
||||
|
||||
def _attach_child_index(
|
||||
connection: "psycopg.Connection[tuple[object, ...]]",
|
||||
schema: str,
|
||||
parent_index: str,
|
||||
child: _Relation,
|
||||
index: RequestLogIndex,
|
||||
) -> bool:
|
||||
from psycopg import sql
|
||||
|
||||
child_index: Final = index.partition_index_name(child.name)
|
||||
built: Final = (
|
||||
build_index_on_partitioned_table(connection, child.schema, index, child.name, child_index)
|
||||
if child.partitioned
|
||||
else _build_leaf_index(connection, child.schema, child.name, child_index, index)
|
||||
)
|
||||
if not built:
|
||||
return False
|
||||
|
||||
def attach() -> bool:
|
||||
connection.execute(
|
||||
sql.SQL("ALTER INDEX {} ATTACH PARTITION {}").format(
|
||||
sql.Identifier(schema, parent_index), sql.Identifier(child.schema, child_index)
|
||||
)
|
||||
)
|
||||
logger.info("Attached index %s on partition %s to %s", child_index, child.name, parent_index)
|
||||
return True
|
||||
|
||||
return _under_migration_lock(connection, attach)
|
||||
|
|
@ -5,6 +5,7 @@ import re
|
|||
import shutil
|
||||
import subprocess
|
||||
import tempfile
|
||||
import threading
|
||||
import time
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass, replace
|
||||
|
|
@ -13,6 +14,7 @@ from typing import TYPE_CHECKING, Final, Optional
|
|||
|
||||
from litellm_proxy_extras import prisma_toolchain
|
||||
from litellm_proxy_extras._logging import logger
|
||||
from litellm_proxy_extras.migration_lock import held_migration_lock
|
||||
from litellm_proxy_extras.prisma_toolchain import (
|
||||
PRISMA_COMMAND_TIMEOUT_ENV_VAR,
|
||||
PRISMA_MIGRATE_DEPLOY_TIMEOUT_ENV_VAR,
|
||||
|
|
@ -24,6 +26,7 @@ from litellm_proxy_extras.replica_identity import (
|
|||
REPLICA_IDENTITY_FULL_ENV_VAR,
|
||||
apply_replica_identity_full,
|
||||
)
|
||||
from litellm_proxy_extras.request_log_indexes import ensure_request_log_indexes, filter_request_log_index_diff
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import psycopg
|
||||
|
|
@ -433,6 +436,21 @@ class ProxyExtrasDBManager:
|
|||
return True
|
||||
return False
|
||||
|
||||
@staticmethod
|
||||
def _filter_migration_job_owned_drift(diff_sql: str, partitioned: bool | None = None) -> str:
|
||||
"""The drift script without the indexes the migration job builds (the schema
|
||||
declares them, the migrations deliberately do not) and, when LiteLLM_SpendLogs
|
||||
is partitioned, without its primary-key rewrite and partitioning artifacts."""
|
||||
without_indexes: Final = filter_request_log_index_diff(diff_sql)
|
||||
is_partitioned: Final = ProxyExtrasDBManager.spend_logs_is_partitioned() if partitioned is None else partitioned
|
||||
if not is_partitioned:
|
||||
return without_indexes
|
||||
logger.info(
|
||||
"LiteLLM_SpendLogs is partitioned; removed its primary-key "
|
||||
"rewrite and partitioning artifacts from the drift script"
|
||||
)
|
||||
return filter_partitioned_spend_logs_diff(without_indexes)
|
||||
|
||||
@staticmethod
|
||||
def _resolve_all_migrations(
|
||||
migrations_dir: str, schema_path: str, mark_all_applied: bool = True
|
||||
|
|
@ -513,21 +531,14 @@ class ProxyExtrasDBManager:
|
|||
return
|
||||
logger.info(f"Migration diff created at {diff_sql_path}")
|
||||
|
||||
if ProxyExtrasDBManager.spend_logs_is_partitioned():
|
||||
filtered_sql = filter_partitioned_spend_logs_diff(
|
||||
diff_sql_path.read_text()
|
||||
)
|
||||
diff_sql_path.write_text(filtered_sql)
|
||||
logger.info(
|
||||
"LiteLLM_SpendLogs is partitioned; removed its primary-key "
|
||||
"rewrite and partitioning artifacts from the drift script"
|
||||
)
|
||||
if not filtered_sql.strip():
|
||||
logger.info("Drift script is empty after filtering; nothing to apply")
|
||||
if not mark_all_applied:
|
||||
return
|
||||
ProxyExtrasDBManager._mark_migrations_applied(migrations_dir)
|
||||
filtered_sql: Final = ProxyExtrasDBManager._filter_migration_job_owned_drift(diff_sql_path.read_text())
|
||||
diff_sql_path.write_text(filtered_sql)
|
||||
if not filtered_sql.strip():
|
||||
logger.info("Drift script is empty after filtering; nothing to apply")
|
||||
if not mark_all_applied:
|
||||
return
|
||||
ProxyExtrasDBManager._mark_migrations_applied(migrations_dir)
|
||||
return
|
||||
|
||||
# 2. Run prisma db execute to apply the migration
|
||||
applied_ok = False
|
||||
|
|
@ -800,7 +811,7 @@ class ProxyExtrasDBManager:
|
|||
conn.execute(statement)
|
||||
except psycopg.Error as e:
|
||||
logger.warning(
|
||||
"Could not repair invalid index %s.%s, will retry on the next startup. "
|
||||
"Could not repair invalid index %s.%s, will retry on the next database setup run. "
|
||||
"If this keeps happening, run `%s` by hand as the index owner. Error: %s",
|
||||
index.schema,
|
||||
index.name,
|
||||
|
|
@ -811,16 +822,21 @@ class ProxyExtrasDBManager:
|
|||
logger.info("%s invalid index %s.%s", action, index.schema, index.name)
|
||||
|
||||
@staticmethod
|
||||
def repair_invalid_indexes(lock_timeout: str = "30s") -> bool:
|
||||
def repair_invalid_indexes(
|
||||
lock_timeout: str = "30s",
|
||||
repair: "Callable[[psycopg.Connection[tuple[str, str, str]], _InvalidIndex], None] | None" = None,
|
||||
) -> bool:
|
||||
"""Rebuild LiteLLM indexes an interrupted CREATE INDEX CONCURRENTLY left
|
||||
INVALID (a migration deadlock between replicas is the usual cause; the
|
||||
retried migration skips them because of IF NOT EXISTS). Never raises:
|
||||
returns True when no invalid index remains, False when the repair was
|
||||
skipped or failed and will be retried on the next startup. Looks in the
|
||||
skipped or failed and will be retried on the next database setup run. Looks in the
|
||||
schema DATABASE_URL names, the only URL Prisma migrates through, but
|
||||
connects over DIRECT_URL when set: the session settings, the advisory
|
||||
lock and REINDEX CONCURRENTLY all need one server session, which a
|
||||
transaction pooler does not give."""
|
||||
transaction pooler does not give. Each rebuild holds the migration
|
||||
coordinator lock on its own, like the migration job's index build, so a resolver
|
||||
booting on another replica waits for one index at most."""
|
||||
prisma_url: Final = os.getenv("DATABASE_URL")
|
||||
if not prisma_url:
|
||||
return False
|
||||
|
|
@ -856,20 +872,53 @@ class ProxyExtrasDBManager:
|
|||
if lock_row is None or not lock_row[0]:
|
||||
logger.info("Another replica is already rebuilding the invalid indexes, skipping")
|
||||
return False
|
||||
for index in ProxyExtrasDBManager._invalid_litellm_indexes(conn, schema):
|
||||
ProxyExtrasDBManager._repair_index(conn, index)
|
||||
repair_one: Final = repair or ProxyExtrasDBManager._repair_index
|
||||
repaired: Final = all(
|
||||
ProxyExtrasDBManager._repair_under_migration_lock(conn, schema, index, repair_one)
|
||||
for index in found
|
||||
)
|
||||
if not repaired:
|
||||
return False
|
||||
remaining: Final = ProxyExtrasDBManager._invalid_litellm_indexes(conn, schema)
|
||||
except psycopg.Error as e:
|
||||
logger.warning("Could not check for invalid indexes, will retry on the next startup. Error: %s", e)
|
||||
logger.warning(
|
||||
"Could not check for invalid indexes, will retry on the next database setup run. Error: %s", e
|
||||
)
|
||||
return False
|
||||
return not remaining
|
||||
|
||||
@staticmethod
|
||||
def _repair_under_migration_lock(
|
||||
conn: "psycopg.Connection[tuple[str, str, str]]",
|
||||
schema: str,
|
||||
index: _InvalidIndex,
|
||||
repair: "Callable[[psycopg.Connection[tuple[str, str, str]], _InvalidIndex], None]",
|
||||
) -> bool:
|
||||
"""Rebuild one index under the migration coordinator lock, skipping it when a
|
||||
migration job finished or dropped it in the meantime. False when another process
|
||||
holds the lock, so the check waits for the next database setup run."""
|
||||
with held_migration_lock(conn) as held:
|
||||
if not held:
|
||||
logger.info(
|
||||
"Another process is building indexes under the migration lock, leaving the "
|
||||
"invalid index check to the next database setup run"
|
||||
)
|
||||
return False
|
||||
still_invalid: Final = ProxyExtrasDBManager._invalid_litellm_indexes(conn, schema)
|
||||
if any(found.schema == index.schema and found.name == index.name for found in still_invalid):
|
||||
repair(conn, index)
|
||||
return True
|
||||
|
||||
@staticmethod
|
||||
def _setup_database_v2(use_migrate: bool) -> bool:
|
||||
if not use_migrate:
|
||||
return ProxyExtrasDBManager._run_database_v2(False)
|
||||
from litellm_proxy_extras.migration_lock import migration_environment, migration_lock
|
||||
from litellm_proxy_extras.migration_recovery import baseline_current_schema, recover_completed_migration
|
||||
from litellm_proxy_extras.migration_recovery import (
|
||||
baseline_current_schema,
|
||||
recover_completed_migration,
|
||||
roll_back_failed_inert_migration,
|
||||
)
|
||||
|
||||
database_url: Final = os.environ.get("DATABASE_URL")
|
||||
if not database_url:
|
||||
|
|
@ -884,7 +933,9 @@ class ProxyExtrasDBManager:
|
|||
if not migration.is_file():
|
||||
return False
|
||||
with migration_lock(lock_url) as coordinator:
|
||||
return recover_completed_migration(coordinator, schema, migration)
|
||||
return recover_completed_migration(coordinator, schema, migration) or roll_back_failed_inert_migration(
|
||||
coordinator, schema, migration
|
||||
)
|
||||
|
||||
def baseline_existing(migrations_dir: str) -> None:
|
||||
with migration_lock(lock_url) as coordinator:
|
||||
|
|
@ -1177,13 +1228,16 @@ class ProxyExtrasDBManager:
|
|||
)
|
||||
|
||||
@staticmethod
|
||||
def setup_database(
|
||||
use_migrate: bool = False, use_v2_resolver: bool = False
|
||||
) -> bool:
|
||||
def setup_database(use_migrate: bool = False, use_v2_resolver: bool = False) -> bool:
|
||||
"""
|
||||
Set up the database using either prisma migrate or prisma db push
|
||||
Uses migrations from litellm-proxy-extras package
|
||||
|
||||
The request-log indexes in `REQUEST_LOG_INDEXES` are not built here: the
|
||||
migration job builds them through `run_migration_job`, and a serving proxy that
|
||||
ran the migrations itself starts them through `start_request_log_index_build`
|
||||
once it is ready to serve.
|
||||
|
||||
Args:
|
||||
use_migrate: Whether to use prisma migrate instead of db push
|
||||
use_v2_resolver: Opt into the v2 migration resolver (safer during
|
||||
|
|
@ -1200,10 +1254,48 @@ class ProxyExtrasDBManager:
|
|||
migrated = ProxyExtrasDBManager._run_migrations(
|
||||
use_migrate=use_migrate, use_v2_resolver=use_v2_resolver
|
||||
)
|
||||
if migrated:
|
||||
ProxyExtrasDBManager.repair_invalid_indexes()
|
||||
ProxyExtrasDBManager.apply_replica_identity_full_if_requested()
|
||||
return migrated
|
||||
if not migrated:
|
||||
return False
|
||||
ProxyExtrasDBManager.repair_invalid_indexes()
|
||||
ProxyExtrasDBManager.apply_replica_identity_full_if_requested()
|
||||
return True
|
||||
|
||||
@staticmethod
|
||||
def build_request_log_indexes(build: Callable[[str, str], bool] = ensure_request_log_indexes) -> bool:
|
||||
"""Build the indexes in `REQUEST_LOG_INDEXES` on the writer, in the schema the
|
||||
migrations target. Idempotent and never raises; False when an index is still
|
||||
missing or invalid, so the migration job reports it and gets rerun instead of
|
||||
leaving the table unindexed until the next deploy."""
|
||||
database_url: Final = os.environ.get("DATABASE_URL")
|
||||
if not database_url:
|
||||
return True
|
||||
direct_url: Final = ProxyExtrasDBManager._strip_prisma_query_params(
|
||||
os.environ.get("DIRECT_URL") or database_url
|
||||
)
|
||||
schema: Final = ProxyExtrasDBManager._prisma_schema_param(database_url) or "public"
|
||||
return build(direct_url, schema)
|
||||
|
||||
@staticmethod
|
||||
def run_migration_job(
|
||||
use_migrate: bool = False,
|
||||
use_v2_resolver: bool = False,
|
||||
setup: Callable[[bool, bool], bool] = setup_database,
|
||||
build: Callable[[], bool] = build_request_log_indexes,
|
||||
) -> bool:
|
||||
"""The migration job's whole run: `setup_database`, then the request-log indexes,
|
||||
built synchronously so the job exits only once they are in place. False when the
|
||||
migrations failed or an index could not be built, so the Job is rerun."""
|
||||
return setup(use_migrate, use_v2_resolver) and build()
|
||||
|
||||
@staticmethod
|
||||
def start_request_log_index_build(build: Callable[[], bool] = build_request_log_indexes) -> threading.Thread:
|
||||
"""A serving proxy that ran the migrations itself (schema updates not disabled)
|
||||
builds the request-log indexes on a daemon thread, so a long build never delays
|
||||
readiness. A build that could not finish is logged and picked up by the next boot
|
||||
or the migration job."""
|
||||
thread: Final = threading.Thread(target=build, name="litellm-request-log-indexes", daemon=True)
|
||||
thread.start()
|
||||
return thread
|
||||
|
||||
@staticmethod
|
||||
def _run_migrations(use_migrate: bool, use_v2_resolver: bool) -> bool:
|
||||
|
|
@ -1247,15 +1339,16 @@ class ProxyExtrasDBManager:
|
|||
logger.info("✅ Post-migration sanity check completed")
|
||||
return True
|
||||
except subprocess.CalledProcessError as e:
|
||||
logger.info(f"prisma db error: {e.stderr}, e: {e.stdout}")
|
||||
if "P3009" in e.stderr:
|
||||
stderr: Final = str(e.stderr or "")
|
||||
logger.info(f"prisma db error: {stderr}, e: {e.stdout}")
|
||||
if "P3009" in stderr:
|
||||
# Extract the failed migration name from the error message
|
||||
migration_match = re.search(
|
||||
r"`(\d+_.*)` migration", e.stderr
|
||||
r"`(\d+_.*)` migration", stderr
|
||||
)
|
||||
if migration_match:
|
||||
failed_migration = migration_match.group(1)
|
||||
if ProxyExtrasDBManager._is_idempotent_error(e.stderr):
|
||||
if ProxyExtrasDBManager._is_idempotent_error(stderr):
|
||||
logger.info(
|
||||
f"Migration {failed_migration} failed due to idempotent error (e.g., column already exists), resolving as applied"
|
||||
)
|
||||
|
|
@ -1311,8 +1404,8 @@ class ProxyExtrasDBManager:
|
|||
f"✅ Migration {failed_migration} marked as rolled back... retrying"
|
||||
)
|
||||
elif (
|
||||
"P3005" in e.stderr
|
||||
and "database schema is not empty" in e.stderr
|
||||
"P3005" in stderr
|
||||
and "database schema is not empty" in stderr
|
||||
):
|
||||
logger.info(
|
||||
"Database schema is not empty, creating baseline migration. In read-only file system, please set an environment variable `LITELLM_MIGRATION_DIR` to a writable directory to enable migrations. Learn more - https://docs.litellm.ai/docs/proxy/prod#read-only-file-system"
|
||||
|
|
@ -1326,13 +1419,13 @@ class ProxyExtrasDBManager:
|
|||
)
|
||||
logger.info("✅ All migrations resolved.")
|
||||
return True
|
||||
elif "P3018" in e.stderr:
|
||||
elif "P3018" in stderr:
|
||||
# Check if this is a permission error or idempotent error
|
||||
if ProxyExtrasDBManager._is_permission_error(e.stderr):
|
||||
if ProxyExtrasDBManager._is_permission_error(stderr):
|
||||
# Permission errors should NOT be marked as applied
|
||||
# Extract migration name for logging
|
||||
migration_match = re.search(
|
||||
r"Migration name: (\d+_.*)", e.stderr
|
||||
r"Migration name: (\d+_.*)", stderr
|
||||
)
|
||||
migration_name = (
|
||||
migration_match.group(1)
|
||||
|
|
@ -1342,7 +1435,7 @@ class ProxyExtrasDBManager:
|
|||
|
||||
logger.error(
|
||||
f"❌ Migration {migration_name} failed due to insufficient permissions. "
|
||||
f"Please check database user privileges. Error: {e.stderr}"
|
||||
f"Please check database user privileges. Error: {stderr}"
|
||||
)
|
||||
|
||||
# Mark as rolled back and exit with error
|
||||
|
|
@ -1365,7 +1458,7 @@ class ProxyExtrasDBManager:
|
|||
f"was NOT applied. Please grant necessary database permissions and retry."
|
||||
) from e
|
||||
|
||||
elif ProxyExtrasDBManager._is_idempotent_error(e.stderr):
|
||||
elif ProxyExtrasDBManager._is_idempotent_error(stderr):
|
||||
# Idempotent errors mean the migration has effectively been applied
|
||||
logger.info(
|
||||
"Migration failed due to idempotent error (e.g., column already exists), "
|
||||
|
|
@ -1373,7 +1466,7 @@ class ProxyExtrasDBManager:
|
|||
)
|
||||
# Extract the migration name from the error message
|
||||
migration_match = re.search(
|
||||
r"Migration name: (\d+_.*)", e.stderr
|
||||
r"Migration name: (\d+_.*)", stderr
|
||||
)
|
||||
if migration_match:
|
||||
migration_name = migration_match.group(1)
|
||||
|
|
@ -1422,7 +1515,7 @@ class ProxyExtrasDBManager:
|
|||
logger.warning(
|
||||
f"P3018 error encountered but could not classify "
|
||||
f"as permission or idempotent error. "
|
||||
f"Error: {e.stderr}"
|
||||
f"Error: {stderr}"
|
||||
)
|
||||
raise
|
||||
else:
|
||||
|
|
|
|||
24
litellm-rust/Cargo.lock
generated
24
litellm-rust/Cargo.lock
generated
|
|
@ -4086,9 +4086,11 @@ dependencies = [
|
|||
"litellm-secrets",
|
||||
"litellm-secrets-aws",
|
||||
"litellm-secrets-types",
|
||||
"litellm-storage-clickhouse",
|
||||
"litellm-token-counter",
|
||||
"litellm-traces",
|
||||
"litellm-tracing",
|
||||
"prost",
|
||||
"pyo3",
|
||||
"pyo3-async-runtimes",
|
||||
"qdrant-client",
|
||||
|
|
@ -4288,6 +4290,20 @@ dependencies = [
|
|||
"veil",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "litellm-storage-clickhouse"
|
||||
version = "0.1.0"
|
||||
dependencies = [
|
||||
"flate2",
|
||||
"litellm-http",
|
||||
"rstest",
|
||||
"serde",
|
||||
"serde_json",
|
||||
"thiserror 2.0.19",
|
||||
"tokio",
|
||||
"url",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "litellm-testkit"
|
||||
version = "0.1.0"
|
||||
|
|
@ -4369,19 +4385,22 @@ name = "litellm-traces"
|
|||
version = "0.1.0"
|
||||
dependencies = [
|
||||
"base64 0.22.1",
|
||||
"criterion",
|
||||
"flate2",
|
||||
"litellm-http",
|
||||
"litellm-storage-clickhouse",
|
||||
"opentelemetry-proto",
|
||||
"prost",
|
||||
"rstest",
|
||||
"serde",
|
||||
"serde_json",
|
||||
"sha2 0.10.9",
|
||||
"strum",
|
||||
"testcontainers-modules",
|
||||
"thiserror 2.0.19",
|
||||
"time",
|
||||
"tokio",
|
||||
"url",
|
||||
"wiremock",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
|
@ -4804,6 +4823,7 @@ dependencies = [
|
|||
"js-sys",
|
||||
"pin-project-lite",
|
||||
"thiserror 2.0.19",
|
||||
"tracing",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
|
@ -4818,6 +4838,8 @@ dependencies = [
|
|||
"opentelemetry_sdk 0.33.0",
|
||||
"prost",
|
||||
"serde",
|
||||
"tonic",
|
||||
"tonic-prost",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
|
|
|||
|
|
@ -13,6 +13,7 @@ litellm-config = { path = "crates/config" }
|
|||
litellm-router = { path = "crates/router" }
|
||||
litellm-tracing = { path = "crates/tracing" }
|
||||
litellm-traces = { path = "crates/traces" }
|
||||
litellm-storage-clickhouse = { path = "crates/storage-clickhouse" }
|
||||
litellm-core = { path = "crates/core" }
|
||||
litellm-gateway-mcp = { path = "crates/gateway-mcp" }
|
||||
litellm-gateway = { path = "crates/gateway" }
|
||||
|
|
@ -81,6 +82,7 @@ reqwest = { version = "0.12", default-features = false, features = ["json", "mul
|
|||
qdrant-client = { version = "1.19.0", default-features = false }
|
||||
uuid = { version = "1", features = ["v4"] }
|
||||
rstest = "0.26.1"
|
||||
wiremock = "0.6.5"
|
||||
rstest_reuse = "0.7.0"
|
||||
rustls = { version = "0.23", default-features = false, features = ["ring", "std", "tls12"] }
|
||||
rustify = "=0.7.0"
|
||||
|
|
@ -115,6 +117,8 @@ time = { version = "0.3.53", features = ["parsing"] }
|
|||
criterion = "0.8.2"
|
||||
fancy-regex = "0.19.2"
|
||||
veil = "0.3.0"
|
||||
prost = "0.14.4"
|
||||
opentelemetry-proto = "0.33"
|
||||
|
||||
[profile.release]
|
||||
opt-level = 3
|
||||
|
|
|
|||
|
|
@ -26,4 +26,4 @@ litellm-cache-testing.workspace = true
|
|||
rstest.workspace = true
|
||||
serde_json.workspace = true
|
||||
tokio = { workspace = true, features = ["macros", "rt-multi-thread"] }
|
||||
wiremock = "0.6.5"
|
||||
wiremock.workspace = true
|
||||
|
|
|
|||
|
|
@ -21,4 +21,4 @@ litellm-cache-testing.workspace = true
|
|||
rstest.workspace = true
|
||||
serde_json.workspace = true
|
||||
tokio.workspace = true
|
||||
wiremock = "0.6.5"
|
||||
wiremock.workspace = true
|
||||
|
|
|
|||
|
|
@ -21,4 +21,4 @@ redis = "1.7.0"
|
|||
redis-test = "1.0.4"
|
||||
rstest.workspace = true
|
||||
tokio.workspace = true
|
||||
wiremock = "0.6.5"
|
||||
wiremock.workspace = true
|
||||
|
|
|
|||
|
|
@ -23,6 +23,6 @@ tokio.workspace = true
|
|||
litellm-http = { workspace = true, features = ["test-support"] }
|
||||
litellm-cache-testing.workspace = true
|
||||
rstest.workspace = true
|
||||
wiremock = "0.6.5"
|
||||
wiremock.workspace = true
|
||||
serde_json.workspace = true
|
||||
tokio = { workspace = true, features = ["macros", "rt-multi-thread"] }
|
||||
|
|
|
|||
|
|
@ -47,4 +47,4 @@ litellm-host-native.workspace = true
|
|||
litellm-llms = { workspace = true, features = ["test-support"] }
|
||||
rstest.workspace = true
|
||||
rstest_reuse.workspace = true
|
||||
wiremock = "0.6.5"
|
||||
wiremock.workspace = true
|
||||
|
|
|
|||
|
|
@ -30,4 +30,4 @@ futures-util.workspace = true
|
|||
tokio = { workspace = true, features = ["io-util"] }
|
||||
rstest.workspace = true
|
||||
tower = { version = "0.5.3", features = ["util"] }
|
||||
wiremock = "0.6.5"
|
||||
wiremock.workspace = true
|
||||
|
|
|
|||
57
litellm-rust/crates/host-python/src/conversion_cache.rs
Normal file
57
litellm-rust/crates/host-python/src/conversion_cache.rs
Normal file
|
|
@ -0,0 +1,57 @@
|
|||
use std::collections::{HashMap, hash_map::Entry};
|
||||
|
||||
use pyo3::prelude::*;
|
||||
|
||||
pub struct ToPythonCache<'a, 'py, T> {
|
||||
entries: HashMap<usize, (&'a T, Bound<'py, PyAny>)>,
|
||||
}
|
||||
|
||||
impl<T> Default for ToPythonCache<'_, '_, T> {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
entries: HashMap::new(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl<'a, 'py, T> ToPythonCache<'a, 'py, T> {
|
||||
pub fn get_or_try_insert_with(
|
||||
&mut self,
|
||||
value: &'a T,
|
||||
convert: impl FnOnce(&'a T) -> PyResult<Bound<'py, PyAny>>,
|
||||
) -> PyResult<&Bound<'py, PyAny>> {
|
||||
let identity = std::ptr::from_ref(value) as usize;
|
||||
let entry = match self.entries.entry(identity) {
|
||||
Entry::Occupied(entry) => entry.into_mut(),
|
||||
Entry::Vacant(entry) => entry.insert((value, convert(value)?)),
|
||||
};
|
||||
Ok(&entry.1)
|
||||
}
|
||||
}
|
||||
|
||||
pub struct FromPythonCache<'py, T> {
|
||||
entries: HashMap<usize, (Bound<'py, PyAny>, T)>,
|
||||
}
|
||||
|
||||
impl<T> Default for FromPythonCache<'_, T> {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
entries: HashMap::new(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl<'py, T> FromPythonCache<'py, T> {
|
||||
pub fn get_or_try_insert_with(
|
||||
&mut self,
|
||||
value: &Bound<'py, PyAny>,
|
||||
convert: impl FnOnce(&Bound<'py, PyAny>) -> PyResult<T>,
|
||||
) -> PyResult<&T> {
|
||||
let identity = value.as_ptr() as usize;
|
||||
let entry = match self.entries.entry(identity) {
|
||||
Entry::Occupied(entry) => entry.into_mut(),
|
||||
Entry::Vacant(entry) => entry.insert((value.clone(), convert(value)?)),
|
||||
};
|
||||
Ok(&entry.1)
|
||||
}
|
||||
}
|
||||
|
|
@ -5,6 +5,7 @@
|
|||
|
||||
mod argument;
|
||||
mod binding;
|
||||
mod conversion_cache;
|
||||
mod driver;
|
||||
mod error;
|
||||
mod file_reader;
|
||||
|
|
@ -20,6 +21,7 @@ mod services;
|
|||
|
||||
pub use argument::lookup;
|
||||
pub use binding::PythonBinding;
|
||||
pub use conversion_cache::{FromPythonCache, ToPythonCache};
|
||||
pub use driver::{CallOptions, run_call};
|
||||
pub use error::{InvokeError, missing_state};
|
||||
pub use file_reader::{FileContent, PythonFileReader, py_bytes};
|
||||
|
|
|
|||
121
litellm-rust/crates/host-python/tests/conversion_cache.rs
Normal file
121
litellm-rust/crates/host-python/tests/conversion_cache.rs
Normal file
|
|
@ -0,0 +1,121 @@
|
|||
use std::{cell::Cell, rc::Rc};
|
||||
|
||||
use litellm_host_python::{FromPythonCache, Pythonized, ToPythonCache};
|
||||
use pyo3::{exceptions::PyValueError, prelude::*, types::PyDict};
|
||||
use rstest::{fixture, rstest};
|
||||
|
||||
#[fixture]
|
||||
fn python() {
|
||||
Python::initialize();
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn rust_identity_reuses_python_objects_without_merging_equal_values(#[from(python)] _python: ()) {
|
||||
Python::attach(|py| {
|
||||
let original = Rc::new(vec![1, 2]);
|
||||
let cloned = original.clone();
|
||||
let equal = Rc::new(vec![1, 2]);
|
||||
let mut cache = ToPythonCache::default();
|
||||
let first = cache
|
||||
.get_or_try_insert_with(original.as_ref(), |value| {
|
||||
Pythonized(value).into_pyobject(py)
|
||||
})
|
||||
.unwrap()
|
||||
.clone();
|
||||
let second = cache
|
||||
.get_or_try_insert_with(cloned.as_ref(), |_| panic!("must reuse conversion"))
|
||||
.unwrap()
|
||||
.clone();
|
||||
let third = cache
|
||||
.get_or_try_insert_with(equal.as_ref(), |value| Pythonized(value).into_pyobject(py))
|
||||
.unwrap();
|
||||
assert!(first.is(&second));
|
||||
assert!(!first.is(third));
|
||||
assert!(first.eq(third).unwrap());
|
||||
});
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn python_identity_reuses_rust_values_without_merging_equal_objects(#[from(python)] _python: ()) {
|
||||
Python::attach(|py| {
|
||||
let original = PyDict::new(py);
|
||||
original.set_item("value", 1).unwrap();
|
||||
let equal = original.copy().unwrap();
|
||||
let calls = Cell::new(0);
|
||||
let mut cache = FromPythonCache::default();
|
||||
let convert = |value: &Bound<'_, PyAny>| {
|
||||
calls.set(calls.get() + 1);
|
||||
value.get_item("value")?.extract::<i32>().map(Rc::new)
|
||||
};
|
||||
let first = cache
|
||||
.get_or_try_insert_with(original.as_any(), convert)
|
||||
.unwrap()
|
||||
.clone();
|
||||
let second = cache
|
||||
.get_or_try_insert_with(original.as_any(), convert)
|
||||
.unwrap()
|
||||
.clone();
|
||||
let third = cache
|
||||
.get_or_try_insert_with(equal.as_any(), convert)
|
||||
.unwrap();
|
||||
assert!(Rc::ptr_eq(&first, &second));
|
||||
assert!(!Rc::ptr_eq(&first, third));
|
||||
assert_eq!(&first, third);
|
||||
assert_eq!(calls.get(), 2);
|
||||
});
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn python_sources_stay_alive_until_the_cache_is_dropped(#[from(python)] _python: ()) {
|
||||
Python::attach(|py| {
|
||||
let value = py
|
||||
.eval(pyo3::ffi::c_str!("type('Tracked', (), {})()"), None, None)
|
||||
.unwrap();
|
||||
let weak = py
|
||||
.import("weakref")
|
||||
.unwrap()
|
||||
.call_method1("ref", (&value,))
|
||||
.unwrap();
|
||||
let mut cache = FromPythonCache::default();
|
||||
cache.get_or_try_insert_with(&value, |_| Ok(42)).unwrap();
|
||||
drop(value);
|
||||
assert!(!weak.call0().unwrap().is_none());
|
||||
drop(cache);
|
||||
assert!(weak.call0().unwrap().is_none());
|
||||
});
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::to_python(true)]
|
||||
#[case::from_python(false)]
|
||||
fn failed_conversions_preserve_exceptions_and_can_be_retried(
|
||||
#[from(python)] _python: (),
|
||||
#[case] to_python: bool,
|
||||
) {
|
||||
Python::attach(|py| {
|
||||
let failure = PyValueError::new_err("conversion failed");
|
||||
if to_python {
|
||||
let source = vec![1, 2];
|
||||
let mut cache = ToPythonCache::default();
|
||||
let error = cache
|
||||
.get_or_try_insert_with(&source, |_| Err(failure.clone_ref(py)))
|
||||
.unwrap_err();
|
||||
assert!(error.value(py).is(failure.value(py)));
|
||||
let result = cache
|
||||
.get_or_try_insert_with(&source, |value| Pythonized(value).into_pyobject(py))
|
||||
.unwrap();
|
||||
assert_eq!(result.extract::<Vec<i32>>().unwrap(), source);
|
||||
} else {
|
||||
let source = PyDict::new(py).into_any();
|
||||
let mut cache = FromPythonCache::default();
|
||||
let error = cache
|
||||
.get_or_try_insert_with(&source, |_| Err(failure.clone_ref(py)))
|
||||
.unwrap_err();
|
||||
assert!(error.value(py).is(failure.value(py)));
|
||||
assert_eq!(
|
||||
*cache.get_or_try_insert_with(&source, |_| Ok(42)).unwrap(),
|
||||
42
|
||||
);
|
||||
}
|
||||
});
|
||||
}
|
||||
|
|
@ -22,6 +22,7 @@ tiktoken = ["litellm-token-counter/tiktoken"]
|
|||
fancy-regex.workspace = true
|
||||
litellm-tracing.workspace = true
|
||||
litellm-traces.workspace = true
|
||||
litellm-storage-clickhouse.workspace = true
|
||||
litellm-host.workspace = true
|
||||
bytes.workspace = true
|
||||
futures-util.workspace = true
|
||||
|
|
@ -51,6 +52,7 @@ litellm-llms-types.workspace = true
|
|||
litellm-host-python.workspace = true
|
||||
litellm-token-counter = { path = "../token-counter", default-features = false }
|
||||
pyo3.workspace = true
|
||||
prost.workspace = true
|
||||
pyo3-async-runtimes.workspace = true
|
||||
reqwest.workspace = true
|
||||
redis = { version = "1.7.0", features = ["tls-rustls"] }
|
||||
|
|
@ -72,7 +74,7 @@ futures-util.workspace = true
|
|||
rstest.workspace = true
|
||||
sha2.workspace = true
|
||||
tokio-tungstenite.workspace = true
|
||||
wiremock = "0.6.5"
|
||||
wiremock.workspace = true
|
||||
aws-sdk-secretsmanager = "1.117.0"
|
||||
|
||||
[[bench]]
|
||||
|
|
|
|||
|
|
@ -44,7 +44,7 @@ mod _native {
|
|||
#[pymodule_export]
|
||||
use crate::routes::token_counter::TokenCounter;
|
||||
#[pymodule_export]
|
||||
use crate::routes::traces::{NativeTraceStorage, trace_decode_otlp};
|
||||
use crate::routes::traces::{NativeTraceStorage, trace_decode_otlp, trace_encode_error};
|
||||
#[cfg(feature = "huggingface")]
|
||||
#[pymodule_export]
|
||||
use crate::tokenizer::HuggingFaceEncoding;
|
||||
|
|
@ -111,6 +111,7 @@ mod tests {
|
|||
"NativeDiagnosticProcessor",
|
||||
"NativeTraceStorage",
|
||||
"trace_decode_otlp",
|
||||
"trace_encode_error",
|
||||
"TokenCounter",
|
||||
"Tokenizer",
|
||||
"gil_stats",
|
||||
|
|
|
|||
|
|
@ -1,12 +1,33 @@
|
|||
use std::collections::BTreeMap;
|
||||
|
||||
use litellm_host_python::{FromPythonCache, ToPythonCache};
|
||||
use litellm_http::ClientVariant;
|
||||
use litellm_traces::{Connection, Error, InsertTable, Parameter, ReadQuery};
|
||||
use litellm_storage_clickhouse::Storage;
|
||||
use litellm_traces::{Error, InsertTable, Parameter, ReadQuery, Shared};
|
||||
use prost::Message;
|
||||
use pyo3::{
|
||||
exceptions::{PyOverflowError, PyRuntimeError, PyValueError},
|
||||
prelude::*,
|
||||
types::{PyBytes, PyDict, PyList, PyMapping, PyString},
|
||||
};
|
||||
|
||||
#[derive(Message)]
|
||||
struct OtlpErrorStatus {
|
||||
#[prost(int32, tag = "1")]
|
||||
code: i32,
|
||||
#[prost(string, tag = "2")]
|
||||
message: String,
|
||||
}
|
||||
|
||||
#[pyfunction]
|
||||
pub fn trace_encode_error<'py>(py: Python<'py>, message: &str) -> Bound<'py, PyBytes> {
|
||||
let status = OtlpErrorStatus {
|
||||
code: 0,
|
||||
message: message.to_owned(),
|
||||
};
|
||||
PyBytes::new(py, &status.encode_to_vec())
|
||||
}
|
||||
|
||||
fn map_error(error: Error) -> PyErr {
|
||||
match error {
|
||||
Error::InvalidRow
|
||||
|
|
@ -27,9 +48,7 @@ fn map_error(error: Error) -> PyErr {
|
|||
|
||||
#[pyclass]
|
||||
pub struct NativeTraceStorage {
|
||||
database: String,
|
||||
writer: Connection,
|
||||
reader: Option<Connection>,
|
||||
storage: Storage,
|
||||
}
|
||||
|
||||
#[pymethods]
|
||||
|
|
@ -39,12 +58,7 @@ impl NativeTraceStorage {
|
|||
fn new(database: String, url: &str, reader_url: Option<&str>) -> PyResult<Self> {
|
||||
litellm_traces::schema_statements(&database, 1, 1).map_err(map_error)?;
|
||||
Ok(Self {
|
||||
writer: Connection::writer(url).map_err(map_error)?,
|
||||
reader: reader_url
|
||||
.map(|value| Connection::reader(value, &database))
|
||||
.transpose()
|
||||
.map_err(map_error)?,
|
||||
database,
|
||||
storage: Storage::new(database, url, reader_url).map_err(map_error)?,
|
||||
})
|
||||
}
|
||||
|
||||
|
|
@ -55,8 +69,8 @@ impl NativeTraceStorage {
|
|||
spend_log_retention_days: u32,
|
||||
) -> PyResult<Bound<'py, PyAny>> {
|
||||
let client = crate::http::host_client(py, ClientVariant::NoRedirect)?;
|
||||
let connection = self.writer.clone();
|
||||
let database = self.database.clone();
|
||||
let connection = self.storage.writer().clone();
|
||||
let database = self.storage.database().to_owned();
|
||||
crate::execution::run_async(
|
||||
py,
|
||||
async move {
|
||||
|
|
@ -77,18 +91,17 @@ impl NativeTraceStorage {
|
|||
&self,
|
||||
py: Python<'py>,
|
||||
table: &str,
|
||||
#[pyo3(from_py_with = litellm_host_python::from_py_argument)] rows: Vec<
|
||||
BTreeMap<String, serde_json::Value>,
|
||||
>,
|
||||
#[pyo3(from_py_with = insert_rows_from_py)] rows: Vec<litellm_traces::InsertRow>,
|
||||
) -> PyResult<Bound<'py, PyAny>> {
|
||||
let table = InsertTable::parse(table).map_err(map_error)?;
|
||||
let client = crate::http::host_client(py, ClientVariant::NoRedirect)?;
|
||||
let connection = self.writer.clone();
|
||||
let database = self.database.clone();
|
||||
let connection = self.storage.writer().clone();
|
||||
let database = self.storage.database().to_owned();
|
||||
crate::execution::run_async(
|
||||
py,
|
||||
async move {
|
||||
litellm_traces::insert_rows(&client, &connection, &database, table, rows).await
|
||||
litellm_traces::insert_shared_rows(&client, &connection, &database, table, rows)
|
||||
.await
|
||||
},
|
||||
map_error,
|
||||
)
|
||||
|
|
@ -104,7 +117,7 @@ impl NativeTraceStorage {
|
|||
>,
|
||||
) -> PyResult<Bound<'py, PyAny>> {
|
||||
let query = litellm_traces::LensQuery::parse(name).map_err(map_error)?;
|
||||
let connection = self.reader.clone().ok_or_else(|| {
|
||||
let connection = self.storage.reader().cloned().ok_or_else(|| {
|
||||
PyRuntimeError::new_err("Trace reads require a separate ClickHouse reader URL")
|
||||
})?;
|
||||
let client = crate::http::host_client(py, ClientVariant::NoRedirect)?;
|
||||
|
|
@ -127,7 +140,7 @@ impl NativeTraceStorage {
|
|||
>,
|
||||
) -> PyResult<Bound<'py, PyAny>> {
|
||||
let query = ReadQuery::parse(query).map_err(map_error)?;
|
||||
let connection = self.reader.clone().ok_or_else(|| {
|
||||
let connection = self.storage.reader().cloned().ok_or_else(|| {
|
||||
PyRuntimeError::new_err("Trace reads require a separate ClickHouse reader URL")
|
||||
})?;
|
||||
let client = crate::http::host_client(py, ClientVariant::NoRedirect)?;
|
||||
|
|
@ -146,21 +159,132 @@ pub fn trace_decode_otlp<'py>(
|
|||
py: Python<'py>,
|
||||
body: &[u8],
|
||||
content_type: Option<&str>,
|
||||
content_encoding: Option<&str>,
|
||||
max_decompressed_bytes: usize,
|
||||
) -> PyResult<Bound<'py, PyAny>> {
|
||||
let spans = py
|
||||
.detach(|| {
|
||||
litellm_traces::decode_otlp(
|
||||
body,
|
||||
content_type,
|
||||
content_encoding,
|
||||
max_decompressed_bytes,
|
||||
)
|
||||
})
|
||||
.detach(|| litellm_traces::decode_otlp(body, content_type))
|
||||
.map_err(|error| match error {
|
||||
litellm_traces::DecodeError::TooLarge => PyOverflowError::new_err(error.to_string()),
|
||||
_ => PyValueError::new_err(error.to_string()),
|
||||
})?;
|
||||
litellm_host_python::Pythonized(spans).into_pyobject(py)
|
||||
spans_to_py(py, &spans).map(Bound::into_any)
|
||||
}
|
||||
|
||||
fn insert_rows_from_py(value: &Bound<'_, PyAny>) -> PyResult<Vec<litellm_traces::InsertRow>> {
|
||||
let mut resources = FromPythonCache::default();
|
||||
value
|
||||
.try_iter()?
|
||||
.map(|row| {
|
||||
let row = row?;
|
||||
let mut fields = BTreeMap::new();
|
||||
for item in row.cast::<PyMapping>()?.items()?.iter() {
|
||||
let (key, value): (String, Bound<'_, PyAny>) = item.extract()?;
|
||||
let converted = if matches!(
|
||||
key.as_str(),
|
||||
"ResourceAttributes" | "ScopeName" | "ScopeVersion"
|
||||
) {
|
||||
resources
|
||||
.get_or_try_insert_with(&value, |value| {
|
||||
litellm_host_python::from_py_argument::<serde_json::Value>(value)
|
||||
.map(Shared::new)
|
||||
})?
|
||||
.clone()
|
||||
} else {
|
||||
Shared::new(litellm_host_python::from_py_argument(&value)?)
|
||||
};
|
||||
fields.insert(key, converted);
|
||||
}
|
||||
Ok(fields)
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn spans_to_py<'py>(
|
||||
py: Python<'py>,
|
||||
spans: &[litellm_traces::DecodedSpan],
|
||||
) -> PyResult<Bound<'py, PyList>> {
|
||||
let mut resources = ToPythonCache::default();
|
||||
let mut scopes = ToPythonCache::default();
|
||||
let result = PyList::empty(py);
|
||||
for span in spans {
|
||||
let resource = resources
|
||||
.get_or_try_insert_with(span.resource_attributes.as_ref(), |value| {
|
||||
litellm_host_python::Pythonized(value).into_pyobject(py)
|
||||
})?;
|
||||
let row = PyDict::new(py);
|
||||
row.set_item("trace_id", &span.trace_id)?;
|
||||
row.set_item("span_id", &span.span_id)?;
|
||||
row.set_item("parent_span_id", &span.parent_span_id)?;
|
||||
row.set_item("trace_state", &span.trace_state)?;
|
||||
row.set_item("name", &span.name)?;
|
||||
row.set_item("kind", &span.kind)?;
|
||||
row.set_item("resource_attributes", resource)?;
|
||||
for (key, value) in [
|
||||
("scope_name", &span.scope_name),
|
||||
("scope_version", &span.scope_version),
|
||||
] {
|
||||
let value = scopes.get_or_try_insert_with(value.as_ref(), |value| {
|
||||
Ok(PyString::new(py, value).into_any())
|
||||
})?;
|
||||
row.set_item(key, value)?;
|
||||
}
|
||||
row.set_item("attributes", &span.attributes)?;
|
||||
row.set_item("start_ns", span.start_ns)?;
|
||||
row.set_item("end_ns", span.end_ns)?;
|
||||
row.set_item("status_code", &span.status_code)?;
|
||||
row.set_item("status_message", &span.status_message)?;
|
||||
row.set_item(
|
||||
"events",
|
||||
litellm_host_python::Pythonized(&span.events).into_pyobject(py)?,
|
||||
)?;
|
||||
result.append(row)?;
|
||||
}
|
||||
Ok(result)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use rstest::rstest;
|
||||
|
||||
#[rstest]
|
||||
fn insert_projection_preserves_identity_without_merging_equal_resources() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let resource = PyDict::new(py);
|
||||
resource.set_item("service.name", "shared").unwrap();
|
||||
let equal_resource = resource.copy().unwrap();
|
||||
let rows = PyList::empty(py);
|
||||
for value in [&resource, &resource, &equal_resource] {
|
||||
let row = PyDict::new(py);
|
||||
row.set_item("ResourceAttributes", value).unwrap();
|
||||
rows.append(row).unwrap();
|
||||
}
|
||||
let projected = insert_rows_from_py(rows.as_any()).unwrap();
|
||||
assert!(Shared::shares_storage_with(
|
||||
&projected[0]["ResourceAttributes"],
|
||||
&projected[1]["ResourceAttributes"]
|
||||
));
|
||||
assert!(!Shared::shares_storage_with(
|
||||
&projected[0]["ResourceAttributes"],
|
||||
&projected[2]["ResourceAttributes"]
|
||||
));
|
||||
assert_eq!(projected[0], projected[2]);
|
||||
});
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn shared_conversion_preserves_every_decoded_field() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let spans = litellm_traces::decode_otlp(
|
||||
include_bytes!("../../../../../tests/test_litellm/tracing/fixtures/langsmith_deep_agent_export.json"),
|
||||
Some("application/json"),
|
||||
).unwrap();
|
||||
let expected = litellm_host_python::Pythonized(&spans)
|
||||
.into_pyobject(py)
|
||||
.unwrap();
|
||||
let actual = spans_to_py(py, &spans).unwrap();
|
||||
assert!(actual.eq(expected).unwrap());
|
||||
});
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -21,5 +21,5 @@ aws-credential-types = "1.3.0"
|
|||
base64.workspace = true
|
||||
rstest.workspace = true
|
||||
tokio.workspace = true
|
||||
wiremock = "0.6.5"
|
||||
wiremock.workspace = true
|
||||
tempfile = "3"
|
||||
|
|
|
|||
|
|
@ -20,7 +20,7 @@ percent-encoding = "2.3"
|
|||
|
||||
[dev-dependencies]
|
||||
litellm-http = { workspace = true, features = ["test-support"] }
|
||||
wiremock = "0.6.5"
|
||||
wiremock.workspace = true
|
||||
rstest.workspace = true
|
||||
serde_json.workspace = true
|
||||
sha2.workspace = true
|
||||
|
|
|
|||
|
|
@ -25,6 +25,6 @@ rcgen = "0.14.10"
|
|||
rstest.workspace = true
|
||||
tempfile = "3.27.0"
|
||||
tokio.workspace = true
|
||||
wiremock = "0.6.5"
|
||||
wiremock.workspace = true
|
||||
serde.workspace = true
|
||||
serde_json.workspace = true
|
||||
|
|
|
|||
|
|
@ -28,4 +28,4 @@ reqwest.workspace = true
|
|||
litellm-http = { workspace = true, features = ["test-support"] }
|
||||
google-cloud-auth.workspace = true
|
||||
rstest.workspace = true
|
||||
wiremock = "0.6.5"
|
||||
wiremock.workspace = true
|
||||
|
|
|
|||
|
|
@ -21,4 +21,4 @@ veil.workspace = true
|
|||
rstest.workspace = true
|
||||
tempfile = "3"
|
||||
tokio.workspace = true
|
||||
wiremock = "0.6.5"
|
||||
wiremock.workspace = true
|
||||
|
|
|
|||
|
|
@ -36,7 +36,7 @@ tokio = { workspace = true, features = ["fs"] }
|
|||
[dev-dependencies]
|
||||
litellm-http = { workspace = true, features = ["test-support"] }
|
||||
rstest.workspace = true
|
||||
wiremock = "0.6.5"
|
||||
wiremock.workspace = true
|
||||
tempfile = "3"
|
||||
aws-sdk-kms = "1.120.0"
|
||||
google-cloud-kms-v1 = "1.14.0"
|
||||
|
|
|
|||
20
litellm-rust/crates/storage-clickhouse/Cargo.toml
Normal file
20
litellm-rust/crates/storage-clickhouse/Cargo.toml
Normal file
|
|
@ -0,0 +1,20 @@
|
|||
[package]
|
||||
name = "litellm-storage-clickhouse"
|
||||
version = "0.1.0"
|
||||
description = "Shared ClickHouse connection and HTTP storage for LiteLLM features"
|
||||
edition.workspace = true
|
||||
license.workspace = true
|
||||
repository.workspace = true
|
||||
|
||||
[dependencies]
|
||||
flate2.workspace = true
|
||||
litellm-http.workspace = true
|
||||
serde.workspace = true
|
||||
serde_json.workspace = true
|
||||
thiserror.workspace = true
|
||||
url.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
litellm-http = { workspace = true, features = ["test-support"] }
|
||||
rstest.workspace = true
|
||||
tokio.workspace = true
|
||||
5
litellm-rust/crates/storage-clickhouse/README.md
Normal file
5
litellm-rust/crates/storage-clickhouse/README.md
Normal file
|
|
@ -0,0 +1,5 @@
|
|||
# ClickHouse storage
|
||||
|
||||
`litellm-storage-clickhouse` exports `Storage`, a shared writer connection and optional reader connection for one ClickHouse database. It also exports bounded HTTP read and insert execution
|
||||
|
||||
The crate has no trace tables, OTLP types, or named trace queries. `litellm-traces` supplies those rules and uses this storage for both trace rows and spend rows
|
||||
29
litellm-rust/crates/storage-clickhouse/src/error.rs
Normal file
29
litellm-rust/crates/storage-clickhouse/src/error.rs
Normal file
|
|
@ -0,0 +1,29 @@
|
|||
#[derive(Debug, thiserror::Error)]
|
||||
pub enum Error {
|
||||
#[error("invalid ClickHouse insert row")]
|
||||
InvalidRow,
|
||||
#[error("invalid ClickHouse insert table")]
|
||||
InvalidTable,
|
||||
#[error("invalid ClickHouse HTTP URL")]
|
||||
InvalidUrl,
|
||||
#[error("database must be a nonempty SQL identifier and retention must be positive")]
|
||||
InvalidSchema,
|
||||
#[error("SQL query must not be empty")]
|
||||
EmptySql,
|
||||
#[error("unknown ClickHouse read query")]
|
||||
InvalidQuery,
|
||||
#[error("ClickHouse query failed with HTTP status {0}")]
|
||||
QueryFailed(u16),
|
||||
#[error("ClickHouse insert failed with HTTP status {0}")]
|
||||
InsertFailed(u16),
|
||||
#[error("ClickHouse insert exceeds the encoded size limit")]
|
||||
InsertTooLarge,
|
||||
#[error("ClickHouse schema setup failed with HTTP status {0}")]
|
||||
SchemaFailed(u16),
|
||||
#[error("ClickHouse query exceeded the response size limit")]
|
||||
ResponseTooLarge,
|
||||
#[error("ClickHouse returned an invalid or failed JSON query response")]
|
||||
InvalidResponse,
|
||||
#[error("ClickHouse query transport failed")]
|
||||
Transport,
|
||||
}
|
||||
87
litellm-rust/crates/storage-clickhouse/src/insert.rs
Normal file
87
litellm-rust/crates/storage-clickhouse/src/insert.rs
Normal file
|
|
@ -0,0 +1,87 @@
|
|||
use std::{io::Write, time::Duration};
|
||||
|
||||
use flate2::{Compression, write::GzEncoder};
|
||||
use litellm_http::Client;
|
||||
|
||||
use crate::{Connection, Error, valid_identifier};
|
||||
|
||||
const INSERT_TIMEOUT: Duration = Duration::from_secs(30);
|
||||
|
||||
pub async fn insert_encoded_rows(
|
||||
client: &Client,
|
||||
connection: &Connection,
|
||||
database: &str,
|
||||
table: &str,
|
||||
token: &str,
|
||||
encoded: &str,
|
||||
) -> Result<(), Error> {
|
||||
if !valid_identifier(database) {
|
||||
return Err(Error::InvalidSchema);
|
||||
}
|
||||
if !valid_identifier(table) {
|
||||
return Err(Error::InvalidTable);
|
||||
}
|
||||
let mut encoder = GzEncoder::new(Vec::new(), Compression::default());
|
||||
encoder
|
||||
.write_all(encoded.as_bytes())
|
||||
.map_err(|_| Error::InvalidRow)?;
|
||||
let body = encoder.finish().map_err(|_| Error::InvalidRow)?;
|
||||
insert_compressed_rows(client, connection, database, table, token, body).await
|
||||
}
|
||||
|
||||
pub async fn insert_compressed_rows(
|
||||
client: &Client,
|
||||
connection: &Connection,
|
||||
database: &str,
|
||||
table: &str,
|
||||
token: &str,
|
||||
body: Vec<u8>,
|
||||
) -> Result<(), Error> {
|
||||
if !valid_identifier(database) {
|
||||
return Err(Error::InvalidSchema);
|
||||
}
|
||||
if !valid_identifier(table) {
|
||||
return Err(Error::InvalidTable);
|
||||
}
|
||||
let mut url = connection.url().clone();
|
||||
let existing_pairs: Vec<(String, String)> = url
|
||||
.query_pairs()
|
||||
.filter(|(key, _)| {
|
||||
!matches!(
|
||||
key.as_ref(),
|
||||
"query"
|
||||
| "async_insert"
|
||||
| "async_insert_deduplicate"
|
||||
| "wait_for_async_insert"
|
||||
| "input_format_skip_unknown_fields"
|
||||
| "date_time_input_format"
|
||||
)
|
||||
})
|
||||
.map(|(key, value)| (key.into_owned(), value.into_owned()))
|
||||
.collect();
|
||||
url.query_pairs_mut()
|
||||
.clear()
|
||||
.extend_pairs(existing_pairs)
|
||||
.append_pair(
|
||||
"query",
|
||||
&format!("INSERT INTO `{database}`.{} FORMAT JSONEachRow", table),
|
||||
)
|
||||
.append_pair("insert_deduplication_token", token)
|
||||
.append_pair("async_insert", "1")
|
||||
.append_pair("async_insert_deduplicate", "1")
|
||||
.append_pair("wait_for_async_insert", "1")
|
||||
.append_pair("input_format_skip_unknown_fields", "0")
|
||||
.append_pair("date_time_input_format", "best_effort");
|
||||
let response = client
|
||||
.post(url)
|
||||
.timeout(INSERT_TIMEOUT)
|
||||
.header("Content-Encoding", "gzip")
|
||||
.body(body)
|
||||
.send()
|
||||
.await
|
||||
.map_err(|_| Error::Transport)?;
|
||||
if !response.status().is_success() {
|
||||
return Err(Error::InsertFailed(response.status().as_u16()));
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
127
litellm-rust/crates/storage-clickhouse/src/lib.rs
Normal file
127
litellm-rust/crates/storage-clickhouse/src/lib.rs
Normal file
|
|
@ -0,0 +1,127 @@
|
|||
mod error;
|
||||
mod insert;
|
||||
mod read;
|
||||
|
||||
pub use error::Error;
|
||||
pub use insert::{insert_compressed_rows, insert_encoded_rows};
|
||||
pub use read::{Parameter, execute_read};
|
||||
use url::Url;
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct Connection {
|
||||
url: Url,
|
||||
}
|
||||
|
||||
impl Connection {
|
||||
pub fn parse(value: &str) -> Result<Self, Error> {
|
||||
let url = Url::parse(value).map_err(|_| Error::InvalidUrl)?;
|
||||
if !matches!(url.scheme(), "http" | "https") || url.host().is_none() {
|
||||
return Err(Error::InvalidUrl);
|
||||
}
|
||||
Ok(Self { url })
|
||||
}
|
||||
|
||||
pub fn configured(
|
||||
url: &str,
|
||||
database: &str,
|
||||
user: &str,
|
||||
password: &str,
|
||||
) -> Result<Self, Error> {
|
||||
let mut connection = Self::parse(url)?;
|
||||
connection
|
||||
.url
|
||||
.set_username(user)
|
||||
.map_err(|_| Error::InvalidUrl)?;
|
||||
connection
|
||||
.url
|
||||
.set_password(Some(password))
|
||||
.map_err(|_| Error::InvalidUrl)?;
|
||||
let pairs: Vec<_> = connection
|
||||
.url
|
||||
.query_pairs()
|
||||
.filter(|(key, _)| !matches!(key.as_ref(), "database" | "user" | "password"))
|
||||
.map(|(key, value)| (key.into_owned(), value.into_owned()))
|
||||
.collect();
|
||||
connection
|
||||
.url
|
||||
.query_pairs_mut()
|
||||
.clear()
|
||||
.extend_pairs(pairs)
|
||||
.append_pair("database", database);
|
||||
Ok(connection)
|
||||
}
|
||||
|
||||
pub fn writer(url: &str) -> Result<Self, Error> {
|
||||
let mut connection = Self::parse(url)?;
|
||||
let pairs: Vec<_> = connection
|
||||
.url
|
||||
.query_pairs()
|
||||
.filter(|(key, _)| !matches!(key.as_ref(), "database" | "readonly" | "query"))
|
||||
.map(|(key, value)| (key.into_owned(), value.into_owned()))
|
||||
.collect();
|
||||
connection.url.query_pairs_mut().clear().extend_pairs(pairs);
|
||||
Ok(connection)
|
||||
}
|
||||
|
||||
pub fn reader(url: &str, database: &str) -> Result<Self, Error> {
|
||||
let mut connection = Self::parse(url)?;
|
||||
let pairs: Vec<_> = connection
|
||||
.url
|
||||
.query_pairs()
|
||||
.filter(|(key, _)| key != "database")
|
||||
.map(|(key, value)| (key.into_owned(), value.into_owned()))
|
||||
.collect();
|
||||
connection
|
||||
.url
|
||||
.query_pairs_mut()
|
||||
.clear()
|
||||
.extend_pairs(pairs)
|
||||
.append_pair("database", database);
|
||||
Ok(connection)
|
||||
}
|
||||
|
||||
pub fn url(&self) -> &Url {
|
||||
&self.url
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct Storage {
|
||||
database: String,
|
||||
writer: Connection,
|
||||
reader: Option<Connection>,
|
||||
}
|
||||
|
||||
impl Storage {
|
||||
pub fn new(database: String, url: &str, reader_url: Option<&str>) -> Result<Self, Error> {
|
||||
if !valid_identifier(&database) {
|
||||
return Err(Error::InvalidSchema);
|
||||
}
|
||||
Ok(Self {
|
||||
writer: Connection::writer(url)?,
|
||||
reader: reader_url
|
||||
.map(|value| Connection::reader(value, &database))
|
||||
.transpose()?,
|
||||
database,
|
||||
})
|
||||
}
|
||||
|
||||
pub fn database(&self) -> &str {
|
||||
&self.database
|
||||
}
|
||||
|
||||
pub fn writer(&self) -> &Connection {
|
||||
&self.writer
|
||||
}
|
||||
|
||||
pub fn reader(&self) -> Option<&Connection> {
|
||||
self.reader.as_ref()
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn valid_identifier(value: &str) -> bool {
|
||||
!value.is_empty()
|
||||
&& value
|
||||
.bytes()
|
||||
.all(|c| c.is_ascii_alphanumeric() || c == b'_')
|
||||
}
|
||||
113
litellm-rust/crates/storage-clickhouse/src/read.rs
Normal file
113
litellm-rust/crates/storage-clickhouse/src/read.rs
Normal file
|
|
@ -0,0 +1,113 @@
|
|||
use std::{collections::BTreeMap, time::Duration};
|
||||
|
||||
use litellm_http::Client;
|
||||
use serde::Deserialize;
|
||||
|
||||
use crate::{Connection, Error};
|
||||
|
||||
const MAX_RESPONSE_BYTES: usize = 4 * 1024 * 1024;
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
#[serde(untagged)]
|
||||
pub enum Parameter {
|
||||
Text(String),
|
||||
Integer(i64),
|
||||
Strings(Vec<String>),
|
||||
}
|
||||
|
||||
impl Parameter {
|
||||
fn encoded(&self) -> String {
|
||||
match self {
|
||||
Self::Text(value) => escaped(value),
|
||||
Self::Integer(value) => value.to_string(),
|
||||
Self::Strings(values) => format!(
|
||||
"[{}]",
|
||||
values
|
||||
.iter()
|
||||
.map(|value| format!("'{}'", escaped(value).replace('\'', "\\'")))
|
||||
.collect::<Vec<_>>()
|
||||
.join(",")
|
||||
),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn escaped(value: &str) -> String {
|
||||
value
|
||||
.replace('\\', "\\\\")
|
||||
.replace('\t', "\\t")
|
||||
.replace('\n', "\\n")
|
||||
.replace('\r', "\\r")
|
||||
.replace('\0', "\\0")
|
||||
}
|
||||
|
||||
pub async fn execute_read(
|
||||
client: &Client,
|
||||
connection: &Connection,
|
||||
sql: &str,
|
||||
parameters: &BTreeMap<String, Parameter>,
|
||||
) -> Result<String, Error> {
|
||||
if sql.trim().is_empty() {
|
||||
return Err(Error::EmptySql);
|
||||
}
|
||||
|
||||
let mut url = connection.url().clone();
|
||||
|
||||
let existing_pairs: Vec<(String, String)> = url
|
||||
.query_pairs()
|
||||
.filter(|(key, _)| {
|
||||
!key.starts_with("param_")
|
||||
&& !matches!(
|
||||
key.as_ref(),
|
||||
"query"
|
||||
| "readonly"
|
||||
| "default_format"
|
||||
| "max_result_rows"
|
||||
| "result_overflow_mode"
|
||||
| "max_execution_time"
|
||||
| "wait_end_of_query"
|
||||
)
|
||||
})
|
||||
.map(|(key, value)| (key.into_owned(), value.into_owned()))
|
||||
.collect();
|
||||
url.query_pairs_mut()
|
||||
.clear()
|
||||
.extend_pairs(existing_pairs)
|
||||
.append_pair("readonly", "1")
|
||||
.append_pair("max_result_rows", "1000")
|
||||
.append_pair("result_overflow_mode", "throw")
|
||||
.append_pair("max_execution_time", "10")
|
||||
.append_pair("wait_end_of_query", "1")
|
||||
.append_pair("default_format", "JSON");
|
||||
|
||||
url.query_pairs_mut().extend_pairs(
|
||||
parameters
|
||||
.iter()
|
||||
.map(|(name, value)| (format!("param_{name}"), value.encoded())),
|
||||
);
|
||||
|
||||
let request = client
|
||||
.post(url)
|
||||
.timeout(Duration::from_secs(15))
|
||||
.body(sql.to_owned());
|
||||
let mut response = request.send().await.map_err(|_| Error::Transport)?;
|
||||
if !response.status().is_success() {
|
||||
return Err(Error::QueryFailed(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() > MAX_RESPONSE_BYTES {
|
||||
return Err(Error::ResponseTooLarge);
|
||||
}
|
||||
body.extend_from_slice(&chunk);
|
||||
}
|
||||
|
||||
let json: serde_json::Value =
|
||||
serde_json::from_slice(&body).map_err(|_| Error::InvalidResponse)?;
|
||||
if json.get("exception").is_some() || !json.get("data").is_some_and(serde_json::Value::is_array)
|
||||
{
|
||||
return Err(Error::InvalidResponse);
|
||||
}
|
||||
String::from_utf8(body).map_err(|_| Error::InvalidResponse)
|
||||
}
|
||||
34
litellm-rust/crates/storage-clickhouse/tests/connection.rs
Normal file
34
litellm-rust/crates/storage-clickhouse/tests/connection.rs
Normal file
|
|
@ -0,0 +1,34 @@
|
|||
use litellm_storage_clickhouse::{Connection, Storage};
|
||||
use rstest::rstest;
|
||||
|
||||
#[rstest]
|
||||
#[case::http("http://localhost:8123", true)]
|
||||
#[case::https("https://localhost:8443", true)]
|
||||
#[case::tcp("tcp://localhost:9000", false)]
|
||||
#[case::missing_host("http://", false)]
|
||||
fn accepts_only_clickhouse_http_urls(#[case] value: &str, #[case] expected: bool) {
|
||||
assert_eq!(Connection::parse(value).is_ok(), expected);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::writer_only(None, false)]
|
||||
#[case::separate_reader(Some("http://localhost:8124"), true)]
|
||||
fn storage_exports_writer_and_optional_reader(
|
||||
#[case] reader_url: Option<&str>,
|
||||
#[case] has_reader: bool,
|
||||
) {
|
||||
let storage = Storage::new("litellm".to_owned(), "http://localhost:8123", reader_url)
|
||||
.expect("valid ClickHouse URLs");
|
||||
|
||||
assert_eq!(storage.database(), "litellm");
|
||||
assert_eq!(storage.writer().url().host_str(), Some("localhost"));
|
||||
assert_eq!(storage.writer().url().port(), Some(8123));
|
||||
assert_eq!(storage.reader().is_some(), has_reader);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::empty("")]
|
||||
#[case::injection("db; DROP DATABASE default")]
|
||||
fn storage_rejects_invalid_database(#[case] database: &str) {
|
||||
assert!(Storage::new(database.to_owned(), "http://localhost:8123", None).is_err());
|
||||
}
|
||||
34
litellm-rust/crates/storage-clickhouse/tests/transport.rs
Normal file
34
litellm-rust/crates/storage-clickhouse/tests/transport.rs
Normal file
|
|
@ -0,0 +1,34 @@
|
|||
use std::collections::BTreeMap;
|
||||
|
||||
use litellm_http::Client;
|
||||
use litellm_storage_clickhouse::{Connection, Error, execute_read, insert_encoded_rows};
|
||||
use rstest::rstest;
|
||||
|
||||
#[rstest]
|
||||
#[case::invalid_database("db; DROP DATABASE default", "spend_logs", true)]
|
||||
#[case::invalid_table("litellm", "spend_logs; DROP TABLE otel_traces", false)]
|
||||
#[tokio::test]
|
||||
async fn insert_rejects_invalid_identifiers(
|
||||
#[case] database: &str,
|
||||
#[case] table: &str,
|
||||
#[case] invalid_database: bool,
|
||||
) {
|
||||
let client = Client::no_redirect_for_test();
|
||||
let connection = Connection::writer("http://localhost:8123").expect("valid URL");
|
||||
let result = insert_encoded_rows(&client, &connection, database, table, "token", "{}").await;
|
||||
|
||||
assert!(matches!(&result, Err(Error::InvalidSchema)) == invalid_database);
|
||||
assert!(matches!(&result, Err(Error::InvalidTable)) == !invalid_database);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn read_rejects_empty_sql() {
|
||||
let client = Client::no_redirect_for_test();
|
||||
let connection = Connection::reader("http://localhost:8123", "litellm").expect("valid URL");
|
||||
|
||||
assert!(matches!(
|
||||
execute_read(&client, &connection, " ", &BTreeMap::new()).await,
|
||||
Err(Error::EmptySql)
|
||||
));
|
||||
}
|
||||
|
|
@ -1,4 +1,4 @@
|
|||
- Rust owns OTLP wire decoding, ClickHouse schema, row encoding, named reads, connection validation and transport
|
||||
- Keep OTLP decoding, trace schema, row encoding and named query selection here. Generic ClickHouse connections and HTTP execution belong in `litellm-storage-clickhouse`
|
||||
- Keep this crate independent of Python; PyO3 conversion and public Python exceptions belong in `python-bridge`
|
||||
- Keep the SQL migrations here as the only ClickHouse schema definition
|
||||
- Use typed query parameters and a dedicated SELECT-only reader with server-side limits
|
||||
|
|
|
|||
|
|
@ -8,18 +8,25 @@ repository.workspace = true
|
|||
[dependencies]
|
||||
base64.workspace = true
|
||||
flate2.workspace = true
|
||||
opentelemetry-proto = { version = "0.33.0", default-features = false, features = ["gen-tonic-messages", "trace", "with-serde"] }
|
||||
prost = "0.14.4"
|
||||
opentelemetry-proto = { workspace = true, features = ["gen-tonic-messages", "trace", "with-serde"] }
|
||||
prost.workspace = true
|
||||
time = { workspace = true, features = ["formatting"] }
|
||||
litellm-http.workspace = true
|
||||
litellm-storage-clickhouse.workspace = true
|
||||
sha2.workspace = true
|
||||
serde.workspace = true
|
||||
serde = { workspace = true, features = ["rc"] }
|
||||
serde_json.workspace = true
|
||||
strum.workspace = true
|
||||
thiserror.workspace = true
|
||||
url.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
criterion.workspace = true
|
||||
litellm-http = { workspace = true, features = ["test-support"] }
|
||||
rstest.workspace = true
|
||||
testcontainers-modules = { version = "0.15.0", features = ["clickhouse"] }
|
||||
tokio.workspace = true
|
||||
wiremock.workspace = true
|
||||
|
||||
[[bench]]
|
||||
name = "resource-fanout"
|
||||
harness = false
|
||||
|
|
|
|||
39
litellm-rust/crates/traces/benches/resource-fanout.rs
Normal file
39
litellm-rust/crates/traces/benches/resource-fanout.rs
Normal file
|
|
@ -0,0 +1,39 @@
|
|||
use std::{collections::BTreeMap, hint::black_box, time::Duration};
|
||||
|
||||
use criterion::{BenchmarkId, Criterion, Throughput, criterion_group, criterion_main};
|
||||
use litellm_traces::Shared;
|
||||
|
||||
fn fanout<T: Clone>(resource: &T, spans: usize) -> Vec<T> {
|
||||
(0..spans).map(|_| resource.clone()).collect()
|
||||
}
|
||||
|
||||
fn resource_fanout(c: &mut Criterion) {
|
||||
let mut group = c.benchmark_group("resource_fanout");
|
||||
for (attribute_bytes, spans) in [(256, 1), (256, 64), (8192, 1024), (16384, 1024)] {
|
||||
let attributes = BTreeMap::from([
|
||||
("service.name".to_owned(), "benchmark".to_owned()),
|
||||
("payload".to_owned(), "x".repeat(attribute_bytes)),
|
||||
]);
|
||||
let owned = Box::new(attributes.clone());
|
||||
let shared = Shared::new(attributes);
|
||||
let case = format!("{attribute_bytes}B_{spans}_spans");
|
||||
group.throughput(Throughput::Elements(spans as u64));
|
||||
group.bench_with_input(BenchmarkId::new("owned", &case), &owned, |b, resource| {
|
||||
b.iter(|| black_box(fanout(black_box(resource), spans)));
|
||||
});
|
||||
group.bench_with_input(BenchmarkId::new("shared", &case), &shared, |b, resource| {
|
||||
b.iter(|| black_box(fanout(black_box(resource), spans)));
|
||||
});
|
||||
}
|
||||
group.finish();
|
||||
}
|
||||
|
||||
criterion_group! {
|
||||
name = benches;
|
||||
config = Criterion::default()
|
||||
.sample_size(20)
|
||||
.warm_up_time(Duration::from_secs(1))
|
||||
.measurement_time(Duration::from_secs(2));
|
||||
targets = resource_fanout
|
||||
}
|
||||
criterion_main!(benches);
|
||||
13
litellm-rust/crates/traces/query/span_error.sql
Normal file
13
litellm-rust/crates/traces/query/span_error.sql
Normal file
|
|
@ -0,0 +1,13 @@
|
|||
SELECT SpanId AS span_id,
|
||||
substringUTF8(StatusMessage, {error_offset:UInt64} + 1, 16384) AS message,
|
||||
lengthUTF8(StatusMessage) AS total_chars,
|
||||
hex(SHA256(StatusMessage)) AS version
|
||||
FROM otel_traces
|
||||
WHERE TraceId = {trace_id:String} AND SpanId = {span_id:String}
|
||||
AND (empty({team_ids:Array(String)}) OR TeamId IN {team_ids:Array(String)})
|
||||
AND ({api_key_hash:String} = '' OR ApiKeyHash = {api_key_hash:String})
|
||||
AND ({trace_ref:String} = '' OR
|
||||
hex(SHA256(concat(TeamId, char(0), ApiKeyHash, char(0), TraceId))) = {trace_ref:String})
|
||||
AND ({error_version:String} = '' OR hex(SHA256(StatusMessage)) = {error_version:String})
|
||||
ORDER BY Timestamp, EngineReceivedMs, StatusMessage
|
||||
LIMIT 1
|
||||
|
|
@ -1,6 +1,7 @@
|
|||
SELECT o.SpanId AS span_id, o.ParentSpanId AS parent_span_id, o.SpanName AS name,
|
||||
o.ObservationType AS type, o.AgentName AS agent, o.StatusCode AS status,
|
||||
o.StatusMessage AS status_message,
|
||||
substringUTF8(o.StatusMessage, 1, 128) AS status_message,
|
||||
lengthUTF8(o.StatusMessage) > 128 AS error_truncated,
|
||||
toUnixTimestamp64Nano(o.Timestamp) AS start_ns, o.Duration AS duration_ns,
|
||||
o.ServiceName AS service, o.InputPreview AS input_preview, o.Model AS model,
|
||||
o.InputTokens AS input_tokens, o.OutputTokens AS output_tokens,
|
||||
|
|
@ -12,5 +13,5 @@ WHERE o.TraceId = {trace_id:String}
|
|||
AND ({api_key_hash:String} = '' OR o.ApiKeyHash = {api_key_hash:String})
|
||||
AND ({trace_ref:String} = '' OR
|
||||
hex(SHA256(concat(o.TeamId, char(0), o.ApiKeyHash, char(0), o.TraceId))) = {trace_ref:String})
|
||||
ORDER BY o.Timestamp
|
||||
ORDER BY o.Timestamp, o.EngineReceivedMs, o.StatusMessage
|
||||
LIMIT 1 BY o.SpanId
|
||||
|
|
|
|||
|
|
@ -1,37 +1,7 @@
|
|||
#[derive(Debug, thiserror::Error)]
|
||||
pub enum Error {
|
||||
#[error("invalid ClickHouse insert row")]
|
||||
InvalidRow,
|
||||
#[error("invalid ClickHouse insert table")]
|
||||
InvalidTable,
|
||||
#[error("invalid ClickHouse HTTP URL")]
|
||||
InvalidUrl,
|
||||
#[error("database must be a nonempty SQL identifier and retention must be positive")]
|
||||
InvalidSchema,
|
||||
#[error("SQL query must not be empty")]
|
||||
EmptySql,
|
||||
#[error("unknown ClickHouse read query")]
|
||||
InvalidQuery,
|
||||
#[error("ClickHouse query failed with HTTP status {0}")]
|
||||
QueryFailed(u16),
|
||||
#[error("ClickHouse insert failed with HTTP status {0}")]
|
||||
InsertFailed(u16),
|
||||
#[error("ClickHouse insert exceeds the encoded size limit")]
|
||||
InsertTooLarge,
|
||||
#[error("ClickHouse schema setup failed with HTTP status {0}")]
|
||||
SchemaFailed(u16),
|
||||
#[error("ClickHouse query exceeded the response size limit")]
|
||||
ResponseTooLarge,
|
||||
#[error("ClickHouse returned an invalid or failed JSON query response")]
|
||||
InvalidResponse,
|
||||
#[error("ClickHouse query transport failed")]
|
||||
Transport,
|
||||
}
|
||||
|
||||
#[derive(Debug, thiserror::Error)]
|
||||
pub enum DecodeError {
|
||||
#[error("invalid OTLP trace payload")]
|
||||
InvalidPayload,
|
||||
#[error("OTLP trace payload exceeds the decompressed size limit")]
|
||||
#[error("OTLP trace payload exceeds the decoding budget")]
|
||||
TooLarge,
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,4 +1,10 @@
|
|||
use std::{collections::BTreeMap, io::Write, time::Duration};
|
||||
use std::{
|
||||
borrow::Cow,
|
||||
collections::BTreeMap,
|
||||
io::{BufWriter, Write},
|
||||
};
|
||||
|
||||
use serde::{Serialize, Serializer, ser::SerializeMap};
|
||||
|
||||
use flate2::{Compression, write::GzEncoder};
|
||||
use litellm_http::Client;
|
||||
|
|
@ -6,10 +12,11 @@ use serde_json::Value;
|
|||
use sha2::{Digest, Sha256};
|
||||
use time::{OffsetDateTime, format_description::well_known::Rfc3339};
|
||||
|
||||
use crate::{Connection, Error};
|
||||
use crate::{Connection, Error, Shared};
|
||||
|
||||
const MAX_INSERT_BYTES: usize = 64 * 1024 * 1024;
|
||||
const INSERT_TIMEOUT: Duration = Duration::from_secs(30);
|
||||
|
||||
pub type InsertRow = BTreeMap<String, Shared<Value>>;
|
||||
|
||||
pub enum InsertTable {
|
||||
OtelTraces,
|
||||
|
|
@ -39,125 +46,176 @@ pub async fn insert_rows(
|
|||
database: &str,
|
||||
table: InsertTable,
|
||||
rows: Vec<BTreeMap<String, Value>>,
|
||||
) -> Result<(), Error> {
|
||||
insert_shared_rows(client, connection, database, table, shared_rows(rows)).await
|
||||
}
|
||||
|
||||
pub async fn insert_shared_rows(
|
||||
client: &Client,
|
||||
connection: &Connection,
|
||||
database: &str,
|
||||
table: InsertTable,
|
||||
rows: Vec<InsertRow>,
|
||||
) -> Result<(), Error> {
|
||||
if rows.is_empty() {
|
||||
return Ok(());
|
||||
}
|
||||
let token = format!(
|
||||
"{:x}",
|
||||
Sha256::digest(encode_rows_with_limit(rows.clone(), MAX_INSERT_BYTES)?.as_bytes())
|
||||
);
|
||||
let received_ms = OffsetDateTime::now_utc().unix_timestamp_nanos() / 1_000_000;
|
||||
let rows = rows
|
||||
.into_iter()
|
||||
let received_ms = (OffsetDateTime::now_utc().unix_timestamp_nanos() / 1_000_000) as u64;
|
||||
let (token, body) = prepare_insert(&rows, received_ms, MAX_INSERT_BYTES)?;
|
||||
litellm_storage_clickhouse::insert_compressed_rows(
|
||||
client,
|
||||
connection,
|
||||
database,
|
||||
table.name(),
|
||||
&token,
|
||||
body,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
fn shared_rows(rows: Vec<BTreeMap<String, Value>>) -> Vec<InsertRow> {
|
||||
rows.into_iter()
|
||||
.map(|row| {
|
||||
row.into_iter()
|
||||
.filter(|(key, _)| key != "EngineReceivedMs")
|
||||
.chain(std::iter::once((
|
||||
"EngineReceivedMs".to_owned(),
|
||||
Value::from(received_ms as u64),
|
||||
)))
|
||||
.map(|(key, value)| (key, Shared::new(value)))
|
||||
.collect()
|
||||
})
|
||||
.collect();
|
||||
let encoded = encode_rows_with_limit(rows, MAX_INSERT_BYTES)?;
|
||||
let mut encoder = GzEncoder::new(Vec::new(), Compression::default());
|
||||
encoder
|
||||
.write_all(encoded.as_bytes())
|
||||
.map_err(|_| Error::InvalidRow)?;
|
||||
let body = encoder.finish().map_err(|_| Error::InvalidRow)?;
|
||||
let mut url = connection.url().clone();
|
||||
let existing_pairs: Vec<(String, String)> = url
|
||||
.query_pairs()
|
||||
.filter(|(key, _)| {
|
||||
!matches!(
|
||||
key.as_ref(),
|
||||
"query"
|
||||
| "async_insert"
|
||||
| "async_insert_deduplicate"
|
||||
| "wait_for_async_insert"
|
||||
| "input_format_skip_unknown_fields"
|
||||
| "date_time_input_format"
|
||||
)
|
||||
})
|
||||
.map(|(key, value)| (key.into_owned(), value.into_owned()))
|
||||
.collect();
|
||||
url.query_pairs_mut()
|
||||
.clear()
|
||||
.extend_pairs(existing_pairs)
|
||||
.append_pair(
|
||||
"query",
|
||||
&format!(
|
||||
"INSERT INTO `{database}`.{} FORMAT JSONEachRow",
|
||||
table.name()
|
||||
),
|
||||
)
|
||||
.append_pair("insert_deduplication_token", &token)
|
||||
.append_pair("async_insert", "1")
|
||||
.append_pair("async_insert_deduplicate", "1")
|
||||
.append_pair("wait_for_async_insert", "1")
|
||||
.append_pair("input_format_skip_unknown_fields", "0")
|
||||
.append_pair("date_time_input_format", "best_effort");
|
||||
let response = client
|
||||
.post(url)
|
||||
.timeout(INSERT_TIMEOUT)
|
||||
.header("Content-Encoding", "gzip")
|
||||
.body(body)
|
||||
.send()
|
||||
.await
|
||||
.map_err(|_| Error::Transport)?;
|
||||
if !response.status().is_success() {
|
||||
return Err(Error::InsertFailed(response.status().as_u16()));
|
||||
}
|
||||
Ok(())
|
||||
.collect()
|
||||
}
|
||||
|
||||
pub fn encode_rows(rows: Vec<BTreeMap<String, Value>>) -> Result<String, Error> {
|
||||
encode_rows_with_limit(rows, usize::MAX)
|
||||
}
|
||||
|
||||
fn encode_rows_with_limit(
|
||||
rows: Vec<BTreeMap<String, Value>>,
|
||||
limit: usize,
|
||||
) -> Result<String, Error> {
|
||||
let mut body = Vec::new();
|
||||
for row in rows {
|
||||
let encoded = row
|
||||
.into_iter()
|
||||
.map(|(name, value)| insert_value(&name, value).map(|value| (name, value)))
|
||||
.collect::<Result<BTreeMap<_, _>, _>>()?;
|
||||
let record = serde_json::to_vec(&encoded).map_err(|_| Error::InvalidRow)?;
|
||||
let size = body
|
||||
.len()
|
||||
.checked_add(record.len())
|
||||
.and_then(|size| size.checked_add(usize::from(!body.is_empty())))
|
||||
.ok_or(Error::InsertTooLarge)?;
|
||||
if size > limit {
|
||||
return Err(Error::InsertTooLarge);
|
||||
}
|
||||
if !body.is_empty() {
|
||||
body.push(b'\n');
|
||||
}
|
||||
body.extend_from_slice(&record);
|
||||
}
|
||||
let body = write_rows(&shared_rows(rows), None, Vec::new(), usize::MAX)?;
|
||||
String::from_utf8(body).map_err(|_| Error::InvalidRow)
|
||||
}
|
||||
|
||||
fn insert_value(name: &str, value: Value) -> Result<Value, Error> {
|
||||
fn prepare_insert(
|
||||
rows: &[InsertRow],
|
||||
received_ms: u64,
|
||||
limit: usize,
|
||||
) -> Result<(String, Vec<u8>), Error> {
|
||||
let hash = write_rows(rows, None, HashWriter(Sha256::new()), limit)?;
|
||||
let token = format!("{:x}", hash.0.finalize());
|
||||
let encoder = write_rows(
|
||||
rows,
|
||||
Some(received_ms),
|
||||
BufWriter::new(GzEncoder::new(Vec::new(), Compression::default())),
|
||||
limit,
|
||||
)?;
|
||||
let body = encoder
|
||||
.into_inner()
|
||||
.map_err(|_| Error::InvalidRow)?
|
||||
.finish()
|
||||
.map_err(|_| Error::InvalidRow)?;
|
||||
Ok((token, body))
|
||||
}
|
||||
|
||||
struct HashWriter(Sha256);
|
||||
|
||||
impl Write for HashWriter {
|
||||
fn write(&mut self, bytes: &[u8]) -> std::io::Result<usize> {
|
||||
self.0.update(bytes);
|
||||
Ok(bytes.len())
|
||||
}
|
||||
|
||||
fn flush(&mut self) -> std::io::Result<()> {
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
struct LimitedWriter<W> {
|
||||
inner: W,
|
||||
remaining: usize,
|
||||
exceeded: bool,
|
||||
}
|
||||
|
||||
impl<W: Write> Write for LimitedWriter<W> {
|
||||
fn write(&mut self, bytes: &[u8]) -> std::io::Result<usize> {
|
||||
if bytes.len() > self.remaining {
|
||||
self.exceeded = true;
|
||||
return Err(std::io::Error::other(Error::InsertTooLarge));
|
||||
}
|
||||
let written = self.inner.write(bytes)?;
|
||||
self.remaining -= written;
|
||||
Ok(written)
|
||||
}
|
||||
|
||||
fn flush(&mut self) -> std::io::Result<()> {
|
||||
self.inner.flush()
|
||||
}
|
||||
}
|
||||
|
||||
fn write_rows<W: Write>(
|
||||
rows: &[InsertRow],
|
||||
received_ms: Option<u64>,
|
||||
writer: W,
|
||||
limit: usize,
|
||||
) -> Result<W, Error> {
|
||||
let mut writer = LimitedWriter {
|
||||
inner: writer,
|
||||
remaining: limit,
|
||||
exceeded: false,
|
||||
};
|
||||
for (index, row) in rows.iter().enumerate() {
|
||||
let result = (|| {
|
||||
if index != 0 {
|
||||
writer.write_all(b"\n").map_err(serde_json::Error::io)?;
|
||||
}
|
||||
serde_json::to_writer(&mut writer, &EncodedRow { row, received_ms })
|
||||
})();
|
||||
if result.is_err() {
|
||||
return Err(if writer.exceeded {
|
||||
Error::InsertTooLarge
|
||||
} else {
|
||||
Error::InvalidRow
|
||||
});
|
||||
}
|
||||
}
|
||||
Ok(writer.inner)
|
||||
}
|
||||
|
||||
struct EncodedRow<'a> {
|
||||
row: &'a InsertRow,
|
||||
received_ms: Option<u64>,
|
||||
}
|
||||
|
||||
impl Serialize for EncodedRow<'_> {
|
||||
fn serialize<S: Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
|
||||
let mut map = serializer.serialize_map(None)?;
|
||||
let mut received_ms = self.received_ms;
|
||||
for (name, value) in self.row {
|
||||
if name.as_str() >= "EngineReceivedMs"
|
||||
&& let Some(timestamp) = received_ms.take()
|
||||
{
|
||||
map.serialize_entry("EngineReceivedMs", ×tamp)?;
|
||||
}
|
||||
if name == "EngineReceivedMs" && self.received_ms.is_some() {
|
||||
continue;
|
||||
}
|
||||
let value = insert_value(name, value).map_err(serde::ser::Error::custom)?;
|
||||
map.serialize_entry(name, &value)?;
|
||||
}
|
||||
if let Some(timestamp) = received_ms {
|
||||
map.serialize_entry("EngineReceivedMs", ×tamp)?;
|
||||
}
|
||||
map.end()
|
||||
}
|
||||
}
|
||||
|
||||
fn insert_value<'a>(name: &str, value: &'a Value) -> Result<Cow<'a, Value>, Error> {
|
||||
let multiplier = match name {
|
||||
"Timestamp" => 1,
|
||||
"start_time" | "end_time" | "completion_start_time" => 1_000_000,
|
||||
_ => return Ok(value),
|
||||
_ => return Ok(Cow::Borrowed(value)),
|
||||
};
|
||||
if name == "completion_start_time" && value.is_null() {
|
||||
return Ok(value);
|
||||
return Ok(Cow::Borrowed(value));
|
||||
}
|
||||
let timestamp = value.as_i64().ok_or(Error::InvalidRow)?;
|
||||
let datetime = OffsetDateTime::from_unix_timestamp_nanos(i128::from(timestamp) * multiplier)
|
||||
.map_err(|_| Error::InvalidRow)?;
|
||||
datetime
|
||||
.format(&Rfc3339)
|
||||
.map(Value::String)
|
||||
.map(|value| Cow::Owned(Value::String(value)))
|
||||
.map_err(|_| Error::InvalidRow)
|
||||
}
|
||||
|
||||
|
|
@ -168,20 +226,83 @@ mod tests {
|
|||
use rstest::rstest;
|
||||
use serde_json::json;
|
||||
|
||||
use super::encode_rows_with_limit;
|
||||
use super::{shared_rows, write_rows};
|
||||
use crate::Error;
|
||||
|
||||
#[rstest]
|
||||
fn encoded_limit_counts_utf8_bytes_across_rows() {
|
||||
let rows = vec![
|
||||
let rows = shared_rows(vec![
|
||||
BTreeMap::from([("Input".to_owned(), json!("雪"))]),
|
||||
BTreeMap::from([("Input".to_owned(), json!("雪"))]),
|
||||
];
|
||||
let encoded = encode_rows_with_limit(rows.clone(), usize::MAX).expect("valid rows");
|
||||
]);
|
||||
let encoded = write_rows(&rows, None, Vec::new(), usize::MAX).expect("valid rows");
|
||||
|
||||
assert!(encode_rows_with_limit(rows.clone(), encoded.len()).is_ok());
|
||||
assert!(write_rows(&rows, None, Vec::new(), encoded.len()).is_ok());
|
||||
assert!(matches!(
|
||||
encode_rows_with_limit(rows, encoded.len() - 1),
|
||||
write_rows(&rows, None, Vec::new(), encoded.len() - 1),
|
||||
Err(Error::InsertTooLarge)
|
||||
));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::absent(None)]
|
||||
#[case::submitted(Some(123))]
|
||||
fn streamed_insert_preserves_token_and_stamps_receive_time(#[case] submitted: Option<u64>) {
|
||||
use flate2::read::GzDecoder;
|
||||
use sha2::{Digest, Sha256};
|
||||
use std::io::Read;
|
||||
let mut row = BTreeMap::from([
|
||||
("ApiKeyHash".into(), json!("key")),
|
||||
("ResourceAttributes".into(), json!({"message": "雪\n\""})),
|
||||
("Timestamp".into(), json!(1_234_567_890)),
|
||||
]);
|
||||
if let Some(value) = submitted {
|
||||
row.insert("EngineReceivedMs".into(), json!(value));
|
||||
}
|
||||
let legacy = match submitted {
|
||||
Some(_) => {
|
||||
"{\"ApiKeyHash\":\"key\",\"EngineReceivedMs\":123,\"ResourceAttributes\":{\"message\":\"雪\\n\\\"\"},\"Timestamp\":\"1970-01-01T00:00:01.23456789Z\"}"
|
||||
}
|
||||
None => {
|
||||
"{\"ApiKeyHash\":\"key\",\"ResourceAttributes\":{\"message\":\"雪\\n\\\"\"},\"Timestamp\":\"1970-01-01T00:00:01.23456789Z\"}"
|
||||
}
|
||||
};
|
||||
let rows = shared_rows(vec![row.clone(), row]);
|
||||
let (token, body) = super::prepare_insert(&rows, 456, 4096).unwrap();
|
||||
assert_eq!(
|
||||
token,
|
||||
format!("{:x}", Sha256::digest(format!("{legacy}\n{legacy}")))
|
||||
);
|
||||
let mut decoded = String::new();
|
||||
GzDecoder::new(body.as_slice())
|
||||
.read_to_string(&mut decoded)
|
||||
.unwrap();
|
||||
let expected = json!({
|
||||
"ApiKeyHash": "key", "EngineReceivedMs": 456,
|
||||
"ResourceAttributes": {"message": "雪\n\""},
|
||||
"Timestamp": "1970-01-01T00:00:01.23456789Z",
|
||||
});
|
||||
assert_eq!(
|
||||
decoded
|
||||
.lines()
|
||||
.map(|line| serde_json::from_str::<serde_json::Value>(line).unwrap())
|
||||
.collect::<Vec<_>>(),
|
||||
vec![expected.clone(), expected]
|
||||
);
|
||||
assert_eq!(
|
||||
rows[0]
|
||||
.get("EngineReceivedMs")
|
||||
.map(|value| value.as_u64().unwrap()),
|
||||
submitted
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn stamped_insert_enforces_the_encoded_limit() {
|
||||
let rows = shared_rows(vec![BTreeMap::new()]);
|
||||
assert!(super::prepare_insert(&rows, 1, 22).is_ok());
|
||||
assert!(matches!(
|
||||
super::prepare_insert(&rows, 1, 21),
|
||||
Err(Error::InsertTooLarge)
|
||||
));
|
||||
}
|
||||
|
|
|
|||
|
|
@ -2,89 +2,13 @@ mod error;
|
|||
mod insert;
|
||||
mod otlp;
|
||||
mod schema;
|
||||
mod shared;
|
||||
mod sql;
|
||||
|
||||
pub use error::{DecodeError, Error};
|
||||
pub use insert::{InsertTable, encode_rows, insert_rows};
|
||||
pub use error::DecodeError;
|
||||
pub use insert::{InsertRow, InsertTable, encode_rows, insert_rows, insert_shared_rows};
|
||||
pub use litellm_storage_clickhouse::{Connection, Error, Parameter, execute_read};
|
||||
pub use otlp::{DecodedSpan, decode_otlp};
|
||||
pub use schema::{ensure_schema, schema_statements};
|
||||
pub use sql::{LensQuery, Parameter, ReadQuery, execute_named_read, execute_read};
|
||||
use url::Url;
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct Connection {
|
||||
url: Url,
|
||||
}
|
||||
|
||||
impl Connection {
|
||||
pub fn parse(value: &str) -> Result<Self, Error> {
|
||||
let url = Url::parse(value).map_err(|_| Error::InvalidUrl)?;
|
||||
if !matches!(url.scheme(), "http" | "https") || url.host().is_none() {
|
||||
return Err(Error::InvalidUrl);
|
||||
}
|
||||
Ok(Self { url })
|
||||
}
|
||||
|
||||
pub fn configured(
|
||||
url: &str,
|
||||
database: &str,
|
||||
user: &str,
|
||||
password: &str,
|
||||
) -> Result<Self, Error> {
|
||||
let mut connection = Self::parse(url)?;
|
||||
connection
|
||||
.url
|
||||
.set_username(user)
|
||||
.map_err(|_| Error::InvalidUrl)?;
|
||||
connection
|
||||
.url
|
||||
.set_password(Some(password))
|
||||
.map_err(|_| Error::InvalidUrl)?;
|
||||
let pairs: Vec<_> = connection
|
||||
.url
|
||||
.query_pairs()
|
||||
.filter(|(key, _)| !matches!(key.as_ref(), "database" | "user" | "password"))
|
||||
.map(|(key, value)| (key.into_owned(), value.into_owned()))
|
||||
.collect();
|
||||
connection
|
||||
.url
|
||||
.query_pairs_mut()
|
||||
.clear()
|
||||
.extend_pairs(pairs)
|
||||
.append_pair("database", database);
|
||||
Ok(connection)
|
||||
}
|
||||
|
||||
pub fn writer(url: &str) -> Result<Self, Error> {
|
||||
let mut connection = Self::parse(url)?;
|
||||
let pairs: Vec<_> = connection
|
||||
.url
|
||||
.query_pairs()
|
||||
.filter(|(key, _)| !matches!(key.as_ref(), "database" | "readonly" | "query"))
|
||||
.map(|(key, value)| (key.into_owned(), value.into_owned()))
|
||||
.collect();
|
||||
connection.url.query_pairs_mut().clear().extend_pairs(pairs);
|
||||
Ok(connection)
|
||||
}
|
||||
|
||||
pub fn reader(url: &str, database: &str) -> Result<Self, Error> {
|
||||
let mut connection = Self::parse(url)?;
|
||||
let pairs: Vec<_> = connection
|
||||
.url
|
||||
.query_pairs()
|
||||
.filter(|(key, _)| key != "database")
|
||||
.map(|(key, value)| (key.into_owned(), value.into_owned()))
|
||||
.collect();
|
||||
connection
|
||||
.url
|
||||
.query_pairs_mut()
|
||||
.clear()
|
||||
.extend_pairs(pairs)
|
||||
.append_pair("database", database);
|
||||
Ok(connection)
|
||||
}
|
||||
|
||||
pub fn url(&self) -> &Url {
|
||||
&self.url
|
||||
}
|
||||
}
|
||||
pub use shared::{Shared, SharedIdentity};
|
||||
pub use sql::{LensQuery, ReadQuery, execute_named_read};
|
||||
|
|
|
|||
|
|
@ -1,221 +0,0 @@
|
|||
use std::{collections::BTreeMap, io::Read};
|
||||
|
||||
use base64::Engine;
|
||||
use flate2::read::GzDecoder;
|
||||
use opentelemetry_proto::tonic::{
|
||||
collector::trace::v1::ExportTraceServiceRequest,
|
||||
common::v1::{AnyValue, KeyValue, any_value::Value as AttributeValue},
|
||||
trace::v1::{Span, span::SpanKind, status::StatusCode},
|
||||
};
|
||||
use prost::Message;
|
||||
use serde::Serialize;
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::DecodeError;
|
||||
|
||||
#[derive(Serialize)]
|
||||
pub struct DecodedEvent {
|
||||
pub name: String,
|
||||
pub attributes: BTreeMap<String, String>,
|
||||
}
|
||||
|
||||
#[derive(Serialize)]
|
||||
pub struct DecodedSpan {
|
||||
pub trace_id: String,
|
||||
pub span_id: String,
|
||||
pub parent_span_id: String,
|
||||
pub trace_state: String,
|
||||
pub name: String,
|
||||
pub kind: String,
|
||||
pub resource_attributes: BTreeMap<String, String>,
|
||||
pub scope_name: String,
|
||||
pub scope_version: String,
|
||||
pub attributes: BTreeMap<String, String>,
|
||||
pub start_ns: u64,
|
||||
pub end_ns: u64,
|
||||
pub status_code: String,
|
||||
pub status_message: String,
|
||||
pub events: Vec<DecodedEvent>,
|
||||
}
|
||||
|
||||
pub fn decode_otlp(
|
||||
body: &[u8],
|
||||
content_type: Option<&str>,
|
||||
content_encoding: Option<&str>,
|
||||
max_decompressed_bytes: usize,
|
||||
) -> Result<Vec<DecodedSpan>, DecodeError> {
|
||||
let payload = if content_encoding == Some("gzip") || body.starts_with(&[0x1f, 0x8b]) {
|
||||
let limit = u64::try_from(max_decompressed_bytes).map_err(|_| DecodeError::TooLarge)?;
|
||||
let mut decoded = Vec::new();
|
||||
GzDecoder::new(body)
|
||||
.take(limit + 1)
|
||||
.read_to_end(&mut decoded)
|
||||
.map_err(|_| DecodeError::InvalidPayload)?;
|
||||
decoded
|
||||
} else {
|
||||
body.to_vec()
|
||||
};
|
||||
if payload.len() > max_decompressed_bytes {
|
||||
return Err(DecodeError::TooLarge);
|
||||
}
|
||||
let request = if content_type.is_some_and(|value| value.contains("json")) {
|
||||
let value: Value =
|
||||
serde_json::from_slice(&payload).map_err(|_| DecodeError::InvalidPayload)?;
|
||||
serde_json::from_value(normalize_json_ids(value)?)
|
||||
.map_err(|_| DecodeError::InvalidPayload)?
|
||||
} else {
|
||||
ExportTraceServiceRequest::decode(payload.as_slice())
|
||||
.map_err(|_| DecodeError::InvalidPayload)?
|
||||
};
|
||||
Ok(request
|
||||
.resource_spans
|
||||
.into_iter()
|
||||
.flat_map(|resource_spans| {
|
||||
let resource_attributes = attributes(
|
||||
resource_spans
|
||||
.resource
|
||||
.map(|resource| resource.attributes)
|
||||
.unwrap_or_default(),
|
||||
);
|
||||
resource_spans
|
||||
.scope_spans
|
||||
.into_iter()
|
||||
.flat_map(move |scope_spans| {
|
||||
let scope = scope_spans.scope.unwrap_or_default();
|
||||
let resource_attributes = resource_attributes.clone();
|
||||
scope_spans.spans.into_iter().map(move |span| {
|
||||
decoded_span(span, &resource_attributes, &scope.name, &scope.version)
|
||||
})
|
||||
})
|
||||
})
|
||||
.collect())
|
||||
}
|
||||
|
||||
fn normalize_json_ids(value: Value) -> Result<Value, DecodeError> {
|
||||
match value {
|
||||
Value::Object(fields) => fields
|
||||
.into_iter()
|
||||
.map(|(name, value)| {
|
||||
let normalized = if matches!(name.as_str(), "traceId" | "spanId" | "parentSpanId") {
|
||||
let encoded = value.as_str().ok_or(DecodeError::InvalidPayload)?;
|
||||
let bytes = base64::engine::general_purpose::STANDARD
|
||||
.decode(encoded)
|
||||
.map_err(|_| DecodeError::InvalidPayload)?;
|
||||
Value::String(hex_bytes(&bytes))
|
||||
} else if name == "kind" && value.is_string() {
|
||||
let kind = SpanKind::from_str_name(value.as_str().unwrap_or_default())
|
||||
.ok_or(DecodeError::InvalidPayload)?;
|
||||
Value::from(kind as i32)
|
||||
} else if name == "code" && value.is_string() {
|
||||
let code = StatusCode::from_str_name(value.as_str().unwrap_or_default())
|
||||
.ok_or(DecodeError::InvalidPayload)?;
|
||||
Value::from(code as i32)
|
||||
} else {
|
||||
normalize_json_ids(value)?
|
||||
};
|
||||
Ok((name, normalized))
|
||||
})
|
||||
.collect::<Result<serde_json::Map<_, _>, _>>()
|
||||
.map(Value::Object),
|
||||
Value::Array(values) => values
|
||||
.into_iter()
|
||||
.map(normalize_json_ids)
|
||||
.collect::<Result<Vec<_>, _>>()
|
||||
.map(Value::Array),
|
||||
value => Ok(value),
|
||||
}
|
||||
}
|
||||
|
||||
fn hex_bytes(bytes: &[u8]) -> String {
|
||||
bytes.iter().map(|byte| format!("{byte:02x}")).collect()
|
||||
}
|
||||
|
||||
fn decoded_span(
|
||||
span: Span,
|
||||
resource_attributes: &BTreeMap<String, String>,
|
||||
scope_name: &str,
|
||||
scope_version: &str,
|
||||
) -> DecodedSpan {
|
||||
let status = span.status.unwrap_or_default();
|
||||
DecodedSpan {
|
||||
trace_id: hex_bytes(&span.trace_id),
|
||||
span_id: hex_bytes(&span.span_id),
|
||||
parent_span_id: hex_bytes(&span.parent_span_id),
|
||||
trace_state: span.trace_state,
|
||||
name: span.name,
|
||||
kind: SpanKind::try_from(span.kind)
|
||||
.unwrap_or(SpanKind::Unspecified)
|
||||
.as_str_name()
|
||||
.to_owned(),
|
||||
resource_attributes: resource_attributes.clone(),
|
||||
scope_name: scope_name.to_owned(),
|
||||
scope_version: scope_version.to_owned(),
|
||||
attributes: attributes(span.attributes),
|
||||
start_ns: span.start_time_unix_nano,
|
||||
end_ns: span.end_time_unix_nano,
|
||||
status_code: StatusCode::try_from(status.code)
|
||||
.unwrap_or(StatusCode::Unset)
|
||||
.as_str_name()
|
||||
.to_owned(),
|
||||
status_message: status.message,
|
||||
events: span
|
||||
.events
|
||||
.into_iter()
|
||||
.map(|event| DecodedEvent {
|
||||
name: event.name,
|
||||
attributes: attributes(event.attributes),
|
||||
})
|
||||
.collect(),
|
||||
}
|
||||
}
|
||||
|
||||
fn attributes(values: Vec<KeyValue>) -> BTreeMap<String, String> {
|
||||
values
|
||||
.into_iter()
|
||||
.map(|entry| {
|
||||
(
|
||||
entry.key,
|
||||
entry.value.as_ref().map(attribute_text).unwrap_or_default(),
|
||||
)
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn attribute_text(value: &AnyValue) -> String {
|
||||
match value.value.as_ref() {
|
||||
Some(AttributeValue::StringValue(value)) => value.clone(),
|
||||
Some(AttributeValue::BoolValue(value)) => value.to_string(),
|
||||
Some(AttributeValue::IntValue(value)) => value.to_string(),
|
||||
Some(AttributeValue::DoubleValue(value)) => {
|
||||
serde_json::to_string(value).unwrap_or_default()
|
||||
}
|
||||
Some(AttributeValue::BytesValue(value)) => String::from_utf8_lossy(value).into_owned(),
|
||||
Some(AttributeValue::ArrayValue(value)) => format!(
|
||||
"[{}]",
|
||||
value
|
||||
.values
|
||||
.iter()
|
||||
.map(|value| serde_json::to_string(&attribute_text(value)).unwrap_or_default())
|
||||
.collect::<Vec<_>>()
|
||||
.join(", ")
|
||||
),
|
||||
Some(AttributeValue::KvlistValue(value)) => format!(
|
||||
"{{{}}}",
|
||||
value
|
||||
.values
|
||||
.iter()
|
||||
.map(|entry| format!(
|
||||
"{}: {}",
|
||||
serde_json::to_string(&entry.key).unwrap_or_default(),
|
||||
serde_json::to_string(
|
||||
&entry.value.as_ref().map(attribute_text).unwrap_or_default()
|
||||
)
|
||||
.unwrap_or_default()
|
||||
))
|
||||
.collect::<Vec<_>>()
|
||||
.join(", ")
|
||||
),
|
||||
Some(AttributeValue::StringValueStrindex(value)) => value.to_string(),
|
||||
None => String::new(),
|
||||
}
|
||||
}
|
||||
101
litellm-rust/crates/traces/src/otlp/attributes.rs
Normal file
101
litellm-rust/crates/traces/src/otlp/attributes.rs
Normal file
|
|
@ -0,0 +1,101 @@
|
|||
use std::{collections::BTreeMap, io::Write};
|
||||
|
||||
use opentelemetry_proto::tonic::common::v1::{
|
||||
AnyValue, KeyValue, any_value::Value as AttributeValue,
|
||||
};
|
||||
use serde::{
|
||||
Serialize, Serializer,
|
||||
ser::{SerializeMap, SerializeSeq},
|
||||
};
|
||||
|
||||
use super::limits::{Budget, MAX_ATTRIBUTES};
|
||||
use crate::DecodeError;
|
||||
|
||||
struct AttributeWriter<'a> {
|
||||
body: Vec<u8>,
|
||||
budget: &'a mut Budget,
|
||||
}
|
||||
|
||||
impl Write for AttributeWriter<'_> {
|
||||
fn write(&mut self, bytes: &[u8]) -> std::io::Result<usize> {
|
||||
self.budget
|
||||
.consume(bytes.len())
|
||||
.map_err(std::io::Error::other)?;
|
||||
self.body.extend_from_slice(bytes);
|
||||
Ok(bytes.len())
|
||||
}
|
||||
fn flush(&mut self) -> std::io::Result<()> {
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn attributes(
|
||||
values: Vec<KeyValue>,
|
||||
budget: &mut Budget,
|
||||
) -> Result<BTreeMap<String, String>, DecodeError> {
|
||||
if values.len() > MAX_ATTRIBUTES {
|
||||
return Err(DecodeError::TooLarge);
|
||||
}
|
||||
values
|
||||
.into_iter()
|
||||
.map(|entry| {
|
||||
budget.consume(entry.key.len() + 96)?;
|
||||
let text = match entry.value {
|
||||
Some(AnyValue {
|
||||
value: Some(AttributeValue::StringValue(value)),
|
||||
}) => {
|
||||
budget.consume(value.len())?;
|
||||
value
|
||||
}
|
||||
Some(AnyValue {
|
||||
value: Some(AttributeValue::BytesValue(value)),
|
||||
}) => {
|
||||
budget.consume(value.len().saturating_mul(3))?;
|
||||
String::from_utf8_lossy(&value).into_owned()
|
||||
}
|
||||
value => {
|
||||
let mut writer = AttributeWriter {
|
||||
body: Vec::new(),
|
||||
budget,
|
||||
};
|
||||
serde_json::to_writer(&mut writer, &AttributeJson(value.as_ref()))
|
||||
.map_err(|_| DecodeError::TooLarge)?;
|
||||
String::from_utf8(writer.body).map_err(|_| DecodeError::InvalidPayload)?
|
||||
}
|
||||
};
|
||||
Ok((entry.key, text))
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
struct AttributeJson<'a>(Option<&'a AnyValue>);
|
||||
|
||||
impl Serialize for AttributeJson<'_> {
|
||||
fn serialize<S: Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
|
||||
match self.0.and_then(|value| value.value.as_ref()) {
|
||||
Some(AttributeValue::StringValue(value)) => serializer.serialize_str(value),
|
||||
Some(AttributeValue::BoolValue(value)) => serializer.serialize_bool(*value),
|
||||
Some(AttributeValue::IntValue(value)) => serializer.serialize_i64(*value),
|
||||
Some(AttributeValue::DoubleValue(value)) => serializer.serialize_f64(*value),
|
||||
Some(AttributeValue::BytesValue(value)) => {
|
||||
serializer.serialize_str(&String::from_utf8_lossy(value))
|
||||
}
|
||||
Some(AttributeValue::ArrayValue(value)) => {
|
||||
let mut sequence = serializer.serialize_seq(Some(value.values.len()))?;
|
||||
for entry in &value.values {
|
||||
sequence.serialize_element(&AttributeJson(Some(entry)))?;
|
||||
}
|
||||
sequence.end()
|
||||
}
|
||||
Some(AttributeValue::KvlistValue(value)) => {
|
||||
let mut map = serializer.serialize_map(Some(value.values.len()))?;
|
||||
for entry in &value.values {
|
||||
map.serialize_entry(&entry.key, &AttributeJson(entry.value.as_ref()))?;
|
||||
}
|
||||
map.end()
|
||||
}
|
||||
Some(AttributeValue::StringValueStrindex(value)) => serializer.serialize_i32(*value),
|
||||
None => serializer.serialize_unit(),
|
||||
}
|
||||
}
|
||||
}
|
||||
212
litellm-rust/crates/traces/src/otlp/limits.rs
Normal file
212
litellm-rust/crates/traces/src/otlp/limits.rs
Normal file
|
|
@ -0,0 +1,212 @@
|
|||
use std::fmt;
|
||||
|
||||
use prost::encoding::{DecodeContext, WireType, decode_key, decode_varint, skip_field};
|
||||
use serde::de::{DeserializeSeed, MapAccess, SeqAccess, Visitor};
|
||||
|
||||
use crate::{DecodeError, Shared};
|
||||
|
||||
pub(super) const MAX_DEPTH: usize = 32;
|
||||
pub(super) const MAX_NODES: usize = 65_536;
|
||||
pub(super) const MAX_SPANS: usize = 4_096;
|
||||
pub(super) const MAX_ATTRIBUTES: usize = 256;
|
||||
pub(super) const MAX_EVENTS: usize = 256;
|
||||
pub(super) const MAX_DECODED_SPAN_BYTES: usize = 16 * 1024 * 1024;
|
||||
|
||||
pub(super) fn json_preflight(payload: &[u8]) -> Result<(), DecodeError> {
|
||||
let mut nodes = 0;
|
||||
let mut exceeded = false;
|
||||
let mut decoder = serde_json::Deserializer::from_slice(payload);
|
||||
let result = JsonBudget {
|
||||
nodes: &mut nodes,
|
||||
exceeded: &mut exceeded,
|
||||
depth: 0,
|
||||
}
|
||||
.deserialize(&mut decoder)
|
||||
.and_then(|()| decoder.end());
|
||||
if exceeded {
|
||||
return Err(DecodeError::TooLarge);
|
||||
}
|
||||
result.map_err(|_| DecodeError::InvalidPayload)
|
||||
}
|
||||
|
||||
struct JsonBudget<'a> {
|
||||
nodes: &'a mut usize,
|
||||
exceeded: &'a mut bool,
|
||||
depth: usize,
|
||||
}
|
||||
|
||||
impl<'de> DeserializeSeed<'de> for JsonBudget<'_> {
|
||||
type Value = ();
|
||||
|
||||
fn deserialize<D: serde::Deserializer<'de>>(self, decoder: D) -> Result<(), D::Error> {
|
||||
*self.nodes += 1;
|
||||
if *self.nodes > MAX_NODES || self.depth > MAX_DEPTH {
|
||||
*self.exceeded = true;
|
||||
return Err(serde::de::Error::custom("OTLP structure exceeds budget"));
|
||||
}
|
||||
decoder.deserialize_any(self)
|
||||
}
|
||||
}
|
||||
|
||||
impl<'de> Visitor<'de> for JsonBudget<'_> {
|
||||
type Value = ();
|
||||
|
||||
fn expecting(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
formatter.write_str("OTLP JSON")
|
||||
}
|
||||
fn visit_bool<E: serde::de::Error>(self, _: bool) -> Result<(), E> {
|
||||
Ok(())
|
||||
}
|
||||
fn visit_i64<E: serde::de::Error>(self, _: i64) -> Result<(), E> {
|
||||
Ok(())
|
||||
}
|
||||
fn visit_u64<E: serde::de::Error>(self, _: u64) -> Result<(), E> {
|
||||
Ok(())
|
||||
}
|
||||
fn visit_f64<E: serde::de::Error>(self, _: f64) -> Result<(), E> {
|
||||
Ok(())
|
||||
}
|
||||
fn visit_str<E: serde::de::Error>(self, _: &str) -> Result<(), E> {
|
||||
Ok(())
|
||||
}
|
||||
fn visit_unit<E: serde::de::Error>(self) -> Result<(), E> {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn visit_seq<A: SeqAccess<'de>>(self, mut sequence: A) -> Result<(), A::Error> {
|
||||
while sequence
|
||||
.next_element_seed(JsonBudget {
|
||||
nodes: self.nodes,
|
||||
exceeded: self.exceeded,
|
||||
depth: self.depth + 1,
|
||||
})?
|
||||
.is_some()
|
||||
{}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn visit_map<A: MapAccess<'de>>(self, mut map: A) -> Result<(), A::Error> {
|
||||
while map
|
||||
.next_key_seed(JsonBudget {
|
||||
nodes: self.nodes,
|
||||
exceeded: self.exceeded,
|
||||
depth: self.depth + 1,
|
||||
})?
|
||||
.is_some()
|
||||
{
|
||||
map.next_value_seed(JsonBudget {
|
||||
nodes: self.nodes,
|
||||
exceeded: self.exceeded,
|
||||
depth: self.depth + 1,
|
||||
})?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy)]
|
||||
enum MessageKind {
|
||||
Export,
|
||||
ResourceSpans,
|
||||
Resource,
|
||||
ScopeSpans,
|
||||
Scope,
|
||||
Span,
|
||||
Event,
|
||||
Link,
|
||||
Status,
|
||||
KeyValue,
|
||||
AnyValue,
|
||||
Array,
|
||||
KvList,
|
||||
}
|
||||
|
||||
impl MessageKind {
|
||||
fn child(self, tag: u32) -> Option<Self> {
|
||||
match (self, tag) {
|
||||
(Self::Export, 1) => Some(Self::ResourceSpans),
|
||||
(Self::ResourceSpans, 1) => Some(Self::Resource),
|
||||
(Self::ResourceSpans, 2) => Some(Self::ScopeSpans),
|
||||
(Self::Resource, 1)
|
||||
| (Self::Scope, 3)
|
||||
| (Self::Span, 9)
|
||||
| (Self::Event, 3)
|
||||
| (Self::Link, 4)
|
||||
| (Self::KvList, 1) => Some(Self::KeyValue),
|
||||
(Self::ScopeSpans, 1) => Some(Self::Scope),
|
||||
(Self::ScopeSpans, 2) => Some(Self::Span),
|
||||
(Self::Span, 11) => Some(Self::Event),
|
||||
(Self::Span, 13) => Some(Self::Link),
|
||||
(Self::Span, 15) => Some(Self::Status),
|
||||
(Self::KeyValue, 2) | (Self::Array, 1) => Some(Self::AnyValue),
|
||||
(Self::AnyValue, 5) => Some(Self::Array),
|
||||
(Self::AnyValue, 6) => Some(Self::KvList),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn protobuf_preflight(payload: &[u8]) -> Result<(), DecodeError> {
|
||||
scan_message(payload, MessageKind::Export, 0, &mut 0)
|
||||
}
|
||||
|
||||
fn scan_message(
|
||||
mut payload: &[u8],
|
||||
kind: MessageKind,
|
||||
depth: usize,
|
||||
nodes: &mut usize,
|
||||
) -> Result<(), DecodeError> {
|
||||
if depth > MAX_DEPTH {
|
||||
return Err(DecodeError::TooLarge);
|
||||
}
|
||||
while !payload.is_empty() {
|
||||
*nodes += 1;
|
||||
if *nodes > MAX_NODES {
|
||||
return Err(DecodeError::TooLarge);
|
||||
}
|
||||
let (tag, wire) = decode_key(&mut payload).map_err(|_| DecodeError::InvalidPayload)?;
|
||||
if let (WireType::LengthDelimited, Some(child)) = (wire, kind.child(tag)) {
|
||||
let length = decode_varint(&mut payload).map_err(|_| DecodeError::InvalidPayload)?;
|
||||
let length = usize::try_from(length).map_err(|_| DecodeError::InvalidPayload)?;
|
||||
let (message, rest) = payload
|
||||
.split_at_checked(length)
|
||||
.ok_or(DecodeError::InvalidPayload)?;
|
||||
scan_message(message, child, depth + 1, nodes)?;
|
||||
payload = rest;
|
||||
} else {
|
||||
skip_field(wire, tag, &mut payload, DecodeContext::default())
|
||||
.map_err(|_| DecodeError::InvalidPayload)?;
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub(super) struct Budget {
|
||||
remaining: usize,
|
||||
}
|
||||
|
||||
impl Budget {
|
||||
pub(super) fn new(remaining: usize) -> Self {
|
||||
Self { remaining }
|
||||
}
|
||||
|
||||
pub(super) fn clone_shared<T: Clone>(
|
||||
&mut self,
|
||||
value: &Shared<T>,
|
||||
allocated_bytes: impl FnOnce(&T) -> usize,
|
||||
) -> Result<Shared<T>, DecodeError> {
|
||||
let cloned = value.clone();
|
||||
if !value.shares_storage_with(&cloned) {
|
||||
self.consume(allocated_bytes(value))?;
|
||||
}
|
||||
Ok(cloned)
|
||||
}
|
||||
|
||||
pub(super) fn consume(&mut self, bytes: usize) -> Result<(), DecodeError> {
|
||||
self.remaining = self
|
||||
.remaining
|
||||
.checked_sub(bytes)
|
||||
.ok_or(DecodeError::TooLarge)?;
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
42
litellm-rust/crates/traces/src/otlp/mod.rs
Normal file
42
litellm-rust/crates/traces/src/otlp/mod.rs
Normal file
|
|
@ -0,0 +1,42 @@
|
|||
mod attributes;
|
||||
mod limits;
|
||||
mod span;
|
||||
mod wire;
|
||||
|
||||
use serde::Serialize;
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
use crate::{DecodeError, Shared};
|
||||
|
||||
#[derive(Serialize)]
|
||||
pub struct DecodedEvent {
|
||||
pub name: String,
|
||||
pub attributes: BTreeMap<String, String>,
|
||||
}
|
||||
|
||||
#[derive(Serialize)]
|
||||
pub struct DecodedSpan {
|
||||
pub trace_id: String,
|
||||
pub span_id: String,
|
||||
pub parent_span_id: String,
|
||||
pub trace_state: String,
|
||||
pub name: String,
|
||||
pub kind: String,
|
||||
pub resource_attributes: Shared<BTreeMap<String, String>>,
|
||||
pub scope_name: Shared<String>,
|
||||
pub scope_version: Shared<String>,
|
||||
pub attributes: BTreeMap<String, String>,
|
||||
pub start_ns: u64,
|
||||
pub end_ns: u64,
|
||||
pub status_code: String,
|
||||
pub status_message: String,
|
||||
pub events: Vec<DecodedEvent>,
|
||||
}
|
||||
|
||||
pub fn decode_otlp(
|
||||
body: &[u8],
|
||||
content_type: Option<&str>,
|
||||
) -> Result<Vec<DecodedSpan>, DecodeError> {
|
||||
let request = wire::decode(body, content_type)?;
|
||||
span::flatten(request)
|
||||
}
|
||||
166
litellm-rust/crates/traces/src/otlp/span.rs
Normal file
166
litellm-rust/crates/traces/src/otlp/span.rs
Normal file
|
|
@ -0,0 +1,166 @@
|
|||
use std::collections::BTreeMap;
|
||||
|
||||
use opentelemetry_proto::tonic::{
|
||||
collector::trace::v1::ExportTraceServiceRequest,
|
||||
trace::v1::{ResourceSpans, ScopeSpans, Span, span::SpanKind, status::StatusCode},
|
||||
};
|
||||
|
||||
use super::{
|
||||
DecodedEvent, DecodedSpan,
|
||||
attributes::attributes,
|
||||
limits::{Budget, MAX_ATTRIBUTES, MAX_DECODED_SPAN_BYTES, MAX_EVENTS, MAX_SPANS},
|
||||
};
|
||||
use crate::{DecodeError, Shared};
|
||||
|
||||
pub(super) fn flatten(request: ExportTraceServiceRequest) -> Result<Vec<DecodedSpan>, DecodeError> {
|
||||
let mut budget = Budget::new(MAX_DECODED_SPAN_BYTES);
|
||||
let mut spans = Vec::new();
|
||||
for resource in request.resource_spans {
|
||||
append_resource(resource, &mut budget, &mut spans)?;
|
||||
}
|
||||
Ok(spans)
|
||||
}
|
||||
|
||||
fn append_resource(
|
||||
resource: ResourceSpans,
|
||||
budget: &mut Budget,
|
||||
spans: &mut Vec<DecodedSpan>,
|
||||
) -> Result<(), DecodeError> {
|
||||
let attributes = Shared::new(attributes(
|
||||
resource
|
||||
.resource
|
||||
.map(|resource| resource.attributes)
|
||||
.unwrap_or_default(),
|
||||
budget,
|
||||
)?);
|
||||
for scope in resource.scope_spans {
|
||||
append_scope(scope, &attributes, budget, spans)?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn append_scope(
|
||||
scope_spans: ScopeSpans,
|
||||
resource: &Shared<BTreeMap<String, String>>,
|
||||
budget: &mut Budget,
|
||||
spans: &mut Vec<DecodedSpan>,
|
||||
) -> Result<(), DecodeError> {
|
||||
let scope = scope_spans.scope.unwrap_or_default();
|
||||
if scope.attributes.len() > MAX_ATTRIBUTES {
|
||||
return Err(DecodeError::TooLarge);
|
||||
}
|
||||
budget.consume(scope.name.len() + scope.version.len())?;
|
||||
let scope_name: Shared<String> = scope.name.into();
|
||||
let scope_version: Shared<String> = scope.version.into();
|
||||
for span in scope_spans.spans {
|
||||
if spans.len() >= MAX_SPANS {
|
||||
return Err(DecodeError::TooLarge);
|
||||
}
|
||||
validate_span(&span)?;
|
||||
budget.consume(
|
||||
span.name.len()
|
||||
+ span.trace_state.len()
|
||||
+ span
|
||||
.status
|
||||
.as_ref()
|
||||
.map_or(0, |status| status.message.len())
|
||||
+ size_of::<DecodedSpan>()
|
||||
+ 128,
|
||||
)?;
|
||||
spans.push(decoded_span(
|
||||
span,
|
||||
resource,
|
||||
&scope_name,
|
||||
&scope_version,
|
||||
budget,
|
||||
)?);
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn valid_id(value: &[u8], length: usize) -> bool {
|
||||
value.len() == length && value.iter().any(|byte| *byte != 0)
|
||||
}
|
||||
|
||||
fn validate_span(span: &Span) -> Result<(), DecodeError> {
|
||||
if !valid_id(&span.trace_id, 16)
|
||||
|| !valid_id(&span.span_id, 8)
|
||||
|| (!span.parent_span_id.is_empty() && !valid_id(&span.parent_span_id, 8))
|
||||
|| span.start_time_unix_nano > i64::MAX as u64
|
||||
|| span.end_time_unix_nano > i64::MAX as u64
|
||||
|| span.end_time_unix_nano < span.start_time_unix_nano
|
||||
|| span
|
||||
.links
|
||||
.iter()
|
||||
.any(|link| !valid_id(&link.trace_id, 16) || !valid_id(&link.span_id, 8))
|
||||
{
|
||||
return Err(DecodeError::InvalidPayload);
|
||||
}
|
||||
if span.events.len() > MAX_EVENTS
|
||||
|| span.links.len() > MAX_EVENTS
|
||||
|| span.attributes.len() > MAX_ATTRIBUTES
|
||||
|| span
|
||||
.links
|
||||
.iter()
|
||||
.any(|link| link.attributes.len() > MAX_ATTRIBUTES)
|
||||
|| span
|
||||
.events
|
||||
.iter()
|
||||
.any(|event| event.attributes.len() > MAX_ATTRIBUTES)
|
||||
{
|
||||
return Err(DecodeError::TooLarge);
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn hex_bytes(bytes: &[u8]) -> String {
|
||||
bytes.iter().map(|byte| format!("{byte:02x}")).collect()
|
||||
}
|
||||
|
||||
fn decoded_span(
|
||||
span: Span,
|
||||
resource_attributes: &Shared<BTreeMap<String, String>>,
|
||||
scope_name: &Shared<String>,
|
||||
scope_version: &Shared<String>,
|
||||
budget: &mut Budget,
|
||||
) -> Result<DecodedSpan, DecodeError> {
|
||||
let status = span.status.unwrap_or_default();
|
||||
Ok(DecodedSpan {
|
||||
trace_id: hex_bytes(&span.trace_id),
|
||||
span_id: hex_bytes(&span.span_id),
|
||||
parent_span_id: hex_bytes(&span.parent_span_id),
|
||||
trace_state: span.trace_state,
|
||||
name: span.name,
|
||||
kind: SpanKind::try_from(span.kind)
|
||||
.unwrap_or(SpanKind::Unspecified)
|
||||
.as_str_name()
|
||||
.to_owned(),
|
||||
resource_attributes: budget.clone_shared(resource_attributes, |attributes| {
|
||||
attributes
|
||||
.iter()
|
||||
.map(|(key, value)| key.len() + value.len() + 96)
|
||||
.sum()
|
||||
})?,
|
||||
scope_name: budget.clone_shared(scope_name, String::len)?,
|
||||
scope_version: budget.clone_shared(scope_version, String::len)?,
|
||||
attributes: attributes(span.attributes, budget)?,
|
||||
start_ns: span.start_time_unix_nano,
|
||||
end_ns: span.end_time_unix_nano,
|
||||
status_code: StatusCode::try_from(status.code)
|
||||
.unwrap_or(StatusCode::Unset)
|
||||
.as_str_name()
|
||||
.to_owned(),
|
||||
status_message: status.message,
|
||||
events: span
|
||||
.events
|
||||
.into_iter()
|
||||
.map(|event| {
|
||||
budget.consume(event.name.len() + 96)?;
|
||||
Ok(DecodedEvent {
|
||||
name: event.name,
|
||||
attributes: attributes(event.attributes, budget)?,
|
||||
})
|
||||
})
|
||||
.collect::<Result<Vec<_>, DecodeError>>()?,
|
||||
})
|
||||
}
|
||||
43
litellm-rust/crates/traces/src/otlp/wire.rs
Normal file
43
litellm-rust/crates/traces/src/otlp/wire.rs
Normal file
|
|
@ -0,0 +1,43 @@
|
|||
use opentelemetry_proto::tonic::collector::trace::v1::ExportTraceServiceRequest;
|
||||
use prost::Message;
|
||||
|
||||
use super::limits::{json_preflight, protobuf_preflight};
|
||||
use crate::DecodeError;
|
||||
|
||||
#[derive(strum::EnumString)]
|
||||
#[strum(ascii_case_insensitive)]
|
||||
enum OtlpMediaType {
|
||||
#[strum(serialize = "application/json")]
|
||||
Json,
|
||||
#[strum(
|
||||
serialize = "application/x-protobuf",
|
||||
serialize = "application/protobuf"
|
||||
)]
|
||||
Protobuf,
|
||||
}
|
||||
|
||||
pub(super) fn decode(
|
||||
body: &[u8],
|
||||
content_type: Option<&str>,
|
||||
) -> Result<ExportTraceServiceRequest, DecodeError> {
|
||||
let media_type = content_type
|
||||
.unwrap_or("application/x-protobuf")
|
||||
.split(';')
|
||||
.next()
|
||||
.unwrap_or_default()
|
||||
.trim()
|
||||
.parse::<OtlpMediaType>()
|
||||
.map_err(|_| DecodeError::InvalidPayload)?;
|
||||
|
||||
let request = match media_type {
|
||||
OtlpMediaType::Json => {
|
||||
json_preflight(body)?;
|
||||
serde_json::from_slice(body).map_err(|_| DecodeError::InvalidPayload)?
|
||||
}
|
||||
OtlpMediaType::Protobuf => {
|
||||
protobuf_preflight(body)?;
|
||||
ExportTraceServiceRequest::decode(body).map_err(|_| DecodeError::InvalidPayload)?
|
||||
}
|
||||
};
|
||||
Ok(request)
|
||||
}
|
||||
46
litellm-rust/crates/traces/src/shared.rs
Normal file
46
litellm-rust/crates/traces/src/shared.rs
Normal file
|
|
@ -0,0 +1,46 @@
|
|||
use std::ops::Deref;
|
||||
|
||||
use serde::Serialize;
|
||||
|
||||
type Storage<T> = std::sync::Arc<T>;
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Serialize)]
|
||||
#[serde(transparent)]
|
||||
pub struct Shared<T>(Storage<T>);
|
||||
|
||||
#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)]
|
||||
pub struct SharedIdentity(usize);
|
||||
|
||||
impl<T> Shared<T> {
|
||||
pub fn new(value: T) -> Self {
|
||||
Self(Storage::new(value))
|
||||
}
|
||||
|
||||
pub fn identity(&self) -> SharedIdentity {
|
||||
SharedIdentity(std::ptr::from_ref(self.as_ref()) as usize)
|
||||
}
|
||||
|
||||
pub fn shares_storage_with(&self, other: &Self) -> bool {
|
||||
self.identity() == other.identity()
|
||||
}
|
||||
}
|
||||
|
||||
impl<T> From<T> for Shared<T> {
|
||||
fn from(value: T) -> Self {
|
||||
Self::new(value)
|
||||
}
|
||||
}
|
||||
|
||||
impl<T> AsRef<T> for Shared<T> {
|
||||
fn as_ref(&self) -> &T {
|
||||
self.0.as_ref()
|
||||
}
|
||||
}
|
||||
|
||||
impl<T> Deref for Shared<T> {
|
||||
type Target = T;
|
||||
|
||||
fn deref(&self) -> &T {
|
||||
self.as_ref()
|
||||
}
|
||||
}
|
||||
|
|
@ -1,17 +1,14 @@
|
|||
use std::{collections::BTreeMap, time::Duration};
|
||||
|
||||
use serde::Deserialize;
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
use litellm_http::Client;
|
||||
|
||||
use crate::{Connection, Error};
|
||||
|
||||
const MAX_RESPONSE_BYTES: usize = 4 * 1024 * 1024;
|
||||
use crate::{Connection, Error, Parameter, execute_read};
|
||||
|
||||
pub enum ReadQuery {
|
||||
ListTraces,
|
||||
TraceSpans,
|
||||
SpanDetail,
|
||||
SpanError,
|
||||
SpendByResponseIds,
|
||||
}
|
||||
|
||||
|
|
@ -21,6 +18,7 @@ impl ReadQuery {
|
|||
"list_traces" => Ok(Self::ListTraces),
|
||||
"trace_spans" => Ok(Self::TraceSpans),
|
||||
"span_detail" => Ok(Self::SpanDetail),
|
||||
"span_error" => Ok(Self::SpanError),
|
||||
"spend_by_response_ids" => Ok(Self::SpendByResponseIds),
|
||||
_ => Err(Error::InvalidQuery),
|
||||
}
|
||||
|
|
@ -31,116 +29,12 @@ impl ReadQuery {
|
|||
Self::ListTraces => include_str!("../query/list_traces.sql"),
|
||||
Self::TraceSpans => include_str!("../query/trace_spans.sql"),
|
||||
Self::SpanDetail => include_str!("../query/span_detail.sql"),
|
||||
Self::SpanError => include_str!("../query/span_error.sql"),
|
||||
Self::SpendByResponseIds => include_str!("../query/spend_by_response_ids.sql"),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
#[serde(untagged)]
|
||||
pub enum Parameter {
|
||||
Text(String),
|
||||
Integer(i64),
|
||||
Strings(Vec<String>),
|
||||
}
|
||||
|
||||
impl Parameter {
|
||||
fn encoded(&self) -> String {
|
||||
match self {
|
||||
Self::Text(value) => escaped(value),
|
||||
Self::Integer(value) => value.to_string(),
|
||||
Self::Strings(values) => format!(
|
||||
"[{}]",
|
||||
values
|
||||
.iter()
|
||||
.map(|value| format!("'{}'", escaped(value).replace('\'', "\\'")))
|
||||
.collect::<Vec<_>>()
|
||||
.join(",")
|
||||
),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn escaped(value: &str) -> String {
|
||||
value
|
||||
.replace('\\', "\\\\")
|
||||
.replace('\t', "\\t")
|
||||
.replace('\n', "\\n")
|
||||
.replace('\r', "\\r")
|
||||
.replace('\0', "\\0")
|
||||
}
|
||||
|
||||
pub async fn execute_read(
|
||||
client: &Client,
|
||||
connection: &Connection,
|
||||
sql: &str,
|
||||
parameters: &BTreeMap<String, Parameter>,
|
||||
) -> Result<String, Error> {
|
||||
if sql.trim().is_empty() {
|
||||
return Err(Error::EmptySql);
|
||||
}
|
||||
|
||||
let mut url = connection.url().clone();
|
||||
|
||||
let existing_pairs: Vec<(String, String)> = url
|
||||
.query_pairs()
|
||||
.filter(|(key, _)| {
|
||||
!key.starts_with("param_")
|
||||
&& !matches!(
|
||||
key.as_ref(),
|
||||
"query"
|
||||
| "readonly"
|
||||
| "default_format"
|
||||
| "max_result_rows"
|
||||
| "result_overflow_mode"
|
||||
| "max_execution_time"
|
||||
| "wait_end_of_query"
|
||||
)
|
||||
})
|
||||
.map(|(key, value)| (key.into_owned(), value.into_owned()))
|
||||
.collect();
|
||||
url.query_pairs_mut()
|
||||
.clear()
|
||||
.extend_pairs(existing_pairs)
|
||||
.append_pair("readonly", "1")
|
||||
.append_pair("max_result_rows", "1000")
|
||||
.append_pair("result_overflow_mode", "throw")
|
||||
.append_pair("max_execution_time", "10")
|
||||
.append_pair("wait_end_of_query", "1")
|
||||
.append_pair("default_format", "JSON");
|
||||
|
||||
url.query_pairs_mut().extend_pairs(
|
||||
parameters
|
||||
.iter()
|
||||
.map(|(name, value)| (format!("param_{name}"), value.encoded())),
|
||||
);
|
||||
|
||||
let request = client
|
||||
.post(url)
|
||||
.timeout(Duration::from_secs(15))
|
||||
.body(sql.to_owned());
|
||||
let mut response = request.send().await.map_err(|_| Error::Transport)?;
|
||||
if !response.status().is_success() {
|
||||
return Err(Error::QueryFailed(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() > MAX_RESPONSE_BYTES {
|
||||
return Err(Error::ResponseTooLarge);
|
||||
}
|
||||
body.extend_from_slice(&chunk);
|
||||
}
|
||||
|
||||
let json: serde_json::Value =
|
||||
serde_json::from_slice(&body).map_err(|_| Error::InvalidResponse)?;
|
||||
if json.get("exception").is_some() || !json.get("data").is_some_and(serde_json::Value::is_array)
|
||||
{
|
||||
return Err(Error::InvalidResponse);
|
||||
}
|
||||
String::from_utf8(body).map_err(|_| Error::InvalidResponse)
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy)]
|
||||
pub enum LensQuery {
|
||||
Sample,
|
||||
|
|
|
|||
|
|
@ -1,8 +1,107 @@
|
|||
use std::collections::BTreeMap;
|
||||
use std::{
|
||||
collections::BTreeMap,
|
||||
io::{BufRead, BufReader},
|
||||
};
|
||||
|
||||
use litellm_traces::encode_rows;
|
||||
use rstest::rstest;
|
||||
use flate2::read::GzDecoder;
|
||||
use litellm_http::Client;
|
||||
use litellm_traces::{
|
||||
Connection, Error, InsertRow, InsertTable, Shared, encode_rows, insert_shared_rows,
|
||||
};
|
||||
use rstest::{fixture, rstest};
|
||||
use serde_json::{Value, json};
|
||||
use wiremock::{
|
||||
Mock, MockServer, ResponseTemplate,
|
||||
matchers::{header, method},
|
||||
};
|
||||
|
||||
#[fixture]
|
||||
fn shared_rows(#[default(16 * 1024)] attribute_bytes: usize) -> Vec<InsertRow> {
|
||||
let resource = Shared::new(json!({"shared": "x".repeat(attribute_bytes)}));
|
||||
(0..1024)
|
||||
.map(|index| {
|
||||
BTreeMap::from([
|
||||
("ResourceAttributes".into(), resource.clone()),
|
||||
("SpanId".into(), Shared::new(json!(format!("{index:016x}")))),
|
||||
("Timestamp".into(), Shared::new(json!(1))),
|
||||
])
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::one_request(1)]
|
||||
#[case::concurrent_requests(2)]
|
||||
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
|
||||
async fn shared_fanout_survives_gzip_insert_over_http(
|
||||
shared_rows: Vec<InsertRow>,
|
||||
#[case] concurrency: usize,
|
||||
) {
|
||||
let server = MockServer::start().await;
|
||||
Mock::given(method("POST"))
|
||||
.and(header("Content-Encoding", "gzip"))
|
||||
.respond_with(ResponseTemplate::new(200))
|
||||
.expect(concurrency as u64)
|
||||
.mount(&server)
|
||||
.await;
|
||||
let client = Client::no_redirect_for_test();
|
||||
let connection = Connection::parse(&server.uri()).unwrap();
|
||||
let expected_resource = shared_rows[0]["ResourceAttributes"].clone();
|
||||
let expected_count = shared_rows.len();
|
||||
let mut requests = tokio::task::JoinSet::new();
|
||||
for _ in 0..concurrency {
|
||||
let client = client.clone();
|
||||
let connection = connection.clone();
|
||||
let rows = shared_rows.clone();
|
||||
requests.spawn(async move {
|
||||
insert_shared_rows(
|
||||
&client,
|
||||
&connection,
|
||||
"traces",
|
||||
InsertTable::OtelTraces,
|
||||
rows,
|
||||
)
|
||||
.await
|
||||
});
|
||||
}
|
||||
while let Some(result) = requests.join_next().await {
|
||||
result.unwrap().unwrap();
|
||||
}
|
||||
let received = server.received_requests().await.unwrap();
|
||||
assert_eq!(received.len(), concurrency);
|
||||
for request in received {
|
||||
let decoder = GzDecoder::new(request.body.as_slice());
|
||||
let mut count = 0;
|
||||
for (index, line) in BufReader::new(decoder).lines().enumerate() {
|
||||
let row: Value = serde_json::from_str(&line.unwrap()).unwrap();
|
||||
assert_eq!(&row["ResourceAttributes"], expected_resource.as_ref());
|
||||
assert_eq!(row["SpanId"], format!("{index:016x}"));
|
||||
assert_eq!(row["Timestamp"], "1970-01-01T00:00:00.000000001Z");
|
||||
assert!(row["EngineReceivedMs"].as_u64().unwrap() > 0);
|
||||
count += 1;
|
||||
}
|
||||
assert_eq!(count, expected_count);
|
||||
}
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn shared_fanout_over_insert_limit_never_reaches_http(
|
||||
#[with(64 * 1024)] shared_rows: Vec<InsertRow>,
|
||||
) {
|
||||
let server = MockServer::start().await;
|
||||
let connection = Connection::parse(&server.uri()).unwrap();
|
||||
let result = insert_shared_rows(
|
||||
&Client::no_redirect_for_test(),
|
||||
&connection,
|
||||
"traces",
|
||||
InsertTable::OtelTraces,
|
||||
shared_rows,
|
||||
)
|
||||
.await;
|
||||
assert!(matches!(result, Err(Error::InsertTooLarge)));
|
||||
assert!(server.received_requests().await.unwrap().is_empty());
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::span("Timestamp", json!(1_234_567_890), json!("1970-01-01T00:00:01.23456789Z"))]
|
||||
|
|
|
|||
|
|
@ -835,3 +835,148 @@ async fn lens_content_keeps_output_visible_after_long_input(
|
|||
assert_eq!(recovered, original);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::ascii(10, format!("ParentCommand: {}", "x".repeat(460_000)))]
|
||||
#[case::multibyte(1_000, "\u{1f9ea}".repeat(1_024))]
|
||||
#[case::escaped(1_000, "\0\n\"\\".repeat(1_024))]
|
||||
#[tokio::test]
|
||||
async fn trace_error_previews_preserve_paginated_diagnostics(
|
||||
#[future(awt)] database: TestResult<ClickHouseDatabase>,
|
||||
#[case] span_count: usize,
|
||||
#[case] message: String,
|
||||
) -> TestResult {
|
||||
let database = database?;
|
||||
let writer = Connection::writer(&database.url)?;
|
||||
ensure_schema(&database.client, &writer, "trace_test", 7, 14).await?;
|
||||
let timestamp = time::OffsetDateTime::now_utc().unix_timestamp_nanos() as i64;
|
||||
let rows = (0..span_count)
|
||||
.map(|index| {
|
||||
serde_json::from_value(serde_json::json!({
|
||||
"Timestamp": timestamp + index as i64, "TraceId": "diagnostic-trace",
|
||||
"SpanId": format!("span-{index}"), "SpanName": "tool",
|
||||
"StatusCode": "STATUS_CODE_ERROR", "StatusMessage": message,
|
||||
}))
|
||||
})
|
||||
.collect::<Result<Vec<BTreeMap<String, serde_json::Value>>, _>>()?;
|
||||
insert_rows(&database, "otel_traces", rows).await?;
|
||||
let reader = Connection::reader(&database.url, "trace_test")?;
|
||||
let mut parameters = BTreeMap::from([
|
||||
(
|
||||
"trace_id".into(),
|
||||
Parameter::Text("diagnostic-trace".into()),
|
||||
),
|
||||
("team_ids".into(), Parameter::Strings(vec![])),
|
||||
("api_key_hash".into(), Parameter::Text(String::new())),
|
||||
("trace_ref".into(), Parameter::Text(String::new())),
|
||||
]);
|
||||
let body = execute_named_read(
|
||||
&database.client,
|
||||
&reader,
|
||||
ReadQuery::TraceSpans,
|
||||
¶meters,
|
||||
)
|
||||
.await?;
|
||||
let response: serde_json::Value = serde_json::from_str(&body)?;
|
||||
let spans = response["data"].as_array().expect("trace spans");
|
||||
assert_eq!(spans.len(), span_count);
|
||||
let prefix: String = message.chars().take(128).collect();
|
||||
assert!(!prefix.is_empty());
|
||||
assert!(
|
||||
spans
|
||||
.iter()
|
||||
.all(|span| span["status_message"] == prefix && span["error_truncated"] == 1)
|
||||
);
|
||||
parameters.insert("span_id".into(), Parameter::Text("span-0".into()));
|
||||
parameters.insert("error_version".into(), Parameter::Text(String::new()));
|
||||
let mut recovered = String::new();
|
||||
loop {
|
||||
parameters.insert(
|
||||
"error_offset".into(),
|
||||
Parameter::Integer(recovered.chars().count() as i64),
|
||||
);
|
||||
let body = execute_named_read(&database.client, &reader, ReadQuery::SpanError, ¶meters)
|
||||
.await?;
|
||||
assert!(body.len() < 128 * 1024);
|
||||
let response: serde_json::Value = serde_json::from_str(&body)?;
|
||||
let chunk = response["data"][0]["message"]
|
||||
.as_str()
|
||||
.expect("diagnostic chunk");
|
||||
assert!(!chunk.is_empty());
|
||||
recovered.push_str(chunk);
|
||||
let version = response["data"][0]["version"]
|
||||
.as_str()
|
||||
.expect("diagnostic version");
|
||||
parameters.insert("error_version".into(), Parameter::Text(version.into()));
|
||||
if recovered.chars().count() >= message.chars().count() {
|
||||
break;
|
||||
}
|
||||
}
|
||||
assert_eq!(recovered, message);
|
||||
parameters.insert(
|
||||
"api_key_hash".into(),
|
||||
Parameter::Text("unrelated-key".into()),
|
||||
);
|
||||
let denied =
|
||||
execute_named_read(&database.client, &reader, ReadQuery::SpanError, ¶meters).await?;
|
||||
assert_eq!(
|
||||
serde_json::from_str::<serde_json::Value>(&denied)?["data"],
|
||||
serde_json::json!([])
|
||||
);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::different_start(1, 0)]
|
||||
#[case::different_receive(0, 1)]
|
||||
#[case::tied_timestamps(0, 0)]
|
||||
#[tokio::test]
|
||||
async fn duplicate_span_preview_matches_diagnostic(
|
||||
#[future(awt)] database: TestResult<ClickHouseDatabase>,
|
||||
#[case] start_delta: i64,
|
||||
#[case] receive_delta: i64,
|
||||
) -> TestResult {
|
||||
let database = database?;
|
||||
let writer = Connection::writer(&database.url)?;
|
||||
ensure_schema(&database.client, &writer, "trace_test", 7, 14).await?;
|
||||
let timestamp = time::OffsetDateTime::now_utc().unix_timestamp_nanos() as i64;
|
||||
let message = "a".repeat(200);
|
||||
let rows = [
|
||||
(start_delta, receive_delta, "z".repeat(200)),
|
||||
(0, 0, message.clone()),
|
||||
]
|
||||
.into_iter()
|
||||
.map(|(start_delta, receive_delta, message)| {
|
||||
serde_json::from_value(serde_json::json!({
|
||||
"Timestamp": timestamp + start_delta, "EngineReceivedMs": 100 + receive_delta,
|
||||
"TraceId": "duplicate-trace", "SpanId": "duplicate-span", "StatusMessage": message,
|
||||
}))
|
||||
})
|
||||
.collect::<Result<Vec<BTreeMap<String, serde_json::Value>>, _>>()?;
|
||||
insert_rows(&database, "otel_traces", rows).await?;
|
||||
let reader = Connection::reader(&database.url, "trace_test")?;
|
||||
let parameters = BTreeMap::from([
|
||||
("trace_id".into(), Parameter::Text("duplicate-trace".into())),
|
||||
("span_id".into(), Parameter::Text("duplicate-span".into())),
|
||||
("team_ids".into(), Parameter::Strings(vec![])),
|
||||
("api_key_hash".into(), Parameter::Text(String::new())),
|
||||
("trace_ref".into(), Parameter::Text(String::new())),
|
||||
("error_version".into(), Parameter::Text(String::new())),
|
||||
("error_offset".into(), Parameter::Integer(0)),
|
||||
]);
|
||||
let preview = execute_named_read(
|
||||
&database.client,
|
||||
&reader,
|
||||
ReadQuery::TraceSpans,
|
||||
¶meters,
|
||||
)
|
||||
.await?;
|
||||
let diagnostic =
|
||||
execute_named_read(&database.client, &reader, ReadQuery::SpanError, ¶meters).await?;
|
||||
let preview: serde_json::Value = serde_json::from_str(&preview)?;
|
||||
let diagnostic: serde_json::Value = serde_json::from_str(&diagnostic)?;
|
||||
assert_eq!(preview["data"].as_array().unwrap().len(), 1);
|
||||
assert_eq!(preview["data"][0]["status_message"], message[..128]);
|
||||
assert_eq!(diagnostic["data"][0]["message"], message);
|
||||
Ok(())
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,33 +1,19 @@
|
|||
use flate2::{Compression, write::GzEncoder};
|
||||
use litellm_traces::Shared;
|
||||
use litellm_traces::decode_otlp;
|
||||
use rstest::rstest;
|
||||
use std::io::Write;
|
||||
|
||||
const FIXTURE: &[u8] = include_bytes!(
|
||||
"../../../../tests/test_litellm/tracing/fixtures/langsmith_deep_agent_export.json"
|
||||
);
|
||||
|
||||
#[rstest]
|
||||
#[case::json(FIXTURE, Some("application/json"), None)]
|
||||
#[case::gzip_json(FIXTURE, Some("application/json"), Some("gzip"))]
|
||||
fn decodes_neutral_spans(
|
||||
#[case] body: &[u8],
|
||||
#[case] content_type: Option<&str>,
|
||||
#[case] content_encoding: Option<&str>,
|
||||
) {
|
||||
let payload = if content_encoding == Some("gzip") {
|
||||
let mut encoder = GzEncoder::new(Vec::new(), Compression::default());
|
||||
encoder.write_all(body).expect("gzip input");
|
||||
encoder.finish().expect("gzip payload")
|
||||
} else {
|
||||
body.to_vec()
|
||||
};
|
||||
let spans = decode_otlp(&payload, content_type, content_encoding, 8 * 1024 * 1024)
|
||||
.expect("valid OTLP export");
|
||||
#[case::json(FIXTURE, Some("application/json"))]
|
||||
fn decodes_neutral_spans(#[case] body: &[u8], #[case] content_type: Option<&str>) {
|
||||
let spans = decode_otlp(body, content_type).expect("valid OTLP export");
|
||||
assert_eq!(spans.len(), 6);
|
||||
assert_eq!(spans[0].trace_id, "4bad42b84e9de3ba46fc870185f8f023");
|
||||
assert_eq!(spans[0].resource_attributes["service.name"], "agent-demo");
|
||||
assert_eq!(spans[0].scope_name, "langsmith");
|
||||
assert_eq!(spans[0].scope_name.as_ref(), "langsmith");
|
||||
assert!(
|
||||
spans
|
||||
.iter()
|
||||
|
|
@ -36,12 +22,322 @@ fn decodes_neutral_spans(
|
|||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::invalid(b"not protobuf", None, 8 * 1024 * 1024)]
|
||||
#[case::too_large(FIXTURE, Some("application/json"), 1)]
|
||||
fn rejects_invalid_or_oversized_payload(
|
||||
#[case] body: &[u8],
|
||||
#[case] content_type: Option<&str>,
|
||||
#[case] limit: usize,
|
||||
) {
|
||||
assert!(decode_otlp(body, content_type, None, limit).is_err());
|
||||
fn accepts_trace_larger_than_eight_mib(mut span: opentelemetry_proto::tonic::trace::v1::Span) {
|
||||
use prost::Message;
|
||||
|
||||
span.name = "x".repeat(9 * 1024 * 1024);
|
||||
let body = request_with(span).encode_to_vec();
|
||||
let decoded = decode_otlp(&body, None).expect("16 MiB default accepts a 9 MiB trace");
|
||||
assert_eq!(decoded[0].name.len(), 9 * 1024 * 1024);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn rejects_invalid_payload() {
|
||||
assert!(decode_otlp(b"not protobuf", None).is_err());
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn decoder_does_not_enforce_the_http_body_limit() {
|
||||
let body = format!("{{\"ignored\":\"{}\"}}", "x".repeat(16 * 1024 * 1024 + 1));
|
||||
assert!(
|
||||
decode_otlp(body.as_bytes(), Some("application/json"))
|
||||
.unwrap()
|
||||
.is_empty()
|
||||
);
|
||||
}
|
||||
|
||||
fn request_with(
|
||||
span: opentelemetry_proto::tonic::trace::v1::Span,
|
||||
) -> opentelemetry_proto::tonic::collector::trace::v1::ExportTraceServiceRequest {
|
||||
use opentelemetry_proto::tonic::{
|
||||
collector::trace::v1::ExportTraceServiceRequest,
|
||||
trace::v1::{ResourceSpans, ScopeSpans},
|
||||
};
|
||||
ExportTraceServiceRequest {
|
||||
resource_spans: vec![ResourceSpans {
|
||||
scope_spans: vec![ScopeSpans {
|
||||
spans: vec![span],
|
||||
..Default::default()
|
||||
}],
|
||||
..Default::default()
|
||||
}],
|
||||
}
|
||||
}
|
||||
|
||||
#[rstest::fixture]
|
||||
fn span() -> opentelemetry_proto::tonic::trace::v1::Span {
|
||||
opentelemetry_proto::tonic::trace::v1::Span {
|
||||
trace_id: vec![1; 16],
|
||||
span_id: vec![2; 8],
|
||||
start_time_unix_nano: 1,
|
||||
end_time_unix_nano: 2,
|
||||
..Default::default()
|
||||
}
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn standard_json_and_protobuf_preserve_the_same_identifiers(
|
||||
span: opentelemetry_proto::tonic::trace::v1::Span,
|
||||
) {
|
||||
use prost::Message;
|
||||
let request = request_with(span);
|
||||
let json = serde_json::to_vec(&request).unwrap();
|
||||
let binary = request.encode_to_vec();
|
||||
let json_spans = decode_otlp(&json, Some("application/json; charset=utf-8")).unwrap();
|
||||
let binary_spans = decode_otlp(&binary, Some("application/x-protobuf")).unwrap();
|
||||
assert_eq!(
|
||||
serde_json::to_value(&json_spans).unwrap(),
|
||||
serde_json::to_value(&binary_spans).unwrap()
|
||||
);
|
||||
assert_eq!(json_spans[0].trace_id, "01".repeat(16));
|
||||
assert_eq!(json_spans[0].span_id, "02".repeat(8));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::json("APPLICATION/JSON; charset=utf-8", b"{}")]
|
||||
#[case::protobuf("application/x-protobuf; charset=binary", b"")]
|
||||
#[case::protobuf_alias("APPLICATION/PROTOBUF", b"")]
|
||||
fn supported_content_types_select_the_decoder(#[case] content_type: &str, #[case] body: &[u8]) {
|
||||
assert!(decode_otlp(body, Some(content_type)).is_ok());
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::missing_content_type(None)]
|
||||
#[case::unsupported_content_type(Some("text/plain"))]
|
||||
fn content_type_defaults_to_protobuf_and_rejects_unknown_values(
|
||||
#[case] content_type: Option<&str>,
|
||||
) {
|
||||
let result = decode_otlp(b"", content_type);
|
||||
assert_eq!(result.is_ok(), content_type.is_none());
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::short_trace(vec![1; 15], vec![2;8], 1, 2)]
|
||||
#[case::zero_trace(vec![0; 16], vec![2;8], 1, 2)]
|
||||
#[case::short_span(vec![1; 16], vec![2;7], 1, 2)]
|
||||
#[case::timestamp_overflow(vec![1;16], vec![2;8], i64::MAX as u64 + 1, i64::MAX as u64 + 1)]
|
||||
#[case::negative_duration(vec![1;16], vec![2;8], 3, 2)]
|
||||
fn rejects_ids_and_timestamps_that_cannot_be_stored(
|
||||
#[case] trace_id: Vec<u8>,
|
||||
#[case] span_id: Vec<u8>,
|
||||
#[case] start: u64,
|
||||
#[case] end: u64,
|
||||
) {
|
||||
use prost::Message;
|
||||
let span = opentelemetry_proto::tonic::trace::v1::Span {
|
||||
trace_id,
|
||||
span_id,
|
||||
start_time_unix_nano: start,
|
||||
end_time_unix_nano: end,
|
||||
..Default::default()
|
||||
};
|
||||
assert!(matches!(
|
||||
decode_otlp(&request_with(span).encode_to_vec(), None),
|
||||
Err(litellm_traces::DecodeError::InvalidPayload)
|
||||
));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn resource_fanout_shares_one_allocation(span: opentelemetry_proto::tonic::trace::v1::Span) {
|
||||
use opentelemetry_proto::tonic::{
|
||||
common::v1::{AnyValue, KeyValue, any_value::Value},
|
||||
resource::v1::Resource,
|
||||
};
|
||||
use prost::Message;
|
||||
let mut request = request_with(span.clone());
|
||||
request.resource_spans[0].resource = Some(Resource {
|
||||
attributes: vec![KeyValue {
|
||||
key: "shared".into(),
|
||||
value: Some(AnyValue {
|
||||
value: Some(Value::StringValue("x".repeat(16 * 1024))),
|
||||
}),
|
||||
..Default::default()
|
||||
}],
|
||||
..Default::default()
|
||||
});
|
||||
request.resource_spans[0].scope_spans[0].spans = vec![span; 1024];
|
||||
let second_scope = request.resource_spans[0].scope_spans[0].clone();
|
||||
request.resource_spans[0].scope_spans.push(second_scope);
|
||||
request
|
||||
.resource_spans
|
||||
.push(request.resource_spans[0].clone());
|
||||
let body = request.encode_to_vec();
|
||||
let decoded = decode_otlp(&body, None).expect("shared resources do not expand with span count");
|
||||
assert_eq!(decoded.len(), 4096);
|
||||
assert!(decoded[..2048].iter().all(|span| {
|
||||
Shared::shares_storage_with(&span.resource_attributes, &decoded[0].resource_attributes)
|
||||
}));
|
||||
assert!(!Shared::shares_storage_with(
|
||||
&decoded[0].resource_attributes,
|
||||
&decoded[2048].resource_attributes
|
||||
));
|
||||
assert_eq!(
|
||||
*decoded[0].resource_attributes,
|
||||
*decoded[2048].resource_attributes
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn nested_values_are_serialized_once(span: opentelemetry_proto::tonic::trace::v1::Span) {
|
||||
use opentelemetry_proto::tonic::common::v1::{
|
||||
AnyValue, ArrayValue, KeyValue, any_value::Value,
|
||||
};
|
||||
use prost::Message;
|
||||
let nested = (0..8).fold(
|
||||
AnyValue {
|
||||
value: Some(Value::StringValue("quoted \"value\"".into())),
|
||||
},
|
||||
|child, _| AnyValue {
|
||||
value: Some(Value::ArrayValue(ArrayValue {
|
||||
values: vec![child],
|
||||
})),
|
||||
},
|
||||
);
|
||||
let mut request = request_with(span);
|
||||
request.resource_spans[0].scope_spans[0].spans[0].attributes = vec![KeyValue {
|
||||
key: "nested".into(),
|
||||
value: Some(nested),
|
||||
..Default::default()
|
||||
}];
|
||||
let spans = decode_otlp(&request.encode_to_vec(), None).unwrap();
|
||||
let expected = (0..8).fold(serde_json::json!("quoted \"value\""), |child, _| {
|
||||
serde_json::json!([child])
|
||||
});
|
||||
assert_eq!(
|
||||
serde_json::from_str::<serde_json::Value>(&spans[0].attributes["nested"]).unwrap(),
|
||||
expected
|
||||
);
|
||||
assert!(spans[0].attributes["nested"].len() < 64);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::nesting(format!("{}0{}", "[".repeat(40), "]".repeat(40)).into_bytes())]
|
||||
#[case::nodes(format!("[{}]", vec!["0"; 65537].join(",")).into_bytes())]
|
||||
fn rejects_json_structure_before_building_a_tree(#[case] body: Vec<u8>) {
|
||||
assert!(matches!(
|
||||
decode_otlp(&body, Some("application/json")),
|
||||
Err(litellm_traces::DecodeError::TooLarge)
|
||||
));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::depth(40, 1)]
|
||||
#[case::nodes(0, 65537)]
|
||||
fn protobuf_preflight_rejects_expansion_before_prost_allocates(
|
||||
span: opentelemetry_proto::tonic::trace::v1::Span,
|
||||
#[case] depth: usize,
|
||||
#[case] count: usize,
|
||||
) {
|
||||
use opentelemetry_proto::tonic::common::v1::{
|
||||
AnyValue, ArrayValue, KeyValue, any_value::Value,
|
||||
};
|
||||
use prost::Message;
|
||||
let value = (0..depth).fold(
|
||||
AnyValue {
|
||||
value: Some(Value::BoolValue(true)),
|
||||
},
|
||||
|child, _| AnyValue {
|
||||
value: Some(Value::ArrayValue(ArrayValue {
|
||||
values: vec![child],
|
||||
})),
|
||||
},
|
||||
);
|
||||
let mut request = request_with(span);
|
||||
request.resource_spans[0].scope_spans[0].spans[0].attributes = vec![KeyValue {
|
||||
key: "deep".into(),
|
||||
value: Some(value),
|
||||
..Default::default()
|
||||
}];
|
||||
request.resource_spans = vec![request.resource_spans[0].clone(); count];
|
||||
let body = request.encode_to_vec();
|
||||
assert!(matches!(
|
||||
decode_otlp(&body, None),
|
||||
Err(litellm_traces::DecodeError::TooLarge)
|
||||
));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn scope_fanout_shares_name_and_version(span: opentelemetry_proto::tonic::trace::v1::Span) {
|
||||
use opentelemetry_proto::tonic::common::v1::InstrumentationScope;
|
||||
use prost::Message;
|
||||
let mut request = request_with(span.clone());
|
||||
request.resource_spans[0].scope_spans[0].scope = Some(InstrumentationScope {
|
||||
name: "n".repeat(16 * 1024),
|
||||
version: "v".repeat(16 * 1024),
|
||||
..Default::default()
|
||||
});
|
||||
request.resource_spans[0].scope_spans[0].spans = vec![span; 1024];
|
||||
let decoded = decode_otlp(&request.encode_to_vec(), None).unwrap();
|
||||
assert!(
|
||||
decoded
|
||||
.iter()
|
||||
.all(|span| Shared::shares_storage_with(&span.scope_name, &decoded[0].scope_name))
|
||||
);
|
||||
assert!(
|
||||
decoded.iter().all(|span| Shared::shares_storage_with(
|
||||
&span.scope_version,
|
||||
&decoded[0].scope_version
|
||||
))
|
||||
);
|
||||
assert_eq!(decoded[0].scope_name.len(), 16 * 1024);
|
||||
assert_eq!(decoded[0].scope_version.len(), 16 * 1024);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn unique_attribute_expansion_still_respects_decoded_budget(
|
||||
span: opentelemetry_proto::tonic::trace::v1::Span,
|
||||
) {
|
||||
use opentelemetry_proto::tonic::common::v1::{AnyValue, KeyValue, any_value::Value};
|
||||
use prost::Message;
|
||||
let mut request = request_with(span.clone());
|
||||
request.resource_spans[0].scope_spans[0].spans = (0..1024)
|
||||
.map(|index| {
|
||||
let mut span = span.clone();
|
||||
span.attributes = vec![KeyValue {
|
||||
key: "unique".into(),
|
||||
value: Some(AnyValue {
|
||||
value: Some(Value::StringValue(format!(
|
||||
"{index:04}{}",
|
||||
"x".repeat(16_300)
|
||||
))),
|
||||
}),
|
||||
..Default::default()
|
||||
}];
|
||||
span
|
||||
})
|
||||
.collect();
|
||||
let body = request.encode_to_vec();
|
||||
assert!(body.len() < 16 * 1024 * 1024);
|
||||
assert!(matches!(
|
||||
decode_otlp(&body, None),
|
||||
Err(litellm_traces::DecodeError::TooLarge)
|
||||
));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn escaped_attribute_expansion_is_bounded_below_four_mib(
|
||||
span: opentelemetry_proto::tonic::trace::v1::Span,
|
||||
) {
|
||||
use opentelemetry_proto::tonic::common::v1::{
|
||||
AnyValue, ArrayValue, KeyValue, any_value::Value,
|
||||
};
|
||||
use prost::Message;
|
||||
let mut request = request_with(span);
|
||||
request.resource_spans[0].scope_spans[0].spans[0].attributes = vec![KeyValue {
|
||||
key: "escaped".into(),
|
||||
value: Some(AnyValue {
|
||||
value: Some(Value::ArrayValue(ArrayValue {
|
||||
values: vec![AnyValue {
|
||||
value: Some(Value::StringValue("\0".repeat(3 * 1024 * 1024))),
|
||||
}],
|
||||
})),
|
||||
}),
|
||||
..Default::default()
|
||||
}];
|
||||
let body = request.encode_to_vec();
|
||||
assert!(body.len() < 4 * 1024 * 1024);
|
||||
assert!(matches!(
|
||||
decode_otlp(&body, None),
|
||||
Err(litellm_traces::DecodeError::TooLarge)
|
||||
));
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,11 +0,0 @@
|
|||
use litellm_traces::Connection;
|
||||
use rstest::rstest;
|
||||
|
||||
#[rstest]
|
||||
#[case::http("http://localhost:8123", true)]
|
||||
#[case::https("https://localhost:8443", true)]
|
||||
#[case::tcp("tcp://localhost:9000", false)]
|
||||
#[case::missing_host("http://", false)]
|
||||
fn accepts_only_clickhouse_http_urls(#[case] value: &str, #[case] expected: bool) {
|
||||
assert_eq!(Connection::parse(value).is_ok(), expected);
|
||||
}
|
||||
23
litellm-rust/crates/traces/tests/shared.rs
Normal file
23
litellm-rust/crates/traces/tests/shared.rs
Normal file
|
|
@ -0,0 +1,23 @@
|
|||
use litellm_traces::Shared;
|
||||
use rstest::rstest;
|
||||
|
||||
#[rstest]
|
||||
fn clones_preserve_values_and_serialize_transparently() {
|
||||
let original = Shared::new(vec!["value".to_owned()]);
|
||||
let cloned = original.clone();
|
||||
assert_eq!(cloned.as_ref(), original.as_ref());
|
||||
assert_eq!(
|
||||
serde_json::to_value(&cloned).unwrap(),
|
||||
serde_json::json!(["value"])
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn clones_share_storage_without_merging_equal_values() {
|
||||
let original = Shared::new("value".to_owned());
|
||||
let cloned = original.clone();
|
||||
let equal = Shared::new("value".to_owned());
|
||||
assert!(original.shares_storage_with(&cloned));
|
||||
assert!(!original.shares_storage_with(&equal));
|
||||
assert_eq!(*original, *equal);
|
||||
}
|
||||
|
|
@ -663,7 +663,7 @@ azure_anthropic_models: Set = set()
|
|||
azure_text_models: Set = set()
|
||||
anyscale_models: Set = set()
|
||||
cerebras_models: Set = set()
|
||||
nadir_models: Set = set() # mutable-ok: provider registry, filled from model_cost at import like every sibling provider
|
||||
nadir_models: Set = set()
|
||||
galadriel_models: Set = set()
|
||||
nvidia_nim_models: Set = set()
|
||||
nvidia_riva_models: Set = set()
|
||||
|
|
@ -697,7 +697,7 @@ recraft_models: Set = set()
|
|||
cometapi_models: Set = set()
|
||||
oci_models: Set = set()
|
||||
vercel_ai_gateway_models: Set = set()
|
||||
edenai_models: Set = set() # mutable-ok: filled from the price map at import, like the sibling provider sets
|
||||
edenai_models: Set = set()
|
||||
volcengine_models: Set = set()
|
||||
wandb_models: Set = set(WANDB_MODELS)
|
||||
ovhcloud_models: Set = set()
|
||||
|
|
@ -2282,6 +2282,24 @@ if TYPE_CHECKING:
|
|||
# Track if async client cleanup has been registered (for lazy loading)
|
||||
_async_client_cleanup_registered = False
|
||||
|
||||
# litellm.agent() entrypoints, resolved lazily from litellm.harness by __getattr__.
|
||||
_AGENT_EXPORTS: Final = frozenset(
|
||||
{
|
||||
"agent",
|
||||
"aagent",
|
||||
"agent_session",
|
||||
"aagent_session",
|
||||
"agent_resume",
|
||||
"aagent_resume",
|
||||
"agent_capabilities",
|
||||
"Harness",
|
||||
"ClaudeCodeOptions",
|
||||
"CodexOptions",
|
||||
"OpenCodeOptions",
|
||||
"DeepAgentsOptions",
|
||||
}
|
||||
)
|
||||
|
||||
# Eager loading for backwards compatibility with VCR and other HTTP recording tools
|
||||
# When LITELLM_DISABLE_LAZY_LOADING is set, lazy-loaded attributes are loaded at import time
|
||||
# For now, this only affects encoding (tiktoken) as it was the only reported issue
|
||||
|
|
@ -2315,6 +2333,13 @@ def __getattr__(name: str) -> Any:
|
|||
handler_func: Final = registry[name]
|
||||
return handler_func(name)
|
||||
|
||||
# litellm.agent() and friends: imported on first access (not needed for completion calls)
|
||||
if name == "harness" or name in _AGENT_EXPORTS:
|
||||
import importlib
|
||||
|
||||
harness_module = importlib.import_module("litellm.harness")
|
||||
return harness_module if name == "harness" else getattr(harness_module, name)
|
||||
|
||||
# Lazy load encoding from main.py to avoid heavy tiktoken import
|
||||
if name == "encoding":
|
||||
from ._lazy_imports import get_litellm_globals
|
||||
|
|
|
|||
|
|
@ -352,13 +352,9 @@ def _replace_string_leaves(value: object, values: Iterator[str]) -> object:
|
|||
if isinstance(value, str):
|
||||
return next(values)
|
||||
if isinstance(value, dict):
|
||||
return { # mutable-ok: LogRecord extras must keep JSON dict shape for handlers
|
||||
key: _replace_string_leaves(child, values) for key, child in value.items()
|
||||
}
|
||||
return {key: _replace_string_leaves(child, values) for key, child in value.items()}
|
||||
if isinstance(value, list):
|
||||
return [ # mutable-ok: LogRecord extras must keep JSON list shape for handlers
|
||||
_replace_string_leaves(child, values) for child in value
|
||||
]
|
||||
return [_replace_string_leaves(child, values) for child in value]
|
||||
if isinstance(value, tuple):
|
||||
return tuple(_replace_string_leaves(child, values) for child in value)
|
||||
return value
|
||||
|
|
@ -368,13 +364,9 @@ def _sort_processed_sets(original: object, processed: object) -> object:
|
|||
if isinstance(original, set) and isinstance(processed, list):
|
||||
return sorted(processed)
|
||||
if isinstance(original, dict) and isinstance(processed, dict):
|
||||
return { # mutable-ok: sorting nested sets must preserve the surrounding JSON dict
|
||||
key: _sort_processed_sets(original.get(key), value) for key, value in processed.items()
|
||||
}
|
||||
return {key: _sort_processed_sets(original.get(key), value) for key, value in processed.items()}
|
||||
if isinstance(original, list) and isinstance(processed, list):
|
||||
return [ # mutable-ok: sorting nested sets must preserve the surrounding JSON list
|
||||
_sort_processed_sets(before, after) for before, after in zip(original, processed)
|
||||
]
|
||||
return [_sort_processed_sets(before, after) for before, after in zip(original, processed)]
|
||||
if isinstance(original, tuple) and isinstance(processed, tuple):
|
||||
return tuple(_sort_processed_sets(before, after) for before, after in zip(original, processed))
|
||||
return processed
|
||||
|
|
|
|||
|
|
@ -233,7 +233,7 @@ def _coerce_redis_kwargs_types(
|
|||
"socket_keepalive": bool,
|
||||
}
|
||||
)
|
||||
result: Final = dict(redis_kwargs) # mutable-ok: per-key try/except coercion below needs to drop individual keys
|
||||
result: Final = dict(redis_kwargs)
|
||||
for key, value in redis_kwargs.items():
|
||||
if not isinstance(value, str):
|
||||
continue
|
||||
|
|
@ -803,7 +803,7 @@ def _credential_provider_auth_kwargs(redis_kwargs: dict) -> dict:
|
|||
|
||||
superseded: Final = frozenset({"redis_connect_func", "username", "password"})
|
||||
kept: Final = ((k, v) for k, v in redis_kwargs.items() if k not in superseded)
|
||||
return dict(kept, credential_provider=credential_provider) # mutable-ok: the branches below mutate these kwargs
|
||||
return dict(kept, credential_provider=credential_provider)
|
||||
|
||||
|
||||
def get_redis_client(**env_overrides):
|
||||
|
|
|
|||
|
|
@ -145,7 +145,7 @@ def _a2a_cost_params(litellm_params: Mapping[str, object] | None) -> Mapping[str
|
|||
|
||||
|
||||
def _card_http_kwargs(extra_headers: dict[str, str] | None) -> dict[str, object] | None:
|
||||
return {"headers": extra_headers} if extra_headers else None # mutable-ok: a2a-sdk's get_agent_card takes a dict
|
||||
return {"headers": extra_headers} if extra_headers else None
|
||||
|
||||
|
||||
def _agent_card_path(litellm_params: Mapping[str, object]) -> str | None:
|
||||
|
|
@ -612,7 +612,7 @@ def _build_streaming_logging_obj(
|
|||
logging_obj.model_call_details["agent_id"] = agent_id
|
||||
|
||||
_request_context: Final = (("metadata", metadata), ("proxy_server_request", proxy_server_request))
|
||||
_litellm_params: Final = dict( # mutable-ok: Logging.litellm_params is declared as a dict
|
||||
_litellm_params: Final = dict(
|
||||
(*_a2a_cost_params(litellm_params).items(), *((key, value) for key, value in _request_context if value))
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -127,7 +127,7 @@ async def _handle_completed_batch(
|
|||
return BatchCostUsageResult(
|
||||
cost=0.0,
|
||||
usage=Usage(prompt_tokens=0, completion_tokens=0, total_tokens=0),
|
||||
models=[], # mutable-ok: no output file means no model was ever priced; BatchCostUsageResult.models requires list[str]
|
||||
models=[],
|
||||
successful_requests=0,
|
||||
failed_requests=await count_error_file_failed_requests(
|
||||
batch, custom_llm_provider=custom_llm_provider, litellm_params=litellm_params
|
||||
|
|
|
|||
|
|
@ -103,10 +103,10 @@ async def claim_affinity_pin(
|
|||
try:
|
||||
claim_script: Final = redis_cache.async_register_script(_CLAIM_PIN_SCRIPT)
|
||||
args: Final = (
|
||||
json.dumps(dict(pin_value)), # mutable-ok: JSON serialization requires dict, not a generic Mapping
|
||||
json.dumps(dict(pin_value)),
|
||||
int(ttl_seconds),
|
||||
*(
|
||||
(json.dumps(tuple(dict(value) for value in eligible_values)),) # mutable-ok: JSON requires dict
|
||||
(json.dumps(tuple(dict(value) for value in eligible_values)),)
|
||||
if eligible_values is not None
|
||||
else ()
|
||||
),
|
||||
|
|
|
|||
|
|
@ -702,14 +702,14 @@ class LLMCachingHandler:
|
|||
)
|
||||
merged: Final = EmbeddingResponse(
|
||||
model=cached.model,
|
||||
data=[ # mutable-ok: EmbeddingResponse.data is a pydantic list field
|
||||
data=[
|
||||
item
|
||||
if item is not None
|
||||
else Embedding(embedding=next(fresh_items)["embedding"], index=position, object="embedding")
|
||||
for position, item in enumerate(cached.data)
|
||||
],
|
||||
usage=merged_usage,
|
||||
hidden_params={ # mutable-ok: EmbeddingResponse._hidden_params is a mutable dict field
|
||||
hidden_params={
|
||||
**cached._hidden_params,
|
||||
"cache_hit": True,
|
||||
},
|
||||
|
|
|
|||
|
|
@ -252,9 +252,7 @@ class DualCache(BaseCache):
|
|||
if value is not None:
|
||||
self.in_memory_cache.set_cache(key, value, **self._backfill_kwargs(kwargs))
|
||||
|
||||
return list( # mutable-ok: public list contract
|
||||
redis_result.get(key) if value is None else value for key, value in zip(keys, result)
|
||||
)
|
||||
return list(redis_result.get(key) if value is None else value for key, value in zip(keys, result))
|
||||
except Exception as e:
|
||||
log_redis_failure(
|
||||
verbose_logger, logging.ERROR, "LiteLLM Cache: exception in batch_get_cache", e, with_traceback=True
|
||||
|
|
@ -329,8 +327,8 @@ class DualCache(BaseCache):
|
|||
def reserve_redis_batch_reads(self, keys: Sequence[str]) -> tuple[list[str], dict[str, float | None]]:
|
||||
"""Reserve the memory-missed keys whose throttled Redis reads are due, as a batch read would."""
|
||||
if self.redis_cache is None:
|
||||
return [], {} # mutable-ok: API contract returns an empty list and dictionary
|
||||
key_list: Final = list(keys) # mutable-ok: batch_get_cache takes a list
|
||||
return [], {}
|
||||
key_list: Final = list(keys)
|
||||
memory: Final = self.in_memory_cache
|
||||
in_memory_result: Final = (
|
||||
None
|
||||
|
|
@ -386,7 +384,7 @@ class DualCache(BaseCache):
|
|||
|
||||
async def declare_batch_get(self, keys: Sequence[str], batch: RedisBatch) -> DeclaredBatchRead:
|
||||
pending: Final = await self._prepare_batch_get(
|
||||
list(keys), # mutable-ok: the shared batch read takes a list
|
||||
list(keys),
|
||||
local_only=False,
|
||||
throttle_redis=False,
|
||||
)
|
||||
|
|
@ -627,7 +625,7 @@ class DualCache(BaseCache):
|
|||
parent_otel_span: Span | None = None,
|
||||
) -> None:
|
||||
batch: Final = None if self.redis_cache is None else active_post_call_redis_batch(self.redis_cache)
|
||||
operations: Final = list(increment_list) # mutable-ok: both increment pipelines take a list
|
||||
operations: Final = list(increment_list)
|
||||
if batch is None:
|
||||
await self.async_increment_cache_pipeline(operations, parent_otel_span=parent_otel_span)
|
||||
return
|
||||
|
|
|
|||
|
|
@ -238,7 +238,7 @@ class EvictedClientCloser:
|
|||
the front rather than having to be searched for.
|
||||
"""
|
||||
with self._queue_lock:
|
||||
bucket: Final = self._buckets.setdefault(_bucket_key(pending), deque()) # mutable-ok: FIFO by design
|
||||
bucket: Final = self._buckets.setdefault(_bucket_key(pending), deque())
|
||||
while bucket and bucket[0].client_ref() is None:
|
||||
bucket.popleft()
|
||||
self._pending_count -= 1
|
||||
|
|
|
|||
|
|
@ -33,7 +33,7 @@ from litellm.types.services import ServiceTypes
|
|||
|
||||
_T = TypeVar("_T")
|
||||
_ScriptArg = str | bytes | int | float
|
||||
SettledHook = Callable[[asyncio.Future[_T]], Awaitable[None] | None] # mutable-ok: Callable params
|
||||
SettledHook = Callable[[asyncio.Future[_T]], Awaitable[None] | None]
|
||||
POST_CALL_FLUSH_DEADLINE_SECONDS: Final = 1.0
|
||||
|
||||
|
||||
|
|
@ -139,7 +139,7 @@ class _MGet(_Op[Mapping[str, object]]):
|
|||
)
|
||||
|
||||
async def run_alone(self) -> Mapping[str, object]:
|
||||
found: Mapping[str, object] = await self._redis_cache.async_batch_get_cache(key_list=list(self._keys)) # pyright: ignore[reportUnknownMemberType, reportUnknownVariableType] # untyped cache API # mutable-ok: the cache API takes a list
|
||||
found: Mapping[str, object] = await self._redis_cache.async_batch_get_cache(key_list=list(self._keys)) # pyright: ignore[reportUnknownMemberType, reportUnknownVariableType] # untyped cache API
|
||||
if any(key not in found for key in self._keys):
|
||||
raise ConnectionError("batch get did not return every key")
|
||||
return found
|
||||
|
|
|
|||
|
|
@ -123,13 +123,13 @@ def _reasoning_input_items(msg: "AllMessageValues") -> list[dict[str, object]]:
|
|||
blocks are the fallback for turns that arrived over another API surface.
|
||||
"""
|
||||
items: Final = _get_reasoning_items(msg)
|
||||
stored: Final = [_reasoning_item_to_response_input(item) for item in items] # mutable-ok: API message payload
|
||||
stored: Final = [_reasoning_item_to_response_input(item) for item in items]
|
||||
if stored:
|
||||
return stored
|
||||
raw_blocks: Final = msg.get("thinking_blocks") or ()
|
||||
blocks: Final = cast("Iterable[ChatCompletionThinkingBlock]", raw_blocks) # cast-ok: untyped client json
|
||||
replayed: Final = responses_reasoning_items_from_thinking_blocks(blocks)
|
||||
return [dict(item) for item in replayed] # mutable-ok: API message payload
|
||||
return [dict(item) for item in replayed]
|
||||
|
||||
|
||||
def _build_reasoning_item(
|
||||
|
|
@ -441,7 +441,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
input_items.extend(_reasoning_input_items(msg))
|
||||
if content:
|
||||
input_items.append(
|
||||
{ # mutable-ok: API message payload
|
||||
{
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"content": self._convert_content_to_responses_format(content, "assistant"),
|
||||
|
|
@ -475,7 +475,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
if role == "assistant":
|
||||
input_items.extend(_reasoning_input_items(msg))
|
||||
input_items.append(
|
||||
{ # mutable-ok: API message payload
|
||||
{
|
||||
"type": "message",
|
||||
"role": role,
|
||||
"content": self._convert_content_to_responses_format(content, cast(str, role)),
|
||||
|
|
@ -531,11 +531,11 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
) -> "ResponseText":
|
||||
existing: Final = cast( # cast-ok: text field is a ResponseText | dict[str, Any] | None union
|
||||
"dict[str, object]",
|
||||
dict(responses_api_request).get("text") or {}, # mutable-ok: one-shot merge seed
|
||||
dict(responses_api_request).get("text") or {},
|
||||
)
|
||||
return cast( # cast-ok: merged mapping is a valid ResponseText shape
|
||||
"ResponseText",
|
||||
{**existing, **update}, # mutable-ok: one-shot merged payload
|
||||
{**existing, **update},
|
||||
)
|
||||
|
||||
def _build_sanitized_litellm_params(self, litellm_params: dict) -> dict[str, object]:
|
||||
|
|
@ -1506,7 +1506,7 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator):
|
|||
# tool call; per-stream callers already received it via
|
||||
# output_item.added and the argument delta events
|
||||
return ModelResponseStream(
|
||||
choices=[ # mutable-ok: ModelResponseStream coerces only list choices
|
||||
choices=[
|
||||
StreamingChoices(
|
||||
index=0,
|
||||
delta=Delta(
|
||||
|
|
@ -1612,7 +1612,7 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator):
|
|||
)
|
||||
],
|
||||
usage=usage,
|
||||
provider_specific_fields=dict(provider_metadata) or None, # mutable-ok: field is typed dict
|
||||
provider_specific_fields=dict(provider_metadata) or None,
|
||||
**(
|
||||
MappingProxyType({"service_tier": served_service_tier})
|
||||
if isinstance(served_service_tier, str)
|
||||
|
|
|
|||
|
|
@ -52,10 +52,10 @@ CLICKHOUSE_MAX_BUFFERED_ROWS: Final = get_env_int("CLICKHOUSE_MAX_BUFFERED_ROWS"
|
|||
CLICKHOUSE_MAX_RETRIES: Final = get_env_int("CLICKHOUSE_MAX_RETRIES", 3)
|
||||
AGENT_TRACING_RETENTION_DAYS: Final = get_env_int("AGENT_TRACING_RETENTION_DAYS", 30)
|
||||
AGENT_TRACING_SPEND_LOG_RETENTION_DAYS: Final = get_env_int("AGENT_TRACING_SPEND_LOG_RETENTION_DAYS", 90)
|
||||
OTLP_MAX_BODY_BYTES: Final = get_env_int("OTLP_MAX_BODY_BYTES", 8 * 1024 * 1024)
|
||||
OTLP_MAX_BODY_BYTES: Final = get_env_int("OTLP_MAX_BODY_BYTES", 16 * 1024 * 1024)
|
||||
OTLP_MAX_ATTRIBUTE_VALUE_BYTES: Final = get_env_int("OTLP_MAX_ATTRIBUTE_VALUE_BYTES", 64 * 1024)
|
||||
OTLP_RETRY_AFTER_SECONDS: Final = get_env_int("OTLP_RETRY_AFTER_SECONDS", 2)
|
||||
OTLP_OFFLOAD_DECODE_BYTES: Final = get_env_int("OTLP_OFFLOAD_DECODE_BYTES", 256 * 1024)
|
||||
OTLP_MAX_CONCURRENT_INGESTS: Final = get_env_int("OTLP_MAX_CONCURRENT_INGESTS", 2)
|
||||
AGENT_TRACING_INPUT_PREVIEW_CHARS: Final = get_env_int("AGENT_TRACING_INPUT_PREVIEW_CHARS", 240)
|
||||
AGENT_TRACING_LIST_PAGE_SIZE: Final = get_env_int("AGENT_TRACING_LIST_PAGE_SIZE", 50)
|
||||
DEFAULT_S3_FLUSH_INTERVAL_SECONDS: Final = int(os.getenv("DEFAULT_S3_FLUSH_INTERVAL_SECONDS", 10))
|
||||
|
|
@ -2162,6 +2162,17 @@ MCP_SPEND_LOG_MODEL_PREFIX: Final[str] = "MCP: "
|
|||
PTU_SENTINEL_API_KEY: Final[str] = "__ptu_flat_cost__"
|
||||
PTU_ROLLUP_JOB_ID: Final[str] = "ptu_flat_cost_rollup_job"
|
||||
PTU_ROLLUP_LOCK_TTL_SECONDS: Final[int] = 900
|
||||
USAGE_TOP_API_KEYS_DEFAULT: Final[int] = 100
|
||||
USAGE_TOP_API_KEYS_MAX: Final[int] = 1000
|
||||
USAGE_KEY_PAGE_DEFAULT: Final[int] = 50
|
||||
USAGE_KEY_PAGE_MAX: Final[int] = 100
|
||||
USAGE_KEY_SEARCH_DEFAULT: Final[int] = 100
|
||||
USAGE_KEY_SEARCH_MAX: Final[int] = 100
|
||||
USAGE_MODEL_TOP_KEYS_DEFAULT: Final[int] = 5
|
||||
USAGE_MODEL_TOP_KEYS_MAX: Final[int] = 100
|
||||
USAGE_CACHE_LEAKAGE_KEYS_DEFAULT: Final[int] = 20
|
||||
USAGE_CACHE_LEAKAGE_KEYS_MAX: Final[int] = 100
|
||||
USAGE_EXPORT_BATCH_SIZE: Final[int] = 1000
|
||||
# Furthest back the catch-up pass looks for unpriced PTU days when a deployment
|
||||
# declares no ptu_effective_from, bounding the scan for an open-ended window.
|
||||
PTU_ROLLUP_MAX_BACKFILL_DAYS: Final[int] = 90
|
||||
|
|
@ -2199,3 +2210,26 @@ EMPTY_MAPPING: Final = MappingProxyType({})
|
|||
|
||||
# API endpoint for breached password k-anonymity search
|
||||
HIBP_RANGE_API_BASE: Final = "https://api.pwnedpasswords.com/range"
|
||||
|
||||
# litellm.harness defaults
|
||||
HARNESS_ENDPOINT_HOST: Final = "127.0.0.1"
|
||||
HARNESS_ENDPOINT_STARTUP_TIMEOUT_SECONDS: Final = 10.0
|
||||
HARNESS_ENDPOINT_REQUEST_TIMEOUT_SECONDS: Final = 600.0
|
||||
HARNESS_SESSION_TOKEN_BYTES: Final = 32
|
||||
HARNESS_MAX_DIFF_BYTES: Final = 256 * 1024
|
||||
HARNESS_STDERR_TAIL_LINES: Final = 40
|
||||
HARNESS_STREAM_READ_CHUNK_BYTES: Final = 64 * 1024
|
||||
HARNESS_EVENT_QUEUE_MAX_SIZE: Final = 1024
|
||||
HARNESS_PROCESS_KILL_GRACE_SECONDS: Final = 5.0
|
||||
HARNESS_SNAPSHOT_SKIP_DIRS: Final = frozenset(
|
||||
{
|
||||
".git",
|
||||
"node_modules",
|
||||
".venv",
|
||||
"venv",
|
||||
"__pycache__",
|
||||
".mypy_cache",
|
||||
".pytest_cache",
|
||||
".ruff_cache",
|
||||
}
|
||||
)
|
||||
|
|
|
|||
|
|
@ -2874,9 +2874,7 @@ class ResponsesWebSocketTokenUsageProcessor(BaseTokenUsageProcessor):
|
|||
collected_usage_objects: Final = ResponsesWebSocketTokenUsageProcessor.collect_usage_from_responses_ws_results(
|
||||
results
|
||||
)
|
||||
return ResponsesWebSocketTokenUsageProcessor.combine_usage_objects(
|
||||
list(collected_usage_objects) # mutable-ok: combine_usage_objects requires a list parameter
|
||||
)
|
||||
return ResponsesWebSocketTokenUsageProcessor.combine_usage_objects(list(collected_usage_objects))
|
||||
|
||||
|
||||
_TRANSCRIPTION_COMPLETED_EVENT_TYPE: Final = "conversation.item.input_audio_transcription.completed"
|
||||
|
|
|
|||
|
|
@ -1136,7 +1136,7 @@ class MCPClient:
|
|||
async def _list_resource_templates_operation(session: ClientSession) -> ListResourceTemplatesResult:
|
||||
capabilities: Final = session.server_capabilities
|
||||
if capabilities is not None and capabilities.resources is None:
|
||||
return ListResourceTemplatesResult(resource_templates=[]) # mutable-ok: MCP result payload
|
||||
return ListResourceTemplatesResult(resource_templates=[])
|
||||
try:
|
||||
return ListResourceTemplatesResult(
|
||||
resource_templates=await self._list_optional_pages(
|
||||
|
|
@ -1150,7 +1150,7 @@ class MCPClient:
|
|||
verbose_logger.debug(
|
||||
"MCP client list_resource_templates is unsupported by %s: %s", self.server_url or "stdio", error
|
||||
)
|
||||
return ListResourceTemplatesResult(resource_templates=[]) # mutable-ok: MCP result payload
|
||||
return ListResourceTemplatesResult(resource_templates=[])
|
||||
|
||||
try:
|
||||
result: Final = await self.run_with_session(_list_resource_templates_operation)
|
||||
|
|
|
|||
|
|
@ -171,9 +171,7 @@ async def load_mcp_tools(
|
|||
"""
|
||||
tools: Final = await list_tools_with_pagination(session)
|
||||
if format == "openai":
|
||||
return [ # mutable-ok: public API returns a list
|
||||
transform_mcp_tool_to_openai_tool(mcp_tool=tool) for tool in tools
|
||||
]
|
||||
return [transform_mcp_tool_to_openai_tool(mcp_tool=tool) for tool in tools]
|
||||
return tools
|
||||
|
||||
|
||||
|
|
|
|||
98
litellm/harness/__init__.py
Normal file
98
litellm/harness/__init__.py
Normal file
|
|
@ -0,0 +1,98 @@
|
|||
"""Agent harnesses: run Claude Code, Codex, OpenCode or Deep Agents on any LiteLLM model.
|
||||
|
||||
The entrypoints live on the top-level package:
|
||||
|
||||
import litellm
|
||||
from litellm import Harness, sandbox
|
||||
|
||||
result = litellm.agent(
|
||||
Harness.CLAUDE_CODE,
|
||||
"fix the failing test",
|
||||
sandbox=sandbox.local("."),
|
||||
model="litellm_proxy/claude-sonnet-4-5", # a model group on your AI Gateway
|
||||
)
|
||||
|
||||
This module holds the types you get back: events, Result, State, errors.
|
||||
"""
|
||||
|
||||
from litellm.harness.errors import (
|
||||
CapabilityUnsupported,
|
||||
HarnessError,
|
||||
HarnessInstallFailed,
|
||||
OptionsMismatch,
|
||||
OutputInvalid,
|
||||
SandboxError,
|
||||
SessionClosed,
|
||||
StateIncompatible,
|
||||
)
|
||||
from litellm.harness.options import (
|
||||
ClaudeCodeOptions,
|
||||
CodexOptions,
|
||||
DeepAgentsOptions,
|
||||
OpenCodeOptions,
|
||||
)
|
||||
from litellm.harness.runtime import (
|
||||
AsyncEventStream,
|
||||
AsyncSession,
|
||||
aagent,
|
||||
aagent_resume,
|
||||
aagent_session,
|
||||
agent_capabilities,
|
||||
)
|
||||
from litellm.harness.sync import EventStream, Session, agent, agent_resume, agent_session
|
||||
from litellm.harness.types import (
|
||||
Approval,
|
||||
Capabilities,
|
||||
Compaction,
|
||||
Done,
|
||||
Event,
|
||||
FileChange,
|
||||
Harness,
|
||||
Reasoning,
|
||||
Result,
|
||||
State,
|
||||
Text,
|
||||
ToolCall,
|
||||
ToolResult,
|
||||
Usage,
|
||||
)
|
||||
|
||||
__all__ = (
|
||||
"Approval",
|
||||
"AsyncEventStream",
|
||||
"AsyncSession",
|
||||
"Capabilities",
|
||||
"CapabilityUnsupported",
|
||||
"ClaudeCodeOptions",
|
||||
"CodexOptions",
|
||||
"Compaction",
|
||||
"DeepAgentsOptions",
|
||||
"Done",
|
||||
"Event",
|
||||
"EventStream",
|
||||
"FileChange",
|
||||
"Harness",
|
||||
"HarnessError",
|
||||
"HarnessInstallFailed",
|
||||
"OpenCodeOptions",
|
||||
"OptionsMismatch",
|
||||
"OutputInvalid",
|
||||
"Reasoning",
|
||||
"Result",
|
||||
"SandboxError",
|
||||
"Session",
|
||||
"SessionClosed",
|
||||
"State",
|
||||
"StateIncompatible",
|
||||
"Text",
|
||||
"ToolCall",
|
||||
"ToolResult",
|
||||
"Usage",
|
||||
"aagent",
|
||||
"aagent_resume",
|
||||
"aagent_session",
|
||||
"agent",
|
||||
"agent_capabilities",
|
||||
"agent_resume",
|
||||
"agent_session",
|
||||
)
|
||||
62
litellm/harness/context.py
Normal file
62
litellm/harness/context.py
Normal file
|
|
@ -0,0 +1,62 @@
|
|||
"""Per-session state shared by the runtime, handlers and harness configs."""
|
||||
|
||||
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 pydantic import BaseModel
|
||||
|
||||
from litellm.harness.options import HarnessOptions
|
||||
from litellm.harness.sandbox.base import Sandbox
|
||||
from litellm.harness.types import Approval, Harness, PermissionMode
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.harness.endpoint import ModelEndpoint
|
||||
|
||||
ApprovalHandler: TypeAlias = Callable[
|
||||
[Approval], bool | Awaitable[bool] # mutable-ok: Callable parameter list in a type alias, not a runtime collection
|
||||
]
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class GatewayTarget:
|
||||
"""Resolved LiteLLM AI Gateway for `litellm_proxy/` models. Internal, not exported."""
|
||||
|
||||
api_base: str
|
||||
api_key: str
|
||||
|
||||
|
||||
@dataclass
|
||||
class SessionContext:
|
||||
"""Everything a handler and config need for a session. Owned by the runtime."""
|
||||
|
||||
harness: Harness
|
||||
sandbox: Sandbox
|
||||
session_id: str
|
||||
# Model name as sent to the runtime (litellm_proxy/ prefix already stripped).
|
||||
model: str | None = None
|
||||
gateway: GatewayTarget | None = None
|
||||
api_key: str | None = None
|
||||
api_base: str | None = None
|
||||
endpoint: ModelEndpoint | None = None
|
||||
instructions: str | None = None
|
||||
tools: Sequence[Callable[..., Any]] = ()
|
||||
skills: Sequence[str] = ()
|
||||
disable_tools: Sequence[str] = ()
|
||||
permissions: PermissionMode = "full"
|
||||
on_approval: ApprovalHandler | None = None
|
||||
output: type[BaseModel] | None = None
|
||||
max_turns: int | None = None
|
||||
timeout: float | None = None
|
||||
metadata: Mapping[str, Any] = field(default_factory=dict)
|
||||
options: HarnessOptions | None = None
|
||||
# Set by the handler after each turn.
|
||||
final_text: str = ""
|
||||
output_json: str | None = None
|
||||
# Usage for in-process harnesses that call LiteLLM directly (no model endpoint).
|
||||
input_tokens: int = 0
|
||||
output_tokens: int = 0
|
||||
cost: float = 0.0
|
||||
calls: int = 0
|
||||
689
litellm/harness/endpoint.py
Normal file
689
litellm/harness/endpoint.py
Normal file
|
|
@ -0,0 +1,689 @@
|
|||
"""Per-session local model endpoint every CLI harness talks to.
|
||||
|
||||
The runtime inside the sandbox points its Anthropic / OpenAI base URL at this endpoint and
|
||||
authenticates with a random per-session token. The endpoint either reverse-proxies to a LiteLLM
|
||||
AI Gateway (gateway mode) or calls the LiteLLM SDK directly (SDK mode), and counts usage + cost.
|
||||
|
||||
starlette and uvicorn are optional: they are imported only when an endpoint starts.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import contextlib
|
||||
import itertools
|
||||
import json
|
||||
import logging
|
||||
import secrets
|
||||
from collections.abc import AsyncIterable, AsyncIterator, Mapping
|
||||
from dataclasses import dataclass
|
||||
from types import MappingProxyType, ModuleType
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
import httpx
|
||||
import openai
|
||||
|
||||
import litellm
|
||||
from litellm.constants import (
|
||||
DEFAULT_POLLING_INTERVAL,
|
||||
HARNESS_ENDPOINT_HOST,
|
||||
HARNESS_ENDPOINT_REQUEST_TIMEOUT_SECONDS,
|
||||
HARNESS_ENDPOINT_STARTUP_TIMEOUT_SECONDS,
|
||||
HARNESS_PROCESS_KILL_GRACE_SECONDS,
|
||||
HARNESS_SESSION_TOKEN_BYTES,
|
||||
)
|
||||
from litellm.harness.context import GatewayTarget
|
||||
from litellm.harness.errors import HarnessError, HarnessInstallFailed
|
||||
from litellm.harness.types import Harness, Usage
|
||||
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
|
||||
from litellm.types.llms.custom_http import httpxSpecialProvider
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from starlette.applications import Starlette
|
||||
from starlette.requests import Request
|
||||
from starlette.responses import Response
|
||||
from uvicorn import Server
|
||||
|
||||
verbose_logger: Final = logging.getLogger("LiteLLM")
|
||||
|
||||
MISSING_DEPS_MESSAGE = "litellm.harness needs starlette and uvicorn: pip install starlette uvicorn"
|
||||
|
||||
ROUTE_MESSAGES = "messages"
|
||||
ROUTE_CHAT = "chat/completions"
|
||||
ROUTE_RESPONSES = "responses"
|
||||
POST_ROUTES = (ROUTE_MESSAGES, ROUTE_CHAT, ROUTE_RESPONSES)
|
||||
ROUTE_PREFIXES: Final = ("", "/v1")
|
||||
|
||||
# What an SDK call or its stream can raise: LiteLLM maps provider failures onto openai's
|
||||
# exception hierarchy; transport errors, bad request kwargs and unserializable chunks remain.
|
||||
SDK_ERRORS: Final = (openai.OpenAIError, httpx.HTTPError, HarnessError, ValueError, TypeError)
|
||||
|
||||
HOP_BY_HOP_HEADERS = frozenset(
|
||||
{
|
||||
"connection",
|
||||
"keep-alive",
|
||||
"proxy-authenticate",
|
||||
"proxy-authorization",
|
||||
"te",
|
||||
"trailer",
|
||||
"trailers",
|
||||
"transfer-encoding",
|
||||
"upgrade",
|
||||
"host",
|
||||
"content-length",
|
||||
}
|
||||
)
|
||||
DROPPED_REQUEST_HEADERS = HOP_BY_HOP_HEADERS | frozenset(
|
||||
(
|
||||
"authorization",
|
||||
"x-api-key",
|
||||
"accept-encoding",
|
||||
)
|
||||
)
|
||||
DROPPED_RESPONSE_HEADERS = HOP_BY_HOP_HEADERS | frozenset(("content-encoding",))
|
||||
COST_HEADER = "x-litellm-response-cost"
|
||||
SSE_MEDIA_TYPE = "text/event-stream"
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _ServerDeps:
|
||||
uvicorn: ModuleType
|
||||
applications: ModuleType
|
||||
routing: ModuleType
|
||||
responses: ModuleType
|
||||
|
||||
|
||||
def _load_server_deps() -> _ServerDeps:
|
||||
"""Import starlette + uvicorn on demand; they are not litellm dependencies."""
|
||||
try:
|
||||
import uvicorn
|
||||
from starlette import applications, responses, routing
|
||||
except ImportError as e:
|
||||
raise HarnessInstallFailed(MISSING_DEPS_MESSAGE) from e
|
||||
return _ServerDeps(
|
||||
uvicorn=uvicorn,
|
||||
applications=applications,
|
||||
routing=routing,
|
||||
responses=responses,
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class UsageTracker:
|
||||
"""Running token + cost totals for one session."""
|
||||
|
||||
input_tokens: int = 0
|
||||
output_tokens: int = 0
|
||||
cost: float = 0.0
|
||||
calls: int = 0
|
||||
|
||||
def add(self, input_tokens: int = 0, output_tokens: int = 0, cost: float = 0.0) -> None:
|
||||
self.input_tokens += input_tokens
|
||||
self.output_tokens += output_tokens
|
||||
self.cost += cost
|
||||
self.calls += 1
|
||||
|
||||
def snapshot(self) -> Usage:
|
||||
return Usage(
|
||||
input_tokens=self.input_tokens,
|
||||
output_tokens=self.output_tokens,
|
||||
calls=self.calls,
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Usage parsing
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _as_int(value: object) -> int:
|
||||
if isinstance(value, bool):
|
||||
return 0
|
||||
if isinstance(value, (int, float)):
|
||||
return int(value)
|
||||
return 0
|
||||
|
||||
|
||||
def usage_from_mapping(usage: object) -> tuple[int, int]:
|
||||
"""(input, output) from a usage dict using OpenAI or Anthropic/Responses field names."""
|
||||
if not isinstance(usage, Mapping):
|
||||
return 0, 0
|
||||
input_tokens = usage.get("input_tokens", usage.get("prompt_tokens"))
|
||||
output_tokens = usage.get("output_tokens", usage.get("completion_tokens"))
|
||||
return _as_int(input_tokens), _as_int(output_tokens)
|
||||
|
||||
|
||||
def usage_from_body(body: object) -> tuple[int, int]:
|
||||
"""Usage from a non-streaming JSON response body."""
|
||||
if not isinstance(body, Mapping):
|
||||
return 0, 0
|
||||
if isinstance(body.get("usage"), Mapping):
|
||||
return usage_from_mapping(body["usage"])
|
||||
response = body.get("response")
|
||||
if isinstance(response, Mapping):
|
||||
return usage_from_mapping(response.get("usage"))
|
||||
return 0, 0
|
||||
|
||||
|
||||
class SSEUsageParser:
|
||||
"""Collects token usage from an SSE byte stream as it passes through."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.input_tokens = 0
|
||||
self.output_tokens = 0
|
||||
self._buffer = b""
|
||||
|
||||
def feed(self, chunk: bytes) -> None:
|
||||
self._buffer += chunk
|
||||
*lines, self._buffer = self._buffer.split(b"\n")
|
||||
for line in lines:
|
||||
self._feed_line(line)
|
||||
|
||||
def close(self) -> None:
|
||||
if self._buffer:
|
||||
self._feed_line(self._buffer)
|
||||
self._buffer = b""
|
||||
|
||||
def _feed_line(self, line: bytes) -> None:
|
||||
text = line.strip()
|
||||
if not text.startswith(b"data:"):
|
||||
return
|
||||
payload = text[len(b"data:") :].strip()
|
||||
if not payload or payload == b"[DONE]":
|
||||
return
|
||||
try:
|
||||
event = json.loads(payload)
|
||||
except ValueError:
|
||||
return
|
||||
if isinstance(event, Mapping):
|
||||
self.absorb(event)
|
||||
|
||||
def absorb(self, event: Mapping[str, Any]) -> None:
|
||||
event_type = event.get("type")
|
||||
if event_type == "message_start":
|
||||
self._absorb_message_start(event)
|
||||
elif event_type == "message_delta":
|
||||
self._absorb_message_delta(event)
|
||||
elif event_type == "response.completed":
|
||||
self._absorb_response_completed(event)
|
||||
elif isinstance(event.get("usage"), Mapping):
|
||||
self._set(*usage_from_mapping(event["usage"]))
|
||||
|
||||
def _absorb_message_start(self, event: Mapping[str, Any]) -> None:
|
||||
message = event.get("message")
|
||||
if isinstance(message, Mapping):
|
||||
self._set(*usage_from_mapping(message.get("usage")))
|
||||
|
||||
def _absorb_message_delta(self, event: Mapping[str, Any]) -> None:
|
||||
# message_delta output_tokens is cumulative for the whole message.
|
||||
self._set(*usage_from_mapping(event.get("usage")))
|
||||
|
||||
def _absorb_response_completed(self, event: Mapping[str, Any]) -> None:
|
||||
response = event.get("response")
|
||||
if isinstance(response, Mapping):
|
||||
self._set(*usage_from_mapping(response.get("usage")))
|
||||
|
||||
def _set(self, input_tokens: int, output_tokens: int) -> None:
|
||||
if input_tokens:
|
||||
self.input_tokens = input_tokens
|
||||
if output_tokens:
|
||||
self.output_tokens = output_tokens
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Cost + helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def compute_cost(model: str | None, input_tokens: int, output_tokens: int) -> float:
|
||||
"""Cost from LiteLLM's price map. Never raises; unknown models cost 0.0."""
|
||||
if not model or not (input_tokens or output_tokens):
|
||||
return 0.0
|
||||
try:
|
||||
prompt_cost, completion_cost = litellm.cost_per_token(
|
||||
model=model, prompt_tokens=input_tokens, completion_tokens=output_tokens
|
||||
)
|
||||
return float(prompt_cost) + float(completion_cost)
|
||||
except Exception: # accounting must never break a call; the price-map lookup raises bare Exception
|
||||
verbose_logger.debug("harness endpoint: cost lookup failed for %s", model, exc_info=True)
|
||||
return 0.0
|
||||
|
||||
|
||||
def header_cost(headers: Mapping[str, str]) -> float | None:
|
||||
raw = headers.get(COST_HEADER)
|
||||
if raw is None:
|
||||
return None
|
||||
try:
|
||||
return float(raw)
|
||||
except (TypeError, ValueError):
|
||||
return None
|
||||
|
||||
|
||||
def hidden_cost(response: object) -> float | None:
|
||||
hidden = getattr(response, "_hidden_params", None)
|
||||
if not isinstance(hidden, Mapping):
|
||||
return None
|
||||
try:
|
||||
cost = hidden.get("response_cost")
|
||||
return None if cost is None else float(cost)
|
||||
except (TypeError, ValueError):
|
||||
return None
|
||||
|
||||
|
||||
def extract_token(headers: Mapping[str, str]) -> str | None:
|
||||
auth = headers.get("authorization") or ""
|
||||
if auth.lower().startswith("bearer "):
|
||||
return auth[len("bearer ") :].strip()
|
||||
return headers.get("x-api-key")
|
||||
|
||||
|
||||
def gateway_headers(
|
||||
incoming: Mapping[str, str],
|
||||
gateway: GatewayTarget,
|
||||
harness: Harness,
|
||||
metadata: Mapping[str, Any] | None,
|
||||
) -> Mapping[str, str]:
|
||||
"""Incoming headers minus hop-by-hop/auth/x-litellm-*, plus gateway auth, tags, metadata."""
|
||||
kept = (
|
||||
(name, value)
|
||||
for name, value in incoming.items()
|
||||
if name.lower() not in DROPPED_REQUEST_HEADERS and not name.lower().startswith("x-litellm-")
|
||||
)
|
||||
metadata_json = json.dumps(dict(metadata), default=str) if metadata else None # mutable-ok: for json.dumps
|
||||
metadata_header = (("x-litellm-spend-logs-metadata", metadata_json),) if metadata_json is not None else ()
|
||||
added = (
|
||||
("authorization", f"Bearer {gateway.api_key}"),
|
||||
("x-litellm-tags", f"harness,{harness.value}"),
|
||||
*metadata_header,
|
||||
)
|
||||
return MappingProxyType(dict(itertools.chain(kept, added)))
|
||||
|
||||
|
||||
def response_headers(upstream: Mapping[str, str]) -> Mapping[str, str]:
|
||||
return MappingProxyType(
|
||||
{name: value for name, value in upstream.items() if name.lower() not in DROPPED_RESPONSE_HEADERS}
|
||||
)
|
||||
|
||||
|
||||
def sanitize(message: str, secret_values: tuple[str | None, ...]) -> str:
|
||||
for value in secret_values:
|
||||
if value:
|
||||
message = message.replace(value, "***")
|
||||
return message
|
||||
|
||||
|
||||
def error_status(exc: BaseException) -> int:
|
||||
status = getattr(exc, "status_code", None)
|
||||
if isinstance(status, int) and 400 <= status <= 599:
|
||||
return status
|
||||
return 500
|
||||
|
||||
|
||||
def error_body(exc: BaseException, message: str) -> dict[str, Any]: # mutable-ok: JSONResponse body
|
||||
return {"error": {"type": type(exc).__name__, "message": message}} # mutable-ok: JSONResponse body
|
||||
|
||||
|
||||
def to_jsonable(obj: object) -> object:
|
||||
if hasattr(obj, "model_dump"):
|
||||
return obj.model_dump(mode="json", exclude_none=True)
|
||||
if isinstance(obj, Mapping):
|
||||
return dict(obj) # mutable-ok: plain-dict copy so json.dumps can serialize any Mapping
|
||||
return obj
|
||||
|
||||
|
||||
def encode_anthropic_chunk(chunk: object) -> bytes:
|
||||
if isinstance(chunk, bytes):
|
||||
return chunk
|
||||
if isinstance(chunk, str):
|
||||
return chunk.encode()
|
||||
data = to_jsonable(chunk)
|
||||
event_type = data.get("type", "message") if isinstance(data, Mapping) else "message"
|
||||
return f"event: {event_type}\ndata: {json.dumps(data)}\n\n".encode()
|
||||
|
||||
|
||||
def encode_chat_chunk(chunk: object) -> bytes:
|
||||
if hasattr(chunk, "model_dump_json"):
|
||||
return f"data: {chunk.model_dump_json()}\n\n".encode()
|
||||
return f"data: {json.dumps(to_jsonable(chunk))}\n\n".encode()
|
||||
|
||||
|
||||
def encode_responses_chunk(chunk: object) -> bytes:
|
||||
data = to_jsonable(chunk)
|
||||
event_type = data.get("type", "message") if isinstance(data, Mapping) else "message"
|
||||
return f"event: {event_type}\ndata: {json.dumps(data)}\n\n".encode()
|
||||
|
||||
|
||||
STREAM_ENCODERS = MappingProxyType(
|
||||
{
|
||||
ROUTE_MESSAGES: encode_anthropic_chunk,
|
||||
ROUTE_CHAT: encode_chat_chunk,
|
||||
ROUTE_RESPONSES: encode_responses_chunk,
|
||||
}
|
||||
)
|
||||
STREAM_TRAILERS = MappingProxyType({ROUTE_CHAT: b"data: [DONE]\n\n"})
|
||||
|
||||
|
||||
def route_of(path: str) -> str:
|
||||
stripped = path.strip("/")
|
||||
stripped = stripped.removeprefix("v1/")
|
||||
return stripped
|
||||
|
||||
|
||||
def _noop() -> None:
|
||||
return None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# ModelEndpoint
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class ModelEndpoint:
|
||||
"""Local HTTP endpoint for one harness session. Use as an async context manager."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
harness: Harness,
|
||||
model: str | None,
|
||||
gateway: GatewayTarget | None,
|
||||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
metadata: Mapping[str, Any] | None = None,
|
||||
*,
|
||||
client: httpx.AsyncClient | None = None,
|
||||
) -> None:
|
||||
self.harness = harness
|
||||
self.model = model
|
||||
self.gateway = gateway
|
||||
self.api_key = api_key
|
||||
self.api_base = api_base
|
||||
self.metadata: Mapping[str, Any] = MappingProxyType(dict(metadata or ()))
|
||||
self.token = secrets.token_urlsafe(HARNESS_SESSION_TOKEN_BYTES)
|
||||
self.usage = UsageTracker()
|
||||
self.port = 0
|
||||
# Injected client (tests); production uses LiteLLM's shared cached client.
|
||||
self._injected_client = client
|
||||
self._deps: _ServerDeps | None = None
|
||||
self._client: httpx.AsyncClient | None = None
|
||||
self._server: Any = None
|
||||
self._task: asyncio.Task[None] | None = None
|
||||
|
||||
@property
|
||||
def url(self) -> str:
|
||||
return f"http://{HARNESS_ENDPOINT_HOST}:{self.port}"
|
||||
|
||||
# -- lifecycle ----------------------------------------------------------
|
||||
|
||||
async def __aenter__(self) -> ModelEndpoint:
|
||||
await self.start()
|
||||
return self
|
||||
|
||||
async def __aexit__(self, *exc_info: object) -> None:
|
||||
await self.stop()
|
||||
|
||||
async def start(self) -> None:
|
||||
self._deps = _load_server_deps()
|
||||
if self.gateway is not None:
|
||||
self._client = self._gateway_client()
|
||||
self._server = self._build_server(self._deps)
|
||||
self._task = asyncio.create_task(self._server.serve())
|
||||
try:
|
||||
await asyncio.wait_for(self._wait_started(), HARNESS_ENDPOINT_STARTUP_TIMEOUT_SECONDS)
|
||||
except BaseException:
|
||||
await self.stop()
|
||||
raise
|
||||
self.port = self._server.servers[0].sockets[0].getsockname()[1]
|
||||
|
||||
async def stop(self) -> None:
|
||||
if self._server is not None:
|
||||
self._server.should_exit = True
|
||||
if self._task is not None:
|
||||
with contextlib.suppress(BaseException):
|
||||
await self._task
|
||||
self._task = None
|
||||
# Never close the client: the shared cached one may still serve other requests,
|
||||
# and an injected one belongs to its caller.
|
||||
self._client = None
|
||||
|
||||
def _gateway_client(self) -> httpx.AsyncClient:
|
||||
"""LiteLLM's shared cached async client, unless one was injected."""
|
||||
if self._injected_client is not None:
|
||||
return self._injected_client
|
||||
handler = get_async_httpx_client(
|
||||
llm_provider=httpxSpecialProvider.AgentHarness,
|
||||
params={ # mutable-ok: get_async_httpx_client takes a dict params argument
|
||||
"timeout": HARNESS_ENDPOINT_REQUEST_TIMEOUT_SECONDS
|
||||
},
|
||||
)
|
||||
return handler.client
|
||||
|
||||
async def _wait_started(self) -> None:
|
||||
while not self._server.started:
|
||||
if self._task is not None and self._task.done():
|
||||
raise HarnessError("harness model endpoint failed to start")
|
||||
await asyncio.sleep(DEFAULT_POLLING_INTERVAL)
|
||||
|
||||
def _build_server(self, deps: _ServerDeps) -> Server:
|
||||
config = deps.uvicorn.Config(
|
||||
self._build_app(deps),
|
||||
host=HARNESS_ENDPOINT_HOST,
|
||||
port=0,
|
||||
log_config=None,
|
||||
log_level="warning",
|
||||
access_log=False,
|
||||
lifespan="off",
|
||||
timeout_graceful_shutdown=HARNESS_PROCESS_KILL_GRACE_SECONDS,
|
||||
)
|
||||
server = deps.uvicorn.Server(config)
|
||||
# Never touch the host process's signal handlers.
|
||||
if hasattr(server, "capture_signals"):
|
||||
server.capture_signals = contextlib.nullcontext
|
||||
if hasattr(server, "install_signal_handlers"):
|
||||
server.install_signal_handlers = _noop
|
||||
return server
|
||||
|
||||
def _build_app(self, deps: _ServerDeps) -> Starlette:
|
||||
Route = deps.routing.Route
|
||||
post_routes = tuple(
|
||||
Route(
|
||||
f"{prefix}/{route}",
|
||||
self._handle,
|
||||
methods=["POST"], # mutable-ok: Starlette Route takes a methods list
|
||||
)
|
||||
for prefix, route in itertools.product(ROUTE_PREFIXES, POST_ROUTES)
|
||||
)
|
||||
get_routes = tuple(
|
||||
Route(f"{prefix}/models", self._models, methods=["GET"]) # mutable-ok: Starlette Route takes a methods list
|
||||
for prefix in ROUTE_PREFIXES
|
||||
)
|
||||
return deps.applications.Starlette(
|
||||
routes=[*post_routes, *get_routes] # mutable-ok: Starlette takes a routes list
|
||||
)
|
||||
|
||||
# -- request handling ---------------------------------------------------
|
||||
|
||||
@property
|
||||
def _responses(self) -> ModuleType:
|
||||
if self._deps is None:
|
||||
raise HarnessError("harness model endpoint is not started")
|
||||
return self._deps.responses
|
||||
|
||||
def _authorized(self, request: Request) -> bool:
|
||||
token = extract_token(request.headers)
|
||||
return token is not None and secrets.compare_digest(token.encode(), self.token.encode())
|
||||
|
||||
def _json(self, body: object, status_code: int = 200) -> Response:
|
||||
return self._responses.JSONResponse(body, status_code=status_code)
|
||||
|
||||
def _unauthorized(self) -> Response:
|
||||
return self._json(
|
||||
{"error": {"type": "authentication_error", "message": "invalid token"}}, # mutable-ok: JSONResponse body
|
||||
401,
|
||||
)
|
||||
|
||||
def _error(self, exc: BaseException, status_code: int | None = None) -> Response:
|
||||
message = sanitize(str(exc), self._secrets())
|
||||
return self._json(error_body(exc, message), status_code or error_status(exc))
|
||||
|
||||
def _secrets(self) -> tuple[str | None, ...]:
|
||||
gateway_key = self.gateway.api_key if self.gateway else None
|
||||
return (gateway_key, self.api_key, self.token)
|
||||
|
||||
async def _models(self, request: Request) -> Response:
|
||||
if not self._authorized(request):
|
||||
return self._unauthorized()
|
||||
entry = {"id": self.model, "object": "model", "created": 0, "owned_by": "litellm"} # mutable-ok: JSON body
|
||||
data = (entry,) if self.model else ()
|
||||
return self._json({"object": "list", "data": data}) # mutable-ok: JSON response body for Starlette JSONResponse
|
||||
|
||||
async def _handle(self, request: Request) -> Response:
|
||||
if not self._authorized(request):
|
||||
return self._unauthorized()
|
||||
try:
|
||||
body = json.loads(await request.body())
|
||||
except ValueError as e:
|
||||
return self._error(e, 400)
|
||||
if not isinstance(body, dict):
|
||||
return self._error(ValueError("request body must be a JSON object"), 400)
|
||||
route = route_of(request.url.path)
|
||||
if self.gateway is not None:
|
||||
return await self._forward(request, route, body)
|
||||
return await self._call_sdk(route, body)
|
||||
|
||||
def _cost_model(self, body: Mapping[str, Any]) -> str | None:
|
||||
model = self.model or body.get("model")
|
||||
return model if isinstance(model, str) else None
|
||||
|
||||
def _record(
|
||||
self,
|
||||
model: str | None,
|
||||
input_tokens: int,
|
||||
output_tokens: int,
|
||||
cost: float | None,
|
||||
) -> None:
|
||||
if cost is None:
|
||||
cost = compute_cost(model, input_tokens, output_tokens)
|
||||
self.usage.add(input_tokens, output_tokens, cost)
|
||||
|
||||
# -- gateway mode -------------------------------------------------------
|
||||
|
||||
async def _forward(self, request: Request, route: str, body: Mapping[str, Any]) -> Response:
|
||||
if self._client is None or self.gateway is None:
|
||||
raise HarnessError("gateway client is not started")
|
||||
if self.model:
|
||||
body = {**body, "model": self.model} # mutable-ok: JSON request body re-sent upstream via httpx json=
|
||||
upstream_request = self._client.build_request(
|
||||
"POST",
|
||||
f"{self.gateway.api_base}/v1/{route}",
|
||||
json=body,
|
||||
headers=gateway_headers(request.headers, self.gateway, self.harness, self.metadata),
|
||||
)
|
||||
try:
|
||||
upstream = await self._client.send(upstream_request, stream=True)
|
||||
except httpx.HTTPError as e:
|
||||
return self._error(e, 502)
|
||||
return self._responses.StreamingResponse(
|
||||
self._relay(upstream, self._cost_model(body)),
|
||||
status_code=upstream.status_code,
|
||||
headers=response_headers(upstream.headers),
|
||||
)
|
||||
|
||||
async def _relay(self, upstream: httpx.Response, model: str | None) -> AsyncIterator[bytes]:
|
||||
is_sse = SSE_MEDIA_TYPE in upstream.headers.get("content-type", "")
|
||||
parser = SSEUsageParser()
|
||||
collected = bytearray()
|
||||
try:
|
||||
async for chunk in upstream.aiter_bytes():
|
||||
if is_sse:
|
||||
parser.feed(chunk)
|
||||
else:
|
||||
collected.extend(chunk)
|
||||
yield chunk
|
||||
finally:
|
||||
await upstream.aclose()
|
||||
if upstream.status_code < 400:
|
||||
self._record_relayed(upstream, model, parser, is_sse, bytes(collected))
|
||||
|
||||
def _record_relayed(
|
||||
self,
|
||||
upstream: httpx.Response,
|
||||
model: str | None,
|
||||
parser: SSEUsageParser,
|
||||
is_sse: bool,
|
||||
collected: bytes,
|
||||
) -> None:
|
||||
if is_sse:
|
||||
parser.close()
|
||||
tokens = (parser.input_tokens, parser.output_tokens)
|
||||
else:
|
||||
try:
|
||||
tokens = usage_from_body(json.loads(collected))
|
||||
except ValueError:
|
||||
tokens = (0, 0)
|
||||
self._record(model, tokens[0], tokens[1], header_cost(upstream.headers))
|
||||
|
||||
# -- SDK mode -----------------------------------------------------------
|
||||
|
||||
def _sdk_kwargs(
|
||||
self, body: Mapping[str, Any]
|
||||
) -> dict[str, Any]: # mutable-ok: SDK call kwargs, mutated by _invoke_sdk then splatted
|
||||
kwargs: dict[str, Any] = {**body} # mutable-ok: SDK call kwargs built from the JSON body, then overridden
|
||||
if self.model:
|
||||
kwargs["model"] = self.model
|
||||
if self.api_key:
|
||||
kwargs["api_key"] = self.api_key
|
||||
if self.api_base:
|
||||
kwargs["api_base"] = self.api_base
|
||||
return kwargs
|
||||
|
||||
async def _invoke_sdk(
|
||||
self,
|
||||
route: str,
|
||||
kwargs: dict[str, Any], # mutable-ok: injects stream_options into the SDK kwargs
|
||||
) -> object:
|
||||
if route == ROUTE_MESSAGES:
|
||||
return await litellm.anthropic.messages.acreate(**kwargs)
|
||||
if route == ROUTE_CHAT:
|
||||
if kwargs.get("stream"):
|
||||
stream_options = kwargs.get("stream_options") or {} # mutable-ok: empty default for a JSON field
|
||||
kwargs["stream_options"] = { # mutable-ok: JSON field sent to litellm.acompletion
|
||||
"include_usage": True,
|
||||
**stream_options,
|
||||
}
|
||||
return await litellm.acompletion(**kwargs)
|
||||
return await litellm.aresponses(**kwargs)
|
||||
|
||||
async def _call_sdk(self, route: str, body: Mapping[str, Any]) -> Response:
|
||||
kwargs = self._sdk_kwargs(body)
|
||||
model = self._cost_model(kwargs)
|
||||
try:
|
||||
response = await self._invoke_sdk(route, kwargs)
|
||||
except SDK_ERRORS as e:
|
||||
verbose_logger.debug("harness endpoint: SDK call failed: %s", type(e).__name__)
|
||||
return self._error(e)
|
||||
if kwargs.get("stream") and isinstance(response, AsyncIterable):
|
||||
return self._responses.StreamingResponse(
|
||||
self._sdk_stream(route, response, model), media_type=SSE_MEDIA_TYPE
|
||||
)
|
||||
data = to_jsonable(response)
|
||||
input_tokens, output_tokens = usage_from_body(data)
|
||||
self._record(model, input_tokens, output_tokens, hidden_cost(response))
|
||||
return self._json(data)
|
||||
|
||||
async def _sdk_stream(self, route: str, iterator: AsyncIterable[object], model: str | None) -> AsyncIterator[bytes]:
|
||||
encode = STREAM_ENCODERS[route]
|
||||
parser = SSEUsageParser()
|
||||
try:
|
||||
async for chunk in iterator:
|
||||
encoded = encode(chunk)
|
||||
parser.feed(encoded)
|
||||
yield encoded
|
||||
trailer = STREAM_TRAILERS.get(route)
|
||||
if trailer:
|
||||
yield trailer
|
||||
except SDK_ERRORS as e:
|
||||
message = sanitize(str(e), self._secrets())
|
||||
yield f"event: error\ndata: {json.dumps(error_body(e, message))}\n\n".encode()
|
||||
finally:
|
||||
parser.close()
|
||||
self._record(model, parser.input_tokens, parser.output_tokens, None)
|
||||
45
litellm/harness/errors.py
Normal file
45
litellm/harness/errors.py
Normal file
|
|
@ -0,0 +1,45 @@
|
|||
"""Exceptions raised by litellm.harness."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.harness.types import Result
|
||||
|
||||
|
||||
class HarnessError(Exception):
|
||||
"""Base class for every litellm.harness error."""
|
||||
|
||||
|
||||
class CapabilityUnsupported(HarnessError):
|
||||
"""The harness cannot do what was asked. Raised before the runtime starts."""
|
||||
|
||||
|
||||
class OptionsMismatch(HarnessError):
|
||||
"""Options for a different harness, or a native option LiteLLM manages itself."""
|
||||
|
||||
|
||||
class HarnessInstallFailed(HarnessError):
|
||||
"""The runtime is missing from the sandbox or failed to start."""
|
||||
|
||||
|
||||
class SandboxError(HarnessError):
|
||||
"""The sandbox failed to start, run a command, or reach the host."""
|
||||
|
||||
|
||||
class SessionClosed(HarnessError):
|
||||
"""A turn was started on a session that is closed or detached."""
|
||||
|
||||
|
||||
class StateIncompatible(HarnessError):
|
||||
"""resume() was given a State from another harness or an unreadable version."""
|
||||
|
||||
|
||||
class OutputInvalid(HarnessError):
|
||||
"""The final answer did not validate against output=."""
|
||||
|
||||
def __init__(self, message: str, raw: str, result: Result | None = None) -> None:
|
||||
super().__init__(message)
|
||||
self.raw = raw
|
||||
self.result = result
|
||||
35
litellm/harness/handlers/__init__.py
Normal file
35
litellm/harness/handlers/__init__.py
Normal file
|
|
@ -0,0 +1,35 @@
|
|||
"""Handlers run a harness config: CLI runtimes as subprocesses, Deep Agents in-process."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from litellm.harness.errors import HarnessError
|
||||
from litellm.harness.handlers.base import BaseHarnessHandler
|
||||
from litellm.harness.types import Harness, require_harness
|
||||
from litellm.llms.base_llm.harness.transformation import (
|
||||
BaseCLIHarnessConfig,
|
||||
BaseHarnessConfig,
|
||||
)
|
||||
from litellm.utils import ProviderConfigManager
|
||||
|
||||
|
||||
def get_harness_config(harness: Harness) -> BaseHarnessConfig:
|
||||
config = ProviderConfigManager.get_provider_harness_config(require_harness(harness))
|
||||
if config is None:
|
||||
raise HarnessError(f"No harness config registered for Harness.{harness.name}")
|
||||
return config
|
||||
|
||||
|
||||
def get_harness_handler(config: BaseHarnessConfig) -> BaseHarnessHandler:
|
||||
"""The handler that knows how to run this kind of config."""
|
||||
if isinstance(config, BaseCLIHarnessConfig):
|
||||
from litellm.harness.handlers.cli_handler import CLIHarnessHandler
|
||||
|
||||
return CLIHarnessHandler(config)
|
||||
if config.harness is Harness.DEEPAGENTS:
|
||||
from litellm.harness.handlers.deepagents_handler import DeepAgentsHandler
|
||||
|
||||
return DeepAgentsHandler(config)
|
||||
raise HarnessError(f"No handler for Harness.{config.harness.name}")
|
||||
|
||||
|
||||
__all__ = ("BaseHarnessHandler", "get_harness_config", "get_harness_handler")
|
||||
42
litellm/harness/handlers/base.py
Normal file
42
litellm/harness/handlers/base.py
Normal file
|
|
@ -0,0 +1,42 @@
|
|||
"""The handler interface the runtime drives. A handler owns I/O for one session."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from collections.abc import AsyncIterator
|
||||
from typing import Any
|
||||
|
||||
from litellm.harness.context import SessionContext
|
||||
from litellm.harness.errors import CapabilityUnsupported
|
||||
from litellm.harness.types import Event
|
||||
from litellm.llms.base_llm.harness.transformation import BaseHarnessConfig
|
||||
|
||||
|
||||
class BaseHarnessHandler(ABC):
|
||||
"""Runs one harness session. The config decides what to run; the handler runs it."""
|
||||
|
||||
def __init__(self, config: BaseHarnessConfig) -> None:
|
||||
self.config = config
|
||||
|
||||
@abstractmethod
|
||||
async def start(self, ctx: SessionContext) -> None:
|
||||
"""Prepare the runtime (config files, skills, agent build). Called again after an interrupt."""
|
||||
|
||||
@abstractmethod
|
||||
def turn(self, ctx: SessionContext, prompt: str) -> AsyncIterator[Event]:
|
||||
"""Run one turn and yield events (never Done). Sets ctx.final_text / ctx.output_json."""
|
||||
|
||||
@abstractmethod
|
||||
async def stop(self, ctx: SessionContext) -> None:
|
||||
"""Stop anything this handler started. Safe to call twice."""
|
||||
|
||||
@abstractmethod
|
||||
def native_session_id(self) -> str | None:
|
||||
"""The runtime's own session id, for State / resume."""
|
||||
|
||||
@abstractmethod
|
||||
async def resume(self, ctx: SessionContext, native_session_id: str) -> None:
|
||||
"""Continue the runtime's own session on the next turn."""
|
||||
|
||||
async def history(self, ctx: SessionContext) -> list[dict[str, Any]]: # mutable-ok: public history() API shape
|
||||
raise CapabilityUnsupported(f"Harness.{self.config.harness.name} does not expose history")
|
||||
161
litellm/harness/handlers/cli_handler.py
Normal file
161
litellm/harness/handlers/cli_handler.py
Normal file
|
|
@ -0,0 +1,161 @@
|
|||
"""
|
||||
Generic handler for CLI harnesses (Claude Code, Codex, OpenCode).
|
||||
|
||||
The config (`litellm/llms/<harness>/harness/transformation.py`) says what to run and how to
|
||||
read it; this handler does every sandbox and process operation: binary check, private dir,
|
||||
config files, persisted dirs, skills, spawning the turn, streaming stdout lines into the
|
||||
config's parser, collecting stderr, and killing the process on early exit.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import os
|
||||
from collections import deque
|
||||
from collections.abc import AsyncIterator, Sequence
|
||||
from typing import Any, Final
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.constants import HARNESS_STDERR_TAIL_LINES, HARNESS_STREAM_READ_CHUNK_BYTES
|
||||
from litellm.harness.context import SessionContext
|
||||
from litellm.harness.errors import HarnessInstallFailed, SandboxError
|
||||
from litellm.harness.handlers.base import BaseHarnessHandler
|
||||
from litellm.harness.sandbox.base import Process, Sandbox
|
||||
from litellm.harness.types import Event
|
||||
from litellm.llms.base_llm.harness.transformation import BaseCLIHarnessConfig, HarnessSessionSetup
|
||||
from litellm.llms.base_llm.harness.utils import decode_json_line, read_skill_files
|
||||
|
||||
# Link <private_dir>/<dir> to a LiteLLM-owned cache dir so a later session can resume.
|
||||
PERSIST_DIR_SCRIPT: Final = (
|
||||
'd="${HOME:-/tmp}/.cache/litellm-harness/$2"; mkdir -p "$d" && mkdir -p "$(dirname "$1")" && ln -sfn "$d" "$1"'
|
||||
)
|
||||
|
||||
|
||||
async def iter_stream_lines(stream: asyncio.StreamReader) -> AsyncIterator[bytes]:
|
||||
"""Newline-delimited lines without StreamReader's 64KiB readline limit."""
|
||||
buffer = b""
|
||||
while True:
|
||||
chunk = await stream.read(HARNESS_STREAM_READ_CHUNK_BYTES)
|
||||
if not chunk:
|
||||
break
|
||||
buffer += chunk
|
||||
*lines, buffer = buffer.split(b"\n")
|
||||
for line in lines:
|
||||
yield line
|
||||
if buffer:
|
||||
yield buffer
|
||||
|
||||
|
||||
async def drain_stderr(stream: asyncio.StreamReader, tail: deque[str]) -> None: # mutable-ok: stderr ring
|
||||
async for line in iter_stream_lines(stream):
|
||||
tail.append(line.decode("utf-8", errors="replace"))
|
||||
|
||||
|
||||
async def send_stdin(proc: Process, data: str) -> None:
|
||||
if proc.stdin is None:
|
||||
raise SandboxError("harness process has no stdin")
|
||||
proc.stdin.write(data.encode("utf-8"))
|
||||
await proc.stdin.drain()
|
||||
proc.stdin.close()
|
||||
|
||||
|
||||
async def private_dir_for(sandbox: Sandbox) -> str:
|
||||
tempdir = getattr(sandbox, "tempdir", None)
|
||||
if tempdir is None:
|
||||
raise SandboxError(f"{type(sandbox).__name__} has no tempdir(); CLI harnesses need a private config dir")
|
||||
path: str = await tempdir()
|
||||
return path
|
||||
|
||||
|
||||
def sandbox_path(private_dir: str, path: str) -> str:
|
||||
return path if path.startswith("/") else f"{private_dir}/{path}"
|
||||
|
||||
|
||||
async def persist_dir(sandbox: Sandbox, link_path: str, cache_subpath: str) -> None:
|
||||
script_args: Final = ("-c", PERSIST_DIR_SCRIPT, "sh", link_path, cache_subpath)
|
||||
cmd: Final = ["sh", *script_args] # mutable-ok: Sandbox.run takes list[str]
|
||||
run = await sandbox.run(cmd)
|
||||
if run.exit_code != 0:
|
||||
verbose_logger.debug(
|
||||
"harness: could not persist %s, resume across sessions disabled: %s", cache_subpath, run.stderr.strip()
|
||||
)
|
||||
|
||||
|
||||
async def copy_skills(sandbox: Sandbox, skills: Sequence[str], skills_root: str) -> None:
|
||||
for skill in skills:
|
||||
name = os.path.basename(os.path.realpath(os.fspath(skill)))
|
||||
for rel, data in await asyncio.to_thread(read_skill_files, skill):
|
||||
await sandbox.write(f"{skills_root}/{name}/{rel.replace(os.sep, '/')}", data)
|
||||
|
||||
|
||||
class CLIHarnessHandler(BaseHarnessHandler):
|
||||
config: BaseCLIHarnessConfig
|
||||
|
||||
def __init__(self, config: BaseCLIHarnessConfig) -> None:
|
||||
super().__init__(config)
|
||||
self._private_dir: str | None = None
|
||||
self._setup: HarnessSessionSetup | None = None
|
||||
self._native_id: str | None = None
|
||||
self._proc: Process | None = None
|
||||
|
||||
async def start(self, ctx: SessionContext) -> None:
|
||||
self.config.validate_environment(ctx)
|
||||
binary = self.config.get_binary()
|
||||
if not await ctx.sandbox.which(binary):
|
||||
raise HarnessInstallFailed(
|
||||
f"`{binary}` was not found on PATH in the sandbox. Install it with: {self.config.get_install_hint()}"
|
||||
)
|
||||
private_dir = await private_dir_for(ctx.sandbox)
|
||||
setup = self.config.transform_session_setup(ctx, private_dir)
|
||||
for link, cache_subpath in setup.persisted_dirs:
|
||||
await persist_dir(ctx.sandbox, sandbox_path(private_dir, link), cache_subpath)
|
||||
for rel_path, data in setup.files.items():
|
||||
await ctx.sandbox.write(sandbox_path(private_dir, rel_path), data)
|
||||
if ctx.skills and setup.skills_dir:
|
||||
await copy_skills(ctx.sandbox, tuple(ctx.skills), sandbox_path(private_dir, setup.skills_dir))
|
||||
self._private_dir = private_dir
|
||||
self._setup = setup
|
||||
|
||||
async def turn(self, ctx: SessionContext, prompt: str) -> AsyncIterator[Event]:
|
||||
if self._setup is None or self._private_dir is None:
|
||||
raise RuntimeError("CLIHarnessHandler.turn() called before start()")
|
||||
request = self.config.transform_turn_request(ctx, self._setup, self._private_dir, prompt, self._native_id)
|
||||
argv: Final = list(request.argv) # mutable-ok: Sandbox.exec takes list[str]
|
||||
proc = await ctx.sandbox.exec(argv, env=request.env, cwd=request.cwd)
|
||||
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()
|
||||
exit_code: int | None = None
|
||||
try:
|
||||
await send_stdin(proc, request.stdin)
|
||||
async for raw in iter_stream_lines(proc.stdout):
|
||||
line = decode_json_line(raw)
|
||||
if line is None:
|
||||
continue
|
||||
for event in self.config.transform_stream_line(line, state):
|
||||
yield event
|
||||
self._native_id = self.config.get_native_session_id(state) or self._native_id
|
||||
exit_code = await proc.wait()
|
||||
await stderr_task
|
||||
finally:
|
||||
self._proc = None
|
||||
if exit_code is None:
|
||||
# Consumer stopped early, timed out or errored: don't leave the runtime running.
|
||||
await proc.kill()
|
||||
if not stderr_task.done():
|
||||
stderr_task.cancel()
|
||||
response = self.config.transform_turn_response(ctx, state, exit_code, tuple(tail))
|
||||
ctx.final_text = response.final_text
|
||||
ctx.output_json = response.output_json
|
||||
|
||||
async def stop(self, ctx: SessionContext) -> None:
|
||||
proc, self._proc = self._proc, None
|
||||
if proc is not None:
|
||||
await proc.kill()
|
||||
|
||||
def native_session_id(self) -> str | None:
|
||||
return self._native_id
|
||||
|
||||
async def resume(self, ctx: SessionContext, native_session_id: str) -> None:
|
||||
self._native_id = native_session_id
|
||||
261
litellm/harness/handlers/deepagents_handler.py
Normal file
261
litellm/harness/handlers/deepagents_handler.py
Normal file
|
|
@ -0,0 +1,261 @@
|
|||
"""
|
||||
In-process handler for Deep Agents.
|
||||
|
||||
Deep Agents is a Python library, so there is no process or model endpoint: the handler
|
||||
builds the agent with a LiteLLM chat model, streams the LangGraph run, turns interrupts into
|
||||
Approval events and counts usage. Translation lives in
|
||||
`litellm/llms/deepagents/harness/transformation.py`.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import importlib
|
||||
from collections.abc import AsyncIterator, Mapping
|
||||
from dataclasses import dataclass
|
||||
from types import MappingProxyType, ModuleType
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from litellm.harness.context import SessionContext
|
||||
from litellm.harness.errors import HarnessError, HarnessInstallFailed
|
||||
from litellm.harness.handlers.base import BaseHarnessHandler
|
||||
from litellm.harness.handlers.cli_handler import copy_skills
|
||||
from litellm.harness.options import DeepAgentsOptions
|
||||
from litellm.harness.types import Approval, Event
|
||||
from litellm.llms.deepagents.harness.transformation import (
|
||||
EXECUTE_TOOLS,
|
||||
INSTALL_HINT,
|
||||
SKILLS_DIR,
|
||||
WRITE_TOOLS,
|
||||
TurnState,
|
||||
approval_requests,
|
||||
blocked_tools,
|
||||
chat_model_kwargs,
|
||||
decision,
|
||||
final_ai_text,
|
||||
interrupt_config,
|
||||
interrupts_in,
|
||||
normalized_tool_name,
|
||||
recursion_limit,
|
||||
stream_events,
|
||||
structured_json,
|
||||
update_events,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from langchain_core.callbacks import BaseCallbackHandler
|
||||
from langchain_core.language_models import BaseChatModel
|
||||
from langchain_core.runnables import RunnableConfig
|
||||
from langgraph.checkpoint.base import BaseCheckpointSaver
|
||||
from langgraph.graph.state import CompiledStateGraph
|
||||
from langgraph.types import Command
|
||||
|
||||
from litellm.llms.base_llm.harness.transformation import BaseHarnessConfig
|
||||
|
||||
_MODEL_NODE = "model"
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class DeepAgentsDeps:
|
||||
"""The optional-dependency entrypoints this handler uses."""
|
||||
|
||||
create_deep_agent: Any
|
||||
chat_litellm: Any
|
||||
checkpointer_cls: Any
|
||||
command_cls: Any
|
||||
subagent_defaults: Mapping[str, Any]
|
||||
convert_to_openai_messages: Any
|
||||
backend: ModuleType
|
||||
|
||||
|
||||
def load_deps() -> DeepAgentsDeps:
|
||||
"""Import deepagents + langchain-litellm, or raise HarnessInstallFailed."""
|
||||
try:
|
||||
deepagents = importlib.import_module("deepagents")
|
||||
subagents = importlib.import_module("deepagents.middleware.subagents")
|
||||
chat = importlib.import_module("langchain_litellm")
|
||||
memory = importlib.import_module("langgraph.checkpoint.memory")
|
||||
lg_types = importlib.import_module("langgraph.types")
|
||||
messages = importlib.import_module("langchain_core.messages")
|
||||
backend = importlib.import_module("litellm.llms.deepagents.harness.sandbox_backend")
|
||||
except ImportError as e:
|
||||
raise HarnessInstallFailed(f"{INSTALL_HINT} ({e})") from e
|
||||
return DeepAgentsDeps(
|
||||
create_deep_agent=deepagents.create_deep_agent,
|
||||
chat_litellm=chat.ChatLiteLLM,
|
||||
checkpointer_cls=memory.InMemorySaver,
|
||||
command_cls=lg_types.Command,
|
||||
subagent_defaults=subagents.GENERAL_PURPOSE_SUBAGENT,
|
||||
convert_to_openai_messages=messages.convert_to_openai_messages,
|
||||
backend=backend,
|
||||
)
|
||||
|
||||
|
||||
_SHARED_CHECKPOINTER: dict[str, Any] = {} # mutable-ok: process-wide lazy singleton slot for the in-memory checkpointer
|
||||
|
||||
|
||||
def shared_checkpointer(deps: DeepAgentsDeps) -> BaseCheckpointSaver:
|
||||
"""One in-memory checkpointer per process, so resume() works across sessions in-process."""
|
||||
saver = _SHARED_CHECKPOINTER.get("saver")
|
||||
if saver is None:
|
||||
saver = deps.checkpointer_cls()
|
||||
_SHARED_CHECKPOINTER["saver"] = saver
|
||||
return saver
|
||||
|
||||
|
||||
def build_chat_model(ctx: SessionContext, deps: DeepAgentsDeps) -> BaseChatModel:
|
||||
"""The LangChain chat model for this session. Tests monkeypatch this."""
|
||||
return deps.chat_litellm(**chat_model_kwargs(ctx))
|
||||
|
||||
|
||||
class DeepAgentsHandler(BaseHarnessHandler):
|
||||
def __init__(self, config: BaseHarnessConfig) -> None:
|
||||
super().__init__(config)
|
||||
self._deps: DeepAgentsDeps | None = None
|
||||
self._agent: Any = None
|
||||
self._thread_id: str | None = None
|
||||
self._skip_tools: frozenset[str] = frozenset()
|
||||
|
||||
async def start(self, ctx: SessionContext) -> None:
|
||||
self.config.validate_environment(ctx)
|
||||
deps = load_deps()
|
||||
self._deps = deps
|
||||
blocked = blocked_tools(ctx.permissions, ctx.disable_tools)
|
||||
backend = deps.backend.SandboxBackend(
|
||||
ctx.sandbox,
|
||||
loop=asyncio.get_running_loop(),
|
||||
writable=WRITE_TOOLS.isdisjoint(blocked),
|
||||
allow_execute=EXECUTE_TOOLS.isdisjoint(blocked),
|
||||
)
|
||||
self._agent = deps.create_deep_agent(
|
||||
model=build_chat_model(ctx, deps),
|
||||
tools=list(ctx.tools), # mutable-ok: deepagents create_deep_agent(tools=) takes a list
|
||||
system_prompt=ctx.instructions,
|
||||
middleware=self._middleware(deps, blocked),
|
||||
subagents=self._subagents(ctx, deps, blocked),
|
||||
skills=await self._install_skills(ctx),
|
||||
backend=backend,
|
||||
interrupt_on=interrupt_config(ctx.permissions, blocked),
|
||||
response_format=ctx.output,
|
||||
checkpointer=shared_checkpointer(deps),
|
||||
)
|
||||
self._skip_tools = frozenset({ctx.output.__name__}) if ctx.output is not None else frozenset()
|
||||
if self._thread_id is None:
|
||||
self._thread_id = ctx.session_id
|
||||
|
||||
async def stop(self, ctx: SessionContext) -> None:
|
||||
self._agent = None
|
||||
|
||||
def native_session_id(self) -> str | None:
|
||||
return self._thread_id
|
||||
|
||||
async def resume(self, ctx: SessionContext, native_session_id: str) -> None:
|
||||
self._thread_id = native_session_id
|
||||
|
||||
async def history(
|
||||
self, ctx: SessionContext
|
||||
) -> list[dict[str, Any]]: # mutable-ok: BaseHarnessHandler.history API returns OpenAI message dicts
|
||||
agent, deps = self._require_agent()
|
||||
snapshot = await agent.aget_state(self._run_config(ctx, None))
|
||||
messages = (snapshot.values or MappingProxyType({})).get("messages") or ()
|
||||
converted: list[dict[str, Any]] = deps.convert_to_openai_messages( # mutable-ok: LangChain returns a list
|
||||
messages
|
||||
)
|
||||
return converted
|
||||
|
||||
async def turn(self, ctx: SessionContext, prompt: str) -> AsyncIterator[Event]:
|
||||
agent, deps = self._require_agent()
|
||||
run_config = self._run_config(ctx, deps.backend.UsageCallback(ctx, ctx.model))
|
||||
user_message = {"role": "user", "content": prompt} # mutable-ok: LangGraph input message dict
|
||||
payload: dict[str, object] | Command = {"messages": [user_message]} # mutable-ok: LangGraph input state
|
||||
while True:
|
||||
state = TurnState()
|
||||
async for event in self._stream_pass(agent, payload, run_config, state):
|
||||
yield event
|
||||
if not state.interrupts:
|
||||
break
|
||||
resume: dict[str, Any] = {} # mutable-ok: Command(resume=) payload, filled per answered approval
|
||||
for interrupt in state.interrupts:
|
||||
decisions: list[dict[str, Any]] = [] # mutable-ok: HITL decisions collected across awaited approvals
|
||||
for request in approval_requests(getattr(interrupt, "value", None)):
|
||||
approval = Approval(
|
||||
tool=normalized_tool_name(str(request.get("name") or "")),
|
||||
input=dict(request.get("args") or ()), # mutable-ok: Approval.input is a public dict field
|
||||
)
|
||||
yield approval
|
||||
decisions.append(decision(*await approval.wait()))
|
||||
resume[interrupt.id] = {"decisions": decisions} # mutable-ok: LangGraph HITL resume payload
|
||||
payload = deps.command_cls(resume=resume)
|
||||
await self._finish_turn(ctx, agent, run_config)
|
||||
|
||||
async def _stream_pass(
|
||||
self,
|
||||
agent: CompiledStateGraph,
|
||||
payload: dict[str, object] | Command, # mutable-ok: LangGraph astream input type
|
||||
run_config: RunnableConfig,
|
||||
state: TurnState,
|
||||
) -> AsyncIterator[Event]:
|
||||
async for part in agent.astream(
|
||||
payload,
|
||||
run_config,
|
||||
stream_mode=["messages", "updates"], # mutable-ok: LangGraph stream_mode takes a list
|
||||
):
|
||||
# A list stream_mode yields (mode, chunk) tuples; LangGraph's overloads don't say so.
|
||||
if not isinstance(part, tuple) or len(part) != 2:
|
||||
continue
|
||||
mode, chunk = part
|
||||
if mode == "messages":
|
||||
message, meta = chunk
|
||||
if isinstance(meta, Mapping) and meta.get("langgraph_node") == _MODEL_NODE:
|
||||
for event in stream_events(message):
|
||||
yield event
|
||||
elif mode == "updates":
|
||||
state.interrupts = (*state.interrupts, *interrupts_in(chunk))
|
||||
for event in update_events(chunk, self._skip_tools):
|
||||
yield event
|
||||
|
||||
async def _finish_turn(self, ctx: SessionContext, agent: CompiledStateGraph, run_config: RunnableConfig) -> None:
|
||||
snapshot = await agent.aget_state(run_config)
|
||||
values = snapshot.values or MappingProxyType({})
|
||||
ctx.final_text = final_ai_text(values.get("messages") or ())
|
||||
if ctx.output is not None:
|
||||
ctx.output_json = structured_json(values.get("structured_response"))
|
||||
|
||||
def _require_agent(self) -> tuple[Any, DeepAgentsDeps]:
|
||||
if self._agent is None or self._deps is None:
|
||||
raise HarnessError("Deep Agents session is not started")
|
||||
return self._agent, self._deps
|
||||
|
||||
def _run_config(self, ctx: SessionContext, usage_callback: BaseCallbackHandler | None) -> RunnableConfig:
|
||||
run_config: RunnableConfig = {
|
||||
"configurable": {"thread_id": self._thread_id or ctx.session_id},
|
||||
"recursion_limit": recursion_limit(ctx),
|
||||
}
|
||||
if usage_callback is not None:
|
||||
run_config["callbacks"] = [usage_callback] # mutable-ok: LangChain RunnableConfig.callbacks is a list
|
||||
return run_config
|
||||
|
||||
@staticmethod
|
||||
def _middleware(deps: DeepAgentsDeps, blocked: frozenset[str]) -> list[Any]: # mutable-ok: deepagents API
|
||||
filters = (deps.backend.ToolFilterMiddleware(blocked),) if blocked else ()
|
||||
return list(filters) # mutable-ok: deepagents create_deep_agent(middleware=) takes a list
|
||||
|
||||
def _subagents(
|
||||
self, ctx: SessionContext, deps: DeepAgentsDeps, blocked: frozenset[str]
|
||||
) -> list[Any]: # mutable-ok: deepagents create_deep_agent(subagents=) takes a list
|
||||
"""User subagents, plus a general-purpose one that honours disable_tools when set."""
|
||||
options = ctx.options if isinstance(ctx.options, DeepAgentsOptions) else None
|
||||
user_subagents = tuple(options.subagents) if options is not None else ()
|
||||
has_general = any(
|
||||
isinstance(s, Mapping) and s.get("name") == deps.subagent_defaults["name"] for s in user_subagents
|
||||
)
|
||||
spec = {**deps.subagent_defaults, "middleware": self._middleware(deps, blocked)} # mutable-ok: SubAgent dict
|
||||
general = (spec,) if blocked and not has_general else ()
|
||||
return [*general, *user_subagents] # mutable-ok: deepagents create_deep_agent(subagents=) takes a list
|
||||
|
||||
@staticmethod
|
||||
async def _install_skills(ctx: SessionContext) -> list[str] | None: # mutable-ok: deepagents skills= takes a list
|
||||
if not ctx.skills:
|
||||
return None
|
||||
await copy_skills(ctx.sandbox, ctx.skills, f"{ctx.sandbox.workdir}/{SKILLS_DIR}")
|
||||
return [f"/{SKILLS_DIR}/"] # mutable-ok: deepagents create_deep_agent(skills=) takes a list
|
||||
37
litellm/harness/options.py
Normal file
37
litellm/harness/options.py
Normal file
|
|
@ -0,0 +1,37 @@
|
|||
"""Typed per-harness options. Settings that only make sense for one runtime live here."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Mapping, Sequence
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, Literal
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ClaudeCodeOptions:
|
||||
config: Mapping[str, Any] = field(default_factory=dict)
|
||||
env: Mapping[str, str] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class CodexOptions:
|
||||
reasoning_effort: Literal["low", "medium", "high", "xhigh"] | None = None
|
||||
web_search: bool = False
|
||||
config: Mapping[str, Any] = 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)
|
||||
env: Mapping[str, str] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class DeepAgentsOptions:
|
||||
subagents: Sequence[Any] = ()
|
||||
recursion_limit: int | None = None
|
||||
|
||||
|
||||
HarnessOptions = ClaudeCodeOptions | CodexOptions | OpenCodeOptions | DeepAgentsOptions
|
||||
1070
litellm/harness/runtime.py
Normal file
1070
litellm/harness/runtime.py
Normal file
File diff suppressed because it is too large
Load diff
25
litellm/harness/sandbox/__init__.py
Normal file
25
litellm/harness/sandbox/__init__.py
Normal file
|
|
@ -0,0 +1,25 @@
|
|||
"""Sandboxes for litellm.harness: where the runtime runs and which files it can touch."""
|
||||
|
||||
from litellm.harness.sandbox.base import CompletedRun, Process, Sandbox
|
||||
from litellm.harness.sandbox.docker import DockerSandbox, docker
|
||||
from litellm.harness.sandbox.local import LocalSandbox, local
|
||||
from litellm.harness.sandbox.snapshot import (
|
||||
build_file_changes,
|
||||
capture_text_contents,
|
||||
diff_snapshots,
|
||||
snapshot_local,
|
||||
)
|
||||
|
||||
__all__ = (
|
||||
"CompletedRun",
|
||||
"DockerSandbox",
|
||||
"LocalSandbox",
|
||||
"Process",
|
||||
"Sandbox",
|
||||
"build_file_changes",
|
||||
"capture_text_contents",
|
||||
"diff_snapshots",
|
||||
"docker",
|
||||
"local",
|
||||
"snapshot_local",
|
||||
)
|
||||
62
litellm/harness/sandbox/base.py
Normal file
62
litellm/harness/sandbox/base.py
Normal file
|
|
@ -0,0 +1,62 @@
|
|||
"""The Sandbox protocol: where a harness runtime runs and which files it can touch."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from collections.abc import Mapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
from typing import Protocol, runtime_checkable
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class CompletedRun:
|
||||
stdout: str
|
||||
stderr: str
|
||||
exit_code: int
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class Process(Protocol):
|
||||
stdin: asyncio.StreamWriter | None
|
||||
stdout: asyncio.StreamReader
|
||||
stderr: asyncio.StreamReader
|
||||
|
||||
async def wait(self) -> int: ...
|
||||
|
||||
async def kill(self) -> None: ...
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class Sandbox(Protocol):
|
||||
workdir: str
|
||||
|
||||
async def exec(
|
||||
self,
|
||||
cmd: Sequence[str],
|
||||
*,
|
||||
env: Mapping[str, str] | None = None,
|
||||
cwd: str | None = None,
|
||||
) -> Process: ...
|
||||
|
||||
async def run(
|
||||
self,
|
||||
cmd: Sequence[str],
|
||||
*,
|
||||
env: Mapping[str, str] | None = None,
|
||||
cwd: str | None = None,
|
||||
timeout: float | None = None,
|
||||
) -> CompletedRun: ...
|
||||
|
||||
async def read(self, path: str) -> bytes: ...
|
||||
|
||||
async def write(self, path: str, data: bytes) -> None: ...
|
||||
|
||||
def host_url(self, port: int) -> str: ...
|
||||
|
||||
async def which(self, binary: str) -> str | None: ...
|
||||
|
||||
async def snapshot(self) -> Mapping[str, str]: ...
|
||||
|
||||
async def tempdir(self) -> str: ...
|
||||
|
||||
async def close(self) -> None: ...
|
||||
298
litellm/harness/sandbox/docker.py
Normal file
298
litellm/harness/sandbox/docker.py
Normal file
|
|
@ -0,0 +1,298 @@
|
|||
"""DockerSandbox: run the harness runtime inside a container via the docker CLI."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import os
|
||||
import posixpath
|
||||
import shutil
|
||||
from collections.abc import Mapping, Sequence
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
|
||||
from litellm.constants import HARNESS_SNAPSHOT_SKIP_DIRS
|
||||
from litellm.harness.errors import SandboxError
|
||||
from litellm.harness.sandbox.base import CompletedRun
|
||||
from litellm.harness.sandbox.local import SubprocessHandle, collect_output
|
||||
from litellm.harness.sandbox.snapshot import HARNESS_SNAPSHOT_MAX_FILE_BYTES
|
||||
|
||||
DOCKER_HOST_ALIAS: Final = "host.docker.internal"
|
||||
_SHA256_HEX_LEN: Final = 64
|
||||
_WRITE_SCRIPT: Final = 'mkdir -p "$(dirname "$1")" && cat > "$1"'
|
||||
_WHICH_SCRIPT: Final = 'command -v "$1"'
|
||||
|
||||
|
||||
def _snapshot_script() -> str:
|
||||
prune = " -o ".join(f"-name '{name}'" for name in sorted(HARNESS_SNAPSHOT_SKIP_DIRS))
|
||||
return (
|
||||
'cd "$1" && find . -type d \\( '
|
||||
+ prune
|
||||
+ " \\) -prune -o -type f -size -"
|
||||
+ f"{HARNESS_SNAPSHOT_MAX_FILE_BYTES + 1}c"
|
||||
+ " -exec sha256sum {} +"
|
||||
)
|
||||
|
||||
|
||||
def parse_sha256sum(output: str) -> Mapping[str, str]:
|
||||
"""Parse `sha256sum` lines ("<hex> ./rel/path") into {rel/path: hex}."""
|
||||
return MappingProxyType(
|
||||
{
|
||||
line[_SHA256_HEX_LEN + 2 :].removeprefix("./"): line[:_SHA256_HEX_LEN]
|
||||
for line in output.splitlines()
|
||||
if len(line) > _SHA256_HEX_LEN + 2
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
class DockerSandbox:
|
||||
"""Sandbox backed by a long-lived `sleep infinity` container."""
|
||||
|
||||
# Harness configs read this to skip a runtime's own nested OS sandbox.
|
||||
is_container = True
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
image: str,
|
||||
mounts: Mapping[str | os.PathLike[str], str] | None = None,
|
||||
workdir: str = "/workspace",
|
||||
env: Mapping[str, str] | None = None,
|
||||
name: str | None = None,
|
||||
) -> None:
|
||||
if not image:
|
||||
raise SandboxError("docker sandbox needs an image")
|
||||
if not posixpath.isabs(workdir):
|
||||
raise SandboxError(f"docker workdir must be absolute: {workdir}")
|
||||
self.image = image
|
||||
self.workdir: str = posixpath.normpath(workdir)
|
||||
self.mounts: Mapping[str, str] = MappingProxyType(
|
||||
{os.path.abspath(os.fspath(host)): container for host, container in (mounts.items() if mounts else ())}
|
||||
)
|
||||
self.env: Mapping[str, str] = MappingProxyType(dict(env or ()))
|
||||
self.name = name
|
||||
self.container_id: str | None = None
|
||||
self._start_lock = asyncio.Lock()
|
||||
self._processes: set[SubprocessHandle] = set() # mutable-ok: live-process registry (add/discard)
|
||||
self._closed = False
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"DockerSandbox({self.image!r}, workdir={self.workdir!r})"
|
||||
|
||||
# -- docker CLI plumbing (tests monkeypatch these two) ---------------------
|
||||
|
||||
def _docker_binary(self) -> str:
|
||||
binary = shutil.which("docker")
|
||||
if binary is None:
|
||||
raise SandboxError(
|
||||
"docker sandbox requires the `docker` CLI on PATH; install Docker or use sandbox.local(path)"
|
||||
)
|
||||
return binary
|
||||
|
||||
async def _spawn(self, args: Sequence[str]) -> SubprocessHandle:
|
||||
"""Start `docker <args>` with stdin/stdout/stderr pipes."""
|
||||
try:
|
||||
proc = await asyncio.create_subprocess_exec(
|
||||
self._docker_binary(),
|
||||
*args,
|
||||
stdin=asyncio.subprocess.PIPE,
|
||||
stdout=asyncio.subprocess.PIPE,
|
||||
stderr=asyncio.subprocess.PIPE,
|
||||
)
|
||||
except (FileNotFoundError, PermissionError) as exc:
|
||||
raise SandboxError(f"could not run docker: {exc}") from exc
|
||||
return SubprocessHandle(proc)
|
||||
|
||||
async def _docker(
|
||||
self,
|
||||
args: Sequence[str],
|
||||
*,
|
||||
input: bytes | None = None,
|
||||
timeout: float | None = None,
|
||||
) -> tuple[int, bytes, bytes]:
|
||||
"""Run `docker <args>` to completion; returns (exit_code, stdout, stderr)."""
|
||||
handle = await self._spawn(args)
|
||||
try:
|
||||
return await asyncio.wait_for(_communicate(handle, input), timeout)
|
||||
except asyncio.TimeoutError:
|
||||
await handle.kill()
|
||||
raise SandboxError(f"docker {args[0]} timed out after {timeout}s")
|
||||
|
||||
# -- command construction --------------------------------------------------
|
||||
|
||||
def run_args(
|
||||
self,
|
||||
) -> list[str]: # mutable-ok: argv is returned as a list, the shape callers and tests compare against
|
||||
name_args = ("--name", self.name) if self.name else ()
|
||||
mount_args = tuple(
|
||||
arg for host, container in self.mounts.items() for arg in ("-v", f"{host}:{container}")
|
||||
) # comprehension-ok: flattens (flag, value) pairs into argv
|
||||
env_args = tuple(
|
||||
arg for key, value in self.env.items() for arg in ("-e", f"{key}={value}")
|
||||
) # comprehension-ok: flattens (flag, value) pairs into argv
|
||||
return [ # mutable-ok: argv is returned as a list, the shape callers and tests compare against
|
||||
"run",
|
||||
"-d",
|
||||
"--rm",
|
||||
f"--add-host={DOCKER_HOST_ALIAS}:host-gateway",
|
||||
*name_args,
|
||||
*mount_args,
|
||||
*env_args,
|
||||
"-w",
|
||||
self.workdir,
|
||||
self.image,
|
||||
"sleep",
|
||||
"infinity",
|
||||
]
|
||||
|
||||
def exec_args(
|
||||
self,
|
||||
container_id: str,
|
||||
cmd: Sequence[str],
|
||||
*,
|
||||
env: Mapping[str, str] | None = None,
|
||||
cwd: str | None = None,
|
||||
) -> list[str]: # mutable-ok: argv is returned as a list, the shape callers and tests compare against
|
||||
env_args = tuple(
|
||||
arg for key, value in (env.items() if env else ()) for arg in ("-e", f"{key}={value}")
|
||||
) # comprehension-ok: flattens (flag, value) pairs into argv
|
||||
return [ # mutable-ok: argv is returned as a list, the shape callers and tests compare against
|
||||
"exec",
|
||||
"-i",
|
||||
"-w",
|
||||
self.container_path(cwd or self.workdir),
|
||||
*env_args,
|
||||
container_id,
|
||||
*cmd,
|
||||
]
|
||||
|
||||
def container_path(self, path: str) -> str:
|
||||
"""Absolute container path; relative paths resolve against workdir."""
|
||||
joined = path if posixpath.isabs(path) else posixpath.join(self.workdir, path)
|
||||
return posixpath.normpath(joined)
|
||||
|
||||
# -- lifecycle ---------------------------------------------------------------
|
||||
|
||||
async def start(self) -> str:
|
||||
"""Start the container if needed and return its id."""
|
||||
if self._closed:
|
||||
raise SandboxError("sandbox is closed")
|
||||
async with self._start_lock:
|
||||
if self.container_id is not None:
|
||||
return self.container_id
|
||||
code, out, err = await self._docker(self.run_args())
|
||||
if code != 0:
|
||||
raise SandboxError(f"docker run {self.image} failed ({code}): {err.decode(errors='replace').strip()}")
|
||||
container_id = out.decode().strip()
|
||||
if not container_id:
|
||||
raise SandboxError("docker run returned no container id")
|
||||
self.container_id = container_id
|
||||
return container_id
|
||||
|
||||
async def _exec_capture(self, cmd: Sequence[str], *, input: bytes | None = None) -> tuple[int, bytes, bytes]:
|
||||
container_id = await self.start()
|
||||
return await self._docker(self.exec_args(container_id, cmd), input=input)
|
||||
|
||||
# -- Sandbox protocol --------------------------------------------------------
|
||||
|
||||
async def exec(
|
||||
self,
|
||||
cmd: Sequence[str],
|
||||
*,
|
||||
env: Mapping[str, str] | None = None,
|
||||
cwd: str | None = None,
|
||||
) -> SubprocessHandle:
|
||||
if not cmd:
|
||||
raise SandboxError("exec() needs a non-empty command")
|
||||
container_id = await self.start()
|
||||
handle = await self._spawn(self.exec_args(container_id, cmd, env=env, cwd=cwd))
|
||||
self._processes.add(handle)
|
||||
return handle
|
||||
|
||||
async def run(
|
||||
self,
|
||||
cmd: Sequence[str],
|
||||
*,
|
||||
env: Mapping[str, str] | None = None,
|
||||
cwd: str | None = None,
|
||||
timeout: float | None = None,
|
||||
) -> CompletedRun:
|
||||
handle = await self.exec(cmd, env=env, cwd=cwd)
|
||||
try:
|
||||
return await collect_output(handle, cmd, timeout)
|
||||
finally:
|
||||
self._processes.discard(handle)
|
||||
|
||||
async def read(self, path: str) -> bytes:
|
||||
target = self.container_path(path)
|
||||
code, out, err = await self._exec_capture(("cat", target))
|
||||
if code != 0:
|
||||
raise SandboxError(f"could not read {target}: {err.decode(errors='replace').strip()}")
|
||||
return out
|
||||
|
||||
async def write(self, path: str, data: bytes) -> None:
|
||||
target = self.container_path(path)
|
||||
code, _, err = await self._exec_capture(("sh", "-c", _WRITE_SCRIPT, "sh", target), input=data)
|
||||
if code != 0:
|
||||
raise SandboxError(f"could not write {target}: {err.decode(errors='replace').strip()}")
|
||||
|
||||
def host_url(self, port: int) -> str:
|
||||
return f"http://{DOCKER_HOST_ALIAS}:{port}"
|
||||
|
||||
async def which(self, binary: str) -> str | None:
|
||||
code, out, _ = await self._exec_capture(("sh", "-lc", _WHICH_SCRIPT, "sh", binary))
|
||||
found = out.decode(errors="replace").strip()
|
||||
return found if code == 0 and found else None
|
||||
|
||||
async def tempdir(self) -> str:
|
||||
"""A fresh `mktemp -d` directory inside the container."""
|
||||
code, out, err = await self._exec_capture(("mktemp", "-d"))
|
||||
path = out.decode(errors="replace").strip()
|
||||
if code != 0 or not path:
|
||||
raise SandboxError(f"mktemp -d failed: {err.decode(errors='replace').strip()}")
|
||||
return path
|
||||
|
||||
async def snapshot(self) -> Mapping[str, str]:
|
||||
code, out, err = await self._exec_capture(("sh", "-c", _snapshot_script(), "sh", self.workdir))
|
||||
if code != 0:
|
||||
raise SandboxError(f"snapshot failed: {err.decode(errors='replace').strip()}")
|
||||
return parse_sha256sum(out.decode("utf-8", errors="replace"))
|
||||
|
||||
async def close(self) -> None:
|
||||
if self._closed:
|
||||
return
|
||||
self._closed = True
|
||||
live = tuple(h for h in self._processes if h.returncode is None)
|
||||
await asyncio.gather(*(h.kill() for h in live), return_exceptions=True)
|
||||
self._processes.clear()
|
||||
if self.container_id is not None:
|
||||
container_id, self.container_id = self.container_id, None
|
||||
await self._docker(
|
||||
["rm", "-f", container_id] # mutable-ok: argv list, the shape _spawn records and tests assert on
|
||||
)
|
||||
|
||||
async def __aenter__(self) -> DockerSandbox:
|
||||
await self.start()
|
||||
return self
|
||||
|
||||
async def __aexit__(self, *exc_info: object) -> None:
|
||||
await self.close()
|
||||
|
||||
|
||||
async def _communicate(handle: SubprocessHandle, data: bytes | None) -> tuple[int, bytes, bytes]:
|
||||
if handle.stdin is not None:
|
||||
if data:
|
||||
handle.stdin.write(data)
|
||||
await handle.stdin.drain()
|
||||
handle.stdin.close()
|
||||
stdout, stderr = await asyncio.gather(handle.stdout.read(), handle.stderr.read())
|
||||
return await handle.wait(), stdout, stderr
|
||||
|
||||
|
||||
def docker(
|
||||
image: str,
|
||||
mounts: Mapping[str | os.PathLike[str], str] | None = None,
|
||||
workdir: str = "/workspace",
|
||||
env: Mapping[str, str] | None = None,
|
||||
name: str | None = None,
|
||||
) -> DockerSandbox:
|
||||
"""Sandbox in a new container of `image`, started lazily on first use."""
|
||||
return DockerSandbox(image, mounts=mounts, workdir=workdir, env=env, name=name)
|
||||
277
litellm/harness/sandbox/local.py
Normal file
277
litellm/harness/sandbox/local.py
Normal file
|
|
@ -0,0 +1,277 @@
|
|||
"""LocalSandbox: run the harness runtime as a subprocess on this machine."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import itertools
|
||||
import os
|
||||
import shutil
|
||||
import signal
|
||||
import tempfile
|
||||
from collections.abc import Mapping, Sequence
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
|
||||
from litellm.constants import HARNESS_PROCESS_KILL_GRACE_SECONDS
|
||||
from litellm.harness.errors import SandboxError
|
||||
from litellm.harness.sandbox.base import CompletedRun
|
||||
from litellm.harness.sandbox.snapshot import snapshot_local
|
||||
|
||||
_SECRET_PREFIXES: Final = (
|
||||
"ANTHROPIC_",
|
||||
"OPENAI_",
|
||||
"LITELLM_",
|
||||
"AZURE_",
|
||||
"AWS_",
|
||||
"GEMINI_",
|
||||
"CODEX_",
|
||||
"CURSOR_",
|
||||
"VERTEX",
|
||||
# A parent Claude Code session's socket/session vars make a child `claude` attach to
|
||||
# the parent's login instead of the harness token.
|
||||
"CLAUDE_CODE_",
|
||||
"CLAUDE_PID",
|
||||
"CLAUDECODE",
|
||||
)
|
||||
_SECRET_NAMES: Final = frozenset({"GOOGLE_API_KEY", "GOOGLE_APPLICATION_CREDENTIALS"})
|
||||
_SECRET_SUBSTRINGS: Final = ("API_KEY", "TOKEN", "SECRET")
|
||||
_TEMPDIR_PREFIX: Final = "litellm-harness-"
|
||||
|
||||
|
||||
def is_secret_env_name(name: str) -> bool:
|
||||
"""True if an env var name looks like a provider credential."""
|
||||
upper = name.upper()
|
||||
if upper in _SECRET_NAMES or upper.startswith(_SECRET_PREFIXES):
|
||||
return True
|
||||
return any(part in upper for part in _SECRET_SUBSTRINGS)
|
||||
|
||||
|
||||
def filtered_environ(
|
||||
base: Mapping[str, str] | None = None,
|
||||
extra: Mapping[str, str] | None = None,
|
||||
) -> Mapping[str, str]:
|
||||
"""base (default os.environ) without provider secrets, then extra on top."""
|
||||
source = os.environ if base is None else base
|
||||
kept = ((k, v) for k, v in source.items() if not is_secret_env_name(k))
|
||||
overlay = extra.items() if extra else ()
|
||||
return MappingProxyType(dict(itertools.chain(kept, overlay)))
|
||||
|
||||
|
||||
def _signal_process(proc: asyncio.subprocess.Process, sig: int) -> None:
|
||||
try:
|
||||
os.killpg(proc.pid, sig)
|
||||
except (ProcessLookupError, PermissionError, OSError):
|
||||
try:
|
||||
proc.send_signal(sig)
|
||||
except ProcessLookupError:
|
||||
pass
|
||||
|
||||
|
||||
class SubprocessHandle:
|
||||
"""Process-protocol wrapper around an asyncio subprocess."""
|
||||
|
||||
def __init__(self, proc: asyncio.subprocess.Process) -> None:
|
||||
if proc.stdout is None or proc.stderr is None:
|
||||
raise SandboxError("subprocess was started without stdout/stderr pipes")
|
||||
self._proc = proc
|
||||
self.stdin: asyncio.StreamWriter | None = proc.stdin
|
||||
self.stdout: asyncio.StreamReader = proc.stdout
|
||||
self.stderr: asyncio.StreamReader = proc.stderr
|
||||
|
||||
@property
|
||||
def pid(self) -> int:
|
||||
return self._proc.pid
|
||||
|
||||
@property
|
||||
def returncode(self) -> int | None:
|
||||
return self._proc.returncode
|
||||
|
||||
async def wait(self) -> int:
|
||||
return await self._proc.wait()
|
||||
|
||||
async def kill(self) -> None:
|
||||
"""SIGTERM, wait HARNESS_PROCESS_KILL_GRACE_SECONDS, then SIGKILL."""
|
||||
if self._proc.returncode is not None:
|
||||
return
|
||||
_signal_process(self._proc, signal.SIGTERM)
|
||||
try:
|
||||
await asyncio.wait_for(self._proc.wait(), timeout=HARNESS_PROCESS_KILL_GRACE_SECONDS)
|
||||
return
|
||||
except asyncio.TimeoutError:
|
||||
pass
|
||||
_signal_process(self._proc, signal.SIGKILL)
|
||||
await self._proc.wait()
|
||||
|
||||
|
||||
async def _read_all(handle: SubprocessHandle) -> tuple[bytes, bytes, int]:
|
||||
if handle.stdin is not None:
|
||||
handle.stdin.close()
|
||||
stdout, stderr = await asyncio.gather(handle.stdout.read(), handle.stderr.read())
|
||||
exit_code = await handle.wait()
|
||||
return stdout, stderr, exit_code
|
||||
|
||||
|
||||
async def collect_output(handle: SubprocessHandle, cmd: Sequence[str], timeout: float | None) -> CompletedRun:
|
||||
"""Close stdin, read stdout/stderr to EOF; kill and raise SandboxError on timeout."""
|
||||
try:
|
||||
stdout, stderr, code = await asyncio.wait_for(_read_all(handle), timeout)
|
||||
except asyncio.TimeoutError:
|
||||
await handle.kill()
|
||||
raise SandboxError(f"command timed out after {timeout}s: {cmd[0]}")
|
||||
return CompletedRun(
|
||||
stdout=stdout.decode("utf-8", errors="replace"),
|
||||
stderr=stderr.decode("utf-8", errors="replace"),
|
||||
exit_code=code,
|
||||
)
|
||||
|
||||
|
||||
def _is_within(path: str, root: str) -> bool:
|
||||
return path == root or path.startswith(root.rstrip(os.sep) + os.sep)
|
||||
|
||||
|
||||
class LocalSandbox:
|
||||
"""Sandbox backed by the local filesystem and asyncio subprocesses."""
|
||||
|
||||
def __init__(self, path: str | os.PathLike[str]) -> None:
|
||||
resolved = os.path.realpath(os.path.abspath(os.fspath(path)))
|
||||
if not os.path.isdir(resolved):
|
||||
raise SandboxError(f"sandbox path does not exist or is not a directory: {resolved}")
|
||||
self.workdir: str = resolved
|
||||
self._processes: set[SubprocessHandle] = set() # mutable-ok: live-process registry (add/discard)
|
||||
self._tempdirs: list[str] = [] # mutable-ok: tempdirs created on demand by tempdir(), removed on close()
|
||||
self._closed = False
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"LocalSandbox({self.workdir!r})"
|
||||
|
||||
def _check_open(self) -> None:
|
||||
if self._closed:
|
||||
raise SandboxError("sandbox is closed")
|
||||
|
||||
def _allowed_roots(self) -> tuple[str, ...]:
|
||||
return (self.workdir, *self._tempdirs)
|
||||
|
||||
def resolve_path(self, path: str) -> str:
|
||||
"""Absolute real path for path; SandboxError if it escapes the sandbox."""
|
||||
joined = path if os.path.isabs(path) else os.path.join(self.workdir, path)
|
||||
real = os.path.realpath(joined)
|
||||
if not any(_is_within(real, root) for root in self._allowed_roots()):
|
||||
raise SandboxError(f"path escapes the sandbox: {path}")
|
||||
return real
|
||||
|
||||
def _resolve_cwd(self, cwd: str | None) -> str:
|
||||
if cwd is None:
|
||||
return self.workdir
|
||||
resolved = self.resolve_path(cwd)
|
||||
if not os.path.isdir(resolved):
|
||||
raise SandboxError(f"cwd is not a directory: {cwd}")
|
||||
return resolved
|
||||
|
||||
def child_env(self, env: Mapping[str, str] | None = None) -> Mapping[str, str]:
|
||||
return filtered_environ(extra=env)
|
||||
|
||||
async def exec(
|
||||
self,
|
||||
cmd: Sequence[str],
|
||||
*,
|
||||
env: Mapping[str, str] | None = None,
|
||||
cwd: str | None = None,
|
||||
) -> SubprocessHandle:
|
||||
self._check_open()
|
||||
if not cmd:
|
||||
raise SandboxError("exec() needs a non-empty command")
|
||||
try:
|
||||
proc = await asyncio.create_subprocess_exec(
|
||||
*cmd,
|
||||
stdin=asyncio.subprocess.PIPE,
|
||||
stdout=asyncio.subprocess.PIPE,
|
||||
stderr=asyncio.subprocess.PIPE,
|
||||
cwd=self._resolve_cwd(cwd),
|
||||
env=self.child_env(env),
|
||||
start_new_session=True,
|
||||
)
|
||||
except (FileNotFoundError, PermissionError) as exc:
|
||||
raise SandboxError(f"could not start {cmd[0]}: {exc}") from exc
|
||||
handle = SubprocessHandle(proc)
|
||||
self._processes.add(handle)
|
||||
return handle
|
||||
|
||||
async def run(
|
||||
self,
|
||||
cmd: Sequence[str],
|
||||
*,
|
||||
env: Mapping[str, str] | None = None,
|
||||
cwd: str | None = None,
|
||||
timeout: float | None = None,
|
||||
) -> CompletedRun:
|
||||
handle = await self.exec(cmd, env=env, cwd=cwd)
|
||||
try:
|
||||
return await collect_output(handle, cmd, timeout)
|
||||
finally:
|
||||
self._processes.discard(handle)
|
||||
|
||||
async def read(self, path: str) -> bytes:
|
||||
self._check_open()
|
||||
resolved = self.resolve_path(path)
|
||||
try:
|
||||
return await asyncio.to_thread(_read_bytes, resolved)
|
||||
except OSError as exc:
|
||||
raise SandboxError(f"could not read {path}: {exc}") from exc
|
||||
|
||||
async def write(self, path: str, data: bytes) -> None:
|
||||
self._check_open()
|
||||
resolved = self.resolve_path(path)
|
||||
try:
|
||||
await asyncio.to_thread(_write_bytes, resolved, data)
|
||||
except OSError as exc:
|
||||
raise SandboxError(f"could not write {path}: {exc}") from exc
|
||||
|
||||
def host_url(self, port: int) -> str:
|
||||
return f"http://127.0.0.1:{port}"
|
||||
|
||||
async def which(self, binary: str) -> str | None:
|
||||
return shutil.which(binary, path=self.child_env().get("PATH"))
|
||||
|
||||
async def tempdir(self) -> str:
|
||||
"""A private temp dir (e.g. for CODEX_HOME), removed on close()."""
|
||||
self._check_open()
|
||||
path = os.path.realpath(tempfile.mkdtemp(prefix=_TEMPDIR_PREFIX))
|
||||
self._tempdirs.append(path)
|
||||
return path
|
||||
|
||||
async def snapshot(self) -> Mapping[str, str]:
|
||||
self._check_open()
|
||||
return await snapshot_local(self.workdir)
|
||||
|
||||
async def close(self) -> None:
|
||||
if self._closed:
|
||||
return
|
||||
self._closed = True
|
||||
live = tuple(h for h in self._processes if h.returncode is None)
|
||||
await asyncio.gather(*(h.kill() for h in live), return_exceptions=True)
|
||||
self._processes.clear()
|
||||
for path in self._tempdirs:
|
||||
shutil.rmtree(path, ignore_errors=True)
|
||||
self._tempdirs.clear()
|
||||
|
||||
async def __aenter__(self) -> LocalSandbox:
|
||||
return self
|
||||
|
||||
async def __aexit__(self, *exc_info: object) -> None:
|
||||
await self.close()
|
||||
|
||||
|
||||
def _read_bytes(path: str) -> bytes:
|
||||
with open(path, "rb") as fh:
|
||||
return fh.read()
|
||||
|
||||
|
||||
def _write_bytes(path: str, data: bytes) -> None:
|
||||
os.makedirs(os.path.dirname(path), exist_ok=True)
|
||||
with open(path, "wb") as fh:
|
||||
fh.write(data)
|
||||
|
||||
|
||||
def local(path: str | os.PathLike[str]) -> LocalSandbox:
|
||||
"""Sandbox rooted at an existing local directory."""
|
||||
return LocalSandbox(path)
|
||||
183
litellm/harness/sandbox/snapshot.py
Normal file
183
litellm/harness/sandbox/snapshot.py
Normal file
|
|
@ -0,0 +1,183 @@
|
|||
"""Workspace snapshots and FileChange construction.
|
||||
|
||||
A snapshot maps a workspace-relative POSIX path to the sha256 of its contents.
|
||||
Diffing two snapshots tells us which files a turn created, modified or deleted;
|
||||
`build_file_changes` turns that into `FileChange` events with unified diffs for
|
||||
small text files.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import difflib
|
||||
import functools
|
||||
import hashlib
|
||||
import os
|
||||
from collections.abc import Iterator, Mapping
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Final
|
||||
|
||||
from litellm.constants import HARNESS_MAX_DIFF_BYTES, HARNESS_SNAPSHOT_SKIP_DIRS
|
||||
from litellm.harness.errors import HarnessError
|
||||
from litellm.harness.types import FileChange, FileChangeKind
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.harness.sandbox.base import Sandbox
|
||||
|
||||
# Files larger than this are left out of snapshots entirely.
|
||||
HARNESS_SNAPSHOT_MAX_FILE_BYTES: Final = 50 * 1024 * 1024
|
||||
# Upper bound on bytes read by capture_text_contents() for one turn.
|
||||
HARNESS_SNAPSHOT_MAX_TOTAL_BYTES: Final = 16 * 1024 * 1024
|
||||
_HASH_CHUNK_BYTES: Final = 1024 * 1024
|
||||
_NO_NEWLINE_MARKER: Final = "\\ No newline at end of file\n"
|
||||
|
||||
|
||||
def _hash_file(path: str) -> str:
|
||||
digest = hashlib.sha256()
|
||||
with open(path, "rb") as fh:
|
||||
for chunk in iter(functools.partial(fh.read, _HASH_CHUNK_BYTES), b""):
|
||||
digest.update(chunk)
|
||||
return digest.hexdigest()
|
||||
|
||||
|
||||
def _hash_entry(root: str, dirpath: str, filename: str) -> tuple[str, str] | None:
|
||||
full = os.path.join(dirpath, filename)
|
||||
try:
|
||||
info = os.lstat(full)
|
||||
except OSError:
|
||||
return None
|
||||
if not os.path.isfile(full) or os.path.islink(full):
|
||||
return None
|
||||
if info.st_size > HARNESS_SNAPSHOT_MAX_FILE_BYTES:
|
||||
return None
|
||||
try:
|
||||
digest = _hash_file(full)
|
||||
except OSError:
|
||||
return None
|
||||
rel = os.path.relpath(full, root).replace(os.sep, "/")
|
||||
return rel, digest
|
||||
|
||||
|
||||
def _walk_entries(root: str) -> Iterator[tuple[str, str]]:
|
||||
for dirpath, dirnames, filenames in os.walk(root, followlinks=False):
|
||||
dirnames[:] = [ # mutable-ok: os.walk prunes only via in-place mutation of its dirnames list
|
||||
d for d in dirnames if d not in HARNESS_SNAPSHOT_SKIP_DIRS
|
||||
]
|
||||
for filename in filenames:
|
||||
entry = _hash_entry(root, dirpath, filename)
|
||||
if entry is not None:
|
||||
yield entry
|
||||
|
||||
|
||||
def snapshot_local_sync(root: str) -> Mapping[str, str]:
|
||||
"""Hash every regular file under root. Symlinks are never followed."""
|
||||
return MappingProxyType(dict(_walk_entries(root)))
|
||||
|
||||
|
||||
async def snapshot_local(root: str) -> Mapping[str, str]:
|
||||
"""Async wrapper around snapshot_local_sync (runs in a worker thread)."""
|
||||
return await asyncio.to_thread(snapshot_local_sync, root)
|
||||
|
||||
|
||||
def _change_kind(path: str, before: Mapping[str, str], after: Mapping[str, str]) -> FileChangeKind | None:
|
||||
if path not in before:
|
||||
return "created"
|
||||
if path not in after:
|
||||
return "deleted"
|
||||
if before[path] != after[path]:
|
||||
return "modified"
|
||||
return None
|
||||
|
||||
|
||||
def diff_snapshots(
|
||||
before: Mapping[str, str], after: Mapping[str, str]
|
||||
) -> list[tuple[str, FileChangeKind]]: # mutable-ok: public sandbox helper; callers compare against a list
|
||||
"""Return (path, kind) for every changed file, sorted by path."""
|
||||
kinds = ((path, _change_kind(path, before, after)) for path in sorted(frozenset(before) | frozenset(after)))
|
||||
return [ # mutable-ok: public sandbox helper returns a list
|
||||
(path, kind) for path, kind in kinds if kind is not None
|
||||
]
|
||||
|
||||
|
||||
def _as_text(data: bytes) -> str | None:
|
||||
if len(data) > HARNESS_MAX_DIFF_BYTES or b"\0" in data:
|
||||
return None
|
||||
try:
|
||||
return data.decode("utf-8")
|
||||
except UnicodeDecodeError:
|
||||
return None
|
||||
|
||||
|
||||
def unified_diff(path: str, old: str | None, new: str | None) -> str:
|
||||
"""Unified diff between two versions of path; None means the file is absent."""
|
||||
from_file = "/dev/null" if old is None else f"a/{path}"
|
||||
to_file = "/dev/null" if new is None else f"b/{path}"
|
||||
lines = difflib.unified_diff(
|
||||
(old or "").splitlines(keepends=True),
|
||||
(new or "").splitlines(keepends=True),
|
||||
fromfile=from_file,
|
||||
tofile=to_file,
|
||||
)
|
||||
return "".join(line if line.endswith("\n") else line + "\n" + _NO_NEWLINE_MARKER for line in lines)
|
||||
|
||||
|
||||
async def _read_or_none(sandbox: Sandbox, path: str) -> bytes | None:
|
||||
try:
|
||||
return await sandbox.read(path)
|
||||
except (HarnessError, OSError):
|
||||
return None
|
||||
|
||||
|
||||
async def capture_text_contents(sandbox: Sandbox, paths_hashes: Mapping[str, str]) -> Mapping[str, bytes]:
|
||||
"""Read small text files before a turn so "modified"/"deleted" diffs can be built.
|
||||
|
||||
Each kept file is <= HARNESS_MAX_DIFF_BYTES; every byte read (kept or not) counts
|
||||
toward HARNESS_SNAPSHOT_MAX_TOTAL_BYTES, after which capture stops.
|
||||
"""
|
||||
captured: dict[str, bytes] = {} # mutable-ok: async accumulator (awaits per read), frozen on return
|
||||
total = 0
|
||||
for path in sorted(paths_hashes):
|
||||
if total >= HARNESS_SNAPSHOT_MAX_TOTAL_BYTES:
|
||||
break
|
||||
data = await _read_or_none(sandbox, path)
|
||||
if data is None:
|
||||
continue
|
||||
total += len(data)
|
||||
if _as_text(data) is not None:
|
||||
captured[path] = data
|
||||
return MappingProxyType(captured)
|
||||
|
||||
|
||||
async def _change_for(
|
||||
sandbox: Sandbox,
|
||||
path: str,
|
||||
kind: FileChangeKind,
|
||||
before_contents: Mapping[str, bytes],
|
||||
) -> FileChange:
|
||||
old_bytes = before_contents.get(path)
|
||||
old = _as_text(old_bytes) if old_bytes is not None else None
|
||||
if kind == "deleted":
|
||||
diff = unified_diff(path, old, None) if old is not None else None
|
||||
return FileChange(path=path, kind=kind, diff=diff)
|
||||
new_bytes = await _read_or_none(sandbox, path)
|
||||
new = _as_text(new_bytes) if new_bytes is not None else None
|
||||
if new is None or (kind == "modified" and old is None):
|
||||
return FileChange(path=path, kind=kind, diff=None)
|
||||
return FileChange(
|
||||
path=path,
|
||||
kind=kind,
|
||||
diff=unified_diff(path, old if kind == "modified" else None, new),
|
||||
)
|
||||
|
||||
|
||||
async def build_file_changes(
|
||||
sandbox: Sandbox,
|
||||
before: Mapping[str, str],
|
||||
after: Mapping[str, str],
|
||||
before_contents: Mapping[str, bytes] | None = None,
|
||||
) -> list[FileChange]: # mutable-ok: feeds the public Result.files list
|
||||
"""FileChange per changed path. diff is None when it cannot be built as text."""
|
||||
contents: Mapping[str, bytes] = before_contents or MappingProxyType({})
|
||||
return [ # mutable-ok: feeds the public Result.files list
|
||||
await _change_for(sandbox, path, kind, contents) for path, kind in diff_snapshots(before, after)
|
||||
]
|
||||
443
litellm/harness/sync.py
Normal file
443
litellm/harness/sync.py
Normal file
|
|
@ -0,0 +1,443 @@
|
|||
"""Sync API for litellm.harness: one daemon event-loop thread runs every async call."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import os
|
||||
import threading
|
||||
from collections.abc import AsyncIterator, Callable, Coroutine, Mapping, Sequence
|
||||
from concurrent.futures import Future
|
||||
from typing import (
|
||||
Any,
|
||||
TypeVar,
|
||||
)
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
from litellm.harness.context import ApprovalHandler
|
||||
from litellm.harness.options import HarnessOptions
|
||||
from litellm.harness.runtime import (
|
||||
AsyncEventStream,
|
||||
AsyncSession,
|
||||
aagent_resume,
|
||||
aagent_session,
|
||||
arun_agent,
|
||||
astream_agent,
|
||||
)
|
||||
from litellm.harness.sandbox.base import Sandbox
|
||||
from litellm.harness.types import (
|
||||
Done,
|
||||
Event,
|
||||
Harness,
|
||||
PermissionMode,
|
||||
Result,
|
||||
State,
|
||||
Usage,
|
||||
)
|
||||
|
||||
T = TypeVar("T")
|
||||
|
||||
IN_LOOP_MESSAGE = (
|
||||
"litellm.{name}() cannot be called from a running event loop; use `await litellm.a{name}(...)` instead"
|
||||
)
|
||||
|
||||
|
||||
class _LoopThread:
|
||||
"""A single background event loop shared by every sync call in the process."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._lock = threading.Lock()
|
||||
self._loop: asyncio.AbstractEventLoop | None = None
|
||||
self._thread: threading.Thread | None = None
|
||||
|
||||
def loop(self) -> asyncio.AbstractEventLoop:
|
||||
with self._lock:
|
||||
if self._loop is None or self._thread is None or not self._thread.is_alive():
|
||||
self._loop = asyncio.new_event_loop()
|
||||
self._thread = threading.Thread(
|
||||
target=self._loop.run_forever,
|
||||
name="litellm-harness-loop",
|
||||
daemon=True,
|
||||
)
|
||||
self._thread.start()
|
||||
return self._loop
|
||||
|
||||
def in_loop_thread(self) -> bool:
|
||||
return self._thread is not None and threading.current_thread() is self._thread
|
||||
|
||||
def submit(self, coro: Coroutine[Any, Any, T]) -> Future[T]:
|
||||
return asyncio.run_coroutine_threadsafe(coro, self.loop())
|
||||
|
||||
|
||||
_LOOP = _LoopThread()
|
||||
|
||||
|
||||
def _ensure_sync_context(name: str) -> None:
|
||||
try:
|
||||
asyncio.get_running_loop()
|
||||
except RuntimeError:
|
||||
return
|
||||
raise RuntimeError(IN_LOOP_MESSAGE.format(name=name))
|
||||
|
||||
|
||||
def run_sync(coro: Coroutine[Any, Any, T], name: str) -> T:
|
||||
"""Run coro on the harness loop thread and block for its result."""
|
||||
try:
|
||||
_ensure_sync_context(name)
|
||||
except RuntimeError:
|
||||
coro.close()
|
||||
raise
|
||||
future = _LOOP.submit(coro)
|
||||
try:
|
||||
return future.result()
|
||||
except KeyboardInterrupt:
|
||||
future.cancel()
|
||||
raise
|
||||
|
||||
|
||||
async def _anext(iterator: AsyncIterator[Event]) -> Event | None:
|
||||
try:
|
||||
return await iterator.__anext__()
|
||||
except StopAsyncIteration:
|
||||
return None
|
||||
|
||||
|
||||
async def _aclose_stream(stream: AsyncEventStream) -> None:
|
||||
await stream.aclose()
|
||||
|
||||
|
||||
class EventStream:
|
||||
"""Sync iterator of events for one turn. `.result` is set once Done is seen."""
|
||||
|
||||
def __init__(self, stream: AsyncEventStream, name: str = "stream") -> None:
|
||||
self._stream = stream
|
||||
self._name = name
|
||||
self._result: Result | None = None
|
||||
self._finished = False
|
||||
|
||||
def __iter__(self) -> EventStream:
|
||||
return self
|
||||
|
||||
def __next__(self) -> Event:
|
||||
if self._finished:
|
||||
raise StopIteration
|
||||
event = run_sync(_anext(self._stream), self._name)
|
||||
if event is None:
|
||||
self._finished = True
|
||||
raise StopIteration
|
||||
if isinstance(event, Done):
|
||||
self._result = event.result
|
||||
return event
|
||||
|
||||
def __enter__(self) -> EventStream:
|
||||
return self
|
||||
|
||||
def __exit__(self, *exc_info: object) -> None:
|
||||
self.close()
|
||||
|
||||
@property
|
||||
def result(self) -> Result | None:
|
||||
return self._result
|
||||
|
||||
def cancel(self) -> None:
|
||||
"""Stop the turn. Iteration still ends with Done(stop_reason='cancelled')."""
|
||||
_LOOP.loop().call_soon_threadsafe(self._stream.cancel)
|
||||
|
||||
def close(self) -> None:
|
||||
"""Abandon the stream and release the session behind it."""
|
||||
if self._finished:
|
||||
return
|
||||
self._finished = True
|
||||
run_sync(_aclose_stream(self._stream), self._name)
|
||||
|
||||
|
||||
class Session:
|
||||
"""Sync multi-turn session. Use as a context manager."""
|
||||
|
||||
def __init__(self, inner: AsyncSession) -> None:
|
||||
self._inner = inner
|
||||
|
||||
@property
|
||||
def aio(self) -> AsyncSession:
|
||||
"""The underlying AsyncSession (runs on the harness loop thread)."""
|
||||
return self._inner
|
||||
|
||||
def start(self) -> Session:
|
||||
run_sync(self._inner.start(), "session")
|
||||
return self
|
||||
|
||||
def __enter__(self) -> Session:
|
||||
return self.start()
|
||||
|
||||
def __exit__(self, *exc_info: object) -> None:
|
||||
self.close()
|
||||
|
||||
def run(self, prompt: str) -> Result:
|
||||
return run_sync(self._inner.arun(prompt), "run")
|
||||
|
||||
def stream(self, prompt: str) -> EventStream:
|
||||
return EventStream(self._inner.astream(prompt))
|
||||
|
||||
def close(self) -> None:
|
||||
run_sync(self._inner.aclose(), "close")
|
||||
|
||||
def detach(self) -> State:
|
||||
return run_sync(self._inner.adetach(), "detach")
|
||||
|
||||
def stop(self) -> State:
|
||||
return run_sync(self._inner.astop(), "stop")
|
||||
|
||||
def history(
|
||||
self,
|
||||
) -> list[dict[str, Any]]: # mutable-ok: public API returns OpenAI-format message dicts from the handler
|
||||
return run_sync(self._inner.history(), "history")
|
||||
|
||||
@property
|
||||
def cost(self) -> float:
|
||||
return self._inner.cost
|
||||
|
||||
@property
|
||||
def usage(self) -> Usage:
|
||||
return self._inner.usage
|
||||
|
||||
@property
|
||||
def results(self) -> list[Result]: # mutable-ok: public property; returns a detached copy of the session's results
|
||||
return list(self._inner.results) # mutable-ok: detached copy so callers cannot mutate the session's accumulator
|
||||
|
||||
@property
|
||||
def session_id(self) -> str:
|
||||
return self._inner.session_id
|
||||
|
||||
|
||||
def _run(
|
||||
harness: Harness,
|
||||
prompt: str,
|
||||
*,
|
||||
sandbox: Sandbox,
|
||||
model: str | None = None,
|
||||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
instructions: str | None = None,
|
||||
tools: Sequence[Callable[..., Any]] = (),
|
||||
skills: Sequence[str | os.PathLike[str]] = (),
|
||||
disable_tools: Sequence[str] = (),
|
||||
permissions: PermissionMode = "full",
|
||||
on_approval: ApprovalHandler | None = None,
|
||||
output: type[BaseModel] | None = None,
|
||||
max_turns: int | None = None,
|
||||
timeout: float | None = None,
|
||||
metadata: Mapping[str, Any] | None = None,
|
||||
options: HarnessOptions | None = None,
|
||||
install: bool = False,
|
||||
) -> Result:
|
||||
"""Run one prompt to completion (blocking) and return the Result."""
|
||||
return run_sync(
|
||||
arun_agent(
|
||||
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,
|
||||
),
|
||||
"agent",
|
||||
)
|
||||
|
||||
|
||||
def _stream(
|
||||
harness: Harness,
|
||||
prompt: str,
|
||||
*,
|
||||
sandbox: Sandbox,
|
||||
model: str | None = None,
|
||||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
instructions: str | None = None,
|
||||
tools: Sequence[Callable[..., Any]] = (),
|
||||
skills: Sequence[str | os.PathLike[str]] = (),
|
||||
disable_tools: Sequence[str] = (),
|
||||
permissions: PermissionMode = "full",
|
||||
on_approval: ApprovalHandler | None = None,
|
||||
output: type[BaseModel] | None = None,
|
||||
max_turns: int | None = None,
|
||||
timeout: float | None = None,
|
||||
metadata: Mapping[str, Any] | None = None,
|
||||
options: HarnessOptions | None = None,
|
||||
install: bool = False,
|
||||
) -> EventStream:
|
||||
"""Stream events for one prompt (sync iterator). Validation errors raise here."""
|
||||
_ensure_sync_context("agent")
|
||||
inner = astream_agent(
|
||||
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,
|
||||
)
|
||||
return EventStream(inner)
|
||||
|
||||
|
||||
def agent_session(
|
||||
harness: Harness,
|
||||
*,
|
||||
sandbox: Sandbox,
|
||||
model: str | None = None,
|
||||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
instructions: str | None = None,
|
||||
tools: Sequence[Callable[..., Any]] = (),
|
||||
skills: Sequence[str | os.PathLike[str]] = (),
|
||||
disable_tools: Sequence[str] = (),
|
||||
permissions: PermissionMode = "full",
|
||||
on_approval: ApprovalHandler | None = None,
|
||||
output: type[BaseModel] | None = None,
|
||||
max_turns: int | None = None,
|
||||
timeout: float | None = None,
|
||||
metadata: Mapping[str, Any] | None = None,
|
||||
options: HarnessOptions | None = None,
|
||||
install: bool = False,
|
||||
) -> Session:
|
||||
"""A multi-turn agent session: `with litellm.agent_session(...) as s: s.run(...)`."""
|
||||
_ensure_sync_context("agent_session")
|
||||
return Session(
|
||||
aagent_session(
|
||||
harness,
|
||||
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,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def agent_resume(
|
||||
state: State | bytes,
|
||||
*,
|
||||
sandbox: Sandbox,
|
||||
model: str | None = None,
|
||||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
instructions: str | None = None,
|
||||
tools: Sequence[Callable[..., Any]] = (),
|
||||
skills: Sequence[str | os.PathLike[str]] = (),
|
||||
disable_tools: Sequence[str] = (),
|
||||
permissions: PermissionMode = "full",
|
||||
on_approval: ApprovalHandler | None = None,
|
||||
output: type[BaseModel] | None = None,
|
||||
max_turns: int | None = None,
|
||||
timeout: float | None = None,
|
||||
metadata: Mapping[str, Any] | None = None,
|
||||
options: HarnessOptions | None = None,
|
||||
install: bool = False,
|
||||
) -> Session:
|
||||
"""Continue a detached or stopped agent session from its State."""
|
||||
_ensure_sync_context("agent_resume")
|
||||
return Session(
|
||||
aagent_resume(
|
||||
state,
|
||||
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,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def agent(
|
||||
harness: Harness,
|
||||
prompt: str,
|
||||
*,
|
||||
sandbox: Sandbox,
|
||||
stream: bool = False,
|
||||
model: str | None = None,
|
||||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
instructions: str | None = None,
|
||||
tools: Sequence[Callable[..., Any]] = (),
|
||||
skills: Sequence[str | os.PathLike[str]] = (),
|
||||
disable_tools: Sequence[str] = (),
|
||||
permissions: PermissionMode = "full",
|
||||
on_approval: ApprovalHandler | None = None,
|
||||
output: type[BaseModel] | None = None,
|
||||
max_turns: int | None = None,
|
||||
timeout: float | None = None,
|
||||
metadata: Mapping[str, Any] | None = None,
|
||||
options: HarnessOptions | None = None,
|
||||
install: bool = False,
|
||||
) -> Result | EventStream:
|
||||
"""Run an agent harness (Claude Code, Codex, OpenCode, Deep Agents) on one prompt.
|
||||
|
||||
Returns a Result. With stream=True it returns an iterator of events instead.
|
||||
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)
|
||||
214
litellm/harness/types.py
Normal file
214
litellm/harness/types.py
Normal file
|
|
@ -0,0 +1,214 @@
|
|||
"""Public types for litellm.harness: the Harness enum, events, results and state."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
from collections.abc import Mapping
|
||||
from dataclasses import dataclass, field
|
||||
from enum import Enum
|
||||
from typing import Any, Literal
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
from litellm.harness.errors import StateIncompatible
|
||||
|
||||
StopReason = Literal["done", "max_turns", "timeout", "cancelled", "runtime_error"]
|
||||
PermissionMode = Literal["read-only", "ask", "edit", "full"]
|
||||
FileChangeKind = Literal["created", "modified", "deleted"]
|
||||
|
||||
STATE_VERSION = 1
|
||||
|
||||
|
||||
class Harness(Enum):
|
||||
"""Supported agent runtimes. A plain Enum on purpose: strings are rejected."""
|
||||
|
||||
CLAUDE_CODE = "claude_code"
|
||||
CODEX = "codex"
|
||||
OPENCODE = "opencode"
|
||||
DEEPAGENTS = "deepagents"
|
||||
|
||||
|
||||
def require_harness(harness: object) -> Harness:
|
||||
"""Return harness if it is a Harness member, else raise TypeError with a hint."""
|
||||
if isinstance(harness, Harness):
|
||||
return harness
|
||||
hint = ""
|
||||
if isinstance(harness, str):
|
||||
normalized = harness.strip().lower().replace("-", "_")
|
||||
for member in Harness:
|
||||
if normalized in (member.value, member.name.lower()):
|
||||
hint = f" Did you mean Harness.{member.name}?"
|
||||
raise TypeError(
|
||||
f"harness must be a litellm.harness.Harness member, got {type(harness).__name__} {harness!r}.{hint}"
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class Usage:
|
||||
input_tokens: int = 0
|
||||
output_tokens: int = 0
|
||||
calls: int = 0
|
||||
|
||||
@property
|
||||
def total_tokens(self) -> int:
|
||||
return self.input_tokens + self.output_tokens
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class Text:
|
||||
delta: str
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class Reasoning:
|
||||
delta: str
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ToolCall:
|
||||
id: str
|
||||
name: str
|
||||
native_name: str
|
||||
input: Mapping[str, Any]
|
||||
builtin: bool = True
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ToolResult:
|
||||
id: str
|
||||
output: str
|
||||
is_error: bool = False
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class FileChange:
|
||||
path: str
|
||||
kind: FileChangeKind
|
||||
diff: str | None = None
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class Compaction:
|
||||
tokens_before: int | None = None
|
||||
tokens_after: int | None = None
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class Approval:
|
||||
"""A request to run a tool. The turn waits until allow() or deny() is called."""
|
||||
|
||||
tool: str
|
||||
input: Mapping[str, Any]
|
||||
_decision: asyncio.Future[tuple[bool, str]] = field(
|
||||
default_factory=lambda: asyncio.get_event_loop().create_future(),
|
||||
compare=False,
|
||||
repr=False,
|
||||
)
|
||||
|
||||
def allow(self) -> None:
|
||||
self._resolve(True, "")
|
||||
|
||||
def deny(self, reason: str = "") -> None:
|
||||
self._resolve(False, reason)
|
||||
|
||||
@property
|
||||
def answered(self) -> bool:
|
||||
return self._decision.done()
|
||||
|
||||
async def wait(self) -> tuple[bool, str]:
|
||||
return await self._decision
|
||||
|
||||
def _resolve(self, allowed: bool, reason: str) -> None:
|
||||
if self._decision.done():
|
||||
return
|
||||
loop = self._decision.get_loop()
|
||||
loop.call_soon_threadsafe(self._set_result, allowed, reason)
|
||||
|
||||
def _set_result(self, allowed: bool, reason: str) -> None:
|
||||
if not self._decision.done():
|
||||
self._decision.set_result((allowed, reason))
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class Result:
|
||||
text: str
|
||||
output: BaseModel | None
|
||||
files: list[FileChange] # mutable-ok: public Result field; users index/iterate it as a list
|
||||
events: list[Event] # mutable-ok: public Result field; users index/iterate it as a list
|
||||
usage: Usage
|
||||
cost: float
|
||||
stop_reason: StopReason
|
||||
session_id: str
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class Done:
|
||||
result: Result
|
||||
|
||||
@property
|
||||
def usage(self) -> Usage:
|
||||
return self.result.usage
|
||||
|
||||
@property
|
||||
def cost(self) -> float:
|
||||
return self.result.cost
|
||||
|
||||
@property
|
||||
def stop_reason(self) -> StopReason:
|
||||
return self.result.stop_reason
|
||||
|
||||
|
||||
Event = Text | Reasoning | ToolCall | ToolResult | FileChange | Compaction | Approval | Done
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class Capabilities:
|
||||
structured_output: bool
|
||||
tool_approval: bool
|
||||
tool_filtering: bool
|
||||
history: bool
|
||||
custom_tools: bool
|
||||
skills: bool
|
||||
resume: bool
|
||||
permission_modes: frozenset[str]
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class State:
|
||||
"""Resume state for a detached or stopped session. Contains no credentials."""
|
||||
|
||||
harness: Harness
|
||||
native_session_id: str | None
|
||||
workdir: str
|
||||
model: str | None = None
|
||||
version: int = STATE_VERSION
|
||||
|
||||
def dumps(self) -> bytes:
|
||||
return json.dumps(
|
||||
{ # mutable-ok: JSON payload serialized immediately by json.dumps
|
||||
"harness": self.harness.value,
|
||||
"native_session_id": self.native_session_id,
|
||||
"workdir": self.workdir,
|
||||
"model": self.model,
|
||||
"version": self.version,
|
||||
}
|
||||
).encode("utf-8")
|
||||
|
||||
@classmethod
|
||||
def loads(cls, data: bytes) -> State:
|
||||
try:
|
||||
raw = json.loads(data.decode("utf-8"))
|
||||
harness = Harness(raw["harness"])
|
||||
version = int(raw["version"])
|
||||
except (ValueError, KeyError, TypeError, UnicodeDecodeError) as e:
|
||||
raise StateIncompatible(f"Unreadable harness state: {e}") from e
|
||||
if version != STATE_VERSION:
|
||||
raise StateIncompatible(f"State version {version} is not supported (expected {STATE_VERSION})")
|
||||
return cls(
|
||||
harness=harness,
|
||||
native_session_id=raw.get("native_session_id"),
|
||||
workdir=raw["workdir"],
|
||||
model=raw.get("model"),
|
||||
version=version,
|
||||
)
|
||||
|
|
@ -1131,7 +1131,7 @@ Model Info:
|
|||
message=message,
|
||||
level=level,
|
||||
alert_type=AlertType.model_deprecation_warnings,
|
||||
alerting_metadata={ # mutable-ok: send_alert takes a dict payload
|
||||
alerting_metadata={
|
||||
"deprecated_count": len(snapshot.deprecated),
|
||||
"imminent_count": len(snapshot.imminent),
|
||||
"upcoming_count": len(snapshot.upcoming),
|
||||
|
|
@ -1245,8 +1245,8 @@ Model Info:
|
|||
try:
|
||||
existing_invitations: Final = TypeAdapter(list[InvitationModel]).validate_python(
|
||||
await InvitationLinkRepository(prisma_client).table.find_many( # pyright: ignore[reportAny] # untyped prisma boundary (any-ok), result validated by TypeAdapter
|
||||
where={"user_id": recipient_user_id}, # mutable-ok: prisma find_many requires a dict where filter
|
||||
order={"created_at": "desc"}, # mutable-ok: prisma find_many requires a dict order arg
|
||||
where={"user_id": recipient_user_id},
|
||||
order={"created_at": "desc"},
|
||||
),
|
||||
from_attributes=True,
|
||||
)
|
||||
|
|
@ -2011,7 +2011,7 @@ Model Info:
|
|||
message="\n\n".join(event.message for event in typed_events),
|
||||
level="High",
|
||||
alert_type=alert_type,
|
||||
alerting_metadata={}, # mutable-ok: send_alert takes a dict payload
|
||||
alerting_metadata={},
|
||||
)
|
||||
for event in typed_events:
|
||||
await self.internal_usage_cache.async_set_cache(
|
||||
|
|
|
|||
|
|
@ -337,7 +337,7 @@ class AzureSentinelLogger(CustomBatchLogger):
|
|||
Raises a NON Blocking verbose_logger.exception if an error occurs
|
||||
"""
|
||||
batch_to_send: Final = tuple(self.log_queue)
|
||||
self.log_queue = [] # mutable-ok: queue ownership is detached before the async send
|
||||
self.log_queue = []
|
||||
try:
|
||||
undelivered: Final = await self._async_send_batch_to_api(
|
||||
log_queue=batch_to_send,
|
||||
|
|
@ -360,7 +360,7 @@ class AzureSentinelLogger(CustomBatchLogger):
|
|||
Sends the batch of audit logs to Azure Monitor Logs Ingestion API
|
||||
"""
|
||||
batch_to_send: Final = tuple(self.audit_log_queue)
|
||||
self.audit_log_queue = [] # mutable-ok: queue ownership is detached before the async send
|
||||
self.audit_log_queue = []
|
||||
try:
|
||||
undelivered: Final = await self._async_send_batch_to_api(
|
||||
log_queue=batch_to_send,
|
||||
|
|
@ -384,7 +384,7 @@ class AzureSentinelLogger(CustomBatchLogger):
|
|||
queue: list[_QueuedPayload],
|
||||
log_type: str,
|
||||
) -> list[_QueuedPayload]:
|
||||
merged: Final = [*undelivered, *queue] # mutable-ok: queue trimming returns a mutable logger queue
|
||||
merged: Final = [*undelivered, *queue]
|
||||
overflow: Final = len(merged) - self.max_queue_size
|
||||
if overflow <= 0:
|
||||
return merged
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue