mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
Merge main into the retry breadcrumb fix
# Conflicts: # .github/workflows/test-unit.yml
This commit is contained in:
commit
9716f2df99
2343 changed files with 105859 additions and 33556 deletions
|
|
@ -17,3 +17,6 @@ rustflags = ["-C", "link-arg=-undefined", "-C", "link-arg=dynamic_lookup"]
|
|||
|
||||
[target.aarch64-apple-darwin]
|
||||
rustflags = ["-C", "link-arg=-undefined", "-C", "link-arg=dynamic_lookup"]
|
||||
|
||||
[env]
|
||||
SQLX_OFFLINE = "true"
|
||||
|
|
|
|||
|
|
@ -141,7 +141,7 @@ commands:
|
|||
node --version
|
||||
npm --version
|
||||
install_rust:
|
||||
description: "Install pinned rustup (1.28.2) and Rust toolchain (1.98.0) with checksum verification. Adds ~/.cargo/bin to PATH. Run this before any `uv sync` or `uv build` of the workspace: the root package builds litellm-rust through maturin, and on an image without cargo maturin fetches an unpinned rustup and a floating toolchain by itself."
|
||||
description: "Install pinned rustup (1.28.2) and Rust toolchain (1.98.0) with checksum verification. Adds ~/.cargo/bin to PATH. Run this before any `uv sync` or `uv build` of the workspace: the root package builds litellm-rust through maturin, and on an image without cargo maturin fetches an unpinned rustup and a floating toolchain by itself. Also restores the dev-profile cargo cache that save_cargo_target writes on main, minus the workspace crates' fingerprints so those always rebuild from the checked-out source."
|
||||
steps:
|
||||
- run:
|
||||
name: Install Rust (rustup 1.28.2, toolchain 1.98.0)
|
||||
|
|
@ -167,9 +167,29 @@ commands:
|
|||
/tmp/rustup-init -y --no-modify-path --profile minimal --default-toolchain 1.98.0
|
||||
rm -f /tmp/rustup-init
|
||||
echo 'export PATH="$HOME/.cargo/bin:$PATH"' >> "$BASH_ENV"
|
||||
echo 'export CARGO_INCREMENTAL=0' >> "$BASH_ENV"
|
||||
export PATH="$HOME/.cargo/bin:$PATH"
|
||||
rustc --version
|
||||
cargo --version
|
||||
{ rustc -vV; cc --version; cat /etc/os-release; } > /tmp/cargo-build-env
|
||||
- restore_cache:
|
||||
keys:
|
||||
- v1-cargo-dev-{{ checksum "/tmp/cargo-build-env" }}-{{ checksum "litellm-rust/Cargo.lock" }}
|
||||
- v1-cargo-dev-{{ checksum "/tmp/cargo-build-env" }}-
|
||||
- run:
|
||||
name: Force a rebuild of the workspace crates restored from the cargo cache
|
||||
command: rm -rf litellm-rust/target/debug/.fingerprint/litellm-*
|
||||
save_cargo_target:
|
||||
steps:
|
||||
- when:
|
||||
condition:
|
||||
equal: [main, << pipeline.git.branch >>]
|
||||
steps:
|
||||
- save_cache:
|
||||
key: v1-cargo-dev-{{ checksum "/tmp/cargo-build-env" }}-{{ checksum "litellm-rust/Cargo.lock" }}
|
||||
paths:
|
||||
- ~/.cargo/registry
|
||||
- ~/project/litellm-rust/target/debug
|
||||
start_postgres:
|
||||
description: "Start a postgres-db container on port 5432 and wait until it accepts connections."
|
||||
parameters:
|
||||
|
|
@ -281,50 +301,11 @@ commands:
|
|||
# `uv sync --package litellm-enterprise` here — that overwrites the
|
||||
# shared .venv and strips out dev/test deps (pytest, prisma, etc.).
|
||||
uv run --no-sync python -c "import litellm_enterprise; print('litellm-enterprise OK:', litellm_enterprise.__file__)"
|
||||
setup_litellm_test_deps:
|
||||
install_windows_toolchain:
|
||||
steps:
|
||||
- checkout
|
||||
- setup_google_dns
|
||||
- install_uv
|
||||
- install_rust
|
||||
- restore_cache:
|
||||
keys:
|
||||
- v3-integration-uv-cache-{{ checksum "uv.lock" }}
|
||||
- run:
|
||||
name: Install Dependencies
|
||||
command: |
|
||||
uv sync --frozen --all-groups --all-extras --python 3.12
|
||||
- setup_litellm_enterprise_pip
|
||||
- save_cache:
|
||||
paths:
|
||||
- ~/.cache/uv
|
||||
key: v3-integration-uv-cache-{{ checksum "uv.lock" }}
|
||||
|
||||
jobs:
|
||||
# Add Windows testing job
|
||||
using_litellm_on_windows:
|
||||
executor:
|
||||
name: win/default
|
||||
shell: powershell.exe
|
||||
working_directory: ~/project
|
||||
environment:
|
||||
UV_PYTHON: "3.11"
|
||||
CARGO_HTTP_MULTIPLEXING: "false"
|
||||
CARGO_NET_RETRY: "5"
|
||||
steps:
|
||||
- checkout
|
||||
- run:
|
||||
name: Install Python
|
||||
command: |
|
||||
choco install python --version=3.11.0 -y --no-progress --force
|
||||
refreshenv
|
||||
python --version
|
||||
environment:
|
||||
CHOCOLATEY_CONFIRM_ALL: "true"
|
||||
- run:
|
||||
name: Install Dependencies
|
||||
environment:
|
||||
UV_HTTP_TIMEOUT: "300"
|
||||
name: Install Rust and uv
|
||||
no_output_timeout: 30m
|
||||
command: |
|
||||
$rustupInit = Join-Path $env:TEMP "rustup-init.exe"
|
||||
$rustupVersion = "1.28.2"
|
||||
|
|
@ -364,6 +345,55 @@ jobs:
|
|||
if (-not (Select-String -Path $PROFILE -SimpleMatch $cargoBin -Quiet)) {
|
||||
Add-Content -Path $PROFILE -Value "`$env:Path = `"$cargoBin;`$env:Path`""
|
||||
}
|
||||
setup_litellm_test_deps:
|
||||
steps:
|
||||
- checkout
|
||||
- setup_google_dns
|
||||
- install_uv
|
||||
- install_rust
|
||||
- restore_cache:
|
||||
keys:
|
||||
- v3-integration-uv-cache-{{ checksum "uv.lock" }}
|
||||
- run:
|
||||
name: Install Dependencies
|
||||
command: |
|
||||
uv sync --frozen --all-groups --all-extras --python 3.12
|
||||
- setup_litellm_enterprise_pip
|
||||
- save_cache:
|
||||
paths:
|
||||
- ~/.cache/uv
|
||||
key: v3-integration-uv-cache-{{ checksum "uv.lock" }}
|
||||
- save_cargo_target
|
||||
|
||||
jobs:
|
||||
# Add Windows testing job
|
||||
using_litellm_on_windows:
|
||||
executor:
|
||||
name: win/default
|
||||
shell: powershell.exe
|
||||
working_directory: ~/project
|
||||
environment:
|
||||
UV_PYTHON: "3.11"
|
||||
CARGO_HTTP_MULTIPLEXING: "false"
|
||||
CARGO_NET_RETRY: "5"
|
||||
steps:
|
||||
- checkout
|
||||
- run:
|
||||
name: Install Python
|
||||
command: |
|
||||
choco install python --version=3.11.0 -y --no-progress --force
|
||||
refreshenv
|
||||
python --version
|
||||
environment:
|
||||
CHOCOLATEY_CONFIRM_ALL: "true"
|
||||
- install_windows_toolchain
|
||||
- run:
|
||||
name: Install Dependencies
|
||||
no_output_timeout: 30m
|
||||
environment:
|
||||
UV_HTTP_TIMEOUT: "300"
|
||||
command: |
|
||||
$env:Path = "$HOME\.cargo\bin;$HOME\.local\bin;$env:Path"
|
||||
for ($attempt = 1; $attempt -le 5; $attempt++) {
|
||||
Write-Host "uv sync attempt $attempt/5"
|
||||
uv sync --frozen --group dev --python 3.11
|
||||
|
|
@ -379,16 +409,68 @@ jobs:
|
|||
name: Run Windows-specific test
|
||||
command: |
|
||||
uv run --no-sync python -m pytest tests/windows_tests/ -v
|
||||
|
||||
windows_release_wheel:
|
||||
executor:
|
||||
name: win/default
|
||||
shell: powershell.exe
|
||||
size: xlarge
|
||||
working_directory: ~/project
|
||||
environment:
|
||||
UV_PYTHON: "3.11"
|
||||
CARGO_HTTP_MULTIPLEXING: "false"
|
||||
CARGO_NET_RETRY: "5"
|
||||
steps:
|
||||
- checkout
|
||||
- run:
|
||||
name: Guard against MAX_PATH-busting packaged wheel paths
|
||||
name: Skip job when no windows-release-relevant files changed
|
||||
shell: bash.exe
|
||||
command: bash .circleci/scripts/path_filter.sh windows-release
|
||||
- run:
|
||||
name: Install Python
|
||||
command: |
|
||||
choco install python --version=3.11.0 -y --no-progress --force
|
||||
refreshenv
|
||||
python --version
|
||||
environment:
|
||||
CHOCOLATEY_CONFIRM_ALL: "true"
|
||||
- install_windows_toolchain
|
||||
- run:
|
||||
name: Record the Rust build environment for the release cargo cache key
|
||||
command: |
|
||||
& "$HOME\.cargo\bin\rustc.exe" -vV | Out-File -Encoding ascii .cargo-build-env
|
||||
- restore_cache:
|
||||
keys:
|
||||
- v1-cargo-release-windows-{{ checksum ".cargo-build-env" }}-{{ checksum "litellm-rust/Cargo.lock" }}
|
||||
- v1-cargo-release-windows-{{ checksum ".cargo-build-env" }}-
|
||||
- run:
|
||||
name: Force a rebuild of the workspace crates restored from the cargo cache
|
||||
command: |
|
||||
$fingerprints = "litellm-rust/target/release/.fingerprint"
|
||||
if (Test-Path $fingerprints) {
|
||||
Get-ChildItem -Path $fingerprints -Filter "litellm-*" | Remove-Item -Recurse -Force
|
||||
}
|
||||
- run:
|
||||
name: Build the release wheel and install it under a worst-case MAX_PATH prefix
|
||||
no_output_timeout: 30m
|
||||
environment:
|
||||
UV_HTTP_TIMEOUT: "300"
|
||||
command: |
|
||||
$env:Path = "$HOME\.cargo\bin;$HOME\.local\bin;$env:Path"
|
||||
cargo --version
|
||||
Get-ChildItem -Path "litellm\rust_bridge" -Filter "_native*" -File -ErrorAction SilentlyContinue | Remove-Item -Force
|
||||
uv build --wheel --out-dir dist
|
||||
uv run --no-sync python tests/windows_tests/check_windows_wheel_install.py
|
||||
if ($LASTEXITCODE -ne 0) {
|
||||
exit $LASTEXITCODE
|
||||
}
|
||||
python tests/windows_tests/check_windows_wheel_install.py
|
||||
- when:
|
||||
condition:
|
||||
equal: [main, << pipeline.git.branch >>]
|
||||
steps:
|
||||
- save_cache:
|
||||
key: v1-cargo-release-windows-{{ checksum ".cargo-build-env" }}-{{ checksum "litellm-rust/Cargo.lock" }}
|
||||
paths:
|
||||
- ~/.cargo/registry
|
||||
- ~/project/litellm-rust/target/release
|
||||
|
||||
base_sdk_install:
|
||||
docker:
|
||||
|
|
@ -416,6 +498,10 @@ jobs:
|
|||
uv venv /tmp/base-sdk --python 3.12
|
||||
VIRTUAL_ENV=/tmp/base-sdk uv pip install dist/*.whl
|
||||
/tmp/base-sdk/bin/python tests/base_sdk_tests/check_base_sdk_install.py
|
||||
- run:
|
||||
name: Guard against MAX_PATH-busting packaged wheel paths
|
||||
command: |
|
||||
python3 tests/windows_tests/check_windows_wheel_install.py --lengths-only
|
||||
|
||||
local_testing_part1:
|
||||
docker:
|
||||
|
|
@ -444,6 +530,7 @@ jobs:
|
|||
paths:
|
||||
- ~/.cache/uv
|
||||
key: v1-uv-cache-{{ checksum "uv.lock" }}
|
||||
- save_cargo_target
|
||||
- run:
|
||||
name: Run prisma ./docker/entrypoint.sh
|
||||
command: |
|
||||
|
|
@ -3118,10 +3205,14 @@ jobs:
|
|||
type: enum
|
||||
enum: [standard, replica]
|
||||
default: standard
|
||||
parallelism:
|
||||
type: integer
|
||||
default: 1
|
||||
machine:
|
||||
image: ubuntu-2204:2024.04.1
|
||||
resource_class: large
|
||||
working_directory: ~/project
|
||||
parallelism: << parameters.parallelism >>
|
||||
steps:
|
||||
- setup_litellm_test_deps
|
||||
- when:
|
||||
|
|
@ -3247,6 +3338,7 @@ jobs:
|
|||
image: ubuntu-2204:2024.04.1
|
||||
resource_class: large
|
||||
working_directory: ~/project
|
||||
parallelism: 4
|
||||
steps:
|
||||
- setup_litellm_test_deps
|
||||
- run:
|
||||
|
|
@ -3256,10 +3348,11 @@ jobs:
|
|||
name: Run unit tests
|
||||
command: |
|
||||
mkdir -p test-results/unit
|
||||
mapfile -t files < <(find tests/unit -name 'test_*.py' | sort)
|
||||
if [ "${#files[@]}" -eq 0 ]; then echo "tests/unit holds no test_*.py files; nothing to run"; exit 0; fi
|
||||
shard="$(find tests/unit -name 'test_*.py' | sort | circleci tests split --split-by=timings --timings-type=filename)"
|
||||
if [ -z "${shard}" ]; then echo "shard ${CIRCLE_NODE_INDEX} received no tests/unit files; nothing to run"; exit 0; fi
|
||||
mapfile -t files < <(printf '%s\n' "${shard}")
|
||||
set +e
|
||||
LITELLM_LOCAL_MODEL_COST_MAP=True uv run --no-sync pytest "${files[@]}" -p no:rerunfailures -p no:pytest-retry --timeout=90 -n 4 --dist=loadscope --tb=short --junitxml=test-results/unit/junit.xml
|
||||
LITELLM_LOCAL_MODEL_COST_MAP=True uv run --no-sync pytest "${files[@]}" -p no:rerunfailures -p no:pytest-retry --timeout=90 -n 4 --dist=loadscope --tb=short -o junit_family=xunit1 --junitxml=test-results/unit/junit.xml
|
||||
status=$?
|
||||
set -e
|
||||
if [ "$status" -eq 5 ]; then echo "pytest collected no tests from tests/unit; passing"; exit 0; fi
|
||||
|
|
@ -3326,23 +3419,17 @@ workflows:
|
|||
name: integration-<< matrix.suite >>
|
||||
matrix:
|
||||
parameters:
|
||||
suite: [management, accounting, database, providers, extensions, mcp, sdk, cost, browser]
|
||||
filters:
|
||||
branches:
|
||||
only:
|
||||
- main
|
||||
- /litellm_.*/
|
||||
suite: [management, accounting, database, providers, mcp, sdk, cost, security, browser]
|
||||
- integration_contracts:
|
||||
name: integration-extensions
|
||||
suite: extensions
|
||||
parallelism: 4
|
||||
- integration_contracts:
|
||||
name: integration-<< matrix.suite >>-replica
|
||||
matrix:
|
||||
parameters:
|
||||
suite: [management, database]
|
||||
mode: [replica]
|
||||
filters:
|
||||
branches:
|
||||
only:
|
||||
- main
|
||||
- /litellm_.*/
|
||||
build_and_test:
|
||||
unless:
|
||||
or:
|
||||
|
|
@ -3350,101 +3437,61 @@ workflows:
|
|||
- not:
|
||||
equal: ["", << pipeline.parameters.routing_parity_base >>]
|
||||
jobs:
|
||||
- using_litellm_on_windows:
|
||||
filters: &main_branches
|
||||
branches:
|
||||
only:
|
||||
- main
|
||||
- /litellm_.*/
|
||||
- unit:
|
||||
filters: *main_branches
|
||||
- using_litellm_on_windows
|
||||
- windows_release_wheel
|
||||
- unit
|
||||
- provider_replay_harness
|
||||
- base_sdk_install:
|
||||
filters: *main_branches
|
||||
- local_testing_part1:
|
||||
filters: *main_branches
|
||||
- local_testing_part2:
|
||||
filters: *main_branches
|
||||
- langfuse_logging_unit_tests:
|
||||
filters: *main_branches
|
||||
- litellm_assistants_api_testing:
|
||||
filters: *main_branches
|
||||
- litellm_router_testing:
|
||||
filters: *main_branches
|
||||
- litellm_router_unit_testing:
|
||||
filters: *main_branches
|
||||
- auth_ui_unit_tests:
|
||||
filters: *main_branches
|
||||
- build_docker_database_image:
|
||||
filters: *main_branches
|
||||
- e2e_ui_testing:
|
||||
filters: *main_branches
|
||||
- e2e_ui_testing_server_root_path:
|
||||
filters: *main_branches
|
||||
- base_sdk_install
|
||||
- local_testing_part1
|
||||
- local_testing_part2
|
||||
- langfuse_logging_unit_tests
|
||||
- litellm_assistants_api_testing
|
||||
- litellm_router_testing
|
||||
- litellm_router_unit_testing
|
||||
- auth_ui_unit_tests
|
||||
- build_docker_database_image
|
||||
- e2e_ui_testing
|
||||
- e2e_ui_testing_server_root_path
|
||||
- build_and_test:
|
||||
requires:
|
||||
- build_docker_database_image
|
||||
filters: *main_branches
|
||||
- e2e_openai_endpoints:
|
||||
requires:
|
||||
- build_docker_database_image
|
||||
filters: *main_branches
|
||||
- proxy_logging_guardrails_model_info_tests:
|
||||
requires:
|
||||
- build_docker_database_image
|
||||
filters: *main_branches
|
||||
- proxy_spend_accuracy_tests:
|
||||
requires:
|
||||
- build_docker_database_image
|
||||
filters: *main_branches
|
||||
- proxy_multi_instance_tests:
|
||||
requires:
|
||||
- build_docker_database_image
|
||||
filters: *main_branches
|
||||
- proxy_store_model_in_db_tests:
|
||||
requires:
|
||||
- build_docker_database_image
|
||||
filters: *main_branches
|
||||
- proxy_build_from_pip_tests:
|
||||
filters: *main_branches
|
||||
- proxy_build_from_pip_tests
|
||||
- proxy_pass_through_endpoint_tests:
|
||||
requires:
|
||||
- build_docker_database_image
|
||||
filters: *main_branches
|
||||
- proxy_e2e_anthropic_messages_tests:
|
||||
requires:
|
||||
- build_docker_database_image
|
||||
filters: *main_branches
|
||||
- llm_translation_testing:
|
||||
filters: *main_branches
|
||||
- realtime_translation_testing:
|
||||
filters: *main_branches
|
||||
- agent_testing:
|
||||
filters: *main_branches
|
||||
- guardrails_testing:
|
||||
filters: *main_branches
|
||||
- google_generate_content_endpoint_testing:
|
||||
filters: *main_branches
|
||||
- llm_responses_api_testing:
|
||||
filters: *main_branches
|
||||
- ocr_testing:
|
||||
filters: *main_branches
|
||||
- search_testing:
|
||||
filters: *main_branches
|
||||
- batches_testing:
|
||||
filters: *main_branches
|
||||
- litellm_utils_testing:
|
||||
filters: *main_branches
|
||||
- pass_through_unit_testing:
|
||||
filters: *main_branches
|
||||
- image_gen_testing:
|
||||
filters: *main_branches
|
||||
- logging_testing:
|
||||
filters: *main_branches
|
||||
- audio_testing:
|
||||
filters: *main_branches
|
||||
- redis_caching_unit_tests:
|
||||
filters: *main_branches
|
||||
- llm_translation_testing
|
||||
- realtime_translation_testing
|
||||
- agent_testing
|
||||
- guardrails_testing
|
||||
- google_generate_content_endpoint_testing
|
||||
- llm_responses_api_testing
|
||||
- ocr_testing
|
||||
- search_testing
|
||||
- batches_testing
|
||||
- litellm_utils_testing
|
||||
- pass_through_unit_testing
|
||||
- image_gen_testing
|
||||
- logging_testing
|
||||
- audio_testing
|
||||
- redis_caching_unit_tests
|
||||
- upload-coverage:
|
||||
requires:
|
||||
- realtime_translation_testing
|
||||
|
|
@ -3469,18 +3516,12 @@ workflows:
|
|||
- db_migration_disable_update_check:
|
||||
requires:
|
||||
- build_docker_database_image
|
||||
filters: *main_branches
|
||||
- installing_litellm_on_python:
|
||||
filters: *main_branches
|
||||
- installing_litellm_on_python_3_13:
|
||||
filters: *main_branches
|
||||
- installing_litellm_on_python_v2_migration_resolver:
|
||||
filters: *main_branches
|
||||
- installing_litellm_on_python
|
||||
- installing_litellm_on_python_3_13
|
||||
- installing_litellm_on_python_v2_migration_resolver
|
||||
- helm_chart_testing:
|
||||
requires:
|
||||
- build_docker_database_image
|
||||
filters: *main_branches
|
||||
- test_bad_database_url:
|
||||
requires:
|
||||
- build_docker_database_image
|
||||
filters: *main_branches
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
#!/usr/bin/env bash
|
||||
set -uo pipefail
|
||||
|
||||
category="${1:?usage: classify_changes.sh <backend|client|ui|provider-harness|cost-map-only|mcp-dependencies>}"
|
||||
category="${1:?usage: classify_changes.sh <backend|client|ui|provider-harness|cost-map-only|mcp-dependencies|windows-release>}"
|
||||
|
||||
has_client=false
|
||||
has_backend=false
|
||||
|
|
@ -9,19 +9,24 @@ has_ci=false
|
|||
has_provider_harness=false
|
||||
has_cost_map=false
|
||||
has_mcp_dependencies=false
|
||||
has_windows_release=false
|
||||
outside_cost_map_set=false
|
||||
while IFS= read -r file || [ -n "$file" ]; do
|
||||
[ -n "$file" ] || continue
|
||||
case "$file" in
|
||||
*.md | *.mdx) : ;;
|
||||
pyproject.toml | */pyproject.toml | uv.lock | uv.toml | .python-version | rust-toolchain.toml | litellm-rust/* | litellm/__init__.py | litellm/proxy/proxy_server.py | litellm/*mcp* | tests/*mcp* | litellm/integrations/arize/* | tests/base_sdk_tests/* | scripts/check_mcp_sdk_install.py | .github/workflows/test-mcp-dependency-resolution.yml | .github/actions/detect-changes/* | .github/actions/setup-uv-with-retries/* | .github/actions/cache-cargo-build/* | .github/scripts/detect_changes.sh | .github/scripts/uv_sync_with_retries.sh | .circleci/scripts/classify_changes.sh | tests/test_litellm/test_circleci_path_filter.py | tests/test_litellm/test_detect_changes.py)
|
||||
pyproject.toml | */pyproject.toml | uv.lock | uv.toml | .python-version | rust-toolchain.toml | litellm-rust/* | litellm/__init__.py | litellm/proxy/proxy_server.py | litellm/*mcp* | tests/*mcp* | litellm/integrations/arize/* | tests/base_sdk_tests/* | scripts/check_mcp_sdk_install.py | .github/workflows/test-mcp-dependency-resolution.yml | .github/actions/detect-changes/* | .github/actions/setup-uv-with-retries/* | .github/actions/cache-cargo-build/* | .github/scripts/detect_changes.sh | .github/scripts/uv_sync_with_retries.sh | .circleci/scripts/classify_changes.sh | tests/unit/test_circleci_path_filter.py | tests/unit/test_detect_changes.py)
|
||||
has_mcp_dependencies=true ;;
|
||||
esac
|
||||
case "$file" in
|
||||
tests/e2e/*/*.py) : ;;
|
||||
tests/e2e/*.py | tests/code_coverage_tests/test_provider_cache.py | tests/code_coverage_tests/test_provider_replay_harness.py | tests/test_litellm/test_circleci_path_filter.py | .circleci/* | pyproject.toml | uv.lock)
|
||||
tests/e2e/*.py | tests/code_coverage_tests/test_provider_cache.py | tests/code_coverage_tests/test_provider_replay_harness.py | tests/unit/test_circleci_path_filter.py | .circleci/* | pyproject.toml | uv.lock)
|
||||
has_provider_harness=true ;;
|
||||
esac
|
||||
case "$file" in
|
||||
litellm-rust/* | litellm/rust_bridge/* | rust-toolchain.toml | pyproject.toml | uv.lock | tests/windows_tests/* | .circleci/*)
|
||||
has_windows_release=true ;;
|
||||
esac
|
||||
case "$file" in
|
||||
ui/* | tests/e2e/ui/*) has_client=true ;;
|
||||
docs/* | *.md | *.mdx) : ;;
|
||||
|
|
@ -46,6 +51,9 @@ case "$category" in
|
|||
provider-harness)
|
||||
[ "$has_provider_harness" = true ] && echo run || echo skip
|
||||
;;
|
||||
windows-release)
|
||||
[ "$has_windows_release" = true ] && echo run || echo skip
|
||||
;;
|
||||
backend)
|
||||
[ "$has_backend" = true ] && echo run || echo skip
|
||||
;;
|
||||
|
|
|
|||
|
|
@ -26,6 +26,7 @@ guard_created=false
|
|||
guard_installed=false
|
||||
guard6_created=false
|
||||
guard6_installed=false
|
||||
egress_cgroup=litellm-integration
|
||||
cleanup() {
|
||||
original_status=$?
|
||||
trap - EXIT INT TERM
|
||||
|
|
@ -47,14 +48,14 @@ cleanup() {
|
|||
fi
|
||||
done
|
||||
if [ "$guard_installed" = true ]; then
|
||||
sudo iptables -D OUTPUT -m owner --uid-owner "$(id -u)" -j integration_only || original_status=1
|
||||
sudo iptables -D OUTPUT -m cgroup --path "$egress_cgroup" -j integration_only || original_status=1
|
||||
fi
|
||||
if [ "$guard_created" = true ]; then
|
||||
sudo iptables -F integration_only || original_status=1
|
||||
sudo iptables -X integration_only || original_status=1
|
||||
fi
|
||||
if [ "$guard6_installed" = true ]; then
|
||||
sudo ip6tables -D OUTPUT -m owner --uid-owner "$(id -u)" -j integration_only || original_status=1
|
||||
sudo ip6tables -D OUTPUT -m cgroup --path "$egress_cgroup" -j integration_only || original_status=1
|
||||
fi
|
||||
if [ "$guard6_created" = true ]; then
|
||||
sudo ip6tables -F integration_only || original_status=1
|
||||
|
|
@ -100,6 +101,8 @@ if [ "$mode" = parity ]; then
|
|||
export INTEGRATION_ROUTING=capture
|
||||
fi
|
||||
|
||||
sudo mkdir -p "/sys/fs/cgroup/$egress_cgroup"
|
||||
echo "$$" | sudo tee "/sys/fs/cgroup/$egress_cgroup/cgroup.procs" > /dev/null
|
||||
sudo iptables -N integration_only
|
||||
guard_created=true
|
||||
sudo iptables -A integration_only -o lo -j ACCEPT
|
||||
|
|
@ -109,13 +112,13 @@ for service in postgres-db redis-cache; do
|
|||
sudo iptables -A integration_only -d "$address" -j ACCEPT
|
||||
done
|
||||
sudo iptables -A integration_only -j REJECT
|
||||
sudo iptables -I OUTPUT 1 -m owner --uid-owner "$(id -u)" -j integration_only
|
||||
sudo iptables -I OUTPUT 1 -m cgroup --path "$egress_cgroup" -j integration_only
|
||||
guard_installed=true
|
||||
sudo ip6tables -N integration_only
|
||||
guard6_created=true
|
||||
sudo ip6tables -A integration_only -o lo -j ACCEPT
|
||||
sudo ip6tables -A integration_only -j REJECT
|
||||
sudo ip6tables -I OUTPUT 1 -m owner --uid-owner "$(id -u)" -j integration_only
|
||||
sudo ip6tables -I OUTPUT 1 -m cgroup --path "$egress_cgroup" -j integration_only
|
||||
guard6_installed=true
|
||||
|
||||
if curl --noproxy '*' --connect-timeout 2 -s http://198.51.100.1 >/dev/null 2>&1; then
|
||||
|
|
@ -209,6 +212,15 @@ if [ "$suite" = browser ]; then
|
|||
exit 0
|
||||
fi
|
||||
|
||||
node_files=()
|
||||
if [ "${CIRCLE_NODE_TOTAL:-1}" -gt 1 ]; then
|
||||
split="$(.venv/bin/python tests/integration/run.py "$suite" --list \
|
||||
| circleci tests split --split-by=timings --timings-type=filename)"
|
||||
read -r -a node_files <<< "$(printf '%s' "$split" | tr '\n' ' ')"
|
||||
test "${#node_files[@]}" -gt 0
|
||||
printf '%s\n' "${node_files[@]}" > "$results/node-files.txt"
|
||||
fi
|
||||
|
||||
env -i PATH="$PATH" HOME="$HOME" PYTHONPATH="$PYTHONPATH" \
|
||||
INTEGRATION_RUN_ID="$integration_identity" \
|
||||
DATABASE_URL="$DATABASE_URL" REDIS_HOST="$REDIS_HOST" REDIS_PORT="$REDIS_PORT" \
|
||||
|
|
@ -222,7 +234,7 @@ env -i PATH="$PATH" HOME="$HOME" PYTHONPATH="$PYTHONPATH" \
|
|||
INTEGRATION_PROXY_DATABASE_URL="$INTEGRATION_PROXY_DATABASE_URL" \
|
||||
INTEGRATION_PROXY_READ_REPLICA_URL="$INTEGRATION_PROXY_READ_REPLICA_URL" \
|
||||
INTEGRATION_ROUTING="$INTEGRATION_ROUTING" \
|
||||
.venv/bin/python tests/integration/run.py "$suite" --results "$results"
|
||||
.venv/bin/python tests/integration/run.py "$suite" --results "$results" "${node_files[@]}"
|
||||
|
||||
if [ "${INTEGRATION_COVERAGE:-0}" = 1 ]; then
|
||||
for covered_pid in "$proxy_pid" "$peer_pid"; do
|
||||
|
|
|
|||
|
|
@ -5,9 +5,14 @@ flag="${1:?usage: unit_selection.sh <codecov flag>}"
|
|||
|
||||
legacy_flags=(
|
||||
caching-local
|
||||
core-utils
|
||||
enterprise-package
|
||||
enterprise-routing
|
||||
integrations
|
||||
llm-other-providers
|
||||
llm-vertex-ai
|
||||
mcp-integration
|
||||
misc
|
||||
proxy-db-auth-checks
|
||||
proxy-db-budgets
|
||||
proxy-db-custom-logging
|
||||
|
|
@ -22,11 +27,13 @@ legacy_flags=(
|
|||
proxy-db-proxy-utils
|
||||
proxy-extras
|
||||
proxy-infra
|
||||
responses-caching-types
|
||||
)
|
||||
|
||||
legacy_paths() {
|
||||
case "$1" in
|
||||
caching-local) echo tests/unit/caching ;;
|
||||
core-utils) echo tests/unit/litellm_core_utils ;;
|
||||
enterprise-package)
|
||||
echo tests/unit/enterprise/integrations
|
||||
echo tests/unit/enterprise/proxy/auth
|
||||
|
|
@ -36,6 +43,10 @@ legacy_paths() {
|
|||
echo tests/unit/enterprise/proxy/test_audit_logging_endpoints.py
|
||||
echo tests/unit/enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py ;;
|
||||
enterprise-routing)
|
||||
echo tests/unit/google_genai
|
||||
echo tests/unit/router_strategy
|
||||
echo tests/unit/router_utils
|
||||
echo tests/unit/proxy/common_utils/test_cache_aware_routing.py
|
||||
echo tests/unit/enterprise/enterprise_callbacks/send_emails
|
||||
echo tests/unit/enterprise/proxy/test_afile_retrieve_returns_unified_id.py
|
||||
echo tests/unit/enterprise/proxy/test_batch_retrieve_input_file_id.py
|
||||
|
|
@ -47,10 +58,34 @@ legacy_paths() {
|
|||
echo tests/unit/enterprise/proxy/test_file_deletion_blocking.py
|
||||
echo tests/unit/enterprise/proxy/test_managed_files_access_check.py
|
||||
echo tests/unit/enterprise/proxy/test_managed_files_hook.py ;;
|
||||
integrations) echo tests/unit/integrations ;;
|
||||
llm-other-providers) find tests/unit/llms -name 'test_*.py' -not -path 'tests/unit/llms/vertex_ai/*' ;;
|
||||
llm-vertex-ai) echo tests/unit/llms/vertex_ai ;;
|
||||
mcp-integration)
|
||||
echo tests/unit/experimental_mcp_client
|
||||
echo tests/unit/proxy/_experimental/mcp_server
|
||||
echo tests/unit/responses/mcp
|
||||
echo tests/mcp_tests/test_proxy_mcp_e2e.py ;;
|
||||
misc)
|
||||
find tests/unit -maxdepth 1 -name 'test_*.py'
|
||||
echo tests/unit/test_router
|
||||
echo tests/unit/a2a_protocol
|
||||
echo tests/unit/batches
|
||||
echo tests/unit/chat_completions
|
||||
echo tests/unit/completion_extras
|
||||
echo tests/unit/containers
|
||||
echo tests/unit/embeddings
|
||||
echo tests/unit/endpoints
|
||||
echo tests/unit/files
|
||||
echo tests/unit/images
|
||||
echo tests/unit/interactions
|
||||
echo tests/unit/messages
|
||||
echo tests/unit/rag
|
||||
echo tests/unit/rerank_api
|
||||
echo tests/unit/rust_bridge
|
||||
echo tests/unit/secret_managers
|
||||
echo tests/unit/vector_stores
|
||||
echo tests/unit/videos ;;
|
||||
proxy-db-auth-checks)
|
||||
echo tests/unit/proxy/auth/test_auth_checks.py
|
||||
echo tests/unit/proxy/auth/test_user_api_key_auth.py
|
||||
|
|
@ -113,6 +148,9 @@ legacy_paths() {
|
|||
proxy-db-proxy-utils) echo tests/unit/proxy/test_proxy_utils.py ;;
|
||||
proxy-extras) echo tests/unit/litellm_proxy_extras ;;
|
||||
proxy-infra) echo tests/unit/gateway ;;
|
||||
responses-caching-types)
|
||||
find tests/unit/responses -name 'test_*.py' -not -path 'tests/unit/responses/mcp/*'
|
||||
echo tests/unit/types ;;
|
||||
*) echo "unit_selection.sh: unknown flag $1" >&2; exit 1 ;;
|
||||
esac
|
||||
}
|
||||
|
|
|
|||
|
|
@ -341,6 +341,7 @@ workflows:
|
|||
flag:
|
||||
- enterprise-package
|
||||
- proxy-infra
|
||||
- responses-caching-types
|
||||
- proxy-db-auth-checks
|
||||
- proxy-db-jwt-and-keys
|
||||
- proxy-db-proxy-server-core
|
||||
|
|
@ -353,6 +354,42 @@ workflows:
|
|||
- proxy-db-endpoints-and-responses
|
||||
base_ref: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.base.ref or "" >>
|
||||
pull_request_url: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.url or "" >>
|
||||
- unit:
|
||||
name: unit-llm-vertex-ai
|
||||
flag: llm-vertex-ai
|
||||
shards: 2
|
||||
workers: 1
|
||||
reruns: 2
|
||||
base_ref: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.base.ref or "" >>
|
||||
pull_request_url: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.url or "" >>
|
||||
- unit:
|
||||
name: unit-llm-other-providers
|
||||
flag: llm-other-providers
|
||||
shards: 3
|
||||
reruns: 2
|
||||
base_ref: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.base.ref or "" >>
|
||||
pull_request_url: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.url or "" >>
|
||||
- unit:
|
||||
name: unit-core-utils
|
||||
flag: core-utils
|
||||
shards: 2
|
||||
reruns: 1
|
||||
base_ref: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.base.ref or "" >>
|
||||
pull_request_url: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.url or "" >>
|
||||
- unit:
|
||||
name: unit-integrations
|
||||
flag: integrations
|
||||
shards: 2
|
||||
reruns: 3
|
||||
base_ref: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.base.ref or "" >>
|
||||
pull_request_url: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.url or "" >>
|
||||
- unit:
|
||||
name: unit-misc
|
||||
flag: misc
|
||||
shards: 2
|
||||
reruns: 2
|
||||
base_ref: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.base.ref or "" >>
|
||||
pull_request_url: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.url or "" >>
|
||||
- unit:
|
||||
name: unit-proxy-db-proxy-utils
|
||||
flag: proxy-db-proxy-utils
|
||||
|
|
|
|||
|
|
@ -42,7 +42,7 @@ case "$subject" in
|
|||
;;
|
||||
esac
|
||||
|
||||
ALLOWED_TYPES="feat|fix|docs|style|refactor|perf|test|build|ci|chore|revert"
|
||||
ALLOWED_TYPES="feat|fix|docs|style|refactor|perf|test|build|ci|chore|revert|security"
|
||||
# Description must not start with an uppercase letter — kept in sync with the
|
||||
# subjectPattern in .github/workflows/conventional-commits.yml so the local
|
||||
# hook is the strictly tighter of the two gates. (Without this guard, a commit
|
||||
|
|
@ -61,7 +61,7 @@ cat >&2 <<EOF
|
|||
Expected: <type>(<scope>)!: <description>
|
||||
(description must start with a lowercase letter)
|
||||
|
||||
Allowed types: feat, fix, docs, style, refactor, perf, test, build, ci, chore, revert
|
||||
Allowed types: feat, fix, docs, style, refactor, perf, test, build, ci, chore, revert, security
|
||||
Examples:
|
||||
feat(router): add weighted round-robin strategy
|
||||
fix(bedrock): decouple STS region from aws_region_name
|
||||
|
|
|
|||
8
.github/CODEOWNERS
vendored
8
.github/CODEOWNERS
vendored
|
|
@ -1,10 +1,2 @@
|
|||
/ui/ @yuneng-berri @ryan-crabbe-berri
|
||||
/litellm/proxy/_experimental/out/ @yuneng-berri @ryan-crabbe-berri
|
||||
/ui/Dockerfile
|
||||
/ui/nginx.conf
|
||||
/ui/litellm-dashboard/src/lib/http/schema.d.ts
|
||||
/ui/litellm-dashboard/tsconfig.tsbuildinfo
|
||||
/model_prices_and_context_window.json @mateo-berri @ryan-crabbe-berri @kerry-berri
|
||||
/litellm/model_prices_and_context_window_backup.json @mateo-berri @ryan-crabbe-berri @kerry-berri
|
||||
/litellm-proxy-extras/litellm_proxy_extras/migrations/ @yuneng-berri @ryan-crabbe-berri
|
||||
/.github/CODEOWNERS @yuneng-berri
|
||||
|
|
|
|||
7
.github/ci-coverage-allowlist.yml
vendored
7
.github/ci-coverage-allowlist.yml
vendored
|
|
@ -111,3 +111,10 @@ dockerfiles:
|
|||
An example image under cookbook/ that is documentation rather than a shipped artifact
|
||||
paths:
|
||||
- cookbook/litellm-ollama-docker-image/Dockerfile
|
||||
- reason: >-
|
||||
The Rust gateway image compiles the whole workspace in release mode, which is too slow for
|
||||
a per-pull-request job while the gateway binary is still being assembled; the Rust lint,
|
||||
clippy, and compile jobs already cover the code it packages. Revisit when the gateway is
|
||||
published
|
||||
paths:
|
||||
- litellm-rust/crates/gateway/Dockerfile
|
||||
|
|
|
|||
18
.github/merge-smoke-tests.json
vendored
18
.github/merge-smoke-tests.json
vendored
|
|
@ -1,15 +1,15 @@
|
|||
{
|
||||
"cases": {
|
||||
"CHAT-JSON": "tests/test_litellm/llms/openai/test_openai.py::test_acompletion_returns_json_reply_over_injected_transport",
|
||||
"CHAT-TEXT-STREAM": "tests/test_litellm/llms/openai/test_openai.py::test_acompletion_streams_text_deltas_over_injected_transport",
|
||||
"CHAT-TOOL-STREAM": "tests/test_litellm/llms/openai/test_openai.py::test_acompletion_streams_tool_call_arguments_over_injected_transport",
|
||||
"CHAT-JSON": "tests/unit/llms/openai/test_openai.py::test_acompletion_returns_json_reply_over_injected_transport",
|
||||
"CHAT-TEXT-STREAM": "tests/unit/llms/openai/test_openai.py::test_acompletion_streams_text_deltas_over_injected_transport",
|
||||
"CHAT-TOOL-STREAM": "tests/unit/llms/openai/test_openai.py::test_acompletion_streams_tool_call_arguments_over_injected_transport",
|
||||
"MODEL-ALLOW": "tests/test_litellm/proxy/auth/test_auth_checks.py::test_can_object_call_model_allows_listed_model_for_key",
|
||||
"MODEL-DENY": "tests/test_litellm/proxy/auth/test_auth_checks.py::test_can_object_call_model_denials_return_forbidden[key-key_model_access_denied]",
|
||||
"COST-EXPLICIT": "tests/test_litellm/test_cost_calculator.py::test_completion_cost_charges_explicit_per_token_rates_over_registered_ones",
|
||||
"COST-ZERO": "tests/test_litellm/test_cost_calculator.py::test_completion_cost_is_zero_when_explicit_rates_are_zero",
|
||||
"LOG-CONTENT-ON": "tests/test_litellm/litellm_core_utils/test_litellm_logging.py::test_standard_logging_payload_keeps_message_content_when_message_logging_is_on",
|
||||
"LOG-CONTENT-OFF": "tests/test_litellm/litellm_core_utils/test_litellm_logging.py::test_standard_logging_payload_redacts_message_content_when_message_logging_is_off",
|
||||
"CALLBACK-SUCCESS": "tests/test_litellm/litellm_core_utils/test_litellm_logging.py::test_async_success_handler_delivers_standard_logging_payload_to_custom_logger",
|
||||
"CALLBACK-FAILURE": "tests/test_litellm/litellm_core_utils/test_litellm_logging.py::test_async_failure_handler_delivers_failure_payload_to_custom_logger"
|
||||
"COST-EXPLICIT": "tests/unit/test_cost_calculator.py::test_completion_cost_charges_explicit_per_token_rates_over_registered_ones",
|
||||
"COST-ZERO": "tests/unit/test_cost_calculator.py::test_completion_cost_is_zero_when_explicit_rates_are_zero",
|
||||
"LOG-CONTENT-ON": "tests/unit/litellm_core_utils/test_litellm_logging.py::test_standard_logging_payload_keeps_message_content_when_message_logging_is_on",
|
||||
"LOG-CONTENT-OFF": "tests/unit/litellm_core_utils/test_litellm_logging.py::test_standard_logging_payload_redacts_message_content_when_message_logging_is_off",
|
||||
"CALLBACK-SUCCESS": "tests/unit/litellm_core_utils/test_litellm_logging.py::test_async_success_handler_delivers_standard_logging_payload_to_custom_logger",
|
||||
"CALLBACK-FAILURE": "tests/unit/litellm_core_utils/test_litellm_logging.py::test_async_failure_handler_delivers_failure_payload_to_custom_logger"
|
||||
}
|
||||
}
|
||||
|
|
|
|||
4
.github/pull_request_template.md
vendored
4
.github/pull_request_template.md
vendored
|
|
@ -54,7 +54,7 @@ After: the same request comes back with real token counts, so the dashboard show
|
|||
|
||||
## Affected release
|
||||
|
||||
<!-- Only for a fix to a regression in a released or rc version (perf, memory, crash, or behavior): name the version it regressed in, e.g. "regression in v1.100.0" or "since v1.101.0-rc.1", and add the `backport-stable` label so the fix is cherry-picked onto the rc line before the stable is tagged. Drop the section otherwise -->
|
||||
<!-- Only for a fix to a regression in a released or rc version (perf, memory, crash, or behavior): name the version it regressed in, e.g. "regression in v1.100.0" or "since v1.101.0-rc.1". Add the `backport-stable` label only when the regression is a P0, meaning its Linear ticket is Urgent (a security hole however narrow, data loss, or a crash or outage for every user on that version), because every labeled PR must be cherry-picked onto the baking rc line before the stable can be tagged; every other regression fix ships in the next rc unlabeled. Drop the section otherwise -->
|
||||
|
||||
## Linear ticket
|
||||
|
||||
|
|
@ -65,7 +65,7 @@ After: the same request comes back with real token counts, so the dashboard show
|
|||
**Please complete all items before asking a LiteLLM maintainer to review your PR**
|
||||
|
||||
- [ ] I have added meaningful tests
|
||||
- [ ] The handful of test files covering my change pass locally, e.g. `uv run pytest tests/test_litellm/<your_test_file>.py -v`. Leave the suites (`make test-unit-*`, `make test-unit`) to CI: it finishes in ~15 minutes where a laptop takes an hour or more
|
||||
- [ ] The handful of test files covering my change pass locally, e.g. `uv run pytest tests/unit/<your_test_file>.py -v`. Leave the suites (`make test-unit-*`, `make test-unit`) to CI: it finishes in ~15 minutes where a laptop takes an hour or more
|
||||
- [ ] My PR passes all required CI/CD checks (e.g., lint, schema.d.ts sync check, etc.)
|
||||
- [ ] My PR's scope is as isolated as possible; it only solves 1 specific problem
|
||||
- [ ] I have received a Greptile **Confidence Score of at least 4/5** before requesting a maintainer review (Greptile reviews automatically once the PR is opened; only comment `@greptileai` to re-request a review after pushing changes)
|
||||
|
|
|
|||
2
.github/scripts/assert_ci_coverage.py
vendored
2
.github/scripts/assert_ci_coverage.py
vendored
|
|
@ -516,7 +516,7 @@ def _integration_ownership(repo_root: pathlib.Path = REPO_ROOT) -> tuple[frozens
|
|||
str(path.relative_to(repo_root))
|
||||
for folders in groups.values()
|
||||
for folder in folders
|
||||
for path in (integration_root / folder).glob("test_*.py")
|
||||
for path in (integration_root / folder).rglob("test_*.py")
|
||||
)
|
||||
browser_manifest: Final = repo_root / "tests/e2e/ui/tests/integrationCritical/expected.json"
|
||||
browser_nodes: Final = json.loads(browser_manifest.read_text()) if browser_manifest.exists() else ()
|
||||
|
|
|
|||
2
.github/scripts/verify_linux_native_wheel.py
vendored
2
.github/scripts/verify_linux_native_wheel.py
vendored
|
|
@ -214,7 +214,7 @@ def main(
|
|||
native_module: Final = load_native_module(native_path)
|
||||
native_module_loads: Final = native_module is not None
|
||||
panic_test_hook_absent: Final = native_module is not None and not hasattr(native_module, "_panic_for_test")
|
||||
native_size_limit: Final = 40_000_000
|
||||
native_size_limit: Final = 45_000_000
|
||||
native_size_within_limit: Final = native_member.file_size <= native_size_limit
|
||||
validations: Final = (
|
||||
(f"Python tag is {EXPECTED_PYTHON_TAG}", python_tag == EXPECTED_PYTHON_TAG),
|
||||
|
|
|
|||
18
.github/workflows/_test-unit-base.yml
vendored
18
.github/workflows/_test-unit-base.yml
vendored
|
|
@ -13,12 +13,13 @@ on:
|
|||
have its path existence-checked like any other token.
|
||||
required: true
|
||||
type: string
|
||||
fork-flag:
|
||||
unit-flag:
|
||||
description: >-
|
||||
Codecov flag of the `.circleci/tests.yml` job that now owns part of
|
||||
this shard. CircleCI does not run on pull requests from forks, so on
|
||||
those events this shard also runs the files
|
||||
`.circleci/scripts/unit_selection.sh` lists for the flag.
|
||||
this shard. The shard also runs the files
|
||||
`.circleci/scripts/unit_selection.sh` lists for the flag, on every
|
||||
event, because the CircleCI pipeline is manual-only while the tests
|
||||
migrate.
|
||||
required: false
|
||||
type: string
|
||||
default: ""
|
||||
|
|
@ -175,8 +176,7 @@ jobs:
|
|||
timeout-minutes: ${{ inputs.timeout-minutes }}
|
||||
env:
|
||||
TEST_PATH: ${{ inputs.test-path }}
|
||||
FORK_FLAG: ${{ inputs.fork-flag }}
|
||||
IS_FORK: ${{ github.event_name == 'pull_request' && github.event.pull_request.head.repo.full_name != github.repository }}
|
||||
UNIT_FLAG: ${{ inputs.unit-flag }}
|
||||
MAX_FAILURES: ${{ inputs.max-failures }}
|
||||
WORKERS: ${{ inputs.workers }}
|
||||
RERUNS: ${{ inputs.reruns }}
|
||||
|
|
@ -186,11 +186,11 @@ jobs:
|
|||
run: |
|
||||
echo "has-coverage=false" >> "$GITHUB_OUTPUT"
|
||||
selection="${TEST_PATH}"
|
||||
if [ "${IS_FORK}" = "true" ] && [ -n "${FORK_FLAG}" ]; then
|
||||
selection="${TEST_PATH} $(bash .circleci/scripts/unit_selection.sh "${FORK_FLAG}" | tr '\n' ' ')"
|
||||
if [ -n "${UNIT_FLAG}" ]; then
|
||||
selection="${TEST_PATH} $(bash .circleci/scripts/unit_selection.sh "${UNIT_FLAG}" | tr '\n' ' ')"
|
||||
fi
|
||||
if [ -z "${selection// /}" ]; then
|
||||
echo "shard selection is empty on this event (CircleCI flag ${FORK_FLAG:-none} owns it); nothing to run"
|
||||
echo "shard selection is empty; nothing to run"
|
||||
exit 0
|
||||
fi
|
||||
pytest_args=()
|
||||
|
|
|
|||
1
.github/workflows/conventional-commits.yml
vendored
1
.github/workflows/conventional-commits.yml
vendored
|
|
@ -41,6 +41,7 @@ jobs:
|
|||
ci
|
||||
chore
|
||||
revert
|
||||
security
|
||||
requireScope: false
|
||||
subjectPattern: ^(?![A-Z]).+$
|
||||
subjectPatternError: |
|
||||
|
|
|
|||
13
.github/workflows/create-rc-branch.yml
vendored
13
.github/workflows/create-rc-branch.yml
vendored
|
|
@ -15,6 +15,8 @@ jobs:
|
|||
runs-on: ubuntu-latest
|
||||
permissions:
|
||||
contents: write
|
||||
outputs:
|
||||
version: ${{ steps.version.outputs.version }}
|
||||
steps:
|
||||
- name: Require main
|
||||
env:
|
||||
|
|
@ -64,3 +66,14 @@ jobs:
|
|||
sha: context.sha,
|
||||
});
|
||||
core.info(`Created branch ${branchName} at ${context.sha}`);
|
||||
|
||||
linear-release:
|
||||
name: Move the Linear release to rc
|
||||
needs: create-rc-branch
|
||||
permissions:
|
||||
contents: read
|
||||
uses: ./.github/workflows/linear-release.yml
|
||||
with:
|
||||
rc_version: ${{ needs.create-rc-branch.outputs.version }}
|
||||
secrets:
|
||||
LINEAR_API_KEY: ${{ secrets.LINEAR_API_KEY }}
|
||||
|
|
|
|||
131
.github/workflows/linear-release.yml
vendored
Normal file
131
.github/workflows/linear-release.yml
vendored
Normal file
|
|
@ -0,0 +1,131 @@
|
|||
name: Linear Release
|
||||
|
||||
on:
|
||||
push:
|
||||
branches:
|
||||
- main
|
||||
- "rc/**"
|
||||
release:
|
||||
types: [published]
|
||||
workflow_call:
|
||||
inputs:
|
||||
rc_version:
|
||||
description: "X.Y.0 release whose rc branch was just cut"
|
||||
required: true
|
||||
type: string
|
||||
secrets:
|
||||
LINEAR_API_KEY:
|
||||
required: true
|
||||
|
||||
permissions: {}
|
||||
|
||||
jobs:
|
||||
linear-release:
|
||||
name: Linear Release
|
||||
if: github.repository == 'BerriAI/litellm'
|
||||
runs-on: ubuntu-latest
|
||||
permissions:
|
||||
contents: read
|
||||
steps:
|
||||
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
|
||||
with:
|
||||
fetch-depth: 0
|
||||
persist-credentials: false
|
||||
|
||||
- name: Plan
|
||||
id: plan
|
||||
env:
|
||||
EVENT: ${{ github.event_name }}
|
||||
REF_NAME: ${{ github.ref_name }}
|
||||
BEFORE: ${{ github.event.before }}
|
||||
CREATED: ${{ github.event.created }}
|
||||
RC_VERSION: ${{ inputs.rc_version }}
|
||||
RELEASE_TAG: ${{ github.event.release.tag_name }}
|
||||
PRERELEASE: ${{ github.event.release.prerelease }}
|
||||
run: |
|
||||
set -euo pipefail
|
||||
sync_base="${BEFORE}"
|
||||
if [ "${CREATED}" = "true" ]; then
|
||||
sync_base=""
|
||||
fi
|
||||
if [ -n "${RC_VERSION}" ]; then
|
||||
echo "version=${RC_VERSION}" >> "$GITHUB_OUTPUT"
|
||||
echo "stage=rc" >> "$GITHUB_OUTPUT"
|
||||
elif [ "${EVENT}" = "release" ]; then
|
||||
if [ "${PRERELEASE}" = "true" ] || ! echo "${RELEASE_TAG}" | grep -qE '^v[0-9]+\.[0-9]+\.0$'; then
|
||||
echo "::notice::${RELEASE_TAG} is not an X.Y.0 stable release; nothing to complete"
|
||||
exit 0
|
||||
fi
|
||||
echo "version=${RELEASE_TAG#v}" >> "$GITHUB_OUTPUT"
|
||||
echo "complete=true" >> "$GITHUB_OUTPUT"
|
||||
elif [ "${REF_NAME}" = "main" ]; then
|
||||
version="$(python3 .github/scripts/read_rc_version.py | cut -d= -f2)"
|
||||
status=0
|
||||
git ls-remote --exit-code --heads origin "rc/${version}" > /dev/null || status=$?
|
||||
case "${status}" in
|
||||
0)
|
||||
IFS=. read -r major minor _ <<< "${version}"
|
||||
version="${major}.$((minor + 1)).0"
|
||||
;;
|
||||
2) ;;
|
||||
*)
|
||||
echo "::error::could not check whether rc/${version} exists (git ls-remote exit ${status})"
|
||||
exit 1
|
||||
;;
|
||||
esac
|
||||
echo "version=${version}" >> "$GITHUB_OUTPUT"
|
||||
echo "sync_base=${sync_base}" >> "$GITHUB_OUTPUT"
|
||||
echo "main=true" >> "$GITHUB_OUTPUT"
|
||||
else
|
||||
echo "version=${REF_NAME#rc/}" >> "$GITHUB_OUTPUT"
|
||||
echo "sync_base=${sync_base}" >> "$GITHUB_OUTPUT"
|
||||
echo "stage=rc" >> "$GITHUB_OUTPUT"
|
||||
fi
|
||||
|
||||
- name: Sync commits into the release
|
||||
if: steps.plan.outputs.sync_base != ''
|
||||
uses: linear/linear-release-action@d4af10092984f9bc6d5efa075b242bdf01333463 # v0.18.0
|
||||
with:
|
||||
access_key: ${{ secrets.LINEAR_API_KEY }}
|
||||
command: sync
|
||||
name: LiteLLM ${{ steps.plan.outputs.version }}
|
||||
version: ${{ steps.plan.outputs.version }}
|
||||
base_ref: ${{ steps.plan.outputs.sync_base }}
|
||||
cli_version: v0.18.0
|
||||
|
||||
- name: Keep the main stage unless the rc branch was cut during this run
|
||||
id: main_stage
|
||||
if: steps.plan.outputs.main == 'true'
|
||||
env:
|
||||
VERSION: ${{ steps.plan.outputs.version }}
|
||||
run: |
|
||||
set -euo pipefail
|
||||
status=0
|
||||
git ls-remote --exit-code --heads origin "rc/${VERSION}" > /dev/null || status=$?
|
||||
case "${status}" in
|
||||
0) echo "::notice::rc/${VERSION} was cut during this run; leaving the release in its rc stage" ;;
|
||||
2) echo "stage=main" >> "$GITHUB_OUTPUT" ;;
|
||||
*)
|
||||
echo "::error::could not check whether rc/${VERSION} exists (git ls-remote exit ${status})"
|
||||
exit 1
|
||||
;;
|
||||
esac
|
||||
|
||||
- name: Move the release to its stage
|
||||
if: steps.plan.outputs.stage != '' || steps.main_stage.outputs.stage != ''
|
||||
uses: linear/linear-release-action@d4af10092984f9bc6d5efa075b242bdf01333463 # v0.18.0
|
||||
with:
|
||||
access_key: ${{ secrets.LINEAR_API_KEY }}
|
||||
command: update
|
||||
stage: ${{ steps.plan.outputs.stage || steps.main_stage.outputs.stage }}
|
||||
version: ${{ steps.plan.outputs.version }}
|
||||
cli_version: v0.18.0
|
||||
|
||||
- name: Complete the release
|
||||
if: steps.plan.outputs.complete == 'true'
|
||||
uses: linear/linear-release-action@d4af10092984f9bc6d5efa075b242bdf01333463 # v0.18.0
|
||||
with:
|
||||
access_key: ${{ secrets.LINEAR_API_KEY }}
|
||||
command: complete
|
||||
version: ${{ steps.plan.outputs.version }}
|
||||
cli_version: v0.18.0
|
||||
3
.github/workflows/test-code-quality.yml
vendored
3
.github/workflows/test-code-quality.yml
vendored
|
|
@ -146,6 +146,9 @@ jobs:
|
|||
- name: check_migrations_no_data_rewrites
|
||||
run: uv run --no-sync python ./tests/code_coverage_tests/check_migrations_no_data_rewrites.py
|
||||
|
||||
- name: check_unbounded_in_lists (fails on findings not in the baseline)
|
||||
run: uv run --no-sync python ./tests/code_coverage_tests/check_unbounded_in_lists.py
|
||||
|
||||
- name: memory_test
|
||||
run: uv run --no-sync python ./tests/code_coverage_tests/memory_test.py
|
||||
|
||||
|
|
|
|||
14
.github/workflows/test-redis-compat.yml
vendored
14
.github/workflows/test-redis-compat.yml
vendored
|
|
@ -8,9 +8,13 @@ on:
|
|||
paths:
|
||||
- "litellm/_redis.py"
|
||||
- "litellm/_redis_credential_provider.py"
|
||||
- "tests/test_litellm/test_redis.py"
|
||||
- "litellm/caching/redis_cache.py"
|
||||
- "litellm/caching/evicted_client_closer.py"
|
||||
- "tests/unit/test_redis.py"
|
||||
- "tests/local_testing/test_caching.py"
|
||||
- "tests/test_litellm/caching/test_redis_connection_pool.py"
|
||||
- "tests/unit/caching/test_redis_connection_pool.py"
|
||||
- "tests/unit/caching/test_redis_cluster_cache.py"
|
||||
- "tests/unit/caching/test_evicted_client_closer.py"
|
||||
- ".github/workflows/test-redis-compat.yml"
|
||||
- "pyproject.toml"
|
||||
- "uv.lock"
|
||||
|
|
@ -80,8 +84,10 @@ jobs:
|
|||
run: |
|
||||
redis-server --version
|
||||
uv run --no-sync pytest \
|
||||
tests/test_litellm/test_redis.py \
|
||||
tests/test_litellm/caching/test_redis_connection_pool.py \
|
||||
tests/unit/test_redis.py \
|
||||
tests/unit/caching/test_redis_connection_pool.py \
|
||||
tests/unit/caching/test_redis_cluster_cache.py \
|
||||
tests/unit/caching/test_evicted_client_closer.py \
|
||||
tests/local_testing/test_caching.py::test_sync_cluster_authenticates_with_azure_credentials \
|
||||
tests/local_testing/test_caching.py::test_sync_cluster_authenticates_with_gcp_credentials \
|
||||
--tb=short -vv \
|
||||
|
|
|
|||
8
.github/workflows/test-rust.yml
vendored
8
.github/workflows/test-rust.yml
vendored
|
|
@ -14,7 +14,6 @@ on:
|
|||
- "litellm/ocr/**"
|
||||
- "litellm/llms/base_llm/ocr/**"
|
||||
- "litellm/llms/custom_httpx/llm_http_handler.py"
|
||||
- "tests/test_litellm/ocr/**"
|
||||
- "tests/test_litellm/conftest.py"
|
||||
- "Makefile"
|
||||
- ".cargo/**"
|
||||
|
|
@ -24,7 +23,7 @@ on:
|
|||
- ".github/actions/setup-uv-with-retries/**"
|
||||
- ".github/scripts/smoke_test_native_wheel.py"
|
||||
- ".github/scripts/verify_linux_native_wheel.py"
|
||||
- "tests/test_litellm/rust_bridge/native_route_wheel_test.py"
|
||||
- "tests/unit/rust_bridge/native_route_wheel_test.py"
|
||||
- ".github/workflows/test-rust.yml"
|
||||
pull_request:
|
||||
branches:
|
||||
|
|
@ -42,7 +41,6 @@ on:
|
|||
- "litellm/ocr/**"
|
||||
- "litellm/llms/base_llm/ocr/**"
|
||||
- "litellm/llms/custom_httpx/llm_http_handler.py"
|
||||
- "tests/test_litellm/ocr/**"
|
||||
- "tests/test_litellm/conftest.py"
|
||||
- "Makefile"
|
||||
- ".cargo/**"
|
||||
|
|
@ -52,7 +50,7 @@ on:
|
|||
- ".github/actions/setup-uv-with-retries/**"
|
||||
- ".github/scripts/smoke_test_native_wheel.py"
|
||||
- ".github/scripts/verify_linux_native_wheel.py"
|
||||
- "tests/test_litellm/rust_bridge/native_route_wheel_test.py"
|
||||
- "tests/unit/rust_bridge/native_route_wheel_test.py"
|
||||
- ".github/workflows/test-rust.yml"
|
||||
|
||||
permissions:
|
||||
|
|
@ -171,7 +169,7 @@ jobs:
|
|||
env:
|
||||
RELEASE_WHEEL_COMMIT_SHA: ${{ github.event.pull_request.head.sha || github.sha }}
|
||||
|
||||
- run: python tests/test_litellm/rust_bridge/native_route_wheel_test.py dist/*.whl
|
||||
- run: python tests/unit/rust_bridge/native_route_wheel_test.py dist/*.whl
|
||||
|
||||
- name: Run pytest tests/test_litellm_rust with the compiled extension
|
||||
run: make test-rust-extension
|
||||
|
|
|
|||
33
.github/workflows/test-unit-proxy-db.yml
vendored
33
.github/workflows/test-unit-proxy-db.yml
vendored
|
|
@ -22,9 +22,10 @@ concurrency:
|
|||
#
|
||||
# `.circleci/tests.yml` runs each group's files on same-repo events under the
|
||||
# `proxy-db-<group>` Codecov flag; `.circleci/scripts/unit_selection.sh` holds
|
||||
# the file lists. CircleCI does not build pull requests from forks, so `fork-flag`
|
||||
# makes the shard run that list there. `test-path` keeps the files that still
|
||||
# reach real providers and never left tests/proxy_unit_tests.
|
||||
# the file lists. That pipeline is manual-only while the tests migrate, so
|
||||
# `unit-flag` makes the shard run that list on every event. `test-path` keeps
|
||||
# the files that still reach real providers and never left
|
||||
# tests/proxy_unit_tests.
|
||||
#
|
||||
# Design targets:
|
||||
# * Every shard runs in <= 7 minutes of wall-clock on the default runner.
|
||||
|
|
@ -78,7 +79,7 @@ jobs:
|
|||
# Must run serially — event-loop conflict with the logging worker.
|
||||
- test-group: key-generation
|
||||
test-path: ""
|
||||
fork-flag: proxy-db-key-generation
|
||||
unit-flag: proxy-db-key-generation
|
||||
workers: 0
|
||||
dist: loadscope
|
||||
timeout: 20
|
||||
|
|
@ -86,13 +87,13 @@ jobs:
|
|||
# ---- auth: split into 2 shards ----
|
||||
- test-group: auth-checks
|
||||
test-path: ""
|
||||
fork-flag: proxy-db-auth-checks
|
||||
unit-flag: proxy-db-auth-checks
|
||||
workers: 4
|
||||
dist: loadscope
|
||||
timeout: 15
|
||||
- test-group: jwt-and-keys
|
||||
test-path: ""
|
||||
fork-flag: proxy-db-jwt-and-keys
|
||||
unit-flag: proxy-db-jwt-and-keys
|
||||
workers: 4
|
||||
dist: loadscope
|
||||
timeout: 15
|
||||
|
|
@ -100,7 +101,7 @@ jobs:
|
|||
# ---- test_proxy_utils.py, single shard, worksteal distribution ----
|
||||
- test-group: proxy-utils
|
||||
test-path: ""
|
||||
fork-flag: proxy-db-proxy-utils
|
||||
unit-flag: proxy-db-proxy-utils
|
||||
workers: 4
|
||||
dist: worksteal
|
||||
timeout: 15
|
||||
|
|
@ -108,13 +109,13 @@ jobs:
|
|||
# ---- proxy server: split into 2 shards ----
|
||||
- test-group: proxy-server-core
|
||||
test-path: "tests/proxy_unit_tests/test_proxy_server_gemini_pass_through.py"
|
||||
fork-flag: proxy-db-proxy-server-core
|
||||
unit-flag: proxy-db-proxy-server-core
|
||||
workers: 4
|
||||
dist: loadscope
|
||||
timeout: 15
|
||||
- test-group: proxy-runtime
|
||||
test-path: ""
|
||||
fork-flag: proxy-db-proxy-runtime
|
||||
unit-flag: proxy-db-proxy-runtime
|
||||
workers: 4
|
||||
dist: loadscope
|
||||
timeout: 15
|
||||
|
|
@ -122,20 +123,20 @@ jobs:
|
|||
# ---- logging: split into 2 shards ----
|
||||
- test-group: custom-logging
|
||||
test-path: "tests/proxy_unit_tests/test_proxy_custom_logger.py"
|
||||
fork-flag: proxy-db-custom-logging
|
||||
unit-flag: proxy-db-custom-logging
|
||||
workers: 4
|
||||
dist: loadscope
|
||||
timeout: 15
|
||||
- test-group: logging-misc
|
||||
test-path: ""
|
||||
fork-flag: proxy-db-logging-misc
|
||||
unit-flag: proxy-db-logging-misc
|
||||
workers: 4
|
||||
dist: loadscope
|
||||
timeout: 15
|
||||
|
||||
- test-group: db-and-spend
|
||||
test-path: ""
|
||||
fork-flag: proxy-db-db-and-spend
|
||||
unit-flag: proxy-db-db-and-spend
|
||||
workers: 4
|
||||
dist: loadscope
|
||||
timeout: 15
|
||||
|
|
@ -143,27 +144,27 @@ jobs:
|
|||
# ---- guardrails + budget + hooks: split into 2 ----
|
||||
- test-group: guardrails-hooks
|
||||
test-path: ""
|
||||
fork-flag: proxy-db-guardrails-hooks
|
||||
unit-flag: proxy-db-guardrails-hooks
|
||||
workers: 4
|
||||
dist: loadscope
|
||||
timeout: 15
|
||||
- test-group: budgets
|
||||
test-path: ""
|
||||
fork-flag: proxy-db-budgets
|
||||
unit-flag: proxy-db-budgets
|
||||
workers: 4
|
||||
dist: loadscope
|
||||
timeout: 15
|
||||
|
||||
- test-group: endpoints-and-responses
|
||||
test-path: "tests/proxy_unit_tests/test_proxy_exception_mapping.py"
|
||||
fork-flag: proxy-db-endpoints-and-responses
|
||||
unit-flag: proxy-db-endpoints-and-responses
|
||||
workers: 4
|
||||
dist: loadscope
|
||||
timeout: 15
|
||||
uses: ./.github/workflows/_test-unit-base.yml
|
||||
with:
|
||||
test-path: ${{ matrix.test-path }}
|
||||
fork-flag: ${{ matrix.fork-flag }}
|
||||
unit-flag: ${{ matrix.unit-flag }}
|
||||
workers: ${{ matrix.workers }}
|
||||
reruns: 2
|
||||
timeout-minutes: ${{ matrix.timeout }}
|
||||
|
|
|
|||
65
.github/workflows/test-unit.yml
vendored
65
.github/workflows/test-unit.yml
vendored
|
|
@ -36,9 +36,9 @@ concurrency:
|
|||
# Folding it in here is a follow-up, together with generalising that guard into
|
||||
# assert_ci_coverage.py.
|
||||
#
|
||||
# `fork-flag` names the `.circleci/tests.yml` job that now runs part of the
|
||||
# shard under the same Codecov flag. CircleCI does not build pull requests from
|
||||
# forks, so the shard still runs those files there and skips them elsewhere.
|
||||
# `unit-flag` names the `.circleci/tests.yml` job that now runs part of the
|
||||
# shard under the same Codecov flag. That pipeline is manual-only while the
|
||||
# tests migrate, so the shard also runs those files on every event.
|
||||
jobs:
|
||||
unit:
|
||||
name: ${{ matrix.shard }}
|
||||
|
|
@ -52,8 +52,8 @@ jobs:
|
|||
include:
|
||||
- shard: mcp-integration
|
||||
artifact-name: mcp-integration
|
||||
test-path: "tests/mcp_tests tests/test_litellm/experimental_mcp_client"
|
||||
fork-flag: mcp-integration
|
||||
test-path: "tests/mcp_tests"
|
||||
unit-flag: mcp-integration
|
||||
workers: 2
|
||||
reruns: 0
|
||||
timeout-minutes: 20
|
||||
|
|
@ -61,7 +61,8 @@ jobs:
|
|||
|
||||
- shard: core-utils
|
||||
artifact-name: core-utils
|
||||
test-path: "tests/test_litellm/litellm_core_utils tests/test_litellm/router"
|
||||
test-path: ""
|
||||
unit-flag: core-utils
|
||||
workers: 2
|
||||
reruns: 1
|
||||
timeout-minutes: 20
|
||||
|
|
@ -69,11 +70,8 @@ jobs:
|
|||
|
||||
- shard: enterprise-routing
|
||||
artifact-name: enterprise-routing
|
||||
test-path: >-
|
||||
tests/test_litellm/google_genai
|
||||
tests/test_litellm/router_utils
|
||||
tests/test_litellm/router_strategy
|
||||
fork-flag: enterprise-routing
|
||||
test-path: ""
|
||||
unit-flag: enterprise-routing
|
||||
workers: 2
|
||||
reruns: 2
|
||||
timeout-minutes: 20
|
||||
|
|
@ -81,7 +79,8 @@ jobs:
|
|||
|
||||
- shard: integrations
|
||||
artifact-name: integrations
|
||||
test-path: "tests/test_litellm/integrations"
|
||||
test-path: ""
|
||||
unit-flag: integrations
|
||||
workers: 2
|
||||
reruns: 3
|
||||
timeout-minutes: 20
|
||||
|
|
@ -89,7 +88,8 @@ jobs:
|
|||
|
||||
- shard: Vertex AI
|
||||
artifact-name: llm-vertex-ai
|
||||
test-path: "tests/test_litellm/llms/vertex_ai"
|
||||
test-path: ""
|
||||
unit-flag: llm-vertex-ai
|
||||
workers: 1
|
||||
reruns: 2
|
||||
timeout-minutes: 20
|
||||
|
|
@ -97,7 +97,8 @@ jobs:
|
|||
|
||||
- shard: All Other Providers
|
||||
artifact-name: llm-other-providers
|
||||
test-path: "tests/test_litellm/llms --ignore=tests/test_litellm/llms/vertex_ai"
|
||||
test-path: ""
|
||||
unit-flag: llm-other-providers
|
||||
workers: 2
|
||||
reruns: 2
|
||||
timeout-minutes: 20
|
||||
|
|
@ -106,26 +107,8 @@ jobs:
|
|||
- shard: misc
|
||||
artifact-name: misc
|
||||
test-path: >-
|
||||
tests/test_litellm/batches
|
||||
tests/test_litellm/secret_managers
|
||||
tests/test_litellm/a2a_protocol
|
||||
tests/test_litellm/chat_completions
|
||||
tests/test_litellm/completion_extras
|
||||
tests/test_litellm/containers
|
||||
tests/test_litellm/endpoints
|
||||
tests/test_litellm/files
|
||||
tests/test_litellm/images
|
||||
tests/test_litellm/interactions
|
||||
tests/test_litellm/messages
|
||||
tests/test_litellm/embeddings
|
||||
tests/test_litellm/ocr
|
||||
tests/test_litellm/passthrough
|
||||
tests/test_litellm/rag
|
||||
tests/test_litellm/rerank_api
|
||||
tests/test_litellm/rust_bridge
|
||||
tests/test_litellm/vector_stores
|
||||
tests/test_litellm/videos
|
||||
tests/test_litellm/test_*.py
|
||||
unit-flag: misc
|
||||
workers: 2
|
||||
reruns: 2
|
||||
timeout-minutes: 20
|
||||
|
|
@ -205,7 +188,7 @@ jobs:
|
|||
tests/test_litellm/proxy/types_utils
|
||||
tests/test_litellm/proxy/logging_endpoints
|
||||
tests/test_litellm/proxy/test_*.py
|
||||
fork-flag: proxy-infra
|
||||
unit-flag: proxy-infra
|
||||
workers: 4
|
||||
reruns: 2
|
||||
timeout-minutes: 20
|
||||
|
|
@ -214,7 +197,7 @@ jobs:
|
|||
- shard: caching-local
|
||||
artifact-name: caching-local
|
||||
test-path: ""
|
||||
fork-flag: caching-local
|
||||
unit-flag: caching-local
|
||||
workers: 2
|
||||
reruns: 2
|
||||
timeout-minutes: 20
|
||||
|
|
@ -223,7 +206,7 @@ jobs:
|
|||
- shard: proxy-extras
|
||||
artifact-name: proxy-extras
|
||||
test-path: ""
|
||||
fork-flag: proxy-extras
|
||||
unit-flag: proxy-extras
|
||||
workers: 2
|
||||
reruns: 2
|
||||
timeout-minutes: 20
|
||||
|
|
@ -232,7 +215,7 @@ jobs:
|
|||
- shard: enterprise-package
|
||||
artifact-name: enterprise-package
|
||||
test-path: ""
|
||||
fork-flag: enterprise-package
|
||||
unit-flag: enterprise-package
|
||||
workers: 4
|
||||
reruns: 2
|
||||
timeout-minutes: 20
|
||||
|
|
@ -240,10 +223,8 @@ jobs:
|
|||
|
||||
- shard: responses-caching-types
|
||||
artifact-name: responses-caching-types
|
||||
test-path: >-
|
||||
tests/test_litellm/responses
|
||||
tests/test_litellm/caching
|
||||
tests/test_litellm/types
|
||||
test-path: ""
|
||||
unit-flag: responses-caching-types
|
||||
workers: 2
|
||||
reruns: 2
|
||||
timeout-minutes: 20
|
||||
|
|
@ -251,7 +232,7 @@ jobs:
|
|||
uses: ./.github/workflows/_test-unit-base.yml
|
||||
with:
|
||||
test-path: ${{ matrix.test-path }}
|
||||
fork-flag: ${{ matrix.fork-flag || '' }}
|
||||
unit-flag: ${{ matrix.unit-flag || '' }}
|
||||
workers: ${{ matrix.workers }}
|
||||
reruns: ${{ matrix.reruns }}
|
||||
timeout-minutes: ${{ matrix.timeout-minutes }}
|
||||
|
|
|
|||
|
|
@ -27,7 +27,7 @@ Never test structure of code only function of it
|
|||
|
||||
A test must only fail when litellm code changes. Never pin facts we don't own (a vendor's price, a third party's field, an upstream default, today's date) as literals or as "X must be absent"; assert the invariant our code guarantees instead, e.g. two rows agree, a value is within range, a field is derived from another. If an outside fact is truly load-bearing, cite its source and date next to the assertion so a reader can tell stale from broken
|
||||
|
||||
`tests/test_litellm/` mirrors `litellm/` in a parallel path (see `tests/test_litellm/readme.md`). Name tests `test_<filename>.py`, but always match the existing test file in the directory you touch — many provider dirs use longer descriptive names (e.g. `test_anthropic_chat_transformation.py`) to avoid ambiguity across sibling folders. For bug fixes, extend the existing mapped test file rather than creating a new one. Only create a new test file for a new feature (provider, endpoint, or transformation module) that has no mapped test yet, following that directory's naming convention (or `test_<filename>.py` if you're the first test there). One focused regression test beats many shallow ones
|
||||
`tests/unit/` mirrors `litellm/` in a parallel path (see `tests/unit/AGENTS.md`). Name tests `test_<filename>.py`, but always match the existing test file in the directory you touch — many provider dirs use longer descriptive names (e.g. `test_anthropic_chat_transformation.py`) to avoid ambiguity across sibling folders. For bug fixes, extend the existing mapped test file rather than creating a new one. Only create a new test file for a new feature (provider, endpoint, or transformation module) that has no mapped test yet, following that directory's naming convention (or `test_<filename>.py` if you're the first test there). One focused regression test beats many shallow ones
|
||||
|
||||
End-to-end tests belong in `tests/e2e/` and must follow the harness conventions documented in that directory's `AGENTS.md`
|
||||
|
||||
|
|
|
|||
|
|
@ -255,7 +255,7 @@ Conventions to follow when touching this layer:
|
|||
| Column vs. field names | Where a model field differs from its DB column (for example `org_id` maps to the `organization_id` column), the repository translates in both directions rather than relying on Pydantic to guess. |
|
||||
| Array mutations | Adds use Prisma's atomic `push` (`add_member`, `add_admin`, `add_models`) to avoid read-modify-write races. Removals fall back to read-modify-write because Prisma has no atomic array remove. |
|
||||
|
||||
To add a new entity, define the model under `litellm/models/`, re-export it from `proxy/_types.py` if existing code imports it from there, and add a repository under `litellm/repositories/` (subclass `BaseRepository` for plain CRUD, or add bespoke methods when the entity needs encryption, archiving, or atomic array updates). Mirror the tests in `tests/test_litellm/repositories/`.
|
||||
To add a new entity, define the model under `litellm/models/`, re-export it from `proxy/_types.py` if existing code imports it from there, and add a repository under `litellm/repositories/` (subclass `BaseRepository` for plain CRUD, or add bespoke methods when the entity needs encryption, archiving, or atomic array updates). Mirror the tests in `tests/unit/repositories/`.
|
||||
|
||||
---
|
||||
|
||||
|
|
@ -336,7 +336,7 @@ Each translation is isolated in its own file, making it easy to test and modify
|
|||
| `/v1/chat/completions` | Gemini | `llms/gemini/chat/transformation.py` |
|
||||
| `/v1/chat/completions` | Vertex AI | `llms/vertex_ai/gemini/transformation.py` |
|
||||
| `/v1/chat/completions` | OpenAI | `llms/openai/chat/gpt_transformation.py` |
|
||||
| `/v1/messages` (passthrough) | Anthropic | `llms/anthropic/experimental_pass_through/messages/transformation.py` |
|
||||
| `/v1/messages` (passthrough) | Anthropic | `llms/anthropic/pass_through/messages/transformation.py` |
|
||||
| `/v1/messages` (passthrough) | Bedrock | `llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py` |
|
||||
| `/v1/messages` (passthrough) | Vertex AI | `llms/vertex_ai/vertex_ai_partner_models/anthropic/experimental_pass_through/transformation.py` |
|
||||
| Passthrough endpoints | All | `proxy/pass_through_endpoints/llm_provider_handlers/` |
|
||||
|
|
|
|||
|
|
@ -14,7 +14,7 @@ Here are the core requirements for any PR submitted to LiteLLM:
|
|||
- [ ] **Add testing** - Adding at least 1 test is a hard requirement - [see details](#adding-testing)
|
||||
- [ ] **Ensure your PR passes all checks**:
|
||||
- [ ] [Linting / Formatting](#running-linting-and-formatting-checks) - `make lint`
|
||||
- [ ] [The tests covering your change](#running-unit-tests) pass, e.g. `uv run pytest tests/test_litellm/<your_test_file>.py -v`. CI runs the full unit test matrix, so you don't need to run the whole suite locally
|
||||
- [ ] [The tests covering your change](#running-unit-tests) pass, e.g. `uv run pytest tests/unit/<your_test_file>.py -v`. CI runs the full unit test matrix, so you don't need to run the whole suite locally
|
||||
|
||||
#### UI PRs
|
||||
|
||||
|
|
@ -72,7 +72,7 @@ make format
|
|||
make lint
|
||||
|
||||
# Run the tests covering your change (CI runs the full suite)
|
||||
uv run pytest tests/test_litellm/<your_test_file>.py -v
|
||||
uv run pytest tests/unit/<your_test_file>.py -v
|
||||
|
||||
# Commit your changes (must follow Conventional Commits — see above)
|
||||
git add .
|
||||
|
|
@ -88,7 +88,7 @@ git push origin feature/your-feature
|
|||
|
||||
### Where to Add Tests
|
||||
|
||||
Add your tests to the [`tests/test_litellm/` directory](https://github.com/BerriAI/litellm/tree/main/tests/test_litellm).
|
||||
Add your tests to the [`tests/unit/` directory](https://github.com/BerriAI/litellm/tree/main/tests/unit).
|
||||
|
||||
- This directory mirrors the structure of the `litellm/` directory
|
||||
- **Only add mocked tests** - no real LLM API calls in this directory
|
||||
|
|
@ -96,10 +96,10 @@ Add your tests to the [`tests/test_litellm/` directory](https://github.com/Berri
|
|||
|
||||
### File Naming Convention
|
||||
|
||||
The `tests/test_litellm/` directory follows the same structure as `litellm/`:
|
||||
The `tests/unit/` directory follows the same structure as `litellm/`:
|
||||
|
||||
- `litellm/proxy/caching_routes.py` → `tests/test_litellm/proxy/test_caching_routes.py`
|
||||
- `litellm/utils.py` → `tests/test_litellm/test_utils.py`
|
||||
- `litellm/utils.py` → `tests/unit/test_utils.py`
|
||||
|
||||
### Example Test
|
||||
|
||||
|
|
@ -125,10 +125,10 @@ def test_your_feature():
|
|||
|
||||
Run the tests covering your change:
|
||||
```bash
|
||||
uv run pytest tests/test_litellm/test_your_file.py -v
|
||||
uv run pytest tests/unit/test_your_file.py -v
|
||||
```
|
||||
|
||||
`tests/test_litellm` holds thousands of tests, so running all of it locally takes a long time. CI runs it as a parallel matrix (`make test-unit-llms`, `make test-unit-proxy-core`, and the other `test-unit-*` targets) on beefier boxes, so if, for whatever reason, you must run the whole suite, it's better to rely on CI to do that.
|
||||
`tests/unit` holds thousands of tests, so running all of it locally takes a long time. CI runs it as a parallel matrix (`make test-unit-llms`, `make test-unit-proxy-core`, and the other `test-unit-*` targets) on beefier boxes, so if, for whatever reason, you must run the whole suite, it's better to rely on CI to do that.
|
||||
|
||||
If you're running broader test suites, proxy tests, or anything that touches PostgreSQL-backed fixtures/plugins, install the full local test environment first:
|
||||
|
||||
|
|
|
|||
22
Makefile
22
Makefile
|
|
@ -4,7 +4,7 @@
|
|||
.PHONY: help test test-unit test-unit-llms test-unit-proxy-guardrails test-unit-proxy-core test-unit-proxy-misc \
|
||||
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 \
|
||||
test-rust-extension rust-sqlx-prepare \
|
||||
info lint lint-inner lint-dev lint-checks format \
|
||||
lint-basedpyright lint-e2e-basedpyright lint-basedpyright-budget-update lint-type-discipline lint-type-discipline-budget-update \
|
||||
lint-ruff-budget lint-ruff-budget-update lint-budget-update lint-gate \
|
||||
|
|
@ -42,7 +42,7 @@ help:
|
|||
@echo " make check-circular-imports - Check for circular imports"
|
||||
@echo " make check-import-safety - Check import safety"
|
||||
@echo " make test - Run all tests"
|
||||
@echo " make test-unit - Run unit tests (tests/test_litellm)"
|
||||
@echo " make test-unit - Run unit tests (tests/unit and tests/test_litellm)"
|
||||
@echo " make test-unit-llms - Run LLM provider tests (~225 files)"
|
||||
@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)"
|
||||
|
|
@ -56,6 +56,7 @@ help:
|
|||
@echo " make test-integration - Run integration tests"
|
||||
@echo " make test-unit-helm - Run helm unit tests"
|
||||
@echo " make test-rust-extension - Build the Rust extension and run its public Python tests"
|
||||
@echo " make rust-sqlx-prepare - Refresh litellm-rust/crates/db/.sqlx against a migrated Postgres container"
|
||||
@echo ""
|
||||
@echo "Heavy targets (check, lint) queue for LITELLM_GATE_SLOTS machine-wide"
|
||||
@echo "slots (default 2; 0 disables) so parallel sessions don't thrash one machine."
|
||||
|
|
@ -301,20 +302,23 @@ test-rust-extension:
|
|||
UV_PROJECT_ENVIRONMENT="$$temporary/venv" $(UV) sync --python 3.12 --frozen --no-install-project --all-groups --all-extras && \
|
||||
$(UV) pip install --python "$$temporary/venv/bin/python" --no-deps "$$1" && \
|
||||
"$$temporary/venv/bin/python" -I -m mypy.stubtest \
|
||||
--mypy-config-file tests/test_litellm/rust_bridge/stubtest.ini \
|
||||
--mypy-config-file tests/unit/rust_bridge/stubtest.ini \
|
||||
litellm.rust_bridge._native && \
|
||||
LITELLM_RUST=1 LITELLM_LOCAL_MODEL_COST_MAP=True \
|
||||
"$$temporary/venv/bin/python" -I -m pytest --import-mode=importlib -m requires_rust_extension tests/test_litellm_rust
|
||||
|
||||
rust-sqlx-prepare:
|
||||
cd litellm-rust && cargo run -p litellm-db-testing --bin sqlx-prepare
|
||||
|
||||
test: install-test-deps
|
||||
$(UV_RUN) pytest tests/
|
||||
|
||||
test-unit: install-test-deps
|
||||
$(UV_RUN) pytest tests/test_litellm -x -vv -n 4
|
||||
$(UV_RUN) pytest tests/unit tests/test_litellm -x -vv -n 4
|
||||
|
||||
# Matrix test targets (matching CI workflow groups)
|
||||
test-unit-llms: install-test-deps
|
||||
$(UV_RUN) pytest tests/test_litellm/llms --tb=short -vv -n 4 --durations=20
|
||||
$(UV_RUN) pytest tests/unit/llms --tb=short -vv -n 4 --durations=20
|
||||
|
||||
test-unit-proxy-guardrails: install-test-deps
|
||||
$(UV_RUN) pytest tests/test_litellm/proxy/guardrails tests/test_litellm/proxy/management_endpoints tests/test_litellm/proxy/management_helpers --tb=short -vv -n 4 --durations=20
|
||||
|
|
@ -326,16 +330,16 @@ test-unit-proxy-misc: install-test-deps
|
|||
$(UV_RUN) pytest tests/test_litellm/proxy/_experimental tests/test_litellm/proxy/agent_endpoints tests/test_litellm/proxy/anthropic_endpoints tests/test_litellm/proxy/common_utils tests/test_litellm/proxy/discovery_endpoints tests/test_litellm/proxy/experimental tests/test_litellm/proxy/google_endpoints tests/test_litellm/proxy/health_endpoints tests/test_litellm/proxy/image_endpoints tests/test_litellm/proxy/middleware tests/test_litellm/proxy/openai_files_endpoint tests/test_litellm/proxy/pass_through_endpoints tests/test_litellm/proxy/prompts tests/test_litellm/proxy/public_endpoints tests/test_litellm/proxy/response_api_endpoints tests/test_litellm/proxy/shutdown tests/test_litellm/proxy/spend_tracking tests/test_litellm/proxy/ui_crud_endpoints tests/test_litellm/proxy/vector_store_endpoints tests/test_litellm/proxy/test_*.py --tb=short -vv -n 4 --durations=20
|
||||
|
||||
test-unit-integrations: install-test-deps
|
||||
$(UV_RUN) pytest tests/test_litellm/integrations --tb=short -vv -n 4 --durations=20
|
||||
$(UV_RUN) pytest tests/unit/integrations --tb=short -vv -n 4 --durations=20
|
||||
|
||||
test-unit-core-utils: install-test-deps
|
||||
$(UV_RUN) pytest tests/test_litellm/litellm_core_utils --tb=short -vv -n 2 --durations=20
|
||||
$(UV_RUN) pytest tests/unit/litellm_core_utils --tb=short -vv -n 2 --durations=20
|
||||
|
||||
test-unit-other: install-test-deps
|
||||
$(UV_RUN) pytest tests/test_litellm/caching tests/test_litellm/responses tests/test_litellm/secret_managers tests/test_litellm/vector_stores tests/test_litellm/a2a_protocol tests/test_litellm/anthropic_interface tests/test_litellm/completion_extras tests/test_litellm/containers tests/unit/enterprise tests/test_litellm/experimental_mcp_client tests/test_litellm/google_genai tests/test_litellm/images tests/test_litellm/interactions tests/test_litellm/passthrough tests/test_litellm/router_strategy tests/test_litellm/router_utils tests/test_litellm/types --tb=short -vv -n 4 --durations=20
|
||||
$(UV_RUN) pytest tests/unit/caching tests/unit/responses tests/unit/secret_managers tests/unit/vector_stores tests/unit/a2a_protocol tests/unit/completion_extras tests/unit/containers tests/unit/enterprise tests/unit/experimental_mcp_client tests/unit/google_genai tests/unit/images tests/unit/interactions tests/unit/router_strategy tests/unit/router_utils tests/unit/types --tb=short -vv -n 4 --durations=20
|
||||
|
||||
test-unit-root: install-test-deps
|
||||
$(UV_RUN) pytest tests/test_litellm/test_*.py --tb=short -vv -n 4 --durations=20
|
||||
$(UV_RUN) pytest tests/unit/test_*.py tests/test_litellm/test_*.py --tb=short -vv -n 4 --durations=20
|
||||
|
||||
# Proxy unit tests (tests/unit/proxy split alphabetically)
|
||||
test-proxy-unit-a: install-test-deps
|
||||
|
|
|
|||
|
|
@ -362,6 +362,7 @@ For MCP OAuth, an upstream may advertise dynamic client registration but refuse
|
|||
| [Recraft (`recraft`)](https://docs.litellm.ai/docs/providers/recraft) | | | | | ✅ | | | | | |
|
||||
| [Replicate (`replicate`)](https://docs.litellm.ai/docs/providers/replicate) | ✅ | ✅ | ✅ | | | | | | | |
|
||||
| [Sagemaker Chat (`sagemaker_chat`)](https://docs.litellm.ai/docs/providers/aws_sagemaker) | ✅ | ✅ | ✅ | | | | | | | |
|
||||
| [Sail (`sail`)](https://docs.litellm.ai/docs/providers/sail) | ✅ | ✅ | ✅ | | | | | | | |
|
||||
| [Sambanova (`sambanova`)](https://docs.litellm.ai/docs/providers/sambanova) | ✅ | ✅ | ✅ | | | | | | | |
|
||||
| [Snowflake (`snowflake`)](https://docs.litellm.ai/docs/providers/snowflake) | ✅ | ✅ | ✅ | | | | | | | |
|
||||
| [Text Completion Codestral (`text-completion-codestral`)](https://docs.litellm.ai/docs/providers/codestral) | ✅ | ✅ | ✅ | | | | | | | |
|
||||
|
|
|
|||
|
|
@ -2,7 +2,7 @@
|
|||
|
||||
Every `litellm_*` metric family the proxy can expose on `/metrics` (136 families across 97 panels), grouped into rows: proxy traffic, latency, spend and tokens, cache, LLM API deployments, key and team rate limits, budgets, guardrails, MCP, managed files and batches, users and teams, the Redis circuit breaker, the spend log cleanup job, and the `prometheus_system` service callback metrics (per-service latency, request and failure rates, spend update queue sizes). Panel titles are the metric names so you can grep the JSON for the metric you care about
|
||||
|
||||
Import `grafana_dashboard.json` from **Dashboards > New > Import** and pick your Prometheus data source when prompted (the `DS_PROMETHEUS` variable). Counters are plotted as `rate()` over `$__rate_interval`, histograms as p50 / p95 / p99, gauges as the raw value grouped by the most useful label. Every query names the metric exactly as the proxy emits it (counters carry the `_total` suffix the Prometheus client adds), and `tests/test_litellm/integrations/test_prometheus_metric_name_consistency.py` fails if a metric is renamed without updating this dashboard
|
||||
Import `grafana_dashboard.json` from **Dashboards > New > Import** and pick your Prometheus data source when prompted (the `DS_PROMETHEUS` variable). Counters are plotted as `rate()` over `$__rate_interval`, histograms as p50 / p95 / p99, gauges as the raw value grouped by the most useful label. Every query names the metric exactly as the proxy emits it (counters carry the `_total` suffix the Prometheus client adds), and `tests/unit/integrations/test_prometheus_metric_name_consistency.py` fails if a metric is renamed without updating this dashboard
|
||||
|
||||
The first eleven rows need only `callbacks: ["prometheus"]`. The last three rows and the `litellm_admission_*` panels are emitted by other subsystems and stay empty until those are on: the service callback row needs `service_callback: ["prometheus_system"]` in `litellm_settings`, the circuit breaker row needs a Redis cache, the cleanup row needs spend log retention, and admission control needs its middleware enabled. Within the base rows, many panels only fill in once the matching feature is in use: budgets need keys, teams, users or orgs with `max_budget` set, cache panels need caching on, guardrail and MCP panels need those features configured, deployment health needs the router with more than one deployment or a failure to record, and `litellm_in_flight_requests` needs traffic at scrape time. An empty panel for a feature you do not use is expected
|
||||
|
||||
|
|
|
|||
|
|
@ -32,6 +32,7 @@ from litellm.integrations.email_templates.key_rotated_email import (
|
|||
from litellm.integrations.email_templates.templates import (
|
||||
MAX_BUDGET_ALERT_EMAIL_TEMPLATE,
|
||||
SOFT_BUDGET_ALERT_EMAIL_TEMPLATE,
|
||||
TEAM_MEMBER_MAX_BUDGET_ALERT_EMAIL_TEMPLATE,
|
||||
TEAM_SOFT_BUDGET_ALERT_EMAIL_TEMPLATE,
|
||||
)
|
||||
from litellm.integrations.email_templates.user_invitation_email import (
|
||||
|
|
@ -48,6 +49,12 @@ from litellm.secret_managers.main import get_secret_bool
|
|||
from litellm.types.integrations.slack_alerting import LITELLM_LOGO_URL
|
||||
|
||||
|
||||
def _max_budget_alert_id(user_info: CallInfo) -> str:
|
||||
if user_info.event_group == Litellm_EntityType.TEAM_MEMBER:
|
||||
return f"team_member:{user_info.user_id}:{user_info.team_id}"
|
||||
return user_info.token or user_info.user_id or "default_id"
|
||||
|
||||
|
||||
def _parse_email_list(raw) -> List[str]:
|
||||
"""Parse emails from a list or comma-separated string."""
|
||||
if isinstance(raw, list):
|
||||
|
|
@ -373,17 +380,31 @@ class BaseEmailLogger(CustomLogger):
|
|||
greeting = html.escape(
|
||||
event.user_email or event.key_alias or event.token or ""
|
||||
)
|
||||
email_html_content = MAX_BUDGET_ALERT_EMAIL_TEMPLATE.format(
|
||||
email_logo_url=email_params.logo_url,
|
||||
recipient_email=greeting,
|
||||
percentage=percentage,
|
||||
spend=spend_str,
|
||||
max_budget=max_budget_str,
|
||||
alert_threshold=alert_threshold_str,
|
||||
base_url=email_params.base_url,
|
||||
email_support_contact=email_params.support_contact,
|
||||
email_footer=email_params.signature,
|
||||
)
|
||||
if event.event_group == Litellm_EntityType.TEAM_MEMBER:
|
||||
email_html_content = TEAM_MEMBER_MAX_BUDGET_ALERT_EMAIL_TEMPLATE.format(
|
||||
email_logo_url=email_params.logo_url,
|
||||
member=html.escape(event.user_email or event.user_id or ""),
|
||||
team_alias=html.escape(event.team_alias or event.team_id or ""),
|
||||
percentage=percentage,
|
||||
spend=spend_str,
|
||||
max_budget=max_budget_str,
|
||||
alert_threshold=alert_threshold_str,
|
||||
base_url=email_params.base_url,
|
||||
email_support_contact=email_params.support_contact,
|
||||
email_footer=email_params.signature,
|
||||
)
|
||||
else:
|
||||
email_html_content = MAX_BUDGET_ALERT_EMAIL_TEMPLATE.format(
|
||||
email_logo_url=email_params.logo_url,
|
||||
recipient_email=greeting,
|
||||
percentage=percentage,
|
||||
spend=spend_str,
|
||||
max_budget=max_budget_str,
|
||||
alert_threshold=alert_threshold_str,
|
||||
base_url=email_params.base_url,
|
||||
email_support_contact=email_params.support_contact,
|
||||
email_footer=email_params.signature,
|
||||
)
|
||||
await self.send_email(
|
||||
from_email=self.DEFAULT_LITELLM_EMAIL,
|
||||
to_email=recipient_emails,
|
||||
|
|
@ -607,7 +628,7 @@ class BaseEmailLogger(CustomLogger):
|
|||
if user_info.spend < threshold_amount:
|
||||
continue
|
||||
|
||||
_id = user_info.token or user_info.user_id or "default_id"
|
||||
_id = _max_budget_alert_id(user_info)
|
||||
_cache_key = (
|
||||
f"email_budget_alerts:max_budget_alert:{threshold_pct}:{_id}"
|
||||
)
|
||||
|
|
@ -618,7 +639,7 @@ class BaseEmailLogger(CustomLogger):
|
|||
emails.append(user_info.user_email)
|
||||
if not emails:
|
||||
verbose_proxy_logger.warning(
|
||||
"No recipients for %d%% threshold on key %s, skipping alert",
|
||||
"No recipients for %d%% threshold on %s, skipping alert",
|
||||
threshold_pct,
|
||||
_id,
|
||||
)
|
||||
|
|
@ -633,7 +654,11 @@ class BaseEmailLogger(CustomLogger):
|
|||
if send_count is not None and send_count > 1:
|
||||
continue
|
||||
|
||||
event_message = f"Max Budget Alert - {threshold_pct}% of Maximum Budget Reached"
|
||||
event_message = (
|
||||
f"Team Member Budget Alert - {threshold_pct}% of Team Member Budget Reached"
|
||||
if user_info.event_group == Litellm_EntityType.TEAM_MEMBER
|
||||
else f"Max Budget Alert - {threshold_pct}% of Maximum Budget Reached"
|
||||
)
|
||||
webhook_event = WebhookEvent(
|
||||
event="max_budget_alert",
|
||||
event_message=event_message,
|
||||
|
|
|
|||
|
|
@ -0,0 +1,2 @@
|
|||
-- AlterTable
|
||||
ALTER TABLE "LiteLLM_MCPServerTable" ADD COLUMN IF NOT EXISTS "pinned_tools" JSONB DEFAULT '{}';
|
||||
|
|
@ -322,6 +322,7 @@ model LiteLLM_MCPServerTable {
|
|||
allowed_tools String[] @default([])
|
||||
tool_name_to_display_name Json? @default("{}")
|
||||
tool_name_to_description Json? @default("{}")
|
||||
pinned_tools Json? @default("{}")
|
||||
extra_headers String[] @default([])
|
||||
static_headers Json? @default("{}")
|
||||
// Admin-configured environment variables interpolated into static_headers
|
||||
|
|
|
|||
22
litellm-rust/.agents/skills/rust-tracing/SKILL.md
Normal file
22
litellm-rust/.agents/skills/rust-tracing/SKILL.md
Normal file
|
|
@ -0,0 +1,22 @@
|
|||
---
|
||||
name: rust-tracing
|
||||
description: Add or change Rust diagnostic tracing in litellm-rust, including route spans, subscriber layers, and Python logger delivery
|
||||
---
|
||||
|
||||
# Rust tracing
|
||||
|
||||
Use upstream `tracing` throughout Rust, including `#[tracing::instrument]`, events, and span propagation. Centralize collection and delivery infrastructure in `crates/tracing`. Direct upstream imports still reach our configured subscriber; re-exporting macros does not control delivery. Do not introduce Rust `log` or `pyo3-log` for this path
|
||||
|
||||
`litellm-tracing` owns shared subscriber layers, span field collection, and diagnostic processing. Keep adapters composable as `tracing_subscriber::Layer`s, with `Logger` providing host setup. Runtime-specific delivery belongs in the host bridge. The Python bridge delivers directly to the existing Python SDK logger, preserving its handlers, filtering, redaction, and request correlation. Keep Python dependencies out of `crates/tracing`
|
||||
|
||||
Hosts configure subscribers. Keep Python execution scoped to its captured dispatch rather than installing a process-wide subscriber. Propagate both span context and dispatch across spawned work and returned streams
|
||||
|
||||
In core, instrument execution shared by native calls and hosted machines. Use consistent route, model, provider, streaming, and outcome fields. Put status recording at shared provider boundaries instead of scattering basic logging through handlers. Keep upstream HTTP status separate from route success
|
||||
|
||||
Use `skip_all` and explicitly selected fields. Basic tracing excludes bodies, credentials, headers, and raw error strings. Avoid automatic `ret` or `err` capture of sensitive values. Keep payload diagnostics separate and subject to existing redaction
|
||||
|
||||
A returned stream retains its route span until exhaustion, error, or drop, with exactly one terminal outcome. Builder construction does not start a trace. Never hold a span entry guard across an await. Diagnostic tracing remains separate from lifecycle callbacks and `CustomLogger` dispatch
|
||||
|
||||
Use `litellm_tracing::sink_layer` to compose a sink with other subscriber layers. It inherits span fields into events and emits span-close summaries with elapsed time. Test observable records, concurrent isolation, dynamic filtering, sensitive-field exclusion, and stream cancellation when changing this behavior
|
||||
|
||||
Consult the [tracing API](https://docs.rs/tracing/latest/tracing/) and [subscriber layers](https://docs.rs/tracing-subscriber/latest/tracing_subscriber/layer/index.html) for implementation details
|
||||
|
|
@ -1,5 +1,7 @@
|
|||
# Rust workspace rules
|
||||
|
||||
For diagnostic tracing changes, follow [.agents/skills/rust-tracing/SKILL.md](.agents/skills/rust-tracing/SKILL.md)
|
||||
|
||||
## Test placement
|
||||
|
||||
- Never create a `tests.rs` (or `test.rs`) file under `src/`, and never `#[path = "tests.rs"] mod tests;`
|
||||
|
|
@ -9,10 +11,16 @@
|
|||
- A test for another crate's item belongs in that crate, not in a downstream one
|
||||
- Never set `autotests = false` or hand-list `[[test]]` targets; every file directly under `tests/` is discovered by cargo, and a shared helper goes in `tests/<name>/mod.rs` or `tests/<subject>/support.rs` so it is not picked up as a test crate of its own
|
||||
|
||||
## Test fixtures and cases
|
||||
|
||||
Use [`#[rstest]`](https://docs.rs/rstest/latest/rstest/attr.rstest.html) for new and updated tests and [`#[fixture]`](https://docs.rs/rstest/latest/rstest/attr.fixture.html) for reusable setup, injected through typed test arguments. Express input variations as named `#[case::name(...)]` cases instead of loops or duplicated tests so each failure identifies its case. Keep behavior assertions in the test body and fixtures focused on setup. Use the workspace `rstest` dependency
|
||||
|
||||
## Error definitions
|
||||
|
||||
- A crate's errors live in `src/error.rs`, defined with `thiserror`, and re-exported from `lib.rs`
|
||||
- Default to one top-level `Error` enum per crate, with one variant per failure mode and a `#[error(...)]` message on each
|
||||
- Put message templates in the variant's `#[error(...)]` declaration. Callers pass only the small typed arguments needed to fill them, never `Error::Variant(format!(...))` or a preformatted message. Keep the smallest set of neutral variants that callers need to distinguish; different wording or providers do not justify new variants
|
||||
- Default to one top-level `Error` enum per crate, with one variant per failure mode and a `#[error(...)]` message on each. A failure mode is something a caller handles differently (phase, status code, retry, a message Python parity pins exactly); failures no caller tells apart share one variant and differ only in its message
|
||||
- Keep shared error enums minimal and provider-neutral. Provider names, credential types, configuration fields, and setup guidance belong in caller-supplied data, not dedicated variants or hardcoded shared messages. Reuse a variant for the same failure mode across providers, such as `MissingApiBase { provider: "Azure", guidance: "..." }`. An exact parity message does not justify a provider-specific variant when caller-supplied context can preserve it
|
||||
- Wrap a lower-level error as a variant with `#[from]` or `#[source]` instead of flattening it to a string
|
||||
- Exception: split into separate types when different functions fail in disjoint ways, especially when different callers see them. A shared enum would force every caller to match variants its function can never return
|
||||
- Name a split type after what went wrong (a unit struct is fine for a single failure mode), not after the function that returns it
|
||||
|
|
|
|||
1943
litellm-rust/Cargo.lock
generated
1943
litellm-rust/Cargo.lock
generated
File diff suppressed because it is too large
Load diff
|
|
@ -9,11 +9,20 @@ license = "MIT"
|
|||
repository = "https://github.com/BerriAI/litellm"
|
||||
|
||||
[workspace.dependencies]
|
||||
litellm-config = { path = "crates/config" }
|
||||
litellm-router = { path = "crates/router" }
|
||||
litellm-tracing = { path = "crates/tracing" }
|
||||
tracing = "0.1"
|
||||
litellm-core = { path = "crates/core" }
|
||||
litellm-gateway-mcp = { path = "crates/gateway-mcp" }
|
||||
litellm-gateway = { path = "crates/gateway" }
|
||||
litellm-gateway-inference = { path = "crates/gateway-inference" }
|
||||
litellm-gateway-auth = { path = "crates/gateway-auth" }
|
||||
litellm-gateway-management = { path = "crates/gateway-management" }
|
||||
litellm-gateway-ui = { path = "crates/gateway-ui" }
|
||||
litellm-coroutine = { path = "crates/coroutine" }
|
||||
litellm-host = { path = "crates/host" }
|
||||
litellm-host-http = { path = "crates/host-http" }
|
||||
litellm-host-native = { path = "crates/host-native" }
|
||||
litellm-callbacks-legacy-python = { path = "crates/callbacks-legacy-python" }
|
||||
litellm-framing = { path = "crates/framer" }
|
||||
litellm-auth = { path = "crates/auth" }
|
||||
|
|
@ -32,6 +41,8 @@ litellm-http = { path = "crates/http" }
|
|||
litellm-llms = { path = "crates/llms" }
|
||||
litellm-types = { path = "crates/types" }
|
||||
litellm-core-utils = { path = "crates/core-utils" }
|
||||
litellm-db = { path = "crates/db" }
|
||||
litellm-db-testing = { path = "crates/db-testing" }
|
||||
litellm-cache = { path = "crates/cache" }
|
||||
litellm-cache-azure-blob = { path = "crates/cache-azure-blob" }
|
||||
litellm-cache-memory = { path = "crates/cache-memory" }
|
||||
|
|
@ -48,7 +59,12 @@ litellm-token-counter-fast = { path = "crates/token-counter-fast" }
|
|||
litellm-token-counter-huggingface = { path = "crates/token-counter-huggingface" }
|
||||
litellm-token-counter-tiktoken = { path = "crates/token-counter-tiktoken" }
|
||||
litellm-host-python = { path = "crates/host-python" }
|
||||
litellm-python-compat = { path = "crates/python-compat" }
|
||||
|
||||
tracing = "0.1"
|
||||
axum = { version = "0.8.9", default-features = false, features = ["http1", "tokio", "multipart"] }
|
||||
axum-login = "0.18.0"
|
||||
tower-sessions = { version = "0.14.0", features = ["memory-store"] }
|
||||
bytes = "1"
|
||||
http = "1"
|
||||
google-cloud-auth = { version = "1.16.0", default-features = false }
|
||||
|
|
@ -57,8 +73,8 @@ hyper-util = { version = "0.1.20", default-features = false, features = ["client
|
|||
proptest = "1.7.0"
|
||||
pyo3 = "0.29.2"
|
||||
pyo3-async-runtimes = { version = "0.29.0", features = ["tokio-runtime"] }
|
||||
pythonize = "0.29.0"
|
||||
rand = "0.8"
|
||||
schemars = "1"
|
||||
reqwest = { version = "0.12", default-features = false, features = ["json", "multipart", "rustls-tls", "http2", "stream"] }
|
||||
qdrant-client = { version = "1.19.0", default-features = false }
|
||||
uuid = { version = "1", features = ["v4"] }
|
||||
|
|
@ -73,6 +89,7 @@ serde = { version = "1.0", features = ["derive"] }
|
|||
serde_json = { version = "1.0", features = ["float_roundtrip"] }
|
||||
serde_with = { version = "=3.16.1", default-features = false, features = ["std", "macros"] }
|
||||
sha2 = "0.10"
|
||||
sqlx = { version = "0.9.0", default-features = false, features = ["json", "macros", "postgres", "runtime-tokio", "chrono", "tls-rustls-ring-native-roots"] }
|
||||
subtle = "2"
|
||||
thiserror = "2.0"
|
||||
tokenizers = { version = "0.23.1", default-features = false, features = ["onig"] }
|
||||
|
|
@ -81,6 +98,12 @@ tokio = { version = "1", features = ["rt-multi-thread", "macros", "time", "net"]
|
|||
tokio-tungstenite = { version = "0.24", default-features = false, features = ["connect", "rustls-tls-native-roots"] }
|
||||
futures-util = { version = "0.3", default-features = false, features = ["sink", "std"] }
|
||||
base64 = "0.22"
|
||||
flate2 = "1"
|
||||
semver = "1"
|
||||
tar = "0.4"
|
||||
target-lexicon = "0.13.5"
|
||||
tempfile = "3"
|
||||
zip = { version = "2", default-features = false, features = ["deflate"] }
|
||||
moka = { version = "0.12.16", features = ["future"] }
|
||||
strum = { version = "0.28.0", features = ["derive"] }
|
||||
url = "2.5.8"
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
# The Tokio runtime is reached only through `host-python/src/execution.rs`, whose fork gate
|
||||
# The Tokio runtime is reached only through `host-python/src/runtime.rs`, whose fork gate
|
||||
# must see every entry. Going around it makes a fork-after-use hang instead of raising.
|
||||
disallowed-methods = [
|
||||
{ path = "pyo3_async_runtimes::tokio::get_runtime", reason = "use litellm_host_python::run_sync / run_sync_value" },
|
||||
|
|
@ -7,4 +7,23 @@ disallowed-methods = [
|
|||
{ path = "pyo3_async_runtimes::tokio::local_future_into_py", reason = "use litellm_host_python::run_async / run_async_value" },
|
||||
{ path = "pyo3_async_runtimes::tokio::run", reason = "use litellm_host_python::run_sync / run_sync_value" },
|
||||
{ path = "pyo3_async_runtimes::tokio::run_until_complete", reason = "use litellm_host_python::run_sync / run_sync_value" },
|
||||
{ path = "reqwest::Client::new", reason = "take litellm_http::Client from HttpClientPool" },
|
||||
{ path = "reqwest::Client::builder", reason = "HttpClientConfig owns client construction" },
|
||||
{ path = "reqwest::ClientBuilder::danger_accept_invalid_certs", reason = "set HttpClientConfig::verify instead" },
|
||||
{ path = "reqwest::ClientBuilder::identity", reason = "set HttpClientConfig::client_certificate instead" },
|
||||
{ path = "reqwest::ClientBuilder::use_preconfigured_tls", reason = "HttpClientConfig owns the TLS configuration" },
|
||||
{ path = "sqlx::query", reason = "use sqlx::query! or query_file! so the SQL is checked against the migrated schema" },
|
||||
{ path = "sqlx::query_as", reason = "use sqlx::query_as! or query_file_as! so the SQL is checked against the migrated schema" },
|
||||
{ path = "sqlx::query_scalar", reason = "use sqlx::query_scalar! so the SQL is checked against the migrated schema" },
|
||||
{ path = "sqlx::query_with", reason = "use sqlx::query! or query_file! so the SQL is checked against the migrated schema" },
|
||||
{ path = "sqlx::query_as_with", reason = "use sqlx::query_as! or query_file_as! so the SQL is checked against the migrated schema" },
|
||||
{ path = "sqlx::query_scalar_with", reason = "use sqlx::query_scalar! so the SQL is checked against the migrated schema" },
|
||||
{ path = "sqlx::raw_sql", reason = "raw_sql is unchecked; use the checked query macros" },
|
||||
]
|
||||
|
||||
# Every outbound client comes from litellm_http::HttpClientPool so it honors the host's TLS,
|
||||
# proxy and timeout settings. Only crates/http builds one.
|
||||
disallowed-types = [
|
||||
{ path = "reqwest::Client", reason = "take litellm_http::Client from HttpClientPool; only crates/http builds one" },
|
||||
{ path = "reqwest::ClientBuilder", reason = "HttpClientConfig owns client construction" },
|
||||
]
|
||||
|
|
|
|||
|
|
@ -22,5 +22,7 @@ aws-types = "1.4.0"
|
|||
aws-smithy-runtime-api = "1.13.0"
|
||||
|
||||
[dev-dependencies]
|
||||
rstest.workspace = true
|
||||
litellm-http = { workspace = true, features = ["test-support"] }
|
||||
reqwest.workspace = true
|
||||
tokio.workspace = true
|
||||
|
|
|
|||
|
|
@ -1,5 +1,4 @@
|
|||
use std::collections::BTreeMap;
|
||||
use std::sync::OnceLock;
|
||||
use std::time::Duration;
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
|
|
@ -26,8 +25,26 @@ use super::constants::{
|
|||
const STATIC_CREDENTIALS_TTL: Duration = Duration::from_secs(3600 - 60);
|
||||
const AMBIENT_CREDENTIALS_TTL: Duration = Duration::from_secs(600);
|
||||
|
||||
static STATIC_CREDENTIALS_CACHE: OnceLock<Cache<String, Credentials>> = OnceLock::new();
|
||||
static AMBIENT_CREDENTIALS_CACHE: OnceLock<Cache<String, Credentials>> = OnceLock::new();
|
||||
#[derive(Clone)]
|
||||
pub struct AwsAuthService {
|
||||
static_credentials: Cache<String, Credentials>,
|
||||
ambient_credentials: Cache<String, Credentials>,
|
||||
}
|
||||
|
||||
impl Default for AwsAuthService {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
static_credentials: Cache::builder()
|
||||
.max_capacity(200)
|
||||
.time_to_live(STATIC_CREDENTIALS_TTL)
|
||||
.build(),
|
||||
ambient_credentials: Cache::builder()
|
||||
.max_capacity(200)
|
||||
.time_to_live(AMBIENT_CREDENTIALS_TTL)
|
||||
.build(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn credential_cache_ttl(flow: &AwsAuthFlow) -> Option<Duration> {
|
||||
match flow {
|
||||
|
|
@ -108,35 +125,19 @@ fn cache_key(config: &AwsAuthConfig, flow: &AwsAuthFlow) -> String {
|
|||
format!("{:x}", hasher.finalize())
|
||||
}
|
||||
|
||||
fn static_credentials_cache() -> &'static Cache<String, Credentials> {
|
||||
STATIC_CREDENTIALS_CACHE.get_or_init(|| {
|
||||
Cache::builder()
|
||||
.max_capacity(200)
|
||||
.time_to_live(STATIC_CREDENTIALS_TTL)
|
||||
.build()
|
||||
})
|
||||
}
|
||||
impl AwsAuthService {
|
||||
fn get_cached_credentials(&self, key: &str) -> Option<Credentials> {
|
||||
self.static_credentials
|
||||
.get(key)
|
||||
.or_else(|| self.ambient_credentials.get(key))
|
||||
}
|
||||
|
||||
fn ambient_credentials_cache() -> &'static Cache<String, Credentials> {
|
||||
AMBIENT_CREDENTIALS_CACHE.get_or_init(|| {
|
||||
Cache::builder()
|
||||
.max_capacity(200)
|
||||
.time_to_live(AMBIENT_CREDENTIALS_TTL)
|
||||
.build()
|
||||
})
|
||||
}
|
||||
|
||||
fn get_cached_credentials(key: &str) -> Option<Credentials> {
|
||||
static_credentials_cache()
|
||||
.get(key)
|
||||
.or_else(|| ambient_credentials_cache().get(key))
|
||||
}
|
||||
|
||||
fn set_cached_credentials(key: String, credentials: Credentials, ttl: Duration) {
|
||||
if ttl == STATIC_CREDENTIALS_TTL {
|
||||
static_credentials_cache().insert(key, credentials);
|
||||
} else {
|
||||
ambient_credentials_cache().insert(key, credentials);
|
||||
fn set_cached_credentials(&self, key: String, credentials: Credentials, ttl: Duration) {
|
||||
if ttl == STATIC_CREDENTIALS_TTL {
|
||||
self.static_credentials.insert(key, credentials);
|
||||
} else {
|
||||
self.ambient_credentials.insert(key, credentials);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -214,66 +215,157 @@ pub fn classify_auth(
|
|||
AwsAuthFlow::DefaultChain
|
||||
}
|
||||
|
||||
pub async fn resolve_credentials(
|
||||
config: AwsAuthConfig,
|
||||
env_lookup: &(dyn Fn(&str) -> Option<String> + Sync),
|
||||
) -> Result<Credentials, Error> {
|
||||
let resolved = config.clone().with_environment(env_lookup);
|
||||
let flow = classify_auth(config, env_lookup);
|
||||
match flow {
|
||||
AwsAuthFlow::SessionToken {
|
||||
access_key_id,
|
||||
secret_access_key,
|
||||
session_token,
|
||||
} => Ok(Credentials::new(
|
||||
access_key_id,
|
||||
secret_access_key,
|
||||
Some(session_token),
|
||||
None,
|
||||
"litellm-static-session",
|
||||
)),
|
||||
AwsAuthFlow::StaticKeys {
|
||||
access_key_id,
|
||||
secret_access_key,
|
||||
region_name,
|
||||
} => {
|
||||
let flow = AwsAuthFlow::StaticKeys {
|
||||
access_key_id: access_key_id.clone(),
|
||||
secret_access_key: secret_access_key.clone(),
|
||||
region_name,
|
||||
};
|
||||
let key = cache_key(&resolved, &flow);
|
||||
if let Some(credentials) = get_cached_credentials(&key) {
|
||||
return Ok(credentials);
|
||||
}
|
||||
let credentials = Credentials::new(
|
||||
impl AwsAuthService {
|
||||
pub async fn resolve_credentials(
|
||||
&self,
|
||||
config: AwsAuthConfig,
|
||||
env_lookup: &(dyn Fn(&str) -> Option<String> + Sync),
|
||||
) -> Result<Credentials, Error> {
|
||||
let resolved = config.clone().with_environment(env_lookup);
|
||||
let flow = classify_auth(config, env_lookup);
|
||||
match flow {
|
||||
AwsAuthFlow::SessionToken {
|
||||
access_key_id,
|
||||
secret_access_key,
|
||||
session_token,
|
||||
} => Ok(Credentials::new(
|
||||
access_key_id,
|
||||
secret_access_key,
|
||||
Some(session_token),
|
||||
None,
|
||||
None,
|
||||
"litellm-static",
|
||||
);
|
||||
set_cached_credentials(
|
||||
key,
|
||||
credentials.clone(),
|
||||
credential_cache_ttl(&flow).unwrap_or(STATIC_CREDENTIALS_TTL),
|
||||
);
|
||||
Ok(credentials)
|
||||
}
|
||||
AwsAuthFlow::Profile { name } => {
|
||||
let provider = aws_config::profile::ProfileFileCredentialsProvider::builder()
|
||||
.profile_name(name)
|
||||
.build();
|
||||
provider
|
||||
.provide_credentials()
|
||||
.await
|
||||
.map_err(|error| Error::AwsProfile(error.to_string()))
|
||||
}
|
||||
AwsAuthFlow::AssumeRole { role, session_name } => {
|
||||
if is_already_running_as_role(&role, &resolved).await? {
|
||||
let ambient_flow = AwsAuthFlow::DefaultChain;
|
||||
let key = cache_key(&resolved, &ambient_flow);
|
||||
if let Some(credentials) = get_cached_credentials(&key) {
|
||||
"litellm-static-session",
|
||||
)),
|
||||
AwsAuthFlow::StaticKeys {
|
||||
access_key_id,
|
||||
secret_access_key,
|
||||
region_name,
|
||||
} => {
|
||||
let flow = AwsAuthFlow::StaticKeys {
|
||||
access_key_id: access_key_id.clone(),
|
||||
secret_access_key: secret_access_key.clone(),
|
||||
region_name,
|
||||
};
|
||||
let key = cache_key(&resolved, &flow);
|
||||
if let Some(credentials) = self.get_cached_credentials(&key) {
|
||||
return Ok(credentials);
|
||||
}
|
||||
let credentials = Credentials::new(
|
||||
access_key_id,
|
||||
secret_access_key,
|
||||
None,
|
||||
None,
|
||||
"litellm-static",
|
||||
);
|
||||
self.set_cached_credentials(
|
||||
key,
|
||||
credentials.clone(),
|
||||
credential_cache_ttl(&flow).unwrap_or(STATIC_CREDENTIALS_TTL),
|
||||
);
|
||||
Ok(credentials)
|
||||
}
|
||||
AwsAuthFlow::Profile { name } => {
|
||||
let provider = aws_config::profile::ProfileFileCredentialsProvider::builder()
|
||||
.profile_name(name)
|
||||
.build();
|
||||
provider
|
||||
.provide_credentials()
|
||||
.await
|
||||
.map_err(|error| Error::AwsProfile(error.to_string()))
|
||||
}
|
||||
AwsAuthFlow::AssumeRole { role, session_name } => {
|
||||
if is_already_running_as_role(&role, &resolved).await? {
|
||||
let ambient_flow = AwsAuthFlow::DefaultChain;
|
||||
let key = cache_key(&resolved, &ambient_flow);
|
||||
if let Some(credentials) = self.get_cached_credentials(&key) {
|
||||
return Ok(credentials);
|
||||
}
|
||||
let provider =
|
||||
aws_config::default_provider::credentials::DefaultCredentialsChain::builder()
|
||||
.build()
|
||||
.await;
|
||||
let credentials = provider
|
||||
.provide_credentials()
|
||||
.await
|
||||
.map_err(|error| Error::AwsDefaultChain(error.to_string()))?;
|
||||
self.set_cached_credentials(
|
||||
key,
|
||||
credentials.clone(),
|
||||
credential_cache_ttl(&ambient_flow).unwrap_or(AMBIENT_CREDENTIALS_TTL),
|
||||
);
|
||||
return Ok(credentials);
|
||||
}
|
||||
let mut loader = aws_config::defaults(aws_config::BehaviorVersion::latest());
|
||||
if let Some(region) = resolved.region_name.clone() {
|
||||
loader = loader.region(aws_types::region::Region::new(region));
|
||||
}
|
||||
if let Some(endpoint) = resolved.sts_endpoint.clone() {
|
||||
loader = loader.endpoint_url(endpoint);
|
||||
}
|
||||
if let (Some(access_key_id), Some(secret_access_key)) =
|
||||
(resolved.access_key_id, resolved.secret_access_key)
|
||||
{
|
||||
loader = loader.credentials_provider(Credentials::new(
|
||||
access_key_id,
|
||||
secret_access_key,
|
||||
resolved.session_token,
|
||||
None,
|
||||
"litellm-role-source",
|
||||
));
|
||||
}
|
||||
let sdk_config = loader.load().await;
|
||||
let builder = aws_config::sts::AssumeRoleProvider::builder(role);
|
||||
let builder = match session_name {
|
||||
Some(name) => builder.session_name(name),
|
||||
None => builder.session_name(default_session_name()),
|
||||
};
|
||||
let builder = match resolved.external_id {
|
||||
Some(id) => builder.external_id(id),
|
||||
None => builder,
|
||||
};
|
||||
let provider = builder.configure(&sdk_config).build().await;
|
||||
provider
|
||||
.provide_credentials()
|
||||
.await
|
||||
.map_err(|error| Error::AwsAssumeRole(error.to_string()))
|
||||
}
|
||||
AwsAuthFlow::WebIdentity {
|
||||
token,
|
||||
role,
|
||||
session_name,
|
||||
} => {
|
||||
let mut loader = aws_config::defaults(aws_config::BehaviorVersion::latest());
|
||||
if let Some(region) = resolved.region_name {
|
||||
loader = loader.region(aws_types::region::Region::new(region));
|
||||
}
|
||||
if let Some(endpoint) = resolved.sts_endpoint {
|
||||
loader = loader.endpoint_url(endpoint);
|
||||
}
|
||||
let sdk_config = loader.load().await;
|
||||
let client = aws_sdk_sts::Client::new(&sdk_config);
|
||||
let response = client
|
||||
.assume_role_with_web_identity()
|
||||
.role_arn(role)
|
||||
.role_session_name(session_name)
|
||||
.web_identity_token(token)
|
||||
.send()
|
||||
.await
|
||||
.map_err(|error| Error::AwsWebIdentity(error.to_string()))?;
|
||||
let credentials = response
|
||||
.credentials()
|
||||
.ok_or(Error::AwsMissingWebIdentityCredentials)?;
|
||||
let expiration = SystemTime::try_from(*credentials.expiration())
|
||||
.map_err(|error| Error::AwsWebIdentityExpiration(error.to_string()))?;
|
||||
Ok(Credentials::new(
|
||||
credentials.access_key_id(),
|
||||
credentials.secret_access_key(),
|
||||
Some(credentials.session_token().to_string()),
|
||||
Some(expiration),
|
||||
"litellm-web-identity",
|
||||
))
|
||||
}
|
||||
AwsAuthFlow::DefaultChain => {
|
||||
let key = cache_key(&resolved, &AwsAuthFlow::DefaultChain);
|
||||
if let Some(credentials) = self.get_cached_credentials(&key) {
|
||||
return Ok(credentials);
|
||||
}
|
||||
let provider =
|
||||
|
|
@ -284,101 +376,14 @@ pub async fn resolve_credentials(
|
|||
.provide_credentials()
|
||||
.await
|
||||
.map_err(|error| Error::AwsDefaultChain(error.to_string()))?;
|
||||
set_cached_credentials(
|
||||
self.set_cached_credentials(
|
||||
key,
|
||||
credentials.clone(),
|
||||
credential_cache_ttl(&ambient_flow).unwrap_or(AMBIENT_CREDENTIALS_TTL),
|
||||
credential_cache_ttl(&AwsAuthFlow::DefaultChain)
|
||||
.unwrap_or(AMBIENT_CREDENTIALS_TTL),
|
||||
);
|
||||
return Ok(credentials);
|
||||
Ok(credentials)
|
||||
}
|
||||
let mut loader = aws_config::defaults(aws_config::BehaviorVersion::latest());
|
||||
if let Some(region) = resolved.region_name.clone() {
|
||||
loader = loader.region(aws_types::region::Region::new(region));
|
||||
}
|
||||
if let Some(endpoint) = resolved.sts_endpoint.clone() {
|
||||
loader = loader.endpoint_url(endpoint);
|
||||
}
|
||||
if let (Some(access_key_id), Some(secret_access_key)) =
|
||||
(resolved.access_key_id, resolved.secret_access_key)
|
||||
{
|
||||
loader = loader.credentials_provider(Credentials::new(
|
||||
access_key_id,
|
||||
secret_access_key,
|
||||
resolved.session_token,
|
||||
None,
|
||||
"litellm-role-source",
|
||||
));
|
||||
}
|
||||
let sdk_config = loader.load().await;
|
||||
let builder = aws_config::sts::AssumeRoleProvider::builder(role);
|
||||
let builder = match session_name {
|
||||
Some(name) => builder.session_name(name),
|
||||
None => builder.session_name(default_session_name()),
|
||||
};
|
||||
let builder = match resolved.external_id {
|
||||
Some(id) => builder.external_id(id),
|
||||
None => builder,
|
||||
};
|
||||
let provider = builder.configure(&sdk_config).build().await;
|
||||
provider
|
||||
.provide_credentials()
|
||||
.await
|
||||
.map_err(|error| Error::AwsAssumeRole(error.to_string()))
|
||||
}
|
||||
AwsAuthFlow::WebIdentity {
|
||||
token,
|
||||
role,
|
||||
session_name,
|
||||
} => {
|
||||
let mut loader = aws_config::defaults(aws_config::BehaviorVersion::latest());
|
||||
if let Some(region) = resolved.region_name {
|
||||
loader = loader.region(aws_types::region::Region::new(region));
|
||||
}
|
||||
if let Some(endpoint) = resolved.sts_endpoint {
|
||||
loader = loader.endpoint_url(endpoint);
|
||||
}
|
||||
let sdk_config = loader.load().await;
|
||||
let client = aws_sdk_sts::Client::new(&sdk_config);
|
||||
let response = client
|
||||
.assume_role_with_web_identity()
|
||||
.role_arn(role)
|
||||
.role_session_name(session_name)
|
||||
.web_identity_token(token)
|
||||
.send()
|
||||
.await
|
||||
.map_err(|error| Error::AwsWebIdentity(error.to_string()))?;
|
||||
let credentials = response
|
||||
.credentials()
|
||||
.ok_or(Error::AwsMissingWebIdentityCredentials)?;
|
||||
let expiration = SystemTime::try_from(*credentials.expiration())
|
||||
.map_err(|error| Error::AwsWebIdentityExpiration(error.to_string()))?;
|
||||
Ok(Credentials::new(
|
||||
credentials.access_key_id(),
|
||||
credentials.secret_access_key(),
|
||||
Some(credentials.session_token().to_string()),
|
||||
Some(expiration),
|
||||
"litellm-web-identity",
|
||||
))
|
||||
}
|
||||
AwsAuthFlow::DefaultChain => {
|
||||
let key = cache_key(&resolved, &AwsAuthFlow::DefaultChain);
|
||||
if let Some(credentials) = get_cached_credentials(&key) {
|
||||
return Ok(credentials);
|
||||
}
|
||||
let provider =
|
||||
aws_config::default_provider::credentials::DefaultCredentialsChain::builder()
|
||||
.build()
|
||||
.await;
|
||||
let credentials = provider
|
||||
.provide_credentials()
|
||||
.await
|
||||
.map_err(|error| Error::AwsDefaultChain(error.to_string()))?;
|
||||
set_cached_credentials(
|
||||
key,
|
||||
credentials.clone(),
|
||||
credential_cache_ttl(&AwsAuthFlow::DefaultChain).unwrap_or(AMBIENT_CREDENTIALS_TTL),
|
||||
);
|
||||
Ok(credentials)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -585,6 +590,37 @@ pub fn aws_auth_config(
|
|||
}
|
||||
}
|
||||
|
||||
/// Where the credentials that sign a request come from, decided when the request is
|
||||
/// prepared and resolved when it is sent.
|
||||
#[derive(Clone, Debug, PartialEq)]
|
||||
pub enum AwsCredentialSource {
|
||||
HostSupplied(Credentials),
|
||||
Chain(AwsAuthConfig),
|
||||
}
|
||||
|
||||
impl AwsCredentialSource {
|
||||
pub fn from_params(
|
||||
optional_params: &Map<String, Value>,
|
||||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> Self {
|
||||
match host_supplied_credentials(optional_params) {
|
||||
Some(credentials) => Self::HostSupplied(credentials),
|
||||
None => Self::Chain(aws_auth_config(optional_params, env_lookup)),
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn resolve(
|
||||
self,
|
||||
auth: &AwsAuthService,
|
||||
env_lookup: &(dyn Fn(&str) -> Option<String> + Sync),
|
||||
) -> Result<Credentials, Error> {
|
||||
match self {
|
||||
Self::HostSupplied(credentials) => Ok(credentials),
|
||||
Self::Chain(config) => auth.resolve_credentials(config, env_lookup).await,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Credentials a host resolved through its own chain and handed down verbatim.
|
||||
///
|
||||
/// A host with its own resolution (LiteLLM's Python `BaseAWSLLM`, which reads
|
||||
|
|
@ -747,17 +783,18 @@ mod tests {
|
|||
|
||||
#[tokio::test]
|
||||
async fn static_credentials_do_not_use_network() {
|
||||
let credentials = resolve_credentials(
|
||||
AwsAuthConfig {
|
||||
access_key_id: Some("ak".into()),
|
||||
secret_access_key: Some("sk".into()),
|
||||
region_name: Some("us-east-1".into()),
|
||||
..Default::default()
|
||||
},
|
||||
&no_env,
|
||||
)
|
||||
.await
|
||||
.expect("static credentials");
|
||||
let credentials = AwsAuthService::default()
|
||||
.resolve_credentials(
|
||||
AwsAuthConfig {
|
||||
access_key_id: Some("ak".into()),
|
||||
secret_access_key: Some("sk".into()),
|
||||
region_name: Some("us-east-1".into()),
|
||||
..Default::default()
|
||||
},
|
||||
&no_env,
|
||||
)
|
||||
.await
|
||||
.expect("static credentials");
|
||||
assert_eq!(credentials.access_key_id(), "ak");
|
||||
assert_eq!(credentials.session_token(), None);
|
||||
}
|
||||
|
|
@ -807,17 +844,67 @@ mod tests {
|
|||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[rstest::rstest]
|
||||
fn cache_round_trip_preserves_credentials() {
|
||||
let auth = AwsAuthService::default();
|
||||
let key = format!("cache-test-{}", std::process::id());
|
||||
let credentials = Credentials::new("cache-ak", "cache-sk", None, None, "test");
|
||||
set_cached_credentials(key.clone(), credentials.clone(), STATIC_CREDENTIALS_TTL);
|
||||
auth.set_cached_credentials(key.clone(), credentials.clone(), STATIC_CREDENTIALS_TTL);
|
||||
assert_eq!(
|
||||
get_cached_credentials(&key).map(|value| value.access_key_id().to_string()),
|
||||
auth.get_cached_credentials(&key)
|
||||
.map(|value| value.access_key_id().to_string()),
|
||||
Some("cache-ak".to_string())
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest::rstest]
|
||||
#[tokio::test]
|
||||
async fn cloned_services_reuse_credentials_but_independent_services_do_not() {
|
||||
let auth = AwsAuthService::default();
|
||||
let config = AwsAuthConfig {
|
||||
access_key_id: Some("configured-key".into()),
|
||||
secret_access_key: Some("configured-secret".into()),
|
||||
region_name: Some("us-east-1".into()),
|
||||
..AwsAuthConfig::default()
|
||||
};
|
||||
let flow = classify_auth(config.clone(), &no_env);
|
||||
let cached = Credentials::new("cached-key", "cached-secret", None, None, "test");
|
||||
auth.set_cached_credentials(
|
||||
cache_key(&config, &flow),
|
||||
cached.clone(),
|
||||
STATIC_CREDENTIALS_TTL,
|
||||
);
|
||||
|
||||
let reused = auth
|
||||
.clone()
|
||||
.resolve_credentials(config.clone(), &no_env)
|
||||
.await
|
||||
.unwrap();
|
||||
let independent = AwsAuthService::default()
|
||||
.resolve_credentials(config.clone(), &no_env)
|
||||
.await
|
||||
.unwrap();
|
||||
let different = AwsAuthConfig {
|
||||
access_key_id: Some("different-key".into()),
|
||||
..config.clone()
|
||||
};
|
||||
let other_identity = auth
|
||||
.resolve_credentials(different.clone(), &no_env)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(reused.access_key_id(), cached.access_key_id());
|
||||
assert_eq!(reused.secret_access_key(), cached.secret_access_key());
|
||||
assert_eq!(
|
||||
Some(independent.access_key_id()),
|
||||
config.access_key_id.as_deref()
|
||||
);
|
||||
assert_eq!(
|
||||
Some(other_identity.access_key_id()),
|
||||
different.access_key_id.as_deref()
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn same_role_comparison_matches_partition_account_and_role() {
|
||||
assert!(same_role_arns(
|
||||
|
|
@ -952,17 +1039,18 @@ mod tests {
|
|||
let body = br#"{"anthropic_version":"bedrock-2023-05-31","max_tokens":1,"messages":[{"role":"user","content":[{"type":"text","text":"ping"}]}]}"#.to_vec();
|
||||
let headers =
|
||||
BTreeMap::from([("Content-Type".to_string(), "application/json".to_string())]);
|
||||
let credentials = resolve_credentials(
|
||||
AwsAuthConfig {
|
||||
access_key_id: Some(access_key_id),
|
||||
secret_access_key: Some(secret_access_key),
|
||||
region_name: Some("us-west-2".to_string()),
|
||||
..Default::default()
|
||||
},
|
||||
&no_env,
|
||||
)
|
||||
.await?;
|
||||
let client = reqwest::Client::new();
|
||||
let credentials = AwsAuthService::default()
|
||||
.resolve_credentials(
|
||||
AwsAuthConfig {
|
||||
access_key_id: Some(access_key_id),
|
||||
secret_access_key: Some(secret_access_key),
|
||||
region_name: Some("us-west-2".to_string()),
|
||||
..Default::default()
|
||||
},
|
||||
&no_env,
|
||||
)
|
||||
.await?;
|
||||
let client = litellm_http::Client::plain_for_test();
|
||||
let mut failures = Vec::new();
|
||||
|
||||
for region in ["us-west-2", "us-east-1"] {
|
||||
|
|
|
|||
|
|
@ -1,13 +1,11 @@
|
|||
use std::{collections::BTreeMap, time::SystemTime};
|
||||
|
||||
use crate::{
|
||||
AwsAuthService, AwsCredentialSource, Error, aws_signature_headers, is_sigv4_computed_header,
|
||||
sign_post,
|
||||
};
|
||||
use aws_credential_types::Credentials;
|
||||
use litellm_http::outbound::{RequestSigner, UnsignedRequest};
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
use crate::{
|
||||
Error, aws_auth_config, aws_signature_headers, host_supplied_credentials,
|
||||
is_sigv4_computed_header, resolve_credentials, sign_post,
|
||||
};
|
||||
|
||||
#[derive(Clone, Debug)]
|
||||
pub struct SigV4Signer {
|
||||
|
|
@ -32,19 +30,17 @@ impl SigV4Signer {
|
|||
}
|
||||
|
||||
pub async fn resolve(
|
||||
auth: &AwsAuthService,
|
||||
region: String,
|
||||
service: &'static str,
|
||||
optional_params: &Map<String, Value>,
|
||||
credentials: AwsCredentialSource,
|
||||
env_lookup: &(dyn Fn(&str) -> Option<String> + Sync),
|
||||
) -> Result<Self, Error> {
|
||||
let credentials = match host_supplied_credentials(optional_params) {
|
||||
Some(credentials) => credentials,
|
||||
None => {
|
||||
resolve_credentials(aws_auth_config(optional_params, env_lookup), env_lookup)
|
||||
.await?
|
||||
}
|
||||
};
|
||||
Ok(Self::new(region, service, credentials))
|
||||
Ok(Self::new(
|
||||
region,
|
||||
service,
|
||||
credentials.resolve(auth, env_lookup).await?,
|
||||
))
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -80,7 +76,7 @@ mod tests {
|
|||
use std::time::{Duration, UNIX_EPOCH};
|
||||
|
||||
use litellm_http::outbound::OutboundRequest;
|
||||
use serde_json::json;
|
||||
use serde_json::{Value, json};
|
||||
|
||||
use super::*;
|
||||
|
||||
|
|
|
|||
|
|
@ -133,7 +133,7 @@ impl NativeAzureTokenAcquirer {
|
|||
let token = credential
|
||||
.get_token(&[scope.as_str()], None)
|
||||
.await
|
||||
.map_err(|error| Error::AzureTokenAcquisition(error.to_string()))?;
|
||||
.map_err(|error| Error::CredentialAcquisition(error.to_string().into()))?;
|
||||
let expires_on = u64::try_from(token.expires_on.unix_timestamp())
|
||||
.ok()
|
||||
.map(|seconds| UNIX_EPOCH + Duration::from_secs(seconds));
|
||||
|
|
@ -250,7 +250,12 @@ fn validate_authority(request: &NativeAzureRequest) -> Result<(), Error> {
|
|||
let Some(authority) = authority else {
|
||||
return Ok(());
|
||||
};
|
||||
let url = url::Url::parse(authority.value()).map_err(|_| Error::InvalidAzureAuthority)?;
|
||||
let url = url::Url::parse(authority.value()).map_err(|_| {
|
||||
Error::InvalidConfiguration(
|
||||
"Azure authority must be an HTTPS origin without credentials, query, or fragment"
|
||||
.into(),
|
||||
)
|
||||
})?;
|
||||
if url.scheme() != "https"
|
||||
|| url.host_str().is_none()
|
||||
|| !url.username().is_empty()
|
||||
|
|
@ -259,7 +264,10 @@ fn validate_authority(request: &NativeAzureRequest) -> Result<(), Error> {
|
|||
|| url.fragment().is_some()
|
||||
|| !matches!(url.path(), "" | "/")
|
||||
{
|
||||
return Err(Error::InvalidAzureAuthority);
|
||||
return Err(Error::InvalidConfiguration(
|
||||
"Azure authority must be an HTTPS origin without credentials, query, or fragment"
|
||||
.into(),
|
||||
));
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
|
@ -368,7 +376,9 @@ fn trusted_source(sources: &[InputSource]) -> InputSource {
|
|||
}
|
||||
|
||||
fn mixed_sources<T>() -> Result<T, Error> {
|
||||
Err(Error::MixedAzureCredentialSources)
|
||||
Err(Error::InvalidConfiguration(
|
||||
"request-controlled Azure auth inputs cannot be combined with host credentials".into(),
|
||||
))
|
||||
}
|
||||
|
||||
fn build_credential(
|
||||
|
|
@ -433,7 +443,12 @@ fn build_credential(
|
|||
NativeAzureRequest::DeveloperTools { .. } => DeveloperToolsCredential::new(None)
|
||||
.map(|credential| credential as Arc<dyn TokenCredential>),
|
||||
}
|
||||
.map_err(|error| Error::AzureCredentialInitialization(error.to_string()))
|
||||
.map_err(|error| {
|
||||
Error::InvalidConfiguration(litellm_auth_types::ErrorDetail::failed(
|
||||
"Azure credential initialization",
|
||||
error,
|
||||
))
|
||||
})
|
||||
}
|
||||
|
||||
fn client_options(
|
||||
|
|
@ -638,7 +653,7 @@ mod tests {
|
|||
assert_eq!(transport.requests.lock().unwrap().len(), 6);
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[rstest::rstest]
|
||||
fn request_authority_requires_request_owned_client_secret_identity() {
|
||||
let error = ValidatedAzureRequest::new(sourced_client_secret(
|
||||
InputSource::Deployment,
|
||||
|
|
@ -647,10 +662,13 @@ mod tests {
|
|||
))
|
||||
.unwrap_err();
|
||||
|
||||
assert!(matches!(
|
||||
assert_eq!(
|
||||
error,
|
||||
litellm_auth_types::Error::MixedAzureCredentialSources
|
||||
));
|
||||
litellm_auth_types::Error::InvalidConfiguration(
|
||||
"request-controlled Azure auth inputs cannot be combined with host credentials"
|
||||
.into()
|
||||
)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
|
@ -665,24 +683,24 @@ mod tests {
|
|||
assert_eq!(request.credential_source(), InputSource::Request);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn authority_is_restricted_to_an_https_origin() {
|
||||
for authority in [
|
||||
"http://login.example",
|
||||
"https://user@login.example",
|
||||
"https://login.example/tenant",
|
||||
"https://login.example?target=other",
|
||||
] {
|
||||
let error = ValidatedAzureRequest::new(sourced_client_secret(
|
||||
InputSource::Deployment,
|
||||
InputSource::Deployment,
|
||||
authority,
|
||||
))
|
||||
.unwrap_err();
|
||||
assert!(matches!(
|
||||
error,
|
||||
litellm_auth_types::Error::InvalidAzureAuthority
|
||||
));
|
||||
}
|
||||
#[rstest::rstest]
|
||||
#[case::http("http://login.example")]
|
||||
#[case::userinfo("https://user@login.example")]
|
||||
#[case::path("https://login.example/tenant")]
|
||||
#[case::query("https://login.example?target=other")]
|
||||
fn authority_is_restricted_to_an_https_origin(#[case] authority: &str) {
|
||||
let error = ValidatedAzureRequest::new(sourced_client_secret(
|
||||
InputSource::Deployment,
|
||||
InputSource::Deployment,
|
||||
authority,
|
||||
))
|
||||
.unwrap_err();
|
||||
assert_eq!(
|
||||
error,
|
||||
litellm_auth_types::Error::InvalidConfiguration(
|
||||
"Azure authority must be an HTTPS origin without credentials, query, or fragment"
|
||||
.into()
|
||||
)
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -91,7 +91,9 @@ impl AzureAuthService {
|
|||
AzureCredentialPlan::Caller(caller) => {
|
||||
let credential = caller.acquire().await?;
|
||||
if credential.secret().expose().is_empty() {
|
||||
return Err(Error::EmptyAzureToken);
|
||||
return Err(Error::EmptyCallerCredential(
|
||||
"Azure AD token provider returned an empty token",
|
||||
));
|
||||
}
|
||||
Ok(Some(Sourced::new(credential, InputSource::Deployment)))
|
||||
}
|
||||
|
|
@ -104,7 +106,11 @@ impl AzureAuthService {
|
|||
} => {
|
||||
let assertion = resolve_reference(inputs, env_lookup, reference.value())
|
||||
.await?
|
||||
.ok_or(Error::UnresolvedOidcReference)?;
|
||||
.ok_or_else(|| {
|
||||
Error::CredentialAcquisition(
|
||||
"Azure OIDC reference did not resolve to a value".into(),
|
||||
)
|
||||
})?;
|
||||
let request = ValidatedAzureRequest::new(NativeAzureRequest::ClientAssertion {
|
||||
tenant_id,
|
||||
client_id,
|
||||
|
|
@ -167,7 +173,7 @@ pub(crate) fn select_auth_plan(
|
|||
.map(|selector| Sourced::new(selector, value.source()))
|
||||
})
|
||||
.transpose()
|
||||
.map_err(|_| Error::InvalidAzureSelector)?;
|
||||
.map_err(|_| Error::InvalidConfiguration("invalid Azure credential selector".into()))?;
|
||||
let federated_token_file = configured_string(
|
||||
&inputs.federated_token_file,
|
||||
AZURE_FEDERATED_TOKEN_FILE_ENV,
|
||||
|
|
@ -257,7 +263,9 @@ fn select_native_plan(
|
|||
let selection_source = selected.source();
|
||||
|
||||
match selected.into_value() {
|
||||
AzureCredentialType::ClientSecretCredential => Err(Error::MissingClientSecretFields),
|
||||
AzureCredentialType::ClientSecretCredential => Err(Error::InvalidConfiguration(
|
||||
"ClientSecretCredential requires tenant_id, client_id, and client_secret".into(),
|
||||
)),
|
||||
AzureCredentialType::WorkloadIdentityCredential => {
|
||||
Ok(AzureCredentialPlan::Native(ValidatedAzureRequest::new(
|
||||
workload_request(tenant_id, client_id, federated_token_file, scope, authority)?,
|
||||
|
|
@ -341,9 +349,17 @@ fn workload_request(
|
|||
authority: Option<Sourced<String>>,
|
||||
) -> Result<NativeAzureRequest, Error> {
|
||||
Ok(NativeAzureRequest::WorkloadIdentity {
|
||||
tenant_id: tenant_id.ok_or(Error::MissingWorkloadTenant)?,
|
||||
client_id: client_id.ok_or(Error::MissingWorkloadClient)?,
|
||||
token_file_path: token_file_path.ok_or(Error::MissingWorkloadTokenFile)?,
|
||||
tenant_id: tenant_id.ok_or_else(|| {
|
||||
Error::InvalidConfiguration("WorkloadIdentityCredential requires tenant_id".into())
|
||||
})?,
|
||||
client_id: client_id.ok_or_else(|| {
|
||||
Error::InvalidConfiguration("WorkloadIdentityCredential requires client_id".into())
|
||||
})?,
|
||||
token_file_path: token_file_path.ok_or_else(|| {
|
||||
Error::InvalidConfiguration(
|
||||
"WorkloadIdentityCredential requires azure_federated_token_file".into(),
|
||||
)
|
||||
})?,
|
||||
scope,
|
||||
authority,
|
||||
})
|
||||
|
|
@ -394,10 +410,11 @@ async fn resolve_reference(
|
|||
.map_or(CredentialLookup::Missing, CredentialLookup::Found),
|
||||
CredentialRef::None => return Ok(None),
|
||||
CredentialRef::File(_) | CredentialRef::Request(_) | CredentialRef::Host(_) => {
|
||||
let resolver = inputs
|
||||
.credential_resolver
|
||||
.as_ref()
|
||||
.ok_or(Error::MissingHostResolver)?;
|
||||
let resolver = inputs.credential_resolver.as_ref().ok_or_else(|| {
|
||||
Error::InvalidConfiguration(
|
||||
"credential reference requires a host credential resolver".into(),
|
||||
)
|
||||
})?;
|
||||
resolver.resolve(reference).await?
|
||||
}
|
||||
};
|
||||
|
|
@ -415,7 +432,9 @@ fn oidc_reference(
|
|||
};
|
||||
let value = token.value().expose();
|
||||
if token.source() == InputSource::Request && value.starts_with("oidc/") {
|
||||
return Err(Error::RequestAzureCredentialReference);
|
||||
return Err(Error::InvalidConfiguration(
|
||||
"request-controlled Azure credential references are not allowed".into(),
|
||||
));
|
||||
}
|
||||
if let Some(name) = value.strip_prefix("oidc/env/") {
|
||||
return non_empty_reference(name, "OIDC environment reference")
|
||||
|
|
@ -437,14 +456,20 @@ fn oidc_reference(
|
|||
)));
|
||||
}
|
||||
if value.starts_with("oidc/") {
|
||||
return Err(Error::UnsupportedOidcReference);
|
||||
return Err(Error::InvalidConfiguration(
|
||||
"unsupported OIDC reference".into(),
|
||||
));
|
||||
}
|
||||
Ok(None)
|
||||
}
|
||||
|
||||
fn non_empty_reference(value: &str, kind: &str) -> Result<String, Error> {
|
||||
if value.is_empty() {
|
||||
return Err(Error::EmptyReference(kind.to_string()));
|
||||
return Err(Error::InvalidConfiguration(
|
||||
litellm_auth_types::ErrorDetail::Empty {
|
||||
subject: kind.into(),
|
||||
},
|
||||
));
|
||||
}
|
||||
Ok(value.to_string())
|
||||
}
|
||||
|
|
@ -493,7 +518,7 @@ mod tests {
|
|||
expires_on: None,
|
||||
})
|
||||
} else {
|
||||
Err(Error::AzureTokenAcquisition(format!("{kind} failed")))
|
||||
Err(Error::CredentialAcquisition(kind.into()))
|
||||
}
|
||||
})
|
||||
}
|
||||
|
|
@ -602,7 +627,7 @@ mod tests {
|
|||
assert!(error.to_string().contains("unsupported OIDC reference"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[rstest::rstest]
|
||||
fn request_oidc_reference_is_rejected_before_lookup() {
|
||||
let params = json!({
|
||||
"azure_ad_token": "oidc/env/ASSERTION",
|
||||
|
|
@ -624,7 +649,12 @@ mod tests {
|
|||
})
|
||||
.unwrap_err();
|
||||
|
||||
assert!(matches!(error, Error::RequestAzureCredentialReference));
|
||||
assert_eq!(
|
||||
error,
|
||||
Error::InvalidConfiguration(
|
||||
"request-controlled Azure credential references are not allowed".into()
|
||||
)
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
|
|
@ -723,6 +753,7 @@ mod tests {
|
|||
assert_eq!(credential.value().secret().expose(), "caller-token");
|
||||
}
|
||||
|
||||
#[rstest::rstest]
|
||||
#[tokio::test]
|
||||
async fn empty_caller_token_is_rejected() {
|
||||
let error = AzureAuthService::default()
|
||||
|
|
@ -730,6 +761,9 @@ mod tests {
|
|||
.await
|
||||
.unwrap_err();
|
||||
|
||||
assert!(matches!(error, Error::EmptyAzureToken));
|
||||
assert_eq!(
|
||||
error,
|
||||
Error::EmptyCallerCredential("Azure AD token provider returned an empty token")
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -117,7 +117,12 @@ fn string_config(
|
|||
None => Ok(ConfigValue::Absent),
|
||||
Some(Value::Null) => Ok(ConfigValue::ExplicitNone(source)),
|
||||
Some(Value::String(value)) => Ok(ConfigValue::Value(Sourced::new(value.clone(), source))),
|
||||
Some(_) => Err(Error::InvalidFieldType(name.to_string())),
|
||||
Some(_) => Err(Error::InvalidConfiguration(
|
||||
litellm_auth_types::ErrorDetail::InvalidType {
|
||||
field: name.into(),
|
||||
expected: "a string or null",
|
||||
},
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -19,3 +19,6 @@ tokio.workspace = true
|
|||
gcp_auth = "0.12.7"
|
||||
google-cloud-auth = { workspace = true, optional = true }
|
||||
http = { workspace = true, optional = true }
|
||||
|
||||
[dev-dependencies]
|
||||
rstest.workspace = true
|
||||
|
|
|
|||
|
|
@ -131,7 +131,7 @@ impl Default for VertexAuth {
|
|||
}
|
||||
|
||||
impl VertexAuth {
|
||||
fn new(loader: Arc<dyn VertexProviderLoader>) -> Self {
|
||||
pub fn new(loader: Arc<dyn VertexProviderLoader>) -> Self {
|
||||
Self {
|
||||
providers: Cache::builder().max_capacity(64).build(),
|
||||
loader,
|
||||
|
|
@ -220,16 +220,16 @@ impl VertexAuth {
|
|||
}
|
||||
}
|
||||
|
||||
trait VertexTokenSource: Send + Sync {
|
||||
pub trait VertexTokenSource: Send + Sync {
|
||||
fn project_id(&self) -> VertexAuthFuture<'_, String>;
|
||||
fn token(&self) -> VertexAuthFuture<'_, String>;
|
||||
}
|
||||
|
||||
trait VertexProviderLoader: Send + Sync {
|
||||
pub trait VertexProviderLoader: Send + Sync {
|
||||
fn load(&self, source: CredentialSource) -> VertexAuthFuture<'_, Arc<dyn VertexTokenSource>>;
|
||||
}
|
||||
|
||||
type VertexAuthFuture<'a, T> = Pin<Box<dyn Future<Output = Result<T, Error>> + Send + 'a>>;
|
||||
pub type VertexAuthFuture<'a, T> = Pin<Box<dyn Future<Output = Result<T, Error>> + Send + 'a>>;
|
||||
|
||||
struct GcpTokenSource(Arc<dyn TokenProvider>);
|
||||
|
||||
|
|
@ -299,13 +299,13 @@ fn validate_request_credentials(configured: &str) -> Result<&str, Error> {
|
|||
.map(str::to_string)
|
||||
});
|
||||
if token_uri.as_deref() != Some(GOOGLE_OAUTH_TOKEN_ENDPOINT) {
|
||||
return Err(Error::RequestVertexTokenEndpoint);
|
||||
return Err(Error::InvalidConfiguration("request-controlled Vertex credentials must use the canonical Google OAuth token endpoint".into()));
|
||||
}
|
||||
Ok(configured)
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug)]
|
||||
enum CredentialSource {
|
||||
pub enum CredentialSource {
|
||||
Inline(SecretValue),
|
||||
Trusted(SecretValue),
|
||||
ApplicationCredentials(String),
|
||||
|
|
@ -376,10 +376,20 @@ fn optional_credentials(
|
|||
.map(SecretValue::new)
|
||||
.map(|value| Sourced::new(value, source))
|
||||
.map(Some)
|
||||
.map_err(|error| Error::InvalidFieldType(format!("{}: {error}", names[0])));
|
||||
.map_err(|error| {
|
||||
Error::InvalidConfiguration(litellm_auth_types::ErrorDetail::failed(
|
||||
"credential serialization",
|
||||
error,
|
||||
))
|
||||
});
|
||||
}
|
||||
Some(_) => {
|
||||
return Err(Error::InvalidFieldType(names[0].to_string()));
|
||||
return Err(Error::InvalidConfiguration(
|
||||
litellm_auth_types::ErrorDetail::InvalidType {
|
||||
field: names[0].into(),
|
||||
expected: "a string or null",
|
||||
},
|
||||
));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -397,7 +407,12 @@ fn optional_string(params: &Map<String, Value>, names: &[&str]) -> Result<Option
|
|||
Some(Value::String(value)) if value.trim().is_empty() => continue,
|
||||
Some(Value::String(value)) => return Ok(Some(value.clone())),
|
||||
Some(_) => {
|
||||
return Err(Error::InvalidFieldType(names[0].to_string()));
|
||||
return Err(Error::InvalidConfiguration(
|
||||
litellm_auth_types::ErrorDetail::InvalidType {
|
||||
field: names[0].into(),
|
||||
expected: "a string or null",
|
||||
},
|
||||
));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -411,7 +426,10 @@ fn non_empty_env(env_lookup: &dyn Fn(&str) -> Option<String>, name: &str) -> Opt
|
|||
}
|
||||
|
||||
fn auth_acquisition_error(error: gcp_auth::Error) -> Error {
|
||||
Error::VertexTokenAcquisition(error.to_string())
|
||||
Error::CredentialAcquisition(litellm_auth_types::ErrorDetail::failed(
|
||||
"Vertex AI credentials",
|
||||
error,
|
||||
))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
|
|
@ -612,20 +630,15 @@ mod tests {
|
|||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn request_credentials_require_canonical_token_endpoint() {
|
||||
assert!(
|
||||
validate_request_credentials(r#"{"token_uri":"https://oauth2.googleapis.com/token"}"#)
|
||||
.is_ok()
|
||||
);
|
||||
assert!(matches!(
|
||||
validate_request_credentials(r#"{"token_uri":"http://127.0.0.1/token"}"#),
|
||||
Err(Error::RequestVertexTokenEndpoint)
|
||||
));
|
||||
assert!(matches!(
|
||||
validate_request_credentials("{}"),
|
||||
Err(Error::RequestVertexTokenEndpoint)
|
||||
));
|
||||
#[rstest::rstest]
|
||||
#[case::canonical_endpoint(r#"{"token_uri":"https://oauth2.googleapis.com/token"}"#, true)]
|
||||
#[case::noncanonical_endpoint(r#"{"token_uri":"http://127.0.0.1/token"}"#, false)]
|
||||
#[case::missing_endpoint("{}", false)]
|
||||
fn request_credentials_require_canonical_token_endpoint(
|
||||
#[case] credentials: &str,
|
||||
#[case] accepted: bool,
|
||||
) {
|
||||
assert_eq!(validate_request_credentials(credentials).is_ok(), accepted);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
|
|
|
|||
|
|
@ -12,4 +12,5 @@ thiserror.workspace = true
|
|||
veil.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
rstest.workspace = true
|
||||
tokio.workspace = true
|
||||
|
|
|
|||
|
|
@ -86,7 +86,9 @@ impl CredentialPlan {
|
|||
Self::Caller(caller) => {
|
||||
let credential = caller.acquire().await?;
|
||||
if credential.secret().expose().is_empty() {
|
||||
return Err(Error::EmptyCallerCredential);
|
||||
return Err(Error::EmptyCallerCredential(
|
||||
"credential caller returned an empty credential",
|
||||
));
|
||||
}
|
||||
Ok(CredentialPlanResolution::Resolved(credential))
|
||||
}
|
||||
|
|
@ -147,10 +149,15 @@ mod tests {
|
|||
|
||||
impl CredentialResolver for FailingResolver {
|
||||
fn resolve<'a>(&'a self, _reference: &'a CredentialRef) -> CredentialLookupFuture<'a> {
|
||||
Box::pin(async { Err(Error::UnresolvedOidcReference) })
|
||||
Box::pin(async {
|
||||
Err(Error::CredentialAcquisition(
|
||||
"host credential lookup failed".into(),
|
||||
))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[rstest::rstest]
|
||||
#[tokio::test]
|
||||
async fn acquisition_failure_is_terminal() {
|
||||
let resolver = CredentialResolverHandle::new(Arc::new(FailingResolver));
|
||||
|
|
@ -161,6 +168,9 @@ mod tests {
|
|||
.await
|
||||
.expect_err("acquisition errors cannot become fallback");
|
||||
|
||||
assert_eq!(error, Error::UnresolvedOidcReference);
|
||||
assert_eq!(
|
||||
error,
|
||||
Error::CredentialAcquisition("host credential lookup failed".into())
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -2,84 +2,16 @@ use thiserror::Error as ThisError;
|
|||
|
||||
#[derive(Clone, Debug, ThisError, PartialEq, Eq)]
|
||||
pub enum Error {
|
||||
#[error("invalid authentication configuration: credential header already exists")]
|
||||
ExistingCredentialHeader,
|
||||
#[error(
|
||||
"invalid authentication configuration: credential plan is not allowed by the provider auth policy"
|
||||
)]
|
||||
DisallowedCredentialPlan,
|
||||
#[error("invalid authentication configuration: credential cannot be empty")]
|
||||
EmptyCredential,
|
||||
#[error("invalid authentication configuration: invalid Azure credential selector")]
|
||||
InvalidAzureSelector,
|
||||
#[error(
|
||||
"invalid authentication configuration: ClientSecretCredential requires tenant_id, client_id, and client_secret"
|
||||
)]
|
||||
MissingClientSecretFields,
|
||||
#[error("invalid authentication configuration: WorkloadIdentityCredential requires tenant_id")]
|
||||
MissingWorkloadTenant,
|
||||
#[error("invalid authentication configuration: WorkloadIdentityCredential requires client_id")]
|
||||
MissingWorkloadClient,
|
||||
#[error(
|
||||
"invalid authentication configuration: WorkloadIdentityCredential requires azure_federated_token_file"
|
||||
)]
|
||||
MissingWorkloadTokenFile,
|
||||
#[error(
|
||||
"invalid authentication configuration: credential reference requires a host credential resolver"
|
||||
)]
|
||||
MissingHostResolver,
|
||||
#[error(
|
||||
"invalid authentication configuration: caller credential plan requires provider-specific inputs"
|
||||
)]
|
||||
MissingCallerInputs,
|
||||
#[error("invalid authentication configuration: credential header {0} already exists")]
|
||||
DuplicateHeader(&'static str),
|
||||
#[error("invalid authentication configuration: {0} must be a string or null")]
|
||||
InvalidFieldType(String),
|
||||
#[error("invalid authentication configuration: unsupported OIDC reference")]
|
||||
UnsupportedOidcReference,
|
||||
#[error("invalid authentication configuration: {0} cannot be empty")]
|
||||
EmptyReference(String),
|
||||
#[error("invalid authentication configuration: Azure credential initialization failed: {0}")]
|
||||
AzureCredentialInitialization(String),
|
||||
#[error(
|
||||
"invalid authentication configuration: Azure authority must be an HTTPS origin without credentials, query, or fragment"
|
||||
)]
|
||||
InvalidAzureAuthority,
|
||||
#[error(
|
||||
"invalid authentication configuration: request-controlled Azure auth inputs cannot be combined with host credentials"
|
||||
)]
|
||||
MixedAzureCredentialSources,
|
||||
#[error(
|
||||
"invalid authentication configuration: request-controlled Azure credential references are not allowed"
|
||||
)]
|
||||
RequestAzureCredentialReference,
|
||||
#[error(
|
||||
"invalid authentication configuration: host credentials cannot be sent to a request-controlled Azure endpoint"
|
||||
)]
|
||||
RequestAzureCredentialDestination,
|
||||
#[error(
|
||||
"invalid authentication configuration: credentials cannot be sent to a request-controlled Vertex AI endpoint"
|
||||
)]
|
||||
RequestVertexCredentialDestination,
|
||||
#[error(
|
||||
"invalid authentication configuration: request-controlled Vertex credentials must use the canonical Google OAuth token endpoint"
|
||||
)]
|
||||
RequestVertexTokenEndpoint,
|
||||
#[error("invalid authentication configuration: {0}")]
|
||||
InvalidConfiguration(#[source] ErrorDetail),
|
||||
#[error("credential acquisition failed: {0}")]
|
||||
AzureTokenAcquisition(String),
|
||||
#[error("credential acquisition failed: Vertex AI credentials: {0}")]
|
||||
VertexTokenAcquisition(String),
|
||||
CredentialAcquisition(#[source] ErrorDetail),
|
||||
#[error("credential caller failed: {0}")]
|
||||
EmptyCallerCredential(&'static str),
|
||||
#[error("{0}")]
|
||||
ProviderAuthentication(String),
|
||||
#[error("credential acquisition failed: {}", .0.iter().map(ToString::to_string).collect::<Vec<_>>().join("; "))]
|
||||
CredentialChain(Vec<Error>),
|
||||
#[error("credential caller failed: credential caller returned an empty credential")]
|
||||
EmptyCallerCredential,
|
||||
#[error("credential caller failed: Azure AD token provider returned an empty token")]
|
||||
EmptyAzureToken,
|
||||
#[error("credential acquisition failed: Azure OIDC reference did not resolve to a value")]
|
||||
UnresolvedOidcReference,
|
||||
#[error(
|
||||
"Missing {provider} API Key - Set `api_key` or the {environment_variable} environment variable"
|
||||
)]
|
||||
|
|
@ -87,34 +19,87 @@ pub enum Error {
|
|||
provider: &'static str,
|
||||
environment_variable: &'static str,
|
||||
},
|
||||
#[error(
|
||||
"Missing {provider} API Base - Set {environment_variable} environment variable or pass api_base parameter"
|
||||
)]
|
||||
#[error("Missing {provider} API Base - {guidance}")]
|
||||
MissingApiBase {
|
||||
provider: &'static str,
|
||||
environment_variable: &'static str,
|
||||
guidance: &'static str,
|
||||
},
|
||||
#[error(
|
||||
"Missing Azure API Base - Set `api_base` or the AZURE_API_BASE environment variable. Expected format: https://<resource-name>.services.ai.azure.com/anthropic"
|
||||
)]
|
||||
MissingAzureApiBase,
|
||||
#[error("invalid authentication header")]
|
||||
InvalidHeader,
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::Error;
|
||||
#[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)]
|
||||
pub enum ErrorDetail {
|
||||
#[error("{0}")]
|
||||
Message(String),
|
||||
#[error("{field} must be {expected}")]
|
||||
InvalidType {
|
||||
field: String,
|
||||
expected: &'static str,
|
||||
},
|
||||
#[error("{subject} cannot be empty")]
|
||||
Empty { subject: String },
|
||||
#[error("credential header {0} already exists")]
|
||||
DuplicateHeader(&'static str),
|
||||
#[error("{operation} failed: {source}")]
|
||||
Failed {
|
||||
operation: &'static str,
|
||||
#[source]
|
||||
source: ErrorSource,
|
||||
},
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn missing_api_key_names_provider_and_environment_variable() {
|
||||
assert_eq!(
|
||||
Error::MissingApiKey {
|
||||
provider: "Anthropic",
|
||||
environment_variable: "ANTHROPIC_API_KEY",
|
||||
}
|
||||
.to_string(),
|
||||
"Missing Anthropic API Key - Set `api_key` or the ANTHROPIC_API_KEY environment variable"
|
||||
);
|
||||
impl ErrorDetail {
|
||||
pub fn failed(
|
||||
operation: &'static str,
|
||||
source: impl std::error::Error + Send + Sync + 'static,
|
||||
) -> Self {
|
||||
Self::Failed {
|
||||
operation,
|
||||
source: ErrorSource::new(source),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl From<String> for ErrorDetail {
|
||||
fn from(message: String) -> Self {
|
||||
Self::Message(message)
|
||||
}
|
||||
}
|
||||
|
||||
impl From<&str> for ErrorDetail {
|
||||
fn from(message: &str) -> Self {
|
||||
Self::Message(message.into())
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug)]
|
||||
pub struct ErrorSource(std::sync::Arc<dyn std::error::Error + Send + Sync>);
|
||||
|
||||
impl std::ops::Deref for ErrorSource {
|
||||
type Target = dyn std::error::Error + Send + Sync;
|
||||
|
||||
fn deref(&self) -> &Self::Target {
|
||||
self.0.as_ref()
|
||||
}
|
||||
}
|
||||
|
||||
impl std::fmt::Display for ErrorSource {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
std::fmt::Display::fmt(&self.0, formatter)
|
||||
}
|
||||
}
|
||||
|
||||
impl ErrorSource {
|
||||
pub fn new(error: impl std::error::Error + Send + Sync + 'static) -> Self {
|
||||
Self(std::sync::Arc::new(error))
|
||||
}
|
||||
}
|
||||
|
||||
impl PartialEq for ErrorSource {
|
||||
fn eq(&self, other: &Self) -> bool {
|
||||
std::sync::Arc::ptr_eq(&self.0, &other.0)
|
||||
}
|
||||
}
|
||||
|
||||
impl Eq for ErrorSource {}
|
||||
|
|
|
|||
|
|
@ -7,7 +7,7 @@ pub enum CredentialPlacement {
|
|||
}
|
||||
|
||||
impl CredentialPlacement {
|
||||
pub fn header_name(self) -> &'static str {
|
||||
pub const fn header_name(self) -> &'static str {
|
||||
match self {
|
||||
Self::Bearer => "Authorization",
|
||||
Self::Header(name) => name,
|
||||
|
|
@ -21,13 +21,17 @@ pub fn apply_credential(
|
|||
placement: CredentialPlacement,
|
||||
) -> Result<Vec<(String, String)>, Error> {
|
||||
if credential.trim().is_empty() {
|
||||
return Err(Error::EmptyCredential);
|
||||
return Err(Error::InvalidConfiguration(
|
||||
"credential cannot be empty".into(),
|
||||
));
|
||||
}
|
||||
if headers
|
||||
.iter()
|
||||
.any(|(name, _)| name.eq_ignore_ascii_case(placement.header_name()))
|
||||
{
|
||||
return Err(Error::DuplicateHeader(placement.header_name()));
|
||||
return Err(Error::InvalidConfiguration(
|
||||
crate::ErrorDetail::DuplicateHeader(placement.header_name()),
|
||||
));
|
||||
}
|
||||
let value = match placement {
|
||||
CredentialPlacement::Bearer => format!("Bearer {credential}"),
|
||||
|
|
@ -40,21 +44,6 @@ pub fn apply_credential(
|
|||
)
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq)]
|
||||
pub enum RequestAuth {
|
||||
Header {
|
||||
name: &'static str,
|
||||
value: String,
|
||||
},
|
||||
Bearer {
|
||||
token: String,
|
||||
},
|
||||
AwsSigV4 {
|
||||
region: String,
|
||||
service: &'static str,
|
||||
},
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{CredentialPlacement, apply_credential};
|
||||
|
|
|
|||
|
|
@ -50,8 +50,8 @@ pub use credential::{
|
|||
CredentialFileRef, CredentialLookup, CredentialLookupFuture, CredentialPlan,
|
||||
CredentialPlanResolution, CredentialRef, CredentialResolver, CredentialResolverHandle,
|
||||
};
|
||||
pub use error::Error;
|
||||
pub use http::{CredentialPlacement, RequestAuth};
|
||||
pub use error::{Error, ErrorDetail, ErrorSource};
|
||||
pub use http::CredentialPlacement;
|
||||
pub use policy::{CredentialPlanKind, CredentialRule, ExistingHeaderBehavior, ProviderAuthPolicy};
|
||||
pub use secret::SecretValue;
|
||||
pub use token::{ResolvedCredential, TokenFuture, TokenProvider, TokenProviderHandle};
|
||||
|
|
|
|||
|
|
@ -47,14 +47,20 @@ impl ProviderAuthPolicy {
|
|||
if self.has_existing_credential(&headers) {
|
||||
return match self.existing_header_behavior {
|
||||
ExistingHeaderBehavior::Preserve => Ok(headers),
|
||||
ExistingHeaderBehavior::Reject => Err(Error::ExistingCredentialHeader),
|
||||
ExistingHeaderBehavior::Reject => Err(Error::InvalidConfiguration(
|
||||
"credential header already exists".into(),
|
||||
)),
|
||||
};
|
||||
}
|
||||
let rule = self
|
||||
.rules
|
||||
.iter()
|
||||
.find(|rule| rule.kind == kind)
|
||||
.ok_or(Error::DisallowedCredentialPlan)?;
|
||||
.ok_or_else(|| {
|
||||
Error::InvalidConfiguration(
|
||||
"credential plan is not allowed by the provider auth policy".into(),
|
||||
)
|
||||
})?;
|
||||
apply_credential(headers, credential.secret().expose(), rule.placement)
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -39,7 +39,33 @@ impl TokenProviderHandle {
|
|||
Self(caller)
|
||||
}
|
||||
|
||||
pub fn from_callback<F, Fut>(acquire: F) -> Self
|
||||
where
|
||||
F: Fn() -> Fut + Send + Sync + 'static,
|
||||
Fut: Future<Output = Result<ResolvedCredential, Error>> + Send + 'static,
|
||||
{
|
||||
Self::new(Arc::new(CallbackTokenProvider(acquire)))
|
||||
}
|
||||
|
||||
pub async fn acquire(&self) -> Result<ResolvedCredential, Error> {
|
||||
self.0.acquire().await
|
||||
}
|
||||
}
|
||||
|
||||
struct CallbackTokenProvider<F>(F);
|
||||
|
||||
impl<F> std::fmt::Debug for CallbackTokenProvider<F> {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter.write_str("CallbackTokenProvider")
|
||||
}
|
||||
}
|
||||
|
||||
impl<F, Fut> TokenProvider for CallbackTokenProvider<F>
|
||||
where
|
||||
F: Fn() -> Fut + Send + Sync,
|
||||
Fut: Future<Output = Result<ResolvedCredential, Error>> + Send + 'static,
|
||||
{
|
||||
fn acquire(&self) -> TokenFuture<'_> {
|
||||
Box::pin((self.0)())
|
||||
}
|
||||
}
|
||||
|
|
|
|||
74
litellm-rust/crates/auth-types/tests/error.rs
Normal file
74
litellm-rust/crates/auth-types/tests/error.rs
Normal file
|
|
@ -0,0 +1,74 @@
|
|||
use litellm_auth_types::Error;
|
||||
use rstest::rstest;
|
||||
|
||||
#[rstest]
|
||||
#[case::api_key(
|
||||
Error::MissingApiKey { provider: "Example", environment_variable: "EXAMPLE_API_KEY" },
|
||||
"Missing Example API Key - Set `api_key` or the EXAMPLE_API_KEY environment variable"
|
||||
)]
|
||||
#[case::another_api_key(
|
||||
Error::MissingApiKey { provider: "Custom", environment_variable: "CUSTOM_KEY" },
|
||||
"Missing Custom API Key - Set `api_key` or the CUSTOM_KEY environment variable"
|
||||
)]
|
||||
#[case::api_base(
|
||||
Error::MissingApiBase { provider: "Example", guidance: "Pass api_base" },
|
||||
"Missing Example API Base - Pass api_base"
|
||||
)]
|
||||
#[case::another_api_base(
|
||||
Error::MissingApiBase { provider: "Custom", guidance: "Set CUSTOM_ENDPOINT" },
|
||||
"Missing Custom API Base - Set CUSTOM_ENDPOINT"
|
||||
)]
|
||||
#[case::configuration(
|
||||
Error::InvalidConfiguration("credential selector is invalid".into()),
|
||||
"invalid authentication configuration: credential selector is invalid"
|
||||
)]
|
||||
#[case::acquisition(
|
||||
Error::CredentialAcquisition("token expired".into()),
|
||||
"credential acquisition failed: token expired"
|
||||
)]
|
||||
#[case::caller(
|
||||
Error::EmptyCallerCredential("empty token"),
|
||||
"credential caller failed: empty token"
|
||||
)]
|
||||
#[case::provider(
|
||||
Error::ProviderAuthentication("provider rejected credentials".into()),
|
||||
"provider rejected credentials"
|
||||
)]
|
||||
#[case::chain(
|
||||
Error::CredentialChain(vec![
|
||||
Error::CredentialAcquisition("token expired".into()),
|
||||
Error::EmptyCallerCredential("empty token"),
|
||||
]),
|
||||
"credential acquisition failed: credential acquisition failed: token expired; credential caller failed: empty token"
|
||||
)]
|
||||
fn display_preserves_failure_phase_and_caller_context(
|
||||
#[case] error: Error,
|
||||
#[case] expected: &str,
|
||||
) {
|
||||
assert_eq!(error.to_string(), expected);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::configuration(true)]
|
||||
#[case::acquisition(false)]
|
||||
fn contextual_failures_keep_the_original_source(#[case] configuration: bool) {
|
||||
use litellm_auth_types::ErrorDetail;
|
||||
|
||||
let detail = ErrorDetail::failed(
|
||||
"test credential",
|
||||
std::io::Error::from(std::io::ErrorKind::PermissionDenied),
|
||||
);
|
||||
let error = if configuration {
|
||||
Error::InvalidConfiguration(detail)
|
||||
} else {
|
||||
Error::CredentialAcquisition(detail)
|
||||
};
|
||||
let source = std::iter::successors(Some(&error as &dyn std::error::Error), |error| {
|
||||
error.source()
|
||||
})
|
||||
.find_map(|error| error.downcast_ref::<std::io::Error>())
|
||||
.expect("the original credential error remains available");
|
||||
assert_eq!(source.kind(), std::io::ErrorKind::PermissionDenied);
|
||||
assert!(error.to_string().contains("test credential failed:"));
|
||||
assert!(error.to_string().ends_with(&source.to_string()));
|
||||
}
|
||||
104
litellm-rust/crates/auth-types/tests/token.rs
Normal file
104
litellm-rust/crates/auth-types/tests/token.rs
Normal file
|
|
@ -0,0 +1,104 @@
|
|||
use std::{
|
||||
error::Error as StdError,
|
||||
future::{Future, poll_fn},
|
||||
sync::{
|
||||
Arc,
|
||||
atomic::{AtomicBool, AtomicUsize, Ordering},
|
||||
},
|
||||
task::Poll,
|
||||
time::{Duration, SystemTime},
|
||||
};
|
||||
|
||||
use litellm_auth_types::{
|
||||
Error, ErrorDetail, ResolvedCredential, SecretValue, TokenProviderHandle,
|
||||
};
|
||||
use rstest::rstest;
|
||||
|
||||
fn credential(index: usize, access_token: bool) -> ResolvedCredential {
|
||||
let token = SecretValue::new(format!("credential-{index}"));
|
||||
if access_token {
|
||||
return ResolvedCredential::AccessToken {
|
||||
token,
|
||||
expires_on: Some(SystemTime::UNIX_EPOCH + Duration::from_secs(index as u64)),
|
||||
};
|
||||
}
|
||||
ResolvedCredential::Static(token)
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::static_secret(false)]
|
||||
#[case::access_token(true)]
|
||||
#[tokio::test]
|
||||
async fn callbacks_acquire_fresh_credentials_on_demand(#[case] access_token: bool) {
|
||||
let calls = Arc::new(AtomicUsize::new(0));
|
||||
let callback_calls = calls.clone();
|
||||
let provider = TokenProviderHandle::from_callback(move || {
|
||||
let index = callback_calls.fetch_add(1, Ordering::SeqCst);
|
||||
async move {
|
||||
tokio::task::yield_now().await;
|
||||
Ok(credential(index, access_token))
|
||||
}
|
||||
});
|
||||
let cloned = provider.clone();
|
||||
|
||||
assert_eq!(calls.load(Ordering::SeqCst), 0);
|
||||
assert_eq!(
|
||||
provider.acquire().await.unwrap(),
|
||||
credential(0, access_token)
|
||||
);
|
||||
assert_eq!(calls.load(Ordering::SeqCst), 1);
|
||||
assert_eq!(cloned.acquire().await.unwrap(), credential(1, access_token));
|
||||
assert_eq!(calls.load(Ordering::SeqCst), 2);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn callback_errors_preserve_the_original_source() {
|
||||
let provider = TokenProviderHandle::from_callback(|| async {
|
||||
Err(Error::CredentialAcquisition(ErrorDetail::failed(
|
||||
"caller credential",
|
||||
std::io::Error::from(std::io::ErrorKind::PermissionDenied),
|
||||
)))
|
||||
});
|
||||
|
||||
let error = provider.acquire().await.unwrap_err();
|
||||
assert!(matches!(error, Error::CredentialAcquisition(_)));
|
||||
let source = std::iter::successors(Some(&error as &(dyn StdError + 'static)), |error| {
|
||||
(*error).source()
|
||||
})
|
||||
.find_map(|error| error.downcast_ref::<std::io::Error>())
|
||||
.unwrap();
|
||||
assert_eq!(source.kind(), std::io::ErrorKind::PermissionDenied);
|
||||
}
|
||||
|
||||
struct Release(Arc<AtomicBool>);
|
||||
|
||||
impl Drop for Release {
|
||||
fn drop(&mut self) {
|
||||
self.0.store(true, Ordering::SeqCst);
|
||||
}
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn cancelling_acquisition_drops_the_callback_future() {
|
||||
let released = Arc::new(AtomicBool::new(false));
|
||||
let callback_released = released.clone();
|
||||
let provider = TokenProviderHandle::from_callback(move || {
|
||||
let released = callback_released.clone();
|
||||
async move {
|
||||
let _release = Release(released);
|
||||
std::future::pending().await
|
||||
}
|
||||
});
|
||||
|
||||
let mut acquisition = Box::pin(provider.acquire());
|
||||
poll_fn(|context| {
|
||||
assert!(acquisition.as_mut().poll(context).is_pending());
|
||||
assert!(!released.load(Ordering::SeqCst));
|
||||
Poll::Ready(())
|
||||
})
|
||||
.await;
|
||||
drop(acquisition);
|
||||
assert!(released.load(Ordering::SeqCst));
|
||||
}
|
||||
|
|
@ -2,6 +2,9 @@
|
|||
|
||||
pub use litellm_auth_types::*;
|
||||
|
||||
mod services;
|
||||
pub use services::AuthServices;
|
||||
|
||||
#[cfg(feature = "aws")]
|
||||
pub use litellm_auth_aws as aws;
|
||||
#[cfg(feature = "azure")]
|
||||
|
|
|
|||
9
litellm-rust/crates/auth/src/services.rs
Normal file
9
litellm-rust/crates/auth/src/services.rs
Normal file
|
|
@ -0,0 +1,9 @@
|
|||
#[derive(Default)]
|
||||
pub struct AuthServices {
|
||||
#[cfg(feature = "aws")]
|
||||
pub aws: litellm_auth_aws::AwsAuthService,
|
||||
#[cfg(feature = "azure")]
|
||||
pub azure: litellm_auth_azure::AzureAuthService,
|
||||
#[cfg(feature = "gcp")]
|
||||
pub gcp: litellm_auth_gcp::VertexAuth,
|
||||
}
|
||||
|
|
@ -6,6 +6,7 @@ license.workspace = true
|
|||
repository.workspace = true
|
||||
|
||||
[dependencies]
|
||||
litellm-http.workspace = true
|
||||
litellm-auth-azure.workspace = true
|
||||
litellm-auth-types.workspace = true
|
||||
litellm-cache.workspace = true
|
||||
|
|
@ -19,6 +20,7 @@ tokio.workspace = true
|
|||
url.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
litellm-http = { workspace = true, features = ["test-support"] }
|
||||
litellm-cache-response.workspace = true
|
||||
litellm-cache-testing.workspace = true
|
||||
rstest.workspace = true
|
||||
|
|
|
|||
|
|
@ -31,7 +31,7 @@ impl<C: CacheCodec> AzureBlobCache<C> {
|
|||
pub async fn connect(
|
||||
account_url: &str,
|
||||
container: &str,
|
||||
http: reqwest::Client,
|
||||
http: litellm_http::Client,
|
||||
codec: C,
|
||||
runtime: Handle,
|
||||
) -> Result<Self, Error> {
|
||||
|
|
|
|||
|
|
@ -8,7 +8,7 @@ use azure_core::{
|
|||
use futures_util::TryStreamExt;
|
||||
|
||||
#[derive(Debug)]
|
||||
pub struct ReqwestTransport(pub reqwest::Client);
|
||||
pub struct ReqwestTransport(pub litellm_http::Client);
|
||||
|
||||
#[async_trait::async_trait]
|
||||
impl HttpClient for ReqwestTransport {
|
||||
|
|
|
|||
|
|
@ -31,7 +31,7 @@ async fn connect(server: &MockServer) -> AzureBlobCache<JsonCodec<Value>> {
|
|||
None,
|
||||
ClientOptions {
|
||||
transport: Some(Transport::new(Arc::new(ReqwestTransport(
|
||||
reqwest::Client::new(),
|
||||
litellm_http::Client::plain_for_test(),
|
||||
)))),
|
||||
..ClientOptions::default()
|
||||
},
|
||||
|
|
|
|||
|
|
@ -9,7 +9,7 @@ repository.workspace = true
|
|||
litellm-cache.workspace = true
|
||||
py_literal = "0.4.0"
|
||||
rand.workspace = true
|
||||
rusqlite = { version = "0.40", features = ["bundled"] }
|
||||
rusqlite = { version = "0.39", features = ["bundled"] }
|
||||
serde-pickle = "1.2"
|
||||
serde_json.workspace = true
|
||||
tokio.workspace = true
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ license.workspace = true
|
|||
repository.workspace = true
|
||||
|
||||
[dependencies]
|
||||
litellm-http.workspace = true
|
||||
futures-util.workspace = true
|
||||
litellm-auth-gcp.workspace = true
|
||||
litellm-auth-types.workspace = true
|
||||
|
|
@ -15,6 +16,7 @@ reqwest.workspace = true
|
|||
tokio.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
litellm-http = { workspace = true, features = ["test-support"] }
|
||||
litellm-cache-testing.workspace = true
|
||||
rstest.workspace = true
|
||||
serde_json.workspace = true
|
||||
|
|
|
|||
|
|
@ -5,8 +5,8 @@ use litellm_cache::{
|
|||
BaseCache, BatchCache, BatchEntry, CacheCodec, DisconnectCache, Error, ExactCacheContext,
|
||||
FlushCache,
|
||||
};
|
||||
use litellm_http::Client;
|
||||
use percent_encoding::{AsciiSet, NON_ALPHANUMERIC, percent_encode};
|
||||
use reqwest::Client;
|
||||
|
||||
use crate::{GcpTokenSource, TokenSource};
|
||||
|
||||
|
|
|
|||
|
|
@ -56,6 +56,7 @@ async fn set_writes_encoded_object_and_headers(#[future(awt)] server: MockServer
|
|||
)]
|
||||
#[case::missing("missing", ResponseTemplate::new(404), Ok(None))]
|
||||
#[case::server_error("server-error", ResponseTemplate::new(500), Err(Error::Unavailable))]
|
||||
#[case::unauthorized("unauthorized", ResponseTemplate::new(401), Err(Error::Unavailable))]
|
||||
#[case::invalid(
|
||||
"invalid",
|
||||
ResponseTemplate::new(200).set_body_string("not json"),
|
||||
|
|
@ -96,7 +97,7 @@ async fn cache_exposes_its_configuration(#[future(awt)] server: MockServer) {
|
|||
path_service_account: Some("/secrets/sa.json".into()),
|
||||
..support::config(&server, Some("folder"))
|
||||
},
|
||||
reqwest::Client::new(),
|
||||
litellm_http::Client::plain_for_test(),
|
||||
litellm_cache::JsonCodec::<Value>::new(),
|
||||
);
|
||||
assert_eq!(cache.bucket_name(), "bucket");
|
||||
|
|
|
|||
|
|
@ -29,7 +29,7 @@ pub fn cache_with_token(
|
|||
) -> JsonGcsCache {
|
||||
GcsCache::with_token_source(
|
||||
config(server, gcs_path),
|
||||
reqwest::Client::new(),
|
||||
litellm_http::Client::plain_for_test(),
|
||||
JsonCodec::new(),
|
||||
token,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ license.workspace = true
|
|||
repository.workspace = true
|
||||
|
||||
[dependencies]
|
||||
litellm-http.workspace = true
|
||||
futures-util.workspace = true
|
||||
litellm-cache.workspace = true
|
||||
qdrant-client = { workspace = true, features = ["serde"] }
|
||||
|
|
@ -17,6 +18,7 @@ tokio.workspace = true
|
|||
uuid.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
litellm-http = { workspace = true, features = ["test-support"] }
|
||||
futures-executor = "0.3"
|
||||
litellm-cache-testing.workspace = true
|
||||
rstest.workspace = true
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
use std::time::Duration;
|
||||
|
||||
use litellm_cache::{Error, semantic::Embedder};
|
||||
use reqwest::Client;
|
||||
use litellm_http::Client;
|
||||
use serde_json::Value;
|
||||
|
||||
pub struct OpenAiEmbedder {
|
||||
|
|
|
|||
|
|
@ -5,6 +5,10 @@ use std::{
|
|||
|
||||
use litellm_cache::{Error, semantic::Embedder};
|
||||
use litellm_cache_qdrant_semantic::{OpenAiEmbedder, OpenAiEmbedderConfig};
|
||||
use litellm_http::{
|
||||
ClientVariant, HttpClientConfig, HttpClientPool, HttpSettings, Resolution,
|
||||
media::PublicDnsResolver,
|
||||
};
|
||||
use rstest::rstest;
|
||||
use serde_json::{Value, json};
|
||||
use tokio::{
|
||||
|
|
@ -104,7 +108,7 @@ fn config(base: String, timeout: Option<Duration>) -> OpenAiEmbedderConfig {
|
|||
async fn posts_embeddings_request_and_parses_vector() {
|
||||
let server = TestHttpServer::response("200 OK", r#"{"data":[{"embedding":[0.1,0.2]}]}"#).await;
|
||||
let embedder = OpenAiEmbedder::new(
|
||||
reqwest::Client::new(),
|
||||
litellm_http::Client::plain_for_test(),
|
||||
config(
|
||||
format!("{}/", server.base_url()),
|
||||
Some(Duration::from_secs(1)),
|
||||
|
|
@ -156,14 +160,17 @@ async fn status_timeout_and_body_errors_are_unavailable(
|
|||
) {
|
||||
let server =
|
||||
TestHttpServer::response_after(status, body, Duration::from_millis(delay_ms)).await;
|
||||
let embedder = OpenAiEmbedder::new(reqwest::Client::new(), config(server.base_url(), timeout));
|
||||
let embedder = OpenAiEmbedder::new(
|
||||
litellm_http::Client::plain_for_test(),
|
||||
config(server.base_url(), timeout),
|
||||
);
|
||||
assert_eq!(embedder.async_embed("hello", None).await, expected);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn sync_embedding_is_unsupported() {
|
||||
let embedder = OpenAiEmbedder::new(
|
||||
reqwest::Client::new(),
|
||||
litellm_http::Client::plain_for_test(),
|
||||
config("http://127.0.0.1:9".to_owned(), None),
|
||||
);
|
||||
assert_eq!(
|
||||
|
|
@ -176,9 +183,12 @@ fn sync_embedding_is_unsupported() {
|
|||
#[tokio::test]
|
||||
async fn uses_the_injected_client() {
|
||||
let server = TestHttpServer::response("200 OK", r#"{"data":[{"embedding":[0.1,0.2]}]}"#).await;
|
||||
let client = reqwest::Client::builder()
|
||||
.user_agent("litellm-embedder-test")
|
||||
.build()
|
||||
let config_with_agent = HttpClientConfig {
|
||||
user_agent: Some("litellm-embedder-test".into()),
|
||||
..Resolution::from(&HttpSettings::default()).config
|
||||
};
|
||||
let client = HttpClientPool::new(Arc::new(PublicDnsResolver))
|
||||
.client(&config_with_agent, ClientVariant::Provider)
|
||||
.unwrap();
|
||||
let embedder = OpenAiEmbedder::new(client, config(server.base_url(), None));
|
||||
assert_eq!(
|
||||
|
|
|
|||
29
litellm-rust/crates/cache-response/AGENTS.md
Normal file
29
litellm-rust/crates/cache-response/AGENTS.md
Normal file
|
|
@ -0,0 +1,29 @@
|
|||
# Response caching
|
||||
|
||||
Design this crate for shared Rust execution used by the Python SDK and the Rust gateway. The Python SDK will remain, with more core execution moving to Rust and Python callbacks staying in Python. The Rust gateway is still evolving and is intended to replace the Python proxy. Keep response-cache policy independent of Python, HTTP serving, and either proxy's configuration format
|
||||
|
||||
Separate what is cached, how a hit is matched, and where entries are stored. Chat Completions, Messages, Responses, and embeddings are API workloads. Exact and semantic matching are lookup behaviors. Memory, Redis, disk, and object stores are storage choices. Embeddings are inference too, so do not use an inference-cache name to imply a category that excludes embeddings. Consult the existing Python cache and caching handler for behavior and compatibility contracts without copying their class structure
|
||||
|
||||
Storage traits, codecs, and backend capabilities belong in `litellm-cache` and the storage crates. Keep storage reusable for value types beyond LLM responses. This crate owns response entries, matching and freshness semantics, the Python-compatible response codec, and deferred-write policy. Core owns route-specific request identity, response encoding and reconstruction, embedding partial-hit orchestration, and stream capture and replay. Boundaries own configuration translation, resource construction, and caller identity
|
||||
|
||||
Construct and inject the response-cache service at the Python bridge or gateway boundary, as with the HTTP client. Reuse it across calls. Core and provider code must not discover cache configuration through Python globals, process configuration, or backend-specific factories
|
||||
|
||||
Keep `ResponseCache<B>` generic over its storage backend. Preserve typed backend contexts and capability bounds internally. Inject an object-safe service into core for runtime backend selection, so storage types do not spread through route and host types. Keep API request and response types statically typed. Add a generic parameter only where it preserves a useful type relationship or capability
|
||||
|
||||
Keep the core service contract narrow. Lookup and store must not require connection testing, ping, flush, deletion, counters, queues, or scripts. Require batch operations where a consumer needs partial hits, and keep management capabilities on their own interfaces. An exact-only adapter must remain explicit about its matching restriction. Supporting semantic matching requires a defined lookup-context and embedding execution contract, not just a renamed trait
|
||||
|
||||
Separate reusable resources from per-call policy. Backend configuration, namespace, default expiry, and entry limits belong to the configured service or backend. Read/write controls, expiry and freshness overrides, and authenticated caller scope belong to the call. Passing call options must not replace or mutate the route's configured service
|
||||
|
||||
Keep cache misses and storage failures distinguishable in return values. Core owns the decision to continue with provider execution after a cache failure. A read can reject an entry for freshness while the backend still retains it. Preserve the timestamp at which a response was produced when writing it later
|
||||
|
||||
Define lookup placement explicitly relative to authorization, deployment and credential resolution, and request-transforming callbacks. Cache identity must account for every input that affects reuse, including API surface and caller scope, while preserving intentional Python caching groups. Preserve existing keys and response formats unless changing them is an explicit migration decision
|
||||
|
||||
Cache normalized provider results before caller-specific response transformations. Hits must still run the applicable response processing, success callbacks, and cache-hit accounting. Keep callback execution in the host. Python cache implementations and semantic embedders that require the caller's task must use the existing host-operation mechanism rather than Python calls from a Rust worker. Preserve legacy fallback until that contract is supported
|
||||
|
||||
Keep unary caching independent of stream-only methods. Store streams only after successful exhaustion and protocol completion. Errors, incomplete streams, cancellation, and oversized entries must not populate the cache. Embedding batches need ordered partial results and reconstruction around the uncached inputs
|
||||
|
||||
Test each contract in its owner: storage capabilities in backend tests, envelopes and freshness here, reuse and replay in core, Python callback and fallback behavior at the bridge, and HTTP behavior at the gateway. Run backend contract checks and Python response-codec fixtures before exposing a new backend
|
||||
|
||||
`ScopedCache` requires an explicit shared or isolated scope at construction. `CacheOptions` has no default sharing policy. Callers may override policy per invocation without replacing the attached service. Versioned native envelopes reject incompatible API surfaces and versions as misses; this envelope is distinct from the legacy Python response codec
|
||||
|
||||
Response storage is not the source of budget or rate-limit coordination dependencies. Keep counters, reservations, and atomic admission operations out of `ResponseCacheService`, including when both services happen to use Redis
|
||||
|
|
@ -13,9 +13,12 @@ serde_json.workspace = true
|
|||
sha2.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
litellm-cache-gcs.workspace = true
|
||||
litellm-http = { workspace = true, features = ["test-support"] }
|
||||
litellm-cache-memory.workspace = true
|
||||
litellm-cache-redis.workspace = true
|
||||
redis = "1.7.0"
|
||||
redis-test = "1.0.4"
|
||||
rstest.workspace = true
|
||||
tokio.workspace = true
|
||||
wiremock = "0.6.5"
|
||||
|
|
|
|||
|
|
@ -1,51 +0,0 @@
|
|||
# Response cache
|
||||
|
||||
`ResponseCache<B>` adds request keys, independent read/write controls, response envelopes, and freshness checks to any `B: BaseCache<Value = CacheEntry>`
|
||||
|
||||
## Ownership
|
||||
|
||||
`litellm-cache` defines typed storage, codec, and capability traits. `BaseCache` is only get, set, TTL, and pipeline writes. Everything else is an optional capability a backend implements only where its Python class defines the method: `DisconnectCache`, `ConnectionCache` (`test_connection`), `PingCache`, `BatchCache`, `DeleteCache`, `FlushCache`, counters, queues, TTL, scan, and scripts. Memory, Redis, disk, S3, GCS, and Azure Blob implement those traits without depending on response policy, so other consumers can store their own value types in the same backends
|
||||
|
||||
Semantic backends (Redis, Valkey, Qdrant) are generic over their embedder and codec, and share one prompt and embedding contract from `litellm_cache::semantic`. They take a `SemanticCacheContext`, so `ResponseCache` drives them the same way it drives exact backends
|
||||
|
||||
`litellm-cache-response` owns response keys, controls, entries, the Python-compatible response codec, and `WriteBuffer`, the backend-neutral deferred-write policy. It has no runtime dependency on a specific cache backend or Python
|
||||
|
||||
`ExactResponseCache` is the object-safe view of a `ResponseCache` over an exact backend. `ConnectionProbe` is the object-safe `test_connection`, implemented only when the backend implements `ConnectionCache`, so a host holds one next to its `ExactResponseCache` and reports the operation as unsupported otherwise, as Python's `BaseCache` does. Lookup, store, batch, and flush never require it
|
||||
|
||||
## Native Rust use
|
||||
|
||||
```rust
|
||||
use std::{sync::Arc, time::Duration};
|
||||
use litellm_cache_memory::InMemoryCache;
|
||||
use litellm_cache_response::{CacheKeyInput, ResponseCache, ResponseCacheRequest};
|
||||
use serde_json::json;
|
||||
|
||||
let cache = ResponseCache::new(Arc::new(InMemoryCache::default()));
|
||||
let request = ResponseCacheRequest::new(CacheKeyInput {
|
||||
preset: Some("example:key".into()),
|
||||
..Default::default()
|
||||
});
|
||||
let now = Duration::from_secs(100);
|
||||
cache.store(&request, json!({"answer": 7}), now)?;
|
||||
assert_eq!(cache.async_lookup(&request, now).await?, Some(json!({"answer": 7})));
|
||||
```
|
||||
|
||||
For Redis, inject `RedisCache::new(url, ttl, ResponseCacheCodec)` instead. Namespaces are optional and existing namespace prefixes are preserved
|
||||
|
||||
Callers supply Unix time for response freshness. Backend TTL uses its own clock. A read can reject an entry through `max_age` even while the backend still retains it
|
||||
|
||||
## Python integration
|
||||
|
||||
The bridge activates backends through the Rust catalog in `litellm/rust_bridge/catalog.py`. Every cache rule ships as `PYTHON_ONLY`, so SDK, Router, and proxy calls stay on Python and construct no native cache resources until a rule is changed
|
||||
|
||||
When a rule selects a backend, the Python `Cache` facade builds the native runtime from its own configuration and routes its storage calls (sync and async lookup and store, and pipelined batch store) to it. Stream replay, embedding partial-hit merging, response reconstruction, and callbacks stay in Python on top of that native store. The Python backend object remains for its direct API
|
||||
|
||||
Object responses are written as they are, and every other response shape is written as a serialized string, which is the pair of shapes Python reads. A string on the wire is therefore always a serialized response, so string-valued responses round trip. Typed backends such as memory never pass through the codec
|
||||
|
||||
Native cache handles must be recreated after fork. Native errors propagate to the host, which owns the existing fail-open and logging policy
|
||||
|
||||
## Adding another backend
|
||||
|
||||
Implement `BaseCache` for the backend with its associated value type and the capability traits its Python class supports, and accept a `CacheCodec` when wire serialization is needed. `ResponseCache<B>` then works without another response implementation
|
||||
|
||||
Run the `litellm-cache-testing` contract checks the backend's capabilities allow, and run response fixtures with `ResponseCacheCodec`, including both Python envelope encodings, before adding a catalog rule
|
||||
|
|
@ -4,6 +4,7 @@ mod codec;
|
|||
mod embedding;
|
||||
mod exact;
|
||||
mod response;
|
||||
mod service;
|
||||
|
||||
pub use buffer::WriteBuffer;
|
||||
pub use caching::{
|
||||
|
|
@ -14,3 +15,8 @@ pub use codec::ResponseCacheCodec;
|
|||
pub use embedding::PartialHits;
|
||||
pub use exact::{ConnectionProbe, ExactResponseCache};
|
||||
pub use response::{ResponseCache, ResponseCacheRequest};
|
||||
|
||||
pub use service::{
|
||||
CacheOptions, CacheScope, ResponseCacheConfig, ResponseCacheService, ResponseEnvelope,
|
||||
ScopedCache,
|
||||
};
|
||||
|
|
|
|||
|
|
@ -7,7 +7,9 @@ use litellm_cache::{
|
|||
};
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::{CacheControls, CacheEntry, CacheKeyInput, PartialHits, cache_key};
|
||||
use crate::{
|
||||
CacheControls, CacheEntry, CacheKeyInput, PartialHits, ResponseCacheConfig, cache_key,
|
||||
};
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct ResponseCacheRequest<C: CacheContext = litellm_cache::ExactCacheContext> {
|
||||
|
|
@ -50,6 +52,7 @@ where
|
|||
B::Context: Default + PartialEq,
|
||||
{
|
||||
backend: Arc<B>,
|
||||
config: ResponseCacheConfig,
|
||||
}
|
||||
|
||||
impl<B> ResponseCache<B>
|
||||
|
|
@ -58,7 +61,18 @@ where
|
|||
B::Context: Default + PartialEq,
|
||||
{
|
||||
pub fn new(backend: Arc<B>) -> Self {
|
||||
Self { backend }
|
||||
Self {
|
||||
backend,
|
||||
config: ResponseCacheConfig::default(),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn with_config(self, config: ResponseCacheConfig) -> Self {
|
||||
Self { config, ..self }
|
||||
}
|
||||
|
||||
pub fn config(&self) -> &ResponseCacheConfig {
|
||||
&self.config
|
||||
}
|
||||
|
||||
pub fn backend(&self) -> &B {
|
||||
|
|
@ -221,7 +235,7 @@ where
|
|||
response: Value,
|
||||
now: Duration,
|
||||
) -> Result<(), Error> {
|
||||
if !request.controls.writes() {
|
||||
if !request.controls.writes() || !self.fits(&response) {
|
||||
return Ok(());
|
||||
}
|
||||
self.backend.set_cache(
|
||||
|
|
@ -240,7 +254,7 @@ where
|
|||
response: Value,
|
||||
now: Duration,
|
||||
) -> Result<(), Error> {
|
||||
if !request.controls.writes() {
|
||||
if !request.controls.writes() || !self.fits(&response) {
|
||||
return Ok(());
|
||||
}
|
||||
self.backend
|
||||
|
|
@ -277,7 +291,7 @@ where
|
|||
) -> Result<(), Error> {
|
||||
let writable = entries
|
||||
.into_iter()
|
||||
.filter(|(request, _, _)| request.controls.writes())
|
||||
.filter(|(request, response, _)| request.controls.writes() && self.fits(response))
|
||||
.map(|(request, response, now)| {
|
||||
(
|
||||
cache_key(&request.key),
|
||||
|
|
@ -312,6 +326,11 @@ where
|
|||
Ok(())
|
||||
}
|
||||
|
||||
fn fits(&self, response: &Value) -> bool {
|
||||
self.config.max_entry_bytes == usize::MAX
|
||||
|| response.to_string().len() <= self.config.max_entry_bytes
|
||||
}
|
||||
|
||||
fn partial_hits(
|
||||
requests: &[ResponseCacheRequest<B::Context>],
|
||||
readable: Vec<(usize, &ResponseCacheRequest<B::Context>)>,
|
||||
|
|
|
|||
177
litellm-rust/crates/cache-response/src/service.rs
Normal file
177
litellm-rust/crates/cache-response/src/service.rs
Normal file
|
|
@ -0,0 +1,177 @@
|
|||
use std::{future::Future, pin::Pin, time::Duration};
|
||||
|
||||
use litellm_cache::{BaseCache, Error, ExactCacheContext};
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::{
|
||||
CacheControls, CacheEntry, CacheKeyField, CacheKeyInput, ResponseCache, ResponseCacheRequest,
|
||||
};
|
||||
|
||||
type CacheFuture<'a, T> = Pin<Box<dyn Future<Output = Result<T, Error>> + Send + 'a>>;
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct ResponseCacheConfig {
|
||||
pub namespace: String,
|
||||
pub max_entry_bytes: usize,
|
||||
}
|
||||
|
||||
impl Default for ResponseCacheConfig {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
namespace: String::new(),
|
||||
max_entry_bytes: usize::MAX,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub trait ResponseCacheService: Send + Sync {
|
||||
fn config(&self) -> &ResponseCacheConfig;
|
||||
|
||||
fn lookup<'a>(
|
||||
&'a self,
|
||||
request: &'a ResponseCacheRequest,
|
||||
now: Duration,
|
||||
) -> CacheFuture<'a, Option<Value>>;
|
||||
|
||||
fn store<'a>(
|
||||
&'a self,
|
||||
request: &'a ResponseCacheRequest,
|
||||
response: Value,
|
||||
now: Duration,
|
||||
) -> CacheFuture<'a, ()>;
|
||||
}
|
||||
|
||||
impl<B> ResponseCacheService for ResponseCache<B>
|
||||
where
|
||||
B: BaseCache<Value = CacheEntry, Context = ExactCacheContext>,
|
||||
{
|
||||
fn config(&self) -> &ResponseCacheConfig {
|
||||
self.config()
|
||||
}
|
||||
|
||||
fn lookup<'a>(
|
||||
&'a self,
|
||||
request: &'a ResponseCacheRequest,
|
||||
now: Duration,
|
||||
) -> CacheFuture<'a, Option<Value>> {
|
||||
Box::pin(self.async_lookup(request, now))
|
||||
}
|
||||
|
||||
fn store<'a>(
|
||||
&'a self,
|
||||
request: &'a ResponseCacheRequest,
|
||||
response: Value,
|
||||
now: Duration,
|
||||
) -> CacheFuture<'a, ()> {
|
||||
Box::pin(self.async_store(request, response, now))
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq)]
|
||||
pub enum CacheScope {
|
||||
Shared,
|
||||
Isolated(String),
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct CacheOptions {
|
||||
pub caching: Option<bool>,
|
||||
pub no_cache: bool,
|
||||
pub no_store: bool,
|
||||
pub ttl: Option<Duration>,
|
||||
pub max_age: Option<Duration>,
|
||||
pub scope: CacheScope,
|
||||
}
|
||||
|
||||
impl CacheOptions {
|
||||
pub fn new(scope: CacheScope) -> Self {
|
||||
Self {
|
||||
caching: None,
|
||||
no_cache: false,
|
||||
no_store: false,
|
||||
ttl: None,
|
||||
max_age: None,
|
||||
scope,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn enabled(&self) -> bool {
|
||||
self.caching != Some(false) && !(self.no_cache && self.no_store)
|
||||
}
|
||||
|
||||
pub fn request(self, namespace: &str, surface: &str, mut input: Value) -> ResponseCacheRequest {
|
||||
input.sort_all_objects();
|
||||
let scope = match self.scope {
|
||||
CacheScope::Shared => String::new(),
|
||||
CacheScope::Isolated(scope) => serde_json::json!(["isolated", scope]).to_string(),
|
||||
};
|
||||
ResponseCacheRequest {
|
||||
key: CacheKeyInput {
|
||||
namespace: Some(format!("{namespace}:inference-v2")),
|
||||
fields: [
|
||||
("surface", surface.to_owned()),
|
||||
("scope", scope),
|
||||
("request", input.to_string()),
|
||||
]
|
||||
.into_iter()
|
||||
.map(|(name, value)| CacheKeyField {
|
||||
name: name.into(),
|
||||
value: Some(value),
|
||||
api_parameter: true,
|
||||
internal_parameter: false,
|
||||
})
|
||||
.collect(),
|
||||
..Default::default()
|
||||
},
|
||||
controls: CacheControls {
|
||||
configured: true,
|
||||
supported_call_type: true,
|
||||
native_backend: true,
|
||||
default_on: true,
|
||||
caching: self.caching,
|
||||
no_cache: self.no_cache,
|
||||
no_store: self.no_store,
|
||||
..Default::default()
|
||||
},
|
||||
context: ExactCacheContext { ttl: self.ttl },
|
||||
max_age: self.max_age,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(serde::Serialize, serde::Deserialize)]
|
||||
pub struct ResponseEnvelope<T> {
|
||||
version: u32,
|
||||
surface: String,
|
||||
output: T,
|
||||
}
|
||||
|
||||
impl<T> ResponseEnvelope<T> {
|
||||
pub fn new(surface: &str, output: T) -> Self {
|
||||
Self {
|
||||
version: 1,
|
||||
surface: surface.into(),
|
||||
output,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn decode(self, surface: &str) -> Option<T> {
|
||||
(self.version == 1 && self.surface == surface).then_some(self.output)
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct ScopedCache {
|
||||
pub service: std::sync::Arc<dyn ResponseCacheService>,
|
||||
pub scope: CacheScope,
|
||||
}
|
||||
|
||||
impl ScopedCache {
|
||||
pub fn new(service: std::sync::Arc<dyn ResponseCacheService>, scope: CacheScope) -> Self {
|
||||
Self { service, scope }
|
||||
}
|
||||
|
||||
pub fn options(&self, overrides: Option<CacheOptions>) -> CacheOptions {
|
||||
overrides.unwrap_or_else(|| CacheOptions::new(self.scope.clone()))
|
||||
}
|
||||
}
|
||||
|
|
@ -18,7 +18,7 @@ use litellm_cache_response::{
|
|||
WriteBuffer, cache_key,
|
||||
};
|
||||
use redis_test::MockCmd;
|
||||
use rstest::rstest;
|
||||
use rstest::{fixture, rstest};
|
||||
use serde_json::{Value, json};
|
||||
use support::{keyed, memory, redis, request};
|
||||
|
||||
|
|
@ -648,3 +648,129 @@ async fn write_buffer_clear_drops_pending_entries(memory: Memory, request: Respo
|
|||
assert_eq!(memory.lookup(&request, now).unwrap(), None);
|
||||
assert_eq!(memory.lookup(&other, now).unwrap(), None);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::python_sync("{'timestamp': 100.0, 'response': '{\"answer\": 7}'}")]
|
||||
#[case::python_async(r#"{"timestamp":100.0,"response":{"answer":7}}"#)]
|
||||
#[case::bare_response(r#"{"answer":7}"#)]
|
||||
#[tokio::test]
|
||||
async fn gcs_reads_python_entries_and_writes_python_compatible_envelopes(
|
||||
#[case] encoded: &str,
|
||||
#[values(false, true)] asynchronous: bool,
|
||||
#[future(awt)] gcs: (wiremock::MockServer, Gcs),
|
||||
) {
|
||||
use wiremock::{
|
||||
Mock, ResponseTemplate,
|
||||
matchers::{body_json, header, method, path, query_param},
|
||||
};
|
||||
|
||||
let (server, cache) = gcs;
|
||||
let response = json!({"answer": 7});
|
||||
Mock::given(method("GET"))
|
||||
.and(path("/storage/v1/b/bucket/o/cache%2Fpython"))
|
||||
.and(query_param("alt", "media"))
|
||||
.and(header("authorization", "Bearer token"))
|
||||
.respond_with(ResponseTemplate::new(200).set_body_string(encoded))
|
||||
.expect(1)
|
||||
.mount(&server)
|
||||
.await;
|
||||
Mock::given(method("POST"))
|
||||
.and(path("/upload/storage/v1/b/bucket/o"))
|
||||
.and(query_param("uploadType", "media"))
|
||||
.and(query_param("name", "cache/native"))
|
||||
.and(header("authorization", "Bearer token"))
|
||||
.and(header("content-type", "application/json"))
|
||||
.and(body_json(json!({"timestamp": 102.0, "response": response})))
|
||||
.respond_with(ResponseTemplate::new(200))
|
||||
.expect(1)
|
||||
.mount(&server)
|
||||
.await;
|
||||
let lookup = if asynchronous {
|
||||
cache
|
||||
.async_lookup(&keyed("python"), Duration::from_secs(102))
|
||||
.await
|
||||
} else {
|
||||
cache.lookup(&keyed("python"), Duration::from_secs(102))
|
||||
};
|
||||
assert_eq!(lookup.unwrap(), Some(response.clone()));
|
||||
let request = ResponseCacheRequest {
|
||||
context: litellm_cache::ExactCacheContext {
|
||||
ttl: Some(Duration::from_secs(12)),
|
||||
},
|
||||
..keyed("native")
|
||||
};
|
||||
let stored = if asynchronous {
|
||||
cache
|
||||
.async_store(&request, response, Duration::from_secs(102))
|
||||
.await
|
||||
} else {
|
||||
cache.store(&request, response, Duration::from_secs(102))
|
||||
};
|
||||
assert_eq!(stored, Ok(()));
|
||||
let requests = server.received_requests().await.unwrap();
|
||||
let upload = requests
|
||||
.iter()
|
||||
.find(|request| request.method.as_str() == "POST")
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
upload.url.query(),
|
||||
Some("uploadType=media&name=cache%2Fnative")
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn gcs_batch_reads_preserve_order_and_treat_invalid_entries_as_misses(
|
||||
#[future(awt)] gcs: (wiremock::MockServer, Gcs),
|
||||
) {
|
||||
use wiremock::{
|
||||
Mock, ResponseTemplate,
|
||||
matchers::{method, path},
|
||||
};
|
||||
|
||||
let (server, cache) = gcs;
|
||||
Mock::given(method("GET"))
|
||||
.and(path("/storage/v1/b/bucket/o/cache%2Fhit"))
|
||||
.respond_with(
|
||||
ResponseTemplate::new(200)
|
||||
.set_body_json(json!({"timestamp": 100.0, "response": {"answer":7}})),
|
||||
)
|
||||
.mount(&server)
|
||||
.await;
|
||||
Mock::given(method("GET"))
|
||||
.and(path("/storage/v1/b/bucket/o/cache%2Finvalid"))
|
||||
.respond_with(ResponseTemplate::new(200).set_body_string("not an entry"))
|
||||
.mount(&server)
|
||||
.await;
|
||||
Mock::given(method("GET"))
|
||||
.and(path("/storage/v1/b/bucket/o/cache%2Fmissing"))
|
||||
.respond_with(ResponseTemplate::new(404))
|
||||
.mount(&server)
|
||||
.await;
|
||||
let requests = [keyed("hit"), keyed("missing"), keyed("invalid")];
|
||||
let partial = cache
|
||||
.async_lookup_batch(&requests, Duration::from_secs(102))
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(partial.values, vec![Some(json!({"answer":7})), None, None]);
|
||||
assert_eq!(partial.missing_indices, vec![1, 2]);
|
||||
}
|
||||
|
||||
type Gcs = ResponseCache<litellm_cache_gcs::GcsCache<litellm_cache_response::ResponseCacheCodec>>;
|
||||
|
||||
#[fixture]
|
||||
async fn gcs() -> (wiremock::MockServer, Gcs) {
|
||||
let server = wiremock::MockServer::start().await;
|
||||
let cache = ResponseCache::new(Arc::new(litellm_cache_gcs::GcsCache::with_token_source(
|
||||
litellm_cache_gcs::GcsConfig {
|
||||
bucket_name: "bucket".into(),
|
||||
gcs_path: Some("cache".into()),
|
||||
path_service_account: None,
|
||||
endpoint: server.uri(),
|
||||
},
|
||||
litellm_http::Client::plain_for_test(),
|
||||
litellm_cache_response::ResponseCacheCodec,
|
||||
Arc::new(litellm_cache_gcs::StaticTokenSource("token".into())),
|
||||
)));
|
||||
(server, cache)
|
||||
}
|
||||
|
|
|
|||
176
litellm-rust/crates/cache-response/tests/service.rs
Normal file
176
litellm-rust/crates/cache-response/tests/service.rs
Normal file
|
|
@ -0,0 +1,176 @@
|
|||
use std::{
|
||||
sync::{
|
||||
Arc,
|
||||
atomic::{AtomicU64, Ordering},
|
||||
},
|
||||
time::Duration,
|
||||
};
|
||||
|
||||
use litellm_cache::ExactCacheContext;
|
||||
use litellm_cache_memory::InMemoryCache;
|
||||
use litellm_cache_response::{
|
||||
CacheEntry, CacheKeyInput, ResponseCache, ResponseCacheConfig, ResponseCacheRequest,
|
||||
ResponseCacheService,
|
||||
};
|
||||
use rstest::rstest;
|
||||
use serde_json::json;
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn service_honors_per_call_expiry_and_freshness() {
|
||||
let clock = Arc::new(AtomicU64::new(0));
|
||||
let cache_clock = clock.clone();
|
||||
let cache: Arc<dyn ResponseCacheService> = Arc::new(ResponseCache::new(Arc::new(
|
||||
InMemoryCache::with_clock(Some(100), Some(Duration::from_secs(60)), move || {
|
||||
Duration::from_secs(cache_clock.load(Ordering::SeqCst))
|
||||
}),
|
||||
)));
|
||||
let request = ResponseCacheRequest {
|
||||
context: ExactCacheContext {
|
||||
ttl: Some(Duration::from_secs(5)),
|
||||
},
|
||||
..ResponseCacheRequest::new(CacheKeyInput {
|
||||
preset: Some("entry".into()),
|
||||
..Default::default()
|
||||
})
|
||||
};
|
||||
cache
|
||||
.store(&request, json!({"answer":7}), Duration::ZERO)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
cache.lookup(&request, Duration::ZERO).await.unwrap(),
|
||||
Some(json!({"answer":7}))
|
||||
);
|
||||
let stale_request = ResponseCacheRequest {
|
||||
max_age: Some(Duration::from_secs(1)),
|
||||
..request.clone()
|
||||
};
|
||||
clock.store(2, Ordering::SeqCst);
|
||||
assert_eq!(
|
||||
cache
|
||||
.lookup(&stale_request, Duration::from_secs(2))
|
||||
.await
|
||||
.unwrap(),
|
||||
None
|
||||
);
|
||||
assert!(
|
||||
cache
|
||||
.lookup(&request, Duration::from_secs(2))
|
||||
.await
|
||||
.unwrap()
|
||||
.is_some()
|
||||
);
|
||||
clock.store(6, Ordering::SeqCst);
|
||||
assert_eq!(
|
||||
cache
|
||||
.lookup(&request, Duration::from_secs(6))
|
||||
.await
|
||||
.unwrap(),
|
||||
None
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn entry_limit_applies_to_sync_async_and_batch_writes() {
|
||||
let storage = Arc::new(InMemoryCache::<CacheEntry>::default());
|
||||
let cache = ResponseCache::new(storage.clone()).with_config(ResponseCacheConfig {
|
||||
namespace: "service-test".into(),
|
||||
max_entry_bytes: json!({"answer":7}).to_string().len(),
|
||||
});
|
||||
let small = json!({"answer":7});
|
||||
let large = json!({"answer":"too large"});
|
||||
let request = |key: &str| {
|
||||
ResponseCacheRequest::new(CacheKeyInput {
|
||||
preset: Some(key.into()),
|
||||
..Default::default()
|
||||
})
|
||||
};
|
||||
cache
|
||||
.store(&request("sync"), large.clone(), Duration::ZERO)
|
||||
.unwrap();
|
||||
cache
|
||||
.async_store(&request("async"), large.clone(), Duration::ZERO)
|
||||
.await
|
||||
.unwrap();
|
||||
cache
|
||||
.async_store_batch(
|
||||
vec![
|
||||
(request("batch-large"), large),
|
||||
(request("batch-small"), small.clone()),
|
||||
],
|
||||
Duration::ZERO,
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
let service: Arc<dyn ResponseCacheService> = Arc::new(cache);
|
||||
service
|
||||
.store(&request("service"), small.clone(), Duration::ZERO)
|
||||
.await
|
||||
.unwrap();
|
||||
for key in ["sync", "async", "batch-large"] {
|
||||
assert!(storage.get_cache(key).unwrap().is_none());
|
||||
}
|
||||
for key in ["batch-small", "service"] {
|
||||
assert_eq!(
|
||||
service.lookup(&request(key), Duration::ZERO).await.unwrap(),
|
||||
Some(small.clone())
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::same_scope("tenant-a", "tenant-a", true)]
|
||||
#[case::different_scope("tenant-a", "tenant-b", false)]
|
||||
#[case::empty_isolated_scope("", "", true)]
|
||||
#[tokio::test]
|
||||
async fn isolated_policy_controls_actual_entry_reuse(
|
||||
#[case] first: &str,
|
||||
#[case] second: &str,
|
||||
#[case] hit: bool,
|
||||
) {
|
||||
use litellm_cache_response::{CacheOptions, CacheScope};
|
||||
let service = ResponseCache::new(Arc::new(InMemoryCache::<CacheEntry>::default()));
|
||||
let request =
|
||||
|scope| CacheOptions::new(scope).request("test", "messages", json!({"prompt":"hello"}));
|
||||
service
|
||||
.async_store(
|
||||
&request(CacheScope::Isolated(first.into())),
|
||||
json!({"answer":7}),
|
||||
Duration::ZERO,
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
service
|
||||
.async_lookup(
|
||||
&request(CacheScope::Isolated(second.into())),
|
||||
Duration::ZERO
|
||||
)
|
||||
.await
|
||||
.unwrap(),
|
||||
hit.then(|| json!({"answer":7}))
|
||||
);
|
||||
assert_eq!(
|
||||
service
|
||||
.async_lookup(&request(CacheScope::Shared), Duration::ZERO)
|
||||
.await
|
||||
.unwrap(),
|
||||
None
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::valid(1, "messages", Some(7))]
|
||||
#[case::unknown_version(2, "messages", None)]
|
||||
#[case::another_surface(1, "responses", None)]
|
||||
fn envelopes_require_a_matching_surface_and_version(
|
||||
#[case] version: u32,
|
||||
#[case] surface: &str,
|
||||
#[case] expected: Option<u32>,
|
||||
) {
|
||||
let envelope: litellm_cache_response::ResponseEnvelope<u32> =
|
||||
serde_json::from_value(json!({"version":version,"surface":surface,"output":7})).unwrap();
|
||||
assert_eq!(envelope.decode("messages"), expected);
|
||||
}
|
||||
|
|
@ -6,6 +6,7 @@ license.workspace = true
|
|||
repository.workspace = true
|
||||
|
||||
[dependencies]
|
||||
litellm-http.workspace = true
|
||||
litellm-cache.workspace = true
|
||||
litellm-auth-aws.workspace = true
|
||||
aws-sdk-s3 = { version = "1.146.1", default-features = false, features = ["rustls", "rt-tokio"] }
|
||||
|
|
@ -19,6 +20,7 @@ reqwest.workspace = true
|
|||
tokio.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
litellm-http = { workspace = true, features = ["test-support"] }
|
||||
litellm-cache-testing.workspace = true
|
||||
rstest.workspace = true
|
||||
wiremock = "0.6.5"
|
||||
|
|
|
|||
|
|
@ -2,10 +2,11 @@ use aws_credential_types::{
|
|||
Credentials as AwsCredentials,
|
||||
provider::{ProvideCredentials, error::CredentialsError, future},
|
||||
};
|
||||
use litellm_auth_aws::{AwsAuthConfig, resolve_credentials};
|
||||
use litellm_auth_aws::{AwsAuthConfig, AwsAuthService};
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct S3Credentials {
|
||||
auth: AwsAuthService,
|
||||
config: AwsAuthConfig,
|
||||
env: fn(&str) -> Option<String>,
|
||||
}
|
||||
|
|
@ -16,7 +17,11 @@ impl S3Credentials {
|
|||
}
|
||||
|
||||
pub fn with_env(config: AwsAuthConfig, env: fn(&str) -> Option<String>) -> Self {
|
||||
Self { config, env }
|
||||
Self {
|
||||
auth: AwsAuthService::default(),
|
||||
config,
|
||||
env,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -38,7 +43,8 @@ impl ProvideCredentials for S3Credentials {
|
|||
"litellm-s3-cache",
|
||||
));
|
||||
}
|
||||
resolve_credentials(self.config.clone(), &self.env)
|
||||
self.auth
|
||||
.resolve_credentials(self.config.clone(), &self.env)
|
||||
.await
|
||||
.map_err(|_| CredentialsError::provider_error("S3 cache authentication failed"))
|
||||
})
|
||||
|
|
|
|||
|
|
@ -42,7 +42,12 @@ pub struct S3Cache<C: CacheCodec> {
|
|||
}
|
||||
|
||||
impl<C: CacheCodec> S3Cache<C> {
|
||||
pub fn new(config: S3CacheConfig, http: reqwest::Client, codec: C, runtime: Handle) -> Self {
|
||||
pub fn new(
|
||||
config: S3CacheConfig,
|
||||
http: litellm_http::Client,
|
||||
codec: C,
|
||||
runtime: Handle,
|
||||
) -> Self {
|
||||
let endpoint_url: Option<String> = config.endpoint.map(|endpoint| endpoint.url);
|
||||
let base = aws_sdk_s3::Config::builder()
|
||||
.behavior_version(BehaviorVersion::latest())
|
||||
|
|
|
|||
|
|
@ -9,7 +9,7 @@ use aws_smithy_runtime_api::client::{
|
|||
use aws_smithy_types::body::SdkBody;
|
||||
|
||||
#[derive(Clone, Debug)]
|
||||
pub(crate) struct ReqwestHttpClient(pub(crate) reqwest::Client);
|
||||
pub(crate) struct ReqwestHttpClient(pub(crate) litellm_http::Client);
|
||||
|
||||
impl HttpClient for ReqwestHttpClient {
|
||||
fn http_connector(
|
||||
|
|
|
|||
|
|
@ -32,7 +32,12 @@ pub fn config(endpoint: &str) -> S3CacheConfig {
|
|||
}
|
||||
|
||||
pub fn cache_with(config: S3CacheConfig, runtime: Handle) -> JsonS3Cache {
|
||||
S3Cache::new(config, reqwest::Client::new(), JsonCodec::new(), runtime)
|
||||
S3Cache::new(
|
||||
config,
|
||||
litellm_http::Client::plain_for_test(),
|
||||
JsonCodec::new(),
|
||||
runtime,
|
||||
)
|
||||
}
|
||||
|
||||
pub fn cache(endpoint: &str) -> JsonS3Cache {
|
||||
|
|
|
|||
|
|
@ -1,19 +1,19 @@
|
|||
- Target invariants, not completion claims
|
||||
- Keep this crate the legacy `@client` wrapper as the native call sees it, and nothing else: the `Logging` contract (`function_setup`, the deployment hooks, `pre_call`/`post_call`, the sync and async success and failure fan-out, the deferred proxy release, the argument sharing those callbacks rely on) plus the kwargs rewrites the wrapper makes on the way in (credential-name inheritance, the budget and retry-count limits)
|
||||
- The driver in `litellm-host-python`, the routes and core see one `PythonLifecycle`; they never learn which Python objects consume a call
|
||||
- Rust drives the call; every litellm Python internal it still borrows is a variant of `LegacyPython`, grouped by subsystem (`Wrapper`, `Logging`, `DeploymentHooks`)
|
||||
- This crate owns compatibility for all existing Python callbacks and loggers, including `CustomLogger`. `mapping.rs` owns the executable call bindings and the inventory of Python-owned hooks. A Python-owned entry records an existing path, never permission to invoke it a second time. The native call adapter preserves the `Logging` contract (`function_setup`, the deployment hooks, `pre_call`/`post_call`, the sync and async success and failure fan-out, the deferred proxy release, the argument sharing those callbacks rely on)
|
||||
- Smell test: if a future callback host (`callbacks-v1-python`, WASM, in-process Rust) could share a piece of this crate, it does not belong here
|
||||
- SDK request policy (credential inheritance, the budget and retry-count limits) is a separate hook supplied by `python-bridge`; compose it after this adapter so logging adopts the final keyword view before policy mutates or rejects it
|
||||
- The driver in `litellm-host-python`, the routes and core see one `PythonCallHooks` using the shared `CallEvent`; they never learn which Python objects consume a call
|
||||
- Every litellm Python internal Rust still borrows is a variant of `LegacyPython`, grouped by subsystem, with its signature pinned in `python_contract.json`
|
||||
- The enum only shrinks: when Rust owns a subsystem, delete its group rather than adding a Rust path beside it
|
||||
- Calling a user's own callback directly is permanent Python surface and gets its own type outside `LegacyPython`
|
||||
- `PublicCall` is the caller's call as `Logging` sees it: the positional arguments, the keyword view as the legacy path rewrites it (setup, deployment hook, prepare) and the bound request object whose attributes back keywords the caller omitted; routes hand it over through `run_legacy_call` and keep no copy
|
||||
- `setup` reuses a `Logging` the caller passed as `litellm_logging_obj` (the proxy and Router are the live cases) and otherwise builds one through `function_setup`, as `@client` does
|
||||
- Either way every phase calls the same `Logging` method the Python path calls; which callbacks run is `Logging`'s decision, never this crate's
|
||||
- `PublicCall` is the caller's call as `Logging` sees it: the positional arguments, the keyword view as the call rewrites it (setup, deployment hook, preflight) and the bound request object backing omitted keywords; shared bridge composition hands it to `LegacyLogging`; routes use the neutral call boundary
|
||||
- `setup` reuses a `Logging` passed as `litellm_logging_obj` (the proxy and Router) and otherwise builds one through `function_setup`; which callbacks run is `Logging`'s decision, never this crate's
|
||||
- Callbacks receive the caller's own objects and may mutate them; this crate alone carries that obligation
|
||||
- Retain complete boundary arguments, opaque unknown values, aliases, omitted/default distinctions and deliberate copies; preserve the established deployment-hook kwargs view
|
||||
- Before `pre_call`, re-alias every body key whose value equals the caller's argument to the caller's own object; this crate compares the two itself, and the argument is resolved by `litellm_host_python::lookup`
|
||||
- Retain independently captured body/header roots from `pre_call` to `post_call`; in-place mutation reaches the wire, envelope field replacement is visible to later callbacks only
|
||||
- A later kind of callback host (WASM, in-process Rust) has none of these obligations, so they stay out of `litellm-host`, `litellm-host-python` and the bridge; the only fact that crosses from the route is the prepared keyword view
|
||||
- Success and failure handlers receive the exact selected public response or exception; logging projections, redaction and snapshots keep their own copy contracts
|
||||
- Ordinary failure-handler errors cannot suppress the other eligible family or replace the mapped provider error; a cancellation ends the call with no further dispatch
|
||||
- Dispatch errors never replay provider work or trigger the opposite outcome; the proxy's acceptance or rejection releases deferred success at most once
|
||||
- Retain complete boundary arguments, opaque values, aliases, omitted/default distinctions and deliberate copies; preserve the deployment-hook kwargs view
|
||||
- Before `pre_call`, re-alias every body key whose value equals the caller's argument to the caller's own object, resolved through `litellm_host_python::lookup`
|
||||
- Retain body/header roots from `pre_call` to `post_call`; in-place mutation reaches the wire, envelope field replacement is visible to later callbacks only
|
||||
- Success and failure handlers receive the exact selected public response or exception
|
||||
- A failure-handler error cannot suppress the other eligible family or replace the mapped provider error; a cancellation ends the call with no further dispatch
|
||||
- Dispatch errors never replay provider work or trigger the opposite outcome; the proxy releases deferred success at most once
|
||||
- Delivery follows the registry, not the callable's type: direct, awaited, executor-submitted, logging-worker and deferred paths stay distinct
|
||||
- Traverse every retained Python edge; `close` is idempotent and restores the correlation context once
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ license.workspace = true
|
|||
repository.workspace = true
|
||||
|
||||
[dependencies]
|
||||
litellm-types.workspace = true
|
||||
litellm-host.workspace = true
|
||||
litellm-host-python.workspace = true
|
||||
|
||||
|
|
|
|||
|
|
@ -6,9 +6,6 @@
|
|||
"start_time",
|
||||
"asynchronous"
|
||||
],
|
||||
"check_limits": [
|
||||
"kwargs"
|
||||
],
|
||||
"finalize": [
|
||||
"response",
|
||||
"logger",
|
||||
|
|
@ -76,11 +73,6 @@
|
|||
],
|
||||
"custom_pricing_fields": [],
|
||||
"is_internal_call": [],
|
||||
"credential_list": [],
|
||||
"warn_unknown_credential": [
|
||||
"name",
|
||||
"loaded"
|
||||
],
|
||||
"before_deployment_call": [
|
||||
"kwargs",
|
||||
"call_type"
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -3,16 +3,13 @@
|
|||
//! lifetime. No other callback host has that obligation, which is why nothing outside
|
||||
//! this crate holds them.
|
||||
|
||||
use litellm_host::{machine::Machine, protocol::Protocol};
|
||||
use litellm_host_python::{ProtocolHost, lookup, run_call};
|
||||
use litellm_host_python::lookup;
|
||||
use pyo3::{
|
||||
gc::{PyTraverseError, PyVisit},
|
||||
prelude::*,
|
||||
types::{PyDict, PyTuple},
|
||||
};
|
||||
|
||||
use crate::{LegacyLogging, LegacySurface};
|
||||
|
||||
pub struct PublicCall {
|
||||
args: Py<PyTuple>,
|
||||
kwargs: Py<PyDict>,
|
||||
|
|
@ -34,12 +31,17 @@ impl PublicCall {
|
|||
})
|
||||
}
|
||||
|
||||
pub fn arguments(&self, py: Python<'_>) -> Py<PyDict> {
|
||||
self.kwargs.clone_ref(py)
|
||||
}
|
||||
|
||||
pub(crate) fn args(&self) -> &Py<PyTuple> {
|
||||
&self.args
|
||||
}
|
||||
|
||||
/// The keyword view the legacy path currently reads: the caller's copy until
|
||||
/// `function_setup`, then each rewrite (setup, deployment hook, prepare) in turn.
|
||||
/// `function_setup`, then each rewrite (setup, deployment hook, the driver's preflight)
|
||||
/// in turn.
|
||||
pub(crate) fn kwargs(&self) -> &Py<PyDict> {
|
||||
&self.kwargs
|
||||
}
|
||||
|
|
@ -63,31 +65,6 @@ impl PublicCall {
|
|||
}
|
||||
}
|
||||
|
||||
/// Runs one native call under the legacy `Logging` contract: the protocol host projects from
|
||||
/// the keyword view the contract prepares, and the contract observes the call.
|
||||
pub fn run_legacy_call<H, M>(
|
||||
py: Python<'_>,
|
||||
surface: LegacySurface,
|
||||
call: PublicCall,
|
||||
machine: M,
|
||||
host: H,
|
||||
asynchronous: bool,
|
||||
) -> PyResult<Py<PyAny>>
|
||||
where
|
||||
H: ProtocolHost + 'static,
|
||||
M: Machine<Protocol = H::Protocol, Complete = <H::Protocol as Protocol>::Response> + 'static,
|
||||
{
|
||||
let arguments = call.kwargs.clone_ref(py);
|
||||
run_call(
|
||||
py,
|
||||
machine,
|
||||
host,
|
||||
Box::new(LegacyLogging::new(py, surface, call, asynchronous)),
|
||||
arguments,
|
||||
asynchronous,
|
||||
)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
|
@ -106,7 +83,7 @@ mod tests {
|
|||
(call, locals)
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[rstest::rstest]
|
||||
fn capture_copies_the_keyword_dict_without_copying_its_values() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
|
|
|
|||
|
|
@ -2,7 +2,7 @@
|
|||
//! the deferred and worker-submitted success paths, and the sync-callbacks-for-async-calls
|
||||
//! duplication. All of it expires with the legacy callback contract.
|
||||
|
||||
use litellm_host::event::{RequestContext, WireRequest};
|
||||
use litellm_host::interceptors::{RequestContext, WireRequest};
|
||||
use litellm_host_python::to_py;
|
||||
use pyo3::{exceptions::PyBaseException, prelude::*, types::PyDict};
|
||||
|
||||
|
|
|
|||
|
|
@ -1,233 +1,24 @@
|
|||
//! The legacy `@client` wrapper as the native call sees it: litellm's `Logging` object, the
|
||||
//! sync and async callback registries it fans out to, the deployment hooks, the deferred
|
||||
//! proxy release, and the kwargs rewrites the wrapper makes on the way in (credential-name
|
||||
//! inheritance, budget and retry-count limits). All of it sits behind one
|
||||
//! [`PythonLifecycle`](litellm_host_python::PythonLifecycle), so the driver, the routes and
|
||||
//! sync and async callback registries it fans out to, the deployment hooks and the deferred
|
||||
//! proxy release. All of it sits behind one
|
||||
//! [`PythonCallHooks`](litellm_host_python::PythonCallHooks), so the driver, the routes and
|
||||
//! core never learn which Python object is on the other end.
|
||||
//!
|
||||
//! Legacy callbacks receive the caller's own objects and may mutate them. [`PublicCall`]
|
||||
//! is where those objects live, and [`run_legacy_call`] is how a route hands them over
|
||||
//! without keeping a copy.
|
||||
//! is where those objects live.
|
||||
|
||||
mod adapter;
|
||||
mod call;
|
||||
mod callbacks;
|
||||
mod deferred;
|
||||
mod logger;
|
||||
mod preparation;
|
||||
mod mapping;
|
||||
mod python;
|
||||
pub(crate) use adapter::LegacyLogging;
|
||||
pub use adapter::{LegacySurface, PassThroughStream};
|
||||
pub use call::{PublicCall, run_legacy_call};
|
||||
pub use adapter::LegacyLogging;
|
||||
pub use call::PublicCall;
|
||||
pub(crate) use callbacks::{LegacyCallbacks, is_internal_call};
|
||||
pub(crate) use logger::{DeploymentHooks, PythonLogger, finalize, setup};
|
||||
pub(crate) use preparation::prepare;
|
||||
pub use mapping::{CallBoundary, CallbackMapping, Dispatch, callback_mappings};
|
||||
|
||||
#[cfg(test)]
|
||||
mod test_support {
|
||||
use std::ffi::CStr;
|
||||
|
||||
use pyo3::prelude::*;
|
||||
use pyo3::types::{PyDict, PyTuple};
|
||||
|
||||
use crate::{LegacyLogging, LegacySurface, PublicCall};
|
||||
|
||||
/// The parameters of every `callbacks_legacy_python` function, as the real module declares them.
|
||||
/// `tests/test_litellm/rust_bridge/test_callbacks_legacy_python.py` pins this file to the Python
|
||||
/// signatures, and [`namespace`] binds every fake call against it.
|
||||
pub(crate) const PYTHON_CONTRACT: &str = include_str!("../python_contract.json");
|
||||
|
||||
/// Stand-ins for `callbacks_legacy_python`, the only Python module the crate calls. Tests
|
||||
/// share one interpreter and run concurrently, so each fake is installed idempotently and
|
||||
/// forwards to the per-test `StubLogger` it is handed (directly, or as `kwargs['logger']`).
|
||||
/// Every fake is bound against the contract first, so a call the real module would reject
|
||||
/// fails here too.
|
||||
const STUBS: &CStr = c"
|
||||
import contextvars
|
||||
import inspect
|
||||
import json
|
||||
import sys
|
||||
import traceback
|
||||
import types
|
||||
|
||||
for name in ('litellm', 'litellm.rust_bridge', 'litellm.rust_bridge.callbacks_legacy_python'):
|
||||
sys.modules.setdefault(name, types.ModuleType(name))
|
||||
|
||||
legacy = sys.modules['litellm.rust_bridge.callbacks_legacy_python']
|
||||
CONTRACT = json.loads(python_contract)
|
||||
|
||||
|
||||
def contracted(name, fake):
|
||||
signature = inspect.Signature(
|
||||
[inspect.Parameter(parameter, inspect.Parameter.POSITIONAL_OR_KEYWORD) for parameter in CONTRACT[name]]
|
||||
)
|
||||
|
||||
def checked(*args, **kwargs):
|
||||
signature.bind(*args, **kwargs)
|
||||
return fake(*args, **kwargs)
|
||||
|
||||
return checked
|
||||
|
||||
|
||||
if not hasattr(legacy, 'is_internal'):
|
||||
legacy.is_internal = contextvars.ContextVar('is_internal_call', default=False)
|
||||
|
||||
FAKES = {
|
||||
'setup': lambda call_type, args, kwargs, start, asynchronous: types.SimpleNamespace(
|
||||
logger=kwargs['logger_factory'](kwargs) if 'logger_factory' in kwargs else kwargs['logger'],
|
||||
kwargs=kwargs,
|
||||
),
|
||||
'check_limits': lambda arguments: arguments['logger'].check_limits(arguments),
|
||||
'finalize': lambda response, logger, kwargs, start, end: logger.record('finalize', response),
|
||||
'update_logging': lambda logger, kwargs, model, optional_params, litellm_params, provider: logger.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model=model,
|
||||
optional_params=optional_params,
|
||||
litellm_params=litellm_params,
|
||||
custom_llm_provider=provider,
|
||||
),
|
||||
'pre_call': lambda logger, input, api_key, additional_args: logger.pre_call(input, api_key, additional_args),
|
||||
'post_call': lambda logger, original_response, api_key, additional_args: logger.post_call(
|
||||
original_response, api_key, additional_args
|
||||
),
|
||||
'defers_async_logging': lambda logger: bool(getattr(logger, '_defer_async_logging', False)),
|
||||
'defer_success': lambda logger, pending: setattr(logger, '_native_pending_logging', pending),
|
||||
'sync_success_for_async_call': lambda logger, response, start, end: logger.handle_sync_success_callbacks_for_async_calls(
|
||||
response, start, end
|
||||
),
|
||||
'failure_handler': lambda logger, error, start, end, asynchronous: (
|
||||
logger.async_failure_handler if asynchronous else logger.failure_handler
|
||||
)(error, ''.join(traceback.format_exception(error)), start, end),
|
||||
'submit_success': lambda logger, response, start, end: logger.record('submit', (response, start, end)),
|
||||
'async_success_handler': lambda logger, response, start, end: logger.async_success_handler(response, start, end),
|
||||
'enqueue_logging': lambda coroutine: coroutine.enqueue(),
|
||||
'restore_context': lambda logger: logger.record('restore', None),
|
||||
'custom_pricing_fields': lambda: ('ocr_cost_per_page',),
|
||||
'is_internal_call': lambda: legacy.is_internal.get(),
|
||||
'credential_list': lambda: [],
|
||||
'warn_unknown_credential': lambda name, loaded: None,
|
||||
'before_deployment_call': lambda kwargs, call_type: kwargs['logger'].hook('pre', kwargs, call_type),
|
||||
'after_deployment_success': lambda kwargs, response, call_type: kwargs['logger'].hook(
|
||||
'success', response, call_type
|
||||
),
|
||||
'after_deployment_failure': lambda kwargs, error, call_type: kwargs['logger'].hook('failure', error, call_type),
|
||||
'stream_opened': lambda logger: logger.record('stream_opened', None),
|
||||
'stream_success': lambda logger, request_body, chunks, start, end, first_chunk: logger.record(
|
||||
'stream_success', list(chunks)
|
||||
),
|
||||
'stream_failure': lambda logger, request_body, chunks, error: logger.record('stream_failure', error),
|
||||
}
|
||||
assert FAKES.keys() == CONTRACT.keys(), sorted(FAKES.keys() ^ CONTRACT.keys())
|
||||
for name, fake in FAKES.items():
|
||||
setattr(legacy, name, contracted(name, fake))
|
||||
|
||||
|
||||
unraisable = sys.modules.setdefault(
|
||||
'litellm_test_unraisable', types.ModuleType('litellm_test_unraisable')
|
||||
)
|
||||
if not hasattr(unraisable, 'events'):
|
||||
unraisable.events = []
|
||||
sys.unraisablehook = lambda event: unraisable.events.append((event.object, event.exc_value))
|
||||
|
||||
|
||||
def unraisable_from(owner):
|
||||
return [error for source, error in unraisable.events if source is owner]
|
||||
|
||||
|
||||
class StubCoroutine:
|
||||
def __init__(self, logger):
|
||||
self.logger = logger
|
||||
|
||||
def enqueue(self):
|
||||
self.logger.record('enqueued', None)
|
||||
self.logger.on_enqueue(self)
|
||||
|
||||
def close(self):
|
||||
self.logger.record('closed', None)
|
||||
|
||||
|
||||
class StubLogger:
|
||||
def __init__(self):
|
||||
self.calls = []
|
||||
self.hooks = {}
|
||||
self.on_enqueue = lambda coroutine: None
|
||||
|
||||
def record(self, name, value):
|
||||
self.calls.append((name, value))
|
||||
|
||||
def names(self):
|
||||
return [name for name, _ in self.calls]
|
||||
|
||||
def hook(self, phase, value, call_type):
|
||||
self.record(phase + '_hook', call_type)
|
||||
return self.hooks.get(phase, lambda value: 'awaitable')(value)
|
||||
|
||||
def check_limits(self, arguments):
|
||||
self.record('check_limits', arguments)
|
||||
|
||||
def failure_handler(self, error, trace, start, end):
|
||||
self.record('failure_handler', error)
|
||||
|
||||
def async_failure_handler(self, error, trace, start, end):
|
||||
self.record('async_failure_handler', error)
|
||||
return 'awaitable'
|
||||
|
||||
def success_handler(self, response, start, end):
|
||||
self.record('success_handler', response)
|
||||
|
||||
def async_success_handler(self, response, start, end):
|
||||
self.record('async_success_handler', response)
|
||||
return StubCoroutine(self)
|
||||
|
||||
def handle_sync_success_callbacks_for_async_calls(self, response, start, end):
|
||||
self.record('sync_success_for_async_call', response)
|
||||
|
||||
|
||||
logger = StubLogger()
|
||||
";
|
||||
|
||||
/// A namespace with the stubs, `StubLogger` and a fresh `logger`, after `script` ran in it.
|
||||
pub(crate) fn namespace<'py>(py: Python<'py>, script: &CStr) -> Bound<'py, PyDict> {
|
||||
let locals = PyDict::new(py);
|
||||
locals.set_item("python_contract", PYTHON_CONTRACT).unwrap();
|
||||
py.run(STUBS, Some(&locals), Some(&locals)).unwrap();
|
||||
py.run(script, Some(&locals), Some(&locals)).unwrap();
|
||||
locals
|
||||
}
|
||||
|
||||
pub(crate) fn run(py: Python<'_>, locals: &Bound<'_, PyDict>, code: &CStr) {
|
||||
py.run(code, Some(locals), Some(locals)).unwrap();
|
||||
}
|
||||
|
||||
pub(crate) fn local<'py>(locals: &Bound<'py, PyDict>, name: &str) -> Bound<'py, PyAny> {
|
||||
locals.get_item(name).unwrap().unwrap()
|
||||
}
|
||||
|
||||
/// A legacy call over the namespace's `kwargs` (or none) and `request` (or `None`).
|
||||
pub(crate) fn legacy_call(
|
||||
py: Python<'_>,
|
||||
locals: &Bound<'_, PyDict>,
|
||||
asynchronous: bool,
|
||||
) -> LegacyLogging {
|
||||
let request = locals
|
||||
.get_item("request")
|
||||
.unwrap()
|
||||
.unwrap_or_else(|| py.None().into_bound(py));
|
||||
let kwargs = locals
|
||||
.get_item("kwargs")
|
||||
.unwrap()
|
||||
.map(|kwargs| kwargs.cast_into::<PyDict>().unwrap())
|
||||
.unwrap_or_else(|| PyDict::new(py));
|
||||
let call = PublicCall::capture(&request, &PyTuple::empty(py), &kwargs).unwrap();
|
||||
LegacyLogging::new(
|
||||
py,
|
||||
LegacySurface {
|
||||
call_type: "test",
|
||||
input_description: "test input",
|
||||
stream: None,
|
||||
},
|
||||
call,
|
||||
asynchronous,
|
||||
)
|
||||
}
|
||||
}
|
||||
mod test_support;
|
||||
|
|
|
|||
288
litellm-rust/crates/callbacks-legacy-python/src/mapping.rs
Normal file
288
litellm-rust/crates/callbacks-legacy-python/src/mapping.rs
Normal file
|
|
@ -0,0 +1,288 @@
|
|||
use litellm_host::{
|
||||
hooks::CallHooks,
|
||||
interceptors::{RawResponse, RequestContext, WireRequest},
|
||||
lifecycle::{ExecutionEvent, FailureOrigin, Timing},
|
||||
};
|
||||
use litellm_host_python::{HookStep, PythonCallEvent, PythonRuntime};
|
||||
use pyo3::{prelude::*, types::PyDict};
|
||||
|
||||
use crate::LegacyLogging;
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||||
pub enum CallBoundary {
|
||||
PrepareArguments,
|
||||
BeforeProviderRequest,
|
||||
AfterProviderResponse,
|
||||
TransformResponse,
|
||||
Succeeded,
|
||||
Failed,
|
||||
StreamOpened,
|
||||
StreamChunk,
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||||
pub enum Dispatch {
|
||||
Call(CallBoundary),
|
||||
Python(&'static str),
|
||||
DeclarationOnly,
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||||
pub struct CallbackMapping {
|
||||
pub callback: &'static str,
|
||||
pub dispatch: Dispatch,
|
||||
}
|
||||
|
||||
struct Binding<H> {
|
||||
boundary: CallBoundary,
|
||||
invoke: H,
|
||||
callbacks: &'static [&'static str],
|
||||
}
|
||||
|
||||
impl<H> Binding<H> {
|
||||
fn mappings(&self) -> impl Iterator<Item = CallbackMapping> {
|
||||
self.callbacks.iter().map(|callback| CallbackMapping {
|
||||
callback,
|
||||
dispatch: Dispatch::Call(self.boundary),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
type Step<T> = PyResult<HookStep<LegacyLogging, T>>;
|
||||
type Prepare = fn(&mut LegacyLogging, Python<'_>, Py<PyDict>, f64) -> Step<Py<PyDict>>;
|
||||
type Before =
|
||||
fn(&mut LegacyLogging, Python<'_>, Box<WireRequest>, &RequestContext) -> Step<Box<WireRequest>>;
|
||||
type After = fn(&mut LegacyLogging, Python<'_>, &RawResponse) -> Step<()>;
|
||||
type Transform = fn(&mut LegacyLogging, Python<'_>, Py<PyAny>, Timing) -> Step<Py<PyAny>>;
|
||||
type Success = fn(&mut LegacyLogging, Python<'_>, Timing, &Py<PyAny>) -> Step<()>;
|
||||
type Failure = fn(&mut LegacyLogging, Python<'_>, Timing, FailureOrigin, &PyErr) -> Step<()>;
|
||||
type Open = fn(&mut LegacyLogging, Python<'_>, &Py<PyAny>) -> PyResult<()>;
|
||||
type Chunk = fn(&mut LegacyLogging, Python<'_>, &Py<PyAny>) -> PyResult<()>;
|
||||
|
||||
const PREPARE: Binding<Prepare> = Binding {
|
||||
boundary: CallBoundary::PrepareArguments,
|
||||
invoke: LegacyLogging::prepare_call,
|
||||
callbacks: &["async_pre_call_deployment_hook"],
|
||||
};
|
||||
|
||||
const BEFORE: Binding<Before> = Binding {
|
||||
boundary: CallBoundary::BeforeProviderRequest,
|
||||
invoke: LegacyLogging::pre_call,
|
||||
callbacks: &["log_pre_api_call", "log_input_event"],
|
||||
};
|
||||
|
||||
const AFTER: Binding<After> = Binding {
|
||||
boundary: CallBoundary::AfterProviderResponse,
|
||||
invoke: LegacyLogging::post_call,
|
||||
callbacks: &["log_post_api_call"],
|
||||
};
|
||||
|
||||
const TRANSFORM: Binding<Transform> = Binding {
|
||||
boundary: CallBoundary::TransformResponse,
|
||||
invoke: LegacyLogging::transform_public_response,
|
||||
callbacks: &["async_post_call_success_deployment_hook"],
|
||||
};
|
||||
|
||||
const SUCCESS: Binding<Success> = Binding {
|
||||
boundary: CallBoundary::Succeeded,
|
||||
invoke: LegacyLogging::succeeded,
|
||||
callbacks: &[
|
||||
"log_success_event",
|
||||
"async_log_success_event",
|
||||
"logging_hook",
|
||||
"async_logging_hook",
|
||||
"redact_standard_logging_payload_from_model_call_details",
|
||||
"log_event",
|
||||
"async_log_event",
|
||||
],
|
||||
};
|
||||
|
||||
const FAILURE: Binding<Failure> = Binding {
|
||||
boundary: CallBoundary::Failed,
|
||||
invoke: LegacyLogging::failed,
|
||||
callbacks: &[
|
||||
"async_post_call_failure_deployment_hook",
|
||||
"log_failure_event",
|
||||
"async_log_failure_event",
|
||||
"log_model_group_rate_limit_error",
|
||||
"log_event",
|
||||
"async_log_event",
|
||||
],
|
||||
};
|
||||
|
||||
const OPEN: Binding<Open> = Binding {
|
||||
boundary: CallBoundary::StreamOpened,
|
||||
invoke: LegacyLogging::stream_opened,
|
||||
callbacks: &[],
|
||||
};
|
||||
|
||||
const CHUNK: Binding<Chunk> = Binding {
|
||||
boundary: CallBoundary::StreamChunk,
|
||||
invoke: LegacyLogging::stream_chunk,
|
||||
callbacks: &[],
|
||||
};
|
||||
|
||||
pub fn callback_mappings() -> impl Iterator<Item = CallbackMapping> {
|
||||
PREPARE
|
||||
.mappings()
|
||||
.chain(BEFORE.mappings())
|
||||
.chain(AFTER.mappings())
|
||||
.chain(TRANSFORM.mappings())
|
||||
.chain(SUCCESS.mappings())
|
||||
.chain(FAILURE.mappings())
|
||||
.chain(OPEN.mappings())
|
||||
.chain(CHUNK.mappings())
|
||||
.chain(PYTHON_CALLBACKS.iter().copied())
|
||||
}
|
||||
|
||||
impl CallHooks<PythonRuntime> for LegacyLogging {
|
||||
fn prepare_arguments(
|
||||
&mut self,
|
||||
py: Python<'_>,
|
||||
arguments: Py<PyDict>,
|
||||
started_at: f64,
|
||||
) -> Step<Py<PyDict>> {
|
||||
(PREPARE.invoke)(self, py, arguments, started_at)
|
||||
}
|
||||
|
||||
fn arguments_prepared(&mut self, py: Python<'_>, arguments: &Py<PyDict>) -> PyResult<()> {
|
||||
self.adopt_arguments(py, arguments);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn before_provider_request(
|
||||
&mut self,
|
||||
py: Python<'_>,
|
||||
wire: Box<WireRequest>,
|
||||
context: &RequestContext,
|
||||
) -> Step<Box<WireRequest>> {
|
||||
(BEFORE.invoke)(self, py, wire, context)
|
||||
}
|
||||
|
||||
fn transform_response(
|
||||
&mut self,
|
||||
py: Python<'_>,
|
||||
response: Py<PyAny>,
|
||||
timing: Timing,
|
||||
) -> Step<Py<PyAny>> {
|
||||
(TRANSFORM.invoke)(self, py, response, timing)
|
||||
}
|
||||
|
||||
fn on_event(&mut self, py: Python<'_>, event: PythonCallEvent<'_>) -> Step<()> {
|
||||
match event {
|
||||
PythonCallEvent::Started { .. } | PythonCallEvent::Cancelled { .. } => {
|
||||
Ok(HookStep::Ready(()))
|
||||
}
|
||||
PythonCallEvent::Execution(ExecutionEvent::ResultReady { facts }) => {
|
||||
self.result_ready(py, &facts)
|
||||
}
|
||||
PythonCallEvent::Execution(ExecutionEvent::ProviderResponseReceived { raw }) => {
|
||||
(AFTER.invoke)(self, py, raw)
|
||||
}
|
||||
PythonCallEvent::Succeeded { timing, response } => {
|
||||
(SUCCESS.invoke)(self, py, timing, response)
|
||||
}
|
||||
PythonCallEvent::Failed {
|
||||
timing,
|
||||
origin,
|
||||
error,
|
||||
} => (FAILURE.invoke)(self, py, timing, origin, error),
|
||||
}
|
||||
}
|
||||
|
||||
fn on_stream_open(&mut self, py: Python<'_>, head: &Py<PyAny>) -> PyResult<()> {
|
||||
(OPEN.invoke)(self, py, head)
|
||||
}
|
||||
|
||||
fn on_stream_chunk(&mut self, py: Python<'_>, chunk: &Py<PyAny>) -> PyResult<()> {
|
||||
(CHUNK.invoke)(self, py, chunk)
|
||||
}
|
||||
}
|
||||
|
||||
macro_rules! python_callbacks {
|
||||
($($dispatch:expr => [$($callback:literal),* $(,)?]),* $(,)?) => {
|
||||
const PYTHON_CALLBACKS: &[CallbackMapping] = &[
|
||||
$($(CallbackMapping { callback: $callback, dispatch: $dispatch },)*)*
|
||||
];
|
||||
};
|
||||
}
|
||||
|
||||
python_callbacks! {
|
||||
Dispatch::Python("litellm.router") => [
|
||||
"async_pre_routing_hook",
|
||||
"async_filter_deployments",
|
||||
"pre_call_check",
|
||||
"async_pre_call_check",
|
||||
],
|
||||
Dispatch::Python("litellm.router_utils.fallback_event_handlers") => [
|
||||
"log_success_fallback_event",
|
||||
"log_failure_fallback_event",
|
||||
],
|
||||
Dispatch::Python("litellm.proxy.utils") => [
|
||||
"async_pre_call_hook",
|
||||
"async_post_call_response_headers_hook",
|
||||
"async_post_call_failure_hook",
|
||||
"async_post_call_success_hook",
|
||||
"async_moderation_hook",
|
||||
"async_post_call_streaming_hook",
|
||||
"async_post_call_streaming_iterator_hook",
|
||||
"async_filter_listed_models",
|
||||
],
|
||||
Dispatch::Python("litellm.litellm_core_utils.litellm_logging") => [
|
||||
"async_get_chat_completion_prompt",
|
||||
"get_chat_completion_prompt",
|
||||
"log_stream_event",
|
||||
"async_log_stream_event",
|
||||
"async_post_mcp_tool_call_hook",
|
||||
],
|
||||
Dispatch::Python("litellm.llms.anthropic.pass_through.messages.handler") => [
|
||||
"async_pre_request_hook",
|
||||
],
|
||||
Dispatch::Python("litellm.litellm_core_utils.streaming_handler") => [
|
||||
"async_post_call_streaming_deployment_hook",
|
||||
],
|
||||
Dispatch::Python("litellm.responses.streaming_iterator") => [
|
||||
"async_post_call_streaming_deployment_hook",
|
||||
],
|
||||
Dispatch::Python("litellm.main") => [
|
||||
"translate_completion_input_params",
|
||||
"translate_completion_output_params",
|
||||
"translate_completion_output_params_streaming",
|
||||
],
|
||||
Dispatch::Python("litellm.integrations.argilla") => ["async_dataset_hook"],
|
||||
Dispatch::Python("litellm.proxy.management_helpers.audit_logs") => ["async_log_audit_log_event"],
|
||||
Dispatch::Python("litellm.llms.custom_httpx.llm_http_handler") => [
|
||||
"async_should_run_agentic_loop",
|
||||
"async_run_agentic_loop",
|
||||
"async_build_agentic_loop_plan",
|
||||
"async_post_agentic_loop_response_hook",
|
||||
"async_agentic_loop_cleanup_hook",
|
||||
"async_should_run_chat_completion_agentic_loop",
|
||||
"async_run_chat_completion_agentic_loop",
|
||||
"async_build_chat_completion_agentic_loop_plan",
|
||||
],
|
||||
Dispatch::Python("litellm.litellm_core_utils.chat_completion_agentic_loop") => [
|
||||
"async_should_run_agentic_loop",
|
||||
"async_run_agentic_loop",
|
||||
"async_build_agentic_loop_plan",
|
||||
"async_post_agentic_loop_response_hook",
|
||||
"async_agentic_loop_cleanup_hook",
|
||||
],
|
||||
Dispatch::Python("litellm.llms.openai.openai") => [
|
||||
"async_should_run_chat_completion_agentic_loop",
|
||||
"async_run_chat_completion_agentic_loop",
|
||||
],
|
||||
Dispatch::Python("litellm.proxy.spend_tracking.cold_storage_handler") => [
|
||||
"get_proxy_server_request_from_cold_storage_with_object_key",
|
||||
],
|
||||
Dispatch::Python("litellm.integrations.custom_logger") => [
|
||||
"truncate_standard_logging_payload_content",
|
||||
"redacts_messages_itself",
|
||||
"handle_callback_failure",
|
||||
"get_callback_env_vars",
|
||||
],
|
||||
Dispatch::DeclarationOnly => [
|
||||
"async_log_pre_api_call",
|
||||
"async_log_input_event",
|
||||
],
|
||||
}
|
||||
|
|
@ -19,18 +19,12 @@ pub(crate) enum LegacyPython {
|
|||
Streaming(Streaming),
|
||||
}
|
||||
|
||||
/// The `@client` wrapper around the call: `function_setup`, limits, credentials,
|
||||
/// response metadata and the correlation context.
|
||||
/// The `@client` wrapper around the call: `function_setup`, response metadata and the
|
||||
/// correlation context.
|
||||
#[derive(Clone, Copy, Debug, IntoStaticStr, PartialEq, Eq, VariantArray)]
|
||||
pub(crate) enum Wrapper {
|
||||
#[strum(serialize = "setup")]
|
||||
Setup,
|
||||
#[strum(serialize = "check_limits")]
|
||||
CheckLimits,
|
||||
#[strum(serialize = "credential_list")]
|
||||
CredentialList,
|
||||
#[strum(serialize = "warn_unknown_credential")]
|
||||
WarnUnknownCredential,
|
||||
#[strum(serialize = "is_internal_call")]
|
||||
IsInternalCall,
|
||||
#[strum(serialize = "finalize")]
|
||||
|
|
|
|||
193
litellm-rust/crates/callbacks-legacy-python/src/test_support.rs
Normal file
193
litellm-rust/crates/callbacks-legacy-python/src/test_support.rs
Normal file
|
|
@ -0,0 +1,193 @@
|
|||
use std::ffi::CStr;
|
||||
|
||||
use pyo3::prelude::*;
|
||||
use pyo3::types::{PyDict, PyTuple};
|
||||
|
||||
use crate::{LegacyLogging, PublicCall};
|
||||
|
||||
/// The parameters of every `callbacks_legacy_python` function, as the real module declares them.
|
||||
/// `tests/unit/rust_bridge/test_callbacks_legacy_python.py` pins this file to the Python
|
||||
/// signatures, and [`namespace`] binds every fake call against it.
|
||||
pub(crate) const PYTHON_CONTRACT: &str = include_str!("../python_contract.json");
|
||||
|
||||
/// Stand-ins for `callbacks_legacy_python`, the only Python module the crate calls. Tests
|
||||
/// share one interpreter and run concurrently, so each fake is installed idempotently and
|
||||
/// forwards to the per-test `StubLogger` it is handed (directly, or as `kwargs['logger']`).
|
||||
/// Every fake is bound against the contract first, so a call the real module would reject
|
||||
/// fails here too.
|
||||
const STUBS: &CStr = c"
|
||||
import contextvars
|
||||
import inspect
|
||||
import json
|
||||
import sys
|
||||
import traceback
|
||||
import types
|
||||
|
||||
for name in ('litellm', 'litellm.rust_bridge', 'litellm.rust_bridge.callbacks_legacy_python'):
|
||||
sys.modules.setdefault(name, types.ModuleType(name))
|
||||
|
||||
legacy = sys.modules['litellm.rust_bridge.callbacks_legacy_python']
|
||||
CONTRACT = json.loads(python_contract)
|
||||
|
||||
|
||||
def contracted(name, fake):
|
||||
signature = inspect.Signature(
|
||||
[inspect.Parameter(parameter, inspect.Parameter.POSITIONAL_OR_KEYWORD) for parameter in CONTRACT[name]]
|
||||
)
|
||||
|
||||
def checked(*args, **kwargs):
|
||||
signature.bind(*args, **kwargs)
|
||||
return fake(*args, **kwargs)
|
||||
|
||||
return checked
|
||||
|
||||
|
||||
if not hasattr(legacy, 'is_internal'):
|
||||
legacy.is_internal = contextvars.ContextVar('is_internal_call', default=False)
|
||||
|
||||
def setup(call_type, args, kwargs, start, asynchronous):
|
||||
logger = kwargs['logger_factory'](kwargs) if 'logger_factory' in kwargs else kwargs['logger']
|
||||
logger.setup_call_type = call_type
|
||||
return types.SimpleNamespace(logger=logger, kwargs=kwargs)
|
||||
|
||||
|
||||
FAKES = {
|
||||
'setup': setup,
|
||||
'finalize': lambda response, logger, kwargs, start, end: logger.record('finalize', response),
|
||||
'update_logging': lambda logger, kwargs, model, optional_params, litellm_params, provider: logger.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model=model,
|
||||
optional_params=optional_params,
|
||||
litellm_params=litellm_params,
|
||||
custom_llm_provider=provider,
|
||||
),
|
||||
'pre_call': lambda logger, input, api_key, additional_args: logger.pre_call(input, api_key, additional_args),
|
||||
'post_call': lambda logger, original_response, api_key, additional_args: logger.post_call(
|
||||
original_response, api_key, additional_args
|
||||
),
|
||||
'defers_async_logging': lambda logger: bool(getattr(logger, '_defer_async_logging', False)),
|
||||
'defer_success': lambda logger, pending: setattr(logger, '_native_pending_logging', pending),
|
||||
'sync_success_for_async_call': lambda logger, response, start, end: logger.handle_sync_success_callbacks_for_async_calls(
|
||||
response, start, end
|
||||
),
|
||||
'failure_handler': lambda logger, error, start, end, asynchronous: (
|
||||
logger.async_failure_handler if asynchronous else logger.failure_handler
|
||||
)(error, ''.join(traceback.format_exception(error)), start, end),
|
||||
'submit_success': lambda logger, response, start, end: logger.record('submit', (response, start, end)),
|
||||
'async_success_handler': lambda logger, response, start, end: logger.async_success_handler(response, start, end),
|
||||
'enqueue_logging': lambda coroutine: coroutine.enqueue(),
|
||||
'restore_context': lambda logger: logger.record('restore', None),
|
||||
'custom_pricing_fields': lambda: ('ocr_cost_per_page',),
|
||||
'is_internal_call': lambda: legacy.is_internal.get(),
|
||||
'before_deployment_call': lambda kwargs, call_type: kwargs['logger'].hook('pre', kwargs, call_type),
|
||||
'after_deployment_success': lambda kwargs, response, call_type: kwargs['logger'].hook(
|
||||
'success', response, call_type
|
||||
),
|
||||
'after_deployment_failure': lambda kwargs, error, call_type: kwargs['logger'].hook('failure', error, call_type),
|
||||
'stream_opened': lambda logger: logger.record('stream_opened', None),
|
||||
'stream_success': lambda logger, url_route, endpoint_type, request_body, chunks, start, end, first_chunk: logger.record(
|
||||
'stream_success', list(chunks)
|
||||
),
|
||||
'stream_failure': lambda logger, endpoint_type, request_body, chunks, error: logger.record('stream_failure', error),
|
||||
}
|
||||
assert FAKES.keys() == CONTRACT.keys(), sorted(FAKES.keys() ^ CONTRACT.keys())
|
||||
for name, fake in FAKES.items():
|
||||
setattr(legacy, name, contracted(name, fake))
|
||||
|
||||
|
||||
unraisable = sys.modules.setdefault(
|
||||
'litellm_test_unraisable', types.ModuleType('litellm_test_unraisable')
|
||||
)
|
||||
if not hasattr(unraisable, 'events'):
|
||||
unraisable.events = []
|
||||
sys.unraisablehook = lambda event: unraisable.events.append((event.object, event.exc_value))
|
||||
|
||||
|
||||
def unraisable_from(owner):
|
||||
return [error for source, error in unraisable.events if source is owner]
|
||||
|
||||
|
||||
class StubCoroutine:
|
||||
def __init__(self, logger):
|
||||
self.logger = logger
|
||||
|
||||
def enqueue(self):
|
||||
self.logger.record('enqueued', None)
|
||||
self.logger.on_enqueue(self)
|
||||
|
||||
def close(self):
|
||||
self.logger.record('closed', None)
|
||||
|
||||
|
||||
class StubLogger:
|
||||
def __init__(self):
|
||||
self.calls = []
|
||||
self.hooks = {}
|
||||
self.on_enqueue = lambda coroutine: None
|
||||
|
||||
def record(self, name, value):
|
||||
self.calls.append((name, value))
|
||||
|
||||
def names(self):
|
||||
return [name for name, _ in self.calls]
|
||||
|
||||
def hook(self, phase, value, call_type):
|
||||
self.record(phase + '_hook', call_type)
|
||||
return self.hooks.get(phase, lambda value: 'awaitable')(value)
|
||||
|
||||
def failure_handler(self, error, trace, start, end):
|
||||
self.record('failure_handler', error)
|
||||
|
||||
def async_failure_handler(self, error, trace, start, end):
|
||||
self.record('async_failure_handler', error)
|
||||
return 'awaitable'
|
||||
|
||||
def success_handler(self, response, start, end):
|
||||
self.record('success_handler', response)
|
||||
|
||||
def async_success_handler(self, response, start, end):
|
||||
self.record('async_success_handler', response)
|
||||
return StubCoroutine(self)
|
||||
|
||||
def handle_sync_success_callbacks_for_async_calls(self, response, start, end):
|
||||
self.record('sync_success_for_async_call', response)
|
||||
|
||||
|
||||
logger = StubLogger()
|
||||
";
|
||||
|
||||
/// A namespace with the stubs, `StubLogger` and a fresh `logger`, after `script` ran in it.
|
||||
pub(crate) fn namespace<'py>(py: Python<'py>, script: &CStr) -> Bound<'py, PyDict> {
|
||||
let locals = PyDict::new(py);
|
||||
locals.set_item("python_contract", PYTHON_CONTRACT).unwrap();
|
||||
py.run(STUBS, Some(&locals), Some(&locals)).unwrap();
|
||||
py.run(script, Some(&locals), Some(&locals)).unwrap();
|
||||
locals
|
||||
}
|
||||
|
||||
pub(crate) fn run(py: Python<'_>, locals: &Bound<'_, PyDict>, code: &CStr) {
|
||||
py.run(code, Some(locals), Some(locals)).unwrap();
|
||||
}
|
||||
|
||||
pub(crate) fn local<'py>(locals: &Bound<'py, PyDict>, name: &str) -> Bound<'py, PyAny> {
|
||||
locals.get_item(name).unwrap().unwrap()
|
||||
}
|
||||
|
||||
/// A legacy call over the namespace's `kwargs` (or none) and `request` (or `None`).
|
||||
pub(crate) fn legacy_call(
|
||||
py: Python<'_>,
|
||||
locals: &Bound<'_, PyDict>,
|
||||
asynchronous: bool,
|
||||
) -> LegacyLogging {
|
||||
let request = locals
|
||||
.get_item("request")
|
||||
.unwrap()
|
||||
.unwrap_or_else(|| py.None().into_bound(py));
|
||||
let kwargs = locals
|
||||
.get_item("kwargs")
|
||||
.unwrap()
|
||||
.map(|kwargs| kwargs.cast_into::<PyDict>().unwrap())
|
||||
.unwrap_or_else(|| PyDict::new(py));
|
||||
let call = PublicCall::capture(&request, &PyTuple::empty(py), &kwargs).unwrap();
|
||||
LegacyLogging::new(py, litellm_types::Operation::Ocr, call, asynchronous)
|
||||
}
|
||||
16
litellm-rust/crates/config/Cargo.toml
Normal file
16
litellm-rust/crates/config/Cargo.toml
Normal file
|
|
@ -0,0 +1,16 @@
|
|||
[package]
|
||||
name = "litellm-config"
|
||||
version = "0.1.0"
|
||||
edition.workspace = true
|
||||
license.workspace = true
|
||||
repository.workspace = true
|
||||
|
||||
[dependencies]
|
||||
litellm-auth-types.workspace = true
|
||||
serde.workspace = true
|
||||
serde_yaml_ng = "0.10.0"
|
||||
thiserror.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
rstest.workspace = true
|
||||
tempfile.workspace = true
|
||||
7
litellm-rust/crates/config/src/error.rs
Normal file
7
litellm-rust/crates/config/src/error.rs
Normal file
|
|
@ -0,0 +1,7 @@
|
|||
#[derive(Debug, thiserror::Error)]
|
||||
pub enum Error {
|
||||
#[error("could not read config")]
|
||||
Read(#[from] std::io::Error),
|
||||
#[error("invalid YAML config")]
|
||||
Parse(#[from] serde_yaml_ng::Error),
|
||||
}
|
||||
61
litellm-rust/crates/config/src/includes.rs
Normal file
61
litellm-rust/crates/config/src/includes.rs
Normal file
|
|
@ -0,0 +1,61 @@
|
|||
use std::{
|
||||
collections::{BTreeSet, VecDeque},
|
||||
path::{Path, PathBuf},
|
||||
};
|
||||
|
||||
use serde::Deserialize;
|
||||
use serde_yaml_ng::{Mapping, Value};
|
||||
|
||||
use crate::Error;
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct Includes {
|
||||
#[serde(default)]
|
||||
include: Vec<String>,
|
||||
}
|
||||
|
||||
fn read(path: &Path) -> Result<Mapping, Error> {
|
||||
Ok(serde_yaml_ng::from_str(&std::fs::read_to_string(path)?)?)
|
||||
}
|
||||
|
||||
fn entries(config: &Mapping, path: &Path) -> Result<Vec<(String, PathBuf)>, Error> {
|
||||
let includes: Includes = serde_yaml_ng::from_value(Value::Mapping(config.clone()))?;
|
||||
Ok(includes
|
||||
.include
|
||||
.into_iter()
|
||||
.map(|entry| (entry, path.to_owned()))
|
||||
.collect())
|
||||
}
|
||||
|
||||
pub(super) fn load(path: &Path) -> Result<Value, Error> {
|
||||
let root = path.canonicalize()?;
|
||||
let mut merged = read(&root)?;
|
||||
let mut pending: VecDeque<_> = entries(&merged, &root)?.into();
|
||||
let mut loaded = BTreeSet::from([root.clone()]);
|
||||
merged.remove(Value::String("include".into()));
|
||||
while let Some((entry, declaring)) = pending.pop_front() {
|
||||
let declared = declaring.parent().unwrap_or(Path::new(".")).join(&entry);
|
||||
let fallback = root.parent().unwrap_or(Path::new(".")).join(&entry);
|
||||
let location = if declared.exists() {
|
||||
declared
|
||||
} else {
|
||||
fallback
|
||||
}
|
||||
.canonicalize()?;
|
||||
if !loaded.insert(location.clone()) {
|
||||
continue;
|
||||
}
|
||||
let mut included = read(&location)?;
|
||||
pending.extend(entries(&included, &location)?);
|
||||
included.remove(Value::String("include".into()));
|
||||
for (key, value) in included {
|
||||
match (merged.get_mut(&key), value) {
|
||||
(Some(Value::Sequence(base)), Value::Sequence(extra)) => base.extend(extra),
|
||||
(_, value) => {
|
||||
merged.insert(key, value);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
Ok(Value::Mapping(merged))
|
||||
}
|
||||
90
litellm-rust/crates/config/src/lib.rs
Normal file
90
litellm-rust/crates/config/src/lib.rs
Normal file
|
|
@ -0,0 +1,90 @@
|
|||
mod error;
|
||||
mod includes;
|
||||
mod mcp;
|
||||
mod model;
|
||||
mod settings;
|
||||
mod value;
|
||||
|
||||
use std::{fmt, path::Path};
|
||||
|
||||
use serde::Deserialize;
|
||||
|
||||
pub use error::Error;
|
||||
pub use mcp::{McpAuth, McpServer, McpTransport};
|
||||
pub use model::{LiteLlmParams, Model};
|
||||
pub use settings::{GeneralSettings, LiteLlmSettings, RouterSettings};
|
||||
pub use value::{AdditionalFields, Flag, NumberOrString, Object, OneOrMany, Value};
|
||||
|
||||
#[derive(Clone, Default, Deserialize)]
|
||||
#[serde(default)]
|
||||
pub struct Config {
|
||||
pub model_list: Box<[Model]>,
|
||||
pub general_settings: GeneralSettings,
|
||||
pub router_settings: RouterSettings,
|
||||
pub litellm_settings: LiteLlmSettings,
|
||||
pub environment_variables: Object,
|
||||
pub callback_settings: Object,
|
||||
pub assistant_settings: Object,
|
||||
pub default_vertex_config: Object,
|
||||
pub mcp_servers: std::collections::BTreeMap<String, McpServer>,
|
||||
pub credential_list: Box<[Object]>,
|
||||
pub guardrails: Box<[Object]>,
|
||||
pub prompts: Box<[Object]>,
|
||||
pub sandbox_tools: Box<[Object]>,
|
||||
pub search_tools: Box<[Object]>,
|
||||
pub files_settings: Box<[Object]>,
|
||||
pub finetune_settings: Box<[Object]>,
|
||||
pub mcp_tools: Box<[Object]>,
|
||||
pub vector_store_registry: Box<[Object]>,
|
||||
pub worker_registry: Box<[Object]>,
|
||||
pub agents: Box<[Object]>,
|
||||
pub agent_list: Box<[Object]>,
|
||||
pub policies: Object,
|
||||
pub policy_attachments: Box<[Object]>,
|
||||
pub include: Box<[String]>,
|
||||
#[serde(flatten)]
|
||||
pub additional_fields: AdditionalFields,
|
||||
}
|
||||
|
||||
impl fmt::Debug for Config {
|
||||
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
formatter
|
||||
.debug_struct("Config")
|
||||
.field("model_list", &self.model_list)
|
||||
.field("general_settings", &self.general_settings)
|
||||
.field("router_settings", &self.router_settings)
|
||||
.field("litellm_settings", &self.litellm_settings)
|
||||
.field("environment_variables", &self.environment_variables)
|
||||
.field("callback_settings", &self.callback_settings)
|
||||
.field("assistant_settings", &self.assistant_settings)
|
||||
.field("default_vertex_config", &self.default_vertex_config)
|
||||
.field("mcp_servers", &self.mcp_servers)
|
||||
.field("credential_list", &self.credential_list)
|
||||
.field("guardrails", &self.guardrails)
|
||||
.field("prompts", &self.prompts)
|
||||
.field("sandbox_tools", &self.sandbox_tools)
|
||||
.field("search_tools", &self.search_tools)
|
||||
.field("files_settings", &self.files_settings)
|
||||
.field("finetune_settings", &self.finetune_settings)
|
||||
.field("mcp_tools", &self.mcp_tools)
|
||||
.field("vector_store_registry", &self.vector_store_registry)
|
||||
.field("worker_registry", &self.worker_registry)
|
||||
.field("agents", &self.agents)
|
||||
.field("agent_list", &self.agent_list)
|
||||
.field("policies", &self.policies)
|
||||
.field("policy_attachments", &self.policy_attachments)
|
||||
.field("include", &self.include)
|
||||
.field("additional_fields", &self.additional_fields.keys())
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
impl Config {
|
||||
pub fn from_yaml(yaml: &str) -> Result<Self, Error> {
|
||||
Ok(serde_yaml_ng::from_str(yaml)?)
|
||||
}
|
||||
|
||||
pub fn load(path: impl AsRef<Path>) -> Result<Self, Error> {
|
||||
Ok(serde_yaml_ng::from_value(includes::load(path.as_ref())?)?)
|
||||
}
|
||||
}
|
||||
97
litellm-rust/crates/config/src/mcp.rs
Normal file
97
litellm-rust/crates/config/src/mcp.rs
Normal file
|
|
@ -0,0 +1,97 @@
|
|||
use std::{collections::BTreeMap, fmt};
|
||||
|
||||
use litellm_auth_types::SecretValue;
|
||||
use serde::Deserialize;
|
||||
|
||||
use crate::Object;
|
||||
|
||||
#[derive(Clone, Default, Deserialize)]
|
||||
#[serde(default)]
|
||||
pub struct McpServer {
|
||||
pub server_id: Option<String>,
|
||||
pub alias: Option<String>,
|
||||
pub description: Option<String>,
|
||||
pub mcp_info: Object,
|
||||
pub transport: McpTransport,
|
||||
pub url: Option<SecretValue>,
|
||||
pub command: Option<String>,
|
||||
pub args: Box<[String]>,
|
||||
pub env: BTreeMap<String, SecretValue>,
|
||||
pub auth_type: Option<McpAuth>,
|
||||
#[serde(alias = "auth_value")]
|
||||
pub authentication_token: Option<SecretValue>,
|
||||
pub static_headers: BTreeMap<String, SecretValue>,
|
||||
pub upstream_token_header: Option<String>,
|
||||
pub allowed_tools: Option<Box<[String]>>,
|
||||
pub timeout: Option<f64>,
|
||||
pub max_concurrent_requests: Option<usize>,
|
||||
#[serde(flatten)]
|
||||
pub unsupported: Object,
|
||||
}
|
||||
|
||||
impl fmt::Debug for McpServer {
|
||||
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
formatter
|
||||
.debug_struct("McpServer")
|
||||
.field("transport", &self.transport)
|
||||
.field("auth_type", &self.auth_type)
|
||||
.field("timeout", &self.timeout)
|
||||
.field("max_concurrent_requests", &self.max_concurrent_requests)
|
||||
.finish_non_exhaustive()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, Default, Deserialize, PartialEq, Eq)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum McpTransport {
|
||||
#[default]
|
||||
Http,
|
||||
Sse,
|
||||
Stdio,
|
||||
}
|
||||
|
||||
impl McpTransport {
|
||||
pub fn as_str(self) -> &'static str {
|
||||
match self {
|
||||
Self::Http => "http",
|
||||
Self::Sse => "sse",
|
||||
Self::Stdio => "stdio",
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, Deserialize, PartialEq, Eq)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum McpAuth {
|
||||
None,
|
||||
ApiKey,
|
||||
BearerToken,
|
||||
Basic,
|
||||
Authorization,
|
||||
Token,
|
||||
Oauth2,
|
||||
AwsSigv4,
|
||||
Oauth2TokenExchange,
|
||||
Oauth2IdJag,
|
||||
TruePassthrough,
|
||||
OauthDelegate,
|
||||
}
|
||||
|
||||
impl McpAuth {
|
||||
pub fn as_str(self) -> &'static str {
|
||||
match self {
|
||||
Self::None => "none",
|
||||
Self::ApiKey => "api_key",
|
||||
Self::BearerToken => "bearer_token",
|
||||
Self::Basic => "basic",
|
||||
Self::Authorization => "authorization",
|
||||
Self::Token => "token",
|
||||
Self::Oauth2 => "oauth2",
|
||||
Self::AwsSigv4 => "aws_sigv4",
|
||||
Self::Oauth2TokenExchange => "oauth2_token_exchange",
|
||||
Self::Oauth2IdJag => "oauth2_id_jag",
|
||||
Self::TruePassthrough => "true_passthrough",
|
||||
Self::OauthDelegate => "oauth_delegate",
|
||||
}
|
||||
}
|
||||
}
|
||||
95
litellm-rust/crates/config/src/model.rs
Normal file
95
litellm-rust/crates/config/src/model.rs
Normal file
|
|
@ -0,0 +1,95 @@
|
|||
use std::fmt;
|
||||
|
||||
use litellm_auth_types::SecretValue;
|
||||
use serde::Deserialize;
|
||||
|
||||
use crate::{AdditionalFields, Flag, NumberOrString, Object};
|
||||
|
||||
#[derive(Clone, Deserialize)]
|
||||
pub struct Model {
|
||||
pub model_name: String,
|
||||
pub litellm_params: LiteLlmParams,
|
||||
#[serde(default)]
|
||||
pub model_info: Object,
|
||||
pub blocked: Option<bool>,
|
||||
#[serde(flatten)]
|
||||
pub additional_fields: AdditionalFields,
|
||||
}
|
||||
|
||||
impl fmt::Debug for Model {
|
||||
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
formatter
|
||||
.debug_struct("Model")
|
||||
.field("model_name", &self.model_name)
|
||||
.field("litellm_params", &self.litellm_params)
|
||||
.field("model_info", &self.model_info)
|
||||
.field("blocked", &self.blocked)
|
||||
.field("additional_fields", &self.additional_fields.keys())
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Deserialize)]
|
||||
pub struct LiteLlmParams {
|
||||
pub model: String,
|
||||
pub api_key: Option<SecretValue>,
|
||||
pub api_base: Option<String>,
|
||||
pub api_version: Option<String>,
|
||||
pub custom_llm_provider: Option<String>,
|
||||
pub timeout: Option<NumberOrString>,
|
||||
pub stream_timeout: Option<NumberOrString>,
|
||||
pub max_retries: Option<NumberOrString>,
|
||||
pub tpm: Option<NumberOrString>,
|
||||
pub rpm: Option<NumberOrString>,
|
||||
pub itpm: Option<NumberOrString>,
|
||||
pub otpm: Option<NumberOrString>,
|
||||
pub max_parallel_requests: Option<u64>,
|
||||
pub organization: Option<serde_yaml_ng::Value>,
|
||||
pub drop_params: Option<Flag>,
|
||||
pub tags: Option<Box<[String]>>,
|
||||
pub tag_regex: Option<Box<[String]>>,
|
||||
pub max_budget: Option<f64>,
|
||||
pub budget_duration: Option<String>,
|
||||
pub default_api_key_tpm_limit: Option<u64>,
|
||||
pub default_api_key_rpm_limit: Option<u64>,
|
||||
pub use_in_pass_through: Option<bool>,
|
||||
pub use_chat_completions_api: Option<bool>,
|
||||
pub litellm_credential_name: Option<String>,
|
||||
pub provider_affinity_header: Option<String>,
|
||||
#[serde(flatten)]
|
||||
pub additional_fields: AdditionalFields,
|
||||
}
|
||||
|
||||
impl fmt::Debug for LiteLlmParams {
|
||||
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
formatter
|
||||
.debug_struct("LiteLlmParams")
|
||||
.field("model", &self.model)
|
||||
.field("api_key", &self.api_key)
|
||||
.field("api_base", &self.api_base)
|
||||
.field("api_version", &self.api_version)
|
||||
.field("custom_llm_provider", &self.custom_llm_provider)
|
||||
.field("timeout", &self.timeout)
|
||||
.field("stream_timeout", &self.stream_timeout)
|
||||
.field("max_retries", &self.max_retries)
|
||||
.field("tpm", &self.tpm)
|
||||
.field("rpm", &self.rpm)
|
||||
.field("itpm", &self.itpm)
|
||||
.field("otpm", &self.otpm)
|
||||
.field("max_parallel_requests", &self.max_parallel_requests)
|
||||
.field("organization", &self.organization)
|
||||
.field("drop_params", &self.drop_params)
|
||||
.field("tags", &self.tags)
|
||||
.field("tag_regex", &self.tag_regex)
|
||||
.field("max_budget", &self.max_budget)
|
||||
.field("budget_duration", &self.budget_duration)
|
||||
.field("default_api_key_tpm_limit", &self.default_api_key_tpm_limit)
|
||||
.field("default_api_key_rpm_limit", &self.default_api_key_rpm_limit)
|
||||
.field("use_in_pass_through", &self.use_in_pass_through)
|
||||
.field("use_chat_completions_api", &self.use_chat_completions_api)
|
||||
.field("litellm_credential_name", &self.litellm_credential_name)
|
||||
.field("provider_affinity_header", &self.provider_affinity_header)
|
||||
.field("additional_fields", &self.additional_fields.keys())
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
221
litellm-rust/crates/config/src/settings.rs
Normal file
221
litellm-rust/crates/config/src/settings.rs
Normal file
|
|
@ -0,0 +1,221 @@
|
|||
use std::fmt;
|
||||
|
||||
use litellm_auth_types::SecretValue;
|
||||
use serde::Deserialize;
|
||||
|
||||
use crate::{AdditionalFields, Flag, NumberOrString, Object, OneOrMany, Value};
|
||||
|
||||
#[derive(Clone, Deserialize)]
|
||||
#[serde(default)]
|
||||
pub struct GeneralSettings {
|
||||
pub completion_model: Option<String>,
|
||||
pub max_in_flight_requests_per_worker: Option<u64>,
|
||||
pub max_queued_requests_per_worker: Option<u64>,
|
||||
pub admission_queue_timeout_seconds: f64,
|
||||
pub master_key: Option<SecretValue>,
|
||||
pub database_url: Option<SecretValue>,
|
||||
pub database_connection_pool_limit: Option<u64>,
|
||||
pub database_connection_timeout: Option<f64>,
|
||||
pub database_connect_timeout: Option<f64>,
|
||||
pub database_socket_timeout: Option<f64>,
|
||||
pub database_max_idle_connection_lifetime: Option<f64>,
|
||||
pub max_parallel_requests: Option<u64>,
|
||||
pub global_max_parallel_requests: Option<u64>,
|
||||
pub max_request_size_mb: Option<u64>,
|
||||
pub max_response_size_mb: Option<u64>,
|
||||
pub proxy_config_reload_interval_seconds: u64,
|
||||
pub background_health_checks: Option<bool>,
|
||||
pub health_check_interval: u64,
|
||||
pub health_check_concurrency: Option<u64>,
|
||||
pub store_model_in_db: Option<bool>,
|
||||
pub forward_client_headers_to_llm_api: Option<bool>,
|
||||
pub cancel_on_disconnect: Option<bool>,
|
||||
pub infer_model_from_keys: Option<bool>,
|
||||
pub enable_public_model_hub: bool,
|
||||
pub dangerously_permit_weak_or_unset_master_key: Option<bool>,
|
||||
pub plugins: Option<Box<[Object]>>,
|
||||
pub coordination_redis: Option<Object>,
|
||||
pub mcp_allowed_hosts: Option<Box<[String]>>,
|
||||
pub mcp_allowed_origins: Box<[String]>,
|
||||
#[serde(flatten)]
|
||||
pub additional_fields: AdditionalFields,
|
||||
}
|
||||
|
||||
impl Default for GeneralSettings {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
completion_model: None,
|
||||
max_in_flight_requests_per_worker: None,
|
||||
max_queued_requests_per_worker: None,
|
||||
admission_queue_timeout_seconds: 1.0,
|
||||
master_key: None,
|
||||
database_url: None,
|
||||
database_connection_pool_limit: Some(10),
|
||||
database_connection_timeout: Some(60.0),
|
||||
database_connect_timeout: None,
|
||||
database_socket_timeout: None,
|
||||
database_max_idle_connection_lifetime: Some(60.0),
|
||||
max_parallel_requests: None,
|
||||
global_max_parallel_requests: None,
|
||||
max_request_size_mb: None,
|
||||
max_response_size_mb: None,
|
||||
proxy_config_reload_interval_seconds: 30,
|
||||
background_health_checks: None,
|
||||
health_check_interval: 300,
|
||||
health_check_concurrency: None,
|
||||
store_model_in_db: None,
|
||||
forward_client_headers_to_llm_api: None,
|
||||
cancel_on_disconnect: None,
|
||||
infer_model_from_keys: None,
|
||||
enable_public_model_hub: false,
|
||||
dangerously_permit_weak_or_unset_master_key: None,
|
||||
plugins: None,
|
||||
coordination_redis: None,
|
||||
mcp_allowed_hosts: None,
|
||||
mcp_allowed_origins: Box::default(),
|
||||
additional_fields: AdditionalFields::new(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl fmt::Debug for GeneralSettings {
|
||||
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
formatter
|
||||
.debug_struct("GeneralSettings")
|
||||
.field("completion_model", &self.completion_model)
|
||||
.field(
|
||||
"max_in_flight_requests_per_worker",
|
||||
&self.max_in_flight_requests_per_worker,
|
||||
)
|
||||
.field(
|
||||
"max_queued_requests_per_worker",
|
||||
&self.max_queued_requests_per_worker,
|
||||
)
|
||||
.field(
|
||||
"admission_queue_timeout_seconds",
|
||||
&self.admission_queue_timeout_seconds,
|
||||
)
|
||||
.field("master_key", &self.master_key)
|
||||
.field("database_url", &self.database_url)
|
||||
.field("store_model_in_db", &self.store_model_in_db)
|
||||
.field("additional_fields", &self.additional_fields.keys())
|
||||
.finish_non_exhaustive()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Default, Deserialize)]
|
||||
#[serde(default)]
|
||||
pub struct RouterSettings {
|
||||
pub routing_strategy: Option<String>,
|
||||
pub routing_strategy_args: Option<Object>,
|
||||
pub routing_groups: Option<Box<[Object]>>,
|
||||
pub retry_policy: Option<Object>,
|
||||
pub model_group_retry_policy: Option<Object>,
|
||||
pub model_group_affinity_config: Option<Object>,
|
||||
pub allowed_fails: Option<u64>,
|
||||
pub cooldown_time: Option<f64>,
|
||||
pub num_retries: Option<u64>,
|
||||
pub timeout: Option<f64>,
|
||||
pub max_retries: Option<u64>,
|
||||
pub retry_after: Option<f64>,
|
||||
pub fallbacks: Option<Box<[Object]>>,
|
||||
pub context_window_fallbacks: Option<Box<[Object]>>,
|
||||
pub model_group_alias: Option<Object>,
|
||||
pub enable_tag_filtering: Option<bool>,
|
||||
pub weights: Option<Object>,
|
||||
pub tag_routing_prefix: Option<String>,
|
||||
pub optional_pre_call_checks: Option<Box<[String]>>,
|
||||
#[serde(flatten)]
|
||||
pub additional_fields: AdditionalFields,
|
||||
}
|
||||
|
||||
impl fmt::Debug for RouterSettings {
|
||||
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
formatter
|
||||
.debug_struct("RouterSettings")
|
||||
.field("routing_strategy", &self.routing_strategy)
|
||||
.field("routing_strategy_args", &self.routing_strategy_args)
|
||||
.field("routing_groups", &self.routing_groups)
|
||||
.field("retry_policy", &self.retry_policy)
|
||||
.field("model_group_retry_policy", &self.model_group_retry_policy)
|
||||
.field(
|
||||
"model_group_affinity_config",
|
||||
&self.model_group_affinity_config,
|
||||
)
|
||||
.field("allowed_fails", &self.allowed_fails)
|
||||
.field("cooldown_time", &self.cooldown_time)
|
||||
.field("num_retries", &self.num_retries)
|
||||
.field("timeout", &self.timeout)
|
||||
.field("max_retries", &self.max_retries)
|
||||
.field("retry_after", &self.retry_after)
|
||||
.field("fallbacks", &self.fallbacks)
|
||||
.field("context_window_fallbacks", &self.context_window_fallbacks)
|
||||
.field("model_group_alias", &self.model_group_alias)
|
||||
.field("enable_tag_filtering", &self.enable_tag_filtering)
|
||||
.field("weights", &self.weights)
|
||||
.field("tag_routing_prefix", &self.tag_routing_prefix)
|
||||
.field("optional_pre_call_checks", &self.optional_pre_call_checks)
|
||||
.field("additional_fields", &self.additional_fields.keys())
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Default, Deserialize)]
|
||||
#[serde(default)]
|
||||
pub struct LiteLlmSettings {
|
||||
pub ssl_verify: Option<Flag>,
|
||||
pub ssl_certificate: Option<String>,
|
||||
pub ssl_security_level: Option<String>,
|
||||
pub ssl_ecdh_curve: Option<String>,
|
||||
pub force_ipv4: Option<bool>,
|
||||
pub http2: Option<bool>,
|
||||
pub aiohttp_trust_env: Option<bool>,
|
||||
pub disable_aiohttp_trust_env: Option<bool>,
|
||||
pub disable_aiohttp_transport: Option<bool>,
|
||||
pub drop_params: Option<Flag>,
|
||||
pub request_timeout: Option<NumberOrString>,
|
||||
pub num_retries: Option<u64>,
|
||||
pub cache: Option<bool>,
|
||||
pub cache_params: Option<Object>,
|
||||
pub callbacks: Option<OneOrMany<Value>>,
|
||||
pub success_callback: Option<OneOrMany<Value>>,
|
||||
pub failure_callback: Option<OneOrMany<Value>>,
|
||||
pub json_logs: Option<bool>,
|
||||
pub set_verbose: Option<bool>,
|
||||
#[serde(flatten)]
|
||||
pub additional_fields: AdditionalFields,
|
||||
}
|
||||
|
||||
impl fmt::Debug for LiteLlmSettings {
|
||||
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
formatter
|
||||
.debug_struct("LiteLlmSettings")
|
||||
.field("drop_params", &self.drop_params)
|
||||
.field("request_timeout", &self.request_timeout)
|
||||
.field("num_retries", &self.num_retries)
|
||||
.field("cache", &self.cache)
|
||||
.field("cache_params", &self.cache_params)
|
||||
.field(
|
||||
"callbacks",
|
||||
&self.callbacks.as_ref().map(|callbacks| callbacks.len()),
|
||||
)
|
||||
.field(
|
||||
"success_callback",
|
||||
&self
|
||||
.success_callback
|
||||
.as_ref()
|
||||
.map(|callbacks| callbacks.len()),
|
||||
)
|
||||
.field(
|
||||
"failure_callback",
|
||||
&self
|
||||
.failure_callback
|
||||
.as_ref()
|
||||
.map(|callbacks| callbacks.len()),
|
||||
)
|
||||
.field("json_logs", &self.json_logs)
|
||||
.field("set_verbose", &self.set_verbose)
|
||||
.field("additional_fields", &self.additional_fields.keys())
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
78
litellm-rust/crates/config/src/value.rs
Normal file
78
litellm-rust/crates/config/src/value.rs
Normal file
|
|
@ -0,0 +1,78 @@
|
|||
use std::{collections::BTreeMap, fmt, ops::Deref};
|
||||
|
||||
use serde::Deserialize;
|
||||
|
||||
pub type Value = serde_yaml_ng::Value;
|
||||
pub type AdditionalFields = BTreeMap<String, Value>;
|
||||
|
||||
#[derive(Clone, Default, Deserialize)]
|
||||
#[serde(transparent)]
|
||||
pub struct Object(BTreeMap<String, Value>);
|
||||
|
||||
impl Object {
|
||||
pub fn get(&self, key: &str) -> Option<&Value> {
|
||||
self.0.get(key)
|
||||
}
|
||||
|
||||
pub fn is_empty(&self) -> bool {
|
||||
self.0.is_empty()
|
||||
}
|
||||
|
||||
pub fn len(&self) -> usize {
|
||||
self.0.len()
|
||||
}
|
||||
}
|
||||
|
||||
impl Deref for Object {
|
||||
type Target = BTreeMap<String, Value>;
|
||||
|
||||
fn deref(&self) -> &Self::Target {
|
||||
&self.0
|
||||
}
|
||||
}
|
||||
|
||||
impl fmt::Debug for Object {
|
||||
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
formatter
|
||||
.debug_struct("Object")
|
||||
.field("keys", &self.0.keys())
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Deserialize, PartialEq)]
|
||||
#[serde(untagged)]
|
||||
pub enum NumberOrString {
|
||||
Number(f64),
|
||||
String(String),
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Deserialize, PartialEq, Eq)]
|
||||
#[serde(untagged)]
|
||||
pub enum Flag {
|
||||
Boolean(bool),
|
||||
String(String),
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Deserialize)]
|
||||
#[serde(untagged)]
|
||||
pub enum OneOrMany<T> {
|
||||
Many(Box<[T]>),
|
||||
One(T),
|
||||
}
|
||||
|
||||
impl<T> OneOrMany<T> {
|
||||
pub fn len(&self) -> usize {
|
||||
match self {
|
||||
Self::Many(values) => values.len(),
|
||||
Self::One(_) => 1,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn is_empty(&self) -> bool {
|
||||
match self {
|
||||
Self::Many(values) => values.is_empty(),
|
||||
Self::One(_) => false,
|
||||
}
|
||||
}
|
||||
}
|
||||
375
litellm-rust/crates/config/tests/config.rs
Normal file
375
litellm-rust/crates/config/tests/config.rs
Normal file
|
|
@ -0,0 +1,375 @@
|
|||
use litellm_config::{Config, Error, Flag, NumberOrString};
|
||||
use rstest::{fixture, rstest};
|
||||
use tempfile::TempDir;
|
||||
|
||||
#[fixture]
|
||||
fn directory() -> TempDir {
|
||||
tempfile::tempdir().unwrap()
|
||||
}
|
||||
|
||||
#[fixture]
|
||||
fn model_list_yaml() -> &'static str {
|
||||
r#"
|
||||
model_list:
|
||||
- model_name: assistant
|
||||
litellm_params:
|
||||
model: anthropic/test-model
|
||||
api_key: os.environ/ANTHROPIC_API_KEY
|
||||
- model_name: local
|
||||
litellm_params:
|
||||
model: test-model
|
||||
api_base: http://localhost:8000/v1
|
||||
custom_llm_provider: openai
|
||||
"#
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn loads_model_list_from_file(directory: TempDir, model_list_yaml: &str) {
|
||||
let path = directory.path().join("config.yaml");
|
||||
std::fs::write(&path, model_list_yaml).unwrap();
|
||||
|
||||
let config = Config::load(path).unwrap();
|
||||
assert_eq!(config.model_list.len(), 2);
|
||||
let anthropic = &config.model_list[0];
|
||||
assert_eq!(anthropic.model_name, "assistant");
|
||||
assert_eq!(anthropic.litellm_params.model, "anthropic/test-model");
|
||||
assert_eq!(
|
||||
anthropic.litellm_params.api_key.as_ref().unwrap().expose(),
|
||||
"os.environ/ANTHROPIC_API_KEY"
|
||||
);
|
||||
assert!(anthropic.litellm_params.api_base.is_none());
|
||||
assert!(anthropic.litellm_params.custom_llm_provider.is_none());
|
||||
let local = &config.model_list[1];
|
||||
assert_eq!(local.model_name, "local");
|
||||
assert_eq!(local.litellm_params.model, "test-model");
|
||||
assert!(local.litellm_params.api_key.is_none());
|
||||
assert_eq!(
|
||||
local.litellm_params.api_base.as_deref(),
|
||||
Some("http://localhost:8000/v1")
|
||||
);
|
||||
assert_eq!(
|
||||
local.litellm_params.custom_llm_provider.as_deref(),
|
||||
Some("openai")
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn config_debug_redacts_api_keys() {
|
||||
let config = Config::from_yaml(
|
||||
"model_list: [{model_name: assistant, litellm_params: {model: anthropic/test-model, api_key: secret-value}}]",
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
config.model_list[0]
|
||||
.litellm_params
|
||||
.api_key
|
||||
.as_ref()
|
||||
.unwrap()
|
||||
.expose(),
|
||||
"secret-value"
|
||||
);
|
||||
assert!(!format!("{config:?}").contains("secret-value"));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::malformed_yaml("model_list: [")]
|
||||
#[case::missing_params("model_list: [{model_name: assistant}]")]
|
||||
#[case::missing_model("model_list: [{model_name: assistant, litellm_params: {api_key: key}}]")]
|
||||
fn rejects_malformed_and_incomplete_config(#[case] yaml: &str) {
|
||||
assert!(matches!(Config::from_yaml(yaml), Err(Error::Parse(_))));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn distinguishes_read_errors_from_parse_errors(directory: TempDir) {
|
||||
assert!(matches!(
|
||||
Config::load(directory.path().join("missing.yaml")),
|
||||
Err(Error::Read(error)) if error.kind() == std::io::ErrorKind::NotFound
|
||||
));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::literal("secret-master-key")]
|
||||
#[case::reference("os.environ/LITELLM_MASTER_KEY")]
|
||||
fn loads_and_redacts_the_master_key(#[case] key: &str) {
|
||||
let config = Config::from_yaml(&format!(
|
||||
"model_list: []\ngeneral_settings:\n master_key: {key}\n"
|
||||
))
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
config
|
||||
.general_settings
|
||||
.master_key
|
||||
.as_ref()
|
||||
.unwrap()
|
||||
.expose(),
|
||||
key
|
||||
);
|
||||
assert!(!format!("{config:?}").contains(key));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn missing_general_settings_has_no_master_key() {
|
||||
let config = Config::from_yaml("model_list: []").unwrap();
|
||||
assert!(config.general_settings.master_key.is_none());
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn empty_config_matches_python_defaults() {
|
||||
let config = Config::from_yaml("{}").unwrap();
|
||||
assert!(config.model_list.is_empty());
|
||||
assert_eq!(config.general_settings.admission_queue_timeout_seconds, 1.0);
|
||||
assert_eq!(
|
||||
config.general_settings.database_connection_pool_limit,
|
||||
Some(10)
|
||||
);
|
||||
assert_eq!(
|
||||
config.general_settings.proxy_config_reload_interval_seconds,
|
||||
30
|
||||
);
|
||||
assert_eq!(config.general_settings.health_check_interval, 300);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn parses_typed_settings_and_preserves_extension_fields() {
|
||||
let config = Config::from_yaml(
|
||||
r#"
|
||||
model_list:
|
||||
- model_name: assistant
|
||||
litellm_params:
|
||||
model: vertex_ai/test-model
|
||||
timeout: os.environ/REQUEST_TIMEOUT
|
||||
tpm: os.environ/TPM_LIMIT
|
||||
rpm: 5
|
||||
drop_params: "true"
|
||||
vertex_project: test-project
|
||||
model_info:
|
||||
mode: chat
|
||||
access_groups: [internal]
|
||||
general_settings:
|
||||
master_key: secret-master-key
|
||||
store_model_in_db: true
|
||||
custom_auth: auth.py
|
||||
router_settings:
|
||||
routing_strategy: simple-shuffle
|
||||
allowed_fails: 2
|
||||
redis_host: cache.internal
|
||||
litellm_settings:
|
||||
drop_params: true
|
||||
cache: true
|
||||
custom_callback_name: audit
|
||||
future_section:
|
||||
enabled: true
|
||||
"#,
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
let model = &config.model_list[0];
|
||||
assert_eq!(
|
||||
model.litellm_params.timeout,
|
||||
Some(NumberOrString::String(
|
||||
"os.environ/REQUEST_TIMEOUT".to_string()
|
||||
))
|
||||
);
|
||||
assert_eq!(
|
||||
model.litellm_params.tpm,
|
||||
Some(NumberOrString::String("os.environ/TPM_LIMIT".to_string()))
|
||||
);
|
||||
assert_eq!(model.litellm_params.rpm, Some(NumberOrString::Number(5.0)));
|
||||
assert_eq!(
|
||||
model.litellm_params.drop_params,
|
||||
Some(Flag::String("true".to_string()))
|
||||
);
|
||||
assert!(
|
||||
model
|
||||
.litellm_params
|
||||
.additional_fields
|
||||
.contains_key("vertex_project")
|
||||
);
|
||||
assert!(model.additional_fields.contains_key("access_groups"));
|
||||
assert_eq!(config.router_settings.allowed_fails, Some(2));
|
||||
assert!(
|
||||
config
|
||||
.router_settings
|
||||
.additional_fields
|
||||
.contains_key("redis_host")
|
||||
);
|
||||
assert!(
|
||||
config
|
||||
.litellm_settings
|
||||
.additional_fields
|
||||
.contains_key("custom_callback_name")
|
||||
);
|
||||
assert!(config.additional_fields.contains_key("future_section"));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn parses_python_config_sections() {
|
||||
let config = Config::from_yaml(
|
||||
r#"
|
||||
environment_variables:
|
||||
REDIS_PORT: 6379
|
||||
callback_settings:
|
||||
otel:
|
||||
message_logging: false
|
||||
assistant_settings:
|
||||
custom_llm_provider: openai
|
||||
credential_list:
|
||||
- credential_name: bedrock
|
||||
credential_values:
|
||||
aws_region_name: us-east-1
|
||||
guardrails:
|
||||
- guardrail_name: pii
|
||||
litellm_params:
|
||||
guardrail: presidio
|
||||
prompts:
|
||||
- prompt_id: support
|
||||
sandbox_tools:
|
||||
- sandbox_tool_name: e2b
|
||||
search_tools:
|
||||
- search_tool_name: web
|
||||
files_settings:
|
||||
- custom_llm_provider: openai
|
||||
finetune_settings:
|
||||
- custom_llm_provider: openai
|
||||
mcp_tools:
|
||||
- name: lookup
|
||||
mcp_servers:
|
||||
docs:
|
||||
url: https://example.test/mcp
|
||||
vector_store_registry:
|
||||
- vector_store_name: docs
|
||||
worker_registry:
|
||||
- worker_id: regional
|
||||
agents:
|
||||
- agent_name: reviewer
|
||||
agent_list:
|
||||
- agent_name: legacy-reviewer
|
||||
policies:
|
||||
safe:
|
||||
guardrails:
|
||||
add: [pii]
|
||||
policy_attachments:
|
||||
- policy_id: safe
|
||||
include:
|
||||
- models.yaml
|
||||
"#,
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(config.environment_variables.len(), 1);
|
||||
assert_eq!(config.credential_list.len(), 1);
|
||||
assert_eq!(config.guardrails.len(), 1);
|
||||
assert_eq!(config.prompts.len(), 1);
|
||||
assert_eq!(config.sandbox_tools.len(), 1);
|
||||
assert_eq!(config.search_tools.len(), 1);
|
||||
assert_eq!(config.files_settings.len(), 1);
|
||||
assert_eq!(config.finetune_settings.len(), 1);
|
||||
assert_eq!(config.mcp_tools.len(), 1);
|
||||
assert_eq!(config.mcp_servers.len(), 1);
|
||||
assert_eq!(config.vector_store_registry.len(), 1);
|
||||
assert_eq!(config.worker_registry.len(), 1);
|
||||
assert_eq!(config.agents.len(), 1);
|
||||
assert_eq!(config.agent_list.len(), 1);
|
||||
assert!(config.policies.contains_key("safe"));
|
||||
assert_eq!(config.policy_attachments.len(), 1);
|
||||
assert_eq!(&*config.include, &["models.yaml"]);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::policy_pipeline("../../../litellm/proxy/example_config_yaml/test_pipeline_config.yaml")]
|
||||
#[case::gateway("../../../tests/e2e/gateway/stage_mirror_ci_config.yml")]
|
||||
fn parses_representative_python_configs(#[case] relative_path: &str) {
|
||||
let path = std::path::Path::new(env!("CARGO_MANIFEST_DIR")).join(relative_path);
|
||||
let config = Config::load(path).unwrap();
|
||||
assert!(!config.model_list.is_empty());
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::one("callbacks: custom_callbacks.logger", 1)]
|
||||
#[case::many("callbacks: [prometheus, otel]", 2)]
|
||||
fn accepts_python_callback_shorthand(#[case] setting: &str, #[case] expected_len: usize) {
|
||||
let config = Config::from_yaml(&format!("litellm_settings:\n {setting}")).unwrap();
|
||||
assert_eq!(
|
||||
config.litellm_settings.callbacks.as_ref().unwrap().len(),
|
||||
expected_len
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::api_key("api_key", "provider-secret")]
|
||||
#[case::provider_extension("aws_secret_access_key", "aws-secret")]
|
||||
#[case::general_extension("custom_auth_secret", "auth-secret")]
|
||||
#[case::router_extension("redis_password", "redis-secret")]
|
||||
#[case::litellm_extension("callback_token", "callback-secret")]
|
||||
#[case::root_extension("private_token", "root-secret")]
|
||||
fn debug_output_does_not_expose_config_values(#[case] field: &str, #[case] secret: &str) {
|
||||
let yaml = match field {
|
||||
"api_key" => format!(
|
||||
"model_list: [{{model_name: assistant, litellm_params: {{model: test, api_key: {secret}}}}}]"
|
||||
),
|
||||
"aws_secret_access_key" => format!(
|
||||
"model_list: [{{model_name: assistant, litellm_params: {{model: test, aws_secret_access_key: {secret}}}}}]"
|
||||
),
|
||||
"custom_auth_secret" => format!("general_settings: {{{field}: {secret}}}"),
|
||||
"redis_password" => format!("router_settings: {{{field}: {secret}}}"),
|
||||
"callback_token" => format!("litellm_settings: {{{field}: {secret}}}"),
|
||||
"private_token" => format!("{field}: {secret}"),
|
||||
_ => unreachable!(),
|
||||
};
|
||||
let config = Config::from_yaml(&yaml).unwrap();
|
||||
assert!(!format!("{config:?}").contains(secret));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn resolves_nested_includes_once_in_breadth_first_order() {
|
||||
let directory = TempDir::new().unwrap();
|
||||
let root = directory.path();
|
||||
std::fs::create_dir(root.join("nested")).unwrap();
|
||||
std::fs::write(root.join("config.yaml"), "include: [nested/first.yaml, second.yaml]\nmodel_list: [{model_name: root, litellm_params: {model: root}}]\n").unwrap();
|
||||
std::fs::write(root.join("nested/first.yaml"), "include: [third.yaml]\nmodel_list: [{model_name: first, litellm_params: {model: first}}]\n").unwrap();
|
||||
std::fs::write(root.join("second.yaml"), "general_settings: {master_key: second}\nmodel_list: [{model_name: second, litellm_params: {model: second}}]\n").unwrap();
|
||||
std::fs::write(root.join("nested/third.yaml"), "include: [../config.yaml]\ngeneral_settings: {master_key: third}\nmodel_list: [{model_name: third, litellm_params: {model: third}}]\n").unwrap();
|
||||
|
||||
let config = Config::load(root.join("config.yaml")).unwrap();
|
||||
assert_eq!(
|
||||
config
|
||||
.model_list
|
||||
.iter()
|
||||
.map(|model| model.model_name.as_str())
|
||||
.collect::<Vec<_>>(),
|
||||
["root", "first", "second", "third"]
|
||||
);
|
||||
assert_eq!(
|
||||
config.general_settings.master_key.unwrap().expose(),
|
||||
"third"
|
||||
);
|
||||
assert!(config.include.is_empty());
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn mcp_config_redacts_nested_credentials_and_preserves_policy_for_validation() {
|
||||
let config = Config::from_yaml("mcp_servers:\n docs:\n url: https://example.test/private-secret/mcp\n authentication_token: upstream-secret\n static_headers: {x-token: header-secret}\n env: {TOKEN: env-secret}\n args: [argument-secret]\n client_secret: oauth-secret\n allowed_tools: [search]\n").unwrap();
|
||||
let server = &config.mcp_servers["docs"];
|
||||
assert_eq!(server.allowed_tools.as_deref().unwrap(), ["search"]);
|
||||
assert!(server.unsupported.contains_key("client_secret"));
|
||||
let debug = format!("{config:?}");
|
||||
for secret in [
|
||||
"private-secret",
|
||||
"upstream-secret",
|
||||
"header-secret",
|
||||
"env-secret",
|
||||
"oauth-secret",
|
||||
"argument-secret",
|
||||
] {
|
||||
assert!(!debug.contains(secret));
|
||||
}
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::transport("transport: invalid")]
|
||||
#[case::auth("auth_type: invalid")]
|
||||
#[case::concurrency("max_concurrent_requests: -1")]
|
||||
#[case::headers("static_headers: {x-token: [not, a, string]}")]
|
||||
fn rejects_invalid_typed_mcp_settings(#[case] setting: &str) {
|
||||
assert!(Config::from_yaml(&format!("mcp_servers:\n docs:\n {setting}\n")).is_err());
|
||||
}
|
||||
|
|
@ -13,6 +13,7 @@ serde.workspace = true
|
|||
serde_json.workspace = true
|
||||
serde_path_to_error = "0.1"
|
||||
serde_with.workspace = true
|
||||
strum.workspace = true
|
||||
thiserror.workspace = true
|
||||
url.workspace = true
|
||||
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue