mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
Merge remote-tracking branch 'origin/main' into litellm_bedrock_grok_chat_completions
Some checks failed
LiteLLM Rust / rust-lint (push) Has been cancelled
LiteLLM Rust / rust-test (push) Has been cancelled
LiteLLM Rust / rust-wheel (push) Has been cancelled
Terraform Provider / gofmt, vet, build, test (push) Has been cancelled
Terraform Provider / Provider endpoints vs proxy OpenAPI schema (push) Has been cancelled
Some checks failed
LiteLLM Rust / rust-lint (push) Has been cancelled
LiteLLM Rust / rust-test (push) Has been cancelled
LiteLLM Rust / rust-wheel (push) Has been cancelled
Terraform Provider / gofmt, vet, build, test (push) Has been cancelled
Terraform Provider / Provider endpoints vs proxy OpenAPI schema (push) Has been cancelled
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> # Conflicts: # litellm/utils.py
This commit is contained in:
commit
ef03d19aae
876 changed files with 57016 additions and 11653 deletions
|
|
@ -17,3 +17,6 @@ rustflags = ["-C", "link-arg=-undefined", "-C", "link-arg=dynamic_lookup"]
|
|||
|
||||
[target.aarch64-apple-darwin]
|
||||
rustflags = ["-C", "link-arg=-undefined", "-C", "link-arg=dynamic_lookup"]
|
||||
|
||||
[env]
|
||||
SQLX_OFFLINE = "true"
|
||||
|
|
|
|||
|
|
@ -141,7 +141,7 @@ commands:
|
|||
node --version
|
||||
npm --version
|
||||
install_rust:
|
||||
description: "Install pinned rustup (1.28.2) and Rust toolchain (1.98.0) with checksum verification. Adds ~/.cargo/bin to PATH. Run this before any `uv sync` or `uv build` of the workspace: the root package builds litellm-rust through maturin, and on an image without cargo maturin fetches an unpinned rustup and a floating toolchain by itself."
|
||||
description: "Install pinned rustup (1.28.2) and Rust toolchain (1.98.0) with checksum verification. Adds ~/.cargo/bin to PATH. Run this before any `uv sync` or `uv build` of the workspace: the root package builds litellm-rust through maturin, and on an image without cargo maturin fetches an unpinned rustup and a floating toolchain by itself. Also restores the dev-profile cargo cache that save_cargo_target writes on main, minus the workspace crates' fingerprints so those always rebuild from the checked-out source."
|
||||
steps:
|
||||
- run:
|
||||
name: Install Rust (rustup 1.28.2, toolchain 1.98.0)
|
||||
|
|
@ -167,9 +167,29 @@ commands:
|
|||
/tmp/rustup-init -y --no-modify-path --profile minimal --default-toolchain 1.98.0
|
||||
rm -f /tmp/rustup-init
|
||||
echo 'export PATH="$HOME/.cargo/bin:$PATH"' >> "$BASH_ENV"
|
||||
echo 'export CARGO_INCREMENTAL=0' >> "$BASH_ENV"
|
||||
export PATH="$HOME/.cargo/bin:$PATH"
|
||||
rustc --version
|
||||
cargo --version
|
||||
{ rustc -vV; cc --version; cat /etc/os-release; } > /tmp/cargo-build-env
|
||||
- restore_cache:
|
||||
keys:
|
||||
- v1-cargo-dev-{{ checksum "/tmp/cargo-build-env" }}-{{ checksum "litellm-rust/Cargo.lock" }}
|
||||
- v1-cargo-dev-{{ checksum "/tmp/cargo-build-env" }}-
|
||||
- run:
|
||||
name: Force a rebuild of the workspace crates restored from the cargo cache
|
||||
command: rm -rf litellm-rust/target/debug/.fingerprint/litellm-*
|
||||
save_cargo_target:
|
||||
steps:
|
||||
- when:
|
||||
condition:
|
||||
equal: [main, << pipeline.git.branch >>]
|
||||
steps:
|
||||
- save_cache:
|
||||
key: v1-cargo-dev-{{ checksum "/tmp/cargo-build-env" }}-{{ checksum "litellm-rust/Cargo.lock" }}
|
||||
paths:
|
||||
- ~/.cargo/registry
|
||||
- ~/project/litellm-rust/target/debug
|
||||
start_postgres:
|
||||
description: "Start a postgres-db container on port 5432 and wait until it accepts connections."
|
||||
parameters:
|
||||
|
|
@ -281,51 +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
|
||||
name: Install Rust and uv
|
||||
no_output_timeout: 30m
|
||||
environment:
|
||||
UV_HTTP_TIMEOUT: "300"
|
||||
command: |
|
||||
$rustupInit = Join-Path $env:TEMP "rustup-init.exe"
|
||||
$rustupVersion = "1.28.2"
|
||||
|
|
@ -365,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
|
||||
|
|
@ -380,17 +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:
|
||||
|
|
@ -418,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:
|
||||
|
|
@ -446,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: |
|
||||
|
|
@ -3120,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:
|
||||
|
|
@ -3249,6 +3338,7 @@ jobs:
|
|||
image: ubuntu-2204:2024.04.1
|
||||
resource_class: large
|
||||
working_directory: ~/project
|
||||
parallelism: 4
|
||||
steps:
|
||||
- setup_litellm_test_deps
|
||||
- run:
|
||||
|
|
@ -3258,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
|
||||
|
|
@ -3328,7 +3419,11 @@ workflows:
|
|||
name: integration-<< matrix.suite >>
|
||||
matrix:
|
||||
parameters:
|
||||
suite: [management, accounting, database, providers, extensions, mcp, sdk, cost, browser]
|
||||
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:
|
||||
|
|
@ -3343,6 +3438,7 @@ workflows:
|
|||
equal: ["", << pipeline.parameters.routing_parity_base >>]
|
||||
jobs:
|
||||
- using_litellm_on_windows
|
||||
- windows_release_wheel
|
||||
- unit
|
||||
- provider_replay_harness
|
||||
- base_sdk_install
|
||||
|
|
|
|||
|
|
@ -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
|
||||
;;
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -46,6 +46,7 @@ legacy_paths() {
|
|||
echo tests/unit/google_genai
|
||||
echo tests/unit/router_strategy
|
||||
echo tests/unit/router_utils
|
||||
echo tests/unit/proxy/common_utils/test_cache_aware_routing.py
|
||||
echo tests/unit/enterprise/enterprise_callbacks/send_emails
|
||||
echo tests/unit/enterprise/proxy/test_afile_retrieve_returns_unified_id.py
|
||||
echo tests/unit/enterprise/proxy/test_batch_retrieve_input_file_id.py
|
||||
|
|
|
|||
|
|
@ -42,7 +42,7 @@ case "$subject" in
|
|||
;;
|
||||
esac
|
||||
|
||||
ALLOWED_TYPES="feat|fix|docs|style|refactor|perf|test|build|ci|chore|revert"
|
||||
ALLOWED_TYPES="feat|fix|docs|style|refactor|perf|test|build|ci|chore|revert|security"
|
||||
# Description must not start with an uppercase letter — kept in sync with the
|
||||
# subjectPattern in .github/workflows/conventional-commits.yml so the local
|
||||
# hook is the strictly tighter of the two gates. (Without this guard, a commit
|
||||
|
|
@ -61,7 +61,7 @@ cat >&2 <<EOF
|
|||
Expected: <type>(<scope>)!: <description>
|
||||
(description must start with a lowercase letter)
|
||||
|
||||
Allowed types: feat, fix, docs, style, refactor, perf, test, build, ci, chore, revert
|
||||
Allowed types: feat, fix, docs, style, refactor, perf, test, build, ci, chore, revert, security
|
||||
Examples:
|
||||
feat(router): add weighted round-robin strategy
|
||||
fix(bedrock): decouple STS region from aws_region_name
|
||||
|
|
|
|||
7
.github/ci-coverage-allowlist.yml
vendored
7
.github/ci-coverage-allowlist.yml
vendored
|
|
@ -111,3 +111,10 @@ dockerfiles:
|
|||
An example image under cookbook/ that is documentation rather than a shipped artifact
|
||||
paths:
|
||||
- cookbook/litellm-ollama-docker-image/Dockerfile
|
||||
- reason: >-
|
||||
The Rust gateway image compiles the whole workspace in release mode, which is too slow for
|
||||
a per-pull-request job while the gateway binary is still being assembled; the Rust lint,
|
||||
clippy, and compile jobs already cover the code it packages. Revisit when the gateway is
|
||||
published
|
||||
paths:
|
||||
- litellm-rust/crates/gateway/Dockerfile
|
||||
|
|
|
|||
2
.github/pull_request_template.md
vendored
2
.github/pull_request_template.md
vendored
|
|
@ -54,7 +54,7 @@ After: the same request comes back with real token counts, so the dashboard show
|
|||
|
||||
## Affected release
|
||||
|
||||
<!-- Only for a fix to a regression in a released or rc version (perf, memory, crash, or behavior): name the version it regressed in, e.g. "regression in v1.100.0" or "since v1.101.0-rc.1", and add the `backport-stable` label so the fix is cherry-picked onto the rc line before the stable is tagged. Drop the section otherwise -->
|
||||
<!-- Only for a fix to a regression in a released or rc version (perf, memory, crash, or behavior): name the version it regressed in, e.g. "regression in v1.100.0" or "since v1.101.0-rc.1". Add the `backport-stable` label only when the regression is a P0, meaning its Linear ticket is Urgent (a security hole however narrow, data loss, or a crash or outage for every user on that version), because every labeled PR must be cherry-picked onto the baking rc line before the stable can be tagged; every other regression fix ships in the next rc unlabeled. Drop the section otherwise -->
|
||||
|
||||
## Linear ticket
|
||||
|
||||
|
|
|
|||
2
.github/scripts/assert_ci_coverage.py
vendored
2
.github/scripts/assert_ci_coverage.py
vendored
|
|
@ -516,7 +516,7 @@ def _integration_ownership(repo_root: pathlib.Path = REPO_ROOT) -> tuple[frozens
|
|||
str(path.relative_to(repo_root))
|
||||
for folders in groups.values()
|
||||
for folder in folders
|
||||
for path in (integration_root / folder).glob("test_*.py")
|
||||
for path in (integration_root / folder).rglob("test_*.py")
|
||||
)
|
||||
browser_manifest: Final = repo_root / "tests/e2e/ui/tests/integrationCritical/expected.json"
|
||||
browser_nodes: Final = json.loads(browser_manifest.read_text()) if browser_manifest.exists() else ()
|
||||
|
|
|
|||
2
.github/scripts/verify_linux_native_wheel.py
vendored
2
.github/scripts/verify_linux_native_wheel.py
vendored
|
|
@ -214,7 +214,7 @@ def main(
|
|||
native_module: Final = load_native_module(native_path)
|
||||
native_module_loads: Final = native_module is not None
|
||||
panic_test_hook_absent: Final = native_module is not None and not hasattr(native_module, "_panic_for_test")
|
||||
native_size_limit: Final = 40_000_000
|
||||
native_size_limit: Final = 45_000_000
|
||||
native_size_within_limit: Final = native_member.file_size <= native_size_limit
|
||||
validations: Final = (
|
||||
(f"Python tag is {EXPECTED_PYTHON_TAG}", python_tag == EXPECTED_PYTHON_TAG),
|
||||
|
|
|
|||
1
.github/workflows/conventional-commits.yml
vendored
1
.github/workflows/conventional-commits.yml
vendored
|
|
@ -41,6 +41,7 @@ jobs:
|
|||
ci
|
||||
chore
|
||||
revert
|
||||
security
|
||||
requireScope: false
|
||||
subjectPattern: ^(?![A-Z]).+$
|
||||
subjectPatternError: |
|
||||
|
|
|
|||
3
.github/workflows/test-code-quality.yml
vendored
3
.github/workflows/test-code-quality.yml
vendored
|
|
@ -146,6 +146,9 @@ jobs:
|
|||
- name: check_migrations_no_data_rewrites
|
||||
run: uv run --no-sync python ./tests/code_coverage_tests/check_migrations_no_data_rewrites.py
|
||||
|
||||
- name: check_unbounded_in_lists (fails on findings not in the baseline)
|
||||
run: uv run --no-sync python ./tests/code_coverage_tests/check_unbounded_in_lists.py
|
||||
|
||||
- name: memory_test
|
||||
run: uv run --no-sync python ./tests/code_coverage_tests/memory_test.py
|
||||
|
||||
|
|
|
|||
|
|
@ -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/` |
|
||||
|
|
|
|||
6
Makefile
6
Makefile
|
|
@ -4,7 +4,7 @@
|
|||
.PHONY: help test test-unit test-unit-llms test-unit-proxy-guardrails test-unit-proxy-core test-unit-proxy-misc \
|
||||
test-unit-integrations test-unit-core-utils test-unit-other test-unit-root \
|
||||
test-proxy-unit-a test-proxy-unit-b test-integration test-unit-helm \
|
||||
test-rust-extension \
|
||||
test-rust-extension rust-sqlx-prepare \
|
||||
info lint lint-inner lint-dev lint-checks format \
|
||||
lint-basedpyright lint-e2e-basedpyright lint-basedpyright-budget-update lint-type-discipline lint-type-discipline-budget-update \
|
||||
lint-ruff-budget lint-ruff-budget-update lint-budget-update lint-gate \
|
||||
|
|
@ -56,6 +56,7 @@ help:
|
|||
@echo " make test-integration - Run integration tests"
|
||||
@echo " make test-unit-helm - Run helm unit tests"
|
||||
@echo " make test-rust-extension - Build the Rust extension and run its public Python tests"
|
||||
@echo " make rust-sqlx-prepare - Refresh litellm-rust/crates/db/.sqlx against a migrated Postgres container"
|
||||
@echo ""
|
||||
@echo "Heavy targets (check, lint) queue for LITELLM_GATE_SLOTS machine-wide"
|
||||
@echo "slots (default 2; 0 disables) so parallel sessions don't thrash one machine."
|
||||
|
|
@ -306,6 +307,9 @@ test-rust-extension:
|
|||
LITELLM_RUST=1 LITELLM_LOCAL_MODEL_COST_MAP=True \
|
||||
"$$temporary/venv/bin/python" -I -m pytest --import-mode=importlib -m requires_rust_extension tests/test_litellm_rust
|
||||
|
||||
rust-sqlx-prepare:
|
||||
cd litellm-rust && cargo run -p litellm-db-testing --bin sqlx-prepare
|
||||
|
||||
test: install-test-deps
|
||||
$(UV_RUN) pytest tests/
|
||||
|
||||
|
|
|
|||
|
|
@ -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) | ✅ | ✅ | ✅ | | | | | | | |
|
||||
|
|
|
|||
37
cookbook/litellm_proxy_server/mcp/README.md
Normal file
37
cookbook/litellm_proxy_server/mcp/README.md
Normal 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
|
||||
22
litellm-rust/.agents/skills/rust-tracing/SKILL.md
Normal file
22
litellm-rust/.agents/skills/rust-tracing/SKILL.md
Normal file
|
|
@ -0,0 +1,22 @@
|
|||
---
|
||||
name: rust-tracing
|
||||
description: Add or change Rust diagnostic tracing in litellm-rust, including route spans, subscriber layers, and Python logger delivery
|
||||
---
|
||||
|
||||
# Rust tracing
|
||||
|
||||
Use upstream `tracing` throughout Rust, including `#[tracing::instrument]`, events, and span propagation. Centralize collection and delivery infrastructure in `crates/tracing`. Direct upstream imports still reach our configured subscriber; re-exporting macros does not control delivery. Do not introduce Rust `log` or `pyo3-log` for this path
|
||||
|
||||
`litellm-tracing` owns shared subscriber layers, span field collection, and diagnostic processing. Keep adapters composable as `tracing_subscriber::Layer`s, with `Logger` providing host setup. Runtime-specific delivery belongs in the host bridge. The Python bridge delivers directly to the existing Python SDK logger, preserving its handlers, filtering, redaction, and request correlation. Keep Python dependencies out of `crates/tracing`
|
||||
|
||||
Hosts configure subscribers. Keep Python execution scoped to its captured dispatch rather than installing a process-wide subscriber. Propagate both span context and dispatch across spawned work and returned streams
|
||||
|
||||
In core, instrument execution shared by native calls and hosted machines. Use consistent route, model, provider, streaming, and outcome fields. Put status recording at shared provider boundaries instead of scattering basic logging through handlers. Keep upstream HTTP status separate from route success
|
||||
|
||||
Use `skip_all` and explicitly selected fields. Basic tracing excludes bodies, credentials, headers, and raw error strings. Avoid automatic `ret` or `err` capture of sensitive values. Keep payload diagnostics separate and subject to existing redaction
|
||||
|
||||
A returned stream retains its route span until exhaustion, error, or drop, with exactly one terminal outcome. Builder construction does not start a trace. Never hold a span entry guard across an await. Diagnostic tracing remains separate from lifecycle callbacks and `CustomLogger` dispatch
|
||||
|
||||
Use `litellm_tracing::sink_layer` to compose a sink with other subscriber layers. It inherits span fields into events and emits span-close summaries with elapsed time. Test observable records, concurrent isolation, dynamic filtering, sensitive-field exclusion, and stream cancellation when changing this behavior
|
||||
|
||||
Consult the [tracing API](https://docs.rs/tracing/latest/tracing/) and [subscriber layers](https://docs.rs/tracing-subscriber/latest/tracing_subscriber/layer/index.html) for implementation details
|
||||
|
|
@ -1,5 +1,7 @@
|
|||
# Rust workspace rules
|
||||
|
||||
For diagnostic tracing changes, follow [.agents/skills/rust-tracing/SKILL.md](.agents/skills/rust-tracing/SKILL.md)
|
||||
|
||||
## Test placement
|
||||
|
||||
- Never create a `tests.rs` (or `test.rs`) file under `src/`, and never `#[path = "tests.rs"] mod tests;`
|
||||
|
|
@ -16,7 +18,9 @@ Use [`#[rstest]`](https://docs.rs/rstest/latest/rstest/attr.rstest.html) for new
|
|||
## Error definitions
|
||||
|
||||
- A crate's errors live in `src/error.rs`, defined with `thiserror`, and re-exported from `lib.rs`
|
||||
- Put message templates in the variant's `#[error(...)]` declaration. Callers pass only the small typed arguments needed to fill them, never `Error::Variant(format!(...))` or a preformatted message. Keep the smallest set of neutral variants that callers need to distinguish; different wording or providers do not justify new variants
|
||||
- Default to one top-level `Error` enum per crate, with one variant per failure mode and a `#[error(...)]` message on each. A failure mode is something a caller handles differently (phase, status code, retry, a message Python parity pins exactly); failures no caller tells apart share one variant and differ only in its message
|
||||
- Keep shared error enums minimal and provider-neutral. Provider names, credential types, configuration fields, and setup guidance belong in caller-supplied data, not dedicated variants or hardcoded shared messages. Reuse a variant for the same failure mode across providers, such as `MissingApiBase { provider: "Azure", guidance: "..." }`. An exact parity message does not justify a provider-specific variant when caller-supplied context can preserve it
|
||||
- Wrap a lower-level error as a variant with `#[from]` or `#[source]` instead of flattening it to a string
|
||||
- Exception: split into separate types when different functions fail in disjoint ways, especially when different callers see them. A shared enum would force every caller to match variants its function can never return
|
||||
- Name a split type after what went wrong (a unit struct is fine for a single failure mode), not after the function that returns it
|
||||
|
|
|
|||
1417
litellm-rust/Cargo.lock
generated
1417
litellm-rust/Cargo.lock
generated
File diff suppressed because it is too large
Load diff
|
|
@ -13,11 +13,16 @@ litellm-config = { path = "crates/config" }
|
|||
litellm-router = { path = "crates/router" }
|
||||
litellm-tracing = { path = "crates/tracing" }
|
||||
litellm-core = { path = "crates/core" }
|
||||
litellm-gateway-mcp = { path = "crates/gateway-mcp" }
|
||||
litellm-gateway = { path = "crates/gateway" }
|
||||
litellm-gateway-inference = { path = "crates/gateway-inference" }
|
||||
litellm-gateway-auth = { path = "crates/gateway-auth" }
|
||||
litellm-gateway-management = { path = "crates/gateway-management" }
|
||||
litellm-gateway-ui = { path = "crates/gateway-ui" }
|
||||
litellm-coroutine = { path = "crates/coroutine" }
|
||||
litellm-host = { path = "crates/host" }
|
||||
litellm-host-http = { path = "crates/host-http" }
|
||||
litellm-host-native = { path = "crates/host-native" }
|
||||
litellm-callbacks-legacy-python = { path = "crates/callbacks-legacy-python" }
|
||||
litellm-framing = { path = "crates/framer" }
|
||||
litellm-auth = { path = "crates/auth" }
|
||||
|
|
@ -36,6 +41,8 @@ litellm-http = { path = "crates/http" }
|
|||
litellm-llms = { path = "crates/llms" }
|
||||
litellm-types = { path = "crates/types" }
|
||||
litellm-core-utils = { path = "crates/core-utils" }
|
||||
litellm-db = { path = "crates/db" }
|
||||
litellm-db-testing = { path = "crates/db-testing" }
|
||||
litellm-cache = { path = "crates/cache" }
|
||||
litellm-cache-azure-blob = { path = "crates/cache-azure-blob" }
|
||||
litellm-cache-memory = { path = "crates/cache-memory" }
|
||||
|
|
@ -56,6 +63,8 @@ litellm-python-compat = { path = "crates/python-compat" }
|
|||
|
||||
tracing = "0.1"
|
||||
axum = { version = "0.8.9", default-features = false, features = ["http1", "tokio", "multipart"] }
|
||||
axum-login = "0.18.0"
|
||||
tower-sessions = { version = "0.14.0", features = ["memory-store"] }
|
||||
bytes = "1"
|
||||
http = "1"
|
||||
google-cloud-auth = { version = "1.16.0", default-features = false }
|
||||
|
|
@ -80,6 +89,7 @@ serde = { version = "1.0", features = ["derive"] }
|
|||
serde_json = { version = "1.0", features = ["float_roundtrip"] }
|
||||
serde_with = { version = "=3.16.1", default-features = false, features = ["std", "macros"] }
|
||||
sha2 = "0.10"
|
||||
sqlx = { version = "0.9.0", default-features = false, features = ["json", "macros", "postgres", "runtime-tokio", "chrono", "tls-rustls-ring-native-roots"] }
|
||||
subtle = "2"
|
||||
thiserror = "2.0"
|
||||
tokenizers = { version = "0.23.1", default-features = false, features = ["onig"] }
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
# The Tokio runtime is reached only through `host-python/src/execution.rs`, whose fork gate
|
||||
# The Tokio runtime is reached only through `host-python/src/runtime.rs`, whose fork gate
|
||||
# must see every entry. Going around it makes a fork-after-use hang instead of raising.
|
||||
disallowed-methods = [
|
||||
{ path = "pyo3_async_runtimes::tokio::get_runtime", reason = "use litellm_host_python::run_sync / run_sync_value" },
|
||||
|
|
@ -12,6 +12,13 @@ disallowed-methods = [
|
|||
{ path = "reqwest::ClientBuilder::danger_accept_invalid_certs", reason = "set HttpClientConfig::verify instead" },
|
||||
{ path = "reqwest::ClientBuilder::identity", reason = "set HttpClientConfig::client_certificate instead" },
|
||||
{ path = "reqwest::ClientBuilder::use_preconfigured_tls", reason = "HttpClientConfig owns the TLS configuration" },
|
||||
{ path = "sqlx::query", reason = "use sqlx::query! or query_file! so the SQL is checked against the migrated schema" },
|
||||
{ path = "sqlx::query_as", reason = "use sqlx::query_as! or query_file_as! so the SQL is checked against the migrated schema" },
|
||||
{ path = "sqlx::query_scalar", reason = "use sqlx::query_scalar! so the SQL is checked against the migrated schema" },
|
||||
{ path = "sqlx::query_with", reason = "use sqlx::query! or query_file! so the SQL is checked against the migrated schema" },
|
||||
{ path = "sqlx::query_as_with", reason = "use sqlx::query_as! or query_file_as! so the SQL is checked against the migrated schema" },
|
||||
{ path = "sqlx::query_scalar_with", reason = "use sqlx::query_scalar! so the SQL is checked against the migrated schema" },
|
||||
{ path = "sqlx::raw_sql", reason = "raw_sql is unchecked; use the checked query macros" },
|
||||
]
|
||||
|
||||
# Every outbound client comes from litellm_http::HttpClientPool so it honors the host's TLS,
|
||||
|
|
|
|||
|
|
@ -133,7 +133,7 @@ impl NativeAzureTokenAcquirer {
|
|||
let token = credential
|
||||
.get_token(&[scope.as_str()], None)
|
||||
.await
|
||||
.map_err(|error| Error::AzureTokenAcquisition(error.to_string()))?;
|
||||
.map_err(|error| Error::CredentialAcquisition(error.to_string().into()))?;
|
||||
let expires_on = u64::try_from(token.expires_on.unix_timestamp())
|
||||
.ok()
|
||||
.map(|seconds| UNIX_EPOCH + Duration::from_secs(seconds));
|
||||
|
|
@ -250,7 +250,12 @@ fn validate_authority(request: &NativeAzureRequest) -> Result<(), Error> {
|
|||
let Some(authority) = authority else {
|
||||
return Ok(());
|
||||
};
|
||||
let url = url::Url::parse(authority.value()).map_err(|_| Error::InvalidAzureAuthority)?;
|
||||
let url = url::Url::parse(authority.value()).map_err(|_| {
|
||||
Error::InvalidConfiguration(
|
||||
"Azure authority must be an HTTPS origin without credentials, query, or fragment"
|
||||
.into(),
|
||||
)
|
||||
})?;
|
||||
if url.scheme() != "https"
|
||||
|| url.host_str().is_none()
|
||||
|| !url.username().is_empty()
|
||||
|
|
@ -259,7 +264,10 @@ fn validate_authority(request: &NativeAzureRequest) -> Result<(), Error> {
|
|||
|| url.fragment().is_some()
|
||||
|| !matches!(url.path(), "" | "/")
|
||||
{
|
||||
return Err(Error::InvalidAzureAuthority);
|
||||
return Err(Error::InvalidConfiguration(
|
||||
"Azure authority must be an HTTPS origin without credentials, query, or fragment"
|
||||
.into(),
|
||||
));
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
|
@ -368,7 +376,9 @@ fn trusted_source(sources: &[InputSource]) -> InputSource {
|
|||
}
|
||||
|
||||
fn mixed_sources<T>() -> Result<T, Error> {
|
||||
Err(Error::MixedAzureCredentialSources)
|
||||
Err(Error::InvalidConfiguration(
|
||||
"request-controlled Azure auth inputs cannot be combined with host credentials".into(),
|
||||
))
|
||||
}
|
||||
|
||||
fn build_credential(
|
||||
|
|
@ -433,7 +443,12 @@ fn build_credential(
|
|||
NativeAzureRequest::DeveloperTools { .. } => DeveloperToolsCredential::new(None)
|
||||
.map(|credential| credential as Arc<dyn TokenCredential>),
|
||||
}
|
||||
.map_err(|error| Error::AzureCredentialInitialization(error.to_string()))
|
||||
.map_err(|error| {
|
||||
Error::InvalidConfiguration(litellm_auth_types::ErrorDetail::failed(
|
||||
"Azure credential initialization",
|
||||
error,
|
||||
))
|
||||
})
|
||||
}
|
||||
|
||||
fn client_options(
|
||||
|
|
@ -638,7 +653,7 @@ mod tests {
|
|||
assert_eq!(transport.requests.lock().unwrap().len(), 6);
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[rstest::rstest]
|
||||
fn request_authority_requires_request_owned_client_secret_identity() {
|
||||
let error = ValidatedAzureRequest::new(sourced_client_secret(
|
||||
InputSource::Deployment,
|
||||
|
|
@ -647,10 +662,13 @@ mod tests {
|
|||
))
|
||||
.unwrap_err();
|
||||
|
||||
assert!(matches!(
|
||||
assert_eq!(
|
||||
error,
|
||||
litellm_auth_types::Error::MixedAzureCredentialSources
|
||||
));
|
||||
litellm_auth_types::Error::InvalidConfiguration(
|
||||
"request-controlled Azure auth inputs cannot be combined with host credentials"
|
||||
.into()
|
||||
)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
|
@ -665,24 +683,24 @@ mod tests {
|
|||
assert_eq!(request.credential_source(), InputSource::Request);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn authority_is_restricted_to_an_https_origin() {
|
||||
for authority in [
|
||||
"http://login.example",
|
||||
"https://user@login.example",
|
||||
"https://login.example/tenant",
|
||||
"https://login.example?target=other",
|
||||
] {
|
||||
let error = ValidatedAzureRequest::new(sourced_client_secret(
|
||||
InputSource::Deployment,
|
||||
InputSource::Deployment,
|
||||
authority,
|
||||
))
|
||||
.unwrap_err();
|
||||
assert!(matches!(
|
||||
error,
|
||||
litellm_auth_types::Error::InvalidAzureAuthority
|
||||
));
|
||||
}
|
||||
#[rstest::rstest]
|
||||
#[case::http("http://login.example")]
|
||||
#[case::userinfo("https://user@login.example")]
|
||||
#[case::path("https://login.example/tenant")]
|
||||
#[case::query("https://login.example?target=other")]
|
||||
fn authority_is_restricted_to_an_https_origin(#[case] authority: &str) {
|
||||
let error = ValidatedAzureRequest::new(sourced_client_secret(
|
||||
InputSource::Deployment,
|
||||
InputSource::Deployment,
|
||||
authority,
|
||||
))
|
||||
.unwrap_err();
|
||||
assert_eq!(
|
||||
error,
|
||||
litellm_auth_types::Error::InvalidConfiguration(
|
||||
"Azure authority must be an HTTPS origin without credentials, query, or fragment"
|
||||
.into()
|
||||
)
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -91,7 +91,9 @@ impl AzureAuthService {
|
|||
AzureCredentialPlan::Caller(caller) => {
|
||||
let credential = caller.acquire().await?;
|
||||
if credential.secret().expose().is_empty() {
|
||||
return Err(Error::EmptyAzureToken);
|
||||
return Err(Error::EmptyCallerCredential(
|
||||
"Azure AD token provider returned an empty token",
|
||||
));
|
||||
}
|
||||
Ok(Some(Sourced::new(credential, InputSource::Deployment)))
|
||||
}
|
||||
|
|
@ -104,7 +106,11 @@ impl AzureAuthService {
|
|||
} => {
|
||||
let assertion = resolve_reference(inputs, env_lookup, reference.value())
|
||||
.await?
|
||||
.ok_or(Error::UnresolvedOidcReference)?;
|
||||
.ok_or_else(|| {
|
||||
Error::CredentialAcquisition(
|
||||
"Azure OIDC reference did not resolve to a value".into(),
|
||||
)
|
||||
})?;
|
||||
let request = ValidatedAzureRequest::new(NativeAzureRequest::ClientAssertion {
|
||||
tenant_id,
|
||||
client_id,
|
||||
|
|
@ -167,7 +173,7 @@ pub(crate) fn select_auth_plan(
|
|||
.map(|selector| Sourced::new(selector, value.source()))
|
||||
})
|
||||
.transpose()
|
||||
.map_err(|_| Error::InvalidAzureSelector)?;
|
||||
.map_err(|_| Error::InvalidConfiguration("invalid Azure credential selector".into()))?;
|
||||
let federated_token_file = configured_string(
|
||||
&inputs.federated_token_file,
|
||||
AZURE_FEDERATED_TOKEN_FILE_ENV,
|
||||
|
|
@ -257,7 +263,9 @@ fn select_native_plan(
|
|||
let selection_source = selected.source();
|
||||
|
||||
match selected.into_value() {
|
||||
AzureCredentialType::ClientSecretCredential => Err(Error::MissingClientSecretFields),
|
||||
AzureCredentialType::ClientSecretCredential => Err(Error::InvalidConfiguration(
|
||||
"ClientSecretCredential requires tenant_id, client_id, and client_secret".into(),
|
||||
)),
|
||||
AzureCredentialType::WorkloadIdentityCredential => {
|
||||
Ok(AzureCredentialPlan::Native(ValidatedAzureRequest::new(
|
||||
workload_request(tenant_id, client_id, federated_token_file, scope, authority)?,
|
||||
|
|
@ -341,9 +349,17 @@ fn workload_request(
|
|||
authority: Option<Sourced<String>>,
|
||||
) -> Result<NativeAzureRequest, Error> {
|
||||
Ok(NativeAzureRequest::WorkloadIdentity {
|
||||
tenant_id: tenant_id.ok_or(Error::MissingWorkloadTenant)?,
|
||||
client_id: client_id.ok_or(Error::MissingWorkloadClient)?,
|
||||
token_file_path: token_file_path.ok_or(Error::MissingWorkloadTokenFile)?,
|
||||
tenant_id: tenant_id.ok_or_else(|| {
|
||||
Error::InvalidConfiguration("WorkloadIdentityCredential requires tenant_id".into())
|
||||
})?,
|
||||
client_id: client_id.ok_or_else(|| {
|
||||
Error::InvalidConfiguration("WorkloadIdentityCredential requires client_id".into())
|
||||
})?,
|
||||
token_file_path: token_file_path.ok_or_else(|| {
|
||||
Error::InvalidConfiguration(
|
||||
"WorkloadIdentityCredential requires azure_federated_token_file".into(),
|
||||
)
|
||||
})?,
|
||||
scope,
|
||||
authority,
|
||||
})
|
||||
|
|
@ -394,10 +410,11 @@ async fn resolve_reference(
|
|||
.map_or(CredentialLookup::Missing, CredentialLookup::Found),
|
||||
CredentialRef::None => return Ok(None),
|
||||
CredentialRef::File(_) | CredentialRef::Request(_) | CredentialRef::Host(_) => {
|
||||
let resolver = inputs
|
||||
.credential_resolver
|
||||
.as_ref()
|
||||
.ok_or(Error::MissingHostResolver)?;
|
||||
let resolver = inputs.credential_resolver.as_ref().ok_or_else(|| {
|
||||
Error::InvalidConfiguration(
|
||||
"credential reference requires a host credential resolver".into(),
|
||||
)
|
||||
})?;
|
||||
resolver.resolve(reference).await?
|
||||
}
|
||||
};
|
||||
|
|
@ -415,7 +432,9 @@ fn oidc_reference(
|
|||
};
|
||||
let value = token.value().expose();
|
||||
if token.source() == InputSource::Request && value.starts_with("oidc/") {
|
||||
return Err(Error::RequestAzureCredentialReference);
|
||||
return Err(Error::InvalidConfiguration(
|
||||
"request-controlled Azure credential references are not allowed".into(),
|
||||
));
|
||||
}
|
||||
if let Some(name) = value.strip_prefix("oidc/env/") {
|
||||
return non_empty_reference(name, "OIDC environment reference")
|
||||
|
|
@ -437,14 +456,20 @@ fn oidc_reference(
|
|||
)));
|
||||
}
|
||||
if value.starts_with("oidc/") {
|
||||
return Err(Error::UnsupportedOidcReference);
|
||||
return Err(Error::InvalidConfiguration(
|
||||
"unsupported OIDC reference".into(),
|
||||
));
|
||||
}
|
||||
Ok(None)
|
||||
}
|
||||
|
||||
fn non_empty_reference(value: &str, kind: &str) -> Result<String, Error> {
|
||||
if value.is_empty() {
|
||||
return Err(Error::EmptyReference(kind.to_string()));
|
||||
return Err(Error::InvalidConfiguration(
|
||||
litellm_auth_types::ErrorDetail::Empty {
|
||||
subject: kind.into(),
|
||||
},
|
||||
));
|
||||
}
|
||||
Ok(value.to_string())
|
||||
}
|
||||
|
|
@ -493,7 +518,7 @@ mod tests {
|
|||
expires_on: None,
|
||||
})
|
||||
} else {
|
||||
Err(Error::AzureTokenAcquisition(format!("{kind} failed")))
|
||||
Err(Error::CredentialAcquisition(kind.into()))
|
||||
}
|
||||
})
|
||||
}
|
||||
|
|
@ -602,7 +627,7 @@ mod tests {
|
|||
assert!(error.to_string().contains("unsupported OIDC reference"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[rstest::rstest]
|
||||
fn request_oidc_reference_is_rejected_before_lookup() {
|
||||
let params = json!({
|
||||
"azure_ad_token": "oidc/env/ASSERTION",
|
||||
|
|
@ -624,7 +649,12 @@ mod tests {
|
|||
})
|
||||
.unwrap_err();
|
||||
|
||||
assert!(matches!(error, Error::RequestAzureCredentialReference));
|
||||
assert_eq!(
|
||||
error,
|
||||
Error::InvalidConfiguration(
|
||||
"request-controlled Azure credential references are not allowed".into()
|
||||
)
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
|
|
@ -723,6 +753,7 @@ mod tests {
|
|||
assert_eq!(credential.value().secret().expose(), "caller-token");
|
||||
}
|
||||
|
||||
#[rstest::rstest]
|
||||
#[tokio::test]
|
||||
async fn empty_caller_token_is_rejected() {
|
||||
let error = AzureAuthService::default()
|
||||
|
|
@ -730,6 +761,9 @@ mod tests {
|
|||
.await
|
||||
.unwrap_err();
|
||||
|
||||
assert!(matches!(error, Error::EmptyAzureToken));
|
||||
assert_eq!(
|
||||
error,
|
||||
Error::EmptyCallerCredential("Azure AD token provider returned an empty token")
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -117,7 +117,12 @@ fn string_config(
|
|||
None => Ok(ConfigValue::Absent),
|
||||
Some(Value::Null) => Ok(ConfigValue::ExplicitNone(source)),
|
||||
Some(Value::String(value)) => Ok(ConfigValue::Value(Sourced::new(value.clone(), source))),
|
||||
Some(_) => Err(Error::InvalidFieldType(name.to_string())),
|
||||
Some(_) => Err(Error::InvalidConfiguration(
|
||||
litellm_auth_types::ErrorDetail::InvalidType {
|
||||
field: name.into(),
|
||||
expected: "a string or null",
|
||||
},
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -19,3 +19,6 @@ tokio.workspace = true
|
|||
gcp_auth = "0.12.7"
|
||||
google-cloud-auth = { workspace = true, optional = true }
|
||||
http = { workspace = true, optional = true }
|
||||
|
||||
[dev-dependencies]
|
||||
rstest.workspace = true
|
||||
|
|
|
|||
|
|
@ -299,7 +299,7 @@ fn validate_request_credentials(configured: &str) -> Result<&str, Error> {
|
|||
.map(str::to_string)
|
||||
});
|
||||
if token_uri.as_deref() != Some(GOOGLE_OAUTH_TOKEN_ENDPOINT) {
|
||||
return Err(Error::RequestVertexTokenEndpoint);
|
||||
return Err(Error::InvalidConfiguration("request-controlled Vertex credentials must use the canonical Google OAuth token endpoint".into()));
|
||||
}
|
||||
Ok(configured)
|
||||
}
|
||||
|
|
@ -376,10 +376,20 @@ fn optional_credentials(
|
|||
.map(SecretValue::new)
|
||||
.map(|value| Sourced::new(value, source))
|
||||
.map(Some)
|
||||
.map_err(|error| Error::InvalidFieldType(format!("{}: {error}", names[0])));
|
||||
.map_err(|error| {
|
||||
Error::InvalidConfiguration(litellm_auth_types::ErrorDetail::failed(
|
||||
"credential serialization",
|
||||
error,
|
||||
))
|
||||
});
|
||||
}
|
||||
Some(_) => {
|
||||
return Err(Error::InvalidFieldType(names[0].to_string()));
|
||||
return Err(Error::InvalidConfiguration(
|
||||
litellm_auth_types::ErrorDetail::InvalidType {
|
||||
field: names[0].into(),
|
||||
expected: "a string or null",
|
||||
},
|
||||
));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -397,7 +407,12 @@ fn optional_string(params: &Map<String, Value>, names: &[&str]) -> Result<Option
|
|||
Some(Value::String(value)) if value.trim().is_empty() => continue,
|
||||
Some(Value::String(value)) => return Ok(Some(value.clone())),
|
||||
Some(_) => {
|
||||
return Err(Error::InvalidFieldType(names[0].to_string()));
|
||||
return Err(Error::InvalidConfiguration(
|
||||
litellm_auth_types::ErrorDetail::InvalidType {
|
||||
field: names[0].into(),
|
||||
expected: "a string or null",
|
||||
},
|
||||
));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -411,7 +426,10 @@ fn non_empty_env(env_lookup: &dyn Fn(&str) -> Option<String>, name: &str) -> Opt
|
|||
}
|
||||
|
||||
fn auth_acquisition_error(error: gcp_auth::Error) -> Error {
|
||||
Error::VertexTokenAcquisition(error.to_string())
|
||||
Error::CredentialAcquisition(litellm_auth_types::ErrorDetail::failed(
|
||||
"Vertex AI credentials",
|
||||
error,
|
||||
))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
|
|
@ -612,20 +630,15 @@ mod tests {
|
|||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn request_credentials_require_canonical_token_endpoint() {
|
||||
assert!(
|
||||
validate_request_credentials(r#"{"token_uri":"https://oauth2.googleapis.com/token"}"#)
|
||||
.is_ok()
|
||||
);
|
||||
assert!(matches!(
|
||||
validate_request_credentials(r#"{"token_uri":"http://127.0.0.1/token"}"#),
|
||||
Err(Error::RequestVertexTokenEndpoint)
|
||||
));
|
||||
assert!(matches!(
|
||||
validate_request_credentials("{}"),
|
||||
Err(Error::RequestVertexTokenEndpoint)
|
||||
));
|
||||
#[rstest::rstest]
|
||||
#[case::canonical_endpoint(r#"{"token_uri":"https://oauth2.googleapis.com/token"}"#, true)]
|
||||
#[case::noncanonical_endpoint(r#"{"token_uri":"http://127.0.0.1/token"}"#, false)]
|
||||
#[case::missing_endpoint("{}", false)]
|
||||
fn request_credentials_require_canonical_token_endpoint(
|
||||
#[case] credentials: &str,
|
||||
#[case] accepted: bool,
|
||||
) {
|
||||
assert_eq!(validate_request_credentials(credentials).is_ok(), accepted);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
|
|
|
|||
|
|
@ -12,4 +12,5 @@ thiserror.workspace = true
|
|||
veil.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
rstest.workspace = true
|
||||
tokio.workspace = true
|
||||
|
|
|
|||
|
|
@ -86,7 +86,9 @@ impl CredentialPlan {
|
|||
Self::Caller(caller) => {
|
||||
let credential = caller.acquire().await?;
|
||||
if credential.secret().expose().is_empty() {
|
||||
return Err(Error::EmptyCallerCredential);
|
||||
return Err(Error::EmptyCallerCredential(
|
||||
"credential caller returned an empty credential",
|
||||
));
|
||||
}
|
||||
Ok(CredentialPlanResolution::Resolved(credential))
|
||||
}
|
||||
|
|
@ -147,10 +149,15 @@ mod tests {
|
|||
|
||||
impl CredentialResolver for FailingResolver {
|
||||
fn resolve<'a>(&'a self, _reference: &'a CredentialRef) -> CredentialLookupFuture<'a> {
|
||||
Box::pin(async { Err(Error::UnresolvedOidcReference) })
|
||||
Box::pin(async {
|
||||
Err(Error::CredentialAcquisition(
|
||||
"host credential lookup failed".into(),
|
||||
))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[rstest::rstest]
|
||||
#[tokio::test]
|
||||
async fn acquisition_failure_is_terminal() {
|
||||
let resolver = CredentialResolverHandle::new(Arc::new(FailingResolver));
|
||||
|
|
@ -161,6 +168,9 @@ mod tests {
|
|||
.await
|
||||
.expect_err("acquisition errors cannot become fallback");
|
||||
|
||||
assert_eq!(error, Error::UnresolvedOidcReference);
|
||||
assert_eq!(
|
||||
error,
|
||||
Error::CredentialAcquisition("host credential lookup failed".into())
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -2,84 +2,16 @@ use thiserror::Error as ThisError;
|
|||
|
||||
#[derive(Clone, Debug, ThisError, PartialEq, Eq)]
|
||||
pub enum Error {
|
||||
#[error("invalid authentication configuration: credential header already exists")]
|
||||
ExistingCredentialHeader,
|
||||
#[error(
|
||||
"invalid authentication configuration: credential plan is not allowed by the provider auth policy"
|
||||
)]
|
||||
DisallowedCredentialPlan,
|
||||
#[error("invalid authentication configuration: credential cannot be empty")]
|
||||
EmptyCredential,
|
||||
#[error("invalid authentication configuration: invalid Azure credential selector")]
|
||||
InvalidAzureSelector,
|
||||
#[error(
|
||||
"invalid authentication configuration: ClientSecretCredential requires tenant_id, client_id, and client_secret"
|
||||
)]
|
||||
MissingClientSecretFields,
|
||||
#[error("invalid authentication configuration: WorkloadIdentityCredential requires tenant_id")]
|
||||
MissingWorkloadTenant,
|
||||
#[error("invalid authentication configuration: WorkloadIdentityCredential requires client_id")]
|
||||
MissingWorkloadClient,
|
||||
#[error(
|
||||
"invalid authentication configuration: WorkloadIdentityCredential requires azure_federated_token_file"
|
||||
)]
|
||||
MissingWorkloadTokenFile,
|
||||
#[error(
|
||||
"invalid authentication configuration: credential reference requires a host credential resolver"
|
||||
)]
|
||||
MissingHostResolver,
|
||||
#[error(
|
||||
"invalid authentication configuration: caller credential plan requires provider-specific inputs"
|
||||
)]
|
||||
MissingCallerInputs,
|
||||
#[error("invalid authentication configuration: credential header {0} already exists")]
|
||||
DuplicateHeader(&'static str),
|
||||
#[error("invalid authentication configuration: {0} must be a string or null")]
|
||||
InvalidFieldType(String),
|
||||
#[error("invalid authentication configuration: unsupported OIDC reference")]
|
||||
UnsupportedOidcReference,
|
||||
#[error("invalid authentication configuration: {0} cannot be empty")]
|
||||
EmptyReference(String),
|
||||
#[error("invalid authentication configuration: Azure credential initialization failed: {0}")]
|
||||
AzureCredentialInitialization(String),
|
||||
#[error(
|
||||
"invalid authentication configuration: Azure authority must be an HTTPS origin without credentials, query, or fragment"
|
||||
)]
|
||||
InvalidAzureAuthority,
|
||||
#[error(
|
||||
"invalid authentication configuration: request-controlled Azure auth inputs cannot be combined with host credentials"
|
||||
)]
|
||||
MixedAzureCredentialSources,
|
||||
#[error(
|
||||
"invalid authentication configuration: request-controlled Azure credential references are not allowed"
|
||||
)]
|
||||
RequestAzureCredentialReference,
|
||||
#[error(
|
||||
"invalid authentication configuration: host credentials cannot be sent to a request-controlled Azure endpoint"
|
||||
)]
|
||||
RequestAzureCredentialDestination,
|
||||
#[error(
|
||||
"invalid authentication configuration: credentials cannot be sent to a request-controlled Vertex AI endpoint"
|
||||
)]
|
||||
RequestVertexCredentialDestination,
|
||||
#[error(
|
||||
"invalid authentication configuration: request-controlled Vertex credentials must use the canonical Google OAuth token endpoint"
|
||||
)]
|
||||
RequestVertexTokenEndpoint,
|
||||
#[error("invalid authentication configuration: {0}")]
|
||||
InvalidConfiguration(#[source] ErrorDetail),
|
||||
#[error("credential acquisition failed: {0}")]
|
||||
AzureTokenAcquisition(String),
|
||||
#[error("credential acquisition failed: Vertex AI credentials: {0}")]
|
||||
VertexTokenAcquisition(String),
|
||||
CredentialAcquisition(#[source] ErrorDetail),
|
||||
#[error("credential caller failed: {0}")]
|
||||
EmptyCallerCredential(&'static str),
|
||||
#[error("{0}")]
|
||||
ProviderAuthentication(String),
|
||||
#[error("credential acquisition failed: {}", .0.iter().map(ToString::to_string).collect::<Vec<_>>().join("; "))]
|
||||
CredentialChain(Vec<Error>),
|
||||
#[error("credential caller failed: credential caller returned an empty credential")]
|
||||
EmptyCallerCredential,
|
||||
#[error("credential caller failed: Azure AD token provider returned an empty token")]
|
||||
EmptyAzureToken,
|
||||
#[error("credential acquisition failed: Azure OIDC reference did not resolve to a value")]
|
||||
UnresolvedOidcReference,
|
||||
#[error(
|
||||
"Missing {provider} API Key - Set `api_key` or the {environment_variable} environment variable"
|
||||
)]
|
||||
|
|
@ -87,34 +19,87 @@ pub enum Error {
|
|||
provider: &'static str,
|
||||
environment_variable: &'static str,
|
||||
},
|
||||
#[error(
|
||||
"Missing {provider} API Base - Set {environment_variable} environment variable or pass api_base parameter"
|
||||
)]
|
||||
#[error("Missing {provider} API Base - {guidance}")]
|
||||
MissingApiBase {
|
||||
provider: &'static str,
|
||||
environment_variable: &'static str,
|
||||
guidance: &'static str,
|
||||
},
|
||||
#[error(
|
||||
"Missing Azure API Base - Set `api_base` or the AZURE_API_BASE environment variable. Expected format: https://<resource-name>.services.ai.azure.com/anthropic"
|
||||
)]
|
||||
MissingAzureApiBase,
|
||||
#[error("invalid authentication header")]
|
||||
InvalidHeader,
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::Error;
|
||||
#[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)]
|
||||
pub enum ErrorDetail {
|
||||
#[error("{0}")]
|
||||
Message(String),
|
||||
#[error("{field} must be {expected}")]
|
||||
InvalidType {
|
||||
field: String,
|
||||
expected: &'static str,
|
||||
},
|
||||
#[error("{subject} cannot be empty")]
|
||||
Empty { subject: String },
|
||||
#[error("credential header {0} already exists")]
|
||||
DuplicateHeader(&'static str),
|
||||
#[error("{operation} failed: {source}")]
|
||||
Failed {
|
||||
operation: &'static str,
|
||||
#[source]
|
||||
source: ErrorSource,
|
||||
},
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn missing_api_key_names_provider_and_environment_variable() {
|
||||
assert_eq!(
|
||||
Error::MissingApiKey {
|
||||
provider: "Anthropic",
|
||||
environment_variable: "ANTHROPIC_API_KEY",
|
||||
}
|
||||
.to_string(),
|
||||
"Missing Anthropic API Key - Set `api_key` or the ANTHROPIC_API_KEY environment variable"
|
||||
);
|
||||
impl ErrorDetail {
|
||||
pub fn failed(
|
||||
operation: &'static str,
|
||||
source: impl std::error::Error + Send + Sync + 'static,
|
||||
) -> Self {
|
||||
Self::Failed {
|
||||
operation,
|
||||
source: ErrorSource::new(source),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl From<String> for ErrorDetail {
|
||||
fn from(message: String) -> Self {
|
||||
Self::Message(message)
|
||||
}
|
||||
}
|
||||
|
||||
impl From<&str> for ErrorDetail {
|
||||
fn from(message: &str) -> Self {
|
||||
Self::Message(message.into())
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug)]
|
||||
pub struct ErrorSource(std::sync::Arc<dyn std::error::Error + Send + Sync>);
|
||||
|
||||
impl std::ops::Deref for ErrorSource {
|
||||
type Target = dyn std::error::Error + Send + Sync;
|
||||
|
||||
fn deref(&self) -> &Self::Target {
|
||||
self.0.as_ref()
|
||||
}
|
||||
}
|
||||
|
||||
impl std::fmt::Display for ErrorSource {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
std::fmt::Display::fmt(&self.0, formatter)
|
||||
}
|
||||
}
|
||||
|
||||
impl ErrorSource {
|
||||
pub fn new(error: impl std::error::Error + Send + Sync + 'static) -> Self {
|
||||
Self(std::sync::Arc::new(error))
|
||||
}
|
||||
}
|
||||
|
||||
impl PartialEq for ErrorSource {
|
||||
fn eq(&self, other: &Self) -> bool {
|
||||
std::sync::Arc::ptr_eq(&self.0, &other.0)
|
||||
}
|
||||
}
|
||||
|
||||
impl Eq for ErrorSource {}
|
||||
|
|
|
|||
|
|
@ -21,13 +21,17 @@ pub fn apply_credential(
|
|||
placement: CredentialPlacement,
|
||||
) -> Result<Vec<(String, String)>, Error> {
|
||||
if credential.trim().is_empty() {
|
||||
return Err(Error::EmptyCredential);
|
||||
return Err(Error::InvalidConfiguration(
|
||||
"credential cannot be empty".into(),
|
||||
));
|
||||
}
|
||||
if headers
|
||||
.iter()
|
||||
.any(|(name, _)| name.eq_ignore_ascii_case(placement.header_name()))
|
||||
{
|
||||
return Err(Error::DuplicateHeader(placement.header_name()));
|
||||
return Err(Error::InvalidConfiguration(
|
||||
crate::ErrorDetail::DuplicateHeader(placement.header_name()),
|
||||
));
|
||||
}
|
||||
let value = match placement {
|
||||
CredentialPlacement::Bearer => format!("Bearer {credential}"),
|
||||
|
|
|
|||
|
|
@ -50,7 +50,7 @@ pub use credential::{
|
|||
CredentialFileRef, CredentialLookup, CredentialLookupFuture, CredentialPlan,
|
||||
CredentialPlanResolution, CredentialRef, CredentialResolver, CredentialResolverHandle,
|
||||
};
|
||||
pub use error::Error;
|
||||
pub use error::{Error, ErrorDetail, ErrorSource};
|
||||
pub use http::CredentialPlacement;
|
||||
pub use policy::{CredentialPlanKind, CredentialRule, ExistingHeaderBehavior, ProviderAuthPolicy};
|
||||
pub use secret::SecretValue;
|
||||
|
|
|
|||
|
|
@ -47,14 +47,20 @@ impl ProviderAuthPolicy {
|
|||
if self.has_existing_credential(&headers) {
|
||||
return match self.existing_header_behavior {
|
||||
ExistingHeaderBehavior::Preserve => Ok(headers),
|
||||
ExistingHeaderBehavior::Reject => Err(Error::ExistingCredentialHeader),
|
||||
ExistingHeaderBehavior::Reject => Err(Error::InvalidConfiguration(
|
||||
"credential header already exists".into(),
|
||||
)),
|
||||
};
|
||||
}
|
||||
let rule = self
|
||||
.rules
|
||||
.iter()
|
||||
.find(|rule| rule.kind == kind)
|
||||
.ok_or(Error::DisallowedCredentialPlan)?;
|
||||
.ok_or_else(|| {
|
||||
Error::InvalidConfiguration(
|
||||
"credential plan is not allowed by the provider auth policy".into(),
|
||||
)
|
||||
})?;
|
||||
apply_credential(headers, credential.secret().expose(), rule.placement)
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -39,7 +39,33 @@ impl TokenProviderHandle {
|
|||
Self(caller)
|
||||
}
|
||||
|
||||
pub fn from_callback<F, Fut>(acquire: F) -> Self
|
||||
where
|
||||
F: Fn() -> Fut + Send + Sync + 'static,
|
||||
Fut: Future<Output = Result<ResolvedCredential, Error>> + Send + 'static,
|
||||
{
|
||||
Self::new(Arc::new(CallbackTokenProvider(acquire)))
|
||||
}
|
||||
|
||||
pub async fn acquire(&self) -> Result<ResolvedCredential, Error> {
|
||||
self.0.acquire().await
|
||||
}
|
||||
}
|
||||
|
||||
struct CallbackTokenProvider<F>(F);
|
||||
|
||||
impl<F> std::fmt::Debug for CallbackTokenProvider<F> {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter.write_str("CallbackTokenProvider")
|
||||
}
|
||||
}
|
||||
|
||||
impl<F, Fut> TokenProvider for CallbackTokenProvider<F>
|
||||
where
|
||||
F: Fn() -> Fut + Send + Sync,
|
||||
Fut: Future<Output = Result<ResolvedCredential, Error>> + Send + 'static,
|
||||
{
|
||||
fn acquire(&self) -> TokenFuture<'_> {
|
||||
Box::pin((self.0)())
|
||||
}
|
||||
}
|
||||
|
|
|
|||
74
litellm-rust/crates/auth-types/tests/error.rs
Normal file
74
litellm-rust/crates/auth-types/tests/error.rs
Normal file
|
|
@ -0,0 +1,74 @@
|
|||
use litellm_auth_types::Error;
|
||||
use rstest::rstest;
|
||||
|
||||
#[rstest]
|
||||
#[case::api_key(
|
||||
Error::MissingApiKey { provider: "Example", environment_variable: "EXAMPLE_API_KEY" },
|
||||
"Missing Example API Key - Set `api_key` or the EXAMPLE_API_KEY environment variable"
|
||||
)]
|
||||
#[case::another_api_key(
|
||||
Error::MissingApiKey { provider: "Custom", environment_variable: "CUSTOM_KEY" },
|
||||
"Missing Custom API Key - Set `api_key` or the CUSTOM_KEY environment variable"
|
||||
)]
|
||||
#[case::api_base(
|
||||
Error::MissingApiBase { provider: "Example", guidance: "Pass api_base" },
|
||||
"Missing Example API Base - Pass api_base"
|
||||
)]
|
||||
#[case::another_api_base(
|
||||
Error::MissingApiBase { provider: "Custom", guidance: "Set CUSTOM_ENDPOINT" },
|
||||
"Missing Custom API Base - Set CUSTOM_ENDPOINT"
|
||||
)]
|
||||
#[case::configuration(
|
||||
Error::InvalidConfiguration("credential selector is invalid".into()),
|
||||
"invalid authentication configuration: credential selector is invalid"
|
||||
)]
|
||||
#[case::acquisition(
|
||||
Error::CredentialAcquisition("token expired".into()),
|
||||
"credential acquisition failed: token expired"
|
||||
)]
|
||||
#[case::caller(
|
||||
Error::EmptyCallerCredential("empty token"),
|
||||
"credential caller failed: empty token"
|
||||
)]
|
||||
#[case::provider(
|
||||
Error::ProviderAuthentication("provider rejected credentials".into()),
|
||||
"provider rejected credentials"
|
||||
)]
|
||||
#[case::chain(
|
||||
Error::CredentialChain(vec![
|
||||
Error::CredentialAcquisition("token expired".into()),
|
||||
Error::EmptyCallerCredential("empty token"),
|
||||
]),
|
||||
"credential acquisition failed: credential acquisition failed: token expired; credential caller failed: empty token"
|
||||
)]
|
||||
fn display_preserves_failure_phase_and_caller_context(
|
||||
#[case] error: Error,
|
||||
#[case] expected: &str,
|
||||
) {
|
||||
assert_eq!(error.to_string(), expected);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::configuration(true)]
|
||||
#[case::acquisition(false)]
|
||||
fn contextual_failures_keep_the_original_source(#[case] configuration: bool) {
|
||||
use litellm_auth_types::ErrorDetail;
|
||||
|
||||
let detail = ErrorDetail::failed(
|
||||
"test credential",
|
||||
std::io::Error::from(std::io::ErrorKind::PermissionDenied),
|
||||
);
|
||||
let error = if configuration {
|
||||
Error::InvalidConfiguration(detail)
|
||||
} else {
|
||||
Error::CredentialAcquisition(detail)
|
||||
};
|
||||
let source = std::iter::successors(Some(&error as &dyn std::error::Error), |error| {
|
||||
error.source()
|
||||
})
|
||||
.find_map(|error| error.downcast_ref::<std::io::Error>())
|
||||
.expect("the original credential error remains available");
|
||||
assert_eq!(source.kind(), std::io::ErrorKind::PermissionDenied);
|
||||
assert!(error.to_string().contains("test credential failed:"));
|
||||
assert!(error.to_string().ends_with(&source.to_string()));
|
||||
}
|
||||
104
litellm-rust/crates/auth-types/tests/token.rs
Normal file
104
litellm-rust/crates/auth-types/tests/token.rs
Normal file
|
|
@ -0,0 +1,104 @@
|
|||
use std::{
|
||||
error::Error as StdError,
|
||||
future::{Future, poll_fn},
|
||||
sync::{
|
||||
Arc,
|
||||
atomic::{AtomicBool, AtomicUsize, Ordering},
|
||||
},
|
||||
task::Poll,
|
||||
time::{Duration, SystemTime},
|
||||
};
|
||||
|
||||
use litellm_auth_types::{
|
||||
Error, ErrorDetail, ResolvedCredential, SecretValue, TokenProviderHandle,
|
||||
};
|
||||
use rstest::rstest;
|
||||
|
||||
fn credential(index: usize, access_token: bool) -> ResolvedCredential {
|
||||
let token = SecretValue::new(format!("credential-{index}"));
|
||||
if access_token {
|
||||
return ResolvedCredential::AccessToken {
|
||||
token,
|
||||
expires_on: Some(SystemTime::UNIX_EPOCH + Duration::from_secs(index as u64)),
|
||||
};
|
||||
}
|
||||
ResolvedCredential::Static(token)
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::static_secret(false)]
|
||||
#[case::access_token(true)]
|
||||
#[tokio::test]
|
||||
async fn callbacks_acquire_fresh_credentials_on_demand(#[case] access_token: bool) {
|
||||
let calls = Arc::new(AtomicUsize::new(0));
|
||||
let callback_calls = calls.clone();
|
||||
let provider = TokenProviderHandle::from_callback(move || {
|
||||
let index = callback_calls.fetch_add(1, Ordering::SeqCst);
|
||||
async move {
|
||||
tokio::task::yield_now().await;
|
||||
Ok(credential(index, access_token))
|
||||
}
|
||||
});
|
||||
let cloned = provider.clone();
|
||||
|
||||
assert_eq!(calls.load(Ordering::SeqCst), 0);
|
||||
assert_eq!(
|
||||
provider.acquire().await.unwrap(),
|
||||
credential(0, access_token)
|
||||
);
|
||||
assert_eq!(calls.load(Ordering::SeqCst), 1);
|
||||
assert_eq!(cloned.acquire().await.unwrap(), credential(1, access_token));
|
||||
assert_eq!(calls.load(Ordering::SeqCst), 2);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn callback_errors_preserve_the_original_source() {
|
||||
let provider = TokenProviderHandle::from_callback(|| async {
|
||||
Err(Error::CredentialAcquisition(ErrorDetail::failed(
|
||||
"caller credential",
|
||||
std::io::Error::from(std::io::ErrorKind::PermissionDenied),
|
||||
)))
|
||||
});
|
||||
|
||||
let error = provider.acquire().await.unwrap_err();
|
||||
assert!(matches!(error, Error::CredentialAcquisition(_)));
|
||||
let source = std::iter::successors(Some(&error as &(dyn StdError + 'static)), |error| {
|
||||
(*error).source()
|
||||
})
|
||||
.find_map(|error| error.downcast_ref::<std::io::Error>())
|
||||
.unwrap();
|
||||
assert_eq!(source.kind(), std::io::ErrorKind::PermissionDenied);
|
||||
}
|
||||
|
||||
struct Release(Arc<AtomicBool>);
|
||||
|
||||
impl Drop for Release {
|
||||
fn drop(&mut self) {
|
||||
self.0.store(true, Ordering::SeqCst);
|
||||
}
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn cancelling_acquisition_drops_the_callback_future() {
|
||||
let released = Arc::new(AtomicBool::new(false));
|
||||
let callback_released = released.clone();
|
||||
let provider = TokenProviderHandle::from_callback(move || {
|
||||
let released = callback_released.clone();
|
||||
async move {
|
||||
let _release = Release(released);
|
||||
std::future::pending().await
|
||||
}
|
||||
});
|
||||
|
||||
let mut acquisition = Box::pin(provider.acquire());
|
||||
poll_fn(|context| {
|
||||
assert!(acquisition.as_mut().poll(context).is_pending());
|
||||
assert!(!released.load(Ordering::SeqCst));
|
||||
Poll::Ready(())
|
||||
})
|
||||
.await;
|
||||
drop(acquisition);
|
||||
assert!(released.load(Ordering::SeqCst));
|
||||
}
|
||||
|
|
@ -9,7 +9,7 @@ repository.workspace = true
|
|||
litellm-cache.workspace = true
|
||||
py_literal = "0.4.0"
|
||||
rand.workspace = true
|
||||
rusqlite = { version = "0.40", features = ["bundled"] }
|
||||
rusqlite = { version = "0.39", features = ["bundled"] }
|
||||
serde-pickle = "1.2"
|
||||
serde_json.workspace = true
|
||||
tokio.workspace = true
|
||||
|
|
|
|||
|
|
@ -1,12 +1,12 @@
|
|||
- Target invariants, not completion claims
|
||||
- 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)
|
||||
- This crate owns compatibility for all existing Python callbacks and loggers, including `CustomLogger`. `mapping.rs` owns the executable call bindings and the inventory of Python-owned hooks. A Python-owned entry records an existing path, never permission to invoke it a second time. The native call adapter preserves the `Logging` contract (`function_setup`, the deployment hooks, `pre_call`/`post_call`, the sync and async success and failure fan-out, the deferred proxy release, the argument sharing those callbacks rely on)
|
||||
- Smell test: if a future callback host (`callbacks-v1-python`, WASM, in-process Rust) could share a piece of this crate, it does not belong here
|
||||
- SDK request policy (credential inheritance, the budget and retry-count limits) is 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
|
||||
- SDK request policy (credential inheritance, the budget and retry-count limits) is a separate hook supplied by `python-bridge`; compose it after this adapter so logging adopts the final keyword view before policy mutates or rejects it
|
||||
- The driver in `litellm-host-python`, the routes and core see one `PythonCallHooks` using the shared `CallEvent`; they never learn which Python objects consume a call
|
||||
- Every litellm Python internal Rust still borrows is a variant of `LegacyPython`, grouped by subsystem, with its signature pinned in `python_contract.json`
|
||||
- The enum only shrinks: when Rust owns a subsystem, delete its group rather than adding a Rust path beside it
|
||||
- Calling a user's own callback directly is permanent Python surface and gets its own type outside `LegacyPython`
|
||||
- `PublicCall` is the caller's call as `Logging` sees it: the positional arguments, the keyword view as the 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
|
||||
- `PublicCall` is the caller's call as `Logging` sees it: the positional arguments, the keyword view as the call rewrites it (setup, deployment hook, preflight) and the bound request object backing omitted keywords; shared bridge composition hands it to `LegacyLogging`; routes use the neutral call boundary
|
||||
- `setup` reuses a `Logging` passed as `litellm_logging_obj` (the proxy and Router) and otherwise builds one through `function_setup`; which callbacks run is `Logging`'s decision, never this crate's
|
||||
- Callbacks receive the caller's own objects and may mutate them; this crate alone carries that obligation
|
||||
- Retain complete boundary arguments, opaque values, aliases, omitted/default distinctions and deliberate copies; preserve the deployment-hook kwargs view
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ license.workspace = true
|
|||
repository.workspace = true
|
||||
|
||||
[dependencies]
|
||||
litellm-types.workspace = true
|
||||
litellm-host.workspace = true
|
||||
litellm-host-python.workspace = true
|
||||
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -3,16 +3,13 @@
|
|||
//! lifetime. No other callback host has that obligation, which is why nothing outside
|
||||
//! this crate holds them.
|
||||
|
||||
use litellm_host::{machine::Machine, protocol::Protocol};
|
||||
use litellm_host_python::{Preflight, ProtocolHost, lookup, run_call};
|
||||
use litellm_host_python::lookup;
|
||||
use pyo3::{
|
||||
gc::{PyTraverseError, PyVisit},
|
||||
prelude::*,
|
||||
types::{PyDict, PyTuple},
|
||||
};
|
||||
|
||||
use crate::{LegacyLogging, LegacySurface};
|
||||
|
||||
pub struct PublicCall {
|
||||
args: Py<PyTuple>,
|
||||
kwargs: Py<PyDict>,
|
||||
|
|
@ -34,6 +31,10 @@ impl PublicCall {
|
|||
})
|
||||
}
|
||||
|
||||
pub fn arguments(&self, py: Python<'_>) -> Py<PyDict> {
|
||||
self.kwargs.clone_ref(py)
|
||||
}
|
||||
|
||||
pub(crate) fn args(&self) -> &Py<PyTuple> {
|
||||
&self.args
|
||||
}
|
||||
|
|
@ -64,34 +65,6 @@ impl PublicCall {
|
|||
}
|
||||
}
|
||||
|
||||
/// Runs one native call under the legacy `Logging` contract: the protocol host projects from
|
||||
/// the keyword view the contract prepares and `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
|
||||
H: ProtocolHost + 'static,
|
||||
M: Machine<Protocol = H::Protocol, Complete = <H::Protocol as Protocol>::Response> + 'static,
|
||||
{
|
||||
let arguments = call.kwargs.clone_ref(py);
|
||||
run_call(
|
||||
py,
|
||||
machine,
|
||||
host,
|
||||
Box::new(LegacyLogging::new(py, surface, call, asynchronous)),
|
||||
preflight,
|
||||
arguments,
|
||||
asynchronous,
|
||||
)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
|
@ -110,7 +83,7 @@ mod tests {
|
|||
(call, locals)
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[rstest::rstest]
|
||||
fn capture_copies_the_keyword_dict_without_copying_its_values() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
|
|
|
|||
|
|
@ -2,7 +2,7 @@
|
|||
//! the deferred and worker-submitted success paths, and the sync-callbacks-for-async-calls
|
||||
//! duplication. All of it expires with the legacy callback contract.
|
||||
|
||||
use litellm_host::event::{RequestContext, WireRequest};
|
||||
use litellm_host::interceptors::{RequestContext, WireRequest};
|
||||
use litellm_host_python::to_py;
|
||||
use pyo3::{exceptions::PyBaseException, prelude::*, types::PyDict};
|
||||
|
||||
|
|
|
|||
|
|
@ -1,26 +1,24 @@
|
|||
//! The legacy `@client` wrapper as the native call sees it: litellm's `Logging` object, the
|
||||
//! sync and async callback registries it fans out to, the deployment hooks 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. The SDK's own request policy
|
||||
//! (credential inheritance, the budget and retry limits) is the driver's preflight, not this
|
||||
//! crate's.
|
||||
//! [`PythonCallHooks`](litellm_host_python::PythonCallHooks), so the driver, the routes and
|
||||
//! core never learn which Python object is on the other end.
|
||||
//!
|
||||
//! Legacy callbacks receive the caller's own objects and may mutate them. [`PublicCall`]
|
||||
//! is where those objects live, and [`run_legacy_call`] is how a route hands them over
|
||||
//! without keeping a copy.
|
||||
//! is where those objects live.
|
||||
|
||||
mod adapter;
|
||||
mod call;
|
||||
mod callbacks;
|
||||
mod deferred;
|
||||
mod logger;
|
||||
mod mapping;
|
||||
mod python;
|
||||
pub(crate) use adapter::LegacyLogging;
|
||||
pub use adapter::{LegacySurface, PassThroughStream};
|
||||
pub use call::{PublicCall, run_legacy_call};
|
||||
pub use adapter::LegacyLogging;
|
||||
pub use call::PublicCall;
|
||||
pub(crate) use callbacks::{LegacyCallbacks, is_internal_call};
|
||||
pub(crate) use logger::{DeploymentHooks, PythonLogger, finalize, setup};
|
||||
pub use mapping::{CallBoundary, CallbackMapping, Dispatch, callback_mappings};
|
||||
|
||||
#[cfg(test)]
|
||||
mod test_support;
|
||||
|
|
|
|||
285
litellm-rust/crates/callbacks-legacy-python/src/mapping.rs
Normal file
285
litellm-rust/crates/callbacks-legacy-python/src/mapping.rs
Normal file
|
|
@ -0,0 +1,285 @@
|
|||
use litellm_host::{
|
||||
hooks::CallHooks,
|
||||
interceptors::{RawResponse, RequestContext, WireRequest},
|
||||
lifecycle::{ExecutionEvent, FailureOrigin, Timing},
|
||||
};
|
||||
use litellm_host_python::{HookStep, PythonCallEvent, PythonRuntime};
|
||||
use pyo3::{prelude::*, types::PyDict};
|
||||
|
||||
use crate::LegacyLogging;
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||||
pub enum CallBoundary {
|
||||
PrepareArguments,
|
||||
BeforeProviderRequest,
|
||||
AfterProviderResponse,
|
||||
TransformResponse,
|
||||
Succeeded,
|
||||
Failed,
|
||||
StreamOpened,
|
||||
StreamChunk,
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||||
pub enum Dispatch {
|
||||
Call(CallBoundary),
|
||||
Python(&'static str),
|
||||
DeclarationOnly,
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||||
pub struct CallbackMapping {
|
||||
pub callback: &'static str,
|
||||
pub dispatch: Dispatch,
|
||||
}
|
||||
|
||||
struct Binding<H> {
|
||||
boundary: CallBoundary,
|
||||
invoke: H,
|
||||
callbacks: &'static [&'static str],
|
||||
}
|
||||
|
||||
impl<H> Binding<H> {
|
||||
fn mappings(&self) -> impl Iterator<Item = CallbackMapping> {
|
||||
self.callbacks.iter().map(|callback| CallbackMapping {
|
||||
callback,
|
||||
dispatch: Dispatch::Call(self.boundary),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
type Step<T> = PyResult<HookStep<LegacyLogging, T>>;
|
||||
type Prepare = fn(&mut LegacyLogging, Python<'_>, Py<PyDict>, f64) -> Step<Py<PyDict>>;
|
||||
type Before =
|
||||
fn(&mut LegacyLogging, Python<'_>, Box<WireRequest>, &RequestContext) -> Step<Box<WireRequest>>;
|
||||
type After = fn(&mut LegacyLogging, Python<'_>, &RawResponse) -> Step<()>;
|
||||
type Transform = fn(&mut LegacyLogging, Python<'_>, Py<PyAny>, Timing) -> Step<Py<PyAny>>;
|
||||
type Success = fn(&mut LegacyLogging, Python<'_>, Timing, &Py<PyAny>) -> Step<()>;
|
||||
type Failure = fn(&mut LegacyLogging, Python<'_>, Timing, FailureOrigin, &PyErr) -> Step<()>;
|
||||
type Open = fn(&mut LegacyLogging, Python<'_>) -> PyResult<()>;
|
||||
type Chunk = fn(&mut LegacyLogging, Python<'_>, &Py<PyAny>) -> PyResult<()>;
|
||||
|
||||
const PREPARE: Binding<Prepare> = Binding {
|
||||
boundary: CallBoundary::PrepareArguments,
|
||||
invoke: LegacyLogging::prepare_call,
|
||||
callbacks: &["async_pre_call_deployment_hook"],
|
||||
};
|
||||
|
||||
const BEFORE: Binding<Before> = Binding {
|
||||
boundary: CallBoundary::BeforeProviderRequest,
|
||||
invoke: LegacyLogging::pre_call,
|
||||
callbacks: &["log_pre_api_call", "log_input_event"],
|
||||
};
|
||||
|
||||
const AFTER: Binding<After> = Binding {
|
||||
boundary: CallBoundary::AfterProviderResponse,
|
||||
invoke: LegacyLogging::post_call,
|
||||
callbacks: &["log_post_api_call"],
|
||||
};
|
||||
|
||||
const TRANSFORM: Binding<Transform> = Binding {
|
||||
boundary: CallBoundary::TransformResponse,
|
||||
invoke: LegacyLogging::transform_public_response,
|
||||
callbacks: &["async_post_call_success_deployment_hook"],
|
||||
};
|
||||
|
||||
const SUCCESS: Binding<Success> = Binding {
|
||||
boundary: CallBoundary::Succeeded,
|
||||
invoke: LegacyLogging::succeeded,
|
||||
callbacks: &[
|
||||
"log_success_event",
|
||||
"async_log_success_event",
|
||||
"logging_hook",
|
||||
"async_logging_hook",
|
||||
"redact_standard_logging_payload_from_model_call_details",
|
||||
"log_event",
|
||||
"async_log_event",
|
||||
],
|
||||
};
|
||||
|
||||
const FAILURE: Binding<Failure> = Binding {
|
||||
boundary: CallBoundary::Failed,
|
||||
invoke: LegacyLogging::failed,
|
||||
callbacks: &[
|
||||
"async_post_call_failure_deployment_hook",
|
||||
"log_failure_event",
|
||||
"async_log_failure_event",
|
||||
"log_model_group_rate_limit_error",
|
||||
"log_event",
|
||||
"async_log_event",
|
||||
],
|
||||
};
|
||||
|
||||
const OPEN: Binding<Open> = Binding {
|
||||
boundary: CallBoundary::StreamOpened,
|
||||
invoke: LegacyLogging::stream_opened,
|
||||
callbacks: &[],
|
||||
};
|
||||
|
||||
const CHUNK: Binding<Chunk> = Binding {
|
||||
boundary: CallBoundary::StreamChunk,
|
||||
invoke: LegacyLogging::stream_chunk,
|
||||
callbacks: &[],
|
||||
};
|
||||
|
||||
pub fn callback_mappings() -> impl Iterator<Item = CallbackMapping> {
|
||||
PREPARE
|
||||
.mappings()
|
||||
.chain(BEFORE.mappings())
|
||||
.chain(AFTER.mappings())
|
||||
.chain(TRANSFORM.mappings())
|
||||
.chain(SUCCESS.mappings())
|
||||
.chain(FAILURE.mappings())
|
||||
.chain(OPEN.mappings())
|
||||
.chain(CHUNK.mappings())
|
||||
.chain(PYTHON_CALLBACKS.iter().copied())
|
||||
}
|
||||
|
||||
impl CallHooks<PythonRuntime> for LegacyLogging {
|
||||
fn prepare_arguments(
|
||||
&mut self,
|
||||
py: Python<'_>,
|
||||
arguments: Py<PyDict>,
|
||||
started_at: f64,
|
||||
) -> Step<Py<PyDict>> {
|
||||
(PREPARE.invoke)(self, py, arguments, started_at)
|
||||
}
|
||||
|
||||
fn arguments_prepared(&mut self, py: Python<'_>, arguments: &Py<PyDict>) -> PyResult<()> {
|
||||
self.adopt_arguments(py, arguments);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn before_provider_request(
|
||||
&mut self,
|
||||
py: Python<'_>,
|
||||
wire: Box<WireRequest>,
|
||||
context: &RequestContext,
|
||||
) -> Step<Box<WireRequest>> {
|
||||
(BEFORE.invoke)(self, py, wire, context)
|
||||
}
|
||||
|
||||
fn transform_response(
|
||||
&mut self,
|
||||
py: Python<'_>,
|
||||
response: Py<PyAny>,
|
||||
timing: Timing,
|
||||
) -> Step<Py<PyAny>> {
|
||||
(TRANSFORM.invoke)(self, py, response, timing)
|
||||
}
|
||||
|
||||
fn on_event(&mut self, py: Python<'_>, event: PythonCallEvent<'_>) -> Step<()> {
|
||||
match event {
|
||||
PythonCallEvent::Started { .. } | PythonCallEvent::Cancelled { .. } => {
|
||||
Ok(HookStep::Ready(()))
|
||||
}
|
||||
PythonCallEvent::Execution(ExecutionEvent::ProviderResponseReceived { raw }) => {
|
||||
(AFTER.invoke)(self, py, raw)
|
||||
}
|
||||
PythonCallEvent::Succeeded { timing, response } => {
|
||||
(SUCCESS.invoke)(self, py, timing, response)
|
||||
}
|
||||
PythonCallEvent::Failed {
|
||||
timing,
|
||||
origin,
|
||||
error,
|
||||
} => (FAILURE.invoke)(self, py, timing, origin, error),
|
||||
}
|
||||
}
|
||||
|
||||
fn on_stream_open(&mut self, py: Python<'_>) -> PyResult<()> {
|
||||
(OPEN.invoke)(self, py)
|
||||
}
|
||||
|
||||
fn on_stream_chunk(&mut self, py: Python<'_>, chunk: &Py<PyAny>) -> PyResult<()> {
|
||||
(CHUNK.invoke)(self, py, chunk)
|
||||
}
|
||||
}
|
||||
|
||||
macro_rules! python_callbacks {
|
||||
($($dispatch:expr => [$($callback:literal),* $(,)?]),* $(,)?) => {
|
||||
const PYTHON_CALLBACKS: &[CallbackMapping] = &[
|
||||
$($(CallbackMapping { callback: $callback, dispatch: $dispatch },)*)*
|
||||
];
|
||||
};
|
||||
}
|
||||
|
||||
python_callbacks! {
|
||||
Dispatch::Python("litellm.router") => [
|
||||
"async_pre_routing_hook",
|
||||
"async_filter_deployments",
|
||||
"pre_call_check",
|
||||
"async_pre_call_check",
|
||||
],
|
||||
Dispatch::Python("litellm.router_utils.fallback_event_handlers") => [
|
||||
"log_success_fallback_event",
|
||||
"log_failure_fallback_event",
|
||||
],
|
||||
Dispatch::Python("litellm.proxy.utils") => [
|
||||
"async_pre_call_hook",
|
||||
"async_post_call_response_headers_hook",
|
||||
"async_post_call_failure_hook",
|
||||
"async_post_call_success_hook",
|
||||
"async_moderation_hook",
|
||||
"async_post_call_streaming_hook",
|
||||
"async_post_call_streaming_iterator_hook",
|
||||
"async_filter_listed_models",
|
||||
],
|
||||
Dispatch::Python("litellm.litellm_core_utils.litellm_logging") => [
|
||||
"async_get_chat_completion_prompt",
|
||||
"get_chat_completion_prompt",
|
||||
"log_stream_event",
|
||||
"async_log_stream_event",
|
||||
"async_post_mcp_tool_call_hook",
|
||||
],
|
||||
Dispatch::Python("litellm.llms.anthropic.pass_through.messages.handler") => [
|
||||
"async_pre_request_hook",
|
||||
],
|
||||
Dispatch::Python("litellm.litellm_core_utils.streaming_handler") => [
|
||||
"async_post_call_streaming_deployment_hook",
|
||||
],
|
||||
Dispatch::Python("litellm.responses.streaming_iterator") => [
|
||||
"async_post_call_streaming_deployment_hook",
|
||||
],
|
||||
Dispatch::Python("litellm.main") => [
|
||||
"translate_completion_input_params",
|
||||
"translate_completion_output_params",
|
||||
"translate_completion_output_params_streaming",
|
||||
],
|
||||
Dispatch::Python("litellm.integrations.argilla") => ["async_dataset_hook"],
|
||||
Dispatch::Python("litellm.proxy.management_helpers.audit_logs") => ["async_log_audit_log_event"],
|
||||
Dispatch::Python("litellm.llms.custom_httpx.llm_http_handler") => [
|
||||
"async_should_run_agentic_loop",
|
||||
"async_run_agentic_loop",
|
||||
"async_build_agentic_loop_plan",
|
||||
"async_post_agentic_loop_response_hook",
|
||||
"async_agentic_loop_cleanup_hook",
|
||||
"async_should_run_chat_completion_agentic_loop",
|
||||
"async_run_chat_completion_agentic_loop",
|
||||
"async_build_chat_completion_agentic_loop_plan",
|
||||
],
|
||||
Dispatch::Python("litellm.litellm_core_utils.chat_completion_agentic_loop") => [
|
||||
"async_should_run_agentic_loop",
|
||||
"async_run_agentic_loop",
|
||||
"async_build_agentic_loop_plan",
|
||||
"async_post_agentic_loop_response_hook",
|
||||
"async_agentic_loop_cleanup_hook",
|
||||
],
|
||||
Dispatch::Python("litellm.llms.openai.openai") => [
|
||||
"async_should_run_chat_completion_agentic_loop",
|
||||
"async_run_chat_completion_agentic_loop",
|
||||
],
|
||||
Dispatch::Python("litellm.proxy.spend_tracking.cold_storage_handler") => [
|
||||
"get_proxy_server_request_from_cold_storage_with_object_key",
|
||||
],
|
||||
Dispatch::Python("litellm.integrations.custom_logger") => [
|
||||
"truncate_standard_logging_payload_content",
|
||||
"redacts_messages_itself",
|
||||
"handle_callback_failure",
|
||||
"get_callback_env_vars",
|
||||
],
|
||||
Dispatch::DeclarationOnly => [
|
||||
"async_log_pre_api_call",
|
||||
"async_log_input_event",
|
||||
],
|
||||
}
|
||||
|
|
@ -3,7 +3,7 @@ use std::ffi::CStr;
|
|||
use pyo3::prelude::*;
|
||||
use pyo3::types::{PyDict, PyTuple};
|
||||
|
||||
use crate::{LegacyLogging, LegacySurface, PublicCall};
|
||||
use crate::{LegacyLogging, PublicCall};
|
||||
|
||||
/// The parameters of every `callbacks_legacy_python` function, as the real module declares them.
|
||||
/// `tests/unit/rust_bridge/test_callbacks_legacy_python.py` pins this file to the Python
|
||||
|
|
@ -45,11 +45,14 @@ def contracted(name, fake):
|
|||
if not hasattr(legacy, 'is_internal'):
|
||||
legacy.is_internal = contextvars.ContextVar('is_internal_call', default=False)
|
||||
|
||||
def setup(call_type, args, kwargs, start, asynchronous):
|
||||
logger = kwargs['logger_factory'](kwargs) if 'logger_factory' in kwargs else kwargs['logger']
|
||||
logger.setup_call_type = call_type
|
||||
return types.SimpleNamespace(logger=logger, kwargs=kwargs)
|
||||
|
||||
|
||||
FAKES = {
|
||||
'setup': lambda call_type, args, kwargs, start, asynchronous: types.SimpleNamespace(
|
||||
logger=kwargs['logger_factory'](kwargs) if 'logger_factory' in kwargs else kwargs['logger'],
|
||||
kwargs=kwargs,
|
||||
),
|
||||
'setup': setup,
|
||||
'finalize': lambda response, logger, kwargs, start, end: logger.record('finalize', response),
|
||||
'update_logging': lambda logger, kwargs, model, optional_params, litellm_params, provider: logger.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
|
|
@ -82,10 +85,10 @@ FAKES = {
|
|||
),
|
||||
'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': lambda logger, url_route, endpoint_type, 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),
|
||||
'stream_failure': lambda logger, endpoint_type, request_body, chunks, error: logger.record('stream_failure', error),
|
||||
}
|
||||
assert FAKES.keys() == CONTRACT.keys(), sorted(FAKES.keys() ^ CONTRACT.keys())
|
||||
for name, fake in FAKES.items():
|
||||
|
|
@ -186,14 +189,5 @@ pub(crate) fn legacy_call(
|
|||
.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,
|
||||
)
|
||||
LegacyLogging::new(py, litellm_types::Operation::Ocr, call, asynchronous)
|
||||
}
|
||||
|
|
|
|||
61
litellm-rust/crates/config/src/includes.rs
Normal file
61
litellm-rust/crates/config/src/includes.rs
Normal file
|
|
@ -0,0 +1,61 @@
|
|||
use std::{
|
||||
collections::{BTreeSet, VecDeque},
|
||||
path::{Path, PathBuf},
|
||||
};
|
||||
|
||||
use serde::Deserialize;
|
||||
use serde_yaml_ng::{Mapping, Value};
|
||||
|
||||
use crate::Error;
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct Includes {
|
||||
#[serde(default)]
|
||||
include: Vec<String>,
|
||||
}
|
||||
|
||||
fn read(path: &Path) -> Result<Mapping, Error> {
|
||||
Ok(serde_yaml_ng::from_str(&std::fs::read_to_string(path)?)?)
|
||||
}
|
||||
|
||||
fn entries(config: &Mapping, path: &Path) -> Result<Vec<(String, PathBuf)>, Error> {
|
||||
let includes: Includes = serde_yaml_ng::from_value(Value::Mapping(config.clone()))?;
|
||||
Ok(includes
|
||||
.include
|
||||
.into_iter()
|
||||
.map(|entry| (entry, path.to_owned()))
|
||||
.collect())
|
||||
}
|
||||
|
||||
pub(super) fn load(path: &Path) -> Result<Value, Error> {
|
||||
let root = path.canonicalize()?;
|
||||
let mut merged = read(&root)?;
|
||||
let mut pending: VecDeque<_> = entries(&merged, &root)?.into();
|
||||
let mut loaded = BTreeSet::from([root.clone()]);
|
||||
merged.remove(Value::String("include".into()));
|
||||
while let Some((entry, declaring)) = pending.pop_front() {
|
||||
let declared = declaring.parent().unwrap_or(Path::new(".")).join(&entry);
|
||||
let fallback = root.parent().unwrap_or(Path::new(".")).join(&entry);
|
||||
let location = if declared.exists() {
|
||||
declared
|
||||
} else {
|
||||
fallback
|
||||
}
|
||||
.canonicalize()?;
|
||||
if !loaded.insert(location.clone()) {
|
||||
continue;
|
||||
}
|
||||
let mut included = read(&location)?;
|
||||
pending.extend(entries(&included, &location)?);
|
||||
included.remove(Value::String("include".into()));
|
||||
for (key, value) in included {
|
||||
match (merged.get_mut(&key), value) {
|
||||
(Some(Value::Sequence(base)), Value::Sequence(extra)) => base.extend(extra),
|
||||
(_, value) => {
|
||||
merged.insert(key, value);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
Ok(Value::Mapping(merged))
|
||||
}
|
||||
|
|
@ -1,40 +1,82 @@
|
|||
mod error;
|
||||
mod includes;
|
||||
mod mcp;
|
||||
mod model;
|
||||
mod settings;
|
||||
mod value;
|
||||
|
||||
use std::path::Path;
|
||||
use std::{fmt, path::Path};
|
||||
|
||||
use litellm_auth_types::SecretValue;
|
||||
use serde::Deserialize;
|
||||
|
||||
pub use error::Error;
|
||||
pub use mcp::{McpAuth, McpServer, McpTransport};
|
||||
pub use model::{LiteLlmParams, Model};
|
||||
pub use settings::{GeneralSettings, LiteLlmSettings, RouterSettings};
|
||||
pub use value::{AdditionalFields, Flag, NumberOrString, Object, OneOrMany, Value};
|
||||
|
||||
#[derive(Clone, Debug, Deserialize)]
|
||||
#[serde(deny_unknown_fields)]
|
||||
#[derive(Clone, Default, Deserialize)]
|
||||
#[serde(default)]
|
||||
pub struct Config {
|
||||
pub model_list: Box<[Model]>,
|
||||
#[serde(default)]
|
||||
pub general_settings: GeneralSettings,
|
||||
pub router_settings: RouterSettings,
|
||||
pub litellm_settings: LiteLlmSettings,
|
||||
pub environment_variables: Object,
|
||||
pub callback_settings: Object,
|
||||
pub assistant_settings: Object,
|
||||
pub default_vertex_config: Object,
|
||||
pub mcp_servers: std::collections::BTreeMap<String, McpServer>,
|
||||
pub credential_list: Box<[Object]>,
|
||||
pub guardrails: Box<[Object]>,
|
||||
pub prompts: Box<[Object]>,
|
||||
pub sandbox_tools: Box<[Object]>,
|
||||
pub search_tools: Box<[Object]>,
|
||||
pub files_settings: Box<[Object]>,
|
||||
pub finetune_settings: Box<[Object]>,
|
||||
pub mcp_tools: Box<[Object]>,
|
||||
pub vector_store_registry: Box<[Object]>,
|
||||
pub worker_registry: Box<[Object]>,
|
||||
pub agents: Box<[Object]>,
|
||||
pub agent_list: Box<[Object]>,
|
||||
pub policies: Object,
|
||||
pub policy_attachments: Box<[Object]>,
|
||||
pub include: Box<[String]>,
|
||||
#[serde(flatten)]
|
||||
pub additional_fields: AdditionalFields,
|
||||
}
|
||||
|
||||
#[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 fmt::Debug for Config {
|
||||
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
formatter
|
||||
.debug_struct("Config")
|
||||
.field("model_list", &self.model_list)
|
||||
.field("general_settings", &self.general_settings)
|
||||
.field("router_settings", &self.router_settings)
|
||||
.field("litellm_settings", &self.litellm_settings)
|
||||
.field("environment_variables", &self.environment_variables)
|
||||
.field("callback_settings", &self.callback_settings)
|
||||
.field("assistant_settings", &self.assistant_settings)
|
||||
.field("default_vertex_config", &self.default_vertex_config)
|
||||
.field("mcp_servers", &self.mcp_servers)
|
||||
.field("credential_list", &self.credential_list)
|
||||
.field("guardrails", &self.guardrails)
|
||||
.field("prompts", &self.prompts)
|
||||
.field("sandbox_tools", &self.sandbox_tools)
|
||||
.field("search_tools", &self.search_tools)
|
||||
.field("files_settings", &self.files_settings)
|
||||
.field("finetune_settings", &self.finetune_settings)
|
||||
.field("mcp_tools", &self.mcp_tools)
|
||||
.field("vector_store_registry", &self.vector_store_registry)
|
||||
.field("worker_registry", &self.worker_registry)
|
||||
.field("agents", &self.agents)
|
||||
.field("agent_list", &self.agent_list)
|
||||
.field("policies", &self.policies)
|
||||
.field("policy_attachments", &self.policy_attachments)
|
||||
.field("include", &self.include)
|
||||
.field("additional_fields", &self.additional_fields.keys())
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
impl Config {
|
||||
|
|
@ -43,6 +85,6 @@ impl Config {
|
|||
}
|
||||
|
||||
pub fn load(path: impl AsRef<Path>) -> Result<Self, Error> {
|
||||
Self::from_yaml(&std::fs::read_to_string(path)?)
|
||||
Ok(serde_yaml_ng::from_value(includes::load(path.as_ref())?)?)
|
||||
}
|
||||
}
|
||||
|
|
|
|||
97
litellm-rust/crates/config/src/mcp.rs
Normal file
97
litellm-rust/crates/config/src/mcp.rs
Normal file
|
|
@ -0,0 +1,97 @@
|
|||
use std::{collections::BTreeMap, fmt};
|
||||
|
||||
use litellm_auth_types::SecretValue;
|
||||
use serde::Deserialize;
|
||||
|
||||
use crate::Object;
|
||||
|
||||
#[derive(Clone, Default, Deserialize)]
|
||||
#[serde(default)]
|
||||
pub struct McpServer {
|
||||
pub server_id: Option<String>,
|
||||
pub alias: Option<String>,
|
||||
pub description: Option<String>,
|
||||
pub mcp_info: Object,
|
||||
pub transport: McpTransport,
|
||||
pub url: Option<SecretValue>,
|
||||
pub command: Option<String>,
|
||||
pub args: Box<[String]>,
|
||||
pub env: BTreeMap<String, SecretValue>,
|
||||
pub auth_type: Option<McpAuth>,
|
||||
#[serde(alias = "auth_value")]
|
||||
pub authentication_token: Option<SecretValue>,
|
||||
pub static_headers: BTreeMap<String, SecretValue>,
|
||||
pub upstream_token_header: Option<String>,
|
||||
pub allowed_tools: Option<Box<[String]>>,
|
||||
pub timeout: Option<f64>,
|
||||
pub max_concurrent_requests: Option<usize>,
|
||||
#[serde(flatten)]
|
||||
pub unsupported: Object,
|
||||
}
|
||||
|
||||
impl fmt::Debug for McpServer {
|
||||
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
formatter
|
||||
.debug_struct("McpServer")
|
||||
.field("transport", &self.transport)
|
||||
.field("auth_type", &self.auth_type)
|
||||
.field("timeout", &self.timeout)
|
||||
.field("max_concurrent_requests", &self.max_concurrent_requests)
|
||||
.finish_non_exhaustive()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, Default, Deserialize, PartialEq, Eq)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum McpTransport {
|
||||
#[default]
|
||||
Http,
|
||||
Sse,
|
||||
Stdio,
|
||||
}
|
||||
|
||||
impl McpTransport {
|
||||
pub fn as_str(self) -> &'static str {
|
||||
match self {
|
||||
Self::Http => "http",
|
||||
Self::Sse => "sse",
|
||||
Self::Stdio => "stdio",
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, Deserialize, PartialEq, Eq)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum McpAuth {
|
||||
None,
|
||||
ApiKey,
|
||||
BearerToken,
|
||||
Basic,
|
||||
Authorization,
|
||||
Token,
|
||||
Oauth2,
|
||||
AwsSigv4,
|
||||
Oauth2TokenExchange,
|
||||
Oauth2IdJag,
|
||||
TruePassthrough,
|
||||
OauthDelegate,
|
||||
}
|
||||
|
||||
impl McpAuth {
|
||||
pub fn as_str(self) -> &'static str {
|
||||
match self {
|
||||
Self::None => "none",
|
||||
Self::ApiKey => "api_key",
|
||||
Self::BearerToken => "bearer_token",
|
||||
Self::Basic => "basic",
|
||||
Self::Authorization => "authorization",
|
||||
Self::Token => "token",
|
||||
Self::Oauth2 => "oauth2",
|
||||
Self::AwsSigv4 => "aws_sigv4",
|
||||
Self::Oauth2TokenExchange => "oauth2_token_exchange",
|
||||
Self::Oauth2IdJag => "oauth2_id_jag",
|
||||
Self::TruePassthrough => "true_passthrough",
|
||||
Self::OauthDelegate => "oauth_delegate",
|
||||
}
|
||||
}
|
||||
}
|
||||
95
litellm-rust/crates/config/src/model.rs
Normal file
95
litellm-rust/crates/config/src/model.rs
Normal file
|
|
@ -0,0 +1,95 @@
|
|||
use std::fmt;
|
||||
|
||||
use litellm_auth_types::SecretValue;
|
||||
use serde::Deserialize;
|
||||
|
||||
use crate::{AdditionalFields, Flag, NumberOrString, Object};
|
||||
|
||||
#[derive(Clone, Deserialize)]
|
||||
pub struct Model {
|
||||
pub model_name: String,
|
||||
pub litellm_params: LiteLlmParams,
|
||||
#[serde(default)]
|
||||
pub model_info: Object,
|
||||
pub blocked: Option<bool>,
|
||||
#[serde(flatten)]
|
||||
pub additional_fields: AdditionalFields,
|
||||
}
|
||||
|
||||
impl fmt::Debug for Model {
|
||||
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
formatter
|
||||
.debug_struct("Model")
|
||||
.field("model_name", &self.model_name)
|
||||
.field("litellm_params", &self.litellm_params)
|
||||
.field("model_info", &self.model_info)
|
||||
.field("blocked", &self.blocked)
|
||||
.field("additional_fields", &self.additional_fields.keys())
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Deserialize)]
|
||||
pub struct LiteLlmParams {
|
||||
pub model: String,
|
||||
pub api_key: Option<SecretValue>,
|
||||
pub api_base: Option<String>,
|
||||
pub api_version: Option<String>,
|
||||
pub custom_llm_provider: Option<String>,
|
||||
pub timeout: Option<NumberOrString>,
|
||||
pub stream_timeout: Option<NumberOrString>,
|
||||
pub max_retries: Option<NumberOrString>,
|
||||
pub tpm: Option<NumberOrString>,
|
||||
pub rpm: Option<NumberOrString>,
|
||||
pub itpm: Option<NumberOrString>,
|
||||
pub otpm: Option<NumberOrString>,
|
||||
pub max_parallel_requests: Option<u64>,
|
||||
pub organization: Option<serde_yaml_ng::Value>,
|
||||
pub drop_params: Option<Flag>,
|
||||
pub tags: Option<Box<[String]>>,
|
||||
pub tag_regex: Option<Box<[String]>>,
|
||||
pub max_budget: Option<f64>,
|
||||
pub budget_duration: Option<String>,
|
||||
pub default_api_key_tpm_limit: Option<u64>,
|
||||
pub default_api_key_rpm_limit: Option<u64>,
|
||||
pub use_in_pass_through: Option<bool>,
|
||||
pub use_chat_completions_api: Option<bool>,
|
||||
pub litellm_credential_name: Option<String>,
|
||||
pub provider_affinity_header: Option<String>,
|
||||
#[serde(flatten)]
|
||||
pub additional_fields: AdditionalFields,
|
||||
}
|
||||
|
||||
impl fmt::Debug for LiteLlmParams {
|
||||
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
formatter
|
||||
.debug_struct("LiteLlmParams")
|
||||
.field("model", &self.model)
|
||||
.field("api_key", &self.api_key)
|
||||
.field("api_base", &self.api_base)
|
||||
.field("api_version", &self.api_version)
|
||||
.field("custom_llm_provider", &self.custom_llm_provider)
|
||||
.field("timeout", &self.timeout)
|
||||
.field("stream_timeout", &self.stream_timeout)
|
||||
.field("max_retries", &self.max_retries)
|
||||
.field("tpm", &self.tpm)
|
||||
.field("rpm", &self.rpm)
|
||||
.field("itpm", &self.itpm)
|
||||
.field("otpm", &self.otpm)
|
||||
.field("max_parallel_requests", &self.max_parallel_requests)
|
||||
.field("organization", &self.organization)
|
||||
.field("drop_params", &self.drop_params)
|
||||
.field("tags", &self.tags)
|
||||
.field("tag_regex", &self.tag_regex)
|
||||
.field("max_budget", &self.max_budget)
|
||||
.field("budget_duration", &self.budget_duration)
|
||||
.field("default_api_key_tpm_limit", &self.default_api_key_tpm_limit)
|
||||
.field("default_api_key_rpm_limit", &self.default_api_key_rpm_limit)
|
||||
.field("use_in_pass_through", &self.use_in_pass_through)
|
||||
.field("use_chat_completions_api", &self.use_chat_completions_api)
|
||||
.field("litellm_credential_name", &self.litellm_credential_name)
|
||||
.field("provider_affinity_header", &self.provider_affinity_header)
|
||||
.field("additional_fields", &self.additional_fields.keys())
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
221
litellm-rust/crates/config/src/settings.rs
Normal file
221
litellm-rust/crates/config/src/settings.rs
Normal file
|
|
@ -0,0 +1,221 @@
|
|||
use std::fmt;
|
||||
|
||||
use litellm_auth_types::SecretValue;
|
||||
use serde::Deserialize;
|
||||
|
||||
use crate::{AdditionalFields, Flag, NumberOrString, Object, OneOrMany, Value};
|
||||
|
||||
#[derive(Clone, Deserialize)]
|
||||
#[serde(default)]
|
||||
pub struct GeneralSettings {
|
||||
pub completion_model: Option<String>,
|
||||
pub max_in_flight_requests_per_worker: Option<u64>,
|
||||
pub max_queued_requests_per_worker: Option<u64>,
|
||||
pub admission_queue_timeout_seconds: f64,
|
||||
pub master_key: Option<SecretValue>,
|
||||
pub database_url: Option<SecretValue>,
|
||||
pub database_connection_pool_limit: Option<u64>,
|
||||
pub database_connection_timeout: Option<f64>,
|
||||
pub database_connect_timeout: Option<f64>,
|
||||
pub database_socket_timeout: Option<f64>,
|
||||
pub database_max_idle_connection_lifetime: Option<f64>,
|
||||
pub max_parallel_requests: Option<u64>,
|
||||
pub global_max_parallel_requests: Option<u64>,
|
||||
pub max_request_size_mb: Option<u64>,
|
||||
pub max_response_size_mb: Option<u64>,
|
||||
pub proxy_config_reload_interval_seconds: u64,
|
||||
pub background_health_checks: Option<bool>,
|
||||
pub health_check_interval: u64,
|
||||
pub health_check_concurrency: Option<u64>,
|
||||
pub store_model_in_db: Option<bool>,
|
||||
pub forward_client_headers_to_llm_api: Option<bool>,
|
||||
pub cancel_on_disconnect: Option<bool>,
|
||||
pub infer_model_from_keys: Option<bool>,
|
||||
pub enable_public_model_hub: bool,
|
||||
pub dangerously_permit_weak_or_unset_master_key: Option<bool>,
|
||||
pub plugins: Option<Box<[Object]>>,
|
||||
pub coordination_redis: Option<Object>,
|
||||
pub mcp_allowed_hosts: Option<Box<[String]>>,
|
||||
pub mcp_allowed_origins: Box<[String]>,
|
||||
#[serde(flatten)]
|
||||
pub additional_fields: AdditionalFields,
|
||||
}
|
||||
|
||||
impl Default for GeneralSettings {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
completion_model: None,
|
||||
max_in_flight_requests_per_worker: None,
|
||||
max_queued_requests_per_worker: None,
|
||||
admission_queue_timeout_seconds: 1.0,
|
||||
master_key: None,
|
||||
database_url: None,
|
||||
database_connection_pool_limit: Some(10),
|
||||
database_connection_timeout: Some(60.0),
|
||||
database_connect_timeout: None,
|
||||
database_socket_timeout: None,
|
||||
database_max_idle_connection_lifetime: Some(60.0),
|
||||
max_parallel_requests: None,
|
||||
global_max_parallel_requests: None,
|
||||
max_request_size_mb: None,
|
||||
max_response_size_mb: None,
|
||||
proxy_config_reload_interval_seconds: 30,
|
||||
background_health_checks: None,
|
||||
health_check_interval: 300,
|
||||
health_check_concurrency: None,
|
||||
store_model_in_db: None,
|
||||
forward_client_headers_to_llm_api: None,
|
||||
cancel_on_disconnect: None,
|
||||
infer_model_from_keys: None,
|
||||
enable_public_model_hub: false,
|
||||
dangerously_permit_weak_or_unset_master_key: None,
|
||||
plugins: None,
|
||||
coordination_redis: None,
|
||||
mcp_allowed_hosts: None,
|
||||
mcp_allowed_origins: Box::default(),
|
||||
additional_fields: AdditionalFields::new(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl fmt::Debug for GeneralSettings {
|
||||
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
formatter
|
||||
.debug_struct("GeneralSettings")
|
||||
.field("completion_model", &self.completion_model)
|
||||
.field(
|
||||
"max_in_flight_requests_per_worker",
|
||||
&self.max_in_flight_requests_per_worker,
|
||||
)
|
||||
.field(
|
||||
"max_queued_requests_per_worker",
|
||||
&self.max_queued_requests_per_worker,
|
||||
)
|
||||
.field(
|
||||
"admission_queue_timeout_seconds",
|
||||
&self.admission_queue_timeout_seconds,
|
||||
)
|
||||
.field("master_key", &self.master_key)
|
||||
.field("database_url", &self.database_url)
|
||||
.field("store_model_in_db", &self.store_model_in_db)
|
||||
.field("additional_fields", &self.additional_fields.keys())
|
||||
.finish_non_exhaustive()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Default, Deserialize)]
|
||||
#[serde(default)]
|
||||
pub struct RouterSettings {
|
||||
pub routing_strategy: Option<String>,
|
||||
pub routing_strategy_args: Option<Object>,
|
||||
pub routing_groups: Option<Box<[Object]>>,
|
||||
pub retry_policy: Option<Object>,
|
||||
pub model_group_retry_policy: Option<Object>,
|
||||
pub model_group_affinity_config: Option<Object>,
|
||||
pub allowed_fails: Option<u64>,
|
||||
pub cooldown_time: Option<f64>,
|
||||
pub num_retries: Option<u64>,
|
||||
pub timeout: Option<f64>,
|
||||
pub max_retries: Option<u64>,
|
||||
pub retry_after: Option<f64>,
|
||||
pub fallbacks: Option<Box<[Object]>>,
|
||||
pub context_window_fallbacks: Option<Box<[Object]>>,
|
||||
pub model_group_alias: Option<Object>,
|
||||
pub enable_tag_filtering: Option<bool>,
|
||||
pub weights: Option<Object>,
|
||||
pub tag_routing_prefix: Option<String>,
|
||||
pub optional_pre_call_checks: Option<Box<[String]>>,
|
||||
#[serde(flatten)]
|
||||
pub additional_fields: AdditionalFields,
|
||||
}
|
||||
|
||||
impl fmt::Debug for RouterSettings {
|
||||
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
formatter
|
||||
.debug_struct("RouterSettings")
|
||||
.field("routing_strategy", &self.routing_strategy)
|
||||
.field("routing_strategy_args", &self.routing_strategy_args)
|
||||
.field("routing_groups", &self.routing_groups)
|
||||
.field("retry_policy", &self.retry_policy)
|
||||
.field("model_group_retry_policy", &self.model_group_retry_policy)
|
||||
.field(
|
||||
"model_group_affinity_config",
|
||||
&self.model_group_affinity_config,
|
||||
)
|
||||
.field("allowed_fails", &self.allowed_fails)
|
||||
.field("cooldown_time", &self.cooldown_time)
|
||||
.field("num_retries", &self.num_retries)
|
||||
.field("timeout", &self.timeout)
|
||||
.field("max_retries", &self.max_retries)
|
||||
.field("retry_after", &self.retry_after)
|
||||
.field("fallbacks", &self.fallbacks)
|
||||
.field("context_window_fallbacks", &self.context_window_fallbacks)
|
||||
.field("model_group_alias", &self.model_group_alias)
|
||||
.field("enable_tag_filtering", &self.enable_tag_filtering)
|
||||
.field("weights", &self.weights)
|
||||
.field("tag_routing_prefix", &self.tag_routing_prefix)
|
||||
.field("optional_pre_call_checks", &self.optional_pre_call_checks)
|
||||
.field("additional_fields", &self.additional_fields.keys())
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Default, Deserialize)]
|
||||
#[serde(default)]
|
||||
pub struct LiteLlmSettings {
|
||||
pub ssl_verify: Option<Flag>,
|
||||
pub ssl_certificate: Option<String>,
|
||||
pub ssl_security_level: Option<String>,
|
||||
pub ssl_ecdh_curve: Option<String>,
|
||||
pub force_ipv4: Option<bool>,
|
||||
pub http2: Option<bool>,
|
||||
pub aiohttp_trust_env: Option<bool>,
|
||||
pub disable_aiohttp_trust_env: Option<bool>,
|
||||
pub disable_aiohttp_transport: Option<bool>,
|
||||
pub drop_params: Option<Flag>,
|
||||
pub request_timeout: Option<NumberOrString>,
|
||||
pub num_retries: Option<u64>,
|
||||
pub cache: Option<bool>,
|
||||
pub cache_params: Option<Object>,
|
||||
pub callbacks: Option<OneOrMany<Value>>,
|
||||
pub success_callback: Option<OneOrMany<Value>>,
|
||||
pub failure_callback: Option<OneOrMany<Value>>,
|
||||
pub json_logs: Option<bool>,
|
||||
pub set_verbose: Option<bool>,
|
||||
#[serde(flatten)]
|
||||
pub additional_fields: AdditionalFields,
|
||||
}
|
||||
|
||||
impl fmt::Debug for LiteLlmSettings {
|
||||
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
formatter
|
||||
.debug_struct("LiteLlmSettings")
|
||||
.field("drop_params", &self.drop_params)
|
||||
.field("request_timeout", &self.request_timeout)
|
||||
.field("num_retries", &self.num_retries)
|
||||
.field("cache", &self.cache)
|
||||
.field("cache_params", &self.cache_params)
|
||||
.field(
|
||||
"callbacks",
|
||||
&self.callbacks.as_ref().map(|callbacks| callbacks.len()),
|
||||
)
|
||||
.field(
|
||||
"success_callback",
|
||||
&self
|
||||
.success_callback
|
||||
.as_ref()
|
||||
.map(|callbacks| callbacks.len()),
|
||||
)
|
||||
.field(
|
||||
"failure_callback",
|
||||
&self
|
||||
.failure_callback
|
||||
.as_ref()
|
||||
.map(|callbacks| callbacks.len()),
|
||||
)
|
||||
.field("json_logs", &self.json_logs)
|
||||
.field("set_verbose", &self.set_verbose)
|
||||
.field("additional_fields", &self.additional_fields.keys())
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
78
litellm-rust/crates/config/src/value.rs
Normal file
78
litellm-rust/crates/config/src/value.rs
Normal file
|
|
@ -0,0 +1,78 @@
|
|||
use std::{collections::BTreeMap, fmt, ops::Deref};
|
||||
|
||||
use serde::Deserialize;
|
||||
|
||||
pub type Value = serde_yaml_ng::Value;
|
||||
pub type AdditionalFields = BTreeMap<String, Value>;
|
||||
|
||||
#[derive(Clone, Default, Deserialize)]
|
||||
#[serde(transparent)]
|
||||
pub struct Object(BTreeMap<String, Value>);
|
||||
|
||||
impl Object {
|
||||
pub fn get(&self, key: &str) -> Option<&Value> {
|
||||
self.0.get(key)
|
||||
}
|
||||
|
||||
pub fn is_empty(&self) -> bool {
|
||||
self.0.is_empty()
|
||||
}
|
||||
|
||||
pub fn len(&self) -> usize {
|
||||
self.0.len()
|
||||
}
|
||||
}
|
||||
|
||||
impl Deref for Object {
|
||||
type Target = BTreeMap<String, Value>;
|
||||
|
||||
fn deref(&self) -> &Self::Target {
|
||||
&self.0
|
||||
}
|
||||
}
|
||||
|
||||
impl fmt::Debug for Object {
|
||||
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
formatter
|
||||
.debug_struct("Object")
|
||||
.field("keys", &self.0.keys())
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Deserialize, PartialEq)]
|
||||
#[serde(untagged)]
|
||||
pub enum NumberOrString {
|
||||
Number(f64),
|
||||
String(String),
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Deserialize, PartialEq, Eq)]
|
||||
#[serde(untagged)]
|
||||
pub enum Flag {
|
||||
Boolean(bool),
|
||||
String(String),
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Deserialize)]
|
||||
#[serde(untagged)]
|
||||
pub enum OneOrMany<T> {
|
||||
Many(Box<[T]>),
|
||||
One(T),
|
||||
}
|
||||
|
||||
impl<T> OneOrMany<T> {
|
||||
pub fn len(&self) -> usize {
|
||||
match self {
|
||||
Self::Many(values) => values.len(),
|
||||
Self::One(_) => 1,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn is_empty(&self) -> bool {
|
||||
match self {
|
||||
Self::Many(values) => values.is_empty(),
|
||||
Self::One(_) => false,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -1,4 +1,4 @@
|
|||
use litellm_config::{Config, Error};
|
||||
use litellm_config::{Config, Error, Flag, NumberOrString};
|
||||
use rstest::{fixture, rstest};
|
||||
use tempfile::TempDir;
|
||||
|
||||
|
|
@ -73,14 +73,9 @@ fn config_debug_redacts_api_keys() {
|
|||
|
||||
#[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) {
|
||||
fn rejects_malformed_and_incomplete_config(#[case] yaml: &str) {
|
||||
assert!(matches!(Config::from_yaml(yaml), Err(Error::Parse(_))));
|
||||
}
|
||||
|
||||
|
|
@ -117,3 +112,264 @@ fn missing_general_settings_has_no_master_key() {
|
|||
let config = Config::from_yaml("model_list: []").unwrap();
|
||||
assert!(config.general_settings.master_key.is_none());
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn empty_config_matches_python_defaults() {
|
||||
let config = Config::from_yaml("{}").unwrap();
|
||||
assert!(config.model_list.is_empty());
|
||||
assert_eq!(config.general_settings.admission_queue_timeout_seconds, 1.0);
|
||||
assert_eq!(
|
||||
config.general_settings.database_connection_pool_limit,
|
||||
Some(10)
|
||||
);
|
||||
assert_eq!(
|
||||
config.general_settings.proxy_config_reload_interval_seconds,
|
||||
30
|
||||
);
|
||||
assert_eq!(config.general_settings.health_check_interval, 300);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn parses_typed_settings_and_preserves_extension_fields() {
|
||||
let config = Config::from_yaml(
|
||||
r#"
|
||||
model_list:
|
||||
- model_name: assistant
|
||||
litellm_params:
|
||||
model: vertex_ai/test-model
|
||||
timeout: os.environ/REQUEST_TIMEOUT
|
||||
tpm: os.environ/TPM_LIMIT
|
||||
rpm: 5
|
||||
drop_params: "true"
|
||||
vertex_project: test-project
|
||||
model_info:
|
||||
mode: chat
|
||||
access_groups: [internal]
|
||||
general_settings:
|
||||
master_key: secret-master-key
|
||||
store_model_in_db: true
|
||||
custom_auth: auth.py
|
||||
router_settings:
|
||||
routing_strategy: simple-shuffle
|
||||
allowed_fails: 2
|
||||
redis_host: cache.internal
|
||||
litellm_settings:
|
||||
drop_params: true
|
||||
cache: true
|
||||
custom_callback_name: audit
|
||||
future_section:
|
||||
enabled: true
|
||||
"#,
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
let model = &config.model_list[0];
|
||||
assert_eq!(
|
||||
model.litellm_params.timeout,
|
||||
Some(NumberOrString::String(
|
||||
"os.environ/REQUEST_TIMEOUT".to_string()
|
||||
))
|
||||
);
|
||||
assert_eq!(
|
||||
model.litellm_params.tpm,
|
||||
Some(NumberOrString::String("os.environ/TPM_LIMIT".to_string()))
|
||||
);
|
||||
assert_eq!(model.litellm_params.rpm, Some(NumberOrString::Number(5.0)));
|
||||
assert_eq!(
|
||||
model.litellm_params.drop_params,
|
||||
Some(Flag::String("true".to_string()))
|
||||
);
|
||||
assert!(
|
||||
model
|
||||
.litellm_params
|
||||
.additional_fields
|
||||
.contains_key("vertex_project")
|
||||
);
|
||||
assert!(model.additional_fields.contains_key("access_groups"));
|
||||
assert_eq!(config.router_settings.allowed_fails, Some(2));
|
||||
assert!(
|
||||
config
|
||||
.router_settings
|
||||
.additional_fields
|
||||
.contains_key("redis_host")
|
||||
);
|
||||
assert!(
|
||||
config
|
||||
.litellm_settings
|
||||
.additional_fields
|
||||
.contains_key("custom_callback_name")
|
||||
);
|
||||
assert!(config.additional_fields.contains_key("future_section"));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn parses_python_config_sections() {
|
||||
let config = Config::from_yaml(
|
||||
r#"
|
||||
environment_variables:
|
||||
REDIS_PORT: 6379
|
||||
callback_settings:
|
||||
otel:
|
||||
message_logging: false
|
||||
assistant_settings:
|
||||
custom_llm_provider: openai
|
||||
credential_list:
|
||||
- credential_name: bedrock
|
||||
credential_values:
|
||||
aws_region_name: us-east-1
|
||||
guardrails:
|
||||
- guardrail_name: pii
|
||||
litellm_params:
|
||||
guardrail: presidio
|
||||
prompts:
|
||||
- prompt_id: support
|
||||
sandbox_tools:
|
||||
- sandbox_tool_name: e2b
|
||||
search_tools:
|
||||
- search_tool_name: web
|
||||
files_settings:
|
||||
- custom_llm_provider: openai
|
||||
finetune_settings:
|
||||
- custom_llm_provider: openai
|
||||
mcp_tools:
|
||||
- name: lookup
|
||||
mcp_servers:
|
||||
docs:
|
||||
url: https://example.test/mcp
|
||||
vector_store_registry:
|
||||
- vector_store_name: docs
|
||||
worker_registry:
|
||||
- worker_id: regional
|
||||
agents:
|
||||
- agent_name: reviewer
|
||||
agent_list:
|
||||
- agent_name: legacy-reviewer
|
||||
policies:
|
||||
safe:
|
||||
guardrails:
|
||||
add: [pii]
|
||||
policy_attachments:
|
||||
- policy_id: safe
|
||||
include:
|
||||
- models.yaml
|
||||
"#,
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(config.environment_variables.len(), 1);
|
||||
assert_eq!(config.credential_list.len(), 1);
|
||||
assert_eq!(config.guardrails.len(), 1);
|
||||
assert_eq!(config.prompts.len(), 1);
|
||||
assert_eq!(config.sandbox_tools.len(), 1);
|
||||
assert_eq!(config.search_tools.len(), 1);
|
||||
assert_eq!(config.files_settings.len(), 1);
|
||||
assert_eq!(config.finetune_settings.len(), 1);
|
||||
assert_eq!(config.mcp_tools.len(), 1);
|
||||
assert_eq!(config.mcp_servers.len(), 1);
|
||||
assert_eq!(config.vector_store_registry.len(), 1);
|
||||
assert_eq!(config.worker_registry.len(), 1);
|
||||
assert_eq!(config.agents.len(), 1);
|
||||
assert_eq!(config.agent_list.len(), 1);
|
||||
assert!(config.policies.contains_key("safe"));
|
||||
assert_eq!(config.policy_attachments.len(), 1);
|
||||
assert_eq!(&*config.include, &["models.yaml"]);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::policy_pipeline("../../../litellm/proxy/example_config_yaml/test_pipeline_config.yaml")]
|
||||
#[case::gateway("../../../tests/e2e/gateway/stage_mirror_ci_config.yml")]
|
||||
fn parses_representative_python_configs(#[case] relative_path: &str) {
|
||||
let path = std::path::Path::new(env!("CARGO_MANIFEST_DIR")).join(relative_path);
|
||||
let config = Config::load(path).unwrap();
|
||||
assert!(!config.model_list.is_empty());
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::one("callbacks: custom_callbacks.logger", 1)]
|
||||
#[case::many("callbacks: [prometheus, otel]", 2)]
|
||||
fn accepts_python_callback_shorthand(#[case] setting: &str, #[case] expected_len: usize) {
|
||||
let config = Config::from_yaml(&format!("litellm_settings:\n {setting}")).unwrap();
|
||||
assert_eq!(
|
||||
config.litellm_settings.callbacks.as_ref().unwrap().len(),
|
||||
expected_len
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::api_key("api_key", "provider-secret")]
|
||||
#[case::provider_extension("aws_secret_access_key", "aws-secret")]
|
||||
#[case::general_extension("custom_auth_secret", "auth-secret")]
|
||||
#[case::router_extension("redis_password", "redis-secret")]
|
||||
#[case::litellm_extension("callback_token", "callback-secret")]
|
||||
#[case::root_extension("private_token", "root-secret")]
|
||||
fn debug_output_does_not_expose_config_values(#[case] field: &str, #[case] secret: &str) {
|
||||
let yaml = match field {
|
||||
"api_key" => format!(
|
||||
"model_list: [{{model_name: assistant, litellm_params: {{model: test, api_key: {secret}}}}}]"
|
||||
),
|
||||
"aws_secret_access_key" => format!(
|
||||
"model_list: [{{model_name: assistant, litellm_params: {{model: test, aws_secret_access_key: {secret}}}}}]"
|
||||
),
|
||||
"custom_auth_secret" => format!("general_settings: {{{field}: {secret}}}"),
|
||||
"redis_password" => format!("router_settings: {{{field}: {secret}}}"),
|
||||
"callback_token" => format!("litellm_settings: {{{field}: {secret}}}"),
|
||||
"private_token" => format!("{field}: {secret}"),
|
||||
_ => unreachable!(),
|
||||
};
|
||||
let config = Config::from_yaml(&yaml).unwrap();
|
||||
assert!(!format!("{config:?}").contains(secret));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn resolves_nested_includes_once_in_breadth_first_order() {
|
||||
let directory = TempDir::new().unwrap();
|
||||
let root = directory.path();
|
||||
std::fs::create_dir(root.join("nested")).unwrap();
|
||||
std::fs::write(root.join("config.yaml"), "include: [nested/first.yaml, second.yaml]\nmodel_list: [{model_name: root, litellm_params: {model: root}}]\n").unwrap();
|
||||
std::fs::write(root.join("nested/first.yaml"), "include: [third.yaml]\nmodel_list: [{model_name: first, litellm_params: {model: first}}]\n").unwrap();
|
||||
std::fs::write(root.join("second.yaml"), "general_settings: {master_key: second}\nmodel_list: [{model_name: second, litellm_params: {model: second}}]\n").unwrap();
|
||||
std::fs::write(root.join("nested/third.yaml"), "include: [../config.yaml]\ngeneral_settings: {master_key: third}\nmodel_list: [{model_name: third, litellm_params: {model: third}}]\n").unwrap();
|
||||
|
||||
let config = Config::load(root.join("config.yaml")).unwrap();
|
||||
assert_eq!(
|
||||
config
|
||||
.model_list
|
||||
.iter()
|
||||
.map(|model| model.model_name.as_str())
|
||||
.collect::<Vec<_>>(),
|
||||
["root", "first", "second", "third"]
|
||||
);
|
||||
assert_eq!(
|
||||
config.general_settings.master_key.unwrap().expose(),
|
||||
"third"
|
||||
);
|
||||
assert!(config.include.is_empty());
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn mcp_config_redacts_nested_credentials_and_preserves_policy_for_validation() {
|
||||
let config = Config::from_yaml("mcp_servers:\n docs:\n url: https://example.test/private-secret/mcp\n authentication_token: upstream-secret\n static_headers: {x-token: header-secret}\n env: {TOKEN: env-secret}\n args: [argument-secret]\n client_secret: oauth-secret\n allowed_tools: [search]\n").unwrap();
|
||||
let server = &config.mcp_servers["docs"];
|
||||
assert_eq!(server.allowed_tools.as_deref().unwrap(), ["search"]);
|
||||
assert!(server.unsupported.contains_key("client_secret"));
|
||||
let debug = format!("{config:?}");
|
||||
for secret in [
|
||||
"private-secret",
|
||||
"upstream-secret",
|
||||
"header-secret",
|
||||
"env-secret",
|
||||
"oauth-secret",
|
||||
"argument-secret",
|
||||
] {
|
||||
assert!(!debug.contains(secret));
|
||||
}
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::transport("transport: invalid")]
|
||||
#[case::auth("auth_type: invalid")]
|
||||
#[case::concurrency("max_concurrent_requests: -1")]
|
||||
#[case::headers("static_headers: {x-token: [not, a, string]}")]
|
||||
fn rejects_invalid_typed_mcp_settings(#[case] setting: &str) {
|
||||
assert!(Config::from_yaml(&format!("mcp_servers:\n docs:\n {setting}\n")).is_err());
|
||||
}
|
||||
|
|
|
|||
|
|
@ -109,6 +109,7 @@ impl IntoIterator for CallArguments {
|
|||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use rstest::rstest;
|
||||
use serde_json::json;
|
||||
|
||||
use super::*;
|
||||
|
|
@ -140,22 +141,31 @@ mod tests {
|
|||
assert_eq!(serde_json::to_value(arguments).unwrap(), original);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn invalid_extra_body_is_rejected_without_coercing_it_to_empty() {
|
||||
for value in [json!(false), json!(0), json!([]), json!("")] {
|
||||
let arguments = serde_json::from_value(json!({"extra_body":value})).unwrap();
|
||||
assert_eq!(
|
||||
compose_body(&arguments, &json!({}), &[]),
|
||||
Err(crate::params::Error::ExtraBody)
|
||||
);
|
||||
}
|
||||
let arguments = serde_json::from_value(json!({"extra_body":null})).unwrap();
|
||||
#[rstest]
|
||||
#[case::boolean(json!(false))]
|
||||
#[case::number(json!(0))]
|
||||
#[case::array(json!([]))]
|
||||
#[case::string(json!(""))]
|
||||
fn invalid_extra_body_is_rejected_without_coercing_it_to_empty(
|
||||
#[case] value: serde_json::Value,
|
||||
) {
|
||||
let arguments = serde_json::from_value(json!({"extra_body":value})).unwrap();
|
||||
assert_eq!(
|
||||
compose_body(&arguments, &json!({}), &[]).unwrap(),
|
||||
json!({})
|
||||
compose_body(&arguments, &json!({}), &[]),
|
||||
Err(crate::params::Error::ExtraBody)
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::null(json!(null), json!({}))]
|
||||
fn null_extra_body_is_coerced_to_empty_object(
|
||||
#[case] value: serde_json::Value,
|
||||
#[case] expected: serde_json::Value,
|
||||
) {
|
||||
let arguments = serde_json::from_value(json!({"extra_body":value})).unwrap();
|
||||
assert_eq!(compose_body(&arguments, &json!({}), &[]).unwrap(), expected);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn typed_views_preserve_missing_and_explicit_null_in_the_source() {
|
||||
#[derive(Deserialize)]
|
||||
|
|
|
|||
|
|
@ -68,27 +68,25 @@ pub fn json_type_name(value: &serde_json::Value) -> &'static str {
|
|||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use rstest::rstest;
|
||||
|
||||
#[test]
|
||||
fn maps_every_reason_the_route_can_observe() {
|
||||
assert_eq!(finish_reason_for("end_turn"), "stop");
|
||||
assert_eq!(finish_reason_for("stop_sequence"), "stop");
|
||||
assert_eq!(finish_reason_for("max_tokens"), "length");
|
||||
assert_eq!(finish_reason_for("refusal"), "content_filter");
|
||||
assert_eq!(finish_reason_for("guardrail_intervened"), "content_filter");
|
||||
// Converse emits these two, and folding them into `stop` would report a
|
||||
// filtered completion as a normal one.
|
||||
assert_eq!(finish_reason_for("content_filtered"), "content_filter");
|
||||
assert_eq!(finish_reason_for("content_filter"), "content_filter");
|
||||
#[rstest]
|
||||
#[case::end_turn("end_turn", "stop")]
|
||||
#[case::stop_sequence("stop_sequence", "stop")]
|
||||
#[case::max_tokens("max_tokens", "length")]
|
||||
#[case::refusal("refusal", "content_filter")]
|
||||
#[case::guardrail_intervened("guardrail_intervened", "content_filter")]
|
||||
#[case::content_filtered("content_filtered", "content_filter")]
|
||||
#[case::content_filter("content_filter", "content_filter")]
|
||||
fn maps_every_reason_the_route_can_observe(#[case] reason: &str, #[case] expected: &str) {
|
||||
assert_eq!(finish_reason_for(reason), expected);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn defaults_an_unmapped_reason_to_stop_like_python() {
|
||||
// Python warns and falls back to `stop` for a reason its own map does
|
||||
// not carry, so only a reason absent from `_FINISH_REASON_MAP` belongs
|
||||
// here.
|
||||
assert_eq!(finish_reason_for("something_new"), "stop");
|
||||
assert_eq!(finish_reason_for(""), "stop");
|
||||
#[rstest]
|
||||
#[case::unknown("something_new")]
|
||||
#[case::empty("")]
|
||||
fn defaults_an_unmapped_reason_to_stop_like_python(#[case] reason: &str) {
|
||||
assert_eq!(finish_reason_for(reason), "stop");
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
|
|
|||
|
|
@ -1,9 +1,26 @@
|
|||
use strum::{EnumString, IntoStaticStr};
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub struct CustomLlmProvider<'a> {
|
||||
pub model: &'a str,
|
||||
pub custom_llm_provider: &'a str,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, EnumString, IntoStaticStr)]
|
||||
#[strum(serialize_all = "snake_case")]
|
||||
pub enum LlmProviders {
|
||||
Anthropic,
|
||||
AwsTextract,
|
||||
AzureAi,
|
||||
Bedrock,
|
||||
Cohere,
|
||||
Mistral,
|
||||
Openai,
|
||||
OpenaiLike,
|
||||
Reducto,
|
||||
VertexAi,
|
||||
}
|
||||
|
||||
pub fn get_custom_llm_provider<'a>(
|
||||
model: &'a str,
|
||||
custom_llm_provider: Option<&'a str>,
|
||||
|
|
|
|||
|
|
@ -129,6 +129,7 @@ fn integral_float(value: f64) -> Option<i64> {
|
|||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use rstest::rstest;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use serde_json::json;
|
||||
use serde_with::serde_as;
|
||||
|
|
@ -144,20 +145,20 @@ mod tests {
|
|||
float: Option<f64>,
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn boolean_tokens_follow_python_string_trimming_without_redis_tokens() {
|
||||
for (input, expected) in [
|
||||
(" True ", Some(true)),
|
||||
("\u{1c}TRUE\u{1f}", Some(true)),
|
||||
("\u{a0}False\u{2003}", Some(false)),
|
||||
("true\u{200b}", None),
|
||||
("yes", None),
|
||||
("1", None),
|
||||
("", None),
|
||||
("unknown", None),
|
||||
] {
|
||||
assert_eq!(parse_str_bool(input), expected, "{input:?}");
|
||||
}
|
||||
#[rstest]
|
||||
#[case::trimmed_true(" True ", Some(true))]
|
||||
#[case::control_whitespace_true("\u{1c}TRUE\u{1f}", Some(true))]
|
||||
#[case::unicode_whitespace_false("\u{a0}False\u{2003}", Some(false))]
|
||||
#[case::zero_width_space("true\u{200b}", None)]
|
||||
#[case::yes("yes", None)]
|
||||
#[case::one("1", None)]
|
||||
#[case::empty("", None)]
|
||||
#[case::unknown("unknown", None)]
|
||||
fn boolean_tokens_follow_python_string_trimming_without_redis_tokens(
|
||||
#[case] input: &str,
|
||||
#[case] expected: Option<bool>,
|
||||
) {
|
||||
assert_eq!(parse_str_bool(input), expected, "{input:?}");
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
|
|
|||
|
|
@ -37,6 +37,22 @@ impl Lookup for ProcessEnvironment {
|
|||
}
|
||||
}
|
||||
|
||||
pub fn resolve_non_empty(
|
||||
value: Option<&str>,
|
||||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
names: &[&str],
|
||||
) -> Option<String> {
|
||||
value
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(str::to_string)
|
||||
.or_else(|| {
|
||||
names
|
||||
.iter()
|
||||
.find_map(|name| env_lookup(name).filter(|value| !value.trim().is_empty()))
|
||||
})
|
||||
}
|
||||
|
||||
pub trait Layer: Default {
|
||||
fn or(self, lower: Self) -> Self;
|
||||
}
|
||||
|
|
|
|||
|
|
@ -72,20 +72,18 @@ impl ApiUrl<Complete> {
|
|||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use rstest::rstest;
|
||||
|
||||
#[test]
|
||||
fn completion_appends_only_the_missing_path_suffix() {
|
||||
for (base, expected) in [
|
||||
("https://example.test", "https://example.test/v1/ocr"),
|
||||
("https://example.test/v1", "https://example.test/v1/ocr"),
|
||||
("https://example.test/v1/ocr", "https://example.test/v1/ocr"),
|
||||
] {
|
||||
let actual = ApiUrl::parse(base)
|
||||
.and_then(|url| url.complete_path(&["v1", "ocr"]))
|
||||
.map(|url| url.into_string())
|
||||
.expect("url builds");
|
||||
assert_eq!(actual, expected);
|
||||
}
|
||||
#[rstest]
|
||||
#[case::root("https://example.test", "https://example.test/v1/ocr")]
|
||||
#[case::version_prefix("https://example.test/v1", "https://example.test/v1/ocr")]
|
||||
#[case::complete("https://example.test/v1/ocr", "https://example.test/v1/ocr")]
|
||||
fn completion_appends_only_the_missing_path_suffix(#[case] base: &str, #[case] expected: &str) {
|
||||
let actual = ApiUrl::parse(base)
|
||||
.and_then(|url| url.complete_path(&["v1", "ocr"]))
|
||||
.map(|url| url.into_string())
|
||||
.expect("url builds");
|
||||
assert_eq!(actual, expected);
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
|
|
|||
33
litellm-rust/crates/core-utils/tests/settings.rs
Normal file
33
litellm-rust/crates/core-utils/tests/settings.rs
Normal file
|
|
@ -0,0 +1,33 @@
|
|||
use litellm_core_utils::settings::resolve_non_empty;
|
||||
use rstest::rstest;
|
||||
|
||||
#[rstest]
|
||||
#[case::explicit_wins(Some(" explicit "), &["FIRST", "SECOND"], Some("explicit"))]
|
||||
#[case::absent_falls_back(None, &["FIRST", "SECOND"], Some(" first "))]
|
||||
#[case::blank_falls_back(Some(" \t "), &["BLANK", "SECOND"], Some("second"))]
|
||||
#[case::skips_missing_and_blank(None, &["MISSING", "BLANK", "SECOND"], Some("second"))]
|
||||
#[case::environment_order(None, &["SECOND", "FIRST"], Some("second"))]
|
||||
#[case::missing(None, &["MISSING", "BLANK"], None)]
|
||||
#[case::no_environment(None, &[], None)]
|
||||
fn resolves_explicit_value_then_first_nonblank_environment_value(
|
||||
#[case] value: Option<&str>,
|
||||
#[case] names: &[&str],
|
||||
#[case] expected: Option<&str>,
|
||||
) {
|
||||
let env = |name: &str| match name {
|
||||
"FIRST" => Some(" first ".to_string()),
|
||||
"SECOND" => Some("second".to_string()),
|
||||
"BLANK" => Some(" \t ".to_string()),
|
||||
_ => None,
|
||||
};
|
||||
assert_eq!(resolve_non_empty(value, &env, names).as_deref(), expected);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn explicit_value_does_not_read_the_environment() {
|
||||
let env = |_: &str| panic!("an explicit value must short-circuit environment lookup");
|
||||
assert_eq!(
|
||||
resolve_non_empty(Some("key"), &env, &["KEY"]).as_deref(),
|
||||
Some("key")
|
||||
);
|
||||
}
|
||||
|
|
@ -1,9 +1,17 @@
|
|||
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
|
||||
litellm-core owns route orchestration. Messages and HTTP Responses return `litellm_host::call::CallOutput`, containing either a completed response or a stream head and chunks. OCR and currently non-streaming Chat Completions return their completed response directly
|
||||
|
||||
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
|
||||
Hosts assemble route objects from shared `CoreResources`, HTTP settings, and secret sources. Each route owns its provider client and authentication dependencies. Gateway routes live for the gateway lifetime; Python assembles routes per call from its settings snapshot
|
||||
|
||||
Chat Completions, Messages, Responses, and OCR execute through their route objects. Calls pass `Interceptors` and an optional `ObservationSender` separately; use `&()` for no hooks and `None` for no observer. Construction does no work; preparation and lifecycle observation begin when the future is polled. Handlers accept `Interceptors`, never a concrete `ChannelInterceptors`. Native observers receive start and terminal events through the shared call runner; a stream retains its lifecycle until exhaustion, error, or drop. Hosted routes leave terminal observation to their driver
|
||||
|
||||
`route.rs` declares the concrete `Protocol` and implements a route method that accepts a typed request and constructs a `litellm_host::call::HostedMachine` with `hosted_call`. The shared call plumbing owns stream opening, delivery, backpressure, and detachment. Request decoding belongs to the boundary before the machine starts. Route closures only supply execution dependencies and route-specific host capabilities such as an OCR token provider. Use `run_hosted` for a native host so detachment is reported as cancellation. Python uses its own shared driver and preserves caller-task callback execution
|
||||
|
||||
Responses WebSocket sessions remain separate from the HTTP call driver because a connection can accept multiple requests while receiving events
|
||||
|
||||
## Crate layering
|
||||
|
||||
For Messages, Responses, Chat Completions, OCR, and other API formats, `core/src/<format>/` owns orchestration. Shared API data contracts belong in `litellm-types`, adapter contracts and shared transformation machinery in `llms/src/base_llm/<format>/`, and provider policy in `llms/src/<provider>/<format>/`. A repeated format directory name does not imply interchangeable responsibilities. Select concrete adapters here, then invoke their contracts instead of applying one provider's policy to every call. Route types describe call envelopes and execution state, not duplicate public payload schemas
|
||||
|
||||
Each crate mirrors one top-level Python package, so a Rust path reads as its Python path with the crate name in place of the package directory. Dependencies only point down:
|
||||
|
||||
- `litellm-types` mirrors `litellm/types/`: pure serde data, no I/O
|
||||
|
|
@ -12,7 +20,7 @@ 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`, 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
|
||||
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::interceptors::Interceptors`. 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
|
||||
|
||||
|
|
|
|||
|
|
@ -18,12 +18,11 @@ litellm-auth-aws.workspace = true
|
|||
litellm-http.workspace = true
|
||||
litellm-llms.workspace = true
|
||||
litellm-tracing.workspace = true
|
||||
tracing.workspace = true
|
||||
moka.workspace = true
|
||||
mime_guess = "2.0.5"
|
||||
rand.workspace = true
|
||||
reqwest.workspace = true
|
||||
rustls.workspace = true
|
||||
rustls-native-certs.workspace = true
|
||||
serde.workspace = true
|
||||
serde_json = { workspace = true, features = ["preserve_order"] }
|
||||
strum.workspace = true
|
||||
|
|
@ -39,6 +38,7 @@ veil.workspace = true
|
|||
[dev-dependencies]
|
||||
litellm-http = { workspace = true, features = ["test-support"] }
|
||||
litellm-auth-gcp.workspace = true
|
||||
litellm-host-native.workspace = true
|
||||
litellm-llms = { workspace = true, features = ["test-support"] }
|
||||
rstest.workspace = true
|
||||
rstest_reuse.workspace = true
|
||||
|
|
|
|||
|
|
@ -15,9 +15,9 @@ pub async fn execute_audio_transcription_provider_call(
|
|||
auth: &litellm_auth::AuthServices,
|
||||
request: ProviderAudioTranscriptionRequest,
|
||||
) -> Result<Value, Error> {
|
||||
let env_lookup = |key: &str| std::env::var(key).ok();
|
||||
let env_lookup = |key: &str| request.secrets.get(key);
|
||||
let authenticated = resolve_auth(auth, request.environment.clone(), &env_lookup).await?;
|
||||
let response = crate::outbound::outbound_request(
|
||||
let outbound = crate::outbound::outbound_request(
|
||||
authenticated,
|
||||
request.url.clone(),
|
||||
&request.body,
|
||||
|
|
@ -26,12 +26,12 @@ pub async fn execute_audio_transcription_provider_call(
|
|||
.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()))
|
||||
})?;
|
||||
)?;
|
||||
let response = crate::outbound::send(outbound, http)
|
||||
.await
|
||||
.map_err(|error| {
|
||||
Error::Transport(litellm_http::transport::Error::Network(error.to_string()))
|
||||
})?;
|
||||
let status = response.status();
|
||||
let text = response.text().await.map_err(|error| {
|
||||
Error::Transport(litellm_http::transport::Error::Network(error.to_string()))
|
||||
|
|
@ -42,8 +42,12 @@ pub async fn execute_audio_transcription_provider_call(
|
|||
body: truncate_error_body(&text),
|
||||
}));
|
||||
}
|
||||
let response_json = serde_json::from_str(&text)
|
||||
.map_err(|error| Error::InvalidResponse(format!("invalid audio response JSON: {error}")))?;
|
||||
let response_json = serde_json::from_str(&text).map_err(|error| {
|
||||
Error::InvalidResponse(litellm_llms::ErrorDetail::invalid(
|
||||
"audio response JSON",
|
||||
error,
|
||||
))
|
||||
})?;
|
||||
Ok(request
|
||||
.config
|
||||
.transform_audio_transcription_response(&request.model, response_json)?
|
||||
|
|
|
|||
|
|
@ -3,18 +3,52 @@ pub use crate::error::RouteError as Error;
|
|||
mod handler;
|
||||
mod prepare;
|
||||
pub use handler::execute_audio_transcription_provider_call;
|
||||
use litellm_http::{ClientVariant, HttpClientConfig};
|
||||
use litellm_auth::AuthServices;
|
||||
use litellm_secrets::source::SecretSource;
|
||||
pub use prepare::prepare_audio_transcription_provider_call;
|
||||
use serde_json::Value;
|
||||
use std::sync::Arc;
|
||||
|
||||
use crate::audio_transcription::types::AudioTranscriptionRequest;
|
||||
|
||||
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
|
||||
#[derive(Clone)]
|
||||
pub struct AudioTranscriptionRoute {
|
||||
http: litellm_http::Client,
|
||||
auth: Arc<AuthServices>,
|
||||
secrets: Arc<dyn SecretSource>,
|
||||
}
|
||||
|
||||
impl AudioTranscriptionRoute {
|
||||
pub fn new(
|
||||
http: litellm_http::Client,
|
||||
auth: Arc<AuthServices>,
|
||||
secrets: Arc<dyn SecretSource>,
|
||||
) -> Self {
|
||||
Self {
|
||||
http,
|
||||
auth,
|
||||
secrets,
|
||||
}
|
||||
}
|
||||
|
||||
#[tracing::instrument(name = "litellm.route", skip_all, fields(
|
||||
route = "audio_transcription",
|
||||
model = request.model,
|
||||
provider,
|
||||
resolved_model,
|
||||
stream = false,
|
||||
outcome
|
||||
))]
|
||||
pub async fn execute(&self, request: AudioTranscriptionRequest<'_>) -> Result<Value, Error> {
|
||||
crate::diagnostic::unary(async {
|
||||
let request =
|
||||
prepare_audio_transcription_provider_call(request, self.secrets.as_ref()).await?;
|
||||
crate::diagnostic::provider(&request.model, &request.custom_llm_provider);
|
||||
let execute: futures_util::future::BoxFuture<'_, Result<Value, Error>> = Box::pin(
|
||||
execute_audio_transcription_provider_call(&self.http, &self.auth, request),
|
||||
);
|
||||
execute.await
|
||||
})
|
||||
.await
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,52 +1,55 @@
|
|||
use litellm_core_utils::get_llm_provider_logic::{CustomLlmProvider, get_custom_llm_provider};
|
||||
use litellm_http::request::string_headers;
|
||||
use litellm_http::request::with_default_headers;
|
||||
use litellm_llms::{
|
||||
base_llm::{
|
||||
audio_transcription::transformation::BaseAudioTranscriptionConfig,
|
||||
auth::{ValidatedEnvironment, with_default_headers},
|
||||
auth::ValidatedEnvironment,
|
||||
},
|
||||
bedrock::audio_transcription::BEDROCK_AUDIO_TRANSCRIPTION_CONFIG,
|
||||
};
|
||||
use litellm_secrets::source::SecretSource;
|
||||
|
||||
use super::Error;
|
||||
use crate::audio_transcription::types::{
|
||||
AudioTranscriptionRequest, ProviderAudioTranscriptionRequest,
|
||||
};
|
||||
use crate::provider::{LlmProviders, resolve_llm_provider};
|
||||
|
||||
fn provider_config(provider: &str) -> Option<&'static dyn BaseAudioTranscriptionConfig> {
|
||||
if provider == "bedrock" {
|
||||
return Some(&BEDROCK_AUDIO_TRANSCRIPTION_CONFIG);
|
||||
fn provider_config(provider: LlmProviders) -> Option<&'static dyn BaseAudioTranscriptionConfig> {
|
||||
match provider {
|
||||
LlmProviders::Bedrock => Some(&BEDROCK_AUDIO_TRANSCRIPTION_CONFIG),
|
||||
LlmProviders::Anthropic
|
||||
| LlmProviders::AwsTextract
|
||||
| LlmProviders::AzureAi
|
||||
| LlmProviders::Cohere
|
||||
| LlmProviders::Mistral
|
||||
| LlmProviders::Openai
|
||||
| LlmProviders::OpenaiLike
|
||||
| LlmProviders::Reducto
|
||||
| LlmProviders::VertexAi => None,
|
||||
}
|
||||
let _ = provider;
|
||||
None
|
||||
}
|
||||
|
||||
pub fn prepare_audio_transcription_provider_call(
|
||||
#[tracing::instrument(name = "litellm.prepare", level = "debug", skip_all)]
|
||||
pub async fn prepare_audio_transcription_provider_call(
|
||||
request: AudioTranscriptionRequest<'_>,
|
||||
secrets: &dyn SecretSource,
|
||||
) -> Result<ProviderAudioTranscriptionRequest, Error> {
|
||||
let provider_info = get_custom_llm_provider(request.model, request.custom_llm_provider)
|
||||
.or_else(|| {
|
||||
request
|
||||
.custom_llm_provider
|
||||
.map(|provider| CustomLlmProvider {
|
||||
model: request.model,
|
||||
custom_llm_provider: provider,
|
||||
})
|
||||
})
|
||||
.ok_or_else(|| {
|
||||
Error::InvalidProvider(
|
||||
"unable to resolve custom_llm_provider for audio transcription request".to_string(),
|
||||
)
|
||||
})?;
|
||||
let provider_info = resolve_llm_provider(
|
||||
request.model,
|
||||
request.custom_llm_provider,
|
||||
"audio transcription",
|
||||
)?;
|
||||
let model = provider_info.model.to_string();
|
||||
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 config = provider_config(provider_info.provider)
|
||||
.ok_or_else(|| Error::InvalidProvider(<&str>::from(provider_info.provider).to_string()))?;
|
||||
let snapshot = secrets.resolve(&config.secret_names()).await?;
|
||||
let env_lookup = |key: &str| snapshot.get(key);
|
||||
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")]),
|
||||
headers: with_default_headers(validated.headers, config.default_headers()),
|
||||
auth: validated.auth,
|
||||
};
|
||||
let url = config.get_complete_url(
|
||||
|
|
@ -60,11 +63,12 @@ pub fn prepare_audio_transcription_provider_call(
|
|||
config.transform_audio_transcription_request(&model, request.audio, filtered_params)?;
|
||||
Ok(ProviderAudioTranscriptionRequest {
|
||||
model,
|
||||
custom_llm_provider: provider_info.custom_llm_provider.to_string(),
|
||||
custom_llm_provider: <&str>::from(provider_info.provider).to_string(),
|
||||
config,
|
||||
url,
|
||||
body: transformed.body,
|
||||
environment,
|
||||
secrets: snapshot,
|
||||
timeout: request.timeout,
|
||||
})
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
use litellm_secrets::source::Secrets;
|
||||
use std::time::Duration;
|
||||
|
||||
use litellm_llms::base_llm::{
|
||||
|
|
@ -24,6 +25,7 @@ pub struct ProviderAudioTranscriptionRequest {
|
|||
pub url: String,
|
||||
pub body: Value,
|
||||
pub environment: ValidatedEnvironment,
|
||||
pub secrets: Secrets,
|
||||
pub timeout: Option<Duration>,
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -3,18 +3,43 @@ use litellm_llms::{
|
|||
anthropic::chat::transformation::ANTHROPIC_CHAT_COMPLETIONS_CONFIG,
|
||||
base_llm::chat::transformation::BaseConfig,
|
||||
bedrock::chat::converse_transformation::BEDROCK_CHAT_COMPLETIONS_CONFIG,
|
||||
openai_like::chat::transformation::OPENAI_LIKE_CHAT_COMPLETIONS_CONFIG,
|
||||
};
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
use super::Error;
|
||||
use crate::provider::LlmProviders;
|
||||
|
||||
const HEADER_CONTEXT: &str = "chat completions";
|
||||
|
||||
pub(super) fn chat_completions_provider_config(provider: &str) -> Option<&'static dyn BaseConfig> {
|
||||
pub(super) enum ChatProvider {
|
||||
Anthropic,
|
||||
Bedrock,
|
||||
OpenaiLike,
|
||||
}
|
||||
|
||||
impl ChatProvider {
|
||||
pub(super) fn config(self) -> &'static dyn BaseConfig {
|
||||
match self {
|
||||
Self::Anthropic => &ANTHROPIC_CHAT_COMPLETIONS_CONFIG,
|
||||
Self::Bedrock => &BEDROCK_CHAT_COMPLETIONS_CONFIG,
|
||||
Self::OpenaiLike => &OPENAI_LIKE_CHAT_COMPLETIONS_CONFIG,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn chat_completions_provider(provider: LlmProviders) -> Option<ChatProvider> {
|
||||
match provider {
|
||||
"anthropic" => Some(&ANTHROPIC_CHAT_COMPLETIONS_CONFIG),
|
||||
"bedrock" => Some(&BEDROCK_CHAT_COMPLETIONS_CONFIG),
|
||||
_ => None,
|
||||
LlmProviders::Anthropic => Some(ChatProvider::Anthropic),
|
||||
LlmProviders::Bedrock => Some(ChatProvider::Bedrock),
|
||||
LlmProviders::OpenaiLike => Some(ChatProvider::OpenaiLike),
|
||||
LlmProviders::AwsTextract
|
||||
| LlmProviders::AzureAi
|
||||
| LlmProviders::Cohere
|
||||
| LlmProviders::Mistral
|
||||
| LlmProviders::Openai
|
||||
| LlmProviders::Reducto
|
||||
| LlmProviders::VertexAi => None,
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -1,10 +1,9 @@
|
|||
use litellm_host::lifecycle::ExecutionEvent;
|
||||
use litellm_host::observation::ObservationSender;
|
||||
use std::time::Duration;
|
||||
|
||||
use litellm_auth::AuthServices;
|
||||
use litellm_host::{
|
||||
event::{MachineEvent, RawResponse, RequestContext, WireRequest},
|
||||
hooks::RouteHooks,
|
||||
};
|
||||
use litellm_host::interceptors::{Interceptors, RawResponse, RequestContext, WireRequest};
|
||||
use litellm_http::{Client, outbound::OutboundRequest, request::truncate_error_body};
|
||||
use litellm_llms::base_llm::{
|
||||
auth::{Authenticated, resolve_auth},
|
||||
|
|
@ -23,7 +22,8 @@ pub(super) async fn execute(
|
|||
http: &Client,
|
||||
auth: &AuthServices,
|
||||
request: ProviderChatCompletionsRequest,
|
||||
hooks: &impl RouteHooks<Error>,
|
||||
interceptors: &impl Interceptors<Error>,
|
||||
observers: Option<&ObservationSender>,
|
||||
) -> Result<ChatCompletionsResponse, Error> {
|
||||
let ProviderChatCompletionsRequest {
|
||||
model,
|
||||
|
|
@ -33,6 +33,7 @@ pub(super) async fn execute(
|
|||
body,
|
||||
optional_params,
|
||||
environment,
|
||||
secrets,
|
||||
timeout,
|
||||
api_key,
|
||||
} = request;
|
||||
|
|
@ -43,9 +44,9 @@ pub(super) async fn execute(
|
|||
secret_fields: Vec::new(),
|
||||
api_key,
|
||||
};
|
||||
let authenticated = resolve_auth(auth, environment, &|key| std::env::var(key).ok()).await?;
|
||||
let wire = hooks
|
||||
.before_send(
|
||||
let authenticated = resolve_auth(auth, environment, &|key| secrets.get(key)).await?;
|
||||
let wire = interceptors
|
||||
.before_provider_request(
|
||||
WireRequest {
|
||||
url,
|
||||
headers: authenticated.headers,
|
||||
|
|
@ -64,7 +65,7 @@ pub(super) async fn execute(
|
|||
timeout,
|
||||
)?;
|
||||
|
||||
let response = outbound.send(http).await.map_err(|err| {
|
||||
let response = crate::outbound::send(outbound, 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.
|
||||
|
|
@ -86,14 +87,22 @@ pub(super) async fn execute(
|
|||
body: truncate_error_body(&text),
|
||||
}));
|
||||
}
|
||||
hooks
|
||||
.emit(MachineEvent::ResponseReceived {
|
||||
raw: RawResponse { body: text.clone() },
|
||||
})
|
||||
.await?;
|
||||
let raw = RawResponse { body: text.clone() };
|
||||
if let Some(observers) = observers {
|
||||
observers.emit(litellm_host::lifecycle::CallEvent::Execution(
|
||||
ExecutionEvent::ProviderResponseReceived { raw: raw.clone() },
|
||||
));
|
||||
}
|
||||
interceptors
|
||||
.after_provider_response(raw)
|
||||
.await
|
||||
.map_err(Error::post_call)?;
|
||||
|
||||
let body: Value = serde_json::from_str(&text).map_err(|err| {
|
||||
Error::InvalidResponse(format!("invalid chat completions response JSON: {err}"))
|
||||
Error::InvalidResponse(litellm_llms::ErrorDetail::invalid(
|
||||
"chat completions response JSON",
|
||||
err,
|
||||
))
|
||||
})?;
|
||||
config
|
||||
.transform_response(&model, ProviderChatResponseData { body })
|
||||
|
|
@ -114,7 +123,7 @@ pub(super) fn as_response_error(err: Error) -> Error {
|
|||
match err {
|
||||
already @ (Error::InvalidResponse(_)
|
||||
| Error::Transport(litellm_http::transport::Error::Http { .. })) => already,
|
||||
other => Error::InvalidResponse(other.to_string()),
|
||||
other => Error::InvalidResponse(other.to_string().into()),
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -163,8 +172,8 @@ mod tests {
|
|||
raw: Mutex<Vec<String>>,
|
||||
}
|
||||
|
||||
impl RouteHooks<Error> for RecordingHooks {
|
||||
async fn before_send(
|
||||
impl Interceptors<Error> for RecordingHooks {
|
||||
async fn before_provider_request(
|
||||
&self,
|
||||
wire: WireRequest,
|
||||
context: RequestContext,
|
||||
|
|
@ -183,8 +192,7 @@ mod tests {
|
|||
})
|
||||
}
|
||||
|
||||
async fn emit(&self, event: MachineEvent) -> Result<(), Error> {
|
||||
let MachineEvent::ResponseReceived { raw } = event;
|
||||
async fn after_provider_response(&self, raw: RawResponse) -> Result<(), Error> {
|
||||
self.raw.lock().unwrap().push(raw.body);
|
||||
Ok(())
|
||||
}
|
||||
|
|
@ -203,6 +211,7 @@ mod tests {
|
|||
timeout: None,
|
||||
})
|
||||
.unwrap(),
|
||||
std::sync::Arc::new(|_: &str| None),
|
||||
)
|
||||
.unwrap()
|
||||
}
|
||||
|
|
@ -217,13 +226,14 @@ mod tests {
|
|||
)
|
||||
.mount(&upstream)
|
||||
.await;
|
||||
let hooks = RecordingHooks::default();
|
||||
let interceptors = RecordingHooks::default();
|
||||
|
||||
execute(
|
||||
&Client::plain_for_test(),
|
||||
&AuthServices::default(),
|
||||
prepared(&upstream.uri()),
|
||||
&hooks,
|
||||
&interceptors,
|
||||
None,
|
||||
)
|
||||
.await
|
||||
.expect("chat completions call succeeds");
|
||||
|
|
@ -234,14 +244,17 @@ mod tests {
|
|||
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()));
|
||||
let [context] =
|
||||
<[RequestContext; 1]>::try_from(interceptors.contexts.into_inner().unwrap())
|
||||
.unwrap_or_else(|seen| {
|
||||
panic!("before_provider_request 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]);
|
||||
assert_eq!(interceptors.raw.into_inner().unwrap(), [ANTHROPIC_MESSAGE]);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
|
|
@ -252,13 +265,14 @@ mod tests {
|
|||
.respond_with(ResponseTemplate::new(500).set_body_string("boom"))
|
||||
.mount(&upstream)
|
||||
.await;
|
||||
let hooks = RecordingHooks::default();
|
||||
let interceptors = RecordingHooks::default();
|
||||
|
||||
let error = execute(
|
||||
&Client::plain_for_test(),
|
||||
&AuthServices::default(),
|
||||
prepared(&upstream.uri()),
|
||||
&hooks,
|
||||
&interceptors,
|
||||
None,
|
||||
)
|
||||
.await
|
||||
.expect_err("the upstream failure fails the call");
|
||||
|
|
@ -267,15 +281,15 @@ mod tests {
|
|||
error,
|
||||
Error::Transport(litellm_http::transport::Error::Http { status: 500, .. })
|
||||
));
|
||||
assert!(hooks.raw.into_inner().unwrap().is_empty());
|
||||
assert!(interceptors.raw.into_inner().unwrap().is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[rstest::rstest]
|
||||
fn response_errors_collapse_to_one_variant_that_can_only_mean_already_sent() {
|
||||
for original in [
|
||||
Error::MissingField("usage"),
|
||||
Error::Unsupported("non-text response content block"),
|
||||
Error::InvalidRequest("whatever".to_string()),
|
||||
Error::InvalidRequest("whatever".to_string().into()),
|
||||
Error::Auth(litellm_auth::Error::InvalidHeader),
|
||||
] {
|
||||
let label = format!("{original:?}");
|
||||
|
|
|
|||
|
|
@ -1,56 +1,85 @@
|
|||
//! The `/chat/completions` call, the Rust equivalent of Python's
|
||||
//! `litellm.completion()`.
|
||||
//!
|
||||
//! [`chat_completions`] is the top-level entrypoint: give it a model, the
|
||||
//! OpenAI-shaped message list, the provider-mapped optional params, and
|
||||
//! credentials, and it resolves the provider, translates the conversation,
|
||||
//! calls the provider, and returns a typed OpenAI-shaped response.
|
||||
|
||||
use litellm_host::observation::ObservationSender;
|
||||
pub mod route;
|
||||
pub mod types;
|
||||
pub use crate::error::RouteError as Error;
|
||||
mod common_utils;
|
||||
pub(crate) mod handler;
|
||||
mod prepare;
|
||||
use litellm_http::{ClientVariant, HttpClientConfig};
|
||||
use litellm_types::utils::ChatCompletionsResponse;
|
||||
use prepare::{parse_messages, prepare_provider_request, resolve_provider_config, resolve_request};
|
||||
use serde_json::{Map, Value};
|
||||
use prepare::{prepare_provider_request, resolve_request};
|
||||
|
||||
use crate::chat_completions::types::ChatCompletionsRequest;
|
||||
use litellm_auth::AuthServices;
|
||||
use litellm_secrets::source::SecretSource;
|
||||
use std::sync::Arc;
|
||||
|
||||
pub async fn chat_completions(
|
||||
resources: &crate::resources::CoreResources,
|
||||
config: &HttpClientConfig,
|
||||
request: ChatCompletionsRequest<'_>,
|
||||
) -> Result<ChatCompletionsResponse, Error> {
|
||||
let http = resources.pool.client(config, ClientVariant::Provider)?;
|
||||
let request = prepare_provider_request(resolve_request(request)?)?;
|
||||
handler::execute(&http, &resources.auth, request, &()).await
|
||||
#[derive(Clone)]
|
||||
pub struct ChatCompletionsRoute {
|
||||
http: litellm_http::Client,
|
||||
auth: Arc<AuthServices>,
|
||||
secrets: Arc<dyn SecretSource>,
|
||||
}
|
||||
|
||||
/// Whether the core would accept this request, without resolving credentials or
|
||||
/// touching the network.
|
||||
///
|
||||
/// A host that keeps the Python implementation asks this first so it can emit
|
||||
/// its pre-call logging exactly once, on whichever path is about to run.
|
||||
/// Returns the decline reason, or `None` when the request is accepted.
|
||||
pub fn chat_completions_decline_reason(
|
||||
model: &str,
|
||||
custom_llm_provider: Option<&str>,
|
||||
messages: Value,
|
||||
optional_params: &Map<String, Value>,
|
||||
) -> Option<&'static str> {
|
||||
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");
|
||||
};
|
||||
if messages.is_empty() {
|
||||
return Some("empty message list");
|
||||
impl ChatCompletionsRoute {
|
||||
pub fn new(
|
||||
http: litellm_http::Client,
|
||||
auth: Arc<AuthServices>,
|
||||
secrets: Arc<dyn SecretSource>,
|
||||
) -> Self {
|
||||
Self {
|
||||
http,
|
||||
auth,
|
||||
secrets,
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn execute(
|
||||
&self,
|
||||
request: ChatCompletionsRequest<'_>,
|
||||
interceptors: &impl litellm_host::interceptors::Interceptors<Error>,
|
||||
observers: Option<ObservationSender>,
|
||||
) -> Result<ChatCompletionsResponse, Error> {
|
||||
litellm_host::lifecycle::observe_unary(
|
||||
observers.clone(),
|
||||
self.run(request, interceptors, observers.as_ref()),
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
#[tracing::instrument(name = "litellm.route", skip_all, fields(
|
||||
route = "chat_completions",
|
||||
model = %request.model,
|
||||
provider,
|
||||
resolved_model,
|
||||
stream = false,
|
||||
outcome
|
||||
))]
|
||||
async fn run(
|
||||
&self,
|
||||
request: ChatCompletionsRequest<'_>,
|
||||
interceptors: &impl litellm_host::interceptors::Interceptors<Error>,
|
||||
observers: Option<&ObservationSender>,
|
||||
) -> Result<ChatCompletionsResponse, Error> {
|
||||
crate::diagnostic::unary(async {
|
||||
let resolved = resolve_request(request)?;
|
||||
let snapshot = self
|
||||
.secrets
|
||||
.resolve(&resolved.config.secret_names())
|
||||
.await?;
|
||||
let prepared = prepare_provider_request(resolved, snapshot)?;
|
||||
crate::diagnostic::provider(&prepared.model, &prepared.custom_llm_provider);
|
||||
let execute: futures_util::future::BoxFuture<
|
||||
'_,
|
||||
Result<ChatCompletionsResponse, Error>,
|
||||
> = Box::pin(handler::execute(
|
||||
&self.http,
|
||||
&self.auth,
|
||||
prepared,
|
||||
interceptors,
|
||||
observers,
|
||||
));
|
||||
execute.await
|
||||
})
|
||||
.await
|
||||
}
|
||||
config
|
||||
.unsupported_reason(&messages, optional_params)
|
||||
.map(|reason| reason.0)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,19 +1,19 @@
|
|||
use litellm_auth::SecretValue;
|
||||
use litellm_core_utils::get_llm_provider_logic::{CustomLlmProvider, get_custom_llm_provider};
|
||||
use litellm_llms::base_llm::{
|
||||
auth::{ValidatedEnvironment, with_default_headers},
|
||||
chat::transformation::BaseConfig,
|
||||
};
|
||||
use litellm_core_utils::settings::Lookup;
|
||||
use litellm_http::request::with_default_headers;
|
||||
use litellm_llms::base_llm::{auth::ValidatedEnvironment, chat::transformation::BaseConfig};
|
||||
use litellm_secrets::source::Secrets;
|
||||
use litellm_types::llms::openai::ChatMessage;
|
||||
use serde_json::Value;
|
||||
|
||||
use super::{
|
||||
Error,
|
||||
common_utils::{chat_completions_provider_config, string_headers},
|
||||
common_utils::{chat_completions_provider, string_headers},
|
||||
};
|
||||
use crate::chat_completions::types::{
|
||||
ChatCompletionsRequest, ProviderChatCompletionsRequest, ResolvedChatCompletionsRequest,
|
||||
};
|
||||
use crate::provider::resolve_llm_provider;
|
||||
|
||||
pub(super) struct ResolvedProvider {
|
||||
pub(super) model: String,
|
||||
|
|
@ -25,30 +25,24 @@ pub(super) fn resolve_provider_config<'a>(
|
|||
model: &'a str,
|
||||
custom_llm_provider: Option<&'a str>,
|
||||
) -> Result<ResolvedProvider, Error> {
|
||||
let provider_info = get_custom_llm_provider(model, custom_llm_provider)
|
||||
.or_else(|| {
|
||||
custom_llm_provider.map(|provider| CustomLlmProvider {
|
||||
model,
|
||||
custom_llm_provider: provider,
|
||||
})
|
||||
})
|
||||
.ok_or_else(|| {
|
||||
Error::InvalidProvider(
|
||||
"unable to resolve custom_llm_provider for chat completions request".to_string(),
|
||||
)
|
||||
})?;
|
||||
let config = chat_completions_provider_config(provider_info.custom_llm_provider)
|
||||
.ok_or_else(|| Error::InvalidProvider(provider_info.custom_llm_provider.to_string()))?;
|
||||
let provider_info = resolve_llm_provider(model, custom_llm_provider, "chat completions")?;
|
||||
let config = chat_completions_provider(provider_info.provider)
|
||||
.ok_or_else(|| Error::InvalidProvider(<&str>::from(provider_info.provider).to_string()))?
|
||||
.config();
|
||||
Ok(ResolvedProvider {
|
||||
model: provider_info.model.to_string(),
|
||||
custom_llm_provider: provider_info.custom_llm_provider.to_string(),
|
||||
custom_llm_provider: <&str>::from(provider_info.provider).to_string(),
|
||||
config,
|
||||
})
|
||||
}
|
||||
|
||||
pub(super) fn parse_messages(messages: Value) -> Result<Vec<ChatMessage>, Error> {
|
||||
serde_json::from_value(messages)
|
||||
.map_err(|err| Error::InvalidRequest(format!("invalid chat completions messages: {err}")))
|
||||
serde_json::from_value(messages).map_err(|err| {
|
||||
Error::InvalidRequest(litellm_llms::ErrorDetail::invalid(
|
||||
"chat completions messages",
|
||||
err,
|
||||
))
|
||||
})
|
||||
}
|
||||
|
||||
pub(super) fn resolve_request(
|
||||
|
|
@ -62,7 +56,7 @@ pub(super) fn resolve_request(
|
|||
let messages = parse_messages(request.messages)?;
|
||||
if messages.is_empty() {
|
||||
return Err(Error::InvalidRequest(
|
||||
"chat completions requires at least one message".to_string(),
|
||||
"chat completions requires at least one message".into(),
|
||||
));
|
||||
}
|
||||
if let Some(reason) = config.unsupported_reason(&messages, &request.optional_params) {
|
||||
|
|
@ -85,8 +79,9 @@ fn validate_environment(
|
|||
request: &ResolvedChatCompletionsRequest<'_>,
|
||||
model: &str,
|
||||
config: &dyn BaseConfig,
|
||||
secrets: &dyn Lookup,
|
||||
) -> Result<ValidatedEnvironment, Error> {
|
||||
let env_lookup = |key: &str| std::env::var(key).ok();
|
||||
let env_lookup = |key: &str| secrets.get(key);
|
||||
let forwarded = string_headers(request.extra_headers.clone())?;
|
||||
let validated = config.validate_environment(
|
||||
forwarded,
|
||||
|
|
@ -101,13 +96,16 @@ fn validate_environment(
|
|||
})
|
||||
}
|
||||
|
||||
#[tracing::instrument(name = "litellm.prepare", level = "debug", skip_all)]
|
||||
pub(super) fn prepare_provider_request(
|
||||
request: ResolvedChatCompletionsRequest<'_>,
|
||||
secrets: Secrets,
|
||||
) -> Result<ProviderChatCompletionsRequest, Error> {
|
||||
let environment = validate_environment(&request, &request.model, request.config)?;
|
||||
let environment =
|
||||
validate_environment(&request, &request.model, request.config, secrets.as_ref())?;
|
||||
let model = request.model;
|
||||
let config = request.config;
|
||||
let env_lookup = |key: &str| std::env::var(key).ok();
|
||||
let env_lookup = |key: &str| secrets.get(key);
|
||||
let url = config.get_complete_url(
|
||||
request.api_base,
|
||||
&model,
|
||||
|
|
@ -125,6 +123,7 @@ pub(super) fn prepare_provider_request(
|
|||
body: transformed.body,
|
||||
optional_params: request.optional_params,
|
||||
environment,
|
||||
secrets,
|
||||
timeout: request.timeout,
|
||||
api_key: request.api_key.map(|key| SecretValue::new(key.to_string())),
|
||||
})
|
||||
|
|
@ -145,7 +144,10 @@ mod tests {
|
|||
fn prepare_chat_completions_call(
|
||||
request: ChatCompletionsRequest<'_>,
|
||||
) -> Result<ProviderChatCompletionsRequest, Error> {
|
||||
prepare_provider_request(resolve_request(request)?)
|
||||
prepare_provider_request(
|
||||
resolve_request(request)?,
|
||||
std::sync::Arc::new(|_: &str| None),
|
||||
)
|
||||
}
|
||||
|
||||
/// The headers as they go on the wire, credential applied.
|
||||
|
|
@ -183,13 +185,10 @@ mod tests {
|
|||
}
|
||||
}
|
||||
|
||||
/// `ProviderChatCompletionsRequest` deliberately has no `Debug` (its headers
|
||||
/// carry resolved credentials), so unwrap the failure case by hand.
|
||||
fn decline(request: ChatCompletionsRequest<'_>) -> Error {
|
||||
match prepare_chat_completions_call(request) {
|
||||
Err(error) => error,
|
||||
Ok(prepared) => panic!("expected a decline, prepared a call to {}", prepared.url),
|
||||
}
|
||||
fn preparation_error(request: ChatCompletionsRequest<'_>) -> Error {
|
||||
prepare_chat_completions_call(request)
|
||||
.err()
|
||||
.expect("request preparation should fail")
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
|
@ -335,24 +334,19 @@ mod tests {
|
|||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn declines_an_unsupported_request_before_resolving_credentials() {
|
||||
let mut call = request(
|
||||
"claude-sonnet-4-5",
|
||||
Some("anthropic"),
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({"stream": true}),
|
||||
);
|
||||
call.api_key = None;
|
||||
// No api_key is set and no env is consulted: the gate must run first, so the
|
||||
// error is the decline rather than a missing-credential error.
|
||||
assert_eq!(decline(call), Error::Unsupported("streaming"));
|
||||
#[rstest::rstest]
|
||||
fn rejects_empty_messages_before_resolving_credentials() {
|
||||
let call = ChatCompletionsRequest {
|
||||
api_key: None,
|
||||
..request("claude-sonnet-4-5", Some("anthropic"), json!([]), json!({}))
|
||||
};
|
||||
assert!(matches!(preparation_error(call), Error::InvalidRequest(_)));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_an_unknown_provider() {
|
||||
assert_eq!(
|
||||
decline(request(
|
||||
preparation_error(request(
|
||||
"openai/gpt-4o",
|
||||
None,
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
|
|
@ -365,7 +359,7 @@ mod tests {
|
|||
#[test]
|
||||
fn rejects_a_model_with_no_resolvable_provider() {
|
||||
assert!(matches!(
|
||||
decline(request(
|
||||
preparation_error(request(
|
||||
"claude-sonnet-4-5",
|
||||
None,
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
|
|
@ -378,16 +372,20 @@ mod tests {
|
|||
#[test]
|
||||
fn rejects_an_empty_or_malformed_message_list() {
|
||||
assert_eq!(
|
||||
decline(request(
|
||||
preparation_error(request(
|
||||
"anthropic/claude-sonnet-4-5",
|
||||
None,
|
||||
json!([]),
|
||||
json!({}),
|
||||
)),
|
||||
Error::InvalidRequest("chat completions requires at least one message".to_string())
|
||||
Error::InvalidRequest(
|
||||
"chat completions requires at least one message"
|
||||
.to_string()
|
||||
.into()
|
||||
)
|
||||
);
|
||||
assert!(matches!(
|
||||
decline(request(
|
||||
preparation_error(request(
|
||||
"anthropic/claude-sonnet-4-5",
|
||||
None,
|
||||
json!("not a list"),
|
||||
|
|
@ -407,7 +405,7 @@ mod tests {
|
|||
);
|
||||
call.extra_headers = Some(Map::from_iter([("x-trace".to_string(), json!(7))]));
|
||||
assert_eq!(
|
||||
decline(call),
|
||||
preparation_error(call),
|
||||
Error::Headers(litellm_http::request::HeaderError {
|
||||
context: "chat completions",
|
||||
name: "x-trace".to_string(),
|
||||
|
|
@ -505,18 +503,17 @@ mod tests {
|
|||
);
|
||||
}
|
||||
|
||||
#[rstest::rstest]
|
||||
#[case::authorization("Authorization")]
|
||||
#[case::amz_date("x-amz-date")]
|
||||
#[case::security_token("x-amz-security-token")]
|
||||
#[case::date("Date")]
|
||||
#[tokio::test]
|
||||
async fn a_forwarded_header_the_signer_computes_declines_to_python() {
|
||||
// Reattaching the caller's copy next to the computed one puts the name on
|
||||
// the wire twice and Bedrock rejects the pair, so a request carrying one
|
||||
// has to go to Python instead of being signed here.
|
||||
for forwarded in [
|
||||
"Authorization",
|
||||
"x-amz-date",
|
||||
"x-amz-security-token",
|
||||
"Date",
|
||||
] {
|
||||
let mut call = request(
|
||||
async fn rejects_a_forwarded_header_the_signer_computes(#[case] forwarded: &str) {
|
||||
let call = ChatCompletionsRequest {
|
||||
api_key: None,
|
||||
extra_headers: Some(Map::from_iter([(forwarded.to_string(), json!("forged"))])),
|
||||
..request(
|
||||
"bedrock/us-east-1/anthropic.claude-v2",
|
||||
None,
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
|
|
@ -525,29 +522,27 @@ mod tests {
|
|||
"aws_access_key_id": "AKIDEXAMPLE",
|
||||
"aws_secret_access_key": "wJalrXUtnFEMI/K7MDENG+bPxRfiCYEXAMPLEKEY"
|
||||
}),
|
||||
);
|
||||
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 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"
|
||||
);
|
||||
}
|
||||
};
|
||||
let prepared = prepare_chat_completions_call(call).expect("prepares");
|
||||
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("conflicting signing headers must fail");
|
||||
assert!(
|
||||
matches!(error, Error::Unsupported(_)),
|
||||
"{forwarded} returned {error:?}"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
|
@ -640,112 +635,4 @@ mod tests {
|
|||
"prepare did not carry the bearer token"
|
||||
);
|
||||
}
|
||||
|
||||
fn decline_reason(
|
||||
model: &str,
|
||||
provider: Option<&str>,
|
||||
messages: Value,
|
||||
params: Value,
|
||||
) -> Option<&'static str> {
|
||||
let params = match params {
|
||||
Value::Object(map) => map,
|
||||
other => panic!("params must be an object, got {other}"),
|
||||
};
|
||||
crate::chat_completions::chat_completions_decline_reason(model, provider, messages, ¶ms)
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn the_gate_accepts_what_prepare_accepts() {
|
||||
assert_eq!(
|
||||
decline_reason(
|
||||
"anthropic/claude-sonnet-4-5",
|
||||
None,
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({"max_tokens": 16}),
|
||||
),
|
||||
None
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn the_gate_declines_without_resolving_credentials_or_calling_out() {
|
||||
assert_eq!(
|
||||
decline_reason(
|
||||
"anthropic/claude-sonnet-4-5",
|
||||
None,
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({"stream": true}),
|
||||
),
|
||||
Some("streaming")
|
||||
);
|
||||
assert_eq!(
|
||||
decline_reason(
|
||||
"openai/gpt-4o",
|
||||
None,
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({}),
|
||||
),
|
||||
Some("provider is not on the rust chat completions path")
|
||||
);
|
||||
assert_eq!(
|
||||
decline_reason(
|
||||
"claude-sonnet-4-5",
|
||||
None,
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({}),
|
||||
),
|
||||
Some("provider is not on the rust chat completions path")
|
||||
);
|
||||
assert_eq!(
|
||||
decline_reason(
|
||||
"anthropic/claude-sonnet-4-5",
|
||||
None,
|
||||
json!("nope"),
|
||||
json!({})
|
||||
),
|
||||
Some("unreadable message list")
|
||||
);
|
||||
assert_eq!(
|
||||
decline_reason("anthropic/claude-sonnet-4-5", None, json!([]), json!({})),
|
||||
Some("empty message list")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn the_gate_agrees_with_prepare_on_every_case_it_accepts() {
|
||||
// A gate that accepts what prepare then declines would make the host emit
|
||||
// its pre-call logging on a path that falls back, so pin the agreement.
|
||||
for (messages, params) in [
|
||||
(
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({"max_tokens": 8}),
|
||||
),
|
||||
(
|
||||
json!([{"role": "system", "content": "s"}, {"role": "user", "content": "hi"}]),
|
||||
json!({"temperature": 0.1}),
|
||||
),
|
||||
(
|
||||
json!([{"role": "user", "content": "hi"}, {"role": "assistant", "content": "yo"}]),
|
||||
json!({}),
|
||||
),
|
||||
] {
|
||||
assert_eq!(
|
||||
decline_reason(
|
||||
"anthropic/claude-sonnet-4-5",
|
||||
None,
|
||||
messages.clone(),
|
||||
params.clone()
|
||||
),
|
||||
None,
|
||||
"gate declined {messages}"
|
||||
);
|
||||
prepare_chat_completions_call(request(
|
||||
"anthropic/claude-sonnet-4-5",
|
||||
None,
|
||||
messages.clone(),
|
||||
params,
|
||||
))
|
||||
.unwrap_or_else(|error| panic!("prepare declined {messages}: {error}"));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
|
|||
52
litellm-rust/crates/core/src/chat_completions/route.rs
Normal file
52
litellm-rust/crates/core/src/chat_completions/route.rs
Normal file
|
|
@ -0,0 +1,52 @@
|
|||
use litellm_host::observation::ObservationSender;
|
||||
use std::convert::Infallible;
|
||||
|
||||
use litellm_host::{
|
||||
call::{CallOutput, HostedMachine, hosted_call},
|
||||
protocol::Protocol,
|
||||
};
|
||||
use litellm_types::utils::ChatCompletionsResponse;
|
||||
|
||||
use super::{
|
||||
ChatCompletionsRoute, Error,
|
||||
types::{ChatCompletionsCall, ChatCompletionsRequest},
|
||||
};
|
||||
|
||||
pub struct ChatCompletions;
|
||||
|
||||
impl Protocol for ChatCompletions {
|
||||
type Response = ChatCompletionsResponse;
|
||||
type Error = Error;
|
||||
type Request = ChatCompletionsCall;
|
||||
type HostCall = Infallible;
|
||||
type Chunk = Infallible;
|
||||
type StreamHead = Infallible;
|
||||
}
|
||||
|
||||
impl ChatCompletionsRoute {
|
||||
pub fn machine(
|
||||
self,
|
||||
call: ChatCompletionsCall,
|
||||
observers: Option<ObservationSender>,
|
||||
) -> HostedMachine<ChatCompletions> {
|
||||
hosted_call(
|
||||
call,
|
||||
observers,
|
||||
move |call: ChatCompletionsCall, _, interceptors, observers| async move {
|
||||
let request = ChatCompletionsRequest {
|
||||
model: &call.model,
|
||||
messages: call.messages,
|
||||
optional_params: call.optional_params,
|
||||
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,
|
||||
timeout: call.timeout,
|
||||
};
|
||||
self.run(request, &interceptors, observers.as_ref())
|
||||
.await
|
||||
.map(CallOutput::Complete)
|
||||
},
|
||||
)
|
||||
}
|
||||
}
|
||||
|
|
@ -1,3 +1,4 @@
|
|||
use litellm_secrets::source::Secrets;
|
||||
use std::time::Duration;
|
||||
|
||||
use litellm_auth::SecretValue;
|
||||
|
|
@ -22,6 +23,32 @@ pub struct ChatCompletionsRequest<'a> {
|
|||
pub timeout: Option<Duration>,
|
||||
}
|
||||
|
||||
pub struct ChatCompletionsCall {
|
||||
pub model: String,
|
||||
pub messages: Value,
|
||||
pub optional_params: 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 timeout: Option<Duration>,
|
||||
}
|
||||
|
||||
impl From<ChatCompletionsRequest<'_>> for ChatCompletionsCall {
|
||||
fn from(request: ChatCompletionsRequest<'_>) -> Self {
|
||||
Self {
|
||||
model: request.model.into(),
|
||||
messages: request.messages,
|
||||
optional_params: request.optional_params,
|
||||
api_key: request.api_key.map(str::to_owned),
|
||||
api_base: request.api_base.map(str::to_owned),
|
||||
custom_llm_provider: request.custom_llm_provider.map(str::to_owned),
|
||||
extra_headers: request.extra_headers,
|
||||
timeout: request.timeout,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub struct ResolvedChatCompletionsRequest<'a> {
|
||||
pub model: String,
|
||||
pub custom_llm_provider: String,
|
||||
|
|
@ -46,6 +73,7 @@ pub struct ProviderChatCompletionsRequest {
|
|||
/// The forwarded and default headers plus how the call authenticates; the credential
|
||||
/// itself is applied when the request is sent.
|
||||
pub environment: ValidatedEnvironment,
|
||||
pub secrets: Secrets,
|
||||
pub timeout: Option<Duration>,
|
||||
pub api_key: Option<SecretValue>,
|
||||
}
|
||||
|
|
|
|||
324
litellm-rust/crates/core/src/diagnostic.rs
Normal file
324
litellm-rust/crates/core/src/diagnostic.rs
Normal file
|
|
@ -0,0 +1,324 @@
|
|||
use std::{
|
||||
future::Future,
|
||||
pin::Pin,
|
||||
task::{Context, Poll},
|
||||
};
|
||||
|
||||
use futures_util::{Stream, stream::BoxStream};
|
||||
use litellm_host::call::CallOutput;
|
||||
use litellm_tracing::Logger;
|
||||
use tracing::Span;
|
||||
|
||||
struct Completion {
|
||||
span: Span,
|
||||
outcome: &'static str,
|
||||
}
|
||||
|
||||
impl Completion {
|
||||
fn new(name: &str) -> Self {
|
||||
let current = Span::current();
|
||||
Self {
|
||||
span: if current
|
||||
.metadata()
|
||||
.is_some_and(|metadata| metadata.name() == name)
|
||||
{
|
||||
current
|
||||
} else {
|
||||
Span::none()
|
||||
},
|
||||
outcome: "cancelled",
|
||||
}
|
||||
}
|
||||
|
||||
fn finish(mut self, outcome: &'static str) {
|
||||
self.outcome = outcome;
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for Completion {
|
||||
fn drop(&mut self) {
|
||||
self.span.record("outcome", self.outcome);
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn provider(model: &str, provider: &str) {
|
||||
let span = Span::current();
|
||||
span.record("resolved_model", model);
|
||||
span.record("provider", provider);
|
||||
}
|
||||
|
||||
pub(crate) async fn unary<R, E>(execute: impl Future<Output = Result<R, E>>) -> Result<R, E> {
|
||||
operation("litellm.route", execute).await
|
||||
}
|
||||
|
||||
pub(crate) async fn operation<R, E>(
|
||||
name: &str,
|
||||
execute: impl Future<Output = Result<R, E>>,
|
||||
) -> Result<R, E> {
|
||||
let completion = Completion::new(name);
|
||||
let result = execute.await;
|
||||
completion.finish(if result.is_ok() { "success" } else { "failure" });
|
||||
result
|
||||
}
|
||||
|
||||
pub(crate) async fn call<R, H, C, E>(
|
||||
execute: impl Future<Output = Result<CallOutput<R, H, C, E>, E>>,
|
||||
) -> Result<CallOutput<R, H, C, E>, E>
|
||||
where
|
||||
C: Send + 'static,
|
||||
E: Send + 'static,
|
||||
{
|
||||
let completion = Completion::new("litellm.route");
|
||||
match execute.await {
|
||||
Err(error) => {
|
||||
completion.finish("failure");
|
||||
Err(error)
|
||||
}
|
||||
Ok(CallOutput::Complete(response)) => {
|
||||
completion.span.record("stream", false);
|
||||
completion.finish("success");
|
||||
Ok(CallOutput::Complete(response))
|
||||
}
|
||||
Ok(CallOutput::Stream { head, chunks }) => {
|
||||
completion.span.record("stream", true);
|
||||
Ok(CallOutput::Stream {
|
||||
head,
|
||||
chunks: Box::pin(TracedStream {
|
||||
state: Some(StreamState {
|
||||
chunks,
|
||||
completion,
|
||||
logger: Logger::current(),
|
||||
}),
|
||||
}),
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
struct StreamState<C, E> {
|
||||
chunks: BoxStream<'static, Result<C, E>>,
|
||||
completion: Completion,
|
||||
logger: Logger,
|
||||
}
|
||||
|
||||
impl<C, E> StreamState<C, E> {
|
||||
fn close(self, outcome: &'static str) {
|
||||
self.logger
|
||||
.scope(|| self.completion.span.in_scope(|| drop(self.chunks)));
|
||||
self.completion.finish(outcome);
|
||||
}
|
||||
}
|
||||
|
||||
struct TracedStream<C, E> {
|
||||
state: Option<StreamState<C, E>>,
|
||||
}
|
||||
|
||||
impl<C, E> Stream for TracedStream<C, E> {
|
||||
type Item = Result<C, E>;
|
||||
|
||||
fn poll_next(mut self: Pin<&mut Self>, context: &mut Context<'_>) -> Poll<Option<Self::Item>> {
|
||||
let Some(state) = self.state.as_mut() else {
|
||||
return Poll::Ready(None);
|
||||
};
|
||||
let next = state.logger.scope(|| {
|
||||
state
|
||||
.completion
|
||||
.span
|
||||
.in_scope(|| state.chunks.as_mut().poll_next(context))
|
||||
});
|
||||
let outcome = match &next {
|
||||
Poll::Ready(None) => "success",
|
||||
Poll::Ready(Some(Err(_))) => "failure",
|
||||
_ => return next,
|
||||
};
|
||||
if let Some(state) = self.state.take() {
|
||||
state.close(outcome);
|
||||
}
|
||||
next
|
||||
}
|
||||
}
|
||||
|
||||
impl<C, E> Drop for TracedStream<C, E> {
|
||||
fn drop(&mut self) {
|
||||
if let Some(state) = self.state.take() {
|
||||
state.close("cancelled");
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::{sync::mpsc, task::Context};
|
||||
|
||||
use futures_util::{StreamExt, task::noop_waker_ref};
|
||||
use litellm_tracing::{Metadata, Record, Sink};
|
||||
use rstest::{fixture, rstest};
|
||||
use serde_json::{Value, json};
|
||||
|
||||
use super::*;
|
||||
|
||||
struct Capture(mpsc::Sender<Value>);
|
||||
|
||||
impl Sink for Capture {
|
||||
fn enabled(&self, metadata: &Metadata<'_>) -> bool {
|
||||
*metadata.level() <= tracing::Level::INFO
|
||||
}
|
||||
fn emit(&self, record: &Record) {
|
||||
self.0.send(Value::Object(record.fields.clone())).unwrap();
|
||||
}
|
||||
}
|
||||
|
||||
#[fixture]
|
||||
fn logger() -> (Logger, mpsc::Receiver<Value>) {
|
||||
let (sender, receiver) = mpsc::channel();
|
||||
(Logger::new(Capture(sender)), receiver)
|
||||
}
|
||||
|
||||
struct Chunks(std::vec::IntoIter<Result<u8, &'static str>>);
|
||||
|
||||
impl Stream for Chunks {
|
||||
type Item = Result<u8, &'static str>;
|
||||
fn poll_next(mut self: Pin<&mut Self>, _: &mut Context<'_>) -> Poll<Option<Self::Item>> {
|
||||
tracing::info!(event = "poll");
|
||||
Poll::Ready(self.0.next())
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for Chunks {
|
||||
fn drop(&mut self) {
|
||||
tracing::info!(event = "drop");
|
||||
}
|
||||
}
|
||||
|
||||
#[tracing::instrument(
|
||||
name = "litellm.route",
|
||||
skip_all,
|
||||
fields(route = "fixture", stream, outcome)
|
||||
)]
|
||||
async fn streamed() -> Result<CallOutput<(), (), u8, &'static str>, &'static str> {
|
||||
call(async {
|
||||
Ok(CallOutput::Stream {
|
||||
head: (),
|
||||
chunks: Box::pin(Chunks(vec![Ok(1), Err("broken"), Ok(2)].into_iter())),
|
||||
})
|
||||
})
|
||||
.await
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn stream_errors_finish_once_and_poll_and_drop_use_the_captured_context(
|
||||
logger: (Logger, mpsc::Receiver<Value>),
|
||||
) {
|
||||
let (logger, records) = logger;
|
||||
let CallOutput::Stream { mut chunks, .. } = logger.instrument(streamed()).await.unwrap()
|
||||
else {
|
||||
panic!()
|
||||
};
|
||||
assert!(records.try_recv().is_err());
|
||||
tokio::spawn(async move {
|
||||
assert_eq!(chunks.next().await, Some(Ok(1)));
|
||||
assert_eq!(chunks.next().await, Some(Err("broken")));
|
||||
assert_eq!(chunks.next().await, None);
|
||||
let emitted = records.try_iter().collect::<Vec<_>>();
|
||||
assert_eq!(emitted.len(), 4);
|
||||
assert!(emitted.iter().all(|record| record["route"] == "fixture"));
|
||||
assert_eq!(emitted[0]["event"], "poll");
|
||||
assert_eq!(emitted[1]["event"], "poll");
|
||||
assert_eq!(emitted[2]["event"], "drop");
|
||||
assert_eq!(emitted[3]["outcome"], "failure");
|
||||
drop(chunks);
|
||||
assert!(records.try_recv().is_err());
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
}
|
||||
|
||||
#[tracing::instrument(name = "litellm.route", skip_all, fields(route = "waiting", outcome))]
|
||||
async fn waiting(streaming: bool) {
|
||||
if streaming {
|
||||
let _: Result<CallOutput<(), (), u8, ()>, ()> = call(std::future::pending()).await;
|
||||
} else {
|
||||
let _: Result<(), ()> = unary(std::future::pending()).await;
|
||||
}
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::unary(false)]
|
||||
#[case::streaming(true)]
|
||||
fn cancellation_before_headers_closes_the_span(
|
||||
logger: (Logger, mpsc::Receiver<Value>),
|
||||
#[case] streaming: bool,
|
||||
) {
|
||||
let (logger, records) = logger;
|
||||
let mut future = Box::pin(logger.instrument(waiting(streaming)));
|
||||
assert!(
|
||||
future
|
||||
.as_mut()
|
||||
.poll(&mut Context::from_waker(noop_waker_ref()))
|
||||
.is_pending()
|
||||
);
|
||||
assert!(records.try_recv().is_err());
|
||||
drop(future);
|
||||
let summary = records.try_recv().unwrap();
|
||||
assert_eq!(summary["outcome"], "cancelled");
|
||||
assert_eq!(summary["route"], "waiting");
|
||||
assert!(records.try_recv().is_err());
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn dropped_stream_teardown_uses_its_original_logger(
|
||||
logger: (Logger, mpsc::Receiver<Value>),
|
||||
) {
|
||||
let (logger, records) = logger;
|
||||
let output = logger.instrument(streamed()).await.unwrap();
|
||||
Logger::default().scope(|| drop(output));
|
||||
let emitted = records.try_iter().collect::<Vec<_>>();
|
||||
assert_eq!(emitted.len(), 2);
|
||||
assert_eq!(
|
||||
emitted[0],
|
||||
json!({"route":"fixture", "stream":true, "event":"drop"})
|
||||
);
|
||||
assert_eq!(emitted[1]["outcome"], "cancelled");
|
||||
}
|
||||
|
||||
#[tracing::instrument(name = "disabled", level = "debug", skip_all, fields(outcome))]
|
||||
async fn disabled_child() {
|
||||
let _: Result<(), ()> = operation("disabled", async { Err(()) }).await;
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn a_filtered_operation_does_not_overwrite_its_parent_outcome(
|
||||
logger: (Logger, mpsc::Receiver<Value>),
|
||||
) {
|
||||
let (logger, records) = logger;
|
||||
logger
|
||||
.instrument(async {
|
||||
let parent = tracing::info_span!("parent", outcome = "original");
|
||||
tracing::Instrument::instrument(disabled_child(), parent).await;
|
||||
})
|
||||
.await;
|
||||
let summary = records.try_recv().unwrap();
|
||||
assert_eq!(summary["span_name"], "parent");
|
||||
assert_eq!(summary["outcome"], "original");
|
||||
assert!(records.try_recv().is_err());
|
||||
}
|
||||
|
||||
#[tracing::instrument(name = "litellm.route", skip_all, fields(stream = true, outcome))]
|
||||
async fn completed() -> Result<CallOutput<(), (), u8, ()>, ()> {
|
||||
call(async { Ok(CallOutput::Complete(())) }).await
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn streaming_mode_reflects_the_returned_output(logger: (Logger, mpsc::Receiver<Value>)) {
|
||||
let (logger, records) = logger;
|
||||
logger.instrument(completed()).await.unwrap();
|
||||
let summary = records.try_recv().unwrap();
|
||||
assert_eq!(summary["stream"], false);
|
||||
assert_eq!(summary["outcome"], "success");
|
||||
assert!(records.try_recv().is_err());
|
||||
}
|
||||
}
|
||||
|
|
@ -22,9 +22,9 @@ pub enum RouteError {
|
|||
#[error("invalid provider: {0}")]
|
||||
InvalidProvider(String),
|
||||
#[error("invalid request: {0}")]
|
||||
InvalidRequest(String),
|
||||
InvalidRequest(#[source] litellm_llms::ErrorDetail),
|
||||
#[error("invalid response: {0}")]
|
||||
InvalidResponse(String),
|
||||
InvalidResponse(#[source] litellm_llms::ErrorDetail),
|
||||
#[error("unsupported by the rust path: {0}")]
|
||||
Unsupported(&'static str),
|
||||
#[error(transparent)]
|
||||
|
|
@ -37,34 +37,23 @@ pub enum RouteError {
|
|||
Http(#[from] litellm_http::Error),
|
||||
#[error(transparent)]
|
||||
Secret(#[from] SecretError),
|
||||
#[error("post-call hook failed: {0}")]
|
||||
PostCallHook(#[source] Arc<RouteError>),
|
||||
}
|
||||
|
||||
/// 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 From<litellm_host::machine::MachineFault> for RouteError {
|
||||
fn from(fault: litellm_host::machine::MachineFault) -> Self {
|
||||
use litellm_host::machine::MachineFault;
|
||||
Self::InvalidRequest(match fault {
|
||||
MachineFault::Abandoned => "host driver was abandoned".into(),
|
||||
MachineFault::Protocol(message) => format!("host {message}").into(),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
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,
|
||||
}
|
||||
pub(crate) fn post_call(error: Self) -> Self {
|
||||
Self::PostCallHook(Arc::new(error))
|
||||
}
|
||||
|
||||
/// The caller's request is what is wrong, as opposed to the environment, the wire, or
|
||||
|
|
@ -78,9 +67,11 @@ impl RouteError {
|
|||
| Self::Unsupported(_)
|
||||
| Self::Headers(_) => true,
|
||||
Self::Auth(error) => !matches!(error, litellm_auth::Error::MissingApiKey { .. }),
|
||||
Self::InvalidResponse(_) | Self::Transport(_) | Self::Http(_) | Self::Secret(_) => {
|
||||
false
|
||||
}
|
||||
Self::InvalidResponse(_)
|
||||
| Self::Transport(_)
|
||||
| Self::Http(_)
|
||||
| Self::Secret(_)
|
||||
| Self::PostCallHook(_) => false,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -124,31 +115,9 @@ 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:?}");
|
||||
}
|
||||
}
|
||||
use super::RouteError;
|
||||
use litellm_llms::{Error as LlmError, ErrorDetail};
|
||||
use rstest::rstest;
|
||||
|
||||
#[test]
|
||||
fn a_missing_api_key_is_the_environment_not_the_request() {
|
||||
|
|
@ -163,4 +132,29 @@ mod tests {
|
|||
assert!(RouteError::InvalidRequest("top_k".into()).is_request());
|
||||
assert!(!RouteError::InvalidResponse("bad json".into()).is_request());
|
||||
}
|
||||
#[rstest]
|
||||
#[case::request(true)]
|
||||
#[case::response(false)]
|
||||
fn contextual_errors_preserve_sources_and_route_classification(#[case] request: bool) {
|
||||
let source = serde_json::from_str::<serde_json::Value>("{").unwrap_err();
|
||||
let source_message = source.to_string();
|
||||
let detail = ErrorDetail::invalid("test payload", source);
|
||||
let error = RouteError::from(if request {
|
||||
LlmError::InvalidRequest(detail)
|
||||
} else {
|
||||
LlmError::InvalidResponse(detail)
|
||||
});
|
||||
assert_eq!(error.is_request(), request);
|
||||
let category = if request { "request" } else { "response" };
|
||||
assert_eq!(
|
||||
error.to_string(),
|
||||
format!("invalid {category}: invalid test payload: {source_message}")
|
||||
);
|
||||
let source = std::iter::successors(Some(&error as &dyn std::error::Error), |error| {
|
||||
error.source()
|
||||
})
|
||||
.find_map(|error| error.downcast_ref::<serde_json::Error>())
|
||||
.expect("the original JSON error remains available");
|
||||
assert_eq!(source.to_string(), source_message);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,3 +1,5 @@
|
|||
mod diagnostic;
|
||||
|
||||
pub mod audio_transcription;
|
||||
pub mod chat_completions;
|
||||
pub mod constants;
|
||||
|
|
@ -5,7 +7,8 @@ pub mod error;
|
|||
pub mod messages;
|
||||
pub mod ocr;
|
||||
mod outbound;
|
||||
mod provider;
|
||||
pub mod resources;
|
||||
pub mod responses;
|
||||
|
||||
pub use error::{Phase, RouteError};
|
||||
pub use error::RouteError;
|
||||
|
|
|
|||
7
litellm-rust/crates/core/src/messages/AGENTS.md
Normal file
7
litellm-rust/crates/core/src/messages/AGENTS.md
Normal file
|
|
@ -0,0 +1,7 @@
|
|||
This directory owns provider-independent Messages call orchestration: the entrypoint, call envelopes, provider selection, credential resolution, transport coordination, hooks, and stream lifecycle. Shared API data contracts belong in `litellm-types::messages`, adapter contracts and execution inputs in `llms/src/base_llm/messages`, and provider implementations in `llms/src/<provider>/messages`
|
||||
|
||||
Select concrete provider adapters and invoke their contracts. Delegate authentication policy, beta selection, payload rewriting, and response interpretation to those adapters. Keep provider policy out of request preparation and transport handlers. Calling a concrete provider helper for every provider is still a policy dependency
|
||||
|
||||
Route types such as `MessagesCall`, prepared requests, and response wrappers containing live streams describe execution. Reuse the shared Messages payload types inside them instead of defining another request or response schema here
|
||||
|
||||
Preserve the order of validation, normalization, caller-requested parameter removal, and provider transformation when that order affects observable behavior. Test provider dispatch, auth precedence, header handling, transformations, and responses through behavior, not source structure
|
||||
|
|
@ -2,19 +2,18 @@ use litellm_http::request::string_headers as shared_string_headers;
|
|||
pub(super) use litellm_http::request::truncate_error_body;
|
||||
use litellm_llms::{
|
||||
anthropic::messages::transformation::ANTHROPIC_MESSAGES_CONFIG,
|
||||
azure_ai::anthropic::messages_transformation::AZURE_ANTHROPIC_MESSAGES_CONFIG,
|
||||
base_llm::anthropic_messages::transformation::BaseAnthropicMessagesConfig,
|
||||
azure_ai::messages::transformation::AZURE_ANTHROPIC_MESSAGES_CONFIG,
|
||||
base_llm::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;
|
||||
use crate::provider::LlmProviders;
|
||||
|
||||
const HEADER_CONTEXT: &str = "messages";
|
||||
|
||||
#[derive(Clone, Copy, Debug, EnumString, IntoStaticStr, PartialEq, Eq)]
|
||||
#[strum(serialize_all = "snake_case")]
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||||
pub(crate) enum MessagesProvider {
|
||||
Anthropic,
|
||||
AzureAi,
|
||||
|
|
@ -23,7 +22,12 @@ pub(crate) enum MessagesProvider {
|
|||
|
||||
impl MessagesProvider {
|
||||
pub(crate) fn as_str(self) -> &'static str {
|
||||
self.into()
|
||||
match self {
|
||||
Self::Anthropic => LlmProviders::Anthropic,
|
||||
Self::AzureAi => LlmProviders::AzureAi,
|
||||
Self::Bedrock => LlmProviders::Bedrock,
|
||||
}
|
||||
.into()
|
||||
}
|
||||
|
||||
pub(crate) fn config(self) -> &'static dyn BaseAnthropicMessagesConfig {
|
||||
|
|
@ -35,6 +39,21 @@ impl MessagesProvider {
|
|||
}
|
||||
}
|
||||
|
||||
pub(crate) fn messages_provider(provider: LlmProviders) -> Option<MessagesProvider> {
|
||||
match provider {
|
||||
LlmProviders::Anthropic => Some(MessagesProvider::Anthropic),
|
||||
LlmProviders::AzureAi => Some(MessagesProvider::AzureAi),
|
||||
LlmProviders::Bedrock => Some(MessagesProvider::Bedrock),
|
||||
LlmProviders::AwsTextract
|
||||
| LlmProviders::Cohere
|
||||
| LlmProviders::Mistral
|
||||
| LlmProviders::Openai
|
||||
| LlmProviders::OpenaiLike
|
||||
| LlmProviders::Reducto
|
||||
| LlmProviders::VertexAi => None,
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn string_headers(
|
||||
extra_headers: Option<Map<String, Value>>,
|
||||
) -> Result<Vec<(String, String)>, Error> {
|
||||
|
|
@ -47,8 +66,9 @@ mod tests {
|
|||
|
||||
use rstest::rstest;
|
||||
|
||||
use super::{MessagesProvider, string_headers, truncate_error_body};
|
||||
use super::{MessagesProvider, messages_provider, string_headers, truncate_error_body};
|
||||
use crate::messages::Error;
|
||||
use crate::provider::LlmProviders;
|
||||
|
||||
#[rstest]
|
||||
#[case::anthropic("anthropic", MessagesProvider::Anthropic)]
|
||||
|
|
@ -58,13 +78,16 @@ mod tests {
|
|||
#[case] name: &str,
|
||||
#[case] provider: MessagesProvider,
|
||||
) {
|
||||
assert_eq!(name.parse::<MessagesProvider>(), Ok(provider));
|
||||
assert_eq!(
|
||||
messages_provider(name.parse::<LlmProviders>().unwrap()),
|
||||
Some(provider)
|
||||
);
|
||||
assert_eq!(provider.as_str(), name);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_without_a_messages_config_is_rejected() {
|
||||
assert!("openai".parse::<MessagesProvider>().is_err());
|
||||
assert_eq!(messages_provider(LlmProviders::Openai), None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
|
|
|||
|
|
@ -1,21 +1,20 @@
|
|||
use litellm_host::lifecycle::ExecutionEvent;
|
||||
use litellm_host::observation::ObservationSender;
|
||||
use std::time::Duration;
|
||||
|
||||
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_host::interceptors::{Interceptors, RawResponse, RequestContext, WireRequest};
|
||||
use litellm_http::transport::Error as TransportError;
|
||||
use litellm_llms::base_llm::{
|
||||
anthropic_messages::{
|
||||
auth::{Authenticated, resolve_auth},
|
||||
messages::{
|
||||
streaming::{ByteStream, StreamDecoder, encode_anthropic_sse},
|
||||
transformation::BaseAnthropicMessagesConfig,
|
||||
},
|
||||
auth::{Authenticated, resolve_auth},
|
||||
};
|
||||
use litellm_tracing::{ByteChunk, debug};
|
||||
use litellm_tracing::ByteChunk;
|
||||
use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse;
|
||||
use serde_json::Value;
|
||||
|
||||
|
|
@ -28,7 +27,8 @@ pub(super) async fn execute(
|
|||
http: &litellm_http::Client,
|
||||
auth: &AuthServices,
|
||||
request: ProviderMessagesRequest,
|
||||
hooks: &impl RouteHooks<Error>,
|
||||
interceptors: &impl Interceptors<Error>,
|
||||
observers: Option<&ObservationSender>,
|
||||
) -> Result<MessagesResponse, Error> {
|
||||
let ProviderMessagesRequest {
|
||||
provider,
|
||||
|
|
@ -47,8 +47,8 @@ pub(super) async fn execute(
|
|||
api_key,
|
||||
};
|
||||
let authenticated = resolve_auth(auth, environment, &|key| std::env::var(key).ok()).await?;
|
||||
let wire = hooks
|
||||
.before_send(
|
||||
let wire = interceptors
|
||||
.before_provider_request(
|
||||
WireRequest {
|
||||
url,
|
||||
headers: authenticated.headers,
|
||||
|
|
@ -58,7 +58,7 @@ pub(super) async fn execute(
|
|||
)
|
||||
.await?;
|
||||
let provider_name = provider.as_str();
|
||||
debug!(provider = provider_name, stream, body = %wire.body, "provider request");
|
||||
log_request_body(provider_name, stream, &wire.body);
|
||||
let response = send(
|
||||
http,
|
||||
Authenticated {
|
||||
|
|
@ -70,11 +70,6 @@ pub(super) async fn execute(
|
|||
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);
|
||||
}
|
||||
|
|
@ -87,19 +82,25 @@ pub(super) async fn execute(
|
|||
));
|
||||
}
|
||||
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?;
|
||||
log_response_body(&text);
|
||||
let raw = RawResponse { body: text.clone() };
|
||||
if let Some(observers) = observers {
|
||||
observers.emit(litellm_host::lifecycle::CallEvent::Execution(
|
||||
ExecutionEvent::ProviderResponseReceived { raw: raw.clone() },
|
||||
));
|
||||
}
|
||||
interceptors
|
||||
.after_provider_response(raw)
|
||||
.await
|
||||
.map_err(Error::post_call)?;
|
||||
decode_response(config, &body.model, &text)
|
||||
.map(|message| MessagesResponse::Message(Box::new(message)))
|
||||
.map(|message| MessagesResponse::Complete(Box::new(message)))
|
||||
}
|
||||
|
||||
fn serialize_failure(err: serde_json::Error) -> Error {
|
||||
Error::InvalidRequest(format!(
|
||||
"failed to serialize Anthropic messages request: {err}"
|
||||
Error::InvalidRequest(litellm_llms::ErrorDetail::failed(
|
||||
"Anthropic messages request serialization",
|
||||
err,
|
||||
))
|
||||
}
|
||||
|
||||
|
|
@ -120,14 +121,14 @@ async fn send(
|
|||
body,
|
||||
Some(timeout.unwrap_or(Duration::from_secs(MESSAGES_TIMEOUT_SECS))),
|
||||
)?;
|
||||
request.send(http).await.map_err(network)
|
||||
crate::outbound::send(request, http).await.map_err(network)
|
||||
}
|
||||
|
||||
async fn provider_error(response: reqwest::Response) -> Error {
|
||||
let status = response.status().as_u16();
|
||||
match response.text().await {
|
||||
Ok(text) => {
|
||||
litellm_tracing::debug!(status, body = text.as_str(), "provider error body");
|
||||
log_error_body(status, &text);
|
||||
Error::Transport(TransportError::Http {
|
||||
status,
|
||||
body: truncate_error_body(&text),
|
||||
|
|
@ -142,8 +143,12 @@ fn decode_response(
|
|||
model: &str,
|
||||
text: &str,
|
||||
) -> Result<AnthropicMessagesResponse, Error> {
|
||||
let response = serde_json::from_str(text)
|
||||
.map_err(|err| Error::InvalidResponse(format!("invalid messages response JSON: {err}")))?;
|
||||
let response = serde_json::from_str(text).map_err(|err| {
|
||||
Error::InvalidResponse(litellm_llms::ErrorDetail::invalid(
|
||||
"messages response JSON",
|
||||
err,
|
||||
))
|
||||
})?;
|
||||
config
|
||||
.transform_anthropic_messages_response(model, response)
|
||||
.map_err(Error::from)
|
||||
|
|
@ -170,7 +175,10 @@ fn streaming_response(
|
|||
.boxed(),
|
||||
Some(decode) => decoded_chunks(response, decode, provider),
|
||||
};
|
||||
MessagesResponse::Stream { headers, chunks }
|
||||
MessagesResponse::Stream {
|
||||
head: super::route::MessagesStreamHead { headers },
|
||||
chunks,
|
||||
}
|
||||
}
|
||||
|
||||
fn decoded_chunks(
|
||||
|
|
@ -194,14 +202,26 @@ fn decoded_chunks(
|
|||
.boxed()
|
||||
}
|
||||
|
||||
fn log_chunk(provider: &str, stage: &str, data: &Bytes) {
|
||||
fn log_request_body(provider: &str, stream: bool, body: &serde_json::Value) {
|
||||
tracing::debug!(provider, stream, body = %body, "provider request");
|
||||
}
|
||||
|
||||
fn log_response_body(body: &str) {
|
||||
tracing::debug!(body, "provider response body");
|
||||
}
|
||||
|
||||
fn log_error_body(status: u16, body: &str) {
|
||||
tracing::debug!(status, body, "provider error body");
|
||||
}
|
||||
|
||||
fn log_chunk(provider: &str, stage: &str, data: &bytes::Bytes) {
|
||||
let chunk = ByteChunk::new(data);
|
||||
debug!(provider, stage, encoding = chunk.encoding(), chunk = %chunk, "stream chunk");
|
||||
tracing::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 litellm_llms::base_llm::messages::streaming::anthropic_sse_event_stream;
|
||||
use rstest::rstest;
|
||||
use wiremock::{Mock, MockServer, ResponseTemplate, matchers::any};
|
||||
|
||||
|
|
@ -212,6 +232,7 @@ mod tests {
|
|||
"data: {\"type\":\"ping\"}\n\n",
|
||||
Some("event: ping\ndata: {\"type\":\"ping\"}\n\n")
|
||||
)]
|
||||
#[rstest::rstest]
|
||||
#[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(
|
||||
|
|
|
|||
|
|
@ -1,27 +1,77 @@
|
|||
//! The Anthropic Messages call, the Rust equivalent of Python's `litellm.messages()`.
|
||||
//!
|
||||
//! [`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.
|
||||
|
||||
use litellm_host::observation::ObservationSender;
|
||||
mod common_utils;
|
||||
mod handler;
|
||||
mod prepare;
|
||||
pub mod route;
|
||||
mod types;
|
||||
|
||||
use litellm_http::{ClientVariant, HttpClientConfig};
|
||||
use litellm_auth::AuthServices;
|
||||
use litellm_secrets::source::SecretSource;
|
||||
use std::sync::Arc;
|
||||
|
||||
pub use crate::error::RouteError as Error;
|
||||
pub use types::{MessagesCall, MessagesResponse, MessagesShaping, messages_body};
|
||||
|
||||
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
|
||||
#[derive(Clone)]
|
||||
pub struct MessagesRoute {
|
||||
http: litellm_http::Client,
|
||||
auth: Arc<AuthServices>,
|
||||
secrets: Arc<dyn SecretSource>,
|
||||
}
|
||||
|
||||
impl MessagesRoute {
|
||||
pub fn new(
|
||||
http: litellm_http::Client,
|
||||
auth: Arc<AuthServices>,
|
||||
secrets: Arc<dyn SecretSource>,
|
||||
) -> Self {
|
||||
Self {
|
||||
http,
|
||||
auth,
|
||||
secrets,
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn execute(
|
||||
&self,
|
||||
call: MessagesCall,
|
||||
interceptors: &impl litellm_host::interceptors::Interceptors<Error>,
|
||||
observers: Option<ObservationSender>,
|
||||
) -> Result<MessagesResponse, Error> {
|
||||
litellm_host::lifecycle::observe_call(
|
||||
observers.clone(),
|
||||
self.run(call, interceptors, observers.as_ref()),
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
#[tracing::instrument(name = "litellm.route", skip_all, fields(
|
||||
route = "messages",
|
||||
model = %call.body.model,
|
||||
provider,
|
||||
resolved_model,
|
||||
stream = call.body.params.stream == Some(true),
|
||||
outcome
|
||||
))]
|
||||
async fn run(
|
||||
&self,
|
||||
call: MessagesCall,
|
||||
interceptors: &impl litellm_host::interceptors::Interceptors<Error>,
|
||||
observers: Option<&ObservationSender>,
|
||||
) -> Result<MessagesResponse, Error> {
|
||||
crate::diagnostic::call(async {
|
||||
let request = prepare::prepare(call, self.secrets.as_ref()).await?;
|
||||
crate::diagnostic::provider(&request.body.model, request.provider.as_str());
|
||||
let execute: futures_util::future::BoxFuture<'_, Result<MessagesResponse, Error>> =
|
||||
Box::pin(handler::execute(
|
||||
&self.http,
|
||||
&self.auth,
|
||||
request,
|
||||
interceptors,
|
||||
observers,
|
||||
));
|
||||
execute.await
|
||||
})
|
||||
.await
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -3,25 +3,21 @@ 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},
|
||||
get_provider_specific_headers::get_provider_specific_headers,
|
||||
settings::Lookup,
|
||||
get_provider_specific_headers::get_provider_specific_headers, settings::Lookup,
|
||||
};
|
||||
use litellm_llms::{
|
||||
anthropic::messages::handler::shape_anthropic_messages_request,
|
||||
base_llm::{
|
||||
anthropic_messages::transformation::MessagesTransformContext,
|
||||
auth::{ValidatedEnvironment, with_default_headers},
|
||||
},
|
||||
use litellm_http::request::with_default_headers;
|
||||
use litellm_llms::base_llm::{
|
||||
auth::ValidatedEnvironment, messages::context::MessagesTransformContext,
|
||||
};
|
||||
use litellm_secrets::source::SecretSource;
|
||||
use litellm_types::llms::anthropic_messages::anthropic_request::AnthropicMessagesRequest;
|
||||
|
||||
use super::{
|
||||
Error, MessagesCall,
|
||||
common_utils::{MessagesProvider, string_headers},
|
||||
common_utils::{MessagesProvider, messages_provider, string_headers},
|
||||
types::invalid_request,
|
||||
};
|
||||
use crate::provider::resolve_llm_provider;
|
||||
|
||||
struct ResolvedProvider {
|
||||
model: String,
|
||||
|
|
@ -38,6 +34,7 @@ pub(super) struct ProviderMessagesRequest {
|
|||
pub(super) api_key: Option<SecretValue>,
|
||||
}
|
||||
|
||||
#[tracing::instrument(name = "litellm.prepare", level = "debug", skip_all)]
|
||||
pub(super) async fn prepare(
|
||||
call: MessagesCall,
|
||||
secrets: &dyn SecretSource,
|
||||
|
|
@ -53,26 +50,11 @@ fn resolve_provider(
|
|||
model: &str,
|
||||
custom_llm_provider: Option<&str>,
|
||||
) -> Result<ResolvedProvider, Error> {
|
||||
let CustomLlmProvider {
|
||||
model,
|
||||
custom_llm_provider: provider,
|
||||
} = get_custom_llm_provider(model, custom_llm_provider)
|
||||
.or_else(|| {
|
||||
custom_llm_provider.map(|provider| CustomLlmProvider {
|
||||
model,
|
||||
custom_llm_provider: provider,
|
||||
})
|
||||
})
|
||||
.ok_or_else(|| {
|
||||
Error::InvalidProvider(
|
||||
"unable to resolve custom_llm_provider for messages request".to_string(),
|
||||
)
|
||||
})?;
|
||||
let provider = provider
|
||||
.parse()
|
||||
.map_err(|_| Error::InvalidProvider(provider.to_string()))?;
|
||||
let resolved = resolve_llm_provider(model, custom_llm_provider, "messages")?;
|
||||
let provider = messages_provider(resolved.provider)
|
||||
.ok_or_else(|| Error::InvalidProvider(<&str>::from(resolved.provider).to_string()))?;
|
||||
Ok(ResolvedProvider {
|
||||
model: model.to_string(),
|
||||
model: resolved.model.to_string(),
|
||||
provider,
|
||||
})
|
||||
}
|
||||
|
|
@ -96,7 +78,7 @@ fn prepare_provider_request(
|
|||
let config = provider.config();
|
||||
let env_lookup = |key: &str| secrets.get(key);
|
||||
|
||||
let sanitized = shape_anthropic_messages_request(
|
||||
let sanitized = config.shape_request(
|
||||
AnthropicMessagesRequest { model, ..body },
|
||||
shaping.reasoning_auto_summary,
|
||||
)?;
|
||||
|
|
@ -468,7 +450,11 @@ mod tests {
|
|||
shaping,
|
||||
),
|
||||
Err(Error::InvalidRequest(
|
||||
"metadata.user_id must be a string, got 123".to_string()
|
||||
litellm_llms::ErrorDetail::InvalidValue {
|
||||
field: "metadata.user_id",
|
||||
expected: "a string",
|
||||
actual: json!(123),
|
||||
}
|
||||
))
|
||||
);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,26 +1,16 @@
|
|||
use std::{
|
||||
convert::Infallible,
|
||||
sync::{Arc, Mutex},
|
||||
};
|
||||
use litellm_host::observation::ObservationSender;
|
||||
use std::convert::Infallible;
|
||||
|
||||
use bytes::Bytes;
|
||||
use futures_util::TryStreamExt;
|
||||
use litellm_host::{
|
||||
host::{Demand, Host},
|
||||
machine::{CallMachine, HostChannel, MachineFault},
|
||||
call::{HostedCompletion, HostedMachine, hosted_call},
|
||||
protocol::Protocol,
|
||||
};
|
||||
use litellm_http::{Client, ClientVariant, HttpClientConfig};
|
||||
use litellm_secrets::source::SecretSource;
|
||||
use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse;
|
||||
|
||||
use super::{Error, MessagesCall, MessagesResponse, handler::execute, prepare::prepare};
|
||||
use super::{Error, MessagesCall};
|
||||
|
||||
pub enum MessagesOutput {
|
||||
Message(Box<AnthropicMessagesResponse>),
|
||||
/// Every chunk already reached the host through `Deliver`.
|
||||
Streamed,
|
||||
}
|
||||
pub type MessagesOutput = HostedCompletion<Box<AnthropicMessagesResponse>>;
|
||||
|
||||
/// The upstream response as the caller sees it at stream hand-off, before any chunk.
|
||||
pub struct MessagesStreamHead {
|
||||
|
|
@ -30,91 +20,28 @@ pub struct MessagesStreamHead {
|
|||
pub struct Messages;
|
||||
|
||||
impl Protocol for Messages {
|
||||
type Response = MessagesOutput;
|
||||
type Response = Box<AnthropicMessagesResponse>;
|
||||
type Error = Error;
|
||||
type Projection = MessagesCall;
|
||||
type Op = Infallible;
|
||||
type Request = MessagesCall;
|
||||
type HostCall = Infallible;
|
||||
type Chunk = Bytes;
|
||||
type StreamHead = MessagesStreamHead;
|
||||
}
|
||||
|
||||
impl From<MachineFault> for Error {
|
||||
fn from(fault: MachineFault) -> Self {
|
||||
Self::InvalidRequest(match fault {
|
||||
MachineFault::Abandoned => "messages host driver was abandoned".into(),
|
||||
MachineFault::Protocol(message) => format!("messages {message}"),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
pub type MessagesHost = HostChannel<Messages>;
|
||||
pub type MessagesMachine = CallMachine<Messages>;
|
||||
|
||||
/// The in-process host for a request already in hand. It answers projection once and
|
||||
/// observes nothing.
|
||||
pub struct LocalMessagesHost {
|
||||
call: Mutex<Option<MessagesCall>>,
|
||||
}
|
||||
|
||||
impl LocalMessagesHost {
|
||||
pub fn new(call: MessagesCall) -> Self {
|
||||
Self {
|
||||
call: Mutex::new(Some(call)),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Host<Messages> for LocalMessagesHost {
|
||||
async fn project(&self) -> Result<MessagesCall, Error> {
|
||||
self.call
|
||||
.lock()
|
||||
.unwrap_or_else(|error| error.into_inner())
|
||||
.take()
|
||||
.ok_or_else(|| Error::InvalidRequest("messages request was already projected".into()))
|
||||
}
|
||||
|
||||
async fn custom_op(&self, op: Infallible) -> Result<(), Error> {
|
||||
match op {}
|
||||
}
|
||||
}
|
||||
|
||||
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))
|
||||
}))
|
||||
}
|
||||
|
||||
/// 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 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)
|
||||
}
|
||||
pub type MessagesMachine = HostedMachine<Messages>;
|
||||
|
||||
impl super::MessagesRoute {
|
||||
pub fn machine(
|
||||
self,
|
||||
request: super::MessagesCall,
|
||||
observers: Option<ObservationSender>,
|
||||
) -> MessagesMachine {
|
||||
hosted_call(
|
||||
request,
|
||||
observers,
|
||||
move |call, _, interceptors, observers| async move {
|
||||
self.run(call, &interceptors, observers.as_ref()).await
|
||||
},
|
||||
)
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,8 +1,8 @@
|
|||
use std::time::Duration;
|
||||
|
||||
use bytes::Bytes;
|
||||
use futures_util::stream::BoxStream;
|
||||
use litellm_llms::anthropic::common_utils::AnthropicModelCapabilities;
|
||||
use litellm_host::call::CallOutput;
|
||||
use litellm_llms::base_llm::messages::context::MessagesModelCapabilities as AnthropicModelCapabilities;
|
||||
use litellm_types::{
|
||||
llms::anthropic_messages::{
|
||||
anthropic_request::AnthropicMessagesRequest, anthropic_response::AnthropicMessagesResponse,
|
||||
|
|
@ -30,16 +30,11 @@ pub fn messages_body(body: Map<String, Value>) -> Result<AnthropicMessagesReques
|
|||
}
|
||||
|
||||
pub(super) fn invalid_request(err: serde_json::Error) -> Error {
|
||||
Error::InvalidRequest(format!("invalid Anthropic messages request: {err}"))
|
||||
Error::InvalidRequest(format!("invalid Anthropic messages request: {err}").into())
|
||||
}
|
||||
|
||||
pub enum MessagesResponse {
|
||||
Message(Box<AnthropicMessagesResponse>),
|
||||
Stream {
|
||||
headers: Vec<(String, String)>,
|
||||
chunks: BoxStream<'static, Result<Bytes, Error>>,
|
||||
},
|
||||
}
|
||||
pub type MessagesResponse =
|
||||
CallOutput<Box<AnthropicMessagesResponse>, super::route::MessagesStreamHead, Bytes, Error>;
|
||||
|
||||
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
|
||||
pub struct MessagesShaping {
|
||||
|
|
@ -55,7 +50,7 @@ pub struct MessagesShaping {
|
|||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use litellm_llms::anthropic::common_utils::SupportedEffortTiers;
|
||||
use litellm_llms::base_llm::messages::context::SupportedEffortTiers;
|
||||
use rstest::rstest;
|
||||
use serde_json::{Value, json};
|
||||
|
||||
|
|
|
|||
|
|
@ -1,15 +1,82 @@
|
|||
use litellm_host::observation::ObservationSender;
|
||||
use std::sync::Arc;
|
||||
|
||||
use litellm_host::interceptors::Interceptors;
|
||||
use litellm_llms::base_llm::ocr::{
|
||||
error::Error, handler::OcrClient, transformation::LiteLLMOcrResponse,
|
||||
};
|
||||
|
||||
use crate::ocr::{
|
||||
route::{LocalOcrHost, ocr_machine},
|
||||
types::LiteLLMOcrRequest,
|
||||
use super::{
|
||||
handler::perform_ocr_request,
|
||||
types::{LiteLLMOcrRequest, OcrDocumentInput, ResolvedOcrRequest},
|
||||
};
|
||||
|
||||
pub async fn perform(
|
||||
client: &OcrClient,
|
||||
request: LiteLLMOcrRequest,
|
||||
) -> Result<LiteLLMOcrResponse, Error> {
|
||||
litellm_host::run::run(ocr_machine(client.clone()), &LocalOcrHost::new(request)).await
|
||||
#[derive(Clone)]
|
||||
pub struct OcrRoute {
|
||||
client: OcrClient,
|
||||
}
|
||||
|
||||
impl OcrRoute {
|
||||
pub fn new(client: OcrClient) -> Self {
|
||||
Self { client }
|
||||
}
|
||||
|
||||
pub async fn execute(
|
||||
&self,
|
||||
request: LiteLLMOcrRequest,
|
||||
interceptors: &impl Interceptors<Error>,
|
||||
observers: Option<ObservationSender>,
|
||||
) -> Result<LiteLLMOcrResponse, Error> {
|
||||
litellm_host::lifecycle::observe_unary(
|
||||
observers.clone(),
|
||||
self.run(request, interceptors, observers.as_ref()),
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
#[tracing::instrument(name = "litellm.route", skip_all, fields(
|
||||
route = "ocr",
|
||||
model = %request.model,
|
||||
resolved_model = %request.model,
|
||||
provider = <&str>::from(request.config.provider()),
|
||||
stream = false,
|
||||
outcome
|
||||
))]
|
||||
pub(super) async fn run(
|
||||
&self,
|
||||
request: LiteLLMOcrRequest,
|
||||
interceptors: &impl Interceptors<Error>,
|
||||
observers: Option<&ObservationSender>,
|
||||
) -> Result<LiteLLMOcrResponse, Error> {
|
||||
crate::diagnostic::unary(async {
|
||||
let caller_document = matches!(&request.document, OcrDocumentInput::Document(_));
|
||||
let prepared = prepare_request_document(request).await?;
|
||||
let execute: futures_util::future::BoxFuture<'_, Result<LiteLLMOcrResponse, Error>> =
|
||||
Box::pin(perform_ocr_request(
|
||||
&self.client,
|
||||
prepared,
|
||||
interceptors,
|
||||
caller_document,
|
||||
observers,
|
||||
));
|
||||
execute.await
|
||||
})
|
||||
.await
|
||||
}
|
||||
}
|
||||
|
||||
#[tracing::instrument(name = "litellm.prepare", level = "debug", skip_all)]
|
||||
async fn prepare_request_document(
|
||||
request: LiteLLMOcrRequest<OcrDocumentInput>,
|
||||
) -> Result<ResolvedOcrRequest, Error> {
|
||||
if let OcrDocumentInput::Document(_) = &request.document {
|
||||
return request.map_document(super::document::prepare_document);
|
||||
}
|
||||
let logger = litellm_tracing::Logger::current();
|
||||
let span = tracing::Span::current();
|
||||
tokio::task::spawn_blocking(move || {
|
||||
logger.scope(|| span.in_scope(|| request.map_document(super::document::prepare_document)))
|
||||
})
|
||||
.await
|
||||
.map_err(|error| Error::DocumentTask(Arc::new(error)))?
|
||||
}
|
||||
|
|
|
|||
|
|
@ -110,6 +110,7 @@ mod tests {
|
|||
use std::collections::BTreeMap as Map;
|
||||
|
||||
use litellm_llms::base_llm::ocr::document::InlineDocument;
|
||||
use rstest::rstest;
|
||||
|
||||
use super::*;
|
||||
|
||||
|
|
@ -135,24 +136,21 @@ mod tests {
|
|||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn file_name_mime_mapping_matches_python() {
|
||||
for (name, expected) in [
|
||||
("document.pdf", "application/pdf"),
|
||||
("image.png", "image/png"),
|
||||
("photo.jpg", "image/jpeg"),
|
||||
("photo.jpeg", "image/jpeg"),
|
||||
("animation.gif", "image/gif"),
|
||||
("image.webp", "image/webp"),
|
||||
("scan.tiff", "image/tiff"),
|
||||
("scan.tif", "image/tiff"),
|
||||
("bitmap.bmp", "image/bmp"),
|
||||
("DOCUMENT.PDF", "application/pdf"),
|
||||
("IMAGE.PNG", "image/png"),
|
||||
("file.unknown-extension", "application/octet-stream"),
|
||||
] {
|
||||
assert_eq!(mime_type_for_name(name), expected);
|
||||
}
|
||||
#[rstest]
|
||||
#[case::pdf("document.pdf", "application/pdf")]
|
||||
#[case::png("image.png", "image/png")]
|
||||
#[case::jpg("photo.jpg", "image/jpeg")]
|
||||
#[case::jpeg("photo.jpeg", "image/jpeg")]
|
||||
#[case::gif("animation.gif", "image/gif")]
|
||||
#[case::webp("image.webp", "image/webp")]
|
||||
#[case::tiff("scan.tiff", "image/tiff")]
|
||||
#[case::tif("scan.tif", "image/tiff")]
|
||||
#[case::bmp("bitmap.bmp", "image/bmp")]
|
||||
#[case::uppercase_pdf("DOCUMENT.PDF", "application/pdf")]
|
||||
#[case::uppercase_png("IMAGE.PNG", "image/png")]
|
||||
#[case::unknown("file.unknown-extension", "application/octet-stream")]
|
||||
fn file_name_mime_mapping_matches_python(#[case] name: &str, #[case] expected: &str) {
|
||||
assert_eq!(mime_type_for_name(name), expected);
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
|
|
|||
|
|
@ -1,6 +1,7 @@
|
|||
use futures_util::future::BoxFuture;
|
||||
use litellm_auth::SecretValue;
|
||||
use litellm_host::event::{MachineEvent, RawResponse, RequestContext, WireRequest};
|
||||
use litellm_host::interceptors::{Interceptors, RawResponse, RequestContext, WireRequest};
|
||||
use litellm_host::lifecycle::ExecutionEvent;
|
||||
use litellm_host::observation::ObservationSender;
|
||||
use litellm_llms::base_llm::ocr::{
|
||||
error::Error,
|
||||
handler::{CallHooks, OcrClient},
|
||||
|
|
@ -8,17 +9,15 @@ use litellm_llms::base_llm::ocr::{
|
|||
};
|
||||
use serde_json::Value;
|
||||
|
||||
use super::{
|
||||
arguments::is_secret_param, prepare::prepare_request, provider_config::OcrConfigKind,
|
||||
route::OcrHost,
|
||||
};
|
||||
use super::{arguments::is_secret_param, prepare::prepare_request, provider_config::OcrConfigKind};
|
||||
use crate::ocr::types::ResolvedOcrRequest;
|
||||
|
||||
pub(crate) async fn perform_ocr_request(
|
||||
client: &OcrClient,
|
||||
request: ResolvedOcrRequest,
|
||||
host: &OcrHost,
|
||||
host: &impl Interceptors<Error>,
|
||||
caller_document: bool,
|
||||
observers: Option<&ObservationSender>,
|
||||
) -> Result<LiteLLMOcrResponse, Error> {
|
||||
request.response_format()?;
|
||||
let config = request.config;
|
||||
|
|
@ -28,56 +27,62 @@ pub(crate) async fn perform_ocr_request(
|
|||
.await
|
||||
.map_err(|error| Error::Secret(std::sync::Arc::new(error)))?;
|
||||
let request = prepare_request(request, caller_document, client, secrets);
|
||||
let hooks = OcrCallHooks::new(host.clone(), &request, config);
|
||||
config.ocr(client, &request, &hooks).await
|
||||
let interceptors = OcrCallHooks::new(host, &request, config, observers);
|
||||
config.ocr(client, &request, &interceptors).await
|
||||
}
|
||||
|
||||
/// Lets provider code reach the host mid-call, filling in the request context only the
|
||||
/// route knows.
|
||||
pub(crate) struct OcrCallHooks {
|
||||
host: OcrHost,
|
||||
model: String,
|
||||
custom_llm_provider: &'static str,
|
||||
optional_params: Value,
|
||||
secret_fields: Vec<String>,
|
||||
api_key: Option<SecretValue>,
|
||||
struct OcrCallHooks<'a, H> {
|
||||
interceptors: &'a H,
|
||||
context: RequestContext,
|
||||
observers: Option<&'a ObservationSender>,
|
||||
}
|
||||
|
||||
impl OcrCallHooks {
|
||||
pub(crate) fn new(host: OcrHost, request: &PreparedOcrRequest, config: OcrConfigKind) -> Self {
|
||||
impl<'a, H> OcrCallHooks<'a, H> {
|
||||
fn new(
|
||||
interceptors: &'a H,
|
||||
request: &PreparedOcrRequest,
|
||||
config: OcrConfigKind,
|
||||
observers: Option<&'a ObservationSender>,
|
||||
) -> Self {
|
||||
Self {
|
||||
host,
|
||||
model: request.model.clone(),
|
||||
custom_llm_provider: config.provider().into(),
|
||||
optional_params: Value::Object(request.optional_params.clone().into()),
|
||||
secret_fields: request
|
||||
.optional_params
|
||||
.keys()
|
||||
.filter(|name| is_secret_param(name))
|
||||
.cloned()
|
||||
.collect(),
|
||||
api_key: request.connection.api_key.clone(),
|
||||
interceptors,
|
||||
observers,
|
||||
context: RequestContext {
|
||||
model: request.model.clone(),
|
||||
custom_llm_provider: <&str>::from(config.provider()).to_owned(),
|
||||
optional_params: Value::Object(request.optional_params.clone().into()),
|
||||
secret_fields: request
|
||||
.optional_params
|
||||
.keys()
|
||||
.filter(|name| is_secret_param(name))
|
||||
.cloned()
|
||||
.collect(),
|
||||
api_key: request.connection.api_key.clone(),
|
||||
},
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl CallHooks<Error> for OcrCallHooks {
|
||||
fn before_send(&self, wire: WireRequest) -> BoxFuture<'_, Result<WireRequest, Error>> {
|
||||
let context = RequestContext {
|
||||
model: self.model.clone(),
|
||||
custom_llm_provider: self.custom_llm_provider.into(),
|
||||
optional_params: self.optional_params.clone(),
|
||||
secret_fields: self.secret_fields.clone(),
|
||||
api_key: self.api_key.clone(),
|
||||
};
|
||||
Box::pin(self.host.before_send(wire, context))
|
||||
impl<H: Interceptors<Error>> CallHooks<Error> for OcrCallHooks<'_, H> {
|
||||
fn before_provider_request(
|
||||
&self,
|
||||
wire: WireRequest,
|
||||
) -> BoxFuture<'_, Result<WireRequest, Error>> {
|
||||
Box::pin(
|
||||
self.interceptors
|
||||
.before_provider_request(wire, self.context.clone()),
|
||||
)
|
||||
}
|
||||
|
||||
fn response_received<'a>(&'a self, body: &'a [u8]) -> BoxFuture<'a, Result<(), Error>> {
|
||||
Box::pin(self.host.emit(MachineEvent::ResponseReceived {
|
||||
raw: RawResponse {
|
||||
body: String::from_utf8_lossy(body).into_owned(),
|
||||
},
|
||||
}))
|
||||
let raw = RawResponse {
|
||||
body: String::from_utf8_lossy(body).into_owned(),
|
||||
};
|
||||
if let Some(observers) = self.observers {
|
||||
observers.emit(litellm_host::lifecycle::CallEvent::Execution(
|
||||
ExecutionEvent::ProviderResponseReceived { raw: raw.clone() },
|
||||
));
|
||||
}
|
||||
Box::pin(self.interceptors.after_provider_response(raw))
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
pub mod arguments;
|
||||
pub mod client;
|
||||
mod client;
|
||||
pub use client::OcrRoute;
|
||||
pub mod document;
|
||||
pub(crate) mod handler;
|
||||
pub(crate) mod prepare;
|
||||
|
|
|
|||
|
|
@ -5,8 +5,8 @@ use litellm_llms::base_llm::ocr::{
|
|||
};
|
||||
use litellm_secrets::source::Secrets;
|
||||
|
||||
use super::provider_config::OcrProvider;
|
||||
use crate::ocr::types::{LiteLLMOcrRequest, ResolvedOcrRequest};
|
||||
use crate::provider::LlmProviders;
|
||||
|
||||
pub(crate) fn prepare_request(
|
||||
request: ResolvedOcrRequest,
|
||||
|
|
@ -16,15 +16,19 @@ pub(crate) fn prepare_request(
|
|||
) -> PreparedOcrRequest {
|
||||
let credentials = request.credentials.clone();
|
||||
let (preferred_api_key_env, api_base_env) = match request.config.provider() {
|
||||
OcrProvider::Mistral => (
|
||||
LlmProviders::Mistral => (
|
||||
Some("MISTRAL_AZURE_API_KEY"),
|
||||
Some("MISTRAL_AZURE_API_BASE"),
|
||||
),
|
||||
OcrProvider::AzureAi => (None, Some("AZURE_AI_API_BASE")),
|
||||
OcrProvider::AwsTextract
|
||||
| OcrProvider::Cohere
|
||||
| OcrProvider::Reducto
|
||||
| OcrProvider::VertexAi => (None, None),
|
||||
LlmProviders::AzureAi => (None, Some("AZURE_AI_API_BASE")),
|
||||
LlmProviders::Anthropic
|
||||
| LlmProviders::AwsTextract
|
||||
| LlmProviders::Bedrock
|
||||
| LlmProviders::Cohere
|
||||
| LlmProviders::Openai
|
||||
| LlmProviders::OpenaiLike
|
||||
| LlmProviders::Reducto
|
||||
| LlmProviders::VertexAi => (None, None),
|
||||
};
|
||||
let secret = |name: &str| secrets.truthy(name);
|
||||
let dynamic_api_key = credentials.dynamic_api_key.or_else(|| {
|
||||
|
|
@ -76,7 +80,7 @@ mod tests {
|
|||
|
||||
use futures_util::future::BoxFuture;
|
||||
use litellm_core_utils::call_arguments::{CallArguments, compose_body, parse_options};
|
||||
use litellm_host::event::WireRequest;
|
||||
use litellm_host::interceptors::WireRequest;
|
||||
use litellm_llms::{
|
||||
base_llm::ocr::{
|
||||
error::Error,
|
||||
|
|
@ -96,11 +100,14 @@ mod tests {
|
|||
wire::{OcrWireRequest, decode_request},
|
||||
};
|
||||
|
||||
/// Stands in for a host with no hooks registered.
|
||||
/// Stands in for a host with no interceptors registered.
|
||||
struct NoHooks;
|
||||
|
||||
impl CallHooks<Error> for NoHooks {
|
||||
fn before_send(&self, wire: WireRequest) -> BoxFuture<'_, Result<WireRequest, Error>> {
|
||||
fn before_provider_request(
|
||||
&self,
|
||||
wire: WireRequest,
|
||||
) -> BoxFuture<'_, Result<WireRequest, Error>> {
|
||||
Box::pin(async move { Ok(wire) })
|
||||
}
|
||||
|
||||
|
|
@ -144,6 +151,7 @@ mod tests {
|
|||
json!({"type": "image_url", "image_url": url})
|
||||
}
|
||||
|
||||
#[rstest::rstest]
|
||||
#[tokio::test]
|
||||
async fn cohere_body_keeps_native_document_fields_and_untyped_overrides() {
|
||||
let request = request(
|
||||
|
|
@ -176,6 +184,7 @@ mod tests {
|
|||
);
|
||||
}
|
||||
|
||||
#[rstest::rstest]
|
||||
#[tokio::test]
|
||||
async fn explicit_null_options_use_defaults_before_http() {
|
||||
let request = request(
|
||||
|
|
@ -199,6 +208,7 @@ mod tests {
|
|||
assert!(body.get("req_format").is_none());
|
||||
}
|
||||
|
||||
#[rstest::rstest]
|
||||
#[tokio::test]
|
||||
async fn direct_and_vertex_mistral_build_the_same_request_and_share_normalization() {
|
||||
let options = json!({
|
||||
|
|
@ -276,7 +286,7 @@ mod tests {
|
|||
pages: Option<Vec<i64>>,
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[rstest::rstest]
|
||||
fn parsed_provider_params_separates_known_and_extra_params() {
|
||||
let arguments: CallArguments = serde_json::from_value(json!({
|
||||
"pages": [0, 2],
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
use crate::provider::LlmProviders;
|
||||
use litellm_core_utils::get_llm_provider_logic::{CustomLlmProvider, get_custom_llm_provider};
|
||||
use litellm_llms::{
|
||||
aws_textract::ocr::{
|
||||
|
|
@ -24,7 +25,6 @@ use litellm_llms::{
|
|||
deepseek_transformation::VertexAIDeepSeekOCRConfig, transformation::VertexAiOcrConfig,
|
||||
},
|
||||
};
|
||||
use strum::{EnumString, IntoStaticStr};
|
||||
|
||||
macro_rules! with_config {
|
||||
($kind:expr, $config:ident => $body:expr) => {
|
||||
|
|
@ -93,16 +93,16 @@ pub(crate) enum OcrConfigKind {
|
|||
}
|
||||
|
||||
impl OcrConfigKind {
|
||||
pub(crate) const fn provider(self) -> OcrProvider {
|
||||
pub(crate) const fn provider(self) -> LlmProviders {
|
||||
match self {
|
||||
Self::AwsTextract | Self::AwsTextractAnalyze => OcrProvider::AwsTextract,
|
||||
Self::Cohere => OcrProvider::Cohere,
|
||||
Self::Mistral => OcrProvider::Mistral,
|
||||
Self::AwsTextract | Self::AwsTextractAnalyze => LlmProviders::AwsTextract,
|
||||
Self::Cohere => LlmProviders::Cohere,
|
||||
Self::Mistral => LlmProviders::Mistral,
|
||||
Self::AzureAi | Self::AzureCohere | Self::AzureDocumentIntelligence => {
|
||||
OcrProvider::AzureAi
|
||||
LlmProviders::AzureAi
|
||||
}
|
||||
Self::ReductoLegacy | Self::ReductoV3 => OcrProvider::Reducto,
|
||||
Self::VertexAi | Self::VertexDeepSeek => OcrProvider::VertexAi,
|
||||
Self::ReductoLegacy | Self::ReductoV3 => LlmProviders::Reducto,
|
||||
Self::VertexAi | Self::VertexDeepSeek => LlmProviders::VertexAi,
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -133,9 +133,9 @@ impl OcrConfigKind {
|
|||
self,
|
||||
client: &OcrClient,
|
||||
request: &PreparedOcrRequest,
|
||||
hooks: &dyn CallHooks<Error>,
|
||||
interceptors: &dyn CallHooks<Error>,
|
||||
) -> Result<LiteLLMOcrResponse, Error> {
|
||||
with_config!(self, config => handler::ocr(&config, client, request, hooks).await)
|
||||
with_config!(self, config => handler::ocr(&config, client, request, interceptors).await)
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -187,17 +187,6 @@ pub fn passthrough_response(
|
|||
.map(Some)
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, EnumString, IntoStaticStr, PartialEq, Eq)]
|
||||
#[strum(serialize_all = "snake_case")]
|
||||
pub(crate) enum OcrProvider {
|
||||
AwsTextract,
|
||||
Cohere,
|
||||
Mistral,
|
||||
AzureAi,
|
||||
Reducto,
|
||||
VertexAi,
|
||||
}
|
||||
|
||||
pub(crate) fn resolve_provider_config(
|
||||
model: &str,
|
||||
custom_llm_provider: Option<&str>,
|
||||
|
|
@ -205,37 +194,45 @@ pub(crate) fn resolve_provider_config(
|
|||
let provider =
|
||||
get_custom_llm_provider(model, custom_llm_provider).unwrap_or(CustomLlmProvider {
|
||||
model,
|
||||
custom_llm_provider: OcrProvider::Mistral.into(),
|
||||
custom_llm_provider: LlmProviders::Mistral.into(),
|
||||
});
|
||||
let ocr_provider = provider
|
||||
let llm_provider = provider
|
||||
.custom_llm_provider
|
||||
.parse::<OcrProvider>()
|
||||
.parse::<LlmProviders>()
|
||||
.map_err(|_| Error::InvalidProvider(provider.custom_llm_provider.to_string()))?;
|
||||
let config = match ocr_provider {
|
||||
OcrProvider::AwsTextract => match TextractOperation::from_model(provider.model)? {
|
||||
let config = match llm_provider {
|
||||
LlmProviders::AwsTextract => match TextractOperation::from_model(provider.model)? {
|
||||
TextractOperation::DetectDocumentText => OcrConfigKind::AwsTextract,
|
||||
TextractOperation::AnalyzeDocument => OcrConfigKind::AwsTextractAnalyze,
|
||||
},
|
||||
OcrProvider::Cohere => OcrConfigKind::Cohere,
|
||||
OcrProvider::Mistral => OcrConfigKind::Mistral,
|
||||
OcrProvider::AzureAi if is_document_intelligence_model(provider.model) => {
|
||||
LlmProviders::Cohere => OcrConfigKind::Cohere,
|
||||
LlmProviders::Mistral => OcrConfigKind::Mistral,
|
||||
LlmProviders::AzureAi if is_document_intelligence_model(provider.model) => {
|
||||
OcrConfigKind::AzureDocumentIntelligence
|
||||
}
|
||||
OcrProvider::AzureAi
|
||||
LlmProviders::AzureAi
|
||||
if provider.model.to_ascii_lowercase().contains("cohere")
|
||||
&& provider.model.to_ascii_lowercase().contains("parse") =>
|
||||
{
|
||||
OcrConfigKind::AzureCohere
|
||||
}
|
||||
OcrProvider::AzureAi => OcrConfigKind::AzureAi,
|
||||
OcrProvider::Reducto if provider.model.eq_ignore_ascii_case("parse-legacy") => {
|
||||
LlmProviders::AzureAi => OcrConfigKind::AzureAi,
|
||||
LlmProviders::Reducto if provider.model.eq_ignore_ascii_case("parse-legacy") => {
|
||||
OcrConfigKind::ReductoLegacy
|
||||
}
|
||||
OcrProvider::Reducto => OcrConfigKind::ReductoV3,
|
||||
OcrProvider::VertexAi if provider.model.to_ascii_lowercase().contains("deepseek") => {
|
||||
LlmProviders::Reducto => OcrConfigKind::ReductoV3,
|
||||
LlmProviders::VertexAi if provider.model.to_ascii_lowercase().contains("deepseek") => {
|
||||
OcrConfigKind::VertexDeepSeek
|
||||
}
|
||||
OcrProvider::VertexAi => OcrConfigKind::VertexAi,
|
||||
LlmProviders::VertexAi => OcrConfigKind::VertexAi,
|
||||
LlmProviders::Anthropic
|
||||
| LlmProviders::Bedrock
|
||||
| LlmProviders::Openai
|
||||
| LlmProviders::OpenaiLike => {
|
||||
return Err(Error::InvalidProvider(
|
||||
provider.custom_llm_provider.to_string(),
|
||||
));
|
||||
}
|
||||
};
|
||||
Ok((provider.model.to_string(), config))
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,25 +1,21 @@
|
|||
use std::sync::{Arc, Mutex};
|
||||
|
||||
use litellm_auth::ResolvedCredential;
|
||||
use litellm_auth::{ResolvedCredential, TokenProviderHandle};
|
||||
use litellm_host::observation::ObservationSender;
|
||||
use litellm_host::{
|
||||
event::{CallEvent, RequestContext, WireRequest},
|
||||
host::Reply,
|
||||
machine::{CallMachine, HostChannel, HostTokenProvider, TokenProtocol},
|
||||
call::{CallOutput, HostedMachine, hosted_call},
|
||||
machine::HostServices,
|
||||
protocol::Protocol,
|
||||
protocol::Reply,
|
||||
};
|
||||
use litellm_llms::base_llm::ocr::{
|
||||
error::Error, handler::OcrClient, transformation::LiteLLMOcrResponse,
|
||||
};
|
||||
use litellm_llms::base_llm::ocr::{error::Error, transformation::LiteLLMOcrResponse};
|
||||
|
||||
use super::handler::perform_ocr_request;
|
||||
use crate::ocr::types::{LiteLLMOcrRequest, OcrDocumentInput, ResolvedOcrRequest};
|
||||
use crate::ocr::types::{LiteLLMOcrRequest, OcrDocumentInput};
|
||||
|
||||
pub enum OcrOp {
|
||||
AcquireAzureAdToken(Reply<ResolvedCredential>),
|
||||
}
|
||||
|
||||
/// The caller's request as the host projects it.
|
||||
pub struct OcrProjection {
|
||||
pub struct OcrCall {
|
||||
pub request: LiteLLMOcrRequest<OcrDocumentInput>,
|
||||
/// The caller passed its own Azure AD token provider, which the host keeps.
|
||||
pub caller_token: bool,
|
||||
|
|
@ -30,134 +26,45 @@ pub struct Ocr;
|
|||
impl Protocol for Ocr {
|
||||
type Response = LiteLLMOcrResponse;
|
||||
type Error = Error;
|
||||
type Projection = OcrProjection;
|
||||
type Op = OcrOp;
|
||||
type Request = OcrCall;
|
||||
type HostCall = OcrOp;
|
||||
type Chunk = std::convert::Infallible;
|
||||
type StreamHead = std::convert::Infallible;
|
||||
}
|
||||
|
||||
impl TokenProtocol for Ocr {
|
||||
fn acquire_token_op(reply: Reply<ResolvedCredential>) -> OcrOp {
|
||||
OcrOp::AcquireAzureAdToken(reply)
|
||||
}
|
||||
}
|
||||
pub type OcrMachine = HostedMachine<Ocr>;
|
||||
|
||||
pub type OcrHost = HostChannel<Ocr>;
|
||||
pub type OcrMachine = CallMachine<Ocr>;
|
||||
|
||||
/// The OCR call as a machine: projection and token acquisition are host operations;
|
||||
/// everything else runs in Rust.
|
||||
pub fn ocr_machine(client: OcrClient) -> OcrMachine {
|
||||
CallMachine::new(move |host| Box::pin(execute(client, host)))
|
||||
}
|
||||
|
||||
async fn execute(client: OcrClient, host: OcrHost) -> Result<LiteLLMOcrResponse, Error> {
|
||||
let OcrProjection {
|
||||
request,
|
||||
caller_token,
|
||||
} = host.project().await?;
|
||||
let request = LiteLLMOcrRequest {
|
||||
azure_ad_token_provider: caller_token
|
||||
.then(|| HostTokenProvider::handle(host.clone()))
|
||||
.or(request.azure_ad_token_provider),
|
||||
..request
|
||||
};
|
||||
let caller_document = matches!(request.document, OcrDocumentInput::Document(_));
|
||||
let request = prepare_request_document(request).await?;
|
||||
perform_ocr_request(&client, request, &host, caller_document).await
|
||||
}
|
||||
|
||||
async fn prepare_request_document(
|
||||
request: LiteLLMOcrRequest<OcrDocumentInput>,
|
||||
) -> Result<ResolvedOcrRequest, Error> {
|
||||
if let OcrDocumentInput::Document(_) = &request.document {
|
||||
return request.map_document(super::document::prepare_document);
|
||||
}
|
||||
tokio::task::spawn_blocking(move || request.map_document(super::document::prepare_document))
|
||||
.await
|
||||
.map_err(|error| Error::DocumentTask(Arc::new(error)))?
|
||||
}
|
||||
|
||||
type BeforeSend =
|
||||
Box<dyn Fn(WireRequest, &RequestContext) -> Result<WireRequest, Error> + Send + Sync>;
|
||||
type Observer = Box<dyn Fn(&CallEvent) + Send + Sync>;
|
||||
|
||||
/// The in-process host for a request that is already in hand: the request answers
|
||||
/// projection, and the optional observer sees and may rewrite the wire request.
|
||||
pub struct LocalOcrHost {
|
||||
request: Mutex<Option<LiteLLMOcrRequest<OcrDocumentInput>>>,
|
||||
before_send: Option<BeforeSend>,
|
||||
observer: Option<Observer>,
|
||||
}
|
||||
|
||||
impl LocalOcrHost {
|
||||
pub fn new(request: LiteLLMOcrRequest<OcrDocumentInput>) -> Self {
|
||||
Self {
|
||||
request: Mutex::new(Some(request)),
|
||||
before_send: None,
|
||||
observer: None,
|
||||
fn caller_token_provider(services: HostServices<Ocr>) -> TokenProviderHandle {
|
||||
TokenProviderHandle::from_callback(move || {
|
||||
let host_services = services.clone();
|
||||
async move {
|
||||
host_services
|
||||
.call(OcrOp::AcquireAzureAdToken)
|
||||
.await
|
||||
.map_err(|error| {
|
||||
litellm_auth::Error::CredentialAcquisition(error.to_string().into())
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
pub fn with_before_send(
|
||||
self,
|
||||
before_send: impl Fn(WireRequest, &RequestContext) -> Result<WireRequest, Error>
|
||||
+ Send
|
||||
+ Sync
|
||||
+ 'static,
|
||||
) -> Self {
|
||||
Self {
|
||||
before_send: Some(Box::new(before_send)),
|
||||
..self
|
||||
}
|
||||
}
|
||||
|
||||
pub fn with_observer(self, observer: impl Fn(&CallEvent) + Send + Sync + 'static) -> Self {
|
||||
Self {
|
||||
observer: Some(Box::new(observer)),
|
||||
..self
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
impl litellm_host::host::Host<Ocr> for LocalOcrHost {
|
||||
async fn project(&self) -> Result<OcrProjection, Error> {
|
||||
self.request
|
||||
.lock()
|
||||
.unwrap_or_else(|error| error.into_inner())
|
||||
.take()
|
||||
.map(|request| OcrProjection {
|
||||
request,
|
||||
caller_token: false,
|
||||
})
|
||||
.ok_or_else(|| Error::InvalidRequest("OCR request was already projected".into()))
|
||||
}
|
||||
|
||||
async fn custom_op(&self, op: OcrOp) -> Result<(), Error> {
|
||||
match op {
|
||||
OcrOp::AcquireAzureAdToken(_) => {
|
||||
Err(Error::Auth(litellm_auth::Error::AzureTokenAcquisition(
|
||||
"OCR host has no Azure AD token provider".into(),
|
||||
)))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
async fn before_send(
|
||||
&self,
|
||||
wire: WireRequest,
|
||||
context: &RequestContext,
|
||||
) -> Result<WireRequest, Error> {
|
||||
match &self.before_send {
|
||||
Some(before_send) => before_send(wire, context),
|
||||
None => Ok(wire),
|
||||
}
|
||||
}
|
||||
|
||||
async fn emit(&self, event: &CallEvent) -> Result<(), Error> {
|
||||
if let Some(observer) = &self.observer {
|
||||
observer(event);
|
||||
}
|
||||
Ok(())
|
||||
impl crate::ocr::OcrRoute {
|
||||
pub fn machine(self, request: OcrCall, observers: Option<ObservationSender>) -> OcrMachine {
|
||||
hosted_call(
|
||||
request,
|
||||
observers,
|
||||
move |projection: OcrCall, services, interceptors, observers| async move {
|
||||
let request = LiteLLMOcrRequest {
|
||||
azure_ad_token_provider: projection
|
||||
.caller_token
|
||||
.then(|| caller_token_provider(services))
|
||||
.or(projection.request.azure_ad_token_provider),
|
||||
..projection.request
|
||||
};
|
||||
self.run(request, &interceptors, observers.as_ref())
|
||||
.await
|
||||
.map(CallOutput::Complete)
|
||||
},
|
||||
)
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -4,6 +4,21 @@ use litellm_http::outbound::OutboundRequest;
|
|||
use litellm_llms::base_llm::auth::Authenticated;
|
||||
use serde_json::Value;
|
||||
|
||||
#[tracing::instrument(
|
||||
name = "litellm.provider.send",
|
||||
level = "debug",
|
||||
skip_all,
|
||||
fields(status)
|
||||
)]
|
||||
pub(crate) async fn send(
|
||||
request: OutboundRequest,
|
||||
client: &litellm_http::Client,
|
||||
) -> Result<reqwest::Response, reqwest::Error> {
|
||||
request.send(client).await.inspect(|response| {
|
||||
tracing::Span::current().record("status", response.status().as_u16());
|
||||
})
|
||||
}
|
||||
|
||||
/// Header credentials are already in `headers`; SigV4 is applied here, over the
|
||||
/// bytes that are sent.
|
||||
pub(crate) fn outbound_request(
|
||||
|
|
|
|||
36
litellm-rust/crates/core/src/provider.rs
Normal file
36
litellm-rust/crates/core/src/provider.rs
Normal file
|
|
@ -0,0 +1,36 @@
|
|||
pub use litellm_core_utils::get_llm_provider_logic::LlmProviders;
|
||||
use litellm_core_utils::get_llm_provider_logic::{CustomLlmProvider, get_custom_llm_provider};
|
||||
|
||||
use crate::error::RouteError as Error;
|
||||
|
||||
#[derive(Debug)]
|
||||
pub(crate) struct ResolvedProvider<'a> {
|
||||
pub(crate) model: &'a str,
|
||||
pub(crate) provider: LlmProviders,
|
||||
}
|
||||
|
||||
pub(crate) fn resolve_llm_provider<'a>(
|
||||
model: &'a str,
|
||||
custom_llm_provider: Option<&'a str>,
|
||||
route: &'static str,
|
||||
) -> Result<ResolvedProvider<'a>, Error> {
|
||||
let CustomLlmProvider {
|
||||
model,
|
||||
custom_llm_provider,
|
||||
} = get_custom_llm_provider(model, custom_llm_provider)
|
||||
.or_else(|| {
|
||||
custom_llm_provider.map(|provider| CustomLlmProvider {
|
||||
model,
|
||||
custom_llm_provider: provider,
|
||||
})
|
||||
})
|
||||
.ok_or_else(|| {
|
||||
Error::InvalidProvider(format!(
|
||||
"unable to resolve custom_llm_provider for {route} request"
|
||||
))
|
||||
})?;
|
||||
let provider = custom_llm_provider
|
||||
.parse()
|
||||
.map_err(|_| Error::InvalidProvider(custom_llm_provider.to_string()))?;
|
||||
Ok(ResolvedProvider { model, provider })
|
||||
}
|
||||
|
|
@ -1,9 +1,7 @@
|
|||
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;
|
||||
use litellm_http::HttpClientPool;
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct CoreResources {
|
||||
|
|
@ -18,21 +16,4 @@ impl CoreResources {
|
|||
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,
|
||||
)
|
||||
}
|
||||
}
|
||||
|
|
|
|||
95
litellm-rust/crates/core/src/responses/handler.rs
Normal file
95
litellm-rust/crates/core/src/responses/handler.rs
Normal file
|
|
@ -0,0 +1,95 @@
|
|||
use litellm_host::lifecycle::ExecutionEvent;
|
||||
use litellm_host::observation::ObservationSender;
|
||||
use std::time::Duration;
|
||||
|
||||
use futures_util::StreamExt;
|
||||
use litellm_host::interceptors::{Interceptors, RawResponse, WireRequest};
|
||||
use litellm_llms::base_llm::auth::{Authenticated, resolve_auth};
|
||||
|
||||
use super::{
|
||||
Error,
|
||||
types::{ProviderResponsesRequest, ResponsesOutput, ResponsesStreamHead},
|
||||
};
|
||||
|
||||
pub(super) async fn execute(
|
||||
http: &litellm_http::Client,
|
||||
auth: &litellm_auth::AuthServices,
|
||||
request: ProviderResponsesRequest,
|
||||
interceptors: &impl Interceptors<Error>,
|
||||
observers: Option<&ObservationSender>,
|
||||
) -> Result<ResponsesOutput, Error> {
|
||||
let authenticated = resolve_auth(auth, request.environment, &|_| None).await?;
|
||||
let wire = interceptors
|
||||
.before_provider_request(
|
||||
WireRequest {
|
||||
url: request.url,
|
||||
headers: authenticated.headers,
|
||||
body: request.body,
|
||||
},
|
||||
request.context,
|
||||
)
|
||||
.await?;
|
||||
let stream = match wire.body.get("stream") {
|
||||
None => false,
|
||||
Some(serde_json::Value::Bool(value)) => *value,
|
||||
Some(_) => return Err(Error::InvalidRequest("stream must be a boolean".into())),
|
||||
};
|
||||
let outbound = crate::outbound::outbound_request(
|
||||
Authenticated {
|
||||
headers: wire.headers,
|
||||
signer: authenticated.signer,
|
||||
},
|
||||
wire.url,
|
||||
&wire.body,
|
||||
Some(request.timeout.unwrap_or(Duration::from_secs(600))),
|
||||
)?;
|
||||
let response = crate::outbound::send(outbound, http)
|
||||
.await
|
||||
.map_err(network)?;
|
||||
let status = response.status().as_u16();
|
||||
if !response.status().is_success() {
|
||||
let body = response.text().await.map_err(network)?;
|
||||
return Err(litellm_http::transport::Error::Http {
|
||||
status,
|
||||
body: litellm_http::request::truncate_error_body(&body),
|
||||
}
|
||||
.into());
|
||||
}
|
||||
if stream {
|
||||
let headers = response
|
||||
.headers()
|
||||
.iter()
|
||||
.filter_map(|(name, value)| Some((name.to_string(), value.to_str().ok()?.to_owned())))
|
||||
.collect();
|
||||
let chunks = response
|
||||
.bytes_stream()
|
||||
.map(|chunk| chunk.map_err(network))
|
||||
.boxed();
|
||||
return Ok(ResponsesOutput::Stream {
|
||||
head: ResponsesStreamHead { headers },
|
||||
chunks,
|
||||
});
|
||||
}
|
||||
let body = response.text().await.map_err(network)?;
|
||||
let raw = RawResponse { body: body.clone() };
|
||||
if let Some(observers) = observers {
|
||||
observers.emit(litellm_host::lifecycle::CallEvent::Execution(
|
||||
ExecutionEvent::ProviderResponseReceived { raw: raw.clone() },
|
||||
));
|
||||
}
|
||||
interceptors
|
||||
.after_provider_response(raw)
|
||||
.await
|
||||
.map_err(Error::post_call)?;
|
||||
let value = serde_json::from_str(&body)
|
||||
.map_err(|error| Error::InvalidResponse(error.to_string().into()))?;
|
||||
request
|
||||
.config
|
||||
.transform_response_api_response(value)
|
||||
.map(ResponsesOutput::Complete)
|
||||
.map_err(Error::from)
|
||||
}
|
||||
|
||||
fn network(error: reqwest::Error) -> Error {
|
||||
litellm_http::transport::Error::Network(error.to_string()).into()
|
||||
}
|
||||
|
|
@ -1,2 +1,82 @@
|
|||
pub use crate::error::RouteError as Error;
|
||||
use litellm_host::observation::ObservationSender;
|
||||
pub mod websocket;
|
||||
|
||||
mod handler;
|
||||
mod prepare;
|
||||
pub mod route;
|
||||
pub mod types;
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use litellm_auth::AuthServices;
|
||||
use litellm_host::interceptors::Interceptors;
|
||||
use litellm_secrets::source::SecretSource;
|
||||
use types::{ResponsesCall, ResponsesOutput};
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct ResponsesRoute {
|
||||
http: litellm_http::Client,
|
||||
auth: Arc<AuthServices>,
|
||||
secrets: Arc<dyn SecretSource>,
|
||||
}
|
||||
|
||||
impl ResponsesRoute {
|
||||
pub fn new(
|
||||
http: litellm_http::Client,
|
||||
auth: Arc<AuthServices>,
|
||||
secrets: Arc<dyn SecretSource>,
|
||||
) -> Self {
|
||||
Self {
|
||||
http,
|
||||
auth,
|
||||
secrets,
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn execute(
|
||||
&self,
|
||||
call: ResponsesCall,
|
||||
interceptors: &impl Interceptors<Error>,
|
||||
observers: Option<ObservationSender>,
|
||||
) -> Result<ResponsesOutput, Error> {
|
||||
litellm_host::lifecycle::observe_call(
|
||||
observers.clone(),
|
||||
self.run(call, interceptors, observers.as_ref()),
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
#[tracing::instrument(name = "litellm.route", skip_all, fields(
|
||||
route = "responses",
|
||||
model = %call.model,
|
||||
provider,
|
||||
resolved_model,
|
||||
stream = call.optional_params.get("stream").and_then(serde_json::Value::as_bool).unwrap_or(false),
|
||||
outcome
|
||||
))]
|
||||
async fn run(
|
||||
&self,
|
||||
call: ResponsesCall,
|
||||
interceptors: &impl Interceptors<Error>,
|
||||
observers: Option<&ObservationSender>,
|
||||
) -> Result<ResponsesOutput, Error> {
|
||||
crate::diagnostic::call(async {
|
||||
let request = prepare::prepare(call, self.secrets.as_ref()).await?;
|
||||
crate::diagnostic::provider(
|
||||
&request.context.model,
|
||||
&request.context.custom_llm_provider,
|
||||
);
|
||||
let execute: futures_util::future::BoxFuture<'_, Result<ResponsesOutput, Error>> =
|
||||
Box::pin(handler::execute(
|
||||
&self.http,
|
||||
&self.auth,
|
||||
request,
|
||||
interceptors,
|
||||
observers,
|
||||
));
|
||||
execute.await
|
||||
})
|
||||
.await
|
||||
}
|
||||
}
|
||||
|
|
|
|||
58
litellm-rust/crates/core/src/responses/prepare.rs
Normal file
58
litellm-rust/crates/core/src/responses/prepare.rs
Normal file
|
|
@ -0,0 +1,58 @@
|
|||
use litellm_host::interceptors::RequestContext;
|
||||
use litellm_llms::{
|
||||
base_llm::responses::transformation::BaseResponsesApiConfig,
|
||||
openai::responses::transformation::OpenAiResponsesApiConfig,
|
||||
};
|
||||
use litellm_secrets::source::SecretSource;
|
||||
use serde_json::Value;
|
||||
|
||||
use super::{
|
||||
Error,
|
||||
types::{ProviderResponsesRequest, ResponsesCall},
|
||||
};
|
||||
|
||||
#[tracing::instrument(name = "litellm.prepare", level = "debug", skip_all)]
|
||||
pub(super) async fn prepare(
|
||||
call: ResponsesCall,
|
||||
secrets: &dyn SecretSource,
|
||||
) -> Result<ProviderResponsesRequest, Error> {
|
||||
let provider = call.custom_llm_provider.as_deref().unwrap_or("openai");
|
||||
if provider != "openai" {
|
||||
return Err(Error::Unsupported("native HTTP responses provider"));
|
||||
}
|
||||
let model = call.model.strip_prefix("openai/").unwrap_or(&call.model);
|
||||
if model.is_empty() || model.contains('/') {
|
||||
return Err(Error::InvalidProvider(call.model));
|
||||
}
|
||||
let config: &'static dyn BaseResponsesApiConfig = &OpenAiResponsesApiConfig;
|
||||
let snapshot = secrets
|
||||
.resolve(config.secret_names(call.api_key.as_deref(), call.api_base.as_deref()))
|
||||
.await?;
|
||||
let lookup = |name: &str| snapshot.get(name);
|
||||
let environment = config.validate_environment(
|
||||
litellm_http::request::string_headers("responses", call.extra_headers)?,
|
||||
call.api_key.as_deref(),
|
||||
&lookup,
|
||||
)?;
|
||||
let context = RequestContext {
|
||||
model: model.into(),
|
||||
custom_llm_provider: provider.into(),
|
||||
optional_params: Value::Object(call.optional_params.clone()),
|
||||
secret_fields: Vec::new(),
|
||||
api_key: match &environment.auth {
|
||||
litellm_llms::base_llm::auth::AuthScheme::Credential { secret, .. } => {
|
||||
Some(secret.clone())
|
||||
}
|
||||
_ => None,
|
||||
},
|
||||
};
|
||||
let body = config.transform_responses_api_request(model, call.input, call.optional_params)?;
|
||||
Ok(ProviderResponsesRequest {
|
||||
url: config.get_complete_url(call.api_base.as_deref(), &lookup),
|
||||
config,
|
||||
environment,
|
||||
body,
|
||||
context,
|
||||
timeout: call.timeout,
|
||||
})
|
||||
}
|
||||
41
litellm-rust/crates/core/src/responses/route.rs
Normal file
41
litellm-rust/crates/core/src/responses/route.rs
Normal file
|
|
@ -0,0 +1,41 @@
|
|||
use litellm_host::observation::ObservationSender;
|
||||
use std::convert::Infallible;
|
||||
|
||||
use bytes::Bytes;
|
||||
use litellm_host::{
|
||||
call::{HostedMachine, hosted_call},
|
||||
protocol::Protocol,
|
||||
};
|
||||
use litellm_types::responses::main::ResponsesApiResponse;
|
||||
|
||||
use super::{
|
||||
Error, ResponsesRoute,
|
||||
types::{ResponsesCall, ResponsesStreamHead},
|
||||
};
|
||||
|
||||
pub struct Responses;
|
||||
|
||||
impl Protocol for Responses {
|
||||
type Response = ResponsesApiResponse;
|
||||
type Error = Error;
|
||||
type Request = ResponsesCall;
|
||||
type HostCall = Infallible;
|
||||
type Chunk = Bytes;
|
||||
type StreamHead = ResponsesStreamHead;
|
||||
}
|
||||
|
||||
impl ResponsesRoute {
|
||||
pub fn machine(
|
||||
self,
|
||||
call: ResponsesCall,
|
||||
observers: Option<ObservationSender>,
|
||||
) -> HostedMachine<Responses> {
|
||||
hosted_call(
|
||||
call,
|
||||
observers,
|
||||
move |call, _, interceptors, observers| async move {
|
||||
self.run(call, &interceptors, observers.as_ref()).await
|
||||
},
|
||||
)
|
||||
}
|
||||
}
|
||||
37
litellm-rust/crates/core/src/responses/types.rs
Normal file
37
litellm-rust/crates/core/src/responses/types.rs
Normal file
|
|
@ -0,0 +1,37 @@
|
|||
use std::time::Duration;
|
||||
|
||||
use bytes::Bytes;
|
||||
use litellm_host::call::CallOutput;
|
||||
use litellm_llms::base_llm::{
|
||||
auth::ValidatedEnvironment, responses::transformation::BaseResponsesApiConfig,
|
||||
};
|
||||
use litellm_types::responses::main::ResponsesApiResponse;
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
use super::Error;
|
||||
|
||||
pub struct ResponsesCall {
|
||||
pub model: String,
|
||||
pub input: Value,
|
||||
pub optional_params: 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 timeout: Option<Duration>,
|
||||
}
|
||||
|
||||
pub struct ResponsesStreamHead {
|
||||
pub headers: Vec<(String, String)>,
|
||||
}
|
||||
|
||||
pub type ResponsesOutput = CallOutput<ResponsesApiResponse, ResponsesStreamHead, Bytes, Error>;
|
||||
|
||||
pub(super) struct ProviderResponsesRequest {
|
||||
pub config: &'static dyn BaseResponsesApiConfig,
|
||||
pub environment: ValidatedEnvironment,
|
||||
pub url: String,
|
||||
pub body: Value,
|
||||
pub context: litellm_host::interceptors::RequestContext,
|
||||
pub timeout: Option<Duration>,
|
||||
}
|
||||
|
|
@ -1,23 +1,13 @@
|
|||
use std::{
|
||||
collections::HashMap,
|
||||
io,
|
||||
sync::{Arc, OnceLock},
|
||||
time::Duration,
|
||||
};
|
||||
use std::{collections::HashMap, sync::Arc, time::Duration};
|
||||
|
||||
use futures_util::{SinkExt, StreamExt};
|
||||
use litellm_http::websocket::{UpstreamWebSocket, connect_upstream};
|
||||
use litellm_types::responses::streaming_websocket::ResponsesWsEventType;
|
||||
use rustls::{ClientConfig, RootCertStore};
|
||||
use tokio::{net::TcpStream, sync::Mutex};
|
||||
use tokio_tungstenite::{
|
||||
Connector, MaybeTlsStream, WebSocketStream, connect_async_tls_with_config,
|
||||
tungstenite::{
|
||||
Message,
|
||||
client::IntoClientRequest,
|
||||
error::TlsError,
|
||||
handshake::client::Response,
|
||||
http::{HeaderName, HeaderValue},
|
||||
},
|
||||
use tokio::sync::Mutex;
|
||||
use tokio_tungstenite::tungstenite::{
|
||||
Message,
|
||||
client::IntoClientRequest,
|
||||
http::{HeaderName, HeaderValue},
|
||||
};
|
||||
|
||||
use super::Error;
|
||||
|
|
@ -33,139 +23,127 @@ pub fn is_terminal_event(event_type: &ResponsesWsEventType) -> bool {
|
|||
)
|
||||
}
|
||||
|
||||
pub type ResponsesUpstreamWs = WebSocketStream<MaybeTlsStream<TcpStream>>;
|
||||
|
||||
static TLS_CONFIG: OnceLock<Arc<ClientConfig>> = OnceLock::new();
|
||||
|
||||
fn build_tls_config() -> Result<ClientConfig, Box<tokio_tungstenite::tungstenite::Error>> {
|
||||
let native = rustls_native_certs::load_native_certs();
|
||||
let mut store = RootCertStore::empty();
|
||||
let (added, _ignored) = store.add_parsable_certificates(native.certs);
|
||||
if added == 0 {
|
||||
return Err(Box::new(tokio_tungstenite::tungstenite::Error::Io(
|
||||
io::Error::other(format!(
|
||||
"no usable native root certificates: {:?}",
|
||||
native.errors
|
||||
)),
|
||||
)));
|
||||
}
|
||||
ClientConfig::builder_with_provider(Arc::new(rustls::crypto::ring::default_provider()))
|
||||
.with_safe_default_protocol_versions()
|
||||
.map(|builder| builder.with_root_certificates(store).with_no_client_auth())
|
||||
.map_err(|error| {
|
||||
Box::new(tokio_tungstenite::tungstenite::Error::Tls(
|
||||
TlsError::Rustls(error),
|
||||
))
|
||||
})
|
||||
}
|
||||
|
||||
fn tls_config() -> Result<Arc<ClientConfig>, Box<tokio_tungstenite::tungstenite::Error>> {
|
||||
if let Some(config) = TLS_CONFIG.get() {
|
||||
return Ok(Arc::clone(config));
|
||||
}
|
||||
let built = Arc::new(build_tls_config()?);
|
||||
Ok(Arc::clone(TLS_CONFIG.get_or_init(|| built)))
|
||||
}
|
||||
|
||||
pub async fn connect_upstream<R>(
|
||||
request: R,
|
||||
) -> Result<(ResponsesUpstreamWs, Response), Box<tokio_tungstenite::tungstenite::Error>>
|
||||
where
|
||||
R: IntoClientRequest + Unpin,
|
||||
{
|
||||
let request = request.into_client_request().map_err(Box::new)?;
|
||||
let connector = match request.uri().scheme_str() {
|
||||
Some("wss") => Some(Connector::Rustls(tls_config()?)),
|
||||
_ => None,
|
||||
};
|
||||
connect_async_tls_with_config(request, None, false, connector)
|
||||
.await
|
||||
.map_err(Box::new)
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct ResponsesWebSocketConnection {
|
||||
socket: Arc<Mutex<Option<ResponsesUpstreamWs>>>,
|
||||
socket: Arc<Mutex<Option<UpstreamWebSocket>>>,
|
||||
}
|
||||
|
||||
impl ResponsesWebSocketConnection {
|
||||
#[tracing::instrument(
|
||||
name = "litellm.websocket.connect_url",
|
||||
level = "debug",
|
||||
skip_all,
|
||||
fields(outcome)
|
||||
)]
|
||||
pub async fn connect_url(
|
||||
url: &str,
|
||||
headers: &HashMap<String, String>,
|
||||
timeout: Option<Duration>,
|
||||
) -> Result<Self, Error> {
|
||||
let mut request = url.into_client_request().map_err(|error| {
|
||||
Error::Transport(litellm_http::transport::Error::Network(error.to_string()))
|
||||
})?;
|
||||
for (name, value) in headers {
|
||||
let header_name = name
|
||||
.parse::<HeaderName>()
|
||||
.map_err(|error| Error::InvalidRequest(error.to_string()))?;
|
||||
let header_value = HeaderValue::from_str(value)
|
||||
.map_err(|error| Error::InvalidRequest(error.to_string()))?;
|
||||
request.headers_mut().insert(header_name, header_value);
|
||||
}
|
||||
let connect = connect_upstream(request);
|
||||
let result = match timeout {
|
||||
Some(timeout) => tokio::time::timeout(timeout, connect).await.map_err(|_| {
|
||||
Error::Transport(litellm_http::transport::Error::Network(
|
||||
"Responses WebSocket connection timed out".into(),
|
||||
))
|
||||
})?,
|
||||
None => connect.await,
|
||||
};
|
||||
let (socket, _) = result.map_err(|error| match *error {
|
||||
tokio_tungstenite::tungstenite::Error::Http(response) => {
|
||||
Error::Transport(litellm_http::transport::Error::Http {
|
||||
status: response.status().as_u16(),
|
||||
body: String::new(),
|
||||
})
|
||||
}
|
||||
other => Error::Transport(litellm_http::transport::Error::Network(other.to_string())),
|
||||
})?;
|
||||
Ok(Self {
|
||||
socket: Arc::new(Mutex::new(Some(socket))),
|
||||
})
|
||||
}
|
||||
|
||||
pub async fn send_text(&self, text: String) -> Result<(), Error> {
|
||||
let mut socket = self.socket.lock().await;
|
||||
let Some(socket) = socket.as_mut() else {
|
||||
return Err(Error::Transport(litellm_http::transport::Error::Network(
|
||||
"Responses WebSocket is closed".into(),
|
||||
)));
|
||||
};
|
||||
socket.send(Message::Text(text)).await.map_err(|error| {
|
||||
Error::Transport(litellm_http::transport::Error::Network(error.to_string()))
|
||||
})
|
||||
}
|
||||
|
||||
pub async fn recv_text(&self) -> Result<Option<String>, Error> {
|
||||
let mut socket = self.socket.lock().await;
|
||||
let Some(socket) = socket.as_mut() else {
|
||||
return Ok(None);
|
||||
};
|
||||
match socket.next().await {
|
||||
Some(Ok(Message::Text(text))) => Ok(Some(text)),
|
||||
Some(Ok(Message::Binary(bytes))) => String::from_utf8(bytes.to_vec())
|
||||
.map(Some)
|
||||
.map_err(|error| Error::InvalidResponse(error.to_string())),
|
||||
Some(Ok(Message::Close(_))) | None => Ok(None),
|
||||
Some(Ok(_)) => Ok(None),
|
||||
Some(Err(error)) => Err(Error::Transport(litellm_http::transport::Error::Network(
|
||||
error.to_string(),
|
||||
))),
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn close(&self) -> Result<(), Error> {
|
||||
let mut socket = self.socket.lock().await;
|
||||
if let Some(socket) = socket.as_mut() {
|
||||
socket.close(None).await.map_err(|error| {
|
||||
crate::diagnostic::operation("litellm.websocket.connect_url", async {
|
||||
let mut request = url.into_client_request().map_err(|error| {
|
||||
Error::Transport(litellm_http::transport::Error::Network(error.to_string()))
|
||||
})?;
|
||||
}
|
||||
*socket = None;
|
||||
Ok(())
|
||||
for (name, value) in headers {
|
||||
let header_name = name
|
||||
.parse::<HeaderName>()
|
||||
.map_err(|error| Error::InvalidRequest(error.to_string().into()))?;
|
||||
let header_value = HeaderValue::from_str(value)
|
||||
.map_err(|error| Error::InvalidRequest(error.to_string().into()))?;
|
||||
request.headers_mut().insert(header_name, header_value);
|
||||
}
|
||||
let connect = connect_upstream(request);
|
||||
let result = match timeout {
|
||||
Some(timeout) => tokio::time::timeout(timeout, connect).await.map_err(|_| {
|
||||
Error::Transport(litellm_http::transport::Error::Network(
|
||||
"Responses WebSocket connection timed out".into(),
|
||||
))
|
||||
})?,
|
||||
None => connect.await,
|
||||
};
|
||||
let (socket, _) = result.map_err(|error| match *error {
|
||||
tokio_tungstenite::tungstenite::Error::Http(response) => {
|
||||
Error::Transport(litellm_http::transport::Error::Http {
|
||||
status: response.status().as_u16(),
|
||||
body: String::new(),
|
||||
})
|
||||
}
|
||||
other => {
|
||||
Error::Transport(litellm_http::transport::Error::Network(other.to_string()))
|
||||
}
|
||||
})?;
|
||||
Ok(Self {
|
||||
socket: Arc::new(Mutex::new(Some(socket))),
|
||||
})
|
||||
})
|
||||
.await
|
||||
}
|
||||
|
||||
#[tracing::instrument(
|
||||
name = "litellm.websocket.send_text",
|
||||
level = "debug",
|
||||
skip_all,
|
||||
fields(outcome)
|
||||
)]
|
||||
pub async fn send_text(&self, text: String) -> Result<(), Error> {
|
||||
crate::diagnostic::operation("litellm.websocket.send_text", async {
|
||||
let mut socket = self.socket.lock().await;
|
||||
let Some(socket) = socket.as_mut() else {
|
||||
return Err(Error::Transport(litellm_http::transport::Error::Network(
|
||||
"Responses WebSocket is closed".into(),
|
||||
)));
|
||||
};
|
||||
socket.send(Message::Text(text)).await.map_err(|error| {
|
||||
Error::Transport(litellm_http::transport::Error::Network(error.to_string()))
|
||||
})
|
||||
})
|
||||
.await
|
||||
}
|
||||
|
||||
#[tracing::instrument(
|
||||
name = "litellm.websocket.recv_text",
|
||||
level = "debug",
|
||||
skip_all,
|
||||
fields(outcome)
|
||||
)]
|
||||
pub async fn recv_text(&self) -> Result<Option<String>, Error> {
|
||||
crate::diagnostic::operation("litellm.websocket.recv_text", async {
|
||||
let mut socket = self.socket.lock().await;
|
||||
let Some(socket) = socket.as_mut() else {
|
||||
return Ok(None);
|
||||
};
|
||||
match socket.next().await {
|
||||
Some(Ok(Message::Text(text))) => Ok(Some(text)),
|
||||
Some(Ok(Message::Binary(bytes))) => String::from_utf8(bytes.to_vec())
|
||||
.map(Some)
|
||||
.map_err(|error| Error::InvalidResponse(error.to_string().into())),
|
||||
Some(Ok(Message::Close(_))) | None => Ok(None),
|
||||
Some(Ok(_)) => Ok(None),
|
||||
Some(Err(error)) => Err(Error::Transport(litellm_http::transport::Error::Network(
|
||||
error.to_string(),
|
||||
))),
|
||||
}
|
||||
})
|
||||
.await
|
||||
}
|
||||
|
||||
#[tracing::instrument(
|
||||
name = "litellm.websocket.close",
|
||||
level = "debug",
|
||||
skip_all,
|
||||
fields(outcome)
|
||||
)]
|
||||
pub async fn close(&self) -> Result<(), Error> {
|
||||
crate::diagnostic::operation("litellm.websocket.close", async {
|
||||
let mut socket = self.socket.lock().await;
|
||||
if let Some(socket) = socket.as_mut() {
|
||||
socket.close(None).await.map_err(|error| {
|
||||
Error::Transport(litellm_http::transport::Error::Network(error.to_string()))
|
||||
})?;
|
||||
}
|
||||
*socket = None;
|
||||
Ok(())
|
||||
})
|
||||
.await
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,6 +1,4 @@
|
|||
use litellm_core::audio_transcription::{
|
||||
Error, audio_transcription, types::AudioTranscriptionRequest,
|
||||
};
|
||||
use litellm_core::audio_transcription::{Error, types::AudioTranscriptionRequest};
|
||||
use rstest::{fixture, rstest};
|
||||
use serde_json::{Map, Value, json};
|
||||
use wiremock::ResponseTemplate;
|
||||
|
|
@ -11,7 +9,7 @@ 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
|
||||
audio_transcription_route().execute(request).await
|
||||
}
|
||||
|
||||
fn transcript_response(text: &str) -> ResponseTemplate {
|
||||
|
|
@ -252,3 +250,31 @@ async fn an_unreadable_success_body_is_an_invalid_response(
|
|||
|
||||
assert!(matches!(error, Error::InvalidResponse(_)), "{error:?}");
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn transcription_records_route_and_resolved_provider(
|
||||
request: AudioTranscriptionRequest<'static>,
|
||||
traces: TraceCapture,
|
||||
) {
|
||||
let upstream = upstream([transcript_response("hello")]).await;
|
||||
let base = upstream.uri();
|
||||
let model = request.model;
|
||||
traces
|
||||
.logger()
|
||||
.instrument(transcribe(AudioTranscriptionRequest {
|
||||
api_base: Some(&base),
|
||||
..request
|
||||
}))
|
||||
.await
|
||||
.unwrap();
|
||||
let summaries = traces.summaries("litellm.route");
|
||||
assert_eq!(summaries.len(), 1);
|
||||
assert_eq!(summaries[0]["route"], "audio_transcription");
|
||||
assert_eq!(summaries[0]["model"], model);
|
||||
assert_eq!(summaries[0]["resolved_model"], model);
|
||||
assert_eq!(summaries[0]["provider"], "bedrock");
|
||||
assert_eq!(summaries[0]["outcome"], "success");
|
||||
assert_eq!(summaries[0]["stream"], false);
|
||||
assert!(!format!("{:?}", traces.records()).contains("secret-key"));
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,8 +1,7 @@
|
|||
use litellm_host::interceptors::RawResponse;
|
||||
use std::time::Duration;
|
||||
|
||||
use litellm_core::chat_completions::{
|
||||
Error, chat_completions, chat_completions_decline_reason, types::ChatCompletionsRequest,
|
||||
};
|
||||
use litellm_core::chat_completions::{Error, types::ChatCompletionsRequest};
|
||||
use litellm_http::transport::Error as TransportError;
|
||||
use litellm_types::utils::ChatCompletionsResponse;
|
||||
use rstest::{fixture, rstest};
|
||||
|
|
@ -15,7 +14,7 @@ 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
|
||||
chat_completions_route().execute(request, &(), None).await
|
||||
}
|
||||
|
||||
fn object(value: Value) -> Map<String, Value> {
|
||||
|
|
@ -156,8 +155,6 @@ async fn bedrock_round_trip_is_signed_and_normalized(request: ChatCompletionsReq
|
|||
assert_eq!(response.usage.total_tokens, 15);
|
||||
}
|
||||
|
||||
/// The provider already answered and billed these, so the host must not retry them on
|
||||
/// its own path: they surface as `InvalidResponse`, never as a pre-send decline.
|
||||
#[rstest]
|
||||
#[case::missing_usage(
|
||||
r#"{"model":"m","content":[{"type":"text","text":"hi"}],"stop_reason":"end_turn"}"#
|
||||
|
|
@ -209,10 +206,9 @@ async fn an_upstream_error_status_keeps_its_code_and_body(
|
|||
);
|
||||
}
|
||||
|
||||
/// Nothing was sent, so nothing was billed and the host can still serve the request.
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn a_connection_that_is_never_established_declines_instead_of_failing(
|
||||
async fn a_connection_that_is_never_established_returns_a_connect_error(
|
||||
request: ChatCompletionsRequest<'static>,
|
||||
) {
|
||||
let error = complete(ChatCompletionsRequest {
|
||||
|
|
@ -230,9 +226,7 @@ async fn a_connection_that_is_never_established_declines_instead_of_failing(
|
|||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn a_timeout_after_sending_is_not_a_pre_send_decline(
|
||||
request: ChatCompletionsRequest<'static>,
|
||||
) {
|
||||
async fn a_timeout_after_sending_returns_a_network_error(request: ChatCompletionsRequest<'static>) {
|
||||
let upstream =
|
||||
upstream([anthropic_response(ANTHROPIC_MESSAGE).set_delay(Duration::from_secs(5))]).await;
|
||||
let base = upstream.uri();
|
||||
|
|
@ -252,74 +246,141 @@ async fn a_timeout_after_sending_is_not_a_pre_send_decline(
|
|||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::accepted("anthropic/claude-sonnet-4-5", None, hi(), json!({"max_tokens": 16}), None)]
|
||||
#[case::accepted_bedrock("bedrock/anthropic.claude-sonnet-4-5", None, hi(), json!({}), None)]
|
||||
#[case::unknown_provider(
|
||||
"gpt-4o",
|
||||
Some("openai"),
|
||||
hi(),
|
||||
json!({}),
|
||||
Some("provider is not on the rust chat completions path")
|
||||
)]
|
||||
#[case::unreadable_messages(
|
||||
"anthropic/claude-sonnet-4-5",
|
||||
None,
|
||||
json!("hi"),
|
||||
json!({}),
|
||||
Some("unreadable message list")
|
||||
)]
|
||||
#[case::empty_messages("anthropic/claude-sonnet-4-5", None, json!([]), json!({}), Some("empty message list"))]
|
||||
#[case::streaming(
|
||||
"anthropic/claude-sonnet-4-5",
|
||||
None,
|
||||
hi(),
|
||||
json!({"stream": true}),
|
||||
Some("streaming")
|
||||
)]
|
||||
#[case::unrecognized_param(
|
||||
"anthropic/claude-sonnet-4-5",
|
||||
None,
|
||||
hi(),
|
||||
json!({"not_a_param": 1}),
|
||||
Some("unrecognized request parameter")
|
||||
)]
|
||||
#[case::opens_on_assistant_turn(
|
||||
"anthropic/claude-sonnet-4-5",
|
||||
None,
|
||||
json!([{"role": "assistant", "content": "hi"}]),
|
||||
json!({}),
|
||||
Some("conversation does not open on a user turn")
|
||||
)]
|
||||
fn decline_reason_names_why_the_core_would_not_serve_the_request(
|
||||
#[case] model: &str,
|
||||
#[case] provider: Option<&str>,
|
||||
#[case] messages: Value,
|
||||
#[case] params: Value,
|
||||
#[case] reason: Option<&str>,
|
||||
#[case::direct(false)]
|
||||
#[case::hosted(true)]
|
||||
#[tokio::test]
|
||||
async fn direct_and_hosted_calls_share_hooks_and_lifecycle(
|
||||
request: ChatCompletionsRequest<'static>,
|
||||
#[case] hosted: bool,
|
||||
) {
|
||||
assert_eq!(
|
||||
chat_completions_decline_reason(model, provider, messages, &object(params)),
|
||||
reason
|
||||
use litellm_core::chat_completions::route::ChatCompletions;
|
||||
use litellm_host::{call::HostedCompletion, lifecycle::CallEvent};
|
||||
|
||||
let upstream = upstream([anthropic_response(ANTHROPIC_MESSAGE)]).await;
|
||||
let base = upstream.uri();
|
||||
let host = RecordingCall::<ChatCompletions>::new(
|
||||
ChatCompletionsRequest {
|
||||
api_base: Some(&base),
|
||||
..request
|
||||
}
|
||||
.into(),
|
||||
);
|
||||
let response = if hosted {
|
||||
let result = litellm_host_native::in_process::run_hosted(
|
||||
chat_completions_route()
|
||||
.machine(host.request().unwrap(), Some(host.events.0.sender.clone())),
|
||||
host.runtime(),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
let HostedCompletion::Complete(response) = result else {
|
||||
panic!("expected a complete response")
|
||||
};
|
||||
response
|
||||
} else {
|
||||
let call = host.request.lock().unwrap().take().unwrap();
|
||||
chat_completions_route()
|
||||
.execute(
|
||||
ChatCompletionsRequest {
|
||||
model: &call.model,
|
||||
messages: call.messages,
|
||||
optional_params: call.optional_params,
|
||||
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,
|
||||
timeout: call.timeout,
|
||||
},
|
||||
&host,
|
||||
Some(host.events.0.sender.clone()),
|
||||
)
|
||||
.await
|
||||
.unwrap()
|
||||
};
|
||||
assert_eq!(
|
||||
response.choices[0].message.content.as_deref(),
|
||||
Some("hello")
|
||||
);
|
||||
assert_eq!(
|
||||
only_request(&upstream).await.header("x-hook"),
|
||||
Some("called")
|
||||
);
|
||||
let events = host.events.0.lock().unwrap();
|
||||
assert!(matches!(
|
||||
&events[..],
|
||||
[
|
||||
CallEvent::Started { .. },
|
||||
CallEvent::Execution(_),
|
||||
CallEvent::Succeeded { .. }
|
||||
]
|
||||
));
|
||||
}
|
||||
|
||||
/// A request the decline check accepts must not be declined by the call itself.
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn a_declined_request_fails_the_call_before_sending(
|
||||
async fn a_post_call_hook_failure_never_looks_safe_to_retry(
|
||||
request: ChatCompletionsRequest<'static>,
|
||||
) {
|
||||
use litellm_host::interceptors::{Interceptors, RequestContext, WireRequest};
|
||||
struct FailingHook;
|
||||
impl Interceptors<Error> for FailingHook {
|
||||
async fn before_provider_request(
|
||||
&self,
|
||||
wire: WireRequest,
|
||||
_: RequestContext,
|
||||
) -> Result<WireRequest, Error> {
|
||||
Ok(wire)
|
||||
}
|
||||
async fn after_provider_response(&self, _: RawResponse) -> Result<(), Error> {
|
||||
Err(Error::InvalidRequest("callback rejected".into()))
|
||||
}
|
||||
}
|
||||
let upstream = upstream([anthropic_response(ANTHROPIC_MESSAGE)]).await;
|
||||
let base = upstream.uri();
|
||||
let error = chat_completions_route()
|
||||
.execute(
|
||||
ChatCompletionsRequest {
|
||||
api_base: Some(&base),
|
||||
..request
|
||||
},
|
||||
&FailingHook,
|
||||
None,
|
||||
)
|
||||
.await
|
||||
.unwrap_err();
|
||||
let Error::PostCallHook(source) = error else {
|
||||
panic!("expected retained callback error")
|
||||
};
|
||||
assert_eq!(*source, Error::InvalidRequest("callback rejected".into()));
|
||||
assert_eq!(received(&upstream).await.len(), 1);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn completed_chat_records_route_and_resolved_provider(
|
||||
request: ChatCompletionsRequest<'static>,
|
||||
traces: TraceCapture,
|
||||
) {
|
||||
let upstream = upstream([anthropic_response(ANTHROPIC_MESSAGE)]).await;
|
||||
let base = upstream.uri();
|
||||
|
||||
let error = complete(ChatCompletionsRequest {
|
||||
optional_params: object(json!({"stream": true})),
|
||||
api_base: Some(&base),
|
||||
..request
|
||||
})
|
||||
.await
|
||||
.expect_err("streaming is declined");
|
||||
|
||||
assert_eq!(error, Error::Unsupported("streaming"));
|
||||
assert!(received(&upstream).await.is_empty());
|
||||
let model = request.model;
|
||||
traces
|
||||
.logger()
|
||||
.instrument(complete(ChatCompletionsRequest {
|
||||
api_base: Some(&base),
|
||||
..request
|
||||
}))
|
||||
.await
|
||||
.unwrap();
|
||||
let summaries = traces.summaries("litellm.route");
|
||||
assert_eq!(summaries.len(), 1);
|
||||
assert_eq!(summaries[0]["route"], "chat_completions");
|
||||
assert_eq!(summaries[0]["model"], model);
|
||||
assert_eq!(summaries[0]["provider"], "anthropic");
|
||||
assert_eq!(
|
||||
summaries[0]["resolved_model"],
|
||||
only_request(&upstream).await.json()["model"]
|
||||
);
|
||||
assert_eq!(summaries[0]["outcome"], "success");
|
||||
assert_eq!(summaries[0]["stream"], false);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,23 +1,24 @@
|
|||
use std::{convert::Infallible, sync::Mutex};
|
||||
use litellm_host::lifecycle::ExecutionEvent;
|
||||
use std::sync::Mutex;
|
||||
|
||||
use litellm_core::messages::route::Messages;
|
||||
use litellm_host::{
|
||||
event::{CallEvent, MachineEvent, RequestContext, WireRequest},
|
||||
host::Host,
|
||||
interceptors::{RequestContext, WireRequest},
|
||||
lifecycle::CallEvent,
|
||||
};
|
||||
use litellm_llms::anthropic::common_utils::AnthropicModelCapabilities;
|
||||
use litellm_llms::base_llm::messages::context::MessagesModelCapabilities as AnthropicModelCapabilities;
|
||||
use rstest::rstest;
|
||||
|
||||
use super::*;
|
||||
|
||||
type Rewrite = Box<dyn Fn(WireRequest) -> Result<WireRequest, Error> + Send + Sync>;
|
||||
|
||||
/// Projects like `LocalMessagesHost`, answers `before_send` through `rewrite`, and keeps
|
||||
/// Projects like `LocalMessagesHost`, answers `before_provider_request` through `rewrite`, and keeps
|
||||
/// every event the driver emits.
|
||||
struct RecordingHost {
|
||||
call: LocalMessagesHost,
|
||||
rewrite: Rewrite,
|
||||
events: Mutex<Vec<CallEvent>>,
|
||||
events: super::support::Observations,
|
||||
optional_params: Mutex<Vec<Value>>,
|
||||
}
|
||||
|
||||
|
|
@ -26,7 +27,7 @@ impl RecordingHost {
|
|||
Self {
|
||||
call: LocalMessagesHost::new(call),
|
||||
rewrite,
|
||||
events: Mutex::new(Vec::new()),
|
||||
events: super::support::Observations::default(),
|
||||
optional_params: Mutex::new(Vec::new()),
|
||||
}
|
||||
}
|
||||
|
|
@ -41,7 +42,7 @@ impl RecordingHost {
|
|||
.unwrap()
|
||||
.iter()
|
||||
.filter_map(|event| match event {
|
||||
CallEvent::Machine(MachineEvent::ResponseReceived { raw }) => {
|
||||
CallEvent::Execution(ExecutionEvent::ProviderResponseReceived { raw }) => {
|
||||
Some(raw.body.clone())
|
||||
}
|
||||
_ => None,
|
||||
|
|
@ -50,19 +51,32 @@ impl RecordingHost {
|
|||
}
|
||||
}
|
||||
|
||||
impl Host<Messages> for RecordingHost {
|
||||
async fn project(&self) -> Result<MessagesCall, Error> {
|
||||
self.call.project().await
|
||||
impl RecordingHost {
|
||||
pub fn request(&self) -> Result<MessagesCall, Error> {
|
||||
self.call.request()
|
||||
}
|
||||
|
||||
async fn custom_op(&self, op: Infallible) -> Result<(), Error> {
|
||||
match op {}
|
||||
pub fn runtime(&self) -> litellm_host_native::in_process::Host<'_, (), Self, ()> {
|
||||
litellm_host_native::in_process::Host {
|
||||
services: &(),
|
||||
interceptors: self,
|
||||
stream: &(),
|
||||
observers: Some(&self.events.sender),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
async fn before_send(
|
||||
impl litellm_host::lifecycle::CallObserver for RecordingHost {
|
||||
fn observe(&self, event: litellm_host::lifecycle::CallEvent) {
|
||||
self.events.sender.emit(event);
|
||||
}
|
||||
}
|
||||
impl litellm_host::interceptors::Interceptors<<Messages as litellm_host::protocol::Protocol>::Error>
|
||||
for RecordingHost
|
||||
{
|
||||
async fn before_provider_request(
|
||||
&self,
|
||||
wire: WireRequest,
|
||||
context: &RequestContext,
|
||||
context: RequestContext,
|
||||
) -> Result<WireRequest, Error> {
|
||||
self.optional_params
|
||||
.lock()
|
||||
|
|
@ -70,15 +84,26 @@ impl Host<Messages> for RecordingHost {
|
|||
.push(context.optional_params.clone());
|
||||
(self.rewrite)(wire)
|
||||
}
|
||||
|
||||
async fn emit(&self, event: &CallEvent) -> Result<(), Error> {
|
||||
self.events.lock().unwrap().push(event.clone());
|
||||
async fn after_provider_response(
|
||||
&self,
|
||||
raw: litellm_host::interceptors::RawResponse,
|
||||
) -> Result<(), <Messages as litellm_host::protocol::Protocol>::Error> {
|
||||
litellm_host::lifecycle::CallObserver::observe(
|
||||
self,
|
||||
litellm_host::lifecycle::CallEvent::Execution(
|
||||
litellm_host::lifecycle::ExecutionEvent::ProviderResponseReceived { raw },
|
||||
),
|
||||
);
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
async fn run_through(host: &RecordingHost) -> Result<MessagesOutput, Error> {
|
||||
litellm_host::run::run(machine(Arc::new(RecordingSecrets::empty())), host).await
|
||||
litellm_host_native::in_process::run_hosted(
|
||||
machine(Arc::new(RecordingSecrets::empty()))(host.request()?),
|
||||
host.runtime(),
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
fn authenticated(call: MessagesCall, api_base: String) -> MessagesCall {
|
||||
|
|
@ -129,8 +154,7 @@ async fn a_before_send_failure_never_sends(call: MessagesCall) {
|
|||
|
||||
let error = run_through(&host)
|
||||
.await
|
||||
.err()
|
||||
.expect("the host failure fails the call");
|
||||
.expect_err("the host failure fails the call");
|
||||
|
||||
assert_eq!(error, Error::InvalidRequest("vetoed by the host".into()));
|
||||
assert!(received(&upstream).await.is_empty());
|
||||
|
|
@ -146,7 +170,7 @@ async fn the_raw_upstream_text_is_emitted_once_for_a_message(call: MessagesCall)
|
|||
|
||||
let output = run_through(&host).await.expect("messages call succeeds");
|
||||
|
||||
assert!(matches!(output, MessagesOutput::Message(_)));
|
||||
assert!(matches!(output, MessagesOutput::Complete(_)));
|
||||
let [emitted] = <[String; 1]>::try_from(host.raw_responses())
|
||||
.unwrap_or_else(|raws| panic!("expected one raw response, got {}", raws.len()));
|
||||
assert_eq!(serde_json::from_str::<Value>(&emitted).unwrap(), raw);
|
||||
|
|
@ -198,6 +222,6 @@ async fn the_request_context_carries_the_shaped_params_without_model_or_messages
|
|||
run_through(&host).await.expect("messages call succeeds");
|
||||
|
||||
let [optional_params] = <[Value; 1]>::try_from(host.optional_params.into_inner().unwrap())
|
||||
.unwrap_or_else(|seen| panic!("before_send runs once, saw {}", seen.len()));
|
||||
.unwrap_or_else(|seen| panic!("before_provider_request runs once, saw {}", seen.len()));
|
||||
assert_eq!(optional_params, json!({"max_tokens": 16}));
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,8 +1,11 @@
|
|||
use std::{sync::Arc, time::Duration};
|
||||
use std::{
|
||||
sync::{Arc, Mutex},
|
||||
time::Duration,
|
||||
};
|
||||
|
||||
use litellm_core::messages::{
|
||||
Error, MessagesCall, MessagesShaping,
|
||||
route::{LocalMessagesHost, MessagesMachine, MessagesOutput, messages_machine},
|
||||
route::{Messages, MessagesMachine, MessagesOutput},
|
||||
};
|
||||
use litellm_http::{HttpSettings, Resolution};
|
||||
use litellm_secrets::source::SecretSource;
|
||||
|
|
@ -95,16 +98,17 @@ 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")
|
||||
fn machine(secrets: Arc<dyn SecretSource>) -> impl FnOnce(MessagesCall) -> MessagesMachine {
|
||||
move |request| messages_route(secrets).machine(request, None)
|
||||
}
|
||||
|
||||
async fn run_with(
|
||||
secrets: Arc<RecordingSecrets>,
|
||||
call: MessagesCall,
|
||||
) -> Result<MessagesOutput, Error> {
|
||||
litellm_host::run::run(machine(secrets), &LocalMessagesHost::new(call)).await
|
||||
let host = LocalMessagesHost::new(call);
|
||||
litellm_host_native::in_process::run_hosted(machine(secrets)(host.request()?), host.runtime())
|
||||
.await
|
||||
}
|
||||
|
||||
/// Runs the route with a secret source that knows nothing, so no environment leaks in.
|
||||
|
|
@ -114,7 +118,69 @@ async fn run(call: MessagesCall) -> Result<MessagesOutput, Error> {
|
|||
|
||||
async fn run_message(call: MessagesCall) -> AnthropicMessagesResponse {
|
||||
match run(call).await.expect("messages call succeeds") {
|
||||
MessagesOutput::Message(message) => *message,
|
||||
MessagesOutput::Streamed => panic!("a non-streaming call returned a stream"),
|
||||
MessagesOutput::Complete(message) => *message,
|
||||
MessagesOutput::StreamEnded | MessagesOutput::Detached => {
|
||||
panic!("a non-streaming call returned a stream")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
struct LocalMessagesHost {
|
||||
call: Mutex<Option<MessagesCall>>,
|
||||
}
|
||||
|
||||
impl LocalMessagesHost {
|
||||
fn new(call: MessagesCall) -> Self {
|
||||
Self {
|
||||
call: Mutex::new(Some(call)),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl LocalMessagesHost {
|
||||
pub fn request(&self) -> Result<MessagesCall, Error> {
|
||||
self.call
|
||||
.lock()
|
||||
.unwrap_or_else(|error| error.into_inner())
|
||||
.take()
|
||||
.ok_or_else(|| Error::InvalidRequest("messages request was already projected".into()))
|
||||
}
|
||||
pub fn runtime(&self) -> litellm_host_native::in_process::Host<'_, (), Self, ()> {
|
||||
litellm_host_native::in_process::Host {
|
||||
services: &(),
|
||||
interceptors: self,
|
||||
stream: &(),
|
||||
observers: None,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl litellm_host::lifecycle::CallObserver for LocalMessagesHost {
|
||||
fn observe(&self, _: litellm_host::lifecycle::CallEvent) {}
|
||||
}
|
||||
impl litellm_host::interceptors::Interceptors<<Messages as litellm_host::protocol::Protocol>::Error>
|
||||
for LocalMessagesHost
|
||||
{
|
||||
async fn before_provider_request(
|
||||
&self,
|
||||
wire: litellm_host::interceptors::WireRequest,
|
||||
_: litellm_host::interceptors::RequestContext,
|
||||
) -> Result<
|
||||
litellm_host::interceptors::WireRequest,
|
||||
<Messages as litellm_host::protocol::Protocol>::Error,
|
||||
> {
|
||||
Ok(wire)
|
||||
}
|
||||
async fn after_provider_response(
|
||||
&self,
|
||||
raw: litellm_host::interceptors::RawResponse,
|
||||
) -> Result<(), <Messages as litellm_host::protocol::Protocol>::Error> {
|
||||
litellm_host::lifecycle::CallObserver::observe(
|
||||
self,
|
||||
litellm_host::lifecycle::CallEvent::Execution(
|
||||
litellm_host::lifecycle::ExecutionEvent::ProviderResponseReceived { raw },
|
||||
),
|
||||
);
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue