Merge remote-tracking branch 'origin/main' into litellm_agent365_fail_open_default

This commit is contained in:
yucheng 2026-09-27 00:01:26 +00:00
commit 4b3022bdcb
1015 changed files with 41388 additions and 10749 deletions

View file

@ -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, 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

View file

@ -1,7 +1,7 @@
#!/usr/bin/env bash
set -uo pipefail
category="${1:?usage: classify_changes.sh <backend|client|ui|provider-harness|cost-map-only|mcp-dependencies>}"
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,6 +9,7 @@ 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
@ -22,6 +23,10 @@ while IFS= read -r file || [ -n "$file" ]; do
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
;;

View file

@ -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

View file

@ -5,6 +5,7 @@ flag="${1:?usage: unit_selection.sh <codecov flag>}"
legacy_flags=(
caching-local
core-utils
enterprise-package
enterprise-routing
integrations
@ -32,6 +33,7 @@ legacy_flags=(
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
@ -42,6 +44,8 @@ legacy_paths() {
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/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
@ -77,6 +81,7 @@ legacy_paths() {
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 ;;
@ -142,7 +147,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) echo tests/unit/types ;;
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
}

View file

@ -369,6 +369,13 @@ workflows:
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

View file

@ -7,9 +7,9 @@
"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/unit/test_cost_calculator.py::test_completion_cost_charges_explicit_per_token_rates_over_registered_ones",
"COST-ZERO": "tests/unit/test_cost_calculator.py::test_completion_cost_is_zero_when_explicit_rates_are_zero",
"LOG-CONTENT-ON": "tests/test_litellm/litellm_core_utils/test_litellm_logging.py::test_standard_logging_payload_keeps_message_content_when_message_logging_is_on",
"LOG-CONTENT-OFF": "tests/test_litellm/litellm_core_utils/test_litellm_logging.py::test_standard_logging_payload_redacts_message_content_when_message_logging_is_off",
"CALLBACK-SUCCESS": "tests/test_litellm/litellm_core_utils/test_litellm_logging.py::test_async_success_handler_delivers_standard_logging_payload_to_custom_logger",
"CALLBACK-FAILURE": "tests/test_litellm/litellm_core_utils/test_litellm_logging.py::test_async_failure_handler_delivers_failure_payload_to_custom_logger"
"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"
}
}

View file

@ -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)

View file

@ -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 ()

View file

@ -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

View file

@ -12,9 +12,9 @@ on:
- "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/test_litellm/caching/test_redis_cluster_cache.py"
- "tests/test_litellm/caching/test_evicted_client_closer.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"
@ -85,9 +85,9 @@ jobs:
redis-server --version
uv run --no-sync pytest \
tests/unit/test_redis.py \
tests/test_litellm/caching/test_redis_connection_pool.py \
tests/test_litellm/caching/test_redis_cluster_cache.py \
tests/test_litellm/caching/test_evicted_client_closer.py \
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 \

View file

@ -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

View file

@ -61,7 +61,8 @@ jobs:
- shard: core-utils
artifact-name: core-utils
test-path: "tests/test_litellm/litellm_core_utils"
test-path: ""
unit-flag: core-utils
workers: 2
reruns: 1
timeout-minutes: 20
@ -69,9 +70,7 @@ jobs:
- shard: enterprise-routing
artifact-name: enterprise-routing
test-path: >-
tests/test_litellm/router_utils
tests/test_litellm/router_strategy
test-path: ""
unit-flag: enterprise-routing
workers: 2
reruns: 2
@ -89,7 +88,7 @@ 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
@ -98,7 +97,7 @@ 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
@ -108,10 +107,6 @@ jobs:
- shard: misc
artifact-name: misc
test-path: >-
tests/test_litellm/interactions
tests/test_litellm/ocr
tests/test_litellm/passthrough
tests/test_litellm/rust_bridge
tests/test_litellm/test_*.py
unit-flag: misc
workers: 2
@ -228,9 +223,7 @@ jobs:
- shard: responses-caching-types
artifact-name: responses-caching-types
test-path: >-
tests/test_litellm/responses
tests/test_litellm/caching
test-path: ""
unit-flag: responses-caching-types
workers: 2
reruns: 2

View file

@ -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`

View file

@ -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/` |

View file

@ -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:

View file

@ -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)"
@ -301,7 +301,7 @@ 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
@ -310,7 +310,7 @@ 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
@ -329,10 +329,10 @@ test-unit-integrations: install-test-deps
$(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/unit/secret_managers tests/unit/vector_stores tests/unit/a2a_protocol tests/test_litellm/anthropic_interface tests/unit/completion_extras tests/unit/containers tests/unit/enterprise tests/unit/experimental_mcp_client tests/unit/google_genai tests/unit/images tests/unit/interactions tests/test_litellm/interactions tests/test_litellm/passthrough tests/test_litellm/router_strategy tests/test_litellm/router_utils tests/unit/types --tb=short -vv -n 4 --durations=20
$(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/unit/test_*.py tests/test_litellm/test_*.py --tb=short -vv -n 4 --durations=20

View file

@ -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) | ✅ | ✅ | ✅ | | | | | | | |

View file

@ -0,0 +1,37 @@
# Publish MCP servers in the AI Hub
Set `litellm_settings.public_mcp_servers` to the concrete IDs of the servers you want listed in the public AI Hub. Pin `server_id` in each configuration entry so the publication list stays stable across deployments
```yaml
mcp_servers:
documentation:
server_id: documentation-mcp
url: https://mcp.example.com/mcp
transport: http
available_on_public_internet: true
litellm_settings:
public_mcp_hub_strict_whitelist: true
public_mcp_servers:
- documentation-mcp
```
Use `documentation-mcp`, the `server_id`, in the publication list. The configuration key `documentation`, display names, and aliases are not publication IDs. Database-created servers use the ID returned by `/v1/mcp/server`
The dashboard's **AI Hub > MCP Hub > Manage MCP Hub Visibility** dialog edits this same list. Its YAML example includes the selected server IDs. With database-backed configuration (`store_model_in_db: true`), a value declared in YAML is owned by that file: edit the file and reload, or remove that key from YAML to let the dashboard manage it in the database. File-backed deployments can save the list directly to their configuration file
To remove all explicit entries, save an empty selection in the dialog or configure:
```yaml
litellm_settings:
public_mcp_hub_strict_whitelist: true
public_mcp_servers: []
```
## Hub listing and network access
The **Hub listing** column in AI Hub identifies servers that appear in `/public/mcp_hub`. The dashboard derives this status from the current registry and publication settings. Setting `mcp_info.is_public` on a server does not publish it; that response field is derived metadata. `mcp_info.is_public_explicit` identifies registered servers included in the explicit publication list
Gateway cards and server details show **All Networks** when `available_on_public_internet` is enabled or the server is explicitly published in `public_mcp_servers`. They show **Internal Only** when both are false. The per-server flag defaults to `true`; explicit publication overrides a disabled flag for compatibility. Older proxies that omit the metadata needed to determine access show **Unknown**. These labels describe allowed client IPs; authentication and tool permissions still apply
The default `public_mcp_hub_strict_whitelist: true` lists only registered servers in `public_mcp_servers`. Legacy mode (`false`) additionally lists registered servers with `available_on_public_internet: true`. In legacy mode, clearing the explicit publication list leaves these automatically listed servers visible. Enable strict mode when the publication list should fully determine hub visibility

View file

@ -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,

View file

@ -9,10 +9,14 @@
- 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
- 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
- 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

192
litellm-rust/Cargo.lock generated
View file

@ -199,7 +199,7 @@ checksum = "ae36dc4177970ef04fde5178d3e2429882def40e57a451f919c098f72baa6cec"
dependencies = [
"proc-macro2",
"quote",
"syn 3.0.0",
"syn 3.0.6",
]
[[package]]
@ -710,14 +710,20 @@ dependencies = [
"http 1.4.2",
"http-body 1.1.0",
"http-body-util",
"hyper 1.10.1",
"hyper-util",
"itoa",
"matchit",
"memchr",
"mime",
"multer",
"percent-encoding",
"pin-project-lite",
"serde_core",
"serde_json",
"serde_path_to_error",
"sync_wrapper",
"tokio",
"tower",
"tower-layer",
"tower-service",
@ -1053,18 +1059,18 @@ dependencies = [
[[package]]
name = "clap"
version = "4.6.6"
version = "4.6.7"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "473c7e07f409a8d772161724aa8db6a765a2532a70f9667eeb7b49d3d02fbdca"
checksum = "aa8876b300ab35ba921adea3dfd70157a46249b33f95c9084ae5709785478946"
dependencies = [
"clap_builder",
]
[[package]]
name = "clap_builder"
version = "4.6.6"
version = "4.6.7"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "7b48fea5a88e9ae728a2dcbedbfc0e730f7d60da42e1cb049a83c9fb8b789889"
checksum = "ec0797fb7aeb1406c84efac526901f7ec3ead2124f946b494e72879d4b54704d"
dependencies = [
"anstyle",
"clap_lex",
@ -1180,7 +1186,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "e75b2483e97a5a7da73ac68a05b629f9c53cff58d8ed1c77866079e18b00dba5"
dependencies = [
"digest 0.10.7",
"spin",
"spin 0.10.1",
]
[[package]]
@ -1581,6 +1587,15 @@ dependencies = [
"serde",
]
[[package]]
name = "encoding_rs"
version = "0.8.35"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "75030f3c4f45dafd7586dd6780965a8c7e8e285a5ecb86713e63a79c5b2766f3"
dependencies = [
"cfg-if",
]
[[package]]
name = "equivalent"
version = "1.0.2"
@ -2816,6 +2831,10 @@ version = "0.12.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "32a66949e030da00e8c7d4434b251670a91556f4144941d37452769c25d58a53"
[[package]]
name = "litellm"
version = "0.0.1"
[[package]]
name = "litellm-auth"
version = "0.1.0"
@ -2840,6 +2859,7 @@ dependencies = [
"litellm-http",
"moka",
"reqwest 0.12.28",
"rstest",
"serde_json",
"sha2 0.10.9",
"thiserror 2.0.19",
@ -2911,6 +2931,7 @@ dependencies = [
"litellm-cache",
"litellm-cache-response",
"litellm-cache-testing",
"litellm-http",
"reqwest 0.12.28",
"rstest",
"serde_json",
@ -2944,6 +2965,7 @@ dependencies = [
"litellm-auth-types",
"litellm-cache",
"litellm-cache-testing",
"litellm-http",
"percent-encoding",
"reqwest 0.12.28",
"rstest",
@ -2970,6 +2992,7 @@ dependencies = [
"futures-util",
"litellm-cache",
"litellm-cache-testing",
"litellm-http",
"qdrant-client",
"reqwest 0.12.28",
"rstest",
@ -3043,6 +3066,7 @@ dependencies = [
"litellm-auth-aws",
"litellm-cache",
"litellm-cache-testing",
"litellm-http",
"reqwest 0.12.28",
"rstest",
"serde_json",
@ -3087,6 +3111,18 @@ dependencies = [
"strum",
]
[[package]]
name = "litellm-config"
version = "0.1.0"
dependencies = [
"litellm-auth-types",
"rstest",
"serde",
"serde_yaml_ng",
"tempfile",
"thiserror 2.0.19",
]
[[package]]
name = "litellm-core"
version = "0.1.0"
@ -3102,6 +3138,7 @@ dependencies = [
"litellm-http",
"litellm-llms",
"litellm-secrets",
"litellm-tracing",
"litellm-types",
"mime_guess",
"moka",
@ -3137,6 +3174,7 @@ dependencies = [
"serde_json",
"serde_path_to_error",
"serde_with",
"strum",
"thiserror 2.0.19",
"url",
]
@ -3173,6 +3211,69 @@ dependencies = [
"tokio-util",
]
[[package]]
name = "litellm-gateway"
version = "0.1.0"
dependencies = [
"axum",
"futures-util",
"http-body-util",
"litellm-config",
"litellm-core",
"litellm-gateway-auth",
"litellm-gateway-inference",
"litellm-http",
"litellm-llms",
"litellm-secrets",
"litellm-tracing",
"rstest",
"serde_json",
"tokio",
"tower",
"tracing",
"uuid",
]
[[package]]
name = "litellm-gateway-auth"
version = "0.1.0"
dependencies = [
"axum",
"futures-util",
"litellm-auth-types",
"litellm-config",
"litellm-secrets",
"rstest",
"sha2 0.10.9",
"subtle",
"thiserror 2.0.19",
"tokio",
"tower",
]
[[package]]
name = "litellm-gateway-inference"
version = "0.1.0"
dependencies = [
"axum",
"base64 0.22.1",
"bytes",
"futures-util",
"litellm-auth",
"litellm-core",
"litellm-http",
"litellm-llms",
"litellm-router",
"litellm-secrets",
"litellm-types",
"rstest",
"serde_json",
"thiserror 2.0.19",
"tokio",
"tower",
"wiremock",
]
[[package]]
name = "litellm-host"
version = "0.1.0"
@ -3207,11 +3308,13 @@ dependencies = [
"http 1.4.2",
"hyper-util",
"litellm-core-utils",
"rcgen",
"reqwest 0.12.28",
"rstest",
"rustls 0.23.42",
"serde",
"serde_json",
"tempfile",
"thiserror 2.0.19",
"tokio",
"veil",
@ -3236,6 +3339,7 @@ dependencies = [
"litellm-framing",
"litellm-host",
"litellm-http",
"litellm-python-compat",
"litellm-secrets",
"litellm-types",
"reqwest 0.12.28",
@ -3257,6 +3361,7 @@ version = "0.1.0"
dependencies = [
"indexmap 2.14.0",
"jsonschema",
"litellm-types",
"rstest",
"schemars 1.2.2",
"serde",
@ -3276,7 +3381,6 @@ dependencies = [
"futures-util",
"litellm-auth",
"litellm-auth-aws",
"litellm-auth-gcp",
"litellm-cache",
"litellm-cache-azure-blob",
"litellm-cache-disk",
@ -3311,6 +3415,7 @@ dependencies = [
"serde_json",
"serde_with",
"sha2 0.10.9",
"strum",
"thiserror 2.0.19",
"tokio",
"tokio-tungstenite",
@ -3334,6 +3439,15 @@ dependencies = [
"thiserror 2.0.19",
]
[[package]]
name = "litellm-router"
version = "0.1.0"
dependencies = [
"litellm-config",
"litellm-core",
"rstest",
]
[[package]]
name = "litellm-secrets"
version = "0.1.0"
@ -3344,6 +3458,7 @@ dependencies = [
"google-cloud-auth",
"google-cloud-kms-v1",
"litellm-core-utils",
"litellm-http",
"litellm-python-compat",
"litellm-secrets-aws",
"litellm-secrets-azure",
@ -3391,6 +3506,7 @@ dependencies = [
"litellm-auth-azure",
"litellm-auth-types",
"litellm-core-utils",
"litellm-http",
"litellm-secrets-types",
"percent-encoding",
"reqwest 0.12.28",
@ -3410,6 +3526,7 @@ version = "0.1.0"
dependencies = [
"base64 0.22.1",
"litellm-core-utils",
"litellm-http",
"litellm-secrets-types",
"litellm-tracing",
"moka",
@ -3438,6 +3555,7 @@ dependencies = [
"litellm-auth-gcp",
"litellm-auth-types",
"litellm-core-utils",
"litellm-http",
"litellm-secrets-types",
"moka",
"percent-encoding",
@ -3479,6 +3597,7 @@ dependencies = [
"rstest",
"serde",
"serde_json",
"strum",
"thiserror 2.0.19",
"tokio",
"veil",
@ -3562,6 +3681,7 @@ dependencies = [
name = "litellm-tracing"
version = "0.1.0"
dependencies = [
"base64 0.22.1",
"fancy-regex 0.19.2",
"percent-encoding",
"rstest",
@ -3576,8 +3696,10 @@ name = "litellm-types"
version = "0.1.0"
dependencies = [
"rstest",
"schemars 1.2.2",
"serde",
"serde_json",
"strum",
]
[[package]]
@ -3745,6 +3867,23 @@ dependencies = [
"syn 2.0.119",
]
[[package]]
name = "multer"
version = "3.1.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "83e87776546dc87511aa5ee218730c92b666d7264ab6ed41f9d215af9cd5224b"
dependencies = [
"bytes",
"encoding_rs",
"futures-util",
"http 1.4.2",
"httparse",
"memchr",
"mime",
"spin 0.9.9",
"version_check",
]
[[package]]
name = "nom"
version = "7.1.3"
@ -4652,7 +4791,7 @@ checksum = "92ecd8964f8453721699a1ed72037b0db49ce2f5a5138486ee89bed6f67cdf3a"
dependencies = [
"proc-macro2",
"quote",
"syn 3.0.0",
"syn 3.0.6",
]
[[package]]
@ -5131,7 +5270,7 @@ dependencies = [
"proc-macro2",
"quote",
"serde_derive_internals",
"syn 3.0.0",
"syn 3.0.6",
]
[[package]]
@ -5219,7 +5358,7 @@ checksum = "e7a5d71263a5a7d47b41f6b3f06ba276f10cc18b0931f1799f710578e2309348"
dependencies = [
"proc-macro2",
"quote",
"syn 3.0.0",
"syn 3.0.6",
]
[[package]]
@ -5230,7 +5369,7 @@ checksum = "f852137cce035d6a4df67ccce505ff6b3e9fd3a10e3e52b24dc71e650bb1a9bd"
dependencies = [
"proc-macro2",
"quote",
"syn 3.0.0",
"syn 3.0.6",
]
[[package]]
@ -5310,6 +5449,19 @@ dependencies = [
"syn 2.0.119",
]
[[package]]
name = "serde_yaml_ng"
version = "0.10.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "7b4db627b98b36d4203a7b458cf3573730f2bb591b28871d916dfa9efabfd41f"
dependencies = [
"indexmap 2.14.0",
"itoa",
"ryu",
"serde",
"unsafe-libyaml",
]
[[package]]
name = "sha1"
version = "0.10.7"
@ -5439,6 +5591,12 @@ dependencies = [
"windows-sys 0.61.2",
]
[[package]]
name = "spin"
version = "0.9.9"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "3763264f6b73151db08c50ff20d7d8a0b8796e021cdea7ceedad07b80155fa0e"
[[package]]
name = "spin"
version = "0.10.1"
@ -5538,9 +5696,9 @@ dependencies = [
[[package]]
name = "syn"
version = "3.0.0"
version = "3.0.6"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "f2fac314a64dc9a36e61a9eb4261a5e9bbfbc922b27e518af97bc32b926cf967"
checksum = "8593e8e72159ed2257d083c7a454a85cbf854f37a0966d8d483aff8c8a3ebcee"
dependencies = [
"proc-macro2",
"quote",
@ -5652,7 +5810,7 @@ checksum = "43cbfe0cf76104d42a574802844187e84a305e531ed54455f11fbde0f10541cd"
dependencies = [
"proc-macro2",
"quote",
"syn 3.0.0",
"syn 3.0.6",
]
[[package]]
@ -6231,6 +6389,12 @@ version = "0.1.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "39ec24b3121d976906ece63c9daad25b85969647682eee313cb5779fdd69e14e"
[[package]]
name = "unsafe-libyaml"
version = "0.2.11"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "673aac59facbab8a9007c7f6108d11f63b603f7cabff99fabf650fea5c32b861"
[[package]]
name = "untrusted"
version = "0.9.0"

View file

@ -9,9 +9,13 @@ 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 = { path = "crates/gateway" }
litellm-gateway-inference = { path = "crates/gateway-inference" }
litellm-gateway-auth = { path = "crates/gateway-auth" }
litellm-coroutine = { path = "crates/coroutine" }
litellm-host = { path = "crates/host" }
litellm-callbacks-legacy-python = { path = "crates/callbacks-legacy-python" }
@ -48,7 +52,10 @@ 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"] }
bytes = "1"
http = "1"
google-cloud-auth = { version = "1.16.0", default-features = false }
@ -57,8 +64,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"] }

View file

@ -7,4 +7,16 @@ 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" },
]
# 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" },
]

View file

@ -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

View file

@ -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"] {

View file

@ -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::*;

View file

@ -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>);
@ -305,7 +305,7 @@ fn validate_request_credentials(configured: &str) -> Result<&str, Error> {
}
#[derive(Clone, Debug)]
enum CredentialSource {
pub enum CredentialSource {
Inline(SecretValue),
Trusted(SecretValue),
ApplicationCredentials(String),

View file

@ -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,
@ -40,21 +40,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};

View file

@ -51,7 +51,7 @@ pub use credential::{
CredentialPlanResolution, CredentialRef, CredentialResolver, CredentialResolverHandle,
};
pub use error::Error;
pub use http::{CredentialPlacement, RequestAuth};
pub use http::CredentialPlacement;
pub use policy::{CredentialPlanKind, CredentialRule, ExistingHeaderBehavior, ProviderAuthPolicy};
pub use secret::SecretValue;
pub use token::{ResolvedCredential, TokenFuture, TokenProvider, TokenProviderHandle};

View file

@ -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")]

View 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,
}

View file

@ -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

View file

@ -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> {

View file

@ -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 {

View file

@ -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()
},

View file

@ -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

View file

@ -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};

View file

@ -96,7 +96,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");

View file

@ -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,
)

View file

@ -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

View file

@ -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 {

View file

@ -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!(

View file

@ -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"

View file

@ -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"))
})

View file

@ -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())

View file

@ -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(

View file

@ -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 {

View file

@ -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)
- This crate is 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)
- 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 the driver's preflight, supplied by `python-bridge`; this crate only adopts the keyword view it produces
- 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`)
- 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; routes hand it over through `run_legacy_call` and keep no copy
- `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

View file

@ -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"

View file

@ -19,7 +19,7 @@ use serde_json::Value;
use crate::{
DeploymentHooks, LegacyCallbacks, PublicCall, PythonLogger,
deferred::{PendingLogging, PendingSuccess},
finalize, is_internal_call, prepare,
finalize, is_internal_call,
python::Streaming,
setup,
};
@ -117,9 +117,13 @@ impl LegacyLogging {
})
}
/// The keyword view the rest of the call reads: a copy, so the deployment hook's own
/// dict is left as the hook returned it, carrying the logger as `@client` injects it.
/// The driver's preflight rewrites this same dict before the host projects from it.
fn prepare(&mut self, py: Python<'_>) -> PyResult<LifecycleStep> {
let prepared = prepare(py, self.call.kwargs().bind(py), self.logger()?)?.unbind();
self.call.set_kwargs(prepared);
let prepared = self.call.kwargs().bind(py).copy()?;
prepared.set_item("litellm_logging_obj", self.logger()?.object(py))?;
self.call.set_kwargs(prepared.unbind());
Ok(LifecycleStep::Arguments(self.call.kwargs().clone_ref(py)))
}
@ -465,6 +469,7 @@ impl PythonLifecycle for LegacyLogging {
error.write_unraisable(py, None);
}
self.body = None;
self.headers = None;
self.context = None;
self.stream = None;
}
@ -482,7 +487,8 @@ impl PythonLifecycle for LegacyLogging {
visit.call(&stream.chunks)?;
visit.call(&stream.first_chunk)?;
}
visit.call(&self.body)
visit.call(&self.body)?;
visit.call(&self.headers)
}
}
@ -580,8 +586,6 @@ assert prepared['document'] is replacement
assert prepared['pages'] is replaced_kwargs['pages']
assert prepared['litellm_logging_obj'] is logger
assert 'litellm_logging_obj' not in replaced_kwargs
[checked] = [value for name, value in logger.calls if name == 'check_limits']
assert checked is prepared
",
);
});
@ -616,8 +620,6 @@ kwargs = {'logger': logger, 'vendor_extension': opaque}
&locals,
c"
assert prepared['vendor_extension'] is opaque
[checked] = [value for name, value in logger.calls if name == 'check_limits']
assert checked['vendor_extension'] is opaque
assert hooked == ([opaque] if asynchronous else []), hooked
",
);
@ -733,45 +735,6 @@ assert all(value is failure for name, value in logger.calls if name.endswith('_h
);
});
}
#[rstest]
#[case::synchronous(false)]
#[case::asynchronous(true)]
fn a_limit_rejected_before_the_call_surfaces_as_the_callers_error(#[case] asynchronous: bool) {
Python::initialize();
Python::attach(|py| {
let locals = namespace(
py,
c"
class BudgetExceeded(Exception):
pass
rejection = BudgetExceeded('over budget')
class LimitedLogger(StubLogger):
def check_limits(self, arguments):
raise rejection
logger = LimitedLogger()
logger.hooks = {'pre': lambda kwargs: kwargs}
kwargs = {'logger': logger}
",
);
let mut logging = legacy_call(py, &locals, asynchronous);
let kwargs = local(&locals, "kwargs")
.cast_into::<PyDict>()
.unwrap()
.unbind();
let result = logging.begin(py, kwargs, 0.0).and_then(|step| match step {
LifecycleStep::Await(_) => {
logging.resume(py, Ok(local(&locals, "kwargs").unbind()))
}
step => Ok(step),
});
let error = result.err().unwrap();
assert!(error.value(py).is(local(&locals, "rejection")));
});
}
}
#[cfg(test)]
@ -782,6 +745,7 @@ mod payload_tests {
use litellm_host::event::{MachineEvent, RawResponse, RequestContext, WireRequest};
use litellm_host_python::{LifecycleEvent, LifecycleStep, PythonLifecycle, to_py};
use proptest::prelude::*;
use pyo3::gc::{PyTraverseError, PyVisit};
use pyo3::prelude::*;
use rstest::rstest;
use serde_json::{Map, Value, json};
@ -871,16 +835,7 @@ check = lambda: None
headers: vec![("x-route".into(), "route".into())],
body,
};
let step = logging.before_send(py, Box::new(wire), &context).unwrap();
let raw = MachineEvent::ResponseReceived {
raw: RawResponse {
body: "raw response".into(),
},
};
assert!(matches!(
logging.emit(py, LifecycleEvent::Machine(&raw)).unwrap(),
LifecycleStep::Done
));
let (_, step) = send_and_receive(py, &mut logging, wire, &context);
run(py, &locals, c"check()");
let LifecycleStep::Wire(wire) = step else {
panic!("before_send did not hand back the wire request");
@ -889,6 +844,134 @@ check = lambda: None
})
}
/// `before_send` over `wire`, then the provider's raw response the way the driver
/// delivers it, so `pre_call` and `post_call` have both seen the retained payload.
fn send_and_receive<'a>(
py: Python<'_>,
logging: &'a mut LegacyLogging,
wire: WireRequest,
context: &RequestContext,
) -> (&'a mut LegacyLogging, LifecycleStep) {
let step = logging.before_send(py, Box::new(wire), context).unwrap();
let raw = MachineEvent::ResponseReceived {
raw: RawResponse {
body: "raw response".into(),
},
};
assert!(matches!(
logging.emit(py, LifecycleEvent::Machine(&raw)).unwrap(),
LifecycleStep::Done
));
(logging, step)
}
fn route_context() -> RequestContext {
RequestContext {
model: "model".into(),
custom_llm_provider: "provider".into(),
optional_params: json!({}),
secret_fields: vec![],
api_key: Some(SecretValue::new("route-key")),
}
}
fn route_wire() -> WireRequest {
WireRequest {
url: "https://provider.invalid/ocr".into(),
headers: vec![("x-route".into(), "route".into())],
body: json!({}),
}
}
/// A Python object owning one `LegacyLogging`, so the interpreter's collector sees the
/// edges the adapter reports and clears them the way the driver's `Execution` does.
#[pyclass(weakref)]
struct Retained {
logging: Option<LegacyLogging>,
}
#[pymethods]
impl Retained {
fn __traverse__(&self, visit: PyVisit<'_>) -> Result<(), PyTraverseError> {
match &self.logging {
Some(logging) => logging.traverse(&visit),
None => Ok(()),
}
}
fn __clear__(slf: &Bound<'_, Self>) {
drop(slf.borrow_mut().logging.take());
}
}
#[test]
fn a_cycle_through_the_retained_headers_is_collected() {
Python::initialize();
Python::attach(|py| {
let locals = namespace(py, PAYLOAD_LOGGER);
let mut logging = LegacyLogging {
logger: Some(PythonLogger::new(local(&locals, "logger").unbind())),
..legacy_call(py, &locals, false)
};
send_and_receive(py, &mut logging, route_wire(), &route_context());
let retained = Py::new(
py,
Retained {
logging: Some(logging),
},
)
.unwrap();
locals.set_item("retained", retained).unwrap();
run(
py,
&locals,
c"
import gc
import weakref
logger.post[2]['headers']['owner'] = retained
logger.pre = logger.post = None
reference = weakref.ref(retained)
del retained
gc.collect()
assert reference() is None
",
);
});
}
#[test]
fn close_releases_the_retained_headers() {
Python::initialize();
Python::attach(|py| {
let locals = namespace(py, PAYLOAD_LOGGER);
let mut logging = LegacyLogging {
logger: Some(PythonLogger::new(local(&locals, "logger").unbind())),
..legacy_call(py, &locals, false)
};
send_and_receive(py, &mut logging, route_wire(), &route_context());
run(
py,
&locals,
c"
import weakref
class Sentinel:
pass
sentinel = Sentinel()
logger.post[2]['headers']['sentinel'] = sentinel
logger.pre = logger.post = None
reference = weakref.ref(sentinel)
del sentinel
assert reference() is not None
",
);
logging.close(py);
run(py, &locals, c"assert reference() is None");
});
}
#[rstest]
#[case::caller_keyword(c"
document = {'type': 'document_url', 'document_url': 'data:application/pdf;base64,YWJj'}

View file

@ -4,7 +4,7 @@
//! this crate holds them.
use litellm_host::{machine::Machine, protocol::Protocol};
use litellm_host_python::{ProtocolHost, lookup, run_call};
use litellm_host_python::{Preflight, ProtocolHost, lookup, run_call};
use pyo3::{
gc::{PyTraverseError, PyVisit},
prelude::*,
@ -39,7 +39,8 @@ impl PublicCall {
}
/// 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
}
@ -64,13 +65,15 @@ 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.
/// the keyword view the contract prepares and `preflight` rewrites, and the contract
/// observes the call.
pub fn run_legacy_call<H, M>(
py: Python<'_>,
surface: LegacySurface,
call: PublicCall,
machine: M,
host: H,
preflight: Preflight,
asynchronous: bool,
) -> PyResult<Py<PyAny>>
where
@ -83,6 +86,7 @@ where
machine,
host,
Box::new(LegacyLogging::new(py, surface, call, asynchronous)),
preflight,
arguments,
asynchronous,
)

View file

@ -1,9 +1,10 @@
//! 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
//! sync and async callback registries it fans out to, the deployment hooks and the deferred
//! proxy release. All of it sits behind one
//! [`PythonLifecycle`](litellm_host_python::PythonLifecycle), so the driver, the routes and
//! core never learn which Python object is on the other end.
//! core never learn which Python object is on the other end. The SDK's own request policy
//! (credential inheritance, the budget and retry limits) is the driver's preflight, not this
//! crate's.
//!
//! 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
@ -14,220 +15,12 @@ mod call;
mod callbacks;
mod deferred;
mod logger;
mod preparation;
mod python;
pub(crate) use adapter::LegacyLogging;
pub use adapter::{LegacySurface, PassThroughStream};
pub use call::{PublicCall, run_legacy_call};
pub(crate) use callbacks::{LegacyCallbacks, is_internal_call};
pub(crate) use logger::{DeploymentHooks, PythonLogger, finalize, setup};
pub(crate) use preparation::prepare;
#[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;

View file

@ -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")]

View file

@ -0,0 +1,199 @@
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/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)
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,
),
'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, 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 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,
)
}

View 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

View 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),
}

View file

@ -0,0 +1,48 @@
mod error;
use std::path::Path;
use litellm_auth_types::SecretValue;
use serde::Deserialize;
pub use error::Error;
#[derive(Clone, Debug, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct Config {
pub model_list: Box<[Model]>,
#[serde(default)]
pub general_settings: GeneralSettings,
}
#[derive(Clone, Debug, Default, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct GeneralSettings {
pub master_key: Option<SecretValue>,
}
#[derive(Clone, Debug, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct Model {
pub model_name: String,
pub litellm_params: LiteLlmParams,
}
#[derive(Clone, Debug, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct LiteLlmParams {
pub model: String,
pub api_key: Option<SecretValue>,
pub api_base: Option<String>,
pub custom_llm_provider: Option<String>,
}
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> {
Self::from_yaml(&std::fs::read_to_string(path)?)
}
}

View file

@ -0,0 +1,119 @@
use litellm_config::{Config, Error};
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_model_list("{}")]
#[case::missing_params("model_list: [{model_name: assistant}]")]
#[case::missing_model("model_list: [{model_name: assistant, litellm_params: {api_key: key}}]")]
#[case::unsupported_settings("model_list: []\ngeneral_settings: {unknown: true}")]
#[case::misspelled_param(
"model_list: [{model_name: assistant, litellm_params: {model: test, api_bsae: url}}]"
)]
fn rejects_malformed_incomplete_and_unsupported_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());
}

View file

@ -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

View file

@ -11,11 +11,13 @@
//! accepts; anything richer is declined upstream by the capability gate.
use litellm_types::llms::openai::{ChatMessage, ChatMessageContent};
use strum::IntoStaticStr;
pub const EMPTY_TEXT_PLACEHOLDER: &str =
"[System: Empty message content sanitised to satisfy protocol]";
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
#[derive(Clone, Copy, Debug, IntoStaticStr, PartialEq, Eq)]
#[strum(serialize_all = "snake_case")]
pub enum TurnRole {
User,
Assistant,
@ -23,10 +25,7 @@ pub enum TurnRole {
impl TurnRole {
pub fn as_str(self) -> &'static str {
match self {
Self::User => "user",
Self::Assistant => "assistant",
}
self.into()
}
}

View file

@ -1,4 +1,6 @@
litellm-core is the LiteLLM SDK in Rust — it makes the LLM call. Each top-level call is a module under `src/<route>/` exposing a public entrypoint named after the route (`messages::messages()`, the Rust equivalent of `litellm.messages()`): you call it and get a typed non-streaming response back.
litellm-core is the LiteLLM SDK in Rust. Each top-level call is a module under `src/<route>/` exposing a public entrypoint named after the route. `messages::messages()` returns `MessagesResponse::Message` for a completed response or `MessagesResponse::Stream { headers, chunks }` when the request sets `stream: true`. The chunks are Anthropic SSE bytes in a `Stream<Item = Result<Bytes, Error>>`. Dropping the stream cancels the call. The Python bridge drives `messages::route::messages_machine()` instead, because Python has to answer the call's operations on its own thread; the gateway and the Rust SDK call the plain entrypoint
A route module has the same five pieces, in the order Python runs them. `types.rs` holds the call, the provider request, and the response. `prepare.rs` resolves the provider and credentials and shapes the request (Python's `validate_environment`, `get_complete_url`, `transform_request`). `handler.rs` resolves auth, offers the wire request to `litellm_host::hooks::RouteHooks::before_send`, sends it, reports the raw response through `emit`, and normalizes the response or stream (`pre_call`, `post`, `post_call`, `transform_response`). `mod.rs` exposes the entrypoint that runs prepare then handler with no hooks (`()`). `route.rs`, where a host needs it, wraps the same two calls in a `CallMachine` whose `HostChannel` is the hooks, and pumps a stream through `open` and `deliver`. A handler takes `&impl RouteHooks<Error>` and never a `HostChannel` directly, so it runs without a coroutine. Keep provider transport and transformation details out of the machine driver
## Crate layering
@ -10,6 +12,16 @@ Each crate mirrors one top-level Python package, so a Rust path reads as its Pyt
- `litellm-llms` mirrors `litellm/llms/`: `base_llm/<api>/transformation.rs`, `<provider>/<api>/transformation.rs`, and `base_llm/ocr/handler.rs` (the OCR request handler)
- `litellm-core` mirrors the route packages (`litellm/ocr/`, `litellm/messages/`, ...): entrypoints, route request types, provider dispatch, the route machine, and hooks
A route module owns the call entrypoint, route request types (`*Request<'a>`), credential fallback, provider dispatch, and the handler glue that runs a provider config. Provider code never imports from core; when it needs the caller's hooks mid-call it goes through `litellm_llms::base_llm::ocr::handler::CallHooks`, which each route implements over its host. Import every item from its canonical path. Never re-export another crate's items or give an item a second public path; the only re-export allowed is a private submodule surfacing its item at its module root (`mod error; pub use error::Error;`). Handlers belong in core or llms, never in a host crate
A route module owns the call entrypoint, route request types (`*Request<'a>`), credential fallback, provider dispatch, and the handler glue that runs a provider config. Provider code never imports from core; when it needs the caller's hooks mid-call it goes through `litellm_llms::base_llm::ocr::handler::CallHooks`, the provider-level hooks OCR implements over its host until it folds into `litellm_host::hooks::RouteHooks`. Import every item from its canonical path. Never re-export another crate's items or give an item a second public path; the only re-export allowed is a private submodule surfacing its item at its module root (`mod error; pub use error::Error;`). Handlers belong in core or llms, never in a host crate
## Error placement
The workspace `Error definitions` rules shape each crate's error; this section decides which crate and module a failure belongs to
A failure is declared once, by the lowest crate that raises it. Every crate above nests that error unchanged (`#[error(transparent)] Auth(#[from] litellm_auth::Error)`) or maps it once at its boundary, as `src/error.rs` does for `litellm_llms::Error`. `RouteError` collects route failures and never re-declares a variant a lower crate raises
Scope follows the concept, not the first caller. An error type under `litellm-llms`'s `<provider>/` is private to that provider: no other provider and nothing in `base_llm` may import it. A failure two providers or two routes can hit, such as wire framing, stream event decoding, or a malformed provider response, belongs to the crate that owns the concept: `litellm-framing` for framing, `litellm_llms::Error` for the transformation layer
`litellm_llms::Error` (`crates/llms/src/error.rs`) is the one transformation error for every provider and API. `base_llm/ocr/error.rs` is the recorded exception until OCR folds into it
Not here: serving HTTP (axum routes, extractors), config file reading, rollout state, databases, or callback execution of any kind. Core runs each route as a machine that yields host operations and call events; which integrations consume those events is the host's business.

View file

@ -13,10 +13,11 @@ litellm-host.workspace = true
bytes.workspace = true
futures-util.workspace = true
base64.workspace = true
litellm-auth.workspace = true
litellm-auth = { workspace = true, features = ["aws", "azure", "gcp"] }
litellm-auth-aws.workspace = true
litellm-http.workspace = true
litellm-llms.workspace = true
litellm-tracing.workspace = true
moka.workspace = true
mime_guess = "2.0.5"
rand.workspace = true
@ -36,6 +37,7 @@ url.workspace = true
veil.workspace = true
[dev-dependencies]
litellm-http = { workspace = true, features = ["test-support"] }
litellm-auth-gcp.workspace = true
litellm-llms = { workspace = true, features = ["test-support"] }
rstest.workspace = true

View file

@ -1,13 +0,0 @@
use std::{sync::OnceLock, time::Duration};
use crate::constants::AUDIO_TRANSCRIPTION_TIMEOUT_SECS;
pub(super) fn http_client() -> &'static reqwest::Client {
static CLIENT: OnceLock<reqwest::Client> = OnceLock::new();
CLIENT.get_or_init(|| {
reqwest::Client::builder()
.timeout(Duration::from_secs(AUDIO_TRANSCRIPTION_TIMEOUT_SECS))
.build()
.unwrap_or_else(|_| reqwest::Client::new())
})
}

View file

@ -1,43 +0,0 @@
use litellm_llms::base_llm::chat::transformation::Error as LlmError;
#[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)]
pub enum Error {
#[error("expected {expected}, got {actual}")]
InvalidType {
expected: &'static str,
actual: &'static str,
},
#[error("missing required field: {0}")]
MissingField(&'static str),
#[error("invalid provider: {0}")]
InvalidProvider(String),
#[error("invalid request: {0}")]
InvalidRequest(String),
#[error("invalid response: {0}")]
InvalidResponse(String),
#[error("unsupported by the rust path: {0}")]
Unsupported(&'static str),
#[error(transparent)]
Auth(#[from] litellm_auth::Error),
#[error(transparent)]
Transport(#[from] litellm_http::transport::Error),
#[error(transparent)]
Headers(#[from] litellm_http::request::HeaderError),
#[error(transparent)]
Http(#[from] litellm_http::Error),
#[error(transparent)]
Aws(#[from] litellm_auth_aws::Error),
}
impl From<LlmError> for Error {
fn from(error: LlmError) -> Self {
match error {
LlmError::InvalidType { expected, actual } => Self::InvalidType { expected, actual },
LlmError::MissingField(field) => Self::MissingField(field),
LlmError::InvalidRequest(message) => Self::InvalidRequest(message),
LlmError::InvalidResponse(message) => Self::InvalidResponse(message),
LlmError::Unsupported(reason) => Self::Unsupported(reason),
LlmError::Auth(error) => Self::Auth(error),
}
}
}

View file

@ -1,22 +1,33 @@
use litellm_http::request::truncate_error_body;
use std::time::Duration;
use litellm_http::{Client, request::truncate_error_body};
use litellm_llms::base_llm::auth::resolve_auth;
use serde_json::Value;
use super::{Error, client::http_client};
use crate::audio_transcription::types::ProviderAudioTranscriptionRequest;
use super::Error;
use crate::{
audio_transcription::types::ProviderAudioTranscriptionRequest,
constants::AUDIO_TRANSCRIPTION_TIMEOUT_SECS,
};
pub async fn execute_audio_transcription_provider_call(
http: &Client,
auth: &litellm_auth::AuthServices,
request: ProviderAudioTranscriptionRequest,
) -> Result<Value, Error> {
let response = crate::outbound::outbound_request::<Error>(
&request.auth,
let env_lookup = |key: &str| std::env::var(key).ok();
let authenticated = resolve_auth(auth, request.environment.clone(), &env_lookup).await?;
let response = crate::outbound::outbound_request(
authenticated,
request.url.clone(),
request.upstream_headers.clone(),
&request.body,
request.timeout,
&request.optional_params,
)
.await?
.send(http_client())
Some(
request
.timeout
.unwrap_or(Duration::from_secs(AUDIO_TRANSCRIPTION_TIMEOUT_SECS)),
),
)?
.send(http)
.await
.map_err(|error| {
Error::Transport(litellm_http::transport::Error::Network(error.to_string()))

View file

@ -1,16 +1,20 @@
mod error;
pub mod types;
pub use error::Error;
mod client;
pub use crate::error::RouteError as Error;
mod handler;
mod prepare;
pub use handler::execute_audio_transcription_provider_call;
use litellm_http::{ClientVariant, HttpClientConfig};
pub use prepare::prepare_audio_transcription_provider_call;
use serde_json::Value;
use crate::audio_transcription::types::AudioTranscriptionRequest;
pub async fn audio_transcription(request: AudioTranscriptionRequest<'_>) -> Result<Value, Error> {
execute_audio_transcription_provider_call(prepare_audio_transcription_provider_call(request)?)
.await
pub async fn audio_transcription(
resources: &crate::resources::CoreResources,
config: &HttpClientConfig,
request: AudioTranscriptionRequest<'_>,
) -> Result<Value, Error> {
let request = prepare_audio_transcription_provider_call(request)?;
let http = resources.pool.client(config, ClientVariant::Provider)?;
execute_audio_transcription_provider_call(&http, &resources.auth, request).await
}

View file

@ -1,7 +1,10 @@
use litellm_core_utils::get_llm_provider_logic::{CustomLlmProvider, get_custom_llm_provider};
use litellm_http::request::{has_header, string_headers};
use litellm_http::request::string_headers;
use litellm_llms::{
base_llm::audio_transcription::transformation::{BaseAudioTranscriptionConfig, RequestAuth},
base_llm::{
audio_transcription::transformation::BaseAudioTranscriptionConfig,
auth::{ValidatedEnvironment, with_default_headers},
},
bedrock::audio_transcription::BEDROCK_AUDIO_TRANSCRIPTION_CONFIG,
};
@ -39,20 +42,13 @@ pub fn prepare_audio_transcription_provider_call(
let config = provider_config(provider_info.custom_llm_provider)
.ok_or_else(|| Error::InvalidProvider(provider_info.custom_llm_provider.to_string()))?;
let env_lookup = |key: &str| std::env::var(key).ok();
let mut headers = string_headers("audio transcription", request.extra_headers)?;
let auth = config.auth_strategy(&model, &request.optional_params, &env_lookup)?;
match &auth {
RequestAuth::Bearer { token } if !has_header(&headers, "authorization") => {
headers.push(("Authorization".to_string(), format!("Bearer {token}")));
}
RequestAuth::Header { name, value } if !has_header(&headers, name) => {
headers.push(((*name).to_string(), value.clone()));
}
RequestAuth::Bearer { .. } | RequestAuth::Header { .. } | RequestAuth::AwsSigV4 { .. } => {}
}
if !has_header(&headers, "content-type") {
headers.push(("Content-Type".to_string(), "application/json".to_string()));
}
let forwarded = string_headers("audio transcription", request.extra_headers)?;
let validated =
config.validate_environment(forwarded, &model, &request.optional_params, &env_lookup)?;
let environment = ValidatedEnvironment {
headers: with_default_headers(validated.headers, &[("Content-Type", "application/json")]),
auth: validated.auth,
};
let url = config.get_complete_url(
request.api_base,
&model,
@ -68,9 +64,7 @@ pub fn prepare_audio_transcription_provider_call(
config,
url,
body: transformed.body,
upstream_headers: headers,
auth,
optional_params: request.optional_params,
environment,
timeout: request.timeout,
})
}

View file

@ -1,7 +1,7 @@
use std::time::Duration;
use litellm_llms::base_llm::audio_transcription::transformation::{
BaseAudioTranscriptionConfig, RequestAuth,
use litellm_llms::base_llm::{
audio_transcription::transformation::BaseAudioTranscriptionConfig, auth::ValidatedEnvironment,
};
use serde_json::{Map, Value};
@ -23,9 +23,7 @@ pub struct ProviderAudioTranscriptionRequest {
pub config: &'static dyn BaseAudioTranscriptionConfig,
pub url: String,
pub body: Value,
pub upstream_headers: Vec<(String, String)>,
pub auth: RequestAuth,
pub optional_params: Map<String, Value>,
pub environment: ValidatedEnvironment,
pub timeout: Option<Duration>,
}

View file

@ -1,14 +0,0 @@
use std::{sync::OnceLock, time::Duration};
use crate::constants::{CHAT_COMPLETIONS_CONNECT_TIMEOUT_SECS, CHAT_COMPLETIONS_TIMEOUT_SECS};
pub(super) fn http_client() -> &'static reqwest::Client {
static CLIENT: OnceLock<reqwest::Client> = OnceLock::new();
CLIENT.get_or_init(|| {
reqwest::Client::builder()
.timeout(Duration::from_secs(CHAT_COMPLETIONS_TIMEOUT_SECS))
.connect_timeout(Duration::from_secs(CHAT_COMPLETIONS_CONNECT_TIMEOUT_SECS))
.build()
.unwrap_or_else(|_| reqwest::Client::new())
})
}

View file

@ -1,43 +0,0 @@
use litellm_llms::base_llm::chat::transformation::Error as LlmError;
#[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)]
pub enum Error {
#[error("expected {expected}, got {actual}")]
InvalidType {
expected: &'static str,
actual: &'static str,
},
#[error("missing required field: {0}")]
MissingField(&'static str),
#[error("invalid provider: {0}")]
InvalidProvider(String),
#[error("invalid request: {0}")]
InvalidRequest(String),
#[error("invalid response: {0}")]
InvalidResponse(String),
#[error("unsupported by the rust path: {0}")]
Unsupported(&'static str),
#[error(transparent)]
Auth(#[from] litellm_auth::Error),
#[error(transparent)]
Transport(#[from] litellm_http::transport::Error),
#[error(transparent)]
Headers(#[from] litellm_http::request::HeaderError),
#[error(transparent)]
Http(#[from] litellm_http::Error),
#[error(transparent)]
Aws(#[from] litellm_auth_aws::Error),
}
impl From<LlmError> for Error {
fn from(error: LlmError) -> Self {
match error {
LlmError::InvalidType { expected, actual } => Self::InvalidType { expected, actual },
LlmError::MissingField(field) => Self::MissingField(field),
LlmError::InvalidRequest(message) => Self::InvalidRequest(message),
LlmError::InvalidResponse(message) => Self::InvalidResponse(message),
LlmError::Unsupported(reason) => Self::Unsupported(reason),
LlmError::Auth(error) => Self::Auth(error),
}
}
}

View file

@ -1,20 +1,70 @@
use litellm_http::{outbound::OutboundRequest, request::truncate_error_body};
use litellm_llms::base_llm::chat::transformation::ProviderChatResponseData;
use std::time::Duration;
use litellm_auth::AuthServices;
use litellm_host::{
event::{MachineEvent, RawResponse, RequestContext, WireRequest},
hooks::RouteHooks,
};
use litellm_http::{Client, outbound::OutboundRequest, request::truncate_error_body};
use litellm_llms::base_llm::{
auth::{Authenticated, resolve_auth},
chat::transformation::ProviderChatResponseData,
};
use litellm_types::utils::ChatCompletionsResponse;
use serde_json::Value;
use super::{Error, client::http_client, prepare::prepare_provider_request};
use crate::chat_completions::types::{
ProviderChatCompletionsRequest, ResolvedChatCompletionsRequest,
use super::Error;
use crate::{
chat_completions::types::ProviderChatCompletionsRequest,
constants::CHAT_COMPLETIONS_TIMEOUT_SECS,
};
pub(super) async fn execute_chat_completions_provider_call(
request: ResolvedChatCompletionsRequest<'_>,
pub(super) async fn execute(
http: &Client,
auth: &AuthServices,
request: ProviderChatCompletionsRequest,
hooks: &impl RouteHooks<Error>,
) -> Result<ChatCompletionsResponse, Error> {
let request = prepare_provider_request(request)?;
let outbound = outbound_request(&request).await?;
let ProviderChatCompletionsRequest {
model,
custom_llm_provider,
config,
url,
body,
optional_params,
environment,
timeout,
api_key,
} = request;
let context = RequestContext {
model: model.clone(),
custom_llm_provider,
optional_params: Value::Object(optional_params),
secret_fields: Vec::new(),
api_key,
};
let authenticated = resolve_auth(auth, environment, &|key| std::env::var(key).ok()).await?;
let wire = hooks
.before_send(
WireRequest {
url,
headers: authenticated.headers,
body,
},
context,
)
.await?;
let outbound = outbound_request(
Authenticated {
headers: wire.headers,
signer: authenticated.signer,
},
wire.url,
&wire.body,
timeout,
)?;
let response = outbound.send(http_client()).await.map_err(|err| {
let response = outbound.send(http).await.map_err(|err| {
// Failing to establish the connection means the request never went out,
// so the host can still serve it. Everything else here, a timeout
// above all, may have reached the provider and been answered.
@ -36,13 +86,17 @@ pub(super) async fn execute_chat_completions_provider_call(
body: truncate_error_body(&text),
}));
}
hooks
.emit(MachineEvent::ResponseReceived {
raw: RawResponse { body: text.clone() },
})
.await?;
let body: Value = serde_json::from_str(&text).map_err(|err| {
Error::InvalidResponse(format!("invalid chat completions response JSON: {err}"))
})?;
request
.config
.transform_response(&request.model, ProviderChatResponseData { body })
config
.transform_response(&model, ProviderChatResponseData { body })
.map_err(Error::from)
.map_err(as_response_error)
}
@ -64,31 +118,157 @@ pub(super) fn as_response_error(err: Error) -> Error {
}
}
pub(super) async fn outbound_request(
request: &ProviderChatCompletionsRequest,
pub(super) fn outbound_request(
authenticated: Authenticated,
url: String,
body: &Value,
timeout: Option<Duration>,
) -> Result<OutboundRequest, Error> {
crate::outbound::outbound_request(
&request.auth,
request.url.clone(),
request.upstream_headers.clone(),
&request.body,
request.timeout,
&request.optional_params,
authenticated,
url,
body,
Some(timeout.unwrap_or(Duration::from_secs(CHAT_COMPLETIONS_TIMEOUT_SECS))),
)
.await
.map_err(|error| match error {
// Python drops the caller's copy and prefers a forwarded Authorization
// over the signature, so leave the request to it.
Error::Http(litellm_http::Error::ComputedHeader(_)) => {
litellm_http::Error::ComputedHeader(_) => {
Error::Unsupported("request forwards a header AWS SigV4 computes")
}
other => other,
other => Error::Http(other),
})
}
#[cfg(test)]
mod tests {
use super::{Error, as_response_error};
use std::sync::Mutex;
use rstest::rstest;
use serde_json::json;
use wiremock::{Mock, MockServer, Request, ResponseTemplate, matchers::any};
use super::*;
use crate::chat_completions::{
prepare::{prepare_provider_request, resolve_request},
types::ChatCompletionsRequest,
};
const ANTHROPIC_MESSAGE: &str = r#"{"id":"msg_1","type":"message","role":"assistant","model":"claude-sonnet-4-5","content":[{"type":"text","text":"hello"}],"stop_reason":"end_turn","stop_sequence":null,"usage":{"input_tokens":11,"output_tokens":4}}"#;
/// Rewrites the outgoing request and records what the call reports back.
#[derive(Default)]
struct RecordingHooks {
contexts: Mutex<Vec<RequestContext>>,
raw: Mutex<Vec<String>>,
}
impl RouteHooks<Error> for RecordingHooks {
async fn before_send(
&self,
wire: WireRequest,
context: RequestContext,
) -> Result<WireRequest, Error> {
self.contexts.lock().unwrap().push(context);
let mut body = wire.body;
body["system"] = json!("added by the host");
Ok(WireRequest {
headers: wire
.headers
.into_iter()
.chain([("x-host".to_string(), "seen".to_string())])
.collect(),
body,
..wire
})
}
async fn emit(&self, event: MachineEvent) -> Result<(), Error> {
let MachineEvent::ResponseReceived { raw } = event;
self.raw.lock().unwrap().push(raw.body);
Ok(())
}
}
fn prepared(api_base: &str) -> ProviderChatCompletionsRequest {
prepare_provider_request(
resolve_request(ChatCompletionsRequest {
model: "anthropic/claude-sonnet-4-5",
messages: json!([{"role": "user", "content": "hi"}]),
optional_params: json!({"max_tokens": 16}).as_object().unwrap().clone(),
api_key: Some("sk-test"),
api_base: Some(api_base),
custom_llm_provider: None,
extra_headers: None,
timeout: None,
})
.unwrap(),
)
.unwrap()
}
#[rstest]
#[tokio::test]
async fn the_hooks_rewrite_the_wire_request_and_see_the_raw_response() {
let upstream = MockServer::start().await;
Mock::given(any())
.respond_with(
ResponseTemplate::new(200).set_body_raw(ANTHROPIC_MESSAGE, "application/json"),
)
.mount(&upstream)
.await;
let hooks = RecordingHooks::default();
execute(
&Client::plain_for_test(),
&AuthServices::default(),
prepared(&upstream.uri()),
&hooks,
)
.await
.expect("chat completions call succeeds");
let [request] = <[Request; 1]>::try_from(upstream.received_requests().await.unwrap())
.unwrap_or_else(|requests| panic!("one request, saw {}", requests.len()));
let sent: Value = serde_json::from_slice(&request.body).unwrap();
assert_eq!(sent["system"], "added by the host");
assert_eq!(request.headers["x-host"], "seen");
assert_eq!(request.headers["x-api-key"], "sk-test");
let [context] = <[RequestContext; 1]>::try_from(hooks.contexts.into_inner().unwrap())
.unwrap_or_else(|seen| panic!("before_send runs once, saw {}", seen.len()));
assert_eq!(
(context.model.as_str(), context.custom_llm_provider.as_str()),
("claude-sonnet-4-5", "anthropic")
);
assert_eq!(context.optional_params, json!({"max_tokens": 16}));
assert_eq!(hooks.raw.into_inner().unwrap(), [ANTHROPIC_MESSAGE]);
}
#[rstest]
#[tokio::test]
async fn an_upstream_failure_is_not_reported_as_a_received_response() {
let upstream = MockServer::start().await;
Mock::given(any())
.respond_with(ResponseTemplate::new(500).set_body_string("boom"))
.mount(&upstream)
.await;
let hooks = RecordingHooks::default();
let error = execute(
&Client::plain_for_test(),
&AuthServices::default(),
prepared(&upstream.uri()),
&hooks,
)
.await
.expect_err("the upstream failure fails the call");
assert!(matches!(
error,
Error::Transport(litellm_http::transport::Error::Http { status: 500, .. })
));
assert!(hooks.raw.into_inner().unwrap().is_empty());
}
#[test]
fn response_errors_collapse_to_one_variant_that_can_only_mean_already_sent() {

View file

@ -6,24 +6,26 @@
//! credentials, and it resolves the provider, translates the conversation,
//! calls the provider, and returns a typed OpenAI-shaped response.
mod error;
pub mod types;
pub use error::Error;
mod client;
pub use crate::error::RouteError as Error;
mod common_utils;
pub(crate) mod handler;
mod prepare;
use handler::execute_chat_completions_provider_call;
use litellm_http::{ClientVariant, HttpClientConfig};
use litellm_types::utils::ChatCompletionsResponse;
use prepare::{parse_messages, resolve_provider_config, resolve_request};
use prepare::{parse_messages, prepare_provider_request, resolve_provider_config, resolve_request};
use serde_json::{Map, Value};
use crate::chat_completions::types::ChatCompletionsRequest;
pub async fn chat_completions(
resources: &crate::resources::CoreResources,
config: &HttpClientConfig,
request: ChatCompletionsRequest<'_>,
) -> Result<ChatCompletionsResponse, Error> {
execute_chat_completions_provider_call(resolve_request(request)?).await
let http = resources.pool.client(config, ClientVariant::Provider)?;
let request = prepare_provider_request(resolve_request(request)?)?;
handler::execute(&http, &resources.auth, request, &()).await
}
/// Whether the core would accept this request, without resolving credentials or
@ -38,9 +40,10 @@ pub fn chat_completions_decline_reason(
messages: Value,
optional_params: &Map<String, Value>,
) -> Option<&'static str> {
let Ok((_, config)) = resolve_provider_config(model, custom_llm_provider) else {
let Ok(resolved) = resolve_provider_config(model, custom_llm_provider) else {
return Some("provider is not on the rust chat completions path");
};
let config = resolved.config;
let Ok(messages) = parse_messages(messages) else {
return Some("unreadable message list");
};

View file

@ -1,6 +1,9 @@
use litellm_auth::SecretValue;
use litellm_core_utils::get_llm_provider_logic::{CustomLlmProvider, get_custom_llm_provider};
use litellm_http::request::has_header;
use litellm_llms::base_llm::chat::transformation::{BaseConfig, RequestAuth};
use litellm_llms::base_llm::{
auth::{ValidatedEnvironment, with_default_headers},
chat::transformation::BaseConfig,
};
use litellm_types::llms::openai::ChatMessage;
use serde_json::Value;
@ -12,10 +15,16 @@ use crate::chat_completions::types::{
ChatCompletionsRequest, ProviderChatCompletionsRequest, ResolvedChatCompletionsRequest,
};
pub(super) struct ResolvedProvider {
pub(super) model: String,
pub(super) custom_llm_provider: String,
pub(super) config: &'static dyn BaseConfig,
}
pub(super) fn resolve_provider_config<'a>(
model: &'a str,
custom_llm_provider: Option<&'a str>,
) -> Result<(String, &'static dyn BaseConfig), Error> {
) -> Result<ResolvedProvider, Error> {
let provider_info = get_custom_llm_provider(model, custom_llm_provider)
.or_else(|| {
custom_llm_provider.map(|provider| CustomLlmProvider {
@ -30,7 +39,11 @@ pub(super) fn resolve_provider_config<'a>(
})?;
let config = chat_completions_provider_config(provider_info.custom_llm_provider)
.ok_or_else(|| Error::InvalidProvider(provider_info.custom_llm_provider.to_string()))?;
Ok((provider_info.model.to_string(), config))
Ok(ResolvedProvider {
model: provider_info.model.to_string(),
custom_llm_provider: provider_info.custom_llm_provider.to_string(),
config,
})
}
pub(super) fn parse_messages(messages: Value) -> Result<Vec<ChatMessage>, Error> {
@ -41,7 +54,11 @@ pub(super) fn parse_messages(messages: Value) -> Result<Vec<ChatMessage>, Error>
pub(super) fn resolve_request(
request: ChatCompletionsRequest<'_>,
) -> Result<ResolvedChatCompletionsRequest<'_>, Error> {
let (model, config) = resolve_provider_config(request.model, request.custom_llm_provider)?;
let ResolvedProvider {
model,
custom_llm_provider,
config,
} = resolve_provider_config(request.model, request.custom_llm_provider)?;
let messages = parse_messages(request.messages)?;
if messages.is_empty() {
return Err(Error::InvalidRequest(
@ -53,6 +70,7 @@ pub(super) fn resolve_request(
}
Ok(ResolvedChatCompletionsRequest {
model,
custom_llm_provider,
config,
messages,
optional_params: request.optional_params,
@ -67,59 +85,26 @@ fn validate_environment(
request: &ResolvedChatCompletionsRequest<'_>,
model: &str,
config: &dyn BaseConfig,
) -> Result<(Vec<(String, String)>, RequestAuth), Error> {
) -> Result<ValidatedEnvironment, Error> {
let env_lookup = |key: &str| std::env::var(key).ok();
let mut headers = string_headers(request.extra_headers.clone())?;
let auth = config.auth(
let forwarded = string_headers(request.extra_headers.clone())?;
let validated = config.validate_environment(
forwarded,
request.api_key,
model,
&request.optional_params,
&env_lookup,
)?;
match &auth {
RequestAuth::Header { name, value } => {
// The deployment's credential replaces whatever the caller forwarded
// under the same name, mirroring Python's
// `{**headers, **anthropic_headers}`: letting a request header win
// would let its sender choose the principal the call bills to.
//
// The exception is a scheme the provider hands off to entirely, such
// as an Anthropic OAuth bearer, where Python drops `x-api-key`
// instead of resolving one. Re-adding it there would put the
// credential into a header the host removed on purpose.
if !config.defers_to_forwarded_auth(&headers) {
headers.retain(|(header, _)| !header.eq_ignore_ascii_case(name));
headers.push(((*name).to_string(), value.clone()));
}
}
RequestAuth::Bearer { token } => {
// Bedrock's `get_request_headers` assigns `headers["Authorization"]`
// unconditionally once a bearer token resolves, so the deployment's
// identity outranks whatever the caller forwarded. Keeping the
// caller's would bill and authorize the call as a different
// principal than the same deployment uses on Python.
//
// The `Header` arm below keeps the opposite precedence on purpose:
// Anthropic's transform honours a forwarded OAuth bearer.
headers.retain(|(name, _)| !name.eq_ignore_ascii_case("authorization"));
headers.push(("authorization".to_string(), format!("Bearer {token}")));
}
// SigV4 signs the serialized body, so the handler adds its headers.
RequestAuth::AwsSigV4 { .. } => {}
}
for (name, value) in config.default_headers() {
if !has_header(&headers, name) {
headers.push(((*name).to_string(), (*value).to_string()));
}
}
Ok((headers, auth))
Ok(ValidatedEnvironment {
headers: with_default_headers(validated.headers, config.default_headers()),
auth: validated.auth,
})
}
pub(super) fn prepare_provider_request(
request: ResolvedChatCompletionsRequest<'_>,
) -> Result<ProviderChatCompletionsRequest, Error> {
let (headers, auth) = validate_environment(&request, &request.model, request.config)?;
let environment = validate_environment(&request, &request.model, request.config)?;
let model = request.model;
let config = request.config;
let env_lookup = |key: &str| std::env::var(key).ok();
@ -134,19 +119,21 @@ pub(super) fn prepare_provider_request(
Ok(ProviderChatCompletionsRequest {
model,
custom_llm_provider: request.custom_llm_provider,
config,
url,
body: transformed.body,
upstream_headers: headers,
auth,
optional_params: request.optional_params,
environment,
timeout: request.timeout,
api_key: request.api_key.map(|key| SecretValue::new(key.to_string())),
})
}
#[cfg(test)]
mod tests {
use litellm_llms::base_llm::chat::transformation::RequestAuth;
use litellm_auth::CredentialPlacement;
use litellm_llms::base_llm::auth::{AuthScheme, resolve_auth};
use serde_json::{Map, Value, json};
use super::{prepare_provider_request, resolve_request};
@ -161,6 +148,20 @@ mod tests {
prepare_provider_request(resolve_request(request)?)
}
/// The headers as they go on the wire, credential applied.
fn wire_headers(prepared: &ProviderChatCompletionsRequest) -> Vec<(String, String)> {
tokio::runtime::Builder::new_current_thread()
.build()
.unwrap()
.block_on(resolve_auth(
&litellm_auth::AuthServices::default(),
prepared.environment.clone(),
&|_| None,
))
.unwrap()
.headers
}
fn request<'a>(
model: &'a str,
provider: Option<&'a str>,
@ -227,19 +228,16 @@ mod tests {
))
.expect("prepares");
assert!(
prepared
.upstream_headers
.contains(&("x-api-key".to_string(), "sk-test".to_string()))
wire_headers(&prepared).contains(&("x-api-key".to_string(), "sk-test".to_string()))
);
assert!(
prepared
.upstream_headers
wire_headers(&prepared)
.contains(&("anthropic-version".to_string(), "2023-06-01".to_string()))
);
assert!(matches!(
prepared.auth,
RequestAuth::Header {
name: "x-api-key",
prepared.environment.auth,
AuthScheme::Credential {
placement: CredentialPlacement::Header("x-api-key"),
..
}
));
@ -261,12 +259,12 @@ mod tests {
json!("sk-caller"),
)]));
let prepared = prepare_chat_completions_call(call).expect("prepares");
let keys: Vec<_> = prepared
.upstream_headers
let headers = wire_headers(&prepared);
let keys: Vec<_> = headers
.iter()
.filter(|(name, _)| name.eq_ignore_ascii_case("x-api-key"))
.collect();
assert_eq!(keys.len(), 1, "got {:?}", prepared.upstream_headers);
assert_eq!(keys.len(), 1, "got {:?}", headers);
assert_eq!(keys[0].1, "sk-test");
}
@ -290,16 +288,14 @@ mod tests {
]));
let prepared = prepare_chat_completions_call(call).expect("prepares");
assert!(
!prepared
.upstream_headers
!wire_headers(&prepared)
.iter()
.any(|(name, value)| name.eq_ignore_ascii_case("x-api-key") && value == "sk-test"),
"the resolved key must not be applied over an OAuth bearer, got {:?}",
prepared.upstream_headers
wire_headers(&prepared)
);
assert!(
prepared
.upstream_headers
wire_headers(&prepared)
.iter()
.any(|(name, value)| name.eq_ignore_ascii_case("authorization")
&& value == "Bearer sk-ant-oat01-token")
@ -322,21 +318,20 @@ mod tests {
("X-Api-Key".to_string(), json!("sk-caller")),
]));
let prepared = prepare_chat_completions_call(call).expect("prepares");
let keys: Vec<_> = prepared
.upstream_headers
let headers = wire_headers(&prepared);
let keys: Vec<_> = headers
.iter()
.filter(|(name, _)| name.eq_ignore_ascii_case("x-api-key"))
.collect();
assert_eq!(keys.len(), 1, "got {:?}", prepared.upstream_headers);
assert_eq!(keys.len(), 1, "got {:?}", headers);
assert_eq!(keys[0].1, "sk-test");
assert!(
prepared
.upstream_headers
wire_headers(&prepared)
.iter()
.any(|(name, value)| name.eq_ignore_ascii_case("authorization")
&& value == "Bearer unrelated"),
"the unrelated authorization must survive, got {:?}",
prepared.upstream_headers
wire_headers(&prepared)
);
}
@ -435,18 +430,16 @@ mod tests {
prepared.url,
"https://bedrock-runtime.us-east-1.amazonaws.com/model/anthropic.claude-v2/converse"
);
assert_eq!(
prepared.auth,
RequestAuth::AwsSigV4 {
region: "us-east-1".to_string(),
service: "bedrock",
}
);
assert!(matches!(
&prepared.environment.auth,
AuthScheme::AwsSigV4 { region, service: "bedrock", .. } if region == "us-east-1"
));
// SigV4 signs the serialized body, so prepare must not have added an
// Authorization header; the handler does it.
// Authorization header; the signer does it over the bytes sent.
assert!(
!prepared
.upstream_headers
.environment
.headers
.iter()
.any(|(name, _)| name.eq_ignore_ascii_case("authorization"))
);
@ -475,9 +468,20 @@ mod tests {
json!("abc-123"),
)]));
let prepared = prepare_chat_completions_call(call).expect("prepares");
let signed = crate::chat_completions::handler::outbound_request(&prepared)
.await
.expect("signs");
let authenticated = resolve_auth(
&litellm_auth::AuthServices::default(),
prepared.environment,
&|_| None,
)
.await
.expect("resolves");
let signed = crate::chat_completions::handler::outbound_request(
authenticated,
prepared.url,
&prepared.body,
prepared.timeout,
)
.expect("signs");
let authorization = signed
.header("authorization")
@ -525,9 +529,20 @@ mod tests {
call.api_key = None;
call.extra_headers = Some(Map::from_iter([(forwarded.to_string(), json!("forged"))]));
let prepared = prepare_chat_completions_call(call).expect("prepares");
let error = crate::chat_completions::handler::outbound_request(&prepared)
.await
.expect_err("{forwarded} should decline instead of being signed");
let authenticated = resolve_auth(
&litellm_auth::AuthServices::default(),
prepared.environment,
&|_| None,
)
.await
.expect("resolves");
let error = crate::chat_completions::handler::outbound_request(
authenticated,
prepared.url,
&prepared.body,
prepared.timeout,
)
.expect_err("{forwarded} should decline instead of being signed");
assert!(
matches!(error, Error::Unsupported(_)),
"{forwarded} declined as {error:?}, which the host would not fall back on"
@ -552,8 +567,8 @@ mod tests {
json!("Bearer caller-supplied"),
)]));
let prepared = prepare_chat_completions_call(call).expect("prepares");
let authorizations: Vec<_> = prepared
.upstream_headers
let headers = wire_headers(&prepared);
let authorizations: Vec<_> = headers
.iter()
.filter(|(name, _)| name.eq_ignore_ascii_case("authorization"))
.map(|(_, value)| value.as_str())
@ -585,16 +600,15 @@ mod tests {
json!("Bearer sk-ant-oat01-forwarded"),
)]));
let prepared = prepare_chat_completions_call(call).expect("prepares");
let keys: Vec<_> = prepared
.upstream_headers
let headers = wire_headers(&prepared);
let keys: Vec<_> = headers
.iter()
.filter(|(name, _)| name.eq_ignore_ascii_case("x-api-key"))
.map(|(_, value)| value.as_str())
.collect();
assert!(keys.is_empty(), "got {:?}", prepared.upstream_headers);
assert!(keys.is_empty(), "got {:?}", headers);
assert!(
prepared
.upstream_headers
wire_headers(&prepared)
.iter()
.any(|(name, value)| name.eq_ignore_ascii_case("authorization")
&& value == "Bearer sk-ant-oat01-forwarded")
@ -613,15 +627,13 @@ mod tests {
json!({"maxTokens": 16}),
))
.expect("prepares");
assert_eq!(
prepared.auth,
RequestAuth::Bearer {
token: "sk-test".to_string()
}
);
assert!(matches!(
&prepared.environment.auth,
AuthScheme::Credential { placement: CredentialPlacement::Bearer, secret }
if secret.expose() == "sk-test"
));
assert!(
prepared
.upstream_headers
wire_headers(&prepared)
.iter()
.any(|(name, value)| name.eq_ignore_ascii_case("authorization")
&& value == "Bearer sk-test"),

View file

@ -1,6 +1,7 @@
use std::time::Duration;
use litellm_llms::base_llm::chat::transformation::{BaseConfig, RequestAuth};
use litellm_auth::SecretValue;
use litellm_llms::base_llm::{auth::ValidatedEnvironment, chat::transformation::BaseConfig};
use litellm_types::llms::openai::ChatMessage;
use serde_json::{Map, Value};
@ -23,6 +24,7 @@ pub struct ChatCompletionsRequest<'a> {
pub struct ResolvedChatCompletionsRequest<'a> {
pub model: String,
pub custom_llm_provider: String,
pub config: &'static dyn BaseConfig,
pub messages: Vec<ChatMessage>,
pub optional_params: Map<String, Value>,
@ -34,11 +36,16 @@ pub struct ResolvedChatCompletionsRequest<'a> {
pub struct ProviderChatCompletionsRequest {
pub model: String,
pub custom_llm_provider: String,
pub config: &'static dyn BaseConfig,
pub url: String,
pub body: Value,
pub upstream_headers: Vec<(String, String)>,
pub auth: RequestAuth,
/// The route's parameters before the provider transformation, reported to the host
/// beside the wire request.
pub optional_params: Map<String, Value>,
/// The forwarded and default headers plus how the call authenticates; the credential
/// itself is applied when the request is sent.
pub environment: ValidatedEnvironment,
pub timeout: Option<Duration>,
pub api_key: Option<SecretValue>,
}

View file

@ -5,20 +5,10 @@ pub const OPENAI_DEFAULT_API_BASE: &str = "https://api.openai.com";
/// timeout from the caller still overrides this on the request builder.
pub(crate) const MESSAGES_TIMEOUT_SECS: u64 = 600;
/// Connect timeout for Anthropic Messages provider calls, in seconds.
pub(crate) const MESSAGES_CONNECT_TIMEOUT_SECS: u64 = 10;
/// Provider name used for Anthropic Messages when a deployment's provider model
/// does not carry an explicit provider prefix.
pub const ANTHROPIC_MESSAGES_PROVIDER: &str = "anthropic";
/// Full-request timeout ceiling for chat completions provider calls, in
/// seconds. Mirrors the Python chat completions default.
pub(crate) const CHAT_COMPLETIONS_TIMEOUT_SECS: u64 = 600;
/// Connect timeout for chat completions provider calls, in seconds.
pub(crate) const CHAT_COMPLETIONS_CONNECT_TIMEOUT_SECS: u64 = 10;
pub(crate) const AUDIO_TRANSCRIPTION_TIMEOUT_SECS: u64 = 600;
/// `object` field every non-streaming chat completion response carries.

View file

@ -1,15 +1,166 @@
use litellm_llms::base_llm::ocr::error::Error as OcrError;
//! One error for every route in this crate. OCR still carries its own, richer enum.
//!
//! A variant is declared by the layer that produces it and nested here as is:
//! credentials by `litellm_auth` (AWS folds into it at that crate's boundary), the wire by
//! `litellm_http`, secrets by `litellm_secrets`. The transformation layer's [`LlmError`]
//! maps onto the same-named variants once, here, so no route re-declares them.
#[derive(Debug, thiserror::Error)]
pub enum Error {
use std::sync::Arc;
use litellm_http::transport::Error as TransportError;
use litellm_llms::Error as LlmError;
#[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)]
pub enum RouteError {
#[error("expected {expected}, got {actual}")]
InvalidType {
expected: &'static str,
actual: &'static str,
},
#[error("missing required field: {0}")]
MissingField(&'static str),
#[error("invalid provider: {0}")]
InvalidProvider(String),
#[error("invalid request: {0}")]
InvalidRequest(String),
#[error("invalid response: {0}")]
InvalidResponse(String),
#[error("unsupported by the rust path: {0}")]
Unsupported(&'static str),
#[error(transparent)]
Ocr(#[from] OcrError),
Auth(#[from] litellm_auth::Error),
#[error(transparent)]
Messages(#[from] crate::messages::Error),
Transport(#[from] TransportError),
#[error(transparent)]
ChatCompletions(#[from] crate::chat_completions::Error),
Headers(#[from] litellm_http::request::HeaderError),
#[error(transparent)]
AudioTranscription(#[from] crate::audio_transcription::Error),
Http(#[from] litellm_http::Error),
#[error(transparent)]
Responses(#[from] crate::responses::Error),
Secret(#[from] SecretError),
}
/// Whether the provider had already been called when the route failed. Before the send, a
/// host may retry on another path; after it, the provider has done the work and billed for it.
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum Phase {
BeforeSend,
AfterSend,
}
impl RouteError {
pub fn phase(&self) -> Phase {
match self {
Self::InvalidResponse(_)
| Self::Transport(TransportError::Http { .. } | TransportError::Network(_)) => {
Phase::AfterSend
}
Self::Transport(TransportError::Connect(_))
| Self::InvalidType { .. }
| Self::MissingField(_)
| Self::InvalidProvider(_)
| Self::InvalidRequest(_)
| Self::Unsupported(_)
| Self::Auth(_)
| Self::Headers(_)
| Self::Http(_)
| Self::Secret(_) => Phase::BeforeSend,
}
}
/// The caller's request is what is wrong, as opposed to the environment, the wire, or
/// the provider's answer.
pub fn is_request(&self) -> bool {
match self {
Self::InvalidType { .. }
| Self::MissingField(_)
| Self::InvalidProvider(_)
| Self::InvalidRequest(_)
| Self::Unsupported(_)
| Self::Headers(_) => true,
Self::Auth(error) => !matches!(error, litellm_auth::Error::MissingApiKey { .. }),
Self::InvalidResponse(_) | Self::Transport(_) | Self::Http(_) | Self::Secret(_) => {
false
}
}
}
}
impl From<LlmError> for RouteError {
fn from(error: LlmError) -> Self {
match error {
LlmError::InvalidType { expected, actual } => Self::InvalidType { expected, actual },
LlmError::MissingField(field) => Self::MissingField(field),
LlmError::InvalidRequest(message) => Self::InvalidRequest(message),
LlmError::InvalidResponse(message) => Self::InvalidResponse(message),
LlmError::Unsupported(reason) => Self::Unsupported(reason),
LlmError::Auth(error) => Self::Auth(error),
}
}
}
#[derive(Clone, Debug, thiserror::Error)]
#[error(transparent)]
pub struct SecretError(Arc<litellm_secrets::Error>);
impl SecretError {
pub fn source_error(&self) -> &litellm_secrets::Error {
&self.0
}
}
impl From<litellm_secrets::Error> for RouteError {
fn from(error: litellm_secrets::Error) -> Self {
Self::Secret(SecretError(Arc::new(error)))
}
}
impl PartialEq for SecretError {
fn eq(&self, other: &Self) -> bool {
Arc::ptr_eq(&self.0, &other.0)
}
}
impl Eq for SecretError {}
#[cfg(test)]
mod tests {
use super::{Phase, RouteError};
use litellm_http::transport::Error as TransportError;
#[test]
fn only_a_provider_answer_or_a_lost_connection_counts_as_after_send() {
let after = [
RouteError::InvalidResponse("bad json".into()),
RouteError::Transport(TransportError::Http {
status: 500,
body: "boom".into(),
}),
RouteError::Transport(TransportError::Network("reset".into())),
];
for error in after {
assert_eq!(error.phase(), Phase::AfterSend, "{error:?}");
}
let before = [
RouteError::Transport(TransportError::Connect("refused".into())),
RouteError::Unsupported("streaming"),
RouteError::Auth(litellm_auth::Error::InvalidHeader),
];
for error in before {
assert_eq!(error.phase(), Phase::BeforeSend, "{error:?}");
}
}
#[test]
fn a_missing_api_key_is_the_environment_not_the_request() {
assert!(
!RouteError::Auth(litellm_auth::Error::MissingApiKey {
provider: "Anthropic",
environment_variable: "ANTHROPIC_API_KEY",
})
.is_request()
);
assert!(RouteError::Auth(litellm_auth::Error::InvalidHeader).is_request());
assert!(RouteError::InvalidRequest("top_k".into()).is_request());
assert!(!RouteError::InvalidResponse("bad json".into()).is_request());
}
}

View file

@ -5,6 +5,7 @@ pub mod error;
pub mod messages;
pub mod ocr;
mod outbound;
pub mod resources;
pub mod responses;
pub use error::Error;
pub use error::{Phase, RouteError};

View file

@ -1,14 +0,0 @@
use std::{sync::OnceLock, time::Duration};
use crate::constants::{MESSAGES_CONNECT_TIMEOUT_SECS, MESSAGES_TIMEOUT_SECS};
pub(super) fn http_client() -> &'static reqwest::Client {
static CLIENT: OnceLock<reqwest::Client> = OnceLock::new();
CLIENT.get_or_init(|| {
reqwest::Client::builder()
.timeout(Duration::from_secs(MESSAGES_TIMEOUT_SECS))
.connect_timeout(Duration::from_secs(MESSAGES_CONNECT_TIMEOUT_SECS))
.build()
.unwrap_or_else(|_| reqwest::Client::new())
})
}

View file

@ -1,23 +1,37 @@
use litellm_http::request::string_headers as shared_string_headers;
pub(super) use litellm_http::request::truncate_error_body;
use litellm_llms::{
anthropic::experimental_pass_through::messages::transformation::ANTHROPIC_MESSAGES_CONFIG,
anthropic::messages::transformation::ANTHROPIC_MESSAGES_CONFIG,
azure_ai::anthropic::messages_transformation::AZURE_ANTHROPIC_MESSAGES_CONFIG,
base_llm::anthropic_messages::transformation::BaseAnthropicMessagesConfig,
bedrock::messages::invoke_transformations::anthropic_claude3_transformation::BEDROCK_ANTHROPIC_MESSAGES_CONFIG,
};
use serde_json::{Map, Value};
use strum::{EnumString, IntoStaticStr};
use super::Error;
const HEADER_CONTEXT: &str = "messages";
pub(super) fn messages_provider_config(
provider: &str,
) -> Option<&'static dyn BaseAnthropicMessagesConfig> {
match provider {
"anthropic" => Some(&ANTHROPIC_MESSAGES_CONFIG),
"azure_ai" => Some(&AZURE_ANTHROPIC_MESSAGES_CONFIG),
_ => None,
#[derive(Clone, Copy, Debug, EnumString, IntoStaticStr, PartialEq, Eq)]
#[strum(serialize_all = "snake_case")]
pub(crate) enum MessagesProvider {
Anthropic,
AzureAi,
Bedrock,
}
impl MessagesProvider {
pub(crate) fn as_str(self) -> &'static str {
self.into()
}
pub(crate) fn config(self) -> &'static dyn BaseAnthropicMessagesConfig {
match self {
Self::Anthropic => &ANTHROPIC_MESSAGES_CONFIG,
Self::AzureAi => &AZURE_ANTHROPIC_MESSAGES_CONFIG,
Self::Bedrock => &BEDROCK_ANTHROPIC_MESSAGES_CONFIG,
}
}
}
@ -31,14 +45,26 @@ pub(super) fn string_headers(
mod tests {
use serde_json::json;
use super::{messages_provider_config, string_headers, truncate_error_body};
use rstest::rstest;
use super::{MessagesProvider, string_headers, truncate_error_body};
use crate::messages::Error;
#[rstest]
#[case::anthropic("anthropic", MessagesProvider::Anthropic)]
#[case::azure_ai("azure_ai", MessagesProvider::AzureAi)]
#[case::bedrock("bedrock", MessagesProvider::Bedrock)]
fn provider_round_trips_through_its_python_name(
#[case] name: &str,
#[case] provider: MessagesProvider,
) {
assert_eq!(name.parse::<MessagesProvider>(), Ok(provider));
assert_eq!(provider.as_str(), name);
}
#[test]
fn provider_config_resolves_anthropic_and_azure_ai() {
assert!(messages_provider_config("anthropic").is_some());
assert!(messages_provider_config("azure_ai").is_some());
assert!(messages_provider_config("openai").is_none());
fn provider_without_a_messages_config_is_rejected() {
assert!("openai".parse::<MessagesProvider>().is_err());
}
#[test]

View file

@ -1,80 +0,0 @@
use std::sync::Arc;
use litellm_llms::base_llm::chat::transformation::Error as LlmError;
#[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)]
pub enum Error {
#[error("invalid provider: {0}")]
InvalidProvider(String),
#[error("missing required field: {0}")]
MissingField(&'static str),
#[error("invalid request: {0}")]
InvalidRequest(String),
#[error("invalid response: {0}")]
InvalidResponse(String),
#[error("unsupported by the Rust messages route: {0}")]
Unsupported(&'static str),
#[error(transparent)]
Auth(#[from] litellm_auth::Error),
#[error(transparent)]
Transport(#[from] litellm_http::transport::Error),
#[error(transparent)]
Headers(#[from] litellm_http::request::HeaderError),
#[error(transparent)]
Secret(#[from] SecretError),
}
#[derive(Clone, Debug, thiserror::Error)]
#[error(transparent)]
pub struct SecretError(Arc<litellm_secrets::Error>);
impl SecretError {
pub fn source_error(&self) -> &litellm_secrets::Error {
&self.0
}
}
impl From<litellm_secrets::Error> for Error {
fn from(error: litellm_secrets::Error) -> Self {
Self::Secret(SecretError(Arc::new(error)))
}
}
impl PartialEq for SecretError {
fn eq(&self, other: &Self) -> bool {
Arc::ptr_eq(&self.0, &other.0)
}
}
impl Eq for SecretError {}
impl From<LlmError> for Error {
fn from(error: LlmError) -> Self {
match error {
error @ LlmError::InvalidType { .. } => Self::InvalidRequest(error.to_string()),
LlmError::MissingField(field) => Self::MissingField(field),
LlmError::InvalidRequest(message) => Self::InvalidRequest(message),
LlmError::InvalidResponse(message) => Self::InvalidResponse(message),
LlmError::Unsupported(reason) => Self::Unsupported(reason),
LlmError::Auth(error) => Self::Auth(error),
}
}
}
impl Error {
pub fn is_request(&self) -> bool {
match self {
Self::InvalidProvider(_)
| Self::MissingField(_)
| Self::InvalidRequest(_)
| Self::Unsupported(_)
| Self::Headers(_) => true,
Self::Auth(error) => !matches!(error, litellm_auth::Error::MissingApiKey { .. }),
_ => false,
}
}
pub fn is_response(&self) -> bool {
matches!(self, Self::InvalidResponse(_))
}
}

View file

@ -1,47 +1,143 @@
use std::time::Duration;
use litellm_http::{request::http_request, transport::Error as TransportError};
use litellm_llms::base_llm::anthropic_messages::transformation::BaseAnthropicMessagesConfig;
use bytes::Bytes;
use futures_util::{StreamExt, TryStreamExt, stream::BoxStream};
use litellm_auth::AuthServices;
use litellm_host::{
event::{MachineEvent, RawResponse, RequestContext, WireRequest},
hooks::RouteHooks,
};
use litellm_http::transport::Error as TransportError;
use litellm_llms::base_llm::{
anthropic_messages::{
streaming::{ByteStream, StreamDecoder, encode_anthropic_sse},
transformation::BaseAnthropicMessagesConfig,
},
auth::{Authenticated, resolve_auth},
};
use litellm_tracing::{ByteChunk, debug};
use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse;
use serde_json::Value;
use super::{Error, client::http_client, common_utils::truncate_error_body};
use super::{
Error, MessagesResponse, common_utils::truncate_error_body, prepare::ProviderMessagesRequest,
};
use crate::{constants::MESSAGES_TIMEOUT_SECS, outbound::outbound_request};
pub(super) fn network(error: reqwest::Error) -> Error {
pub(super) async fn execute(
http: &litellm_http::Client,
auth: &AuthServices,
request: ProviderMessagesRequest,
hooks: &impl RouteHooks<Error>,
) -> Result<MessagesResponse, Error> {
let ProviderMessagesRequest {
provider,
url,
body,
environment,
timeout,
api_key,
} = request;
let stream = body.params.stream == Some(true);
let context = RequestContext {
model: body.model.clone(),
custom_llm_provider: provider.as_str().to_string(),
optional_params: serde_json::to_value(&body.params).map_err(serialize_failure)?,
secret_fields: Vec::new(),
api_key,
};
let authenticated = resolve_auth(auth, environment, &|key| std::env::var(key).ok()).await?;
let wire = hooks
.before_send(
WireRequest {
url,
headers: authenticated.headers,
body: serde_json::to_value(&body).map_err(serialize_failure)?,
},
context,
)
.await?;
let provider_name = provider.as_str();
debug!(provider = provider_name, stream, body = %wire.body, "provider request");
let response = send(
http,
Authenticated {
headers: wire.headers,
signer: authenticated.signer,
},
&wire.url,
&wire.body,
timeout,
)
.await?;
debug!(
provider = provider_name,
status = response.status().as_u16(),
"provider response headers"
);
if !response.status().is_success() {
return Err(provider_error(response).await);
}
let config = provider.config();
if stream {
return Ok(streaming_response(
response,
config.stream_decoder(),
provider_name,
));
}
let text = response.text().await.map_err(network)?;
debug!(body = text.as_str(), "provider response body");
hooks
.emit(MachineEvent::ResponseReceived {
raw: RawResponse { body: text.clone() },
})
.await?;
decode_response(config, &body.model, &text)
.map(|message| MessagesResponse::Message(Box::new(message)))
}
fn serialize_failure(err: serde_json::Error) -> Error {
Error::InvalidRequest(format!(
"failed to serialize Anthropic messages request: {err}"
))
}
fn network(error: reqwest::Error) -> Error {
Error::Transport(TransportError::Network(error.to_string()))
}
pub(super) async fn send(
async fn send(
http: &litellm_http::Client,
authenticated: Authenticated,
url: &str,
headers: &[(String, String)],
body: &Value,
timeout: Option<Duration>,
) -> Result<reqwest::Response, Error> {
let encoded = serde_json::to_vec(body)
.map_err(|err| Error::InvalidRequest(format!("failed to encode messages body: {err}")))?;
let builder = headers.iter().fold(
http_client().post(url).body(encoded),
|builder, (key, value)| builder.header(key, value),
);
let builder = match timeout {
Some(duration) => builder.timeout(duration),
None => builder,
};
http_request(builder).await.map_err(network)
let request = outbound_request(
authenticated,
url.to_string(),
body,
Some(timeout.unwrap_or(Duration::from_secs(MESSAGES_TIMEOUT_SECS))),
)?;
request.send(http).await.map_err(network)
}
pub(super) async fn provider_error(response: reqwest::Response) -> Error {
async fn provider_error(response: reqwest::Response) -> Error {
let status = response.status().as_u16();
match response.text().await {
Ok(text) => Error::Transport(TransportError::Http {
status,
body: truncate_error_body(&text),
}),
Ok(text) => {
litellm_tracing::debug!(status, body = text.as_str(), "provider error body");
Error::Transport(TransportError::Http {
status,
body: truncate_error_body(&text),
})
}
Err(error) => network(error),
}
}
pub(super) fn decode_response(
fn decode_response(
config: &dyn BaseAnthropicMessagesConfig,
model: &str,
text: &str,
@ -52,3 +148,97 @@ pub(super) fn decode_response(
.transform_anthropic_messages_response(model, response)
.map_err(Error::from)
}
fn streaming_response(
response: reqwest::Response,
decoder: Option<StreamDecoder>,
provider: &'static str,
) -> MessagesResponse {
let headers = response
.headers()
.iter()
.filter_map(|(name, value)| Some((name.to_string(), value.to_str().ok()?.to_string())))
.collect();
let chunks = match decoder {
None => futures_util::stream::try_unfold(response, move |mut response| async move {
let chunk = response.chunk().await.map_err(network)?;
Ok(chunk.map(|chunk| {
log_chunk(provider, "provider_response", &chunk);
(chunk, response)
}))
})
.boxed(),
Some(decode) => decoded_chunks(response, decode, provider),
};
MessagesResponse::Stream { headers, chunks }
}
fn decoded_chunks(
response: reqwest::Response,
decode: StreamDecoder,
provider: &'static str,
) -> BoxStream<'static, Result<Bytes, Error>> {
let bytes: ByteStream = response
.bytes_stream()
.inspect_ok(move |chunk| log_chunk(provider, "provider_response", chunk))
.map_err(std::io::Error::other)
.boxed();
futures_util::stream::try_unfold(decode(bytes), move |mut events| async move {
let Some(event) = events.try_next().await? else {
return Ok(None);
};
let chunk = encode_anthropic_sse(&event)?;
log_chunk(provider, "client_response", &chunk);
Ok(Some((chunk, events)))
})
.boxed()
}
fn log_chunk(provider: &str, stage: &str, data: &Bytes) {
let chunk = ByteChunk::new(data);
debug!(provider, stage, encoding = chunk.encoding(), chunk = %chunk, "stream chunk");
}
#[cfg(test)]
mod tests {
use litellm_llms::base_llm::anthropic_messages::streaming::anthropic_sse_event_stream;
use rstest::rstest;
use wiremock::{Mock, MockServer, ResponseTemplate, matchers::any};
use super::*;
#[rstest]
#[case::event(
"data: {\"type\":\"ping\"}\n\n",
Some("event: ping\ndata: {\"type\":\"ping\"}\n\n")
)]
#[case::invalid_event("data: invalid\n\ndata: {\"type\":\"ping\"}\n\n", None)]
#[tokio::test]
async fn decoded_streams_encode_events_and_stop_at_the_first_error(
#[case] body: &'static str,
#[case] expected: Option<&str>,
) {
let upstream = MockServer::start().await;
Mock::given(any())
.respond_with(ResponseTemplate::new(200).set_body_raw(body, "text/event-stream"))
.mount(&upstream)
.await;
let response = litellm_http::Client::plain_for_test()
.get(upstream.uri())
.send()
.await
.unwrap();
let MessagesResponse::Stream { mut chunks, .. } =
streaming_response(response, Some(anthropic_sse_event_stream), "test")
else {
panic!("a streaming response returns chunks");
};
let chunk = chunks.next().await.unwrap();
match expected {
Some(expected) => assert_eq!(chunk.unwrap().as_ref(), expected.as_bytes()),
None => assert!(matches!(chunk, Err(Error::InvalidResponse(_))), "{chunk:?}"),
}
assert!(chunks.next().await.is_none());
}
}

View file

@ -1,48 +1,27 @@
//! The Anthropic Messages call, the Rust equivalent of Python's
//! `litellm.messages()`.
//! The Anthropic Messages call, the Rust equivalent of Python's `litellm.messages()`.
//!
//! [`route`] is the call as a machine a host drives, streaming or not. [`messages`] runs
//! it in process for a caller that already holds the request and wants the message.
//! [`messages`] prepares the provider request and sends it in process. [`route`] runs the
//! same two steps as a machine for a host that answers the call's operations itself.
mod error;
pub mod types;
pub use error::Error;
mod client;
mod common_utils;
mod handler;
mod prepare;
pub mod route;
use std::sync::Arc;
mod types;
use litellm_secrets::source::EnvironmentSecrets;
use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse;
use route::{LocalMessagesHost, MessagesCall, MessagesOutput, messages_machine};
use serde_json::Value;
use litellm_http::{ClientVariant, HttpClientConfig};
use litellm_secrets::source::SecretSource;
use crate::messages::types::MessagesRequest;
pub use crate::error::RouteError as Error;
pub use types::{MessagesCall, MessagesResponse, MessagesShaping, messages_body};
pub async fn messages(request: MessagesRequest<'_>) -> Result<AnthropicMessagesResponse, Error> {
let Value::Object(body) = request.body else {
return Err(Error::InvalidRequest(
"messages body must be an object".into(),
));
};
let call = MessagesCall {
model: request.model.into(),
body,
api_key: request.api_key.map(Into::into),
api_base: request.api_base.map(Into::into),
custom_llm_provider: request.custom_llm_provider.map(Into::into),
extra_headers: request.extra_headers,
provider_specific_header: request.provider_specific_header,
timeout: request.timeout,
shaping: request.shaping,
};
let secrets = Arc::new(EnvironmentSecrets::python_compatible());
match litellm_host::run::run(messages_machine(secrets), &LocalMessagesHost::new(call)).await? {
MessagesOutput::Message(message) => Ok(*message),
MessagesOutput::Streamed => Err(Error::Unsupported(
"streamed responses need a streaming host",
)),
}
pub async fn messages(
resources: &crate::resources::CoreResources,
config: &HttpClientConfig,
secrets: &dyn SecretSource,
call: MessagesCall,
) -> Result<MessagesResponse, Error> {
let http = resources.pool.client(config, ClientVariant::Provider)?;
let request = prepare::prepare(call, secrets).await?;
handler::execute(&http, &resources.auth, request, &()).await
}

View file

@ -1,3 +1,6 @@
use std::time::Duration;
use litellm_auth::SecretValue;
use litellm_core_utils::{
dot_notation_indexing::delete_nested_value,
get_llm_provider_logic::{CustomLlmProvider, get_custom_llm_provider},
@ -5,30 +8,51 @@ use litellm_core_utils::{
settings::Lookup,
};
use litellm_llms::{
anthropic::experimental_pass_through::messages::handler::shape_anthropic_messages_request,
base_llm::anthropic_messages::transformation::{
BaseAnthropicMessagesConfig, MessagesTransformContext,
anthropic::messages::handler::shape_anthropic_messages_request,
base_llm::{
anthropic_messages::transformation::MessagesTransformContext,
auth::{ValidatedEnvironment, with_default_headers},
},
};
use litellm_secrets::source::SecretSource;
use litellm_types::llms::anthropic_messages::anthropic_request::AnthropicMessagesRequest;
use serde_json::{Map, Value};
use super::{
Error,
common_utils::{messages_provider_config, string_headers},
Error, MessagesCall,
common_utils::{MessagesProvider, string_headers},
types::invalid_request,
};
use crate::messages::types::{MessagesRequest, ProviderMessagesRequest};
pub(super) struct ResolvedProvider<'a> {
pub(super) model: &'a str,
pub(super) provider: &'a str,
pub(super) config: &'static dyn BaseAnthropicMessagesConfig,
struct ResolvedProvider {
model: String,
provider: MessagesProvider,
}
pub(super) fn resolve_provider<'a>(
model: &'a str,
custom_llm_provider: Option<&'a str>,
) -> Result<ResolvedProvider<'a>, Error> {
pub(super) struct ProviderMessagesRequest {
pub(super) provider: MessagesProvider,
pub(super) url: String,
pub(super) body: AnthropicMessagesRequest,
pub(super) environment: ValidatedEnvironment,
pub(super) timeout: Option<Duration>,
/// The caller's own credential, reported to the host beside the wire request.
pub(super) api_key: Option<SecretValue>,
}
pub(super) async fn prepare(
call: MessagesCall,
secrets: &dyn SecretSource,
) -> Result<ProviderMessagesRequest, Error> {
let resolved = resolve_provider(&call.body.model, call.custom_llm_provider.as_deref())?;
let secrets = secrets
.resolve(resolved.provider.config().secret_names())
.await?;
prepare_provider_request(call, resolved, secrets.as_ref())
}
fn resolve_provider(
model: &str,
custom_llm_provider: Option<&str>,
) -> Result<ResolvedProvider, Error> {
let CustomLlmProvider {
model,
custom_llm_provider: provider,
@ -44,82 +68,79 @@ pub(super) fn resolve_provider<'a>(
"unable to resolve custom_llm_provider for messages request".to_string(),
)
})?;
let config = messages_provider_config(provider)
.ok_or_else(|| Error::InvalidProvider(provider.to_string()))?;
let provider = provider
.parse()
.map_err(|_| Error::InvalidProvider(provider.to_string()))?;
Ok(ResolvedProvider {
model,
model: model.to_string(),
provider,
config,
})
}
pub(super) fn prepare_provider_request(
request: MessagesRequest<'_>,
resolved: ResolvedProvider<'_>,
fn prepare_provider_request(
call: MessagesCall,
resolved: ResolvedProvider,
secrets: &dyn Lookup,
) -> Result<ProviderMessagesRequest, Error> {
let ResolvedProvider {
model,
provider,
config,
} = resolved;
let model = model.to_string();
let ResolvedProvider { model, provider } = resolved;
let MessagesCall {
body,
api_key,
api_base,
extra_headers,
provider_specific_header,
timeout,
shaping,
..
} = call;
let config = provider.config();
let env_lookup = |key: &str| secrets.get(key);
let typed_request: AnthropicMessagesRequest =
serde_json::from_value(request.body).map_err(invalid_request)?;
let sanitized = shape_anthropic_messages_request(
AnthropicMessagesRequest {
model: model.clone(),
..typed_request
},
request.shaping.reasoning_auto_summary,
AnthropicMessagesRequest { model, ..body },
shaping.reasoning_auto_summary,
)?;
let trimmed =
without_additional_drop_params(sanitized, &request.shaping.additional_drop_params)?;
let trimmed = without_additional_drop_params(sanitized, &shaping.additional_drop_params)?;
let transformed = config.transform_anthropic_messages_request(
trimmed,
&MessagesTransformContext::new(request.shaping.capabilities, request.shaping.drop_params),
&MessagesTransformContext::new(shaping.capabilities, shaping.drop_params),
)?;
let scoped = get_provider_specific_headers(request.provider_specific_header.as_ref(), provider);
let scoped =
get_provider_specific_headers(provider_specific_header.as_ref(), provider.as_str());
let forwarded = string_headers(Some(
request
.extra_headers
.into_iter()
.flatten()
.chain(scoped)
.collect(),
extra_headers.into_iter().flatten().chain(scoped).collect(),
))?;
let authenticated = config.authenticate(forwarded, request.api_key, &env_lookup)?;
let headers = config.request_headers(
with_default_headers(authenticated, config.default_headers()),
&transformed,
);
let validated = config.validate_environment(
forwarded,
api_key.as_deref(),
&transformed.model,
&env_lookup,
)?;
let environment = ValidatedEnvironment {
headers: config.request_headers(
with_default_headers(validated.headers, config.default_headers()),
&transformed,
),
auth: validated.auth,
};
let body = serde_json::to_value(transformed).map_err(|err| {
Error::InvalidRequest(format!(
"failed to serialize Anthropic messages request: {err}"
))
})?;
let url = config.get_complete_url(request.api_base, &model, &env_lookup)?;
let url = if transformed.params.stream == Some(true) {
config.complete_stream_url(api_base.as_deref(), &transformed.model, &env_lookup)?
} else {
config.get_complete_url(api_base.as_deref(), &transformed.model, &env_lookup)?
};
Ok(ProviderMessagesRequest {
provider: provider.to_string(),
model,
config,
provider,
url,
body,
upstream_headers: headers,
timeout: request.timeout,
body: transformed,
environment,
timeout,
api_key: api_key.map(SecretValue::new),
})
}
fn invalid_request(err: serde_json::Error) -> Error {
Error::InvalidRequest(format!("invalid Anthropic messages request: {err}"))
}
fn without_additional_drop_params(
request: AnthropicMessagesRequest,
paths: &[String],
@ -127,64 +148,59 @@ fn without_additional_drop_params(
if paths.is_empty() {
return Ok(request);
}
let Value::Object(fields) = serde_json::to_value(request).map_err(invalid_request)? else {
return Err(Error::InvalidRequest(
"Anthropic messages request did not serialize to an object".to_string(),
));
};
let (required, optional): (Map<String, Value>, Map<String, Value>) = fields
.into_iter()
.partition(|(key, _)| matches!(key.as_str(), "model" | "messages"));
let trimmed = paths.iter().fold(Value::Object(optional), |body, path| {
delete_nested_value(body, path)
});
let merged: Map<String, Value> = required
.into_iter()
.chain(trimmed.as_object().cloned().unwrap_or_default())
.collect();
serde_json::from_value(Value::Object(merged)).map_err(invalid_request)
}
fn with_default_headers(
headers: Vec<(String, String)>,
defaults: &[(&str, &str)],
) -> Vec<(String, String)> {
let missing: Vec<(String, String)> = defaults
let params = serde_json::to_value(request.params).map_err(invalid_request)?;
let trimmed = paths
.iter()
.filter(|(name, _)| {
!headers
.iter()
.any(|(header, _)| header.eq_ignore_ascii_case(name))
})
.map(|(name, value)| ((*name).to_string(), (*value).to_string()))
.collect();
headers.into_iter().chain(missing).collect()
.fold(params, |params, path| delete_nested_value(params, path));
Ok(AnthropicMessagesRequest {
params: serde_json::from_value(trimmed).map_err(invalid_request)?,
..request
})
}
#[cfg(test)]
mod tests {
use litellm_llms::base_llm::auth::resolve_auth;
use litellm_types::utils::ProviderSpecificHeaders;
use rstest::{fixture, rstest};
use serde_json::json;
use serde_json::{Map, Value, json};
use super::*;
use crate::messages::types::MessagesShaping;
use crate::messages::MessagesShaping;
#[fixture]
fn shaping() -> MessagesShaping {
MessagesShaping::default()
}
fn prepare(request: MessagesRequest<'_>) -> Result<ProviderMessagesRequest, Error> {
prepare_with_secrets(request, &|_: &str| None)
fn body(value: Value) -> AnthropicMessagesRequest {
serde_json::from_value(value).unwrap()
}
fn prepare(call: MessagesCall) -> Result<ProviderMessagesRequest, Error> {
prepare_with_secrets(call, &|_: &str| None)
}
fn prepare_with_secrets(
request: MessagesRequest<'_>,
call: MessagesCall,
secrets: &dyn Lookup,
) -> Result<ProviderMessagesRequest, Error> {
let resolved = resolve_provider(request.model, request.custom_llm_provider)?;
prepare_provider_request(request, resolved, secrets)
let resolved = resolve_provider(&call.body.model, call.custom_llm_provider.as_deref())?;
prepare_provider_request(call, resolved, secrets)
}
/// The headers as they go on the wire, credential applied.
fn wire_headers(prepared: &ProviderMessagesRequest) -> Vec<(String, String)> {
tokio::runtime::Builder::new_current_thread()
.build()
.unwrap()
.block_on(resolve_auth(
&litellm_auth::AuthServices::default(),
prepared.environment.clone(),
&|_| None,
))
.unwrap()
.headers
}
#[rstest]
@ -221,12 +237,13 @@ mod tests {
.map(|(_, value)| value.to_string())
};
let prepared = prepare_with_secrets(
MessagesRequest {
model: "claude-test",
body: json!({"model": "claude-test", "messages": [{"role": "user", "content": "hi"}], "max_tokens": 16}),
MessagesCall {
body: body(
json!({"model": "claude-test", "messages": [{"role": "user", "content": "hi"}], "max_tokens": 16}),
),
api_key: None,
api_base: None,
custom_llm_provider: Some("anthropic"),
custom_llm_provider: Some("anthropic".into()),
extra_headers: None,
provider_specific_header: None,
timeout: None,
@ -235,8 +252,8 @@ mod tests {
&lookup,
)
.unwrap();
let auth: Vec<(&str, &str)> = prepared
.upstream_headers
let headers = wire_headers(&prepared);
let auth: Vec<(&str, &str)> = headers
.iter()
.filter(|(name, _)| matches!(name.as_str(), "x-api-key" | "authorization"))
.map(|(name, value)| (name.as_str(), value.as_str()))
@ -247,48 +264,18 @@ mod tests {
);
}
fn prepared_body(body: Value, shaping: MessagesShaping) -> Result<Value, Error> {
prepare(MessagesRequest {
model: "anthropic/claude-test",
body,
api_key: Some("sk-test"),
api_base: Some("https://anthropic.test"),
custom_llm_provider: Some("anthropic"),
fn prepared_body(fields: Value, shaping: MessagesShaping) -> Result<Value, Error> {
prepare(MessagesCall {
body: body(fields),
api_key: Some("sk-test".into()),
api_base: Some("https://anthropic.test".into()),
custom_llm_provider: Some("anthropic".into()),
extra_headers: None,
provider_specific_header: None,
timeout: None,
shaping,
})
.map(|prepared| prepared.body)
}
#[rstest]
#[case::nothing_forwarded(
&[],
&[("x-version", "1"), ("content-type", "application/json")],
&[("x-version", "1"), ("content-type", "application/json")],
)]
#[case::forwarded_header_wins_in_any_case(
&[("X-Version", "custom"), ("x-api-key", "k")],
&[("x-version", "1"), ("content-type", "application/json")],
&[("X-Version", "custom"), ("x-api-key", "k"), ("content-type", "application/json")],
)]
#[case::no_defaults(&[("x-api-key", "k")], &[], &[("x-api-key", "k")])]
fn default_headers_fill_only_missing_names(
#[case] forwarded: &[(&str, &str)],
#[case] defaults: &[(&str, &str)],
#[case] expected: &[(&str, &str)],
) {
let owned = |headers: &[(&str, &str)]| -> Vec<(String, String)> {
headers
.iter()
.map(|(name, value)| ((*name).to_string(), (*value).to_string()))
.collect()
};
assert_eq!(
with_default_headers(owned(forwarded), defaults),
owned(expected)
);
.map(|prepared| serde_json::to_value(prepared.body).unwrap())
}
#[rstest]
@ -380,20 +367,22 @@ mod tests {
{"custom_llm_provider": "anthropic", "extra_headers": {"x-scoped": "anthropic", "x-priority": "scoped"}}
]))
.unwrap();
let prepared = prepare(MessagesRequest {
model,
body: json!({"model": model, "messages": [{"role": "user", "content": "hi"}], "max_tokens": 16}),
api_key: Some("sk-test"),
api_base: Some("https://resource.services.ai.azure.com"),
custom_llm_provider,
extra_headers: Some(serde_json::from_value(json!({"x-priority": "extra"})).unwrap()),
let prepared = prepare(MessagesCall {
body: body(
json!({"model": model, "messages": [{"role": "user", "content": "hi"}], "max_tokens": 16}),
),
api_key: Some("sk-test".into()),
api_base: Some("https://resource.services.ai.azure.com".into()),
custom_llm_provider: custom_llm_provider.map(Into::into),
extra_headers: Some(Map::from_iter([("x-priority".into(), json!("extra"))])),
provider_specific_header: Some(configured),
timeout: None,
shaping,
})
.unwrap();
let caller_headers: Vec<(&str, &str)> = prepared
.upstream_headers
.environment
.headers
.iter()
.filter(|(name, _)| matches!(name.as_str(), "x-priority" | "x-scoped"))
.map(|(name, value)| (name.as_str(), value.as_str()))

View file

@ -1,50 +1,20 @@
use std::{
convert::Infallible,
sync::{Arc, Mutex},
time::Duration,
};
use bytes::Bytes;
use litellm_auth::SecretValue;
use futures_util::TryStreamExt;
use litellm_host::{
event::{MachineEvent, RawResponse, RequestContext, WireRequest},
host::{Demand, Host},
machine::{CallMachine, HostChannel, MachineFault},
protocol::Protocol,
};
use litellm_http::{Client, ClientVariant, HttpClientConfig};
use litellm_secrets::source::SecretSource;
use litellm_types::{
llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse,
utils::ProviderSpecificHeaders,
};
use serde_json::{Map, Value};
use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse;
use super::{
Error,
handler::{decode_response, network, provider_error, send},
prepare::{prepare_provider_request, resolve_provider},
types::{MessagesRequest, MessagesShaping},
};
use crate::constants::ANTHROPIC_MESSAGES_PROVIDER;
/// The caller's request as the host projects it.
pub struct MessagesCall {
pub model: String,
pub body: Map<String, Value>,
pub api_key: Option<String>,
pub api_base: Option<String>,
pub custom_llm_provider: Option<String>,
pub extra_headers: Option<Map<String, Value>>,
pub provider_specific_header: Option<ProviderSpecificHeaders>,
pub timeout: Option<Duration>,
pub shaping: MessagesShaping,
}
impl MessagesCall {
fn streams(&self) -> bool {
self.body.get("stream").and_then(Value::as_bool) == Some(true)
}
}
use super::{Error, MessagesCall, MessagesResponse, handler::execute, prepare::prepare};
pub enum MessagesOutput {
Message(Box<AnthropicMessagesResponse>),
@ -108,98 +78,43 @@ impl Host<Messages> for LocalMessagesHost {
}
}
pub fn messages_machine(secrets: Arc<dyn SecretSource>) -> MessagesMachine {
CallMachine::new(move |host| Box::pin(execute(host, secrets.clone())))
pub fn messages_machine(
resources: &crate::resources::CoreResources,
config: &HttpClientConfig,
secrets: Arc<dyn SecretSource>,
) -> Result<MessagesMachine, litellm_http::Error> {
let http = resources.pool.client(config, ClientVariant::Provider)?;
let auth = resources.auth.clone();
Ok(CallMachine::new(move |host| {
Box::pin(drive(host, http, auth, secrets))
}))
}
async fn execute(
/// The call as its host sees it: projection first, then the same prepare and execute as
/// [`super::messages`], with each chunk of a stream handed over as it arrives.
async fn drive(
host: MessagesHost,
http: Client,
auth: Arc<litellm_auth::AuthServices>,
secrets: Arc<dyn SecretSource>,
) -> Result<MessagesOutput, Error> {
let call = host.project().await?;
let stream = call.streams();
let resolved = resolve_provider(&call.model, call.custom_llm_provider.as_deref())?;
let secrets = secrets.resolve(resolved.config.secret_names()).await?;
let request = prepare_provider_request(
MessagesRequest {
model: &call.model,
body: Value::Object(call.body.clone()),
api_key: call.api_key.as_deref(),
api_base: call.api_base.as_deref(),
custom_llm_provider: call.custom_llm_provider.as_deref(),
extra_headers: call.extra_headers.clone(),
provider_specific_header: call.provider_specific_header.clone(),
timeout: call.timeout,
shaping: call.shaping.clone(),
},
resolved,
secrets.as_ref(),
)?;
if stream && request.provider != ANTHROPIC_MESSAGES_PROVIDER {
return Err(Error::Unsupported("streaming messages for this provider"));
}
let context = RequestContext {
model: request.model.clone(),
custom_llm_provider: request.provider.clone(),
optional_params: Value::Object(
request
.body
.as_object()
.into_iter()
.flatten()
.filter(|(name, _)| !matches!(name.as_str(), "model" | "messages"))
.map(|(name, value)| (name.clone(), value.clone()))
.collect(),
),
secret_fields: Vec::new(),
api_key: call.api_key.clone().map(SecretValue::new),
};
let wire = host
.before_send(
WireRequest {
url: request.url,
headers: request.upstream_headers,
body: request.body,
},
context,
)
.await?;
let response = send(&wire.url, &wire.headers, &wire.body, request.timeout).await?;
if !response.status().is_success() {
return Err(provider_error(response).await);
}
if stream {
return relay(&host, response).await;
}
let text = response.text().await.map_err(network)?;
host.emit(MachineEvent::ResponseReceived {
raw: RawResponse { body: text.clone() },
})
.await?;
decode_response(request.config, &request.model, &text)
.map(|message| MessagesOutput::Message(Box::new(message)))
}
/// Hands each upstream chunk to the caller as it arrives. A caller that stops reading
/// ends the upstream read, and the call completes with what it delivered.
async fn relay(
host: &MessagesHost,
mut response: reqwest::Response,
) -> Result<MessagesOutput, Error> {
let head = MessagesStreamHead {
headers: response
.headers()
.iter()
.filter_map(|(name, value)| Some((name.to_string(), value.to_str().ok()?.to_string())))
.collect(),
};
if host.open(head).await? == Demand::Detached {
return Ok(MessagesOutput::Streamed);
}
while let Some(chunk) = response.chunk().await.map_err(network)? {
if host.deliver(chunk).await? == Demand::Detached {
break;
let request = prepare(call, secrets.as_ref()).await?;
match execute(&http, &auth, request, &host).await? {
MessagesResponse::Message(message) => Ok(MessagesOutput::Message(message)),
MessagesResponse::Stream {
headers,
mut chunks,
} => {
if host.open(MessagesStreamHead { headers }).await? == Demand::Detached {
return Ok(MessagesOutput::Streamed);
}
while let Some(chunk) = chunks.try_next().await? {
if host.deliver(chunk).await? == Demand::Detached {
break;
}
}
Ok(MessagesOutput::Streamed)
}
}
Ok(MessagesOutput::Streamed)
}

View file

@ -1,13 +1,46 @@
use std::time::Duration;
use litellm_llms::{
anthropic::common_utils::AnthropicModelCapabilities,
base_llm::anthropic_messages::transformation::BaseAnthropicMessagesConfig,
use bytes::Bytes;
use futures_util::stream::BoxStream;
use litellm_llms::anthropic::common_utils::AnthropicModelCapabilities;
use litellm_types::{
llms::anthropic_messages::{
anthropic_request::AnthropicMessagesRequest, anthropic_response::AnthropicMessagesResponse,
},
utils::ProviderSpecificHeaders,
};
use litellm_types::utils::ProviderSpecificHeaders;
use serde::{Deserialize, Serialize};
use serde_json::{Map, Value};
use super::Error;
pub struct MessagesCall {
pub body: AnthropicMessagesRequest,
pub api_key: Option<String>,
pub api_base: Option<String>,
pub custom_llm_provider: Option<String>,
pub extra_headers: Option<Map<String, Value>>,
pub provider_specific_header: Option<ProviderSpecificHeaders>,
pub timeout: Option<Duration>,
pub shaping: MessagesShaping,
}
pub fn messages_body(body: Map<String, Value>) -> Result<AnthropicMessagesRequest, Error> {
serde_json::from_value(Value::Object(body)).map_err(invalid_request)
}
pub(super) fn invalid_request(err: serde_json::Error) -> Error {
Error::InvalidRequest(format!("invalid Anthropic messages request: {err}"))
}
pub enum MessagesResponse {
Message(Box<AnthropicMessagesResponse>),
Stream {
headers: Vec<(String, String)>,
chunks: BoxStream<'static, Result<Bytes, Error>>,
},
}
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
pub struct MessagesShaping {
#[serde(default)]
@ -20,33 +53,11 @@ pub struct MessagesShaping {
pub additional_drop_params: Vec<String>,
}
pub struct MessagesRequest<'a> {
pub model: &'a str,
pub body: Value,
pub api_key: Option<&'a str>,
pub api_base: Option<&'a str>,
pub custom_llm_provider: Option<&'a str>,
pub extra_headers: Option<Map<String, Value>>,
pub provider_specific_header: Option<ProviderSpecificHeaders>,
pub timeout: Option<Duration>,
pub shaping: MessagesShaping,
}
pub struct ProviderMessagesRequest {
pub provider: String,
pub model: String,
pub config: &'static dyn BaseAnthropicMessagesConfig,
pub url: String,
pub body: Value,
pub upstream_headers: Vec<(String, String)>,
pub timeout: Option<Duration>,
}
#[cfg(test)]
mod tests {
use litellm_llms::anthropic::common_utils::SupportedEffortTiers;
use rstest::rstest;
use serde_json::json;
use serde_json::{Value, json};
use super::*;

View file

@ -110,7 +110,10 @@ mod tests {
}
fn client() -> OcrClient {
OcrClient::for_test(reqwest::Client::new(), reqwest::Client::new())
OcrClient::for_test(
litellm_http::Client::plain_for_test(),
litellm_http::Client::no_redirect_for_test(),
)
}
fn request(model: &str, base: &str, document: Value, options: Value) -> LiteLLMOcrRequest {

View file

@ -1,30 +1,20 @@
use std::time::Duration;
use litellm_auth::RequestAuth;
use litellm_auth_aws::SigV4Signer;
use litellm_http::outbound::OutboundRequest;
use serde_json::{Map, Value};
use litellm_llms::base_llm::auth::Authenticated;
use serde_json::Value;
/// Header credentials are already in `headers`; SigV4 is applied here, over the
/// bytes that are sent.
pub(crate) async fn outbound_request<E>(
auth: &RequestAuth,
pub(crate) fn outbound_request(
authenticated: Authenticated,
url: String,
headers: Vec<(String, String)>,
body: &Value,
timeout: Option<Duration>,
optional_params: &Map<String, Value>,
) -> Result<OutboundRequest, E>
where
E: From<litellm_http::Error> + From<litellm_auth_aws::Error>,
{
let RequestAuth::AwsSigV4 { region, service } = auth else {
return Ok(OutboundRequest::json(url, headers, body, timeout)?);
};
let env_lookup = |key: &str| std::env::var(key).ok();
let signer =
SigV4Signer::resolve(region.clone(), service, optional_params, &env_lookup).await?;
Ok(OutboundRequest::signed_json(
url, headers, body, timeout, &signer,
)?)
) -> Result<OutboundRequest, litellm_http::Error> {
let Authenticated { headers, signer } = authenticated;
match signer {
None => OutboundRequest::json(url, headers, body, timeout),
Some(signer) => OutboundRequest::signed_json(url, headers, body, timeout, &signer),
}
}

View file

@ -0,0 +1,38 @@
use std::sync::Arc;
use litellm_auth::AuthServices;
use litellm_http::{HttpClientConfig, HttpClientPool, media::UrlPolicy};
use litellm_llms::base_llm::ocr::{handler::OcrClient, settings::OcrSettings};
use litellm_secrets::source::SecretSource;
#[derive(Clone)]
pub struct CoreResources {
pub pool: Arc<HttpClientPool>,
pub auth: Arc<AuthServices>,
}
impl CoreResources {
pub fn new(pool: Arc<HttpClientPool>) -> Self {
Self {
pool,
auth: Arc::new(AuthServices::default()),
}
}
pub fn ocr_client(
&self,
config: &HttpClientConfig,
url_policy: UrlPolicy,
settings: OcrSettings,
secrets: Arc<dyn SecretSource>,
) -> Result<OcrClient, litellm_http::Error> {
OcrClient::new(
&self.pool,
config,
url_policy,
self.auth.clone(),
settings,
secrets,
)
}
}

View file

@ -1,17 +0,0 @@
#[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)]
pub enum Error {
#[error("invalid provider: {0}")]
InvalidProvider(String),
#[error("invalid request: {0}")]
InvalidRequest(String),
#[error("invalid response: {0}")]
InvalidResponse(String),
#[error("routing error: {0}")]
Routing(String),
#[error(transparent)]
Auth(#[from] litellm_auth::Error),
#[error(transparent)]
Transport(#[from] litellm_http::transport::Error),
#[error(transparent)]
Headers(#[from] litellm_http::request::HeaderError),
}

View file

@ -1,3 +1,2 @@
mod error;
pub use error::Error;
pub use crate::error::RouteError as Error;
pub mod websocket;

View file

@ -10,6 +10,10 @@ use support::*;
const MODEL: &str = "mistral.voxtral-mini-3b-2507";
async fn transcribe(request: AudioTranscriptionRequest<'_>) -> Result<Value, Error> {
audio_transcription(&support::resources(), &http_config(), request).await
}
fn transcript_response(text: &str) -> ResponseTemplate {
json_response(json!({"output": {"message": {"content": [{"text": text}]}}}))
}
@ -47,7 +51,7 @@ async fn bedrock_converse_request_is_signed_for_the_requested_region(
let upstream = upstream([transcript_response("hello")]).await;
let base = upstream.uri();
let response = audio_transcription(AudioTranscriptionRequest {
let response = transcribe(AudioTranscriptionRequest {
api_base: Some(&base),
optional_params: aws_params(region),
..request
@ -79,7 +83,7 @@ async fn the_provider_can_come_from_the_model_prefix(request: AudioTranscription
let base = upstream.uri();
let model = format!("bedrock/{MODEL}");
audio_transcription(AudioTranscriptionRequest {
transcribe(AudioTranscriptionRequest {
model: &model,
custom_llm_provider: None,
api_base: Some(&base),
@ -110,7 +114,7 @@ async fn audio_and_transcription_params_reach_the_converse_body(
])
.collect();
audio_transcription(AudioTranscriptionRequest {
transcribe(AudioTranscriptionRequest {
audio: json!({"data": "AQI=", "format": format}),
api_base: Some(&base),
optional_params,
@ -142,7 +146,7 @@ async fn invalid_audio_is_rejected_before_sending(
let upstream = upstream([transcript_response("hello")]).await;
let base = upstream.uri();
let error = audio_transcription(AudioTranscriptionRequest {
let error = transcribe(AudioTranscriptionRequest {
audio,
api_base: Some(&base),
..request
@ -174,7 +178,7 @@ async fn unsupported_providers_are_rejected_before_sending(
#[case] provider: Option<&'static str>,
#[case] reported: &str,
) {
let error = audio_transcription(AudioTranscriptionRequest {
let error = transcribe(AudioTranscriptionRequest {
model,
custom_llm_provider: provider,
api_base: Some(UNREACHABLE_BASE),
@ -189,7 +193,7 @@ async fn unsupported_providers_are_rejected_before_sending(
#[rstest]
#[tokio::test]
async fn a_non_string_extra_header_is_rejected(request: AudioTranscriptionRequest<'static>) {
let error = audio_transcription(AudioTranscriptionRequest {
let error = transcribe(AudioTranscriptionRequest {
extra_headers: Some(Map::from_iter([("x-count".to_string(), json!(3))])),
api_base: Some(UNREACHABLE_BASE),
..request
@ -212,7 +216,7 @@ async fn an_upstream_error_keeps_its_status_and_body(
upstream([ResponseTemplate::new(status).set_body_string("upstream said no")]).await;
let base = upstream.uri();
let error = audio_transcription(AudioTranscriptionRequest {
let error = transcribe(AudioTranscriptionRequest {
api_base: Some(&base),
..request
})
@ -239,7 +243,7 @@ async fn an_unreadable_success_body_is_an_invalid_response(
let upstream = upstream([response]).await;
let base = upstream.uri();
let error = audio_transcription(AudioTranscriptionRequest {
let error = transcribe(AudioTranscriptionRequest {
api_base: Some(&base),
..request
})

View file

@ -4,6 +4,7 @@ use litellm_core::chat_completions::{
Error, chat_completions, chat_completions_decline_reason, types::ChatCompletionsRequest,
};
use litellm_http::transport::Error as TransportError;
use litellm_types::utils::ChatCompletionsResponse;
use rstest::{fixture, rstest};
use serde_json::{Map, Value, json};
use wiremock::ResponseTemplate;
@ -13,6 +14,10 @@ use support::*;
const ANTHROPIC_MESSAGE: &str = r#"{"id":"msg_1","type":"message","role":"assistant","model":"claude-sonnet-4-5-20260101","content":[{"type":"text","text":"hello"}],"stop_reason":"end_turn","stop_sequence":null,"usage":{"input_tokens":11,"output_tokens":4}}"#;
async fn complete(request: ChatCompletionsRequest<'_>) -> Result<ChatCompletionsResponse, Error> {
chat_completions(&support::resources(), &http_config(), request).await
}
fn object(value: Value) -> Map<String, Value> {
let Value::Object(map) = value else {
panic!("expected a json object, got {value}");
@ -50,7 +55,7 @@ async fn anthropic_round_trip_translates_the_conversation_and_normalizes_the_res
let upstream = upstream([anthropic_response(ANTHROPIC_MESSAGE)]).await;
let base = upstream.uri();
let response = chat_completions(ChatCompletionsRequest {
let response = complete(ChatCompletionsRequest {
messages: json!([
{"role": "system", "content": "be terse"},
{"role": "user", "content": "hi"}
@ -90,7 +95,7 @@ async fn the_deployment_key_replaces_a_caller_supplied_x_api_key(
let upstream = upstream([anthropic_response(ANTHROPIC_MESSAGE)]).await;
let base = upstream.uri();
chat_completions(ChatCompletionsRequest {
complete(ChatCompletionsRequest {
api_base: Some(&base),
extra_headers: Some(object(
json!({"x-api-key": "caller-key", "x-trace": "kept"}),
@ -116,7 +121,7 @@ async fn bedrock_round_trip_is_signed_and_normalized(request: ChatCompletionsReq
.await;
let base = upstream.uri();
let response = chat_completions(ChatCompletionsRequest {
let response = complete(ChatCompletionsRequest {
model: "bedrock/anthropic.claude-sonnet-4-5",
optional_params: object(json!({
"aws_access_key_id": "access-key",
@ -167,7 +172,7 @@ async fn a_response_it_cannot_normalize_is_reported_as_already_sent(
let upstream = upstream([anthropic_response(body)]).await;
let base = upstream.uri();
let error = chat_completions(ChatCompletionsRequest {
let error = complete(ChatCompletionsRequest {
api_base: Some(&base),
..request
})
@ -188,7 +193,7 @@ async fn an_upstream_error_status_keeps_its_code_and_body(
let upstream = upstream([ResponseTemplate::new(status).set_body_string("slow down")]).await;
let base = upstream.uri();
let error = chat_completions(ChatCompletionsRequest {
let error = complete(ChatCompletionsRequest {
api_base: Some(&base),
..request
})
@ -210,7 +215,7 @@ async fn an_upstream_error_status_keeps_its_code_and_body(
async fn a_connection_that_is_never_established_declines_instead_of_failing(
request: ChatCompletionsRequest<'static>,
) {
let error = chat_completions(ChatCompletionsRequest {
let error = complete(ChatCompletionsRequest {
api_base: Some(UNREACHABLE_BASE),
..request
})
@ -232,7 +237,7 @@ async fn a_timeout_after_sending_is_not_a_pre_send_decline(
upstream([anthropic_response(ANTHROPIC_MESSAGE).set_delay(Duration::from_secs(5))]).await;
let base = upstream.uri();
let error = chat_completions(ChatCompletionsRequest {
let error = complete(ChatCompletionsRequest {
api_base: Some(&base),
timeout: Some(Duration::from_millis(100)),
..request
@ -307,7 +312,7 @@ async fn a_declined_request_fails_the_call_before_sending(
let upstream = upstream([anthropic_response(ANTHROPIC_MESSAGE)]).await;
let base = upstream.uri();
let error = chat_completions(ChatCompletionsRequest {
let error = complete(ChatCompletionsRequest {
optional_params: object(json!({"stream": true})),
api_base: Some(&base),
..request

View file

@ -78,7 +78,7 @@ impl Host<Messages> for RecordingHost {
}
async fn run_through(host: &RecordingHost) -> Result<MessagesOutput, Error> {
litellm_host::run::run(messages_machine(Arc::new(RecordingSecrets::empty())), host).await
litellm_host::run::run(machine(Arc::new(RecordingSecrets::empty())), host).await
}
fn authenticated(call: MessagesCall, api_base: String) -> MessagesCall {
@ -161,10 +161,10 @@ async fn no_raw_response_is_emitted_for_a_stream_or_a_failure(
#[case] response: ResponseTemplate,
) {
let upstream = upstream([response]).await;
let mut body = call.body.clone();
body.insert("stream".into(), json!(true));
let host =
RecordingHost::passthrough(authenticated(MessagesCall { body, ..call }, upstream.uri()));
let host = RecordingHost::passthrough(authenticated(
with_fields(call, json!({"stream": true})),
upstream.uri(),
));
let _ = run_through(&host).await;
@ -180,15 +180,8 @@ async fn the_request_context_carries_the_shaped_params_without_model_or_messages
call: MessagesCall,
) {
let upstream = upstream([message_response()]).await;
let body: Map<String, Value> = call
.body
.clone()
.into_iter()
.chain([("temperature".to_string(), json!(0.2))])
.collect();
let host = RecordingHost::passthrough(authenticated(
MessagesCall {
body,
shaping: MessagesShaping {
capabilities: AnthropicModelCapabilities {
supports_sampling_params: false,
@ -197,7 +190,7 @@ async fn the_request_context_carries_the_shaped_params_without_model_or_messages
drop_params: true,
..MessagesShaping::default()
},
..call
..with_fields(call, json!({"temperature": 0.2}))
},
upstream.uri(),
));

View file

@ -1,11 +1,14 @@
use std::{sync::Arc, time::Duration};
use litellm_core::messages::{
Error,
route::{LocalMessagesHost, MessagesCall, MessagesOutput, messages_machine},
types::MessagesShaping,
Error, MessagesCall, MessagesShaping,
route::{LocalMessagesHost, MessagesMachine, MessagesOutput, messages_machine},
};
use litellm_http::{HttpSettings, Resolution};
use litellm_secrets::source::SecretSource;
use litellm_types::llms::anthropic_messages::{
anthropic_request::AnthropicMessagesRequest, anthropic_response::AnthropicMessagesResponse,
};
use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse;
use rstest::fixture;
use serde_json::{Map, Value, json};
use wiremock::ResponseTemplate;
@ -29,6 +32,24 @@ fn object(value: Value) -> Map<String, Value> {
map
}
fn body(value: Value) -> AnthropicMessagesRequest {
serde_json::from_value(value).unwrap()
}
fn with_fields(call: MessagesCall, fields: Value) -> MessagesCall {
let current = object(serde_json::to_value(&call.body).unwrap());
MessagesCall {
body: body(Value::Object(
current.into_iter().chain(object(fields)).collect(),
)),
..call
}
}
fn with_model(call: MessagesCall, model: &str) -> MessagesCall {
with_fields(call, json!({"model": model}))
}
fn message_body() -> Value {
json!({
"id": "msg_1",
@ -50,8 +71,7 @@ fn message_response() -> ResponseTemplate {
#[fixture]
fn call() -> MessagesCall {
MessagesCall {
model: MODEL.into(),
body: object(json!({
body: body(json!({
"model": MODEL,
"max_tokens": 16,
"messages": [{"role": "user", "content": "hi"}]
@ -75,11 +95,16 @@ fn headers<'a>(pairs: impl IntoIterator<Item = (&'a str, &'a str)>) -> Option<Ma
)
}
fn machine(secrets: Arc<dyn SecretSource>) -> MessagesMachine {
messages_machine(&support::resources(), &http_config(), secrets)
.expect("default HTTP settings build a client")
}
async fn run_with(
secrets: Arc<RecordingSecrets>,
call: MessagesCall,
) -> Result<MessagesOutput, Error> {
litellm_host::run::run(messages_machine(secrets), &LocalMessagesHost::new(call)).await
litellm_host::run::run(machine(secrets), &LocalMessagesHost::new(call)).await
}
/// Runs the route with a secret source that knows nothing, so no environment leaks in.

View file

@ -1,7 +1,5 @@
use litellm_llms::anthropic::common_utils::{
ANTHROPIC_ADVISOR_TOOL_TYPE, ANTHROPIC_OAUTH_BETA_HEADER, AnthropicModelCapabilities,
SupportedEffortTiers, beta,
};
use litellm_llms::anthropic::common_utils::{AnthropicModelCapabilities, SupportedEffortTiers};
use litellm_types::llms::anthropic::{AnthropicBeta, BetaSet};
use litellm_types::utils::{ProviderSpecificHeader, ProviderSpecificHeaders};
use rstest::rstest;
@ -124,11 +122,10 @@ async fn each_provider_posts_to_its_messages_endpoint(
let upstream = upstream([message_response()]).await;
run_message(MessagesCall {
model: model.into(),
custom_llm_provider: provider.map(Into::into),
api_key: Some("sk".into()),
api_base: Some(format!("{}{base_suffix}", upstream.uri())),
..call
..with_model(call, model)
})
.await;
@ -155,11 +152,10 @@ async fn unsupported_providers_are_rejected_before_sending(
#[case] reported: &str,
) {
let error = run(MessagesCall {
model: model.into(),
custom_llm_provider: provider.map(Into::into),
api_key: Some("sk".into()),
api_base: Some(UNREACHABLE_BASE.into()),
..call
..with_model(call, model)
})
.await
.err()
@ -206,7 +202,7 @@ async fn azure_strips_the_cache_control_scope_anthropic_rejects(call: MessagesCa
custom_llm_provider: Some("azure_ai".into()),
api_key: Some("sk-azure".into()),
api_base: Some(upstream.uri()),
body: object(json!({
body: body(json!({
"model": MODEL,
"max_tokens": 16,
"messages": [{
@ -232,19 +228,15 @@ async fn azure_strips_the_cache_control_scope_anthropic_rejects(call: MessagesCa
#[tokio::test]
async fn additional_drop_params_remove_fields_before_sending(call: MessagesCall) {
let upstream = upstream([message_response()]).await;
let mut body = call.body.clone();
body.insert("temperature".into(), json!(0.5));
body.insert("top_k".into(), json!(3));
run_message(MessagesCall {
api_key: Some("sk".into()),
api_base: Some(upstream.uri()),
body,
shaping: MessagesShaping {
additional_drop_params: vec!["temperature".into()],
..MessagesShaping::default()
},
..call
..with_fields(call, json!({"temperature": 0.5, "top_k": 3}))
})
.await;
@ -253,46 +245,37 @@ async fn additional_drop_params_remove_fields_before_sending(call: MessagesCall)
assert_eq!(sent["top_k"], 3);
}
fn with_fields(call: MessagesCall, fields: Value) -> MessagesCall {
let body: Map<String, Value> = call.body.into_iter().chain(object(fields)).collect();
MessagesCall { body, ..call }
}
fn sent_betas(request: &wiremock::Request) -> Vec<String> {
fn sent_betas(request: &wiremock::Request) -> BetaSet {
let [header] = <[&str; 1]>::try_from(request.header_values("anthropic-beta"))
.unwrap_or_else(|values| panic!("expected one anthropic-beta header, got {values:?}"));
header
.split(',')
.map(str::trim)
.map(str::to_string)
.collect()
header.parse().unwrap()
}
#[rstest]
#[case::structured_output(json!({"output_format": {"type": "json_schema"}}), &[beta::STRUCTURED_OUTPUT])]
#[case::fast_mode(json!({"speed": "fast"}), &[beta::FAST_MODE_2026_02_01])]
#[case::compaction(json!({"compaction": {"enabled": true}}), &[beta::COMPACT_2026_09_04])]
#[case::structured_output(json!({"output_format": {"type": "json_schema"}}), &[AnthropicBeta::StructuredOutputs20251113])]
#[case::fast_mode(json!({"speed": "fast"}), &[AnthropicBeta::FastMode20260201])]
#[case::compaction(json!({"compaction": {"enabled": true}}), &[AnthropicBeta::Compact20260904])]
#[case::context_management_edits(
json!({"context_management": {"edits": [{"type": "clear_tool_uses_20250919"}]}}),
&[beta::CONTEXT_MANAGEMENT_2025_06_27]
&[AnthropicBeta::ContextManagement20250627]
)]
#[case::per_message_output_config(
json!({"messages": [{"role": "user", "content": "hi", "output_config": {"effort": "low"}}]}),
&[beta::PER_TURN_CONTROL_2026_07_01]
&[AnthropicBeta::PerTurnControl20260701]
)]
#[case::advisor_tool(
json!({"tools": [{"type": ANTHROPIC_ADVISOR_TOOL_TYPE, "name": "advisor", "model": MODEL}]}),
&[beta::ADVISOR_TOOL_2026_03_01]
json!({"tools": [{"type": "advisor_20260301", "name": "advisor", "model": MODEL}]}),
&[AnthropicBeta::AdvisorTool20260301]
)]
#[case::several_features_at_once(
json!({"speed": "fast", "output_format": {"type": "json_schema"}}),
&[beta::STRUCTURED_OUTPUT, beta::FAST_MODE_2026_02_01]
&[AnthropicBeta::StructuredOutputs20251113, AnthropicBeta::FastMode20260201]
)]
#[tokio::test]
async fn feature_betas_join_the_callers_betas_in_one_sorted_header(
call: MessagesCall,
#[case] fields: Value,
#[case] features: &[&str],
#[case] features: &[AnthropicBeta],
) {
let upstream = upstream([message_response()]).await;
let capabilities = AnthropicModelCapabilities {
@ -316,12 +299,11 @@ async fn feature_betas_join_the_callers_betas_in_one_sorted_header(
.await;
let sent = sent_betas(&only_request(&upstream).await);
let mut expected: Vec<String> = features
let expected: BetaSet = features
.iter()
.map(|feature| feature.to_string())
.chain(["caller-beta-2025-01-01".to_string()])
.cloned()
.chain([AnthropicBeta::Other("caller-beta-2025-01-01".to_string())])
.collect();
expected.sort();
assert_eq!(sent, expected);
}
@ -342,7 +324,10 @@ async fn an_oauth_key_sends_the_browser_access_header_and_the_oauth_beta(call: M
request.header("anthropic-dangerous-direct-browser-access"),
Some("true")
);
assert_eq!(sent_betas(&request), [ANTHROPIC_OAUTH_BETA_HEADER]);
assert_eq!(
sent_betas(&request),
BetaSet::from_iter([AnthropicBeta::Oauth20250420])
);
assert_eq!(request.header("x-api-key"), None);
}
@ -406,7 +391,6 @@ async fn unsupported_params_are_dropped_under_drop_params_and_rejected_without_i
custom_llm_provider: call.custom_llm_provider.clone(),
extra_headers: None,
provider_specific_header: None,
model: call.model.clone(),
timeout: call.timeout,
},
fields.clone(),
@ -664,10 +648,9 @@ async fn the_provider_prefix_is_stripped_exactly_once(
let upstream = upstream([message_response()]).await;
run_message(MessagesCall {
model: model.into(),
api_key: Some("sk".into()),
api_base: Some(upstream.uri()),
..call
..with_model(call, model)
})
.await;

View file

@ -1,4 +1,7 @@
use litellm_core::messages::{messages, types::MessagesRequest};
use litellm_core::{
Phase,
messages::{MessagesResponse, messages, messages_body},
};
use litellm_http::transport::Error as TransportError;
use rstest::rstest;
@ -154,7 +157,7 @@ async fn an_unreadable_success_body_is_an_invalid_response(
.err()
.expect("an unreadable body fails");
assert!(error.is_response(), "{error:?}");
assert_eq!(error.phase(), Phase::AfterSend, "{error:?}");
}
#[rstest]
@ -175,47 +178,46 @@ async fn a_provider_slower_than_the_timeout_fails_the_call(call: MessagesCall) {
assert!(matches!(error, Error::Transport(_)), "{error:?}");
}
fn facade_request(body: Value, api_base: &str) -> MessagesRequest<'_> {
MessagesRequest {
model: MODEL,
body,
api_key: Some("sk-ant"),
api_base: Some(api_base),
custom_llm_provider: Some("anthropic"),
extra_headers: None,
provider_specific_header: None,
timeout: Some(Duration::from_secs(5)),
shaping: MessagesShaping::default(),
}
}
#[rstest]
#[tokio::test]
async fn the_facade_runs_the_route_in_process() {
async fn the_facade_sends_through_the_injected_http_pool_configuration(call: MessagesCall) {
let upstream = upstream([message_response()]).await;
let base = upstream.uri();
let settings = HttpSettings {
user_agent: Some("host-owned/1".into()),
..HttpSettings::default()
};
let message = messages(facade_request(
json!({"model": MODEL, "max_tokens": 16, "messages": [{"role": "user", "content": "hi"}]}),
&base,
))
let response = messages(
&support::resources(),
&Resolution::from(&settings).config,
&RecordingSecrets::empty(),
MessagesCall {
api_key: Some("sk-ant".into()),
api_base: Some(base),
..call
},
)
.await
.expect("messages request succeeds");
let MessagesResponse::Message(message) = response else {
panic!("a non-streaming request returns a message");
};
assert_eq!(message.id, "msg_1");
assert_eq!(
only_request(&upstream).await.header("x-api-key"),
Some("sk-ant")
);
let sent = only_request(&upstream).await;
assert_eq!(sent.header("x-api-key"), Some("sk-ant"));
assert_eq!(sent.header("user-agent"), Some("host-owned/1"));
}
#[tokio::test]
async fn the_facade_rejects_a_body_that_is_not_an_object() {
let error = messages(facade_request(json!([]), UNREACHABLE_BASE))
.await
.expect_err("a non-object body is rejected");
#[rstest]
#[case::mistyped_param(json!({"model": MODEL, "messages": [], "max_tokens": "16"}))]
#[case::missing_messages(json!({"model": MODEL, "max_tokens": 16}))]
fn a_body_that_does_not_parse_is_an_invalid_request(#[case] raw: Value) {
let error = messages_body(object(raw)).expect_err("the body is rejected");
assert_eq!(
error,
Error::InvalidRequest("messages body must be an object".into())
assert!(
matches!(&error, Error::InvalidRequest(message) if message.starts_with("invalid Anthropic messages request: ")),
"{error:?}"
);
}

View file

@ -1,12 +1,21 @@
use std::{convert::Infallible, sync::Mutex};
use std::{
convert::Infallible,
sync::{Mutex, mpsc},
};
use bytes::Bytes;
use litellm_core::messages::route::{Messages, MessagesStreamHead};
use futures_util::{StreamExt, TryStreamExt};
use litellm_core::messages::{
MessagesResponse, messages,
route::{Messages, MessagesStreamHead},
};
use litellm_host::host::{Demand, Host};
use litellm_tracing::{Logger, Metadata, Record, Sink};
use rstest::rstest;
use tokio::{
io::{AsyncReadExt, AsyncWriteExt},
net::TcpListener,
task::JoinHandle,
};
use super::*;
@ -23,6 +32,20 @@ enum Seen {
Deliver(Bytes),
}
struct TraceSink(mpsc::Sender<(String, Value)>);
impl Sink for TraceSink {
fn enabled(&self, metadata: &Metadata<'_>) -> bool {
metadata.target().starts_with("litellm_core::messages")
}
fn emit(&self, record: &Record) {
self.0
.send((record.message.clone(), Value::Object(record.fields.clone())))
.unwrap();
}
}
/// Projects like `LocalMessagesHost`, records every stream op in the order the route
/// performs it, and detaches after `detach_after` ops.
struct RecordingStreamHost {
@ -69,13 +92,10 @@ impl Host<Messages> for RecordingStreamHost {
}
fn streaming(call: MessagesCall, api_base: String) -> MessagesCall {
let mut body = call.body.clone();
body.insert("stream".into(), json!(true));
MessagesCall {
api_key: Some("sk-ant".into()),
api_base: Some(api_base),
body,
..call
..with_fields(call, json!({"stream": true}))
}
}
@ -87,7 +107,7 @@ fn sse_response() -> ResponseTemplate {
}
async fn stream_through(host: &RecordingStreamHost) -> Result<MessagesOutput, Error> {
litellm_host::run::run(messages_machine(Arc::new(RecordingSecrets::empty())), host).await
litellm_host::run::run(machine(Arc::new(RecordingSecrets::empty())), host).await
}
#[rstest]
@ -123,6 +143,37 @@ async fn upstream_headers_are_on_the_stream_head_before_the_first_chunk(call: Me
assert_eq!(delivered, SSE_BODY.as_bytes());
}
#[rstest]
#[tokio::test]
async fn debug_trace_keeps_provider_input_and_every_stream_chunk(call: MessagesCall) {
let upstream = upstream([sse_response()]).await;
let host = RecordingStreamHost::new(streaming(call, upstream.uri()), usize::MAX);
let (sender, receiver) = mpsc::channel();
Logger::new(TraceSink(sender))
.instrument(stream_through(&host))
.await
.unwrap();
let records: Vec<(String, Value)> = receiver.try_iter().collect();
let request = records
.iter()
.find(|(message, _)| message == "provider request")
.unwrap();
let body: Value = serde_json::from_str(request.1["body"].as_str().unwrap()).unwrap();
assert_eq!(body["messages"][0]["content"], "hi");
assert_eq!(request.1["stream"], true);
let chunks: String = records
.iter()
.filter(|(message, fields)| {
message == "stream chunk" && fields["stage"] == "provider_response"
})
.map(|(_, fields)| fields["chunk"].as_str().unwrap())
.collect();
assert_eq!(chunks, SSE_BODY);
assert!(!format!("{records:?}").contains("sk-ant"));
}
#[rstest]
#[case::at_open(1)]
#[case::after_the_first_chunk(2)]
@ -195,10 +246,10 @@ async fn a_stream_that_ends_without_message_stop_is_relayed_as_is(call: Messages
}
/// Serves one SSE chunk and then holds the connection open without ever finishing.
async fn stalling_upstream() -> String {
async fn stalling_upstream() -> (String, JoinHandle<()>) {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let base = format!("http://{}", listener.local_addr().unwrap());
tokio::spawn(async move {
let connection = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.unwrap();
let mut request = vec![0; 4096];
let _ = socket.read(&mut request).await;
@ -209,15 +260,15 @@ async fn stalling_upstream() -> String {
)
.await
.unwrap();
std::future::pending::<()>().await;
let _ = socket.read_to_end(&mut Vec::new()).await;
});
base
(base, connection)
}
#[rstest]
#[tokio::test]
async fn the_timeout_covers_a_stalled_stream_body(call: MessagesCall) {
let base = stalling_upstream().await;
let (base, connection) = stalling_upstream().await;
let host = RecordingStreamHost::new(
MessagesCall {
timeout: Some(Duration::from_millis(300)),
@ -239,11 +290,150 @@ async fn the_timeout_covers_a_stalled_stream_body(call: MessagesCall) {
"the chunk before the stall reached the caller, saw {} ops",
seen.len()
);
tokio::time::timeout(Duration::from_secs(5), connection)
.await
.expect("timing out closes the upstream connection")
.unwrap();
}
#[rstest]
#[case::anthropic("anthropic")]
#[case::azure_ai("azure_ai")]
#[tokio::test]
async fn the_sdk_returns_stream_headers_and_every_sse_byte(
call: MessagesCall,
#[case] provider: &str,
) {
let upstream = upstream([sse_response()]).await;
let response = messages(
&support::resources(),
&http_config(),
&RecordingSecrets::empty(),
MessagesCall {
custom_llm_provider: Some(provider.into()),
..streaming(call, upstream.uri())
},
)
.await
.unwrap();
let MessagesResponse::Stream { headers, chunks } = response else {
panic!("a streaming request returns a stream");
};
for (name, value) in UPSTREAM_HEADERS {
assert!(headers.contains(&(name.into(), value.into())));
}
let delivered = chunks.try_collect::<Vec<_>>().await.unwrap().concat();
assert_eq!(delivered, SSE_BODY.as_bytes());
assert_eq!(only_request(&upstream).await.json()["stream"], true);
}
#[rstest]
#[tokio::test]
async fn streaming_is_refused_for_providers_that_cannot_stream(call: MessagesCall) {
async fn the_sdk_returns_http_errors_before_opening_a_stream(call: MessagesCall) {
let upstream = upstream([ResponseTemplate::new(429).set_body_string("slow down")]).await;
let error = messages(
&support::resources(),
&http_config(),
&RecordingSecrets::empty(),
streaming(call, upstream.uri()),
)
.await
.err()
.expect("upstream failure is returned by messages()");
assert_eq!(
error,
Error::Transport(litellm_http::transport::Error::Http {
status: 429,
body: "slow down".into(),
})
);
}
#[rstest]
#[case::before_reading(false)]
#[case::after_reading(true)]
#[tokio::test]
async fn dropping_the_sdk_stream_closes_the_unfinished_upstream(
call: MessagesCall,
#[case] read_chunk: bool,
) {
let (base, connection) = stalling_upstream().await;
let response = tokio::time::timeout(
Duration::from_secs(5),
messages(
&support::resources(),
&http_config(),
&RecordingSecrets::empty(),
MessagesCall {
timeout: Some(Duration::from_secs(30)),
..streaming(call, base)
},
),
)
.await
.expect("messages() returns before the upstream finishes")
.unwrap();
let MessagesResponse::Stream { mut chunks, .. } = response else {
panic!("a streaming request returns a stream");
};
if read_chunk {
let chunk = tokio::time::timeout(Duration::from_secs(5), chunks.next())
.await
.expect("the first chunk arrives before the upstream finishes")
.unwrap()
.unwrap();
assert_eq!(chunk.as_ref(), b"event: message_start\ndata: {}\n\n");
}
assert!(!connection.is_finished());
drop(chunks);
tokio::time::timeout(Duration::from_secs(5), connection)
.await
.expect("dropping the stream closes the upstream connection")
.unwrap();
}
#[rstest]
#[tokio::test]
async fn the_sdk_yields_a_body_error_once_after_delivered_chunks(call: MessagesCall) {
let (base, connection) = stalling_upstream().await;
let response = messages(
&support::resources(),
&http_config(),
&RecordingSecrets::empty(),
MessagesCall {
timeout: Some(Duration::from_millis(300)),
..streaming(call, base)
},
)
.await
.unwrap();
let MessagesResponse::Stream { mut chunks, .. } = response else {
panic!("a streaming request returns a stream");
};
assert_eq!(
chunks.next().await.unwrap().unwrap().as_ref(),
b"event: message_start\ndata: {}\n\n"
);
let error = tokio::time::timeout(Duration::from_secs(5), chunks.next())
.await
.expect("the stalled body times out")
.unwrap()
.unwrap_err();
assert!(matches!(error, Error::Transport(_)), "{error:?}");
assert!(chunks.next().await.is_none());
tokio::time::timeout(Duration::from_secs(5), connection)
.await
.expect("the failed stream closes its upstream connection")
.unwrap();
}
#[rstest]
#[tokio::test]
async fn a_host_on_anthropic_sse_is_relayed_byte_for_byte(call: MessagesCall) {
let upstream = upstream([sse_response()]).await;
let host = RecordingStreamHost::new(
MessagesCall {
@ -253,14 +443,17 @@ async fn streaming_is_refused_for_providers_that_cannot_stream(call: MessagesCal
usize::MAX,
);
let error = stream_through(&host)
.await
.err()
.expect("azure streaming is refused");
let outcome = stream_through(&host).await.expect("azure streams");
assert_eq!(
error,
Error::Unsupported("streaming messages for this provider")
);
assert!(received(&upstream).await.is_empty());
assert!(matches!(outcome, MessagesOutput::Streamed));
let seen = host.seen.into_inner().unwrap();
let delivered: Vec<u8> = seen
.iter()
.filter_map(|step| match step {
Seen::Deliver(chunk) => Some(chunk.to_vec()),
Seen::Open(_) => None,
})
.flatten()
.collect();
assert_eq!(delivered, SSE_BODY.as_bytes());
}

View file

@ -4,6 +4,7 @@ use litellm_core::ocr::{
types::LiteLLMOcrRequest,
wire::{OcrWireRequest, decode_request},
};
use litellm_http::Client;
use litellm_llms::base_llm::ocr::{
error::Error,
handler::OcrClient,
@ -37,11 +38,7 @@ fn object(value: Value) -> Map<String, Value> {
}
fn ocr_client() -> OcrClient {
let document_http = reqwest::Client::builder()
.redirect(reqwest::redirect::Policy::none())
.build()
.expect("test document client builds");
OcrClient::for_test(reqwest::Client::new(), document_http)
OcrClient::for_test(Client::plain_for_test(), Client::no_redirect_for_test())
}
async fn perform(request: LiteLLMOcrRequest) -> Result<LiteLLMOcrResponse, Error> {

View file

@ -1,10 +1,6 @@
use std::sync::Arc;
use litellm_auth_gcp::VertexAuth;
use litellm_http::{
HttpClientPool, HttpSettings, Resolution,
media::{PublicDnsResolver, UrlPolicy},
};
use litellm_http::{HttpSettings, Resolution, media::UrlPolicy};
use litellm_llms::{
base_llm::ocr::{
settings::OcrSettings,
@ -184,6 +180,7 @@ async fn missing_credentials_come_from_the_injected_secret_source(
);
}
#[rstest]
#[tokio::test]
async fn the_client_uses_the_injected_http_pool_configuration() {
let upstream = upstream([pages_response()]).await;
@ -191,15 +188,18 @@ async fn the_client_uses_the_injected_http_pool_configuration() {
user_agent: Some("host-owned/1".into()),
..HttpSettings::default()
};
let client = OcrClient::new(
&HttpClientPool::new(Arc::new(PublicDnsResolver)),
&Resolution::from(&settings).config,
UrlPolicy::default(),
VertexAuth::default(),
OcrSettings::default(),
Arc::new(litellm_secrets::source::EnvironmentSecrets::default()),
)
.unwrap();
let client = resources()
.ocr_client(
&Resolution::from(&settings).config,
UrlPolicy::default(),
OcrSettings::default(),
Arc::new(
litellm_secrets::source::EnvironmentSecrets::python_compatible(
litellm_http::Client::plain_for_test(),
),
),
)
.unwrap();
litellm_core::ocr::client::perform(
&client,

View file

@ -0,0 +1,157 @@
mod support;
use std::sync::{
Arc,
atomic::{AtomicUsize, Ordering},
};
use litellm_auth::AuthServices;
use litellm_auth_gcp::{
CredentialSource, VertexAuth, VertexAuthFuture, VertexProviderLoader, VertexTokenSource,
};
use litellm_core::{
ocr::{
client::perform,
wire::{OcrWireRequest, decode_request},
},
resources::CoreResources,
};
use litellm_http::{HttpSettings, Resolution};
use litellm_llms::base_llm::ocr::settings::OcrSettings;
use rstest::{fixture, rstest};
use serde_json::json;
use support::{ReceivedRequest, RecordingSecrets, http_pool, json_response, upstream};
struct TokenSource(String);
impl VertexTokenSource for TokenSource {
fn project_id(&self) -> VertexAuthFuture<'_, String> {
Box::pin(async { Ok(self.0.clone()) })
}
fn token(&self) -> VertexAuthFuture<'_, String> {
Box::pin(async { Ok(self.0.clone()) })
}
}
#[derive(Default)]
struct Loader(AtomicUsize);
impl VertexProviderLoader for Loader {
fn load(&self, source: CredentialSource) -> VertexAuthFuture<'_, Arc<dyn VertexTokenSource>> {
Box::pin(async move {
self.0.fetch_add(1, Ordering::SeqCst);
let identity = match source {
CredentialSource::Trusted(secret) => secret.expose().to_string(),
other => panic!("unexpected credential source: {other:?}"),
};
Ok(Arc::new(TokenSource(identity)) as Arc<dyn VertexTokenSource>)
})
}
}
#[fixture]
fn loader() -> Arc<Loader> {
Arc::new(Loader::default())
}
#[fixture]
fn resources(loader: Arc<Loader>) -> CoreResources {
CoreResources {
auth: Arc::new(AuthServices {
gcp: VertexAuth::new(loader),
..AuthServices::default()
}),
pool: Arc::new(http_pool()),
}
}
#[rstest]
#[case::shared_identity(false, "first-identity", 1)]
#[case::different_identity(false, "second-identity", 2)]
#[case::independent_resources(true, "first-identity", 2)]
#[tokio::test]
async fn auth_survives_per_call_clients_without_freezing_settings_or_secrets(
loader: Arc<Loader>,
#[with(loader.clone())] resources: CoreResources,
#[case] independent: bool,
#[case] second_identity: &str,
#[case] expected_loads: usize,
) {
let response = json_response(json!({"pages": [{"index": 0, "markdown": "hello"}]}));
let upstream = upstream([response.clone(), response]).await;
let second_resources = if independent {
CoreResources {
auth: Arc::new(AuthServices {
gcp: VertexAuth::new(loader.clone()),
..AuthServices::default()
}),
..resources.clone()
}
} else {
resources.clone()
};
for (owner, identity, agent, location) in [
(&resources, "first-identity", "first-agent", "us-central1"),
(
&second_resources,
second_identity,
"second-agent",
"europe-west4",
),
] {
let http = Resolution::from(&HttpSettings {
user_agent: Some(agent.into()),
..HttpSettings::default()
})
.config;
let client = owner
.ocr_client(
&http,
Default::default(),
OcrSettings {
vertex_location: Some(location.into()),
..OcrSettings::default()
},
Arc::new(RecordingSecrets::new([("VERTEXAI_CREDENTIALS", identity)])),
)
.unwrap();
let request = decode_request(OcrWireRequest {
model: "vertex_ai/mistral-ocr-maas".into(),
document: json!({"type": "document_url", "document_url": "data:application/pdf;base64,YWJj"}),
api_key: None,
api_base: Some(upstream.uri()),
custom_llm_provider: None,
extra_headers: None,
optional_params: Default::default(),
input_sources: Default::default(),
timeout_seconds: Some(5.0),
}).unwrap();
let result = perform(&client, request).await.unwrap();
assert!(!result.pages.is_empty());
}
let requests = upstream.received_requests().await.unwrap();
assert_eq!(requests.len(), 2);
for (request, identity, agent, location) in [
(&requests[0], "first-identity", "first-agent", "us-central1"),
(
&requests[1],
second_identity,
"second-agent",
"europe-west4",
),
] {
assert_eq!(
request.header("authorization"),
Some(format!("Bearer {identity}").as_str())
);
assert_eq!(request.header("user-agent"), Some(agent));
assert!(
request
.url
.path()
.contains(&format!("/projects/{identity}/locations/{location}/"))
);
}
assert_eq!(loader.0.load(Ordering::SeqCst), expected_loads);
}

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