mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
Merge branch 'main' into fix/bedrock-converse-redacted-thinking-replay-43009
This commit is contained in:
commit
64bc20a5d0
512 changed files with 30155 additions and 5291 deletions
|
|
@ -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
|
||||
|
|
|
|||
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 ()
|
||||
|
|
|
|||
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/` |
|
||||
|
|
|
|||
|
|
@ -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
|
||||
30
litellm-rust/Cargo.lock
generated
30
litellm-rust/Cargo.lock
generated
|
|
@ -3138,6 +3138,7 @@ dependencies = [
|
|||
"litellm-http",
|
||||
"litellm-llms",
|
||||
"litellm-secrets",
|
||||
"litellm-tracing",
|
||||
"litellm-types",
|
||||
"mime_guess",
|
||||
"moka",
|
||||
|
|
@ -3215,6 +3216,8 @@ name = "litellm-gateway"
|
|||
version = "0.1.0"
|
||||
dependencies = [
|
||||
"axum",
|
||||
"futures-util",
|
||||
"http-body-util",
|
||||
"litellm-config",
|
||||
"litellm-core",
|
||||
"litellm-gateway-auth",
|
||||
|
|
@ -3222,11 +3225,13 @@ dependencies = [
|
|||
"litellm-http",
|
||||
"litellm-llms",
|
||||
"litellm-secrets",
|
||||
"litellm-tracing",
|
||||
"rstest",
|
||||
"serde_json",
|
||||
"tokio",
|
||||
"tower-http 0.7.1",
|
||||
"tower",
|
||||
"tracing",
|
||||
"uuid",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
|
@ -3256,7 +3261,6 @@ dependencies = [
|
|||
"futures-util",
|
||||
"litellm-auth",
|
||||
"litellm-core",
|
||||
"litellm-host",
|
||||
"litellm-http",
|
||||
"litellm-llms",
|
||||
"litellm-router",
|
||||
|
|
@ -3677,6 +3681,7 @@ dependencies = [
|
|||
name = "litellm-tracing"
|
||||
version = "0.1.0"
|
||||
dependencies = [
|
||||
"base64 0.22.1",
|
||||
"fancy-regex 0.19.2",
|
||||
"percent-encoding",
|
||||
"rstest",
|
||||
|
|
@ -4880,7 +4885,7 @@ dependencies = [
|
|||
"tokio-rustls 0.26.4",
|
||||
"tokio-util",
|
||||
"tower",
|
||||
"tower-http 0.6.11",
|
||||
"tower-http",
|
||||
"tower-service",
|
||||
"url",
|
||||
"wasm-bindgen",
|
||||
|
|
@ -4922,7 +4927,7 @@ dependencies = [
|
|||
"tokio-rustls 0.26.4",
|
||||
"tokio-util",
|
||||
"tower",
|
||||
"tower-http 0.6.11",
|
||||
"tower-http",
|
||||
"tower-service",
|
||||
"url",
|
||||
"wasm-bindgen",
|
||||
|
|
@ -6163,23 +6168,6 @@ dependencies = [
|
|||
"url",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "tower-http"
|
||||
version = "0.7.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "08a05a66a4fdd61cbbe0a1d755ffe0ca6aba159dd4820936a0ff8a8278245b9c"
|
||||
dependencies = [
|
||||
"bitflags 2.13.1",
|
||||
"bytes",
|
||||
"http 1.4.2",
|
||||
"http-body 1.1.0",
|
||||
"percent-encoding",
|
||||
"pin-project-lite",
|
||||
"tower-layer",
|
||||
"tower-service",
|
||||
"tracing",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "tower-layer"
|
||||
version = "0.3.3"
|
||||
|
|
|
|||
|
|
@ -7,7 +7,7 @@ pub enum CredentialPlacement {
|
|||
}
|
||||
|
||||
impl CredentialPlacement {
|
||||
pub fn header_name(self) -> &'static str {
|
||||
pub const fn header_name(self) -> &'static str {
|
||||
match self {
|
||||
Self::Bearer => "Authorization",
|
||||
Self::Header(name) => name,
|
||||
|
|
|
|||
|
|
@ -1,4 +1,6 @@
|
|||
litellm-core is the LiteLLM SDK in Rust — it makes the LLM call. Each top-level call is a module under `src/<route>/` exposing a public entrypoint named after the route (`messages::messages()`, the Rust equivalent of `litellm.messages()`): you call it and get a typed non-streaming response back.
|
||||
litellm-core is the LiteLLM SDK in Rust. Each top-level call is a module under `src/<route>/` exposing a public entrypoint named after the route. `messages::messages()` returns `MessagesResponse::Message` for a completed response or `MessagesResponse::Stream { headers, chunks }` when the request sets `stream: true`. The chunks are Anthropic SSE bytes in a `Stream<Item = Result<Bytes, Error>>`. Dropping the stream cancels the call. The Python bridge drives `messages::route::messages_machine()` instead, because Python has to answer the call's operations on its own thread; the gateway and the Rust SDK call the plain entrypoint
|
||||
|
||||
A route module has the same five pieces, in the order Python runs them. `types.rs` holds the call, the provider request, and the response. `prepare.rs` resolves the provider and credentials and shapes the request (Python's `validate_environment`, `get_complete_url`, `transform_request`). `handler.rs` resolves auth, offers the wire request to `litellm_host::hooks::RouteHooks::before_send`, sends it, reports the raw response through `emit`, and normalizes the response or stream (`pre_call`, `post`, `post_call`, `transform_response`). `mod.rs` exposes the entrypoint that runs prepare then handler with no hooks (`()`). `route.rs`, where a host needs it, wraps the same two calls in a `CallMachine` whose `HostChannel` is the hooks, and pumps a stream through `open` and `deliver`. A handler takes `&impl RouteHooks<Error>` and never a `HostChannel` directly, so it runs without a coroutine. Keep provider transport and transformation details out of the machine driver
|
||||
|
||||
## Crate layering
|
||||
|
||||
|
|
@ -10,7 +12,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`, which each route implements over its host. Import every item from its canonical path. Never re-export another crate's items or give an item a second public path; the only re-export allowed is a private submodule surfacing its item at its module root (`mod error; pub use error::Error;`). Handlers belong in core or llms, never in a host crate
|
||||
A route module owns the call entrypoint, route request types (`*Request<'a>`), credential fallback, provider dispatch, and the handler glue that runs a provider config. Provider code never imports from core; when it needs the caller's hooks mid-call it goes through `litellm_llms::base_llm::ocr::handler::CallHooks`, the provider-level hooks OCR implements over its host until it folds into `litellm_host::hooks::RouteHooks`. Import every item from its canonical path. Never re-export another crate's items or give an item a second public path; the only re-export allowed is a private submodule surfacing its item at its module root (`mod error; pub use error::Error;`). Handlers belong in core or llms, never in a host crate
|
||||
|
||||
## Error placement
|
||||
|
||||
|
|
|
|||
|
|
@ -17,6 +17,7 @@ litellm-auth = { workspace = true, features = ["aws", "azure", "gcp"] }
|
|||
litellm-auth-aws.workspace = true
|
||||
litellm-http.workspace = true
|
||||
litellm-llms.workspace = true
|
||||
litellm-tracing.workspace = true
|
||||
moka.workspace = true
|
||||
mime_guess = "2.0.5"
|
||||
rand.workspace = true
|
||||
|
|
|
|||
|
|
@ -3,6 +3,7 @@ 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};
|
||||
|
||||
|
|
@ -14,6 +15,7 @@ pub(super) fn chat_completions_provider_config(provider: &str) -> Option<&'stati
|
|||
match provider {
|
||||
"anthropic" => Some(&ANTHROPIC_CHAT_COMPLETIONS_CONFIG),
|
||||
"bedrock" => Some(&BEDROCK_CHAT_COMPLETIONS_CONFIG),
|
||||
"openai_like" => Some(&OPENAI_LIKE_CHAT_COMPLETIONS_CONFIG),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,23 +1,68 @@
|
|||
use std::time::Duration;
|
||||
|
||||
use litellm_auth::AuthServices;
|
||||
use litellm_host::{
|
||||
event::{MachineEvent, RawResponse, RequestContext, WireRequest},
|
||||
hooks::RouteHooks,
|
||||
};
|
||||
use litellm_http::{Client, outbound::OutboundRequest, request::truncate_error_body};
|
||||
use litellm_llms::base_llm::{auth::resolve_auth, chat::transformation::ProviderChatResponseData};
|
||||
use litellm_llms::base_llm::{
|
||||
auth::{Authenticated, resolve_auth},
|
||||
chat::transformation::ProviderChatResponseData,
|
||||
};
|
||||
use litellm_types::utils::ChatCompletionsResponse;
|
||||
use serde_json::Value;
|
||||
|
||||
use super::{Error, prepare::prepare_provider_request};
|
||||
use super::Error;
|
||||
use crate::{
|
||||
chat_completions::types::{ProviderChatCompletionsRequest, ResolvedChatCompletionsRequest},
|
||||
chat_completions::types::ProviderChatCompletionsRequest,
|
||||
constants::CHAT_COMPLETIONS_TIMEOUT_SECS,
|
||||
};
|
||||
|
||||
pub(super) async fn execute_chat_completions_provider_call(
|
||||
pub(super) async fn execute(
|
||||
http: &Client,
|
||||
auth: &litellm_auth::AuthServices,
|
||||
request: ResolvedChatCompletionsRequest<'_>,
|
||||
auth: &AuthServices,
|
||||
request: ProviderChatCompletionsRequest,
|
||||
hooks: &impl RouteHooks<Error>,
|
||||
) -> Result<ChatCompletionsResponse, Error> {
|
||||
let request = prepare_provider_request(request)?;
|
||||
let outbound = outbound_request(auth, &request).await?;
|
||||
let ProviderChatCompletionsRequest {
|
||||
model,
|
||||
custom_llm_provider,
|
||||
config,
|
||||
url,
|
||||
body,
|
||||
optional_params,
|
||||
environment,
|
||||
timeout,
|
||||
api_key,
|
||||
} = request;
|
||||
let context = RequestContext {
|
||||
model: model.clone(),
|
||||
custom_llm_provider,
|
||||
optional_params: Value::Object(optional_params),
|
||||
secret_fields: Vec::new(),
|
||||
api_key,
|
||||
};
|
||||
let authenticated = resolve_auth(auth, environment, &|key| std::env::var(key).ok()).await?;
|
||||
let wire = hooks
|
||||
.before_send(
|
||||
WireRequest {
|
||||
url,
|
||||
headers: authenticated.headers,
|
||||
body,
|
||||
},
|
||||
context,
|
||||
)
|
||||
.await?;
|
||||
let outbound = outbound_request(
|
||||
Authenticated {
|
||||
headers: wire.headers,
|
||||
signer: authenticated.signer,
|
||||
},
|
||||
wire.url,
|
||||
&wire.body,
|
||||
timeout,
|
||||
)?;
|
||||
|
||||
let response = outbound.send(http).await.map_err(|err| {
|
||||
// Failing to establish the connection means the request never went out,
|
||||
|
|
@ -41,13 +86,17 @@ pub(super) async fn execute_chat_completions_provider_call(
|
|||
body: truncate_error_body(&text),
|
||||
}));
|
||||
}
|
||||
hooks
|
||||
.emit(MachineEvent::ResponseReceived {
|
||||
raw: RawResponse { body: text.clone() },
|
||||
})
|
||||
.await?;
|
||||
|
||||
let body: Value = serde_json::from_str(&text).map_err(|err| {
|
||||
Error::InvalidResponse(format!("invalid chat completions response JSON: {err}"))
|
||||
})?;
|
||||
request
|
||||
.config
|
||||
.transform_response(&request.model, ProviderChatResponseData { body })
|
||||
config
|
||||
.transform_response(&model, ProviderChatResponseData { body })
|
||||
.map_err(Error::from)
|
||||
.map_err(as_response_error)
|
||||
}
|
||||
|
|
@ -69,21 +118,17 @@ pub(super) fn as_response_error(err: Error) -> Error {
|
|||
}
|
||||
}
|
||||
|
||||
pub(super) async fn outbound_request(
|
||||
auth: &litellm_auth::AuthServices,
|
||||
request: &ProviderChatCompletionsRequest,
|
||||
pub(super) fn outbound_request(
|
||||
authenticated: Authenticated,
|
||||
url: String,
|
||||
body: &Value,
|
||||
timeout: Option<Duration>,
|
||||
) -> Result<OutboundRequest, Error> {
|
||||
let env_lookup = |key: &str| std::env::var(key).ok();
|
||||
let authenticated = resolve_auth(auth, request.environment.clone(), &env_lookup).await?;
|
||||
crate::outbound::outbound_request(
|
||||
authenticated,
|
||||
request.url.clone(),
|
||||
&request.body,
|
||||
Some(
|
||||
request
|
||||
.timeout
|
||||
.unwrap_or(Duration::from_secs(CHAT_COMPLETIONS_TIMEOUT_SECS)),
|
||||
),
|
||||
url,
|
||||
body,
|
||||
Some(timeout.unwrap_or(Duration::from_secs(CHAT_COMPLETIONS_TIMEOUT_SECS))),
|
||||
)
|
||||
.map_err(|error| match error {
|
||||
// Python drops the caller's copy and prefers a forwarded Authorization
|
||||
|
|
@ -97,7 +142,133 @@ pub(super) async fn outbound_request(
|
|||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{Error, as_response_error};
|
||||
use std::sync::Mutex;
|
||||
|
||||
use rstest::rstest;
|
||||
use serde_json::json;
|
||||
use wiremock::{Mock, MockServer, Request, ResponseTemplate, matchers::any};
|
||||
|
||||
use super::*;
|
||||
use crate::chat_completions::{
|
||||
prepare::{prepare_provider_request, resolve_request},
|
||||
types::ChatCompletionsRequest,
|
||||
};
|
||||
|
||||
const ANTHROPIC_MESSAGE: &str = r#"{"id":"msg_1","type":"message","role":"assistant","model":"claude-sonnet-4-5","content":[{"type":"text","text":"hello"}],"stop_reason":"end_turn","stop_sequence":null,"usage":{"input_tokens":11,"output_tokens":4}}"#;
|
||||
|
||||
/// Rewrites the outgoing request and records what the call reports back.
|
||||
#[derive(Default)]
|
||||
struct RecordingHooks {
|
||||
contexts: Mutex<Vec<RequestContext>>,
|
||||
raw: Mutex<Vec<String>>,
|
||||
}
|
||||
|
||||
impl RouteHooks<Error> for RecordingHooks {
|
||||
async fn before_send(
|
||||
&self,
|
||||
wire: WireRequest,
|
||||
context: RequestContext,
|
||||
) -> Result<WireRequest, Error> {
|
||||
self.contexts.lock().unwrap().push(context);
|
||||
let mut body = wire.body;
|
||||
body["system"] = json!("added by the host");
|
||||
Ok(WireRequest {
|
||||
headers: wire
|
||||
.headers
|
||||
.into_iter()
|
||||
.chain([("x-host".to_string(), "seen".to_string())])
|
||||
.collect(),
|
||||
body,
|
||||
..wire
|
||||
})
|
||||
}
|
||||
|
||||
async fn emit(&self, event: MachineEvent) -> Result<(), Error> {
|
||||
let MachineEvent::ResponseReceived { raw } = event;
|
||||
self.raw.lock().unwrap().push(raw.body);
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
fn prepared(api_base: &str) -> ProviderChatCompletionsRequest {
|
||||
prepare_provider_request(
|
||||
resolve_request(ChatCompletionsRequest {
|
||||
model: "anthropic/claude-sonnet-4-5",
|
||||
messages: json!([{"role": "user", "content": "hi"}]),
|
||||
optional_params: json!({"max_tokens": 16}).as_object().unwrap().clone(),
|
||||
api_key: Some("sk-test"),
|
||||
api_base: Some(api_base),
|
||||
custom_llm_provider: None,
|
||||
extra_headers: None,
|
||||
timeout: None,
|
||||
})
|
||||
.unwrap(),
|
||||
)
|
||||
.unwrap()
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn the_hooks_rewrite_the_wire_request_and_see_the_raw_response() {
|
||||
let upstream = MockServer::start().await;
|
||||
Mock::given(any())
|
||||
.respond_with(
|
||||
ResponseTemplate::new(200).set_body_raw(ANTHROPIC_MESSAGE, "application/json"),
|
||||
)
|
||||
.mount(&upstream)
|
||||
.await;
|
||||
let hooks = RecordingHooks::default();
|
||||
|
||||
execute(
|
||||
&Client::plain_for_test(),
|
||||
&AuthServices::default(),
|
||||
prepared(&upstream.uri()),
|
||||
&hooks,
|
||||
)
|
||||
.await
|
||||
.expect("chat completions call succeeds");
|
||||
|
||||
let [request] = <[Request; 1]>::try_from(upstream.received_requests().await.unwrap())
|
||||
.unwrap_or_else(|requests| panic!("one request, saw {}", requests.len()));
|
||||
let sent: Value = serde_json::from_slice(&request.body).unwrap();
|
||||
assert_eq!(sent["system"], "added by the host");
|
||||
assert_eq!(request.headers["x-host"], "seen");
|
||||
assert_eq!(request.headers["x-api-key"], "sk-test");
|
||||
let [context] = <[RequestContext; 1]>::try_from(hooks.contexts.into_inner().unwrap())
|
||||
.unwrap_or_else(|seen| panic!("before_send runs once, saw {}", seen.len()));
|
||||
assert_eq!(
|
||||
(context.model.as_str(), context.custom_llm_provider.as_str()),
|
||||
("claude-sonnet-4-5", "anthropic")
|
||||
);
|
||||
assert_eq!(context.optional_params, json!({"max_tokens": 16}));
|
||||
assert_eq!(hooks.raw.into_inner().unwrap(), [ANTHROPIC_MESSAGE]);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn an_upstream_failure_is_not_reported_as_a_received_response() {
|
||||
let upstream = MockServer::start().await;
|
||||
Mock::given(any())
|
||||
.respond_with(ResponseTemplate::new(500).set_body_string("boom"))
|
||||
.mount(&upstream)
|
||||
.await;
|
||||
let hooks = RecordingHooks::default();
|
||||
|
||||
let error = execute(
|
||||
&Client::plain_for_test(),
|
||||
&AuthServices::default(),
|
||||
prepared(&upstream.uri()),
|
||||
&hooks,
|
||||
)
|
||||
.await
|
||||
.expect_err("the upstream failure fails the call");
|
||||
|
||||
assert!(matches!(
|
||||
error,
|
||||
Error::Transport(litellm_http::transport::Error::Http { status: 500, .. })
|
||||
));
|
||||
assert!(hooks.raw.into_inner().unwrap().is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn response_errors_collapse_to_one_variant_that_can_only_mean_already_sent() {
|
||||
|
|
|
|||
|
|
@ -11,10 +11,9 @@ pub use crate::error::RouteError as Error;
|
|||
mod common_utils;
|
||||
pub(crate) mod handler;
|
||||
mod prepare;
|
||||
use handler::execute_chat_completions_provider_call;
|
||||
use litellm_http::{ClientVariant, HttpClientConfig};
|
||||
use litellm_types::utils::ChatCompletionsResponse;
|
||||
use prepare::{parse_messages, resolve_provider_config, resolve_request};
|
||||
use prepare::{parse_messages, prepare_provider_request, resolve_provider_config, resolve_request};
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
use crate::chat_completions::types::ChatCompletionsRequest;
|
||||
|
|
@ -24,9 +23,9 @@ pub async fn chat_completions(
|
|||
config: &HttpClientConfig,
|
||||
request: ChatCompletionsRequest<'_>,
|
||||
) -> Result<ChatCompletionsResponse, Error> {
|
||||
let request = resolve_request(request)?;
|
||||
let http = resources.pool.client(config, ClientVariant::Provider)?;
|
||||
execute_chat_completions_provider_call(&http, &resources.auth, request).await
|
||||
let request = prepare_provider_request(resolve_request(request)?)?;
|
||||
handler::execute(&http, &resources.auth, request, &()).await
|
||||
}
|
||||
|
||||
/// Whether the core would accept this request, without resolving credentials or
|
||||
|
|
@ -41,9 +40,10 @@ pub fn chat_completions_decline_reason(
|
|||
messages: Value,
|
||||
optional_params: &Map<String, Value>,
|
||||
) -> Option<&'static str> {
|
||||
let Ok((_, config)) = resolve_provider_config(model, custom_llm_provider) else {
|
||||
let Ok(resolved) = resolve_provider_config(model, custom_llm_provider) else {
|
||||
return Some("provider is not on the rust chat completions path");
|
||||
};
|
||||
let config = resolved.config;
|
||||
let Ok(messages) = parse_messages(messages) else {
|
||||
return Some("unreadable message list");
|
||||
};
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
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},
|
||||
|
|
@ -14,10 +15,16 @@ use crate::chat_completions::types::{
|
|||
ChatCompletionsRequest, ProviderChatCompletionsRequest, ResolvedChatCompletionsRequest,
|
||||
};
|
||||
|
||||
pub(super) struct ResolvedProvider {
|
||||
pub(super) model: String,
|
||||
pub(super) custom_llm_provider: String,
|
||||
pub(super) config: &'static dyn BaseConfig,
|
||||
}
|
||||
|
||||
pub(super) fn resolve_provider_config<'a>(
|
||||
model: &'a str,
|
||||
custom_llm_provider: Option<&'a str>,
|
||||
) -> Result<(String, &'static dyn BaseConfig), Error> {
|
||||
) -> Result<ResolvedProvider, Error> {
|
||||
let provider_info = get_custom_llm_provider(model, custom_llm_provider)
|
||||
.or_else(|| {
|
||||
custom_llm_provider.map(|provider| CustomLlmProvider {
|
||||
|
|
@ -32,7 +39,11 @@ pub(super) fn resolve_provider_config<'a>(
|
|||
})?;
|
||||
let config = chat_completions_provider_config(provider_info.custom_llm_provider)
|
||||
.ok_or_else(|| Error::InvalidProvider(provider_info.custom_llm_provider.to_string()))?;
|
||||
Ok((provider_info.model.to_string(), config))
|
||||
Ok(ResolvedProvider {
|
||||
model: provider_info.model.to_string(),
|
||||
custom_llm_provider: provider_info.custom_llm_provider.to_string(),
|
||||
config,
|
||||
})
|
||||
}
|
||||
|
||||
pub(super) fn parse_messages(messages: Value) -> Result<Vec<ChatMessage>, Error> {
|
||||
|
|
@ -43,7 +54,11 @@ pub(super) fn parse_messages(messages: Value) -> Result<Vec<ChatMessage>, Error>
|
|||
pub(super) fn resolve_request(
|
||||
request: ChatCompletionsRequest<'_>,
|
||||
) -> Result<ResolvedChatCompletionsRequest<'_>, Error> {
|
||||
let (model, config) = resolve_provider_config(request.model, request.custom_llm_provider)?;
|
||||
let ResolvedProvider {
|
||||
model,
|
||||
custom_llm_provider,
|
||||
config,
|
||||
} = resolve_provider_config(request.model, request.custom_llm_provider)?;
|
||||
let messages = parse_messages(request.messages)?;
|
||||
if messages.is_empty() {
|
||||
return Err(Error::InvalidRequest(
|
||||
|
|
@ -55,6 +70,7 @@ pub(super) fn resolve_request(
|
|||
}
|
||||
Ok(ResolvedChatCompletionsRequest {
|
||||
model,
|
||||
custom_llm_provider,
|
||||
config,
|
||||
messages,
|
||||
optional_params: request.optional_params,
|
||||
|
|
@ -99,15 +115,18 @@ pub(super) fn prepare_provider_request(
|
|||
&env_lookup,
|
||||
)?;
|
||||
let transformed =
|
||||
config.transform_request(&model, request.messages, request.optional_params)?;
|
||||
config.transform_request(&model, request.messages, request.optional_params.clone())?;
|
||||
|
||||
Ok(ProviderChatCompletionsRequest {
|
||||
model,
|
||||
custom_llm_provider: request.custom_llm_provider,
|
||||
config,
|
||||
url,
|
||||
body: transformed.body,
|
||||
optional_params: request.optional_params,
|
||||
environment,
|
||||
timeout: request.timeout,
|
||||
api_key: request.api_key.map(|key| SecretValue::new(key.to_string())),
|
||||
})
|
||||
}
|
||||
|
||||
|
|
@ -449,11 +468,19 @@ mod tests {
|
|||
json!("abc-123"),
|
||||
)]));
|
||||
let prepared = prepare_chat_completions_call(call).expect("prepares");
|
||||
let signed = crate::chat_completions::handler::outbound_request(
|
||||
let authenticated = resolve_auth(
|
||||
&litellm_auth::AuthServices::default(),
|
||||
&prepared,
|
||||
prepared.environment,
|
||||
&|_| None,
|
||||
)
|
||||
.await
|
||||
.expect("resolves");
|
||||
let signed = crate::chat_completions::handler::outbound_request(
|
||||
authenticated,
|
||||
prepared.url,
|
||||
&prepared.body,
|
||||
prepared.timeout,
|
||||
)
|
||||
.expect("signs");
|
||||
|
||||
let authorization = signed
|
||||
|
|
@ -502,11 +529,19 @@ mod tests {
|
|||
call.api_key = None;
|
||||
call.extra_headers = Some(Map::from_iter([(forwarded.to_string(), json!("forged"))]));
|
||||
let prepared = prepare_chat_completions_call(call).expect("prepares");
|
||||
let error = crate::chat_completions::handler::outbound_request(
|
||||
let authenticated = resolve_auth(
|
||||
&litellm_auth::AuthServices::default(),
|
||||
&prepared,
|
||||
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(_)),
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
use std::time::Duration;
|
||||
|
||||
use litellm_auth::SecretValue;
|
||||
use litellm_llms::base_llm::{auth::ValidatedEnvironment, chat::transformation::BaseConfig};
|
||||
use litellm_types::llms::openai::ChatMessage;
|
||||
use serde_json::{Map, Value};
|
||||
|
|
@ -23,6 +24,7 @@ pub struct ChatCompletionsRequest<'a> {
|
|||
|
||||
pub struct ResolvedChatCompletionsRequest<'a> {
|
||||
pub model: String,
|
||||
pub custom_llm_provider: String,
|
||||
pub config: &'static dyn BaseConfig,
|
||||
pub messages: Vec<ChatMessage>,
|
||||
pub optional_params: Map<String, Value>,
|
||||
|
|
@ -34,11 +36,16 @@ pub struct ResolvedChatCompletionsRequest<'a> {
|
|||
|
||||
pub struct ProviderChatCompletionsRequest {
|
||||
pub model: String,
|
||||
pub custom_llm_provider: String,
|
||||
pub config: &'static dyn BaseConfig,
|
||||
pub url: String,
|
||||
pub body: Value,
|
||||
/// The route's parameters before the provider transformation, reported to the host
|
||||
/// beside the wire request.
|
||||
pub optional_params: Map<String, Value>,
|
||||
/// The forwarded and default headers plus how the call authenticates; the credential
|
||||
/// itself is applied when the request is sent.
|
||||
pub environment: ValidatedEnvironment,
|
||||
pub timeout: Option<Duration>,
|
||||
pub api_key: Option<SecretValue>,
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,20 +1,113 @@
|
|||
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_http::transport::Error as TransportError;
|
||||
use litellm_llms::base_llm::{
|
||||
anthropic_messages::transformation::BaseAnthropicMessagesConfig, auth::Authenticated,
|
||||
anthropic_messages::{
|
||||
streaming::{ByteStream, StreamDecoder, encode_anthropic_sse},
|
||||
transformation::BaseAnthropicMessagesConfig,
|
||||
},
|
||||
auth::{Authenticated, resolve_auth},
|
||||
};
|
||||
use litellm_tracing::{ByteChunk, debug};
|
||||
use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse;
|
||||
use serde_json::Value;
|
||||
|
||||
use super::{Error, common_utils::truncate_error_body};
|
||||
use super::{
|
||||
Error, MessagesResponse, common_utils::truncate_error_body, prepare::ProviderMessagesRequest,
|
||||
};
|
||||
use crate::{constants::MESSAGES_TIMEOUT_SECS, outbound::outbound_request};
|
||||
|
||||
pub(super) fn network(error: reqwest::Error) -> Error {
|
||||
pub(super) async fn execute(
|
||||
http: &litellm_http::Client,
|
||||
auth: &AuthServices,
|
||||
request: ProviderMessagesRequest,
|
||||
hooks: &impl RouteHooks<Error>,
|
||||
) -> Result<MessagesResponse, Error> {
|
||||
let ProviderMessagesRequest {
|
||||
provider,
|
||||
url,
|
||||
body,
|
||||
environment,
|
||||
timeout,
|
||||
api_key,
|
||||
} = request;
|
||||
let stream = body.params.stream == Some(true);
|
||||
let context = RequestContext {
|
||||
model: body.model.clone(),
|
||||
custom_llm_provider: provider.as_str().to_string(),
|
||||
optional_params: serde_json::to_value(&body.params).map_err(serialize_failure)?,
|
||||
secret_fields: Vec::new(),
|
||||
api_key,
|
||||
};
|
||||
let authenticated = resolve_auth(auth, environment, &|key| std::env::var(key).ok()).await?;
|
||||
let wire = hooks
|
||||
.before_send(
|
||||
WireRequest {
|
||||
url,
|
||||
headers: authenticated.headers,
|
||||
body: serde_json::to_value(&body).map_err(serialize_failure)?,
|
||||
},
|
||||
context,
|
||||
)
|
||||
.await?;
|
||||
let provider_name = provider.as_str();
|
||||
debug!(provider = provider_name, stream, body = %wire.body, "provider request");
|
||||
let response = send(
|
||||
http,
|
||||
Authenticated {
|
||||
headers: wire.headers,
|
||||
signer: authenticated.signer,
|
||||
},
|
||||
&wire.url,
|
||||
&wire.body,
|
||||
timeout,
|
||||
)
|
||||
.await?;
|
||||
debug!(
|
||||
provider = provider_name,
|
||||
status = response.status().as_u16(),
|
||||
"provider response headers"
|
||||
);
|
||||
if !response.status().is_success() {
|
||||
return Err(provider_error(response).await);
|
||||
}
|
||||
let config = provider.config();
|
||||
if stream {
|
||||
return Ok(streaming_response(
|
||||
response,
|
||||
config.stream_decoder(),
|
||||
provider_name,
|
||||
));
|
||||
}
|
||||
let text = response.text().await.map_err(network)?;
|
||||
debug!(body = text.as_str(), "provider response body");
|
||||
hooks
|
||||
.emit(MachineEvent::ResponseReceived {
|
||||
raw: RawResponse { body: text.clone() },
|
||||
})
|
||||
.await?;
|
||||
decode_response(config, &body.model, &text)
|
||||
.map(|message| MessagesResponse::Message(Box::new(message)))
|
||||
}
|
||||
|
||||
fn serialize_failure(err: serde_json::Error) -> Error {
|
||||
Error::InvalidRequest(format!(
|
||||
"failed to serialize Anthropic messages request: {err}"
|
||||
))
|
||||
}
|
||||
|
||||
fn network(error: reqwest::Error) -> Error {
|
||||
Error::Transport(TransportError::Network(error.to_string()))
|
||||
}
|
||||
|
||||
pub(super) async fn send(
|
||||
async fn send(
|
||||
http: &litellm_http::Client,
|
||||
authenticated: Authenticated,
|
||||
url: &str,
|
||||
|
|
@ -30,18 +123,21 @@ pub(super) async fn send(
|
|||
request.send(http).await.map_err(network)
|
||||
}
|
||||
|
||||
pub(super) async fn provider_error(response: reqwest::Response) -> Error {
|
||||
async fn provider_error(response: reqwest::Response) -> Error {
|
||||
let status = response.status().as_u16();
|
||||
match response.text().await {
|
||||
Ok(text) => Error::Transport(TransportError::Http {
|
||||
status,
|
||||
body: truncate_error_body(&text),
|
||||
}),
|
||||
Ok(text) => {
|
||||
litellm_tracing::debug!(status, body = text.as_str(), "provider error body");
|
||||
Error::Transport(TransportError::Http {
|
||||
status,
|
||||
body: truncate_error_body(&text),
|
||||
})
|
||||
}
|
||||
Err(error) => network(error),
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn decode_response(
|
||||
fn decode_response(
|
||||
config: &dyn BaseAnthropicMessagesConfig,
|
||||
model: &str,
|
||||
text: &str,
|
||||
|
|
@ -52,3 +148,97 @@ pub(super) fn decode_response(
|
|||
.transform_anthropic_messages_response(model, response)
|
||||
.map_err(Error::from)
|
||||
}
|
||||
|
||||
fn streaming_response(
|
||||
response: reqwest::Response,
|
||||
decoder: Option<StreamDecoder>,
|
||||
provider: &'static str,
|
||||
) -> MessagesResponse {
|
||||
let headers = response
|
||||
.headers()
|
||||
.iter()
|
||||
.filter_map(|(name, value)| Some((name.to_string(), value.to_str().ok()?.to_string())))
|
||||
.collect();
|
||||
let chunks = match decoder {
|
||||
None => futures_util::stream::try_unfold(response, move |mut response| async move {
|
||||
let chunk = response.chunk().await.map_err(network)?;
|
||||
Ok(chunk.map(|chunk| {
|
||||
log_chunk(provider, "provider_response", &chunk);
|
||||
(chunk, response)
|
||||
}))
|
||||
})
|
||||
.boxed(),
|
||||
Some(decode) => decoded_chunks(response, decode, provider),
|
||||
};
|
||||
MessagesResponse::Stream { headers, chunks }
|
||||
}
|
||||
|
||||
fn decoded_chunks(
|
||||
response: reqwest::Response,
|
||||
decode: StreamDecoder,
|
||||
provider: &'static str,
|
||||
) -> BoxStream<'static, Result<Bytes, Error>> {
|
||||
let bytes: ByteStream = response
|
||||
.bytes_stream()
|
||||
.inspect_ok(move |chunk| log_chunk(provider, "provider_response", chunk))
|
||||
.map_err(std::io::Error::other)
|
||||
.boxed();
|
||||
futures_util::stream::try_unfold(decode(bytes), move |mut events| async move {
|
||||
let Some(event) = events.try_next().await? else {
|
||||
return Ok(None);
|
||||
};
|
||||
let chunk = encode_anthropic_sse(&event)?;
|
||||
log_chunk(provider, "client_response", &chunk);
|
||||
Ok(Some((chunk, events)))
|
||||
})
|
||||
.boxed()
|
||||
}
|
||||
|
||||
fn log_chunk(provider: &str, stage: &str, data: &Bytes) {
|
||||
let chunk = ByteChunk::new(data);
|
||||
debug!(provider, stage, encoding = chunk.encoding(), chunk = %chunk, "stream chunk");
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use litellm_llms::base_llm::anthropic_messages::streaming::anthropic_sse_event_stream;
|
||||
use rstest::rstest;
|
||||
use wiremock::{Mock, MockServer, ResponseTemplate, matchers::any};
|
||||
|
||||
use super::*;
|
||||
|
||||
#[rstest]
|
||||
#[case::event(
|
||||
"data: {\"type\":\"ping\"}\n\n",
|
||||
Some("event: ping\ndata: {\"type\":\"ping\"}\n\n")
|
||||
)]
|
||||
#[case::invalid_event("data: invalid\n\ndata: {\"type\":\"ping\"}\n\n", None)]
|
||||
#[tokio::test]
|
||||
async fn decoded_streams_encode_events_and_stop_at_the_first_error(
|
||||
#[case] body: &'static str,
|
||||
#[case] expected: Option<&str>,
|
||||
) {
|
||||
let upstream = MockServer::start().await;
|
||||
Mock::given(any())
|
||||
.respond_with(ResponseTemplate::new(200).set_body_raw(body, "text/event-stream"))
|
||||
.mount(&upstream)
|
||||
.await;
|
||||
let response = litellm_http::Client::plain_for_test()
|
||||
.get(upstream.uri())
|
||||
.send()
|
||||
.await
|
||||
.unwrap();
|
||||
let MessagesResponse::Stream { mut chunks, .. } =
|
||||
streaming_response(response, Some(anthropic_sse_event_stream), "test")
|
||||
else {
|
||||
panic!("a streaming response returns chunks");
|
||||
};
|
||||
|
||||
let chunk = chunks.next().await.unwrap();
|
||||
match expected {
|
||||
Some(expected) => assert_eq!(chunk.unwrap().as_ref(), expected.as_bytes()),
|
||||
None => assert!(matches!(chunk, Err(Error::InvalidResponse(_))), "{chunk:?}"),
|
||||
}
|
||||
assert!(chunks.next().await.is_none());
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,39 +1,27 @@
|
|||
//! The Anthropic Messages call, the Rust equivalent of Python's
|
||||
//! `litellm.messages()`.
|
||||
//! The Anthropic Messages call, the Rust equivalent of Python's `litellm.messages()`.
|
||||
//!
|
||||
//! [`route`] is the call as a machine a host drives, streaming or not. [`messages`] runs
|
||||
//! it in process for a caller that already holds the request and wants the message.
|
||||
//! [`messages`] prepares the provider request and sends it in process. [`route`] runs the
|
||||
//! same two steps as a machine for a host that answers the call's operations itself.
|
||||
|
||||
pub mod types;
|
||||
pub use crate::error::RouteError as Error;
|
||||
mod common_utils;
|
||||
mod handler;
|
||||
mod prepare;
|
||||
pub mod route;
|
||||
use std::sync::Arc;
|
||||
mod types;
|
||||
|
||||
use litellm_http::{ClientVariant, HttpClientConfig};
|
||||
use litellm_secrets::source::EnvironmentSecrets;
|
||||
use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse;
|
||||
use route::{LocalMessagesHost, MessagesCall, MessagesOutput, messages_machine};
|
||||
use litellm_secrets::source::SecretSource;
|
||||
|
||||
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<AnthropicMessagesResponse, Error> {
|
||||
let secrets = Arc::new(EnvironmentSecrets::python_compatible(
|
||||
resources.pool.client(config, ClientVariant::Provider)?,
|
||||
));
|
||||
match litellm_host::run::run(
|
||||
messages_machine(resources, config, secrets)?,
|
||||
&LocalMessagesHost::new(call),
|
||||
)
|
||||
.await?
|
||||
{
|
||||
MessagesOutput::Message(message) => Ok(*message),
|
||||
MessagesOutput::Streamed => Err(Error::Unsupported(
|
||||
"streamed responses need a streaming host",
|
||||
)),
|
||||
}
|
||||
) -> 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
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,3 +1,6 @@
|
|||
use std::time::Duration;
|
||||
|
||||
use litellm_auth::SecretValue;
|
||||
use litellm_core_utils::{
|
||||
dot_notation_indexing::delete_nested_value,
|
||||
get_llm_provider_logic::{CustomLlmProvider, get_custom_llm_provider},
|
||||
|
|
@ -11,21 +14,42 @@ use litellm_llms::{
|
|||
auth::{ValidatedEnvironment, with_default_headers},
|
||||
},
|
||||
};
|
||||
use litellm_secrets::source::SecretSource;
|
||||
use litellm_types::llms::anthropic_messages::anthropic_request::AnthropicMessagesRequest;
|
||||
|
||||
use super::{
|
||||
Error,
|
||||
Error, MessagesCall,
|
||||
common_utils::{MessagesProvider, string_headers},
|
||||
route::MessagesCall,
|
||||
types::ProviderMessagesRequest,
|
||||
types::invalid_request,
|
||||
};
|
||||
|
||||
pub(super) struct ResolvedProvider {
|
||||
pub(super) model: String,
|
||||
pub(super) provider: MessagesProvider,
|
||||
struct ResolvedProvider {
|
||||
model: String,
|
||||
provider: MessagesProvider,
|
||||
}
|
||||
|
||||
pub(super) fn resolve_provider(
|
||||
pub(super) struct ProviderMessagesRequest {
|
||||
pub(super) provider: MessagesProvider,
|
||||
pub(super) url: String,
|
||||
pub(super) body: AnthropicMessagesRequest,
|
||||
pub(super) environment: ValidatedEnvironment,
|
||||
pub(super) timeout: Option<Duration>,
|
||||
/// The caller's own credential, reported to the host beside the wire request.
|
||||
pub(super) api_key: Option<SecretValue>,
|
||||
}
|
||||
|
||||
pub(super) async fn prepare(
|
||||
call: MessagesCall,
|
||||
secrets: &dyn SecretSource,
|
||||
) -> Result<ProviderMessagesRequest, Error> {
|
||||
let resolved = resolve_provider(&call.body.model, call.custom_llm_provider.as_deref())?;
|
||||
let secrets = secrets
|
||||
.resolve(resolved.provider.config().secret_names())
|
||||
.await?;
|
||||
prepare_provider_request(call, resolved, secrets.as_ref())
|
||||
}
|
||||
|
||||
fn resolve_provider(
|
||||
model: &str,
|
||||
custom_llm_provider: Option<&str>,
|
||||
) -> Result<ResolvedProvider, Error> {
|
||||
|
|
@ -53,7 +77,7 @@ pub(super) fn resolve_provider(
|
|||
})
|
||||
}
|
||||
|
||||
pub(super) fn prepare_provider_request(
|
||||
fn prepare_provider_request(
|
||||
call: MessagesCall,
|
||||
resolved: ResolvedProvider,
|
||||
secrets: &dyn Lookup,
|
||||
|
|
@ -113,13 +137,10 @@ pub(super) fn prepare_provider_request(
|
|||
body: transformed,
|
||||
environment,
|
||||
timeout,
|
||||
api_key: api_key.map(SecretValue::new),
|
||||
})
|
||||
}
|
||||
|
||||
pub(super) fn invalid_request(err: serde_json::Error) -> Error {
|
||||
Error::InvalidRequest(format!("invalid Anthropic messages request: {err}"))
|
||||
}
|
||||
|
||||
fn without_additional_drop_params(
|
||||
request: AnthropicMessagesRequest,
|
||||
paths: &[String],
|
||||
|
|
@ -145,7 +166,7 @@ mod tests {
|
|||
use serde_json::{Map, Value, json};
|
||||
|
||||
use super::*;
|
||||
use crate::messages::types::MessagesShaping;
|
||||
use crate::messages::MessagesShaping;
|
||||
|
||||
#[fixture]
|
||||
fn shaping() -> MessagesShaping {
|
||||
|
|
|
|||
|
|
@ -1,55 +1,20 @@
|
|||
use std::{
|
||||
convert::Infallible,
|
||||
sync::{Arc, Mutex},
|
||||
time::Duration,
|
||||
};
|
||||
|
||||
use bytes::Bytes;
|
||||
use futures_util::StreamExt;
|
||||
use litellm_auth::SecretValue;
|
||||
use futures_util::TryStreamExt;
|
||||
use litellm_host::{
|
||||
event::{MachineEvent, RawResponse, RequestContext, WireRequest},
|
||||
host::{Demand, Host},
|
||||
machine::{CallMachine, HostChannel, MachineFault},
|
||||
protocol::Protocol,
|
||||
};
|
||||
use litellm_http::{Client, ClientVariant, HttpClientConfig};
|
||||
use litellm_llms::base_llm::{
|
||||
anthropic_messages::streaming::{ByteStream, StreamDecoder, encode_anthropic_sse},
|
||||
auth::{Authenticated, resolve_auth},
|
||||
};
|
||||
use litellm_secrets::source::SecretSource;
|
||||
use litellm_types::{
|
||||
llms::anthropic_messages::{
|
||||
anthropic_request::AnthropicMessagesRequest, anthropic_response::AnthropicMessagesResponse,
|
||||
},
|
||||
utils::ProviderSpecificHeaders,
|
||||
};
|
||||
use serde_json::{Map, Value};
|
||||
use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse;
|
||||
|
||||
use super::{
|
||||
Error,
|
||||
handler::{decode_response, network, provider_error, send},
|
||||
prepare::{invalid_request, prepare_provider_request, resolve_provider},
|
||||
types::MessagesShaping,
|
||||
};
|
||||
|
||||
/// The caller's request as the host projects it.
|
||||
pub struct MessagesCall {
|
||||
pub body: AnthropicMessagesRequest,
|
||||
pub api_key: Option<String>,
|
||||
pub api_base: Option<String>,
|
||||
pub custom_llm_provider: Option<String>,
|
||||
pub extra_headers: Option<Map<String, Value>>,
|
||||
pub provider_specific_header: Option<ProviderSpecificHeaders>,
|
||||
pub timeout: Option<Duration>,
|
||||
pub shaping: MessagesShaping,
|
||||
}
|
||||
|
||||
/// Parses a caller's raw body, failing the way the route fails for any invalid request.
|
||||
pub fn messages_body(body: Map<String, Value>) -> Result<AnthropicMessagesRequest, Error> {
|
||||
serde_json::from_value(Value::Object(body)).map_err(invalid_request)
|
||||
}
|
||||
use super::{Error, MessagesCall, MessagesResponse, handler::execute, prepare::prepare};
|
||||
|
||||
pub enum MessagesOutput {
|
||||
Message(Box<AnthropicMessagesResponse>),
|
||||
|
|
@ -121,134 +86,35 @@ pub fn messages_machine(
|
|||
let http = resources.pool.client(config, ClientVariant::Provider)?;
|
||||
let auth = resources.auth.clone();
|
||||
Ok(CallMachine::new(move |host| {
|
||||
Box::pin(execute(host, http.clone(), auth.clone(), secrets.clone()))
|
||||
Box::pin(drive(host, http, auth, secrets))
|
||||
}))
|
||||
}
|
||||
|
||||
async fn execute(
|
||||
/// The call as its host sees it: projection first, then the same prepare and execute as
|
||||
/// [`super::messages`], with each chunk of a stream handed over as it arrives.
|
||||
async fn drive(
|
||||
host: MessagesHost,
|
||||
http: Client,
|
||||
auth: Arc<litellm_auth::AuthServices>,
|
||||
secrets: Arc<dyn SecretSource>,
|
||||
) -> Result<MessagesOutput, Error> {
|
||||
let call = host.project().await?;
|
||||
let resolved = resolve_provider(&call.body.model, call.custom_llm_provider.as_deref())?;
|
||||
let secrets = secrets
|
||||
.resolve(resolved.provider.config().secret_names())
|
||||
.await?;
|
||||
let api_key = call.api_key.clone().map(SecretValue::new);
|
||||
let request = prepare_provider_request(call, resolved, secrets.as_ref())?;
|
||||
let context = RequestContext {
|
||||
model: request.body.model.clone(),
|
||||
custom_llm_provider: request.provider.as_str().to_string(),
|
||||
optional_params: serde_json::to_value(&request.body.params).map_err(serialize_failure)?,
|
||||
secret_fields: Vec::new(),
|
||||
api_key,
|
||||
};
|
||||
let stream = request.body.params.stream == Some(true);
|
||||
let config = request.provider.config();
|
||||
let body = serde_json::to_value(&request.body).map_err(serialize_failure)?;
|
||||
let env_lookup = |key: &str| std::env::var(key).ok();
|
||||
let authenticated = resolve_auth(&auth, request.environment, &env_lookup).await?;
|
||||
let wire = host
|
||||
.before_send(
|
||||
WireRequest {
|
||||
url: request.url,
|
||||
headers: authenticated.headers,
|
||||
body,
|
||||
},
|
||||
context,
|
||||
)
|
||||
.await?;
|
||||
let response = send(
|
||||
&http,
|
||||
Authenticated {
|
||||
headers: wire.headers,
|
||||
signer: authenticated.signer,
|
||||
},
|
||||
&wire.url,
|
||||
&wire.body,
|
||||
request.timeout,
|
||||
)
|
||||
.await?;
|
||||
if !response.status().is_success() {
|
||||
return Err(provider_error(response).await);
|
||||
}
|
||||
if stream {
|
||||
return relay(&host, response, config.stream_decoder()).await;
|
||||
}
|
||||
let text = response.text().await.map_err(network)?;
|
||||
host.emit(MachineEvent::ResponseReceived {
|
||||
raw: RawResponse { body: text.clone() },
|
||||
})
|
||||
.await?;
|
||||
decode_response(config, &request.body.model, &text)
|
||||
.map(|message| MessagesOutput::Message(Box::new(message)))
|
||||
}
|
||||
|
||||
fn serialize_failure(err: serde_json::Error) -> Error {
|
||||
Error::InvalidRequest(format!(
|
||||
"failed to serialize Anthropic messages request: {err}"
|
||||
))
|
||||
}
|
||||
|
||||
/// Hands each upstream chunk to the caller as it arrives. A caller that stops reading
|
||||
/// ends the upstream read, and the call completes with what it delivered.
|
||||
///
|
||||
/// A host on Anthropic SSE is relayed byte for byte. A host on another wire is decoded into
|
||||
/// Anthropic stream events and re-encoded as Anthropic SSE.
|
||||
async fn relay(
|
||||
host: &MessagesHost,
|
||||
response: reqwest::Response,
|
||||
decoder: Option<StreamDecoder>,
|
||||
) -> Result<MessagesOutput, Error> {
|
||||
let head = MessagesStreamHead {
|
||||
headers: response
|
||||
.headers()
|
||||
.iter()
|
||||
.filter_map(|(name, value)| Some((name.to_string(), value.to_str().ok()?.to_string())))
|
||||
.collect(),
|
||||
};
|
||||
if host.open(head).await? == Demand::Detached {
|
||||
return Ok(MessagesOutput::Streamed);
|
||||
}
|
||||
match decoder {
|
||||
None => relay_bytes(host, response).await,
|
||||
Some(decode) => relay_events(host, response, decode).await,
|
||||
}
|
||||
}
|
||||
|
||||
async fn relay_bytes(
|
||||
host: &MessagesHost,
|
||||
mut response: reqwest::Response,
|
||||
) -> Result<MessagesOutput, Error> {
|
||||
while let Some(chunk) = response.chunk().await.map_err(network)? {
|
||||
if host.deliver(chunk).await? == Demand::Detached {
|
||||
break;
|
||||
let request = prepare(call, secrets.as_ref()).await?;
|
||||
match execute(&http, &auth, request, &host).await? {
|
||||
MessagesResponse::Message(message) => Ok(MessagesOutput::Message(message)),
|
||||
MessagesResponse::Stream {
|
||||
headers,
|
||||
mut chunks,
|
||||
} => {
|
||||
if host.open(MessagesStreamHead { headers }).await? == Demand::Detached {
|
||||
return Ok(MessagesOutput::Streamed);
|
||||
}
|
||||
while let Some(chunk) = chunks.try_next().await? {
|
||||
if host.deliver(chunk).await? == Demand::Detached {
|
||||
break;
|
||||
}
|
||||
}
|
||||
Ok(MessagesOutput::Streamed)
|
||||
}
|
||||
}
|
||||
Ok(MessagesOutput::Streamed)
|
||||
}
|
||||
|
||||
async fn relay_events(
|
||||
host: &MessagesHost,
|
||||
response: reqwest::Response,
|
||||
decode: StreamDecoder,
|
||||
) -> Result<MessagesOutput, Error> {
|
||||
let bytes: ByteStream = futures_util::stream::unfold(response, |mut response| async move {
|
||||
match response.chunk().await {
|
||||
Ok(Some(chunk)) => Some((Ok(chunk), response)),
|
||||
Ok(None) => None,
|
||||
Err(error) => Some((Err(std::io::Error::other(error)), response)),
|
||||
}
|
||||
})
|
||||
.boxed();
|
||||
let mut events = decode(bytes);
|
||||
while let Some(event) = events.next().await {
|
||||
let chunk = encode_anthropic_sse(&event?)?;
|
||||
if host.deliver(chunk).await? == Demand::Detached {
|
||||
break;
|
||||
}
|
||||
}
|
||||
Ok(MessagesOutput::Streamed)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,12 +1,45 @@
|
|||
use std::time::Duration;
|
||||
|
||||
use litellm_llms::{
|
||||
anthropic::common_utils::AnthropicModelCapabilities, base_llm::auth::ValidatedEnvironment,
|
||||
use bytes::Bytes;
|
||||
use futures_util::stream::BoxStream;
|
||||
use litellm_llms::anthropic::common_utils::AnthropicModelCapabilities;
|
||||
use litellm_types::{
|
||||
llms::anthropic_messages::{
|
||||
anthropic_request::AnthropicMessagesRequest, anthropic_response::AnthropicMessagesResponse,
|
||||
},
|
||||
utils::ProviderSpecificHeaders,
|
||||
};
|
||||
use litellm_types::llms::anthropic_messages::anthropic_request::AnthropicMessagesRequest;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
use super::common_utils::MessagesProvider;
|
||||
use super::Error;
|
||||
|
||||
pub struct MessagesCall {
|
||||
pub body: AnthropicMessagesRequest,
|
||||
pub api_key: Option<String>,
|
||||
pub api_base: Option<String>,
|
||||
pub custom_llm_provider: Option<String>,
|
||||
pub extra_headers: Option<Map<String, Value>>,
|
||||
pub provider_specific_header: Option<ProviderSpecificHeaders>,
|
||||
pub timeout: Option<Duration>,
|
||||
pub shaping: MessagesShaping,
|
||||
}
|
||||
|
||||
pub fn messages_body(body: Map<String, Value>) -> Result<AnthropicMessagesRequest, Error> {
|
||||
serde_json::from_value(Value::Object(body)).map_err(invalid_request)
|
||||
}
|
||||
|
||||
pub(super) fn invalid_request(err: serde_json::Error) -> Error {
|
||||
Error::InvalidRequest(format!("invalid Anthropic messages request: {err}"))
|
||||
}
|
||||
|
||||
pub enum MessagesResponse {
|
||||
Message(Box<AnthropicMessagesResponse>),
|
||||
Stream {
|
||||
headers: Vec<(String, String)>,
|
||||
chunks: BoxStream<'static, Result<Bytes, Error>>,
|
||||
},
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
|
||||
pub struct MessagesShaping {
|
||||
|
|
@ -20,16 +53,6 @@ pub struct MessagesShaping {
|
|||
pub additional_drop_params: Vec<String>,
|
||||
}
|
||||
|
||||
pub(crate) struct ProviderMessagesRequest {
|
||||
pub(crate) provider: MessagesProvider,
|
||||
pub(crate) url: String,
|
||||
pub(crate) body: AnthropicMessagesRequest,
|
||||
/// The forwarded, default and feature headers plus how the call authenticates; the
|
||||
/// credential itself is applied when the request is sent.
|
||||
pub(crate) environment: ValidatedEnvironment,
|
||||
pub(crate) timeout: Option<Duration>,
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use litellm_llms::anthropic::common_utils::SupportedEffortTiers;
|
||||
|
|
|
|||
|
|
@ -1,9 +1,8 @@
|
|||
use std::{sync::Arc, time::Duration};
|
||||
|
||||
use litellm_core::messages::{
|
||||
Error,
|
||||
route::{LocalMessagesHost, MessagesCall, MessagesMachine, MessagesOutput, messages_machine},
|
||||
types::MessagesShaping,
|
||||
Error, MessagesCall, MessagesShaping,
|
||||
route::{LocalMessagesHost, MessagesMachine, MessagesOutput, messages_machine},
|
||||
};
|
||||
use litellm_http::{HttpSettings, Resolution};
|
||||
use litellm_secrets::source::SecretSource;
|
||||
|
|
|
|||
|
|
@ -1,7 +1,5 @@
|
|||
use litellm_llms::anthropic::common_utils::{
|
||||
ANTHROPIC_ADVISOR_TOOL_TYPE, ANTHROPIC_OAUTH_BETA_HEADER, AnthropicModelCapabilities,
|
||||
SupportedEffortTiers, beta,
|
||||
};
|
||||
use litellm_llms::anthropic::common_utils::{AnthropicModelCapabilities, SupportedEffortTiers};
|
||||
use litellm_types::llms::anthropic::{AnthropicBeta, BetaSet};
|
||||
use litellm_types::utils::{ProviderSpecificHeader, ProviderSpecificHeaders};
|
||||
use rstest::rstest;
|
||||
|
||||
|
|
@ -247,41 +245,37 @@ async fn additional_drop_params_remove_fields_before_sending(call: MessagesCall)
|
|||
assert_eq!(sent["top_k"], 3);
|
||||
}
|
||||
|
||||
fn sent_betas(request: &wiremock::Request) -> Vec<String> {
|
||||
fn sent_betas(request: &wiremock::Request) -> BetaSet {
|
||||
let [header] = <[&str; 1]>::try_from(request.header_values("anthropic-beta"))
|
||||
.unwrap_or_else(|values| panic!("expected one anthropic-beta header, got {values:?}"));
|
||||
header
|
||||
.split(',')
|
||||
.map(str::trim)
|
||||
.map(str::to_string)
|
||||
.collect()
|
||||
header.parse().unwrap()
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::structured_output(json!({"output_format": {"type": "json_schema"}}), &[beta::STRUCTURED_OUTPUT])]
|
||||
#[case::fast_mode(json!({"speed": "fast"}), &[beta::FAST_MODE_2026_02_01])]
|
||||
#[case::compaction(json!({"compaction": {"enabled": true}}), &[beta::COMPACT_2026_09_04])]
|
||||
#[case::structured_output(json!({"output_format": {"type": "json_schema"}}), &[AnthropicBeta::StructuredOutputs20251113])]
|
||||
#[case::fast_mode(json!({"speed": "fast"}), &[AnthropicBeta::FastMode20260201])]
|
||||
#[case::compaction(json!({"compaction": {"enabled": true}}), &[AnthropicBeta::Compact20260904])]
|
||||
#[case::context_management_edits(
|
||||
json!({"context_management": {"edits": [{"type": "clear_tool_uses_20250919"}]}}),
|
||||
&[beta::CONTEXT_MANAGEMENT_2025_06_27]
|
||||
&[AnthropicBeta::ContextManagement20250627]
|
||||
)]
|
||||
#[case::per_message_output_config(
|
||||
json!({"messages": [{"role": "user", "content": "hi", "output_config": {"effort": "low"}}]}),
|
||||
&[beta::PER_TURN_CONTROL_2026_07_01]
|
||||
&[AnthropicBeta::PerTurnControl20260701]
|
||||
)]
|
||||
#[case::advisor_tool(
|
||||
json!({"tools": [{"type": ANTHROPIC_ADVISOR_TOOL_TYPE, "name": "advisor", "model": MODEL}]}),
|
||||
&[beta::ADVISOR_TOOL_2026_03_01]
|
||||
json!({"tools": [{"type": "advisor_20260301", "name": "advisor", "model": MODEL}]}),
|
||||
&[AnthropicBeta::AdvisorTool20260301]
|
||||
)]
|
||||
#[case::several_features_at_once(
|
||||
json!({"speed": "fast", "output_format": {"type": "json_schema"}}),
|
||||
&[beta::STRUCTURED_OUTPUT, beta::FAST_MODE_2026_02_01]
|
||||
&[AnthropicBeta::StructuredOutputs20251113, AnthropicBeta::FastMode20260201]
|
||||
)]
|
||||
#[tokio::test]
|
||||
async fn feature_betas_join_the_callers_betas_in_one_sorted_header(
|
||||
call: MessagesCall,
|
||||
#[case] fields: Value,
|
||||
#[case] features: &[&str],
|
||||
#[case] features: &[AnthropicBeta],
|
||||
) {
|
||||
let upstream = upstream([message_response()]).await;
|
||||
let capabilities = AnthropicModelCapabilities {
|
||||
|
|
@ -305,12 +299,11 @@ async fn feature_betas_join_the_callers_betas_in_one_sorted_header(
|
|||
.await;
|
||||
|
||||
let sent = sent_betas(&only_request(&upstream).await);
|
||||
let mut expected: Vec<String> = features
|
||||
let expected: BetaSet = features
|
||||
.iter()
|
||||
.map(|feature| feature.to_string())
|
||||
.chain(["caller-beta-2025-01-01".to_string()])
|
||||
.cloned()
|
||||
.chain([AnthropicBeta::Other("caller-beta-2025-01-01".to_string())])
|
||||
.collect();
|
||||
expected.sort();
|
||||
assert_eq!(sent, expected);
|
||||
}
|
||||
|
||||
|
|
@ -331,7 +324,10 @@ async fn an_oauth_key_sends_the_browser_access_header_and_the_oauth_beta(call: M
|
|||
request.header("anthropic-dangerous-direct-browser-access"),
|
||||
Some("true")
|
||||
);
|
||||
assert_eq!(sent_betas(&request), [ANTHROPIC_OAUTH_BETA_HEADER]);
|
||||
assert_eq!(
|
||||
sent_betas(&request),
|
||||
BetaSet::from_iter([AnthropicBeta::Oauth20250420])
|
||||
);
|
||||
assert_eq!(request.header("x-api-key"), None);
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
use litellm_core::{
|
||||
Phase,
|
||||
messages::{messages, route::messages_body},
|
||||
messages::{MessagesResponse, messages, messages_body},
|
||||
};
|
||||
use litellm_http::transport::Error as TransportError;
|
||||
use rstest::rstest;
|
||||
|
|
@ -188,9 +188,10 @@ async fn the_facade_sends_through_the_injected_http_pool_configuration(call: Mes
|
|||
..HttpSettings::default()
|
||||
};
|
||||
|
||||
let message = messages(
|
||||
let response = messages(
|
||||
&support::resources(),
|
||||
&Resolution::from(&settings).config,
|
||||
&RecordingSecrets::empty(),
|
||||
MessagesCall {
|
||||
api_key: Some("sk-ant".into()),
|
||||
api_base: Some(base),
|
||||
|
|
@ -200,6 +201,9 @@ async fn the_facade_sends_through_the_injected_http_pool_configuration(call: Mes
|
|||
.await
|
||||
.expect("messages request succeeds");
|
||||
|
||||
let MessagesResponse::Message(message) = response else {
|
||||
panic!("a non-streaming request returns a message");
|
||||
};
|
||||
assert_eq!(message.id, "msg_1");
|
||||
let sent = only_request(&upstream).await;
|
||||
assert_eq!(sent.header("x-api-key"), Some("sk-ant"));
|
||||
|
|
|
|||
|
|
@ -1,12 +1,21 @@
|
|||
use std::{convert::Infallible, sync::Mutex};
|
||||
use std::{
|
||||
convert::Infallible,
|
||||
sync::{Mutex, mpsc},
|
||||
};
|
||||
|
||||
use bytes::Bytes;
|
||||
use litellm_core::messages::route::{Messages, MessagesStreamHead};
|
||||
use futures_util::{StreamExt, TryStreamExt};
|
||||
use litellm_core::messages::{
|
||||
MessagesResponse, messages,
|
||||
route::{Messages, MessagesStreamHead},
|
||||
};
|
||||
use litellm_host::host::{Demand, Host};
|
||||
use litellm_tracing::{Logger, Metadata, Record, Sink};
|
||||
use rstest::rstest;
|
||||
use tokio::{
|
||||
io::{AsyncReadExt, AsyncWriteExt},
|
||||
net::TcpListener,
|
||||
task::JoinHandle,
|
||||
};
|
||||
|
||||
use super::*;
|
||||
|
|
@ -23,6 +32,20 @@ enum Seen {
|
|||
Deliver(Bytes),
|
||||
}
|
||||
|
||||
struct TraceSink(mpsc::Sender<(String, Value)>);
|
||||
|
||||
impl Sink for TraceSink {
|
||||
fn enabled(&self, metadata: &Metadata<'_>) -> bool {
|
||||
metadata.target().starts_with("litellm_core::messages")
|
||||
}
|
||||
|
||||
fn emit(&self, record: &Record) {
|
||||
self.0
|
||||
.send((record.message.clone(), Value::Object(record.fields.clone())))
|
||||
.unwrap();
|
||||
}
|
||||
}
|
||||
|
||||
/// Projects like `LocalMessagesHost`, records every stream op in the order the route
|
||||
/// performs it, and detaches after `detach_after` ops.
|
||||
struct RecordingStreamHost {
|
||||
|
|
@ -120,6 +143,37 @@ async fn upstream_headers_are_on_the_stream_head_before_the_first_chunk(call: Me
|
|||
assert_eq!(delivered, SSE_BODY.as_bytes());
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn debug_trace_keeps_provider_input_and_every_stream_chunk(call: MessagesCall) {
|
||||
let upstream = upstream([sse_response()]).await;
|
||||
let host = RecordingStreamHost::new(streaming(call, upstream.uri()), usize::MAX);
|
||||
let (sender, receiver) = mpsc::channel();
|
||||
|
||||
Logger::new(TraceSink(sender))
|
||||
.instrument(stream_through(&host))
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let records: Vec<(String, Value)> = receiver.try_iter().collect();
|
||||
let request = records
|
||||
.iter()
|
||||
.find(|(message, _)| message == "provider request")
|
||||
.unwrap();
|
||||
let body: Value = serde_json::from_str(request.1["body"].as_str().unwrap()).unwrap();
|
||||
assert_eq!(body["messages"][0]["content"], "hi");
|
||||
assert_eq!(request.1["stream"], true);
|
||||
let chunks: String = records
|
||||
.iter()
|
||||
.filter(|(message, fields)| {
|
||||
message == "stream chunk" && fields["stage"] == "provider_response"
|
||||
})
|
||||
.map(|(_, fields)| fields["chunk"].as_str().unwrap())
|
||||
.collect();
|
||||
assert_eq!(chunks, SSE_BODY);
|
||||
assert!(!format!("{records:?}").contains("sk-ant"));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::at_open(1)]
|
||||
#[case::after_the_first_chunk(2)]
|
||||
|
|
@ -192,10 +246,10 @@ async fn a_stream_that_ends_without_message_stop_is_relayed_as_is(call: Messages
|
|||
}
|
||||
|
||||
/// Serves one SSE chunk and then holds the connection open without ever finishing.
|
||||
async fn stalling_upstream() -> String {
|
||||
async fn stalling_upstream() -> (String, JoinHandle<()>) {
|
||||
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
|
||||
let base = format!("http://{}", listener.local_addr().unwrap());
|
||||
tokio::spawn(async move {
|
||||
let connection = tokio::spawn(async move {
|
||||
let (mut socket, _) = listener.accept().await.unwrap();
|
||||
let mut request = vec![0; 4096];
|
||||
let _ = socket.read(&mut request).await;
|
||||
|
|
@ -206,15 +260,15 @@ async fn stalling_upstream() -> String {
|
|||
)
|
||||
.await
|
||||
.unwrap();
|
||||
std::future::pending::<()>().await;
|
||||
let _ = socket.read_to_end(&mut Vec::new()).await;
|
||||
});
|
||||
base
|
||||
(base, connection)
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn the_timeout_covers_a_stalled_stream_body(call: MessagesCall) {
|
||||
let base = stalling_upstream().await;
|
||||
let (base, connection) = stalling_upstream().await;
|
||||
let host = RecordingStreamHost::new(
|
||||
MessagesCall {
|
||||
timeout: Some(Duration::from_millis(300)),
|
||||
|
|
@ -236,6 +290,145 @@ async fn the_timeout_covers_a_stalled_stream_body(call: MessagesCall) {
|
|||
"the chunk before the stall reached the caller, saw {} ops",
|
||||
seen.len()
|
||||
);
|
||||
tokio::time::timeout(Duration::from_secs(5), connection)
|
||||
.await
|
||||
.expect("timing out closes the upstream connection")
|
||||
.unwrap();
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::anthropic("anthropic")]
|
||||
#[case::azure_ai("azure_ai")]
|
||||
#[tokio::test]
|
||||
async fn the_sdk_returns_stream_headers_and_every_sse_byte(
|
||||
call: MessagesCall,
|
||||
#[case] provider: &str,
|
||||
) {
|
||||
let upstream = upstream([sse_response()]).await;
|
||||
let response = messages(
|
||||
&support::resources(),
|
||||
&http_config(),
|
||||
&RecordingSecrets::empty(),
|
||||
MessagesCall {
|
||||
custom_llm_provider: Some(provider.into()),
|
||||
..streaming(call, upstream.uri())
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let MessagesResponse::Stream { headers, chunks } = response else {
|
||||
panic!("a streaming request returns a stream");
|
||||
};
|
||||
for (name, value) in UPSTREAM_HEADERS {
|
||||
assert!(headers.contains(&(name.into(), value.into())));
|
||||
}
|
||||
let delivered = chunks.try_collect::<Vec<_>>().await.unwrap().concat();
|
||||
assert_eq!(delivered, SSE_BODY.as_bytes());
|
||||
assert_eq!(only_request(&upstream).await.json()["stream"], true);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn the_sdk_returns_http_errors_before_opening_a_stream(call: MessagesCall) {
|
||||
let upstream = upstream([ResponseTemplate::new(429).set_body_string("slow down")]).await;
|
||||
let error = messages(
|
||||
&support::resources(),
|
||||
&http_config(),
|
||||
&RecordingSecrets::empty(),
|
||||
streaming(call, upstream.uri()),
|
||||
)
|
||||
.await
|
||||
.err()
|
||||
.expect("upstream failure is returned by messages()");
|
||||
|
||||
assert_eq!(
|
||||
error,
|
||||
Error::Transport(litellm_http::transport::Error::Http {
|
||||
status: 429,
|
||||
body: "slow down".into(),
|
||||
})
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::before_reading(false)]
|
||||
#[case::after_reading(true)]
|
||||
#[tokio::test]
|
||||
async fn dropping_the_sdk_stream_closes_the_unfinished_upstream(
|
||||
call: MessagesCall,
|
||||
#[case] read_chunk: bool,
|
||||
) {
|
||||
let (base, connection) = stalling_upstream().await;
|
||||
let response = tokio::time::timeout(
|
||||
Duration::from_secs(5),
|
||||
messages(
|
||||
&support::resources(),
|
||||
&http_config(),
|
||||
&RecordingSecrets::empty(),
|
||||
MessagesCall {
|
||||
timeout: Some(Duration::from_secs(30)),
|
||||
..streaming(call, base)
|
||||
},
|
||||
),
|
||||
)
|
||||
.await
|
||||
.expect("messages() returns before the upstream finishes")
|
||||
.unwrap();
|
||||
|
||||
let MessagesResponse::Stream { mut chunks, .. } = response else {
|
||||
panic!("a streaming request returns a stream");
|
||||
};
|
||||
if read_chunk {
|
||||
let chunk = tokio::time::timeout(Duration::from_secs(5), chunks.next())
|
||||
.await
|
||||
.expect("the first chunk arrives before the upstream finishes")
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
assert_eq!(chunk.as_ref(), b"event: message_start\ndata: {}\n\n");
|
||||
}
|
||||
assert!(!connection.is_finished());
|
||||
drop(chunks);
|
||||
tokio::time::timeout(Duration::from_secs(5), connection)
|
||||
.await
|
||||
.expect("dropping the stream closes the upstream connection")
|
||||
.unwrap();
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn the_sdk_yields_a_body_error_once_after_delivered_chunks(call: MessagesCall) {
|
||||
let (base, connection) = stalling_upstream().await;
|
||||
let response = messages(
|
||||
&support::resources(),
|
||||
&http_config(),
|
||||
&RecordingSecrets::empty(),
|
||||
MessagesCall {
|
||||
timeout: Some(Duration::from_millis(300)),
|
||||
..streaming(call, base)
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let MessagesResponse::Stream { mut chunks, .. } = response else {
|
||||
panic!("a streaming request returns a stream");
|
||||
};
|
||||
assert_eq!(
|
||||
chunks.next().await.unwrap().unwrap().as_ref(),
|
||||
b"event: message_start\ndata: {}\n\n"
|
||||
);
|
||||
let error = tokio::time::timeout(Duration::from_secs(5), chunks.next())
|
||||
.await
|
||||
.expect("the stalled body times out")
|
||||
.unwrap()
|
||||
.unwrap_err();
|
||||
assert!(matches!(error, Error::Transport(_)), "{error:?}");
|
||||
assert!(chunks.next().await.is_none());
|
||||
tokio::time::timeout(Duration::from_secs(5), connection)
|
||||
.await
|
||||
.expect("the failed stream closes its upstream connection")
|
||||
.unwrap();
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
|
|
|
|||
|
|
@ -12,7 +12,6 @@ bytes.workspace = true
|
|||
futures-util.workspace = true
|
||||
litellm-auth.workspace = true
|
||||
litellm-core.workspace = true
|
||||
litellm-host.workspace = true
|
||||
litellm-http.workspace = true
|
||||
litellm-llms.workspace = true
|
||||
litellm-router.workspace = true
|
||||
|
|
@ -20,10 +19,10 @@ litellm-secrets.workspace = true
|
|||
litellm-types.workspace = true
|
||||
serde_json.workspace = true
|
||||
thiserror.workspace = true
|
||||
tokio = { workspace = true, features = ["sync"] }
|
||||
|
||||
[dev-dependencies]
|
||||
futures-util.workspace = true
|
||||
tokio = { workspace = true, features = ["io-util"] }
|
||||
rstest.workspace = true
|
||||
tower = { version = "0.5.3", features = ["util"] }
|
||||
wiremock = "0.6.5"
|
||||
|
|
|
|||
|
|
@ -1,69 +0,0 @@
|
|||
use std::{convert::Infallible, sync::Mutex};
|
||||
|
||||
use bytes::Bytes;
|
||||
use litellm_core::messages::{
|
||||
Error,
|
||||
route::{LocalMessagesHost, Messages, MessagesCall, MessagesStreamHead},
|
||||
};
|
||||
use litellm_host::host::{Demand, Host};
|
||||
use tokio::sync::{mpsc, oneshot};
|
||||
|
||||
/// Hands a streamed response to the HTTP body: the head once, then each chunk. A dropped
|
||||
/// receiver means the client went away, which detaches the call.
|
||||
pub(super) struct ChannelHost {
|
||||
local: LocalMessagesHost,
|
||||
head: Mutex<Option<oneshot::Sender<MessagesStreamHead>>>,
|
||||
pub(super) chunks: mpsc::Sender<Bytes>,
|
||||
}
|
||||
|
||||
impl ChannelHost {
|
||||
pub(super) fn new(
|
||||
call: MessagesCall,
|
||||
head: oneshot::Sender<MessagesStreamHead>,
|
||||
chunks: mpsc::Sender<Bytes>,
|
||||
) -> Self {
|
||||
Self {
|
||||
local: LocalMessagesHost::new(call),
|
||||
head: Mutex::new(Some(head)),
|
||||
chunks,
|
||||
}
|
||||
}
|
||||
|
||||
fn take_head(&self) -> Option<oneshot::Sender<MessagesStreamHead>> {
|
||||
self.head
|
||||
.lock()
|
||||
.unwrap_or_else(|error| error.into_inner())
|
||||
.take()
|
||||
}
|
||||
|
||||
pub(super) fn opened(&self) -> bool {
|
||||
self.head
|
||||
.lock()
|
||||
.unwrap_or_else(|error| error.into_inner())
|
||||
.is_none()
|
||||
}
|
||||
}
|
||||
|
||||
impl Host<Messages> for ChannelHost {
|
||||
async fn project(&self) -> Result<MessagesCall, Error> {
|
||||
self.local.project().await
|
||||
}
|
||||
|
||||
async fn custom_op(&self, op: Infallible) -> Result<(), Error> {
|
||||
match op {}
|
||||
}
|
||||
|
||||
async fn open(&self, head: MessagesStreamHead) -> Result<Demand, Error> {
|
||||
Ok(match self.take_head().map(|sender| sender.send(head)) {
|
||||
Some(Ok(())) => Demand::More,
|
||||
Some(Err(_)) | None => Demand::Detached,
|
||||
})
|
||||
}
|
||||
|
||||
async fn deliver(&self, chunk: Bytes) -> Result<Demand, Error> {
|
||||
Ok(match self.chunks.send(chunk).await {
|
||||
Ok(()) => Demand::More,
|
||||
Err(_) => Demand::Detached,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
|
@ -1,7 +1,5 @@
|
|||
//! `POST /v1/messages`, as the Python proxy's `anthropic_response` serves it.
|
||||
|
||||
mod host;
|
||||
|
||||
use std::{convert::Infallible, sync::Arc};
|
||||
|
||||
use axum::{
|
||||
|
|
@ -11,13 +9,12 @@ use axum::{
|
|||
http::{HeaderMap, StatusCode, header},
|
||||
response::{IntoResponse, Response},
|
||||
};
|
||||
use host::ChannelHost;
|
||||
use litellm_core::messages::route::{
|
||||
MessagesCall, MessagesOutput, messages_body, messages_machine,
|
||||
use futures_util::{StreamExt, stream::BoxStream};
|
||||
use litellm_core::messages::{
|
||||
Error as RouteError, MessagesCall, MessagesResponse, messages, messages_body,
|
||||
};
|
||||
use litellm_types::utils::{ProviderSpecificHeader, ProviderSpecificHeaders};
|
||||
use serde_json::{Map, Value};
|
||||
use tokio::sync::{mpsc, oneshot};
|
||||
|
||||
use crate::{Deployment, Error, Gateway};
|
||||
|
||||
|
|
@ -55,31 +52,16 @@ async fn handle(gateway: &Gateway, headers: &HeaderMap, body: &[u8]) -> Result<R
|
|||
.get(model_name)
|
||||
.ok_or_else(|| Error::UnknownModel(model_name.to_owned()))?;
|
||||
let call = project(deployment, body, headers)?;
|
||||
let machine = messages_machine(&gateway.resources, &gateway.http, gateway.secrets.clone())
|
||||
.map_err(|error| Error::Route(error.into()))?;
|
||||
|
||||
let (head_sender, head) = oneshot::channel();
|
||||
let (chunk_sender, chunks) = mpsc::channel(1);
|
||||
let host = ChannelHost::new(call, head_sender, chunk_sender);
|
||||
let call = tokio::spawn(async move {
|
||||
let outcome = litellm_host::run::run(machine, &host).await;
|
||||
if let Err(error) = &outcome
|
||||
&& host.opened()
|
||||
{
|
||||
let _ = host
|
||||
.chunks
|
||||
.send(Bytes::from(Error::Route(error.clone()).sse_frame()))
|
||||
.await;
|
||||
}
|
||||
outcome
|
||||
});
|
||||
tokio::select! {
|
||||
biased;
|
||||
Ok(_) = head => Ok(stream(chunks)),
|
||||
joined = call => match joined.map_err(|error| Error::Internal(error.to_string()))?? {
|
||||
MessagesOutput::Message(message) => Ok(Json(message).into_response()),
|
||||
MessagesOutput::Streamed => Err(Error::Internal("the stream ended before it opened".into())),
|
||||
},
|
||||
match messages(
|
||||
&gateway.resources,
|
||||
&gateway.http,
|
||||
gateway.secrets.as_ref(),
|
||||
call,
|
||||
)
|
||||
.await?
|
||||
{
|
||||
MessagesResponse::Message(message) => Ok(Json(message).into_response()),
|
||||
MessagesResponse::Stream { chunks, .. } => Ok(stream(chunks)),
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -123,10 +105,13 @@ fn anthropic_api_headers(headers: &HeaderMap) -> Option<ProviderSpecificHeaders>
|
|||
})
|
||||
}
|
||||
|
||||
fn stream(chunks: mpsc::Receiver<Bytes>) -> Response {
|
||||
let body = futures_util::stream::unfold(chunks, |mut chunks| async move {
|
||||
let chunk = chunks.recv().await?;
|
||||
Some((Ok::<_, Infallible>(chunk), chunks))
|
||||
/// A chunk that fails after the stream opened is delivered as an SSE error frame, since
|
||||
/// the status line already went out; the stream ends on it.
|
||||
fn stream(chunks: BoxStream<'static, Result<Bytes, RouteError>>) -> Response {
|
||||
let body = chunks.map(|chunk| {
|
||||
Ok::<_, Infallible>(
|
||||
chunk.unwrap_or_else(|error| Bytes::from(Error::Route(error).sse_frame())),
|
||||
)
|
||||
});
|
||||
(
|
||||
StatusCode::OK,
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ use axum::{
|
|||
};
|
||||
use rstest::rstest;
|
||||
use serde_json::json;
|
||||
use tokio::io::{AsyncReadExt, AsyncWriteExt};
|
||||
use tower::ServiceExt;
|
||||
use wiremock::{
|
||||
Mock, MockServer, ResponseTemplate,
|
||||
|
|
@ -64,3 +65,57 @@ async fn invalid_messages_stays_an_anthropic_error() {
|
|||
assert_eq!(body["type"], "error");
|
||||
assert_eq!(body["error"]["type"], "invalid_request_error");
|
||||
}
|
||||
|
||||
/// Answers with the SSE head and one event, then drops the connection short of the
|
||||
/// announced body length.
|
||||
async fn truncating_upstream() -> String {
|
||||
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
|
||||
let base = format!("http://{}", listener.local_addr().unwrap());
|
||||
tokio::spawn(async move {
|
||||
let (mut socket, _) = listener.accept().await.unwrap();
|
||||
let mut request = vec![0; 4096];
|
||||
let _ = socket.read(&mut request).await;
|
||||
socket
|
||||
.write_all(
|
||||
format!(
|
||||
"HTTP/1.1 200 OK\r\ncontent-type: text/event-stream\r\ncontent-length: {}\r\n\r\n{FIRST_EVENT}",
|
||||
FIRST_EVENT.len() * 2
|
||||
)
|
||||
.as_bytes(),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
});
|
||||
base
|
||||
}
|
||||
|
||||
const FIRST_EVENT: &str = "event: message_start\ndata: {}\n\n";
|
||||
|
||||
#[tokio::test]
|
||||
async fn a_stream_that_fails_after_opening_ends_with_an_sse_error_frame() {
|
||||
let base = truncating_upstream().await;
|
||||
let request = Request::post("/v1/messages")
|
||||
.header("content-type", "application/json")
|
||||
.body(Body::from(
|
||||
json!({"model": "public/model", "messages": [{"role": "user", "content": "hi"}],
|
||||
"max_tokens": 16, "stream": true})
|
||||
.to_string(),
|
||||
))
|
||||
.unwrap();
|
||||
|
||||
let response = support::app("anthropic/test-model", &base)
|
||||
.oneshot(request)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(response.status(), 200);
|
||||
let body = to_bytes(response.into_body(), 4096).await.unwrap();
|
||||
let text = std::str::from_utf8(&body).unwrap();
|
||||
let frame = text
|
||||
.strip_prefix(FIRST_EVENT)
|
||||
.and_then(|rest| rest.strip_prefix("event: error\ndata: "))
|
||||
.unwrap_or_else(|| panic!("the delivered event then one error frame, got {text:?}"));
|
||||
let error: serde_json::Value = serde_json::from_str(frame.trim_end()).unwrap();
|
||||
assert_eq!(error["type"], "error");
|
||||
assert_eq!(error["error"]["type"], "api_error");
|
||||
}
|
||||
|
|
|
|||
|
|
@ -7,6 +7,7 @@ repository.workspace = true
|
|||
|
||||
[dependencies]
|
||||
axum.workspace = true
|
||||
http-body-util = "0.1"
|
||||
litellm-core.workspace = true
|
||||
litellm-gateway-inference.workspace = true
|
||||
litellm-gateway-auth.workspace = true
|
||||
|
|
@ -14,11 +15,14 @@ litellm-config.workspace = true
|
|||
litellm-http.workspace = true
|
||||
litellm-llms.workspace = true
|
||||
litellm-secrets.workspace = true
|
||||
tower-http = { version = "0.7.1", default-features = false, features = ["trace"] }
|
||||
litellm-tracing.workspace = true
|
||||
serde_json.workspace = true
|
||||
tracing.workspace = true
|
||||
tokio.workspace = true
|
||||
uuid.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
futures-util.workspace = true
|
||||
rstest.workspace = true
|
||||
serde_json.workspace = true
|
||||
tokio = { workspace = true, features = ["sync"] }
|
||||
tower = { version = "0.5", features = ["util"] }
|
||||
|
|
|
|||
|
|
@ -1,7 +1,13 @@
|
|||
use std::sync::Arc;
|
||||
use std::{sync::Arc, time::Instant};
|
||||
|
||||
use axum::{Router, extract::Request};
|
||||
use tower_http::trace::{DefaultOnResponse, TraceLayer};
|
||||
use axum::{
|
||||
Router,
|
||||
body::{Body, Bytes},
|
||||
extract::Request,
|
||||
middleware::Next,
|
||||
response::Response,
|
||||
};
|
||||
use http_body_util::BodyExt;
|
||||
|
||||
use litellm_config::Config;
|
||||
use litellm_core::resources::CoreResources;
|
||||
|
|
@ -12,6 +18,8 @@ use litellm_http::{
|
|||
};
|
||||
use litellm_llms::base_llm::ocr::settings::OcrSettings;
|
||||
use litellm_secrets::source::EnvironmentSecrets;
|
||||
use litellm_tracing::ByteChunk;
|
||||
use uuid::Uuid;
|
||||
|
||||
pub fn build_inference(config: &Config) -> Result<Arc<Gateway>, litellm_http::Error> {
|
||||
let pool = Arc::new(HttpClientPool::new(Arc::new(PublicDnsResolver)));
|
||||
|
|
@ -42,11 +50,125 @@ pub fn router(inference: Arc<Gateway>, config: &Config) -> Router {
|
|||
RequireMasterKey,
|
||||
_,
|
||||
>(auth))
|
||||
.layer(
|
||||
TraceLayer::new_for_http()
|
||||
.make_span_with(|request: &Request| {
|
||||
tracing::info_span!("request", method = %request.method(), path = request.uri().path())
|
||||
})
|
||||
.on_response(DefaultOnResponse::new().level(tracing::Level::INFO)),
|
||||
)
|
||||
.layer(axum::middleware::from_fn(log_request))
|
||||
}
|
||||
|
||||
async fn log_request(request: Request, next: Next) -> Response {
|
||||
let request_id = Uuid::new_v4().to_string();
|
||||
let log_body_chunks = tracing::enabled!(tracing::Level::DEBUG);
|
||||
let method = request.method().clone();
|
||||
let path = request.uri().path().to_owned();
|
||||
let started = Instant::now();
|
||||
let request = if log_body_chunks {
|
||||
request.map(|body| logged_body(body, request_id.clone(), "input"))
|
||||
} else {
|
||||
request
|
||||
};
|
||||
let response = next.run(request).await;
|
||||
tracing::info!(
|
||||
%request_id,
|
||||
%method,
|
||||
%path,
|
||||
status = response.status().as_u16(),
|
||||
time_to_headers_ms = started.elapsed().as_secs_f64() * 1000.0,
|
||||
"response headers"
|
||||
);
|
||||
if log_body_chunks {
|
||||
response.map(|body| logged_body(body, request_id, "output"))
|
||||
} else {
|
||||
response
|
||||
}
|
||||
}
|
||||
|
||||
fn logged_body(body: Body, request_id: String, direction: &'static str) -> Body {
|
||||
Body::new(body.map_frame(move |frame| {
|
||||
if let Some(data) = frame.data_ref() {
|
||||
log_chunk(&request_id, direction, data);
|
||||
}
|
||||
frame
|
||||
}))
|
||||
}
|
||||
|
||||
fn log_chunk(request_id: &str, direction: &str, data: &Bytes) {
|
||||
let chunk = ByteChunk::new(data);
|
||||
tracing::debug!(request_id, direction, encoding = chunk.encoding(), chunk = %chunk, "body chunk");
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::{convert::Infallible, sync::mpsc};
|
||||
|
||||
use axum::{body::to_bytes, http::StatusCode, routing::post};
|
||||
use futures_util::stream;
|
||||
use litellm_tracing::{Logger, Metadata, Record, Sink};
|
||||
use rstest::rstest;
|
||||
use serde_json::{Value, json};
|
||||
use tower::ServiceExt;
|
||||
|
||||
use super::*;
|
||||
|
||||
struct LogSink(mpsc::Sender<Value>);
|
||||
|
||||
impl Sink for LogSink {
|
||||
fn enabled(&self, _: &Metadata<'_>) -> bool {
|
||||
true
|
||||
}
|
||||
|
||||
fn emit(&self, record: &Record) {
|
||||
self.0
|
||||
.send(json!({"message": record.message, "fields": record.fields}))
|
||||
.unwrap();
|
||||
}
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn logs_each_body_chunk_without_changing_streamed_bytes() {
|
||||
let app = Router::new()
|
||||
.route(
|
||||
"/stream",
|
||||
post(|_: Bytes| async {
|
||||
(
|
||||
StatusCode::OK,
|
||||
Body::from_stream(stream::iter([
|
||||
Ok::<_, Infallible>(Bytes::from_static(b"event: first\n\n")),
|
||||
Ok(Bytes::from_static(b"event: second\n\n")),
|
||||
])),
|
||||
)
|
||||
}),
|
||||
)
|
||||
.layer(axum::middleware::from_fn(log_request));
|
||||
let request_chunks = [
|
||||
Ok::<_, Infallible>(Bytes::from_static(b"hello")),
|
||||
Ok(Bytes::from_static(b" world")),
|
||||
];
|
||||
let request = Request::post("/stream")
|
||||
.body(Body::from_stream(stream::iter(request_chunks)))
|
||||
.unwrap();
|
||||
let (sender, receiver) = mpsc::channel();
|
||||
let logger = Logger::new(LogSink(sender));
|
||||
|
||||
let output = logger
|
||||
.instrument(async {
|
||||
let response = app.oneshot(request).await.unwrap();
|
||||
to_bytes(response.into_body(), 1024).await.unwrap()
|
||||
})
|
||||
.await;
|
||||
|
||||
assert_eq!(output, "event: first\n\nevent: second\n\n");
|
||||
let records: Vec<Value> = receiver.try_iter().collect();
|
||||
assert_eq!(records.len(), 5);
|
||||
assert_eq!(records[0]["fields"]["chunk"], "hello");
|
||||
assert_eq!(records[1]["fields"]["chunk"], " world");
|
||||
assert_eq!(records[2]["fields"]["status"], 200);
|
||||
assert_eq!(records[3]["fields"]["chunk"], "event: first\n\n");
|
||||
assert_eq!(records[4]["fields"]["chunk"], "event: second\n\n");
|
||||
let request_id = &records[2]["fields"]["request_id"];
|
||||
assert!(request_id.as_str().is_some());
|
||||
assert!(
|
||||
records
|
||||
.iter()
|
||||
.all(|record| &record["fields"]["request_id"] == request_id)
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,9 +1,46 @@
|
|||
use std::error::Error;
|
||||
use std::{
|
||||
error::Error,
|
||||
time::{SystemTime, UNIX_EPOCH},
|
||||
};
|
||||
|
||||
use litellm_config::Config;
|
||||
use litellm_tracing::{Level, Logger, Metadata, Record, Sink};
|
||||
use serde_json::json;
|
||||
|
||||
struct StderrSink {
|
||||
level: Level,
|
||||
}
|
||||
|
||||
impl Sink for StderrSink {
|
||||
fn enabled(&self, metadata: &Metadata<'_>) -> bool {
|
||||
*metadata.level() <= self.level && metadata.target().starts_with("litellm")
|
||||
}
|
||||
|
||||
fn emit(&self, record: &Record) {
|
||||
let timestamp_ms = SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.unwrap_or_default()
|
||||
.as_millis();
|
||||
eprintln!(
|
||||
"{}",
|
||||
json!({
|
||||
"timestamp_ms": timestamp_ms,
|
||||
"level": record.metadata.level().as_str(),
|
||||
"target": record.metadata.target(),
|
||||
"message": record.message,
|
||||
"fields": record.fields,
|
||||
})
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::main]
|
||||
async fn main() -> Result<(), Box<dyn Error>> {
|
||||
let level = std::env::var("RUST_LOG")
|
||||
.ok()
|
||||
.and_then(|value| value.parse().ok())
|
||||
.unwrap_or(Level::INFO);
|
||||
Logger::new(StderrSink { level }).install_global()?;
|
||||
let config_path = std::env::var("LITELLM_CONFIG").unwrap_or_else(|_| "config.yaml".into());
|
||||
let config = Config::load(config_path)?;
|
||||
let inference = litellm_gateway::build_inference(&config)?;
|
||||
|
|
@ -13,6 +50,8 @@ async fn main() -> Result<(), Box<dyn Error>> {
|
|||
.parse::<u16>()?;
|
||||
let listener = tokio::net::TcpListener::bind((host.as_str(), port)).await?;
|
||||
|
||||
tracing::info!(address = %listener.local_addr()?, models = config.model_list.len(), log_level = %level, "gateway listening");
|
||||
|
||||
axum::serve(listener, litellm_gateway::router(inference, &config)).await?;
|
||||
Ok(())
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,11 +1,31 @@
|
|||
use std::{sync::Arc, time::Duration};
|
||||
use std::{
|
||||
sync::{Arc, mpsc},
|
||||
time::Duration,
|
||||
};
|
||||
|
||||
use axum::{body::Body, http::Request};
|
||||
use litellm_config::Config;
|
||||
use litellm_gateway_inference::{Error, Gateway};
|
||||
use litellm_http::ClientVariant;
|
||||
use litellm_tracing::{Logger, Metadata, Record, Sink};
|
||||
use rstest::{fixture, rstest};
|
||||
use serde_json::{Value, json};
|
||||
use tokio::{net::TcpListener, sync::oneshot, time::timeout};
|
||||
use tower::ServiceExt;
|
||||
|
||||
struct LogSink(mpsc::Sender<Value>);
|
||||
|
||||
impl Sink for LogSink {
|
||||
fn enabled(&self, _: &Metadata<'_>) -> bool {
|
||||
true
|
||||
}
|
||||
|
||||
fn emit(&self, record: &Record) {
|
||||
self.0
|
||||
.send(json!({"message": record.message, "fields": record.fields}))
|
||||
.unwrap();
|
||||
}
|
||||
}
|
||||
|
||||
#[fixture]
|
||||
fn inference() -> Arc<Gateway> {
|
||||
|
|
@ -88,3 +108,36 @@ async fn authenticates_before_serving_mounted_inference_routes(
|
|||
.unwrap()
|
||||
.unwrap();
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn logs_request_outcome_without_credentials_or_query(inference: Arc<Gateway>) {
|
||||
let config =
|
||||
Config::from_yaml("model_list: []\ngeneral_settings:\n master_key: gateway-key\n")
|
||||
.unwrap();
|
||||
let request = Request::builder()
|
||||
.method("POST")
|
||||
.uri("/v1/messages?token=query-secret")
|
||||
.header("authorization", "Bearer header-secret")
|
||||
.body(Body::empty())
|
||||
.unwrap();
|
||||
let (sender, receiver) = mpsc::channel();
|
||||
let logger = Logger::new(LogSink(sender));
|
||||
|
||||
let response = logger
|
||||
.instrument(litellm_gateway::router(inference, &config).oneshot(request))
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(response.status().as_u16(), 401);
|
||||
let record = receiver.try_recv().unwrap();
|
||||
assert_eq!(record["message"], "response headers");
|
||||
assert_eq!(record["fields"]["method"], "POST");
|
||||
assert_eq!(record["fields"]["path"], "/v1/messages");
|
||||
assert_eq!(record["fields"]["status"], 401);
|
||||
assert!(record["fields"]["time_to_headers_ms"].as_f64().unwrap() >= 0.0);
|
||||
assert!(record["fields"]["request_id"].as_str().is_some());
|
||||
assert!(receiver.try_recv().is_err());
|
||||
assert!(!record.to_string().contains("header-secret"));
|
||||
assert!(!record.to_string().contains("query-secret"));
|
||||
}
|
||||
|
|
|
|||
145
litellm-rust/crates/host/src/hooks.rs
Normal file
145
litellm-rust/crates/host/src/hooks.rs
Normal file
|
|
@ -0,0 +1,145 @@
|
|||
use std::future::Future;
|
||||
|
||||
use crate::{
|
||||
event::{MachineEvent, RequestContext, WireRequest},
|
||||
machine::{HostChannel, MachineFault},
|
||||
protocol::Protocol,
|
||||
};
|
||||
|
||||
/// What a route reaches for mid-call: the send-time rewrite and the events it reports.
|
||||
/// Python's `logging_obj.pre_call` and `post_call`, in that order.
|
||||
pub trait RouteHooks<E>: Send + Sync {
|
||||
fn before_send(
|
||||
&self,
|
||||
wire: WireRequest,
|
||||
context: RequestContext,
|
||||
) -> impl Future<Output = Result<WireRequest, E>> + Send;
|
||||
|
||||
fn emit(&self, event: MachineEvent) -> impl Future<Output = Result<(), E>> + Send;
|
||||
}
|
||||
|
||||
/// No host: the wire request goes out as prepared and nothing observes the call.
|
||||
impl<E> RouteHooks<E> for () {
|
||||
async fn before_send(&self, wire: WireRequest, _: RequestContext) -> Result<WireRequest, E> {
|
||||
Ok(wire)
|
||||
}
|
||||
|
||||
async fn emit(&self, _: MachineEvent) -> Result<(), E> {
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
impl<R: Protocol> RouteHooks<R::Error> for HostChannel<R>
|
||||
where
|
||||
R::Error: From<MachineFault>,
|
||||
{
|
||||
async fn before_send(
|
||||
&self,
|
||||
wire: WireRequest,
|
||||
context: RequestContext,
|
||||
) -> Result<WireRequest, R::Error> {
|
||||
HostChannel::before_send(self, wire, context).await
|
||||
}
|
||||
|
||||
async fn emit(&self, event: MachineEvent) -> Result<(), R::Error> {
|
||||
HostChannel::emit(self, event).await
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::convert::Infallible;
|
||||
|
||||
use serde_json::json;
|
||||
|
||||
use super::*;
|
||||
use crate::{
|
||||
event::RawResponse,
|
||||
host::HostOp,
|
||||
machine::{CallMachine, Machine, MachineStep},
|
||||
};
|
||||
|
||||
struct Unit;
|
||||
|
||||
#[derive(Clone, Debug)]
|
||||
struct Fault;
|
||||
|
||||
impl Protocol for Unit {
|
||||
type Response = (WireRequest, ());
|
||||
type Error = Fault;
|
||||
type Projection = ();
|
||||
type Op = Infallible;
|
||||
type Chunk = Infallible;
|
||||
type StreamHead = Infallible;
|
||||
}
|
||||
|
||||
impl From<MachineFault> for Fault {
|
||||
fn from(_: MachineFault) -> Self {
|
||||
Fault
|
||||
}
|
||||
}
|
||||
|
||||
fn wire(url: &str) -> WireRequest {
|
||||
WireRequest {
|
||||
url: url.into(),
|
||||
headers: Vec::new(),
|
||||
body: json!({}),
|
||||
}
|
||||
}
|
||||
|
||||
fn context() -> RequestContext {
|
||||
RequestContext {
|
||||
model: "m".into(),
|
||||
custom_llm_provider: "p".into(),
|
||||
optional_params: json!({}),
|
||||
secret_fields: Vec::new(),
|
||||
api_key: None,
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn the_channel_yields_each_hook_as_its_op_and_returns_the_answer() {
|
||||
let mut machine = CallMachine::<Unit>::new(|channel| {
|
||||
Box::pin(async move {
|
||||
let sent = RouteHooks::before_send(&channel, wire("prepared"), context()).await?;
|
||||
RouteHooks::emit(
|
||||
&channel,
|
||||
MachineEvent::ResponseReceived {
|
||||
raw: RawResponse { body: "raw".into() },
|
||||
},
|
||||
)
|
||||
.await?;
|
||||
Ok((sent, ()))
|
||||
})
|
||||
});
|
||||
|
||||
let Ok(MachineStep::Host(HostOp::BeforeSend { wire, reply, .. })) = machine.resume().await
|
||||
else {
|
||||
panic!("before_send yields BeforeSend");
|
||||
};
|
||||
assert_eq!(wire.url, "prepared");
|
||||
reply.send(WireRequest {
|
||||
url: "rewritten".into(),
|
||||
..*wire
|
||||
});
|
||||
|
||||
let Ok(MachineStep::Host(HostOp::Emit(event, reply))) = machine.resume().await else {
|
||||
panic!("emit yields Emit");
|
||||
};
|
||||
assert!(matches!(event, MachineEvent::ResponseReceived { .. }));
|
||||
reply.send(());
|
||||
|
||||
let Ok(MachineStep::Complete((sent, ()))) = machine.resume().await else {
|
||||
panic!("the call completes with the answers");
|
||||
};
|
||||
assert_eq!(sent.url, "rewritten");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn no_hooks_pass_the_wire_request_through() {
|
||||
let sent = RouteHooks::<Fault>::before_send(&(), wire("prepared"), context())
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(sent.url, "prepared");
|
||||
}
|
||||
}
|
||||
|
|
@ -7,6 +7,7 @@
|
|||
//! may rewrite the wire request before it is sent.
|
||||
|
||||
pub mod event;
|
||||
pub mod hooks;
|
||||
pub mod host;
|
||||
pub mod machine;
|
||||
pub mod protocol;
|
||||
|
|
|
|||
|
|
@ -87,6 +87,20 @@ pub fn has_header(headers: &[(String, String)], name: &str) -> bool {
|
|||
.any(|(key, _)| key.eq_ignore_ascii_case(name))
|
||||
}
|
||||
|
||||
pub fn header_value<'a>(headers: &'a [(String, String)], name: &str) -> Option<&'a str> {
|
||||
headers
|
||||
.iter()
|
||||
.find(|(key, _)| key.eq_ignore_ascii_case(name))
|
||||
.map(|(_, value)| value.as_str())
|
||||
}
|
||||
|
||||
pub fn without_headers(headers: Vec<(String, String)>, names: &[&str]) -> Vec<(String, String)> {
|
||||
headers
|
||||
.into_iter()
|
||||
.filter(|(key, _)| !names.iter().any(|name| key.eq_ignore_ascii_case(name)))
|
||||
.collect()
|
||||
}
|
||||
|
||||
pub fn has_bearer_auth(headers: &[(String, String)]) -> bool {
|
||||
headers.iter().any(|(name, value)| {
|
||||
if !name.eq_ignore_ascii_case("authorization") {
|
||||
|
|
@ -194,6 +208,30 @@ mod tests {
|
|||
assert!(!has_header(&headers, "authorization"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn header_value_reads_the_first_match_in_any_case() {
|
||||
let headers = vec![
|
||||
("X-Api-Key".to_string(), "first".to_string()),
|
||||
("x-api-key".to_string(), "second".to_string()),
|
||||
];
|
||||
assert_eq!(header_value(&headers, "x-API-key"), Some("first"));
|
||||
assert_eq!(header_value(&headers, "authorization"), None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn without_headers_drops_every_casing_of_the_named_headers_and_keeps_order() {
|
||||
let headers = vec![
|
||||
("X-Api-Key".to_string(), "k".to_string()),
|
||||
("anthropic-version".to_string(), "v".to_string()),
|
||||
("AUTHORIZATION".to_string(), "Bearer t".to_string()),
|
||||
("x-api-key".to_string(), "k2".to_string()),
|
||||
];
|
||||
assert_eq!(
|
||||
without_headers(headers, &["x-api-key", "authorization"]),
|
||||
vec![("anthropic-version".to_string(), "v".to_string())]
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn auth_header_detection_is_case_insensitive() {
|
||||
let headers = vec![
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
[package]
|
||||
name = "litellm"
|
||||
version = "0.0.1"
|
||||
edition.workspace = true
|
||||
|
|
|
|||
|
|
@ -4,7 +4,7 @@ use serde_json::Value;
|
|||
use time::OffsetDateTime;
|
||||
use url::Url;
|
||||
|
||||
use crate::{Error, anthropic::messages::transformation::resolve_anthropic_api_base};
|
||||
use crate::{Error, anthropic::common_utils::resolve_anthropic_api_base};
|
||||
|
||||
const BATCHES_PATH_SUFFIX: &str = "/v1/messages/batches";
|
||||
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
use litellm_auth::{CredentialPlacement, SecretValue};
|
||||
use litellm_auth::SecretValue;
|
||||
use litellm_core_utils::{
|
||||
core_helpers::{finish_reason_for, unix_now, usage_from_parts},
|
||||
prompt_templates::factory::{Conversation, build_conversation},
|
||||
|
|
@ -12,9 +12,11 @@ use serde_json::{Map, Value, json};
|
|||
use crate::{
|
||||
Error,
|
||||
anthropic::{
|
||||
ANTHROPIC_OAUTH_TOKEN_PREFIX,
|
||||
chat::handler::ModelResponseIterator,
|
||||
messages::transformation::{complete_anthropic_url, resolve_anthropic_api_key},
|
||||
common_utils::{
|
||||
API_KEY_PLACEMENT, complete_anthropic_url, forwarded_oauth_bearer,
|
||||
resolve_anthropic_api_key,
|
||||
},
|
||||
},
|
||||
base_llm::{
|
||||
anthropic_messages::streaming::anthropic_sse_event_stream,
|
||||
|
|
@ -50,15 +52,6 @@ pub struct AnthropicConfig;
|
|||
|
||||
pub const ANTHROPIC_CHAT_COMPLETIONS_CONFIG: AnthropicConfig = AnthropicConfig;
|
||||
|
||||
fn forwards_oauth_bearer(headers: &[(String, String)]) -> bool {
|
||||
headers.iter().any(|(name, value)| {
|
||||
name.eq_ignore_ascii_case("authorization")
|
||||
&& value
|
||||
.strip_prefix("Bearer ")
|
||||
.is_some_and(|token| token.starts_with(ANTHROPIC_OAUTH_TOKEN_PREFIX))
|
||||
})
|
||||
}
|
||||
|
||||
impl BaseConfig for AnthropicConfig {
|
||||
fn supported_openai_param_mappings(&self) -> &'static [(&'static str, &'static str)] {
|
||||
SUPPORTED_PARAMS
|
||||
|
|
@ -160,14 +153,14 @@ impl BaseConfig for AnthropicConfig {
|
|||
_optional_params: &Map<String, Value>,
|
||||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> Result<ValidatedEnvironment, Error> {
|
||||
if forwards_oauth_bearer(&headers) {
|
||||
if forwarded_oauth_bearer(&headers).is_some() {
|
||||
return Ok(ValidatedEnvironment {
|
||||
headers,
|
||||
auth: AuthScheme::Forwarded,
|
||||
});
|
||||
}
|
||||
let auth = AuthScheme::Credential {
|
||||
placement: CredentialPlacement::Header("x-api-key"),
|
||||
placement: API_KEY_PLACEMENT,
|
||||
secret: SecretValue::new(resolve_anthropic_api_key(api_key, env_lookup)?),
|
||||
};
|
||||
Ok(ValidatedEnvironment { headers, auth })
|
||||
|
|
|
|||
|
|
@ -1,30 +1,33 @@
|
|||
use litellm_types::llms::anthropic_messages::anthropic_request::{
|
||||
AnthropicMessage, ContentBlock, EffortLevel, MessageContent,
|
||||
use litellm_auth::{CredentialPlacement, SecretValue};
|
||||
use litellm_http::request::{has_header, header_value, without_headers};
|
||||
use litellm_types::llms::{
|
||||
anthropic::{AnthropicBeta, BetaSet},
|
||||
anthropic_messages::anthropic_request::{
|
||||
AnthropicMessage, AnthropicTool, ContentBlock, EffortLevel, MessageContent,
|
||||
},
|
||||
};
|
||||
use litellm_types::recognized::Recognized;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::anthropic::ANTHROPIC_OAUTH_TOKEN_PREFIX;
|
||||
use crate::{
|
||||
anthropic::ANTHROPIC_OAUTH_TOKEN_PREFIX,
|
||||
base_llm::auth::{AuthScheme, Headers},
|
||||
};
|
||||
|
||||
pub const ANTHROPIC_OAUTH_BETA_HEADER: &str = "oauth-2025-04-20";
|
||||
pub const ANTHROPIC_ADVISOR_TOOL_TYPE: &str = "advisor_20260301";
|
||||
pub const ANTHROPIC_TOOL_SEARCH_TOOL_TYPES: [&str; 2] = [
|
||||
"tool_search_tool_regex_20251119",
|
||||
"tool_search_tool_bm25_20251119",
|
||||
];
|
||||
pub const ANTHROPIC_API_KEY_ENV: &str = "ANTHROPIC_API_KEY";
|
||||
pub const ANTHROPIC_AUTH_TOKEN_ENV: &str = "ANTHROPIC_AUTH_TOKEN";
|
||||
pub const ENCRYPTED_REASONING_SIGNATURE_PREFIX: &str = "litellm_encrypted_reasoning:";
|
||||
const THOUGHT_SIGNATURE_SEPARATOR: &str = "__thought__";
|
||||
|
||||
pub mod beta {
|
||||
pub const CONTEXT_MANAGEMENT_2025_06_27: &str = "context-management-2025-06-27";
|
||||
pub const COMPACT_2026_01_12: &str = "compact-2026-01-12";
|
||||
pub const COMPACT_2026_09_04: &str = "compact-2026-09-04";
|
||||
pub const STRUCTURED_OUTPUT: &str = "structured-outputs-2025-11-13";
|
||||
pub const ADVANCED_TOOL_USE_2025_11_20: &str = "advanced-tool-use-2025-11-20";
|
||||
pub const FAST_MODE_2026_02_01: &str = "fast-mode-2026-02-01";
|
||||
pub const ADVISOR_TOOL_2026_03_01: &str = "advisor-tool-2026-03-01";
|
||||
pub const PER_TURN_CONTROL_2026_07_01: &str = "per-turn-control-2026-07-01";
|
||||
}
|
||||
const BETA_HEADER: &str = "anthropic-beta";
|
||||
pub const ANTHROPIC_API_BASE_ENV: &str = "ANTHROPIC_API_BASE";
|
||||
pub const ANTHROPIC_BASE_URL_ENV: &str = "ANTHROPIC_BASE_URL";
|
||||
pub const DEFAULT_ANTHROPIC_API_BASE: &str = "https://api.anthropic.com";
|
||||
pub const MESSAGES_PATH_SUFFIX: &str = "/v1/messages";
|
||||
pub const API_KEY_PLACEMENT: CredentialPlacement = CredentialPlacement::Header("x-api-key");
|
||||
const API_KEY_HEADER: &str = API_KEY_PLACEMENT.header_name();
|
||||
const AUTHORIZATION: &str = CredentialPlacement::Bearer.header_name();
|
||||
const DIRECT_BROWSER_ACCESS_HEADER: &str = "anthropic-dangerous-direct-browser-access";
|
||||
|
||||
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq, Serialize, Deserialize)]
|
||||
pub struct SupportedEffortTiers {
|
||||
|
|
@ -111,42 +114,207 @@ impl AnthropicModelCapabilities {
|
|||
}
|
||||
}
|
||||
|
||||
pub fn is_anthropic_oauth_key(value: &str) -> bool {
|
||||
value
|
||||
.strip_prefix("Bearer ")
|
||||
.unwrap_or(value)
|
||||
.starts_with(ANTHROPIC_OAUTH_TOKEN_PREFIX)
|
||||
pub fn non_empty(value: Option<&str>) -> Option<&str> {
|
||||
value.map(str::trim).filter(|value| !value.is_empty())
|
||||
}
|
||||
|
||||
pub fn split_beta_values(header: Option<&str>) -> impl Iterator<Item = String> + '_ {
|
||||
header
|
||||
.into_iter()
|
||||
.flat_map(|value| value.split(','))
|
||||
.map(str::trim)
|
||||
.filter(|piece| !piece.is_empty())
|
||||
pub fn non_empty_env(env_lookup: &dyn Fn(&str) -> Option<String>, name: &str) -> Option<String> {
|
||||
env_lookup(name).filter(|value| !value.trim().is_empty())
|
||||
}
|
||||
|
||||
/// An Anthropic OAuth access token, which authenticates as a bearer instead of an `x-api-key`.
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||||
pub struct OauthToken<'a>(&'a str);
|
||||
|
||||
impl<'a> OauthToken<'a> {
|
||||
/// The raw token, as a caller passes it in `api_key`.
|
||||
pub fn parse(value: &'a str) -> Option<Self> {
|
||||
value
|
||||
.starts_with(ANTHROPIC_OAUTH_TOKEN_PREFIX)
|
||||
.then_some(Self(value))
|
||||
}
|
||||
|
||||
/// A configured key, which users paste either raw or already prefixed with `Bearer `.
|
||||
pub fn parse_key(value: &'a str) -> Option<Self> {
|
||||
Self::parse(value.strip_prefix("Bearer ").unwrap_or(value))
|
||||
}
|
||||
|
||||
pub fn as_str(self) -> &'a str {
|
||||
self.0
|
||||
}
|
||||
|
||||
pub fn into_auth(self) -> AuthScheme {
|
||||
AuthScheme::Credential {
|
||||
placement: CredentialPlacement::Bearer,
|
||||
secret: SecretValue::new(self.0),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Python's `AnthropicModelInfo.get_api_key`: the param, else `ANTHROPIC_API_KEY`.
|
||||
pub fn get_api_key(
|
||||
api_key: Option<&str>,
|
||||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> Option<String> {
|
||||
non_empty(api_key)
|
||||
.map(str::to_string)
|
||||
.or_else(|| non_empty_env(env_lookup, ANTHROPIC_API_KEY_ENV))
|
||||
}
|
||||
|
||||
pub fn join_beta_values(values: impl IntoIterator<Item = String>) -> String {
|
||||
let mut values: Vec<String> = values.into_iter().collect();
|
||||
values.sort();
|
||||
values.dedup();
|
||||
values.join(",")
|
||||
pub fn get_auth_token(env_lookup: &dyn Fn(&str) -> Option<String>) -> Option<String> {
|
||||
non_empty_env(env_lookup, ANTHROPIC_AUTH_TOKEN_ENV)
|
||||
}
|
||||
|
||||
pub fn is_tool_search_used(tools: Option<&[Value]>) -> bool {
|
||||
tools.into_iter().flatten().any(|tool| {
|
||||
tool.get("type")
|
||||
.and_then(Value::as_str)
|
||||
.is_some_and(|tool_type| ANTHROPIC_TOOL_SEARCH_TOOL_TYPES.contains(&tool_type))
|
||||
/// Python's `AnthropicModelInfo.get_auth_header`, naming the credential instead of building
|
||||
/// the header: the key goes in `x-api-key` unless it is an OAuth token, and without a key
|
||||
/// `ANTHROPIC_AUTH_TOKEN` is sent as a bearer.
|
||||
pub fn get_auth_header(
|
||||
api_key: Option<&str>,
|
||||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> Option<AuthScheme> {
|
||||
if let Some(key) = get_api_key(api_key, env_lookup) {
|
||||
return Some(match OauthToken::parse_key(&key) {
|
||||
Some(token) => token.into_auth(),
|
||||
None => AuthScheme::Credential {
|
||||
placement: API_KEY_PLACEMENT,
|
||||
secret: SecretValue::new(key),
|
||||
},
|
||||
});
|
||||
}
|
||||
get_auth_token(env_lookup).map(|token| AuthScheme::Credential {
|
||||
placement: CredentialPlacement::Bearer,
|
||||
secret: SecretValue::new(token),
|
||||
})
|
||||
}
|
||||
|
||||
pub fn has_advisor_tool(tools: Option<&[Value]>) -> bool {
|
||||
pub fn resolve_anthropic_api_key(
|
||||
api_key: Option<&str>,
|
||||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> Result<String, litellm_auth::Error> {
|
||||
get_api_key(api_key, env_lookup).ok_or(litellm_auth::Error::MissingApiKey {
|
||||
provider: "Anthropic",
|
||||
environment_variable: ANTHROPIC_API_KEY_ENV,
|
||||
})
|
||||
}
|
||||
|
||||
/// Whether the caller already forwarded an Anthropic credential, in either header.
|
||||
pub fn has_anthropic_credential(headers: &[(String, String)]) -> bool {
|
||||
has_header(headers, API_KEY_HEADER) || has_header(headers, AUTHORIZATION)
|
||||
}
|
||||
|
||||
pub fn resolve_anthropic_api_base(
|
||||
api_base: Option<&str>,
|
||||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> String {
|
||||
non_empty(api_base)
|
||||
.map(str::to_string)
|
||||
.or_else(|| non_empty_env(env_lookup, ANTHROPIC_API_BASE_ENV))
|
||||
.or_else(|| non_empty_env(env_lookup, ANTHROPIC_BASE_URL_ENV))
|
||||
.unwrap_or_else(|| DEFAULT_ANTHROPIC_API_BASE.to_string())
|
||||
}
|
||||
|
||||
pub fn complete_anthropic_url(
|
||||
api_base: Option<&str>,
|
||||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> String {
|
||||
let api_base = resolve_anthropic_api_base(api_base, env_lookup);
|
||||
|
||||
let api_base = api_base.trim_end_matches('/');
|
||||
if api_base.ends_with(MESSAGES_PATH_SUFFIX) {
|
||||
return api_base.to_string();
|
||||
}
|
||||
format!("{api_base}{MESSAGES_PATH_SUFFIX}")
|
||||
}
|
||||
|
||||
pub fn existing_betas(headers: &[(String, String)]) -> BetaSet {
|
||||
headers
|
||||
.iter()
|
||||
.filter(|(name, _)| name.eq_ignore_ascii_case(BETA_HEADER))
|
||||
.flat_map(|(_, value)| {
|
||||
value
|
||||
.parse::<BetaSet>()
|
||||
.unwrap_or_else(|never| match never {})
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
/// Python's `_merge_beta_headers`, over every casing of the header at once: the union of what
|
||||
/// the caller sent and `added` replaces the header, sorted and deduplicated. Headers without
|
||||
/// any beta value stay as they are.
|
||||
pub fn merge_beta_headers(headers: Headers, added: BetaSet) -> Headers {
|
||||
let merged = existing_betas(&headers).union(added);
|
||||
if merged.is_empty() {
|
||||
return headers;
|
||||
}
|
||||
without_headers(headers, &[BETA_HEADER])
|
||||
.into_iter()
|
||||
.chain([(BETA_HEADER.to_string(), merged.to_string())])
|
||||
.collect()
|
||||
}
|
||||
|
||||
/// The outcome of Python's `optionally_handle_anthropic_oauth`.
|
||||
#[derive(Clone, Debug, PartialEq)]
|
||||
pub enum OauthHandling {
|
||||
/// An OAuth token is the whole credential. The headers carry its companions and no
|
||||
/// longer any `x-api-key` or `authorization`, so the bearer is applied on top.
|
||||
Bearer {
|
||||
headers: Headers,
|
||||
token: SecretValue,
|
||||
},
|
||||
Untouched(Headers),
|
||||
}
|
||||
|
||||
/// The OAuth token a caller forwarded as `Authorization: Bearer sk-ant-oat…`.
|
||||
pub fn forwarded_oauth_bearer(headers: &[(String, String)]) -> Option<OauthToken<'_>> {
|
||||
header_value(headers, AUTHORIZATION)
|
||||
.and_then(|value| value.strip_prefix("Bearer "))
|
||||
.and_then(OauthToken::parse)
|
||||
}
|
||||
|
||||
fn with_oauth_companions(headers: Headers, dropped: &[&str]) -> Headers {
|
||||
merge_beta_headers(
|
||||
without_headers(headers, dropped),
|
||||
BetaSet::from_iter([AnthropicBeta::Oauth20250420]),
|
||||
)
|
||||
.into_iter()
|
||||
.chain([(DIRECT_BROWSER_ACCESS_HEADER.to_string(), "true".to_string())])
|
||||
.collect()
|
||||
}
|
||||
|
||||
pub fn optionally_handle_anthropic_oauth(headers: Headers, api_key: Option<&str>) -> OauthHandling {
|
||||
if let Some(token) =
|
||||
forwarded_oauth_bearer(&headers).map(|token| SecretValue::new(token.as_str()))
|
||||
{
|
||||
return OauthHandling::Bearer {
|
||||
headers: with_oauth_companions(headers, &[API_KEY_HEADER, AUTHORIZATION]),
|
||||
token,
|
||||
};
|
||||
}
|
||||
if let Some(token) = api_key.and_then(OauthToken::parse) {
|
||||
return OauthHandling::Bearer {
|
||||
headers: with_oauth_companions(headers, &[API_KEY_HEADER]),
|
||||
token: SecretValue::new(token.as_str()),
|
||||
};
|
||||
}
|
||||
OauthHandling::Untouched(headers)
|
||||
}
|
||||
|
||||
pub fn is_tool_search_used(tools: Option<&[Recognized<AnthropicTool>]>) -> bool {
|
||||
tools.into_iter().flatten().any(|tool| {
|
||||
matches!(
|
||||
tool,
|
||||
Recognized::Known(
|
||||
AnthropicTool::ToolSearchRegex { .. } | AnthropicTool::ToolSearchBm25 { .. }
|
||||
)
|
||||
)
|
||||
})
|
||||
}
|
||||
|
||||
pub fn has_advisor_tool(tools: Option<&[Recognized<AnthropicTool>]>) -> bool {
|
||||
tools
|
||||
.into_iter()
|
||||
.flatten()
|
||||
.any(|tool| tool.get("type").and_then(Value::as_str) == Some(ANTHROPIC_ADVISOR_TOOL_TYPE))
|
||||
.any(|tool| matches!(tool, Recognized::Known(AnthropicTool::Advisor { .. })))
|
||||
}
|
||||
|
||||
pub fn requires_native_compaction_beta(
|
||||
|
|
@ -521,8 +689,97 @@ mod tests {
|
|||
serde_json::from_value(messages).unwrap()
|
||||
}
|
||||
|
||||
fn tools(value: Option<Value>) -> Option<Vec<Value>> {
|
||||
value.map(|tools| tools.as_array().unwrap().clone())
|
||||
fn tools(value: Option<Value>) -> Option<Vec<Recognized<AnthropicTool>>> {
|
||||
value.map(|tools| serde_json::from_value(tools).unwrap())
|
||||
}
|
||||
|
||||
fn headers(pairs: &[(&str, &str)]) -> Headers {
|
||||
pairs
|
||||
.iter()
|
||||
.map(|(name, value)| (name.to_string(), value.to_string()))
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn betas(values: &[&str]) -> BetaSet {
|
||||
values.join(",").parse().unwrap()
|
||||
}
|
||||
|
||||
fn env(vars: &'static [(&'static str, &'static str)]) -> impl Fn(&str) -> Option<String> {
|
||||
move |name| {
|
||||
vars.iter()
|
||||
.find(|(key, _)| *key == name)
|
||||
.map(|(_, value)| value.to_string())
|
||||
}
|
||||
}
|
||||
|
||||
const BOTH_BASE_ENVS: &[(&str, &str)] = &[
|
||||
(ANTHROPIC_API_BASE_ENV, "https://api-base.example.com"),
|
||||
(ANTHROPIC_BASE_URL_ENV, "https://base-url.example.com"),
|
||||
];
|
||||
|
||||
#[rstest]
|
||||
#[case::public_endpoint_by_default(None, &[], "https://api.anthropic.com")]
|
||||
#[case::explicit_api_base_beats_env(
|
||||
Some("https://explicit.example.com"),
|
||||
BOTH_BASE_ENVS,
|
||||
"https://explicit.example.com"
|
||||
)]
|
||||
#[case::explicit_api_base_is_trimmed(
|
||||
Some(" https://explicit.example.com "),
|
||||
&[],
|
||||
"https://explicit.example.com"
|
||||
)]
|
||||
#[case::blank_api_base_falls_back_to_env(
|
||||
Some(" "),
|
||||
BOTH_BASE_ENVS,
|
||||
"https://api-base.example.com"
|
||||
)]
|
||||
#[case::api_base_env_beats_base_url_env(None, BOTH_BASE_ENVS, "https://api-base.example.com")]
|
||||
#[case::base_url_env_without_api_base_env(
|
||||
None,
|
||||
&[(ANTHROPIC_BASE_URL_ENV, "https://base-url.example.com")],
|
||||
"https://base-url.example.com"
|
||||
)]
|
||||
#[case::blank_api_base_env_falls_back_to_base_url_env(
|
||||
None,
|
||||
&[(ANTHROPIC_API_BASE_ENV, " \t "), (ANTHROPIC_BASE_URL_ENV, "https://base-url.example.com")],
|
||||
"https://base-url.example.com"
|
||||
)]
|
||||
#[case::blank_envs_fall_back_to_public_endpoint(
|
||||
None,
|
||||
&[(ANTHROPIC_API_BASE_ENV, ""), (ANTHROPIC_BASE_URL_ENV, " ")],
|
||||
"https://api.anthropic.com"
|
||||
)]
|
||||
fn api_base_resolution(
|
||||
#[case] api_base: Option<&str>,
|
||||
#[case] vars: &'static [(&'static str, &'static str)],
|
||||
#[case] expected: &str,
|
||||
) {
|
||||
assert_eq!(resolve_anthropic_api_base(api_base, &env(vars)), expected);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::forwarded_api_key(&[("X-Api-Key", "k")], true)]
|
||||
#[case::forwarded_bearer(&[("Authorization", "Bearer t")], true)]
|
||||
#[case::nothing_forwarded(&[("anthropic-version", "2023-06-01")], false)]
|
||||
fn forwarded_credential_is_detected_in_either_header(
|
||||
#[case] forwarded: &[(&str, &str)],
|
||||
#[case] expected: bool,
|
||||
) {
|
||||
let headers: Headers = forwarded
|
||||
.iter()
|
||||
.map(|(name, value)| (name.to_string(), value.to_string()))
|
||||
.collect();
|
||||
assert_eq!(has_anthropic_credential(&headers), expected);
|
||||
}
|
||||
|
||||
fn credential(auth: Option<AuthScheme>) -> Option<(&'static str, String)> {
|
||||
auth.map(|auth| match auth {
|
||||
AuthScheme::Credential { placement, secret } => {
|
||||
(placement.header_name(), secret.expose().to_string())
|
||||
}
|
||||
other => panic!("expected a credential, got {other:?}"),
|
||||
})
|
||||
}
|
||||
|
||||
fn tagged(encrypted: &str) -> String {
|
||||
|
|
@ -1195,50 +1452,268 @@ mod tests {
|
|||
assert_eq!(twice, once);
|
||||
}
|
||||
|
||||
const OAUTH_TOKEN: &str = "sk-ant-oat01-token";
|
||||
const OAUTH_BEARER: &str = "Bearer sk-ant-oat01-token";
|
||||
const REGULAR_KEY: &str = "sk-ant-api03-regular";
|
||||
const OAUTH_BETA: &str = "oauth-2025-04-20";
|
||||
const BROWSER_ACCESS: (&str, &str) = ("anthropic-dangerous-direct-browser-access", "true");
|
||||
|
||||
#[rstest]
|
||||
#[case::no_existing_header(None, "b", "b")]
|
||||
#[case::empty_existing_header(Some(""), "b", "b")]
|
||||
#[case::whitespace_existing_header(Some(" "), "b", "b")]
|
||||
#[case::sorted_after_merge(Some("c,a"), "b", "a,b,c")]
|
||||
#[case::already_present(Some("a,b"), "a", "a,b")]
|
||||
#[case::trimmed_and_deduplicated(Some("b, a ,b"), "c", "a,b,c")]
|
||||
#[case::blank_pieces_skipped(Some("a,,b"), "c", "a,b,c")]
|
||||
fn beta_values_merge_sorted_and_deduplicated(
|
||||
#[case] existing: Option<&str>,
|
||||
#[case] new_beta: &str,
|
||||
#[case] expected: &str,
|
||||
#[case::no_beta_header(&[("x-api-key", "k")], &[], &[("x-api-key", "k")])]
|
||||
#[case::blank_beta_header(&[("Anthropic-Beta", " , "), ("x-api-key", "k")], &[], &[("Anthropic-Beta", " , "), ("x-api-key", "k")])]
|
||||
#[case::added_to_no_header(&[("x-api-key", "k")], &["b"], &[("x-api-key", "k"), ("anthropic-beta", "b")])]
|
||||
#[case::added_to_blank_header(&[("anthropic-beta", " ")], &["b"], &[("anthropic-beta", "b")])]
|
||||
#[case::sorted_after_merge(&[("anthropic-beta", "c,a")], &["b"], &[("anthropic-beta", "a,b,c")])]
|
||||
#[case::already_present(&[("anthropic-beta", "a,b")], &["a"], &[("anthropic-beta", "a,b")])]
|
||||
#[case::existing_normalized_without_additions(
|
||||
&[("Anthropic-Beta", "b, a ,b"), ("x-api-key", "k")],
|
||||
&[],
|
||||
&[("x-api-key", "k"), ("anthropic-beta", "a,b")]
|
||||
)]
|
||||
#[case::every_casing_unioned_into_one_lowercase_header(
|
||||
&[("anthropic-beta", "a"), ("ANTHROPIC-BETA", "c"), ("x-api-key", "k")],
|
||||
&["b"],
|
||||
&[("x-api-key", "k"), ("anthropic-beta", "a,b,c")]
|
||||
)]
|
||||
fn merge_beta_headers_replaces_the_header_with_the_sorted_union(
|
||||
#[case] input: &[(&str, &str)],
|
||||
#[case] added: &[&str],
|
||||
#[case] expected: &[(&str, &str)],
|
||||
) {
|
||||
assert_eq!(
|
||||
join_beta_values(split_beta_values(existing).chain([new_beta.to_string()])),
|
||||
merge_beta_headers(headers(input), betas(added)),
|
||||
headers(expected)
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::raw_token(OAUTH_TOKEN, Some(OAUTH_TOKEN))]
|
||||
#[case::bare_prefix(ANTHROPIC_OAUTH_TOKEN_PREFIX, Some(ANTHROPIC_OAUTH_TOKEN_PREFIX))]
|
||||
#[case::bearer_token(OAUTH_BEARER, None)]
|
||||
#[case::api_key(REGULAR_KEY, None)]
|
||||
#[case::empty("", None)]
|
||||
#[case::uppercase_prefix("sk-ant-OAT01-abc123", None)]
|
||||
#[case::prefix_not_at_start(" sk-ant-oat01-abc123", None)]
|
||||
fn oauth_token_parses_only_the_raw_token(#[case] value: &str, #[case] expected: Option<&str>) {
|
||||
assert_eq!(OauthToken::parse(value).map(OauthToken::as_str), expected);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::raw_token(OAUTH_TOKEN, Some(OAUTH_TOKEN))]
|
||||
#[case::bearer_token(OAUTH_BEARER, Some(OAUTH_TOKEN))]
|
||||
#[case::api_key(REGULAR_KEY, None)]
|
||||
#[case::bearer_api_key("Bearer sk-ant-api01-abc123", None)]
|
||||
#[case::empty("", None)]
|
||||
#[case::shouting_prefix("SK-ANT-OAT01-abc123", None)]
|
||||
#[case::lowercase_bearer("bearer sk-ant-oat01-abc123", None)]
|
||||
#[case::bearer_stripped_once("Bearer Bearer sk-ant-oat01-abc123", None)]
|
||||
fn oauth_key_parses_the_token_behind_an_optional_bearer(
|
||||
#[case] value: &str,
|
||||
#[case] expected: Option<&str>,
|
||||
) {
|
||||
assert_eq!(
|
||||
OauthToken::parse_key(value).map(OauthToken::as_str),
|
||||
expected
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::raw_token("sk-ant-oat01-abc123", true)]
|
||||
#[case::bearer_token("Bearer sk-ant-oat02-xyz789", true)]
|
||||
#[case::bare_prefix(ANTHROPIC_OAUTH_TOKEN_PREFIX, true)]
|
||||
#[case::api_key("sk-ant-api01-abc123", false)]
|
||||
#[case::bearer_api_key("Bearer sk-ant-api01-abc123", false)]
|
||||
#[case::empty("", false)]
|
||||
#[case::uppercase_prefix("sk-ant-OAT01-abc123", false)]
|
||||
#[case::shouting_prefix("SK-ANT-OAT01-abc123", false)]
|
||||
#[case::lowercase_bearer("bearer sk-ant-oat01-abc123", false)]
|
||||
#[case::bearer_stripped_once("Bearer Bearer sk-ant-oat01-abc123", false)]
|
||||
#[case::prefix_not_at_start(" sk-ant-oat01-abc123", false)]
|
||||
fn anthropic_oauth_key_detection(#[case] value: &str, #[case] expected: bool) {
|
||||
assert_eq!(is_anthropic_oauth_key(value), expected);
|
||||
#[case::bearer(&[("authorization", OAUTH_BEARER)], Some(OAUTH_TOKEN))]
|
||||
#[case::uppercase_header(&[("AUTHORIZATION", OAUTH_BEARER)], Some(OAUTH_TOKEN))]
|
||||
#[case::non_oauth_bearer(&[("authorization", "Bearer some-proxy-token")], None)]
|
||||
#[case::token_without_the_bearer_scheme(&[("authorization", OAUTH_TOKEN)], None)]
|
||||
#[case::lowercase_bearer_scheme(&[("authorization", "bearer sk-ant-oat01-token")], None)]
|
||||
#[case::token_in_x_api_key(&[("x-api-key", OAUTH_TOKEN)], None)]
|
||||
#[case::no_headers(&[], None)]
|
||||
fn forwarded_oauth_bearer_reads_the_authorization_header(
|
||||
#[case] forwarded: &[(&str, &str)],
|
||||
#[case] expected: Option<&str>,
|
||||
) {
|
||||
assert_eq!(
|
||||
forwarded_oauth_bearer(&headers(forwarded)).map(OauthToken::as_str),
|
||||
expected
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::regex_tool(Some(json!([{"type": ANTHROPIC_TOOL_SEARCH_TOOL_TYPES[0], "name": "tool_search_tool_regex"}])), true)]
|
||||
#[case::bm25_tool(Some(json!([{"type": ANTHROPIC_TOOL_SEARCH_TOOL_TYPES[1], "name": "tool_search_tool_bm25"}])), true)]
|
||||
#[case::forwarded_bearer_drops_forwarded_and_deployment_keys(
|
||||
&[("X-Api-Key", REGULAR_KEY), ("Authorization", OAUTH_BEARER)],
|
||||
Some(REGULAR_KEY),
|
||||
&[],
|
||||
)]
|
||||
#[case::forwarded_bearer_keeps_unrelated_headers_in_place(
|
||||
&[("anthropic-version", "2023-06-01"), ("authorization", OAUTH_BEARER)],
|
||||
None,
|
||||
&[("anthropic-version", "2023-06-01")],
|
||||
)]
|
||||
#[case::forwarded_bearer_wins_over_an_oauth_api_key(
|
||||
&[("authorization", OAUTH_BEARER)],
|
||||
Some("sk-ant-oat01-deployment"),
|
||||
&[],
|
||||
)]
|
||||
#[case::api_key_alone(&[], Some(OAUTH_TOKEN), &[])]
|
||||
#[case::api_key_removes_a_forwarded_x_api_key(&[("x-api-key", OAUTH_TOKEN)], Some(OAUTH_TOKEN), &[])]
|
||||
#[case::api_key_keeps_a_forwarded_non_oauth_bearer(
|
||||
&[("Authorization", "Bearer some-proxy-token")],
|
||||
Some(OAUTH_TOKEN),
|
||||
&[("Authorization", "Bearer some-proxy-token")],
|
||||
)]
|
||||
fn oauth_token_is_the_whole_credential(
|
||||
#[case] forwarded: &[(&str, &str)],
|
||||
#[case] api_key: Option<&str>,
|
||||
#[case] kept: &[(&str, &str)],
|
||||
) {
|
||||
let expected = kept
|
||||
.iter()
|
||||
.copied()
|
||||
.chain([("anthropic-beta", OAUTH_BETA), BROWSER_ACCESS])
|
||||
.collect::<Vec<_>>();
|
||||
assert_eq!(
|
||||
optionally_handle_anthropic_oauth(headers(forwarded), api_key),
|
||||
OauthHandling::Bearer {
|
||||
headers: headers(&expected),
|
||||
token: SecretValue::new(OAUTH_TOKEN),
|
||||
}
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::forwarded_bearer_merges_a_differently_cased_beta_header(
|
||||
&[("Anthropic-Beta", "web-search-2025-03-05"), ("authorization", OAUTH_BEARER)],
|
||||
None,
|
||||
)]
|
||||
#[case::forwarded_bearer_dedupes_an_existing_oauth_beta(
|
||||
&[("anthropic-beta", "web-search-2025-03-05, oauth-2025-04-20"), ("authorization", OAUTH_BEARER)],
|
||||
None,
|
||||
)]
|
||||
#[case::api_key_merges_the_existing_beta_header(
|
||||
&[("anthropic-beta", " web-search-2025-03-05 ,")],
|
||||
Some(OAUTH_TOKEN),
|
||||
)]
|
||||
#[case::forwarded_bearer_unions_every_beta_header_casing(
|
||||
&[("anthropic-beta", "oauth-2025-04-20"), ("ANTHROPIC-BETA", "web-search-2025-03-05"), ("authorization", OAUTH_BEARER)],
|
||||
None,
|
||||
)]
|
||||
fn oauth_beta_merges_into_existing_betas(
|
||||
#[case] forwarded: &[(&str, &str)],
|
||||
#[case] api_key: Option<&str>,
|
||||
) {
|
||||
assert_eq!(
|
||||
optionally_handle_anthropic_oauth(headers(forwarded), api_key),
|
||||
OauthHandling::Bearer {
|
||||
headers: headers(&[
|
||||
("anthropic-beta", "oauth-2025-04-20,web-search-2025-03-05"),
|
||||
BROWSER_ACCESS,
|
||||
]),
|
||||
token: SecretValue::new(OAUTH_TOKEN),
|
||||
}
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::x_api_key(&[("x-api-key", "caller-key")], Some("sk-other"))]
|
||||
#[case::non_oauth_bearer(&[("Authorization", "Bearer some-proxy-token")], Some(REGULAR_KEY))]
|
||||
#[case::oauth_token_without_the_bearer_scheme(&[("authorization", OAUTH_TOKEN)], None)]
|
||||
#[case::bearer_prefixed_api_key(&[], Some(OAUTH_BEARER))]
|
||||
#[case::nothing(&[], None)]
|
||||
fn without_an_oauth_token_the_headers_are_untouched(
|
||||
#[case] forwarded: &[(&str, &str)],
|
||||
#[case] api_key: Option<&str>,
|
||||
) {
|
||||
assert_eq!(
|
||||
optionally_handle_anthropic_oauth(headers(forwarded), api_key),
|
||||
OauthHandling::Untouched(headers(forwarded))
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::api_key_param(Some("sk-param"), &[], Some(("x-api-key", "sk-param")))]
|
||||
#[case::api_key_param_over_env_key_and_auth_token(
|
||||
Some("sk-param"),
|
||||
&[("ANTHROPIC_API_KEY", "sk-env"), ("ANTHROPIC_AUTH_TOKEN", "env-token")],
|
||||
Some(("x-api-key", "sk-param")),
|
||||
)]
|
||||
#[case::env_key_without_a_param(None, &[("ANTHROPIC_API_KEY", "sk-env")], Some(("x-api-key", "sk-env")))]
|
||||
#[case::env_key_when_the_param_is_blank(Some(" "), &[("ANTHROPIC_API_KEY", "sk-env")], Some(("x-api-key", "sk-env")))]
|
||||
#[case::env_key_over_auth_token(
|
||||
None,
|
||||
&[("ANTHROPIC_API_KEY", "sk-env"), ("ANTHROPIC_AUTH_TOKEN", "env-token")],
|
||||
Some(("x-api-key", "sk-env")),
|
||||
)]
|
||||
#[case::auth_token_as_a_bearer(
|
||||
None,
|
||||
&[("ANTHROPIC_AUTH_TOKEN", "env-token")],
|
||||
Some(("Authorization", "env-token")),
|
||||
)]
|
||||
#[case::auth_token_when_the_env_key_is_blank(
|
||||
None,
|
||||
&[("ANTHROPIC_API_KEY", " \t"), ("ANTHROPIC_AUTH_TOKEN", "env-token")],
|
||||
Some(("Authorization", "env-token")),
|
||||
)]
|
||||
#[case::oauth_param_as_a_bearer(Some(OAUTH_TOKEN), &[], Some(("Authorization", OAUTH_TOKEN)))]
|
||||
#[case::bearer_prefixed_oauth_env_key_as_a_bearer_once(
|
||||
None,
|
||||
&[("ANTHROPIC_API_KEY", OAUTH_BEARER)],
|
||||
Some(("Authorization", OAUTH_TOKEN)),
|
||||
)]
|
||||
#[case::no_credentials(None, &[], None)]
|
||||
#[case::blank_everything(Some(""), &[("ANTHROPIC_API_KEY", " "), ("ANTHROPIC_AUTH_TOKEN", " \t")], None)]
|
||||
fn auth_header_prefers_the_key_then_the_auth_token(
|
||||
#[case] api_key: Option<&str>,
|
||||
#[case] vars: &'static [(&'static str, &'static str)],
|
||||
#[case] expected: Option<(&str, &str)>,
|
||||
) {
|
||||
assert_eq!(
|
||||
credential(get_auth_header(api_key, &env(vars))),
|
||||
expected.map(|(header, secret)| (header, secret.to_string()))
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::param(Some("sk-param"), &[("ANTHROPIC_API_KEY", "sk-env")], Ok("sk-param"))]
|
||||
#[case::blank_param_falls_back_to_env(Some(" "), &[("ANTHROPIC_API_KEY", "sk-env")], Ok("sk-env"))]
|
||||
#[case::env_without_param(None, &[("ANTHROPIC_API_KEY", "sk-env")], Ok("sk-env"))]
|
||||
#[case::blank_env_is_missing(None, &[("ANTHROPIC_API_KEY", " ")], Err(()))]
|
||||
#[case::nothing_is_missing(None, &[], Err(()))]
|
||||
fn api_key_resolution(
|
||||
#[case] api_key: Option<&str>,
|
||||
#[case] vars: &'static [(&'static str, &'static str)],
|
||||
#[case] expected: Result<&str, ()>,
|
||||
) {
|
||||
assert_eq!(
|
||||
resolve_anthropic_api_key(api_key, &env(vars)).map_err(|error| {
|
||||
assert!(matches!(
|
||||
error,
|
||||
litellm_auth::Error::MissingApiKey {
|
||||
provider: "Anthropic",
|
||||
environment_variable: "ANTHROPIC_API_KEY",
|
||||
}
|
||||
));
|
||||
}),
|
||||
expected.map(str::to_string)
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::absent(None, None)]
|
||||
#[case::blank(Some(" \t "), None)]
|
||||
#[case::padded(Some(" value "), Some("value"))]
|
||||
fn non_empty_trims_and_drops_blank_values(
|
||||
#[case] value: Option<&str>,
|
||||
#[case] expected: Option<&str>,
|
||||
) {
|
||||
assert_eq!(non_empty(value), expected);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::regex_tool(Some(json!([{"type": "tool_search_tool_regex_20251119", "name": "tool_search_tool_regex"}])), true)]
|
||||
#[case::bm25_tool(Some(json!([{"type": "tool_search_tool_bm25_20251119", "name": "tool_search_tool_bm25"}])), true)]
|
||||
#[case::after_other_tools(
|
||||
Some(json!([{"name": "get_weather", "input_schema": {}}, {"type": ANTHROPIC_TOOL_SEARCH_TOOL_TYPES[1]}])),
|
||||
Some(json!([{"name": "get_weather", "input_schema": {}}, {"type": "tool_search_tool_bm25_20251119"}])),
|
||||
true
|
||||
)]
|
||||
#[case::function_tool(Some(json!([{"type": "function", "function": {"name": "get_weather"}}])), false)]
|
||||
#[case::name_without_type(Some(json!([{"name": ANTHROPIC_TOOL_SEARCH_TOOL_TYPES[0]}])), false)]
|
||||
#[case::name_without_type(Some(json!([{"name": "tool_search_tool_regex_20251119"}])), false)]
|
||||
#[case::empty_tools(Some(json!([])), false)]
|
||||
#[case::no_tools(None, false)]
|
||||
fn tool_search_detection(#[case] input: Option<Value>, #[case] expected: bool) {
|
||||
|
|
@ -1246,8 +1721,8 @@ mod tests {
|
|||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::advisor_tool(Some(json!([{"type": ANTHROPIC_ADVISOR_TOOL_TYPE, "name": "advisor"}])), true)]
|
||||
#[case::after_other_tools(Some(json!([{"name": "f", "input_schema": {}}, {"type": ANTHROPIC_ADVISOR_TOOL_TYPE}])), true)]
|
||||
#[case::advisor_tool(Some(json!([{"type": "advisor_20260301", "name": "advisor"}])), true)]
|
||||
#[case::after_other_tools(Some(json!([{"name": "f", "input_schema": {}}, {"type": "advisor_20260301"}])), true)]
|
||||
#[case::tool_named_advisor(Some(json!([{"name": "advisor", "input_schema": {}}])), false)]
|
||||
#[case::other_server_tool(Some(json!([{"type": "web_search_20250305", "name": "web_search"}])), false)]
|
||||
#[case::empty_tools(Some(json!([])), false)]
|
||||
|
|
|
|||
|
|
@ -1,677 +0,0 @@
|
|||
use litellm_auth::{CredentialPlacement, SecretValue};
|
||||
use litellm_types::{
|
||||
llms::anthropic_messages::anthropic_request::AnthropicMessagesRequest, recognized::Recognized,
|
||||
};
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::{
|
||||
anthropic::{
|
||||
ANTHROPIC_OAUTH_TOKEN_PREFIX,
|
||||
common_utils::{
|
||||
ANTHROPIC_OAUTH_BETA_HEADER, beta, has_advisor_tool, is_anthropic_oauth_key,
|
||||
is_tool_search_used, join_beta_values, requires_native_compaction_beta,
|
||||
split_beta_values,
|
||||
},
|
||||
},
|
||||
base_llm::{
|
||||
anthropic_messages::transformation::Headers,
|
||||
auth::{AuthScheme, ValidatedEnvironment},
|
||||
},
|
||||
};
|
||||
|
||||
const ANTHROPIC_API_KEY_ENV: &str = "ANTHROPIC_API_KEY";
|
||||
const ANTHROPIC_AUTH_TOKEN_ENV: &str = "ANTHROPIC_AUTH_TOKEN";
|
||||
const BETA_HEADER: &str = "anthropic-beta";
|
||||
const AUTHORIZATION: &str = "authorization";
|
||||
const API_KEY_HEADER: &str = "x-api-key";
|
||||
const DIRECT_BROWSER_ACCESS_HEADER: &str = "anthropic-dangerous-direct-browser-access";
|
||||
|
||||
fn header_value<'a>(headers: &'a [(String, String)], name: &str) -> Option<&'a str> {
|
||||
headers
|
||||
.iter()
|
||||
.find(|(header, _)| header.eq_ignore_ascii_case(name))
|
||||
.map(|(_, value)| value.as_str())
|
||||
}
|
||||
|
||||
fn without(headers: Headers, names: &[&str]) -> Headers {
|
||||
headers
|
||||
.into_iter()
|
||||
.filter(|(header, _)| !names.iter().any(|name| header.eq_ignore_ascii_case(name)))
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn existing_betas(headers: &[(String, String)]) -> impl Iterator<Item = String> + '_ {
|
||||
headers
|
||||
.iter()
|
||||
.filter(|(header, _)| header.eq_ignore_ascii_case(BETA_HEADER))
|
||||
.flat_map(|(_, value)| split_beta_values(Some(value)))
|
||||
}
|
||||
|
||||
/// The OAuth headers Python's `optionally_handle_anthropic_oauth` sets next to the bearer.
|
||||
fn with_oauth_companions(headers: Headers, dropped: &[&str]) -> Headers {
|
||||
let beta =
|
||||
join_beta_values(existing_betas(&headers).chain([ANTHROPIC_OAUTH_BETA_HEADER.to_string()]));
|
||||
without(headers, &[dropped, &[BETA_HEADER]].concat())
|
||||
.into_iter()
|
||||
.chain([
|
||||
(BETA_HEADER.to_string(), beta),
|
||||
(DIRECT_BROWSER_ACCESS_HEADER.to_string(), "true".to_string()),
|
||||
])
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn non_empty(value: Option<&str>) -> Option<&str> {
|
||||
value.map(str::trim).filter(|value| !value.is_empty())
|
||||
}
|
||||
|
||||
fn bearer(token: &str) -> AuthScheme {
|
||||
AuthScheme::Credential {
|
||||
placement: CredentialPlacement::Bearer,
|
||||
secret: SecretValue::new(token),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn validate_environment(
|
||||
headers: Headers,
|
||||
api_key: Option<&str>,
|
||||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> Result<ValidatedEnvironment, litellm_auth::Error> {
|
||||
if let Some(token) = header_value(&headers, AUTHORIZATION)
|
||||
.and_then(|forwarded| forwarded.strip_prefix("Bearer "))
|
||||
.filter(|token| token.starts_with(ANTHROPIC_OAUTH_TOKEN_PREFIX))
|
||||
{
|
||||
let auth = bearer(token);
|
||||
return Ok(ValidatedEnvironment {
|
||||
headers: with_oauth_companions(headers, &[API_KEY_HEADER, AUTHORIZATION]),
|
||||
auth,
|
||||
});
|
||||
}
|
||||
if let Some(key) = api_key.filter(|key| key.starts_with(ANTHROPIC_OAUTH_TOKEN_PREFIX)) {
|
||||
return Ok(ValidatedEnvironment {
|
||||
headers: with_oauth_companions(headers, &[API_KEY_HEADER]),
|
||||
auth: bearer(key),
|
||||
});
|
||||
}
|
||||
if header_value(&headers, API_KEY_HEADER).is_some()
|
||||
|| header_value(&headers, AUTHORIZATION).is_some()
|
||||
{
|
||||
return Ok(ValidatedEnvironment {
|
||||
headers,
|
||||
auth: AuthScheme::Forwarded,
|
||||
});
|
||||
}
|
||||
let resolved_key = non_empty(api_key)
|
||||
.map(str::to_string)
|
||||
.or_else(|| env_lookup(ANTHROPIC_API_KEY_ENV).filter(|value| !value.trim().is_empty()));
|
||||
let auth = match resolved_key {
|
||||
Some(key) if is_anthropic_oauth_key(&key) => bearer(&key),
|
||||
Some(key) => AuthScheme::Credential {
|
||||
placement: CredentialPlacement::Header(API_KEY_HEADER),
|
||||
secret: SecretValue::new(key),
|
||||
},
|
||||
None => match env_lookup(ANTHROPIC_AUTH_TOKEN_ENV).filter(|value| !value.trim().is_empty())
|
||||
{
|
||||
Some(token) => bearer(&token),
|
||||
None => {
|
||||
return Err(litellm_auth::Error::MissingApiKey {
|
||||
provider: "Anthropic",
|
||||
environment_variable: ANTHROPIC_API_KEY_ENV,
|
||||
});
|
||||
}
|
||||
},
|
||||
};
|
||||
Ok(ValidatedEnvironment { headers, auth })
|
||||
}
|
||||
|
||||
fn context_management_betas(
|
||||
context_management: Option<&Value>,
|
||||
) -> impl Iterator<Item = &'static str> {
|
||||
let edits = context_management
|
||||
.and_then(|value| value.get("edits"))
|
||||
.and_then(Value::as_array)
|
||||
.map(Vec::as_slice)
|
||||
.unwrap_or(&[]);
|
||||
let (compact, other) = edits.iter().fold((false, false), |(compact, other), edit| {
|
||||
match edit.get("type").and_then(Value::as_str) {
|
||||
Some("compact_20260112") => (true, other),
|
||||
_ => (compact, true),
|
||||
}
|
||||
});
|
||||
compact
|
||||
.then_some(beta::COMPACT_2026_01_12)
|
||||
.into_iter()
|
||||
.chain(other.then_some(beta::CONTEXT_MANAGEMENT_2025_06_27))
|
||||
}
|
||||
|
||||
fn uses_structured_output(request: &AnthropicMessagesRequest) -> bool {
|
||||
request.params.output_format.is_some()
|
||||
|| request
|
||||
.params
|
||||
.output_config
|
||||
.as_ref()
|
||||
.and_then(Recognized::known)
|
||||
.is_some_and(|config| config.format.is_some())
|
||||
}
|
||||
|
||||
fn messages_carry_output_config(request: &AnthropicMessagesRequest) -> bool {
|
||||
request
|
||||
.messages
|
||||
.iter()
|
||||
.any(|message| message.extra.contains_key("output_config"))
|
||||
}
|
||||
|
||||
pub fn feature_betas(request: &AnthropicMessagesRequest) -> Vec<&'static str> {
|
||||
let tools = request.params.tools.as_deref();
|
||||
[
|
||||
requires_native_compaction_beta(request.params.compaction.as_ref(), &request.messages)
|
||||
.then_some(beta::COMPACT_2026_09_04),
|
||||
uses_structured_output(request).then_some(beta::STRUCTURED_OUTPUT),
|
||||
(request.params.speed.as_deref() == Some("fast")).then_some(beta::FAST_MODE_2026_02_01),
|
||||
messages_carry_output_config(request).then_some(beta::PER_TURN_CONTROL_2026_07_01),
|
||||
has_advisor_tool(tools).then_some(beta::ADVISOR_TOOL_2026_03_01),
|
||||
is_tool_search_used(tools).then_some(beta::ADVANCED_TOOL_USE_2025_11_20),
|
||||
]
|
||||
.into_iter()
|
||||
.flatten()
|
||||
.chain(context_management_betas(
|
||||
request.params.context_management.as_ref(),
|
||||
))
|
||||
.collect()
|
||||
}
|
||||
|
||||
pub fn with_feature_betas(headers: Headers, request: &AnthropicMessagesRequest) -> Headers {
|
||||
let existing = existing_betas(&headers).collect::<Vec<_>>();
|
||||
let features = feature_betas(request);
|
||||
if existing.is_empty() && features.is_empty() {
|
||||
return headers;
|
||||
}
|
||||
let merged = join_beta_values(
|
||||
existing
|
||||
.into_iter()
|
||||
.chain(features.into_iter().map(str::to_string)),
|
||||
);
|
||||
without(headers, &[BETA_HEADER])
|
||||
.into_iter()
|
||||
.chain([(BETA_HEADER.to_string(), merged)])
|
||||
.collect()
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use rstest::{fixture, rstest};
|
||||
use serde_json::json;
|
||||
|
||||
use super::*;
|
||||
use crate::base_llm::auth::resolve_auth;
|
||||
|
||||
const OAUTH_TOKEN: &str = "sk-ant-oat01-token";
|
||||
const OAUTH_BEARER: &str = "Bearer sk-ant-oat01-token";
|
||||
const REGULAR_KEY: &str = "sk-ant-api03-regular";
|
||||
const BROWSER_ACCESS: (&str, &str) = ("anthropic-dangerous-direct-browser-access", "true");
|
||||
|
||||
type Env = &'static [(&'static str, &'static str)];
|
||||
|
||||
fn request(fields: Value) -> AnthropicMessagesRequest {
|
||||
let mut body =
|
||||
json!({"model": "claude", "messages": [{"role": "user", "content": "Hello"}]});
|
||||
body.as_object_mut()
|
||||
.unwrap()
|
||||
.extend(fields.as_object().unwrap().clone());
|
||||
serde_json::from_value(body).unwrap()
|
||||
}
|
||||
|
||||
fn headers(pairs: &[(&str, &str)]) -> Headers {
|
||||
pairs
|
||||
.iter()
|
||||
.map(|(name, value)| (name.to_string(), value.to_string()))
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn betas(values: &[&str]) -> String {
|
||||
values.join(",")
|
||||
}
|
||||
|
||||
#[fixture]
|
||||
fn no_env() -> Env {
|
||||
&[]
|
||||
}
|
||||
|
||||
#[fixture]
|
||||
fn full_env() -> Env {
|
||||
&[
|
||||
("ANTHROPIC_API_KEY", "sk-env"),
|
||||
("ANTHROPIC_AUTH_TOKEN", "env-token"),
|
||||
]
|
||||
}
|
||||
|
||||
fn authenticate_with(
|
||||
forwarded: &[(&str, &str)],
|
||||
api_key: Option<&str>,
|
||||
env: Env,
|
||||
) -> Result<Headers, litellm_auth::Error> {
|
||||
let lookup = |name: &str| {
|
||||
env.iter()
|
||||
.find(|(key, _)| *key == name)
|
||||
.map(|(_, value)| value.to_string())
|
||||
};
|
||||
let validated = validate_environment(headers(forwarded), api_key, &lookup)?;
|
||||
let resolved = tokio::runtime::Builder::new_current_thread()
|
||||
.build()
|
||||
.unwrap()
|
||||
.block_on(resolve_auth(
|
||||
&litellm_auth::AuthServices::default(),
|
||||
validated,
|
||||
&lookup,
|
||||
))
|
||||
.unwrap();
|
||||
Ok(resolved.headers)
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::forwarded_bearer_drops_forwarded_and_deployment_keys(
|
||||
&[("X-Api-Key", REGULAR_KEY), ("Authorization", OAUTH_BEARER)],
|
||||
Some(REGULAR_KEY),
|
||||
OAUTH_BEARER,
|
||||
&[],
|
||||
)]
|
||||
#[case::forwarded_bearer_in_uppercase_authorization_header(
|
||||
&[("AUTHORIZATION", OAUTH_BEARER)],
|
||||
None,
|
||||
OAUTH_BEARER,
|
||||
&[],
|
||||
)]
|
||||
#[case::forwarded_bearer_keeps_unrelated_headers_in_place(
|
||||
&[("anthropic-version", "2023-06-01"), ("authorization", OAUTH_BEARER)],
|
||||
None,
|
||||
OAUTH_BEARER,
|
||||
&[("anthropic-version", "2023-06-01")],
|
||||
)]
|
||||
#[case::forwarded_bearer_wins_over_an_oauth_api_key(
|
||||
&[("authorization", OAUTH_BEARER)],
|
||||
Some("sk-ant-oat01-deployment"),
|
||||
OAUTH_BEARER,
|
||||
&[],
|
||||
)]
|
||||
#[case::api_key_authenticates_as_a_bearer(&[], Some(OAUTH_TOKEN), OAUTH_BEARER, &[])]
|
||||
#[case::api_key_removes_a_forwarded_x_api_key(
|
||||
&[("x-api-key", OAUTH_TOKEN)],
|
||||
Some(OAUTH_TOKEN),
|
||||
OAUTH_BEARER,
|
||||
&[],
|
||||
)]
|
||||
#[case::api_key_replaces_a_forwarded_non_oauth_bearer(
|
||||
&[("Authorization", "Bearer some-proxy-token")],
|
||||
Some(OAUTH_TOKEN),
|
||||
OAUTH_BEARER,
|
||||
&[],
|
||||
)]
|
||||
fn oauth_token_is_the_whole_credential(
|
||||
#[case] forwarded: &[(&str, &str)],
|
||||
#[case] api_key: Option<&str>,
|
||||
#[case] expected_bearer: &str,
|
||||
#[case] kept: &[(&str, &str)],
|
||||
full_env: Env,
|
||||
) {
|
||||
let expected = kept
|
||||
.iter()
|
||||
.copied()
|
||||
.chain([
|
||||
("anthropic-beta", ANTHROPIC_OAUTH_BETA_HEADER),
|
||||
BROWSER_ACCESS,
|
||||
("authorization", expected_bearer),
|
||||
])
|
||||
.collect::<Vec<_>>();
|
||||
assert_eq!(
|
||||
authenticate_with(forwarded, api_key, full_env).unwrap(),
|
||||
headers(&expected)
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::forwarded_bearer_merges_a_differently_cased_beta_header(
|
||||
&[("Anthropic-Beta", "web-search-2025-03-05"), ("authorization", OAUTH_BEARER)],
|
||||
None,
|
||||
)]
|
||||
#[case::forwarded_bearer_dedupes_an_existing_oauth_beta(
|
||||
&[("anthropic-beta", "web-search-2025-03-05, oauth-2025-04-20"), ("authorization", OAUTH_BEARER)],
|
||||
None,
|
||||
)]
|
||||
#[case::api_key_merges_the_existing_beta_header(
|
||||
&[("anthropic-beta", " web-search-2025-03-05 ,")],
|
||||
Some(OAUTH_TOKEN),
|
||||
)]
|
||||
#[case::forwarded_bearer_unions_every_beta_header_casing(
|
||||
&[("anthropic-beta", "oauth-2025-04-20"), ("ANTHROPIC-BETA", "web-search-2025-03-05"), ("authorization", OAUTH_BEARER)],
|
||||
None,
|
||||
)]
|
||||
fn oauth_beta_merges_into_existing_betas(
|
||||
#[case] forwarded: &[(&str, &str)],
|
||||
#[case] api_key: Option<&str>,
|
||||
no_env: Env,
|
||||
) {
|
||||
assert_eq!(
|
||||
authenticate_with(forwarded, api_key, no_env).unwrap(),
|
||||
headers(&[
|
||||
(
|
||||
"anthropic-beta",
|
||||
&betas(&[ANTHROPIC_OAUTH_BETA_HEADER, "web-search-2025-03-05"])
|
||||
),
|
||||
BROWSER_ACCESS,
|
||||
("authorization", OAUTH_BEARER),
|
||||
])
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::x_api_key_over_the_deployment_key(&[("x-api-key", "caller-key")], Some("sk-other"))]
|
||||
#[case::uppercase_x_api_key(&[("X-API-KEY", "caller-key")], None)]
|
||||
#[case::non_oauth_bearer(&[("Authorization", "Bearer some-proxy-token")], None)]
|
||||
#[case::non_oauth_bearer_over_a_regular_api_key(
|
||||
&[("authorization", "Bearer sk-ant-api03-forwarded")],
|
||||
Some(REGULAR_KEY),
|
||||
)]
|
||||
#[case::oauth_token_without_the_bearer_scheme(&[("authorization", OAUTH_TOKEN)], None)]
|
||||
#[case::oauth_token_behind_a_lowercase_bearer_scheme(
|
||||
&[("authorization", "bearer sk-ant-oat01-token")],
|
||||
None,
|
||||
)]
|
||||
fn forwarded_auth_header_is_kept_untouched(
|
||||
#[case] forwarded: &[(&str, &str)],
|
||||
#[case] api_key: Option<&str>,
|
||||
full_env: Env,
|
||||
) {
|
||||
assert_eq!(
|
||||
authenticate_with(forwarded, api_key, full_env).unwrap(),
|
||||
headers(forwarded)
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::api_key_param(Some("sk-param"), &[], ("x-api-key", "sk-param"))]
|
||||
#[case::api_key_param_over_env_key_and_auth_token(
|
||||
Some("sk-param"),
|
||||
&[("ANTHROPIC_API_KEY", "sk-env"), ("ANTHROPIC_AUTH_TOKEN", "env-token")],
|
||||
("x-api-key", "sk-param"),
|
||||
)]
|
||||
#[case::env_key_without_a_param(None, &[("ANTHROPIC_API_KEY", "sk-env")], ("x-api-key", "sk-env"))]
|
||||
#[case::env_key_when_the_param_is_empty(Some(""), &[("ANTHROPIC_API_KEY", "sk-env")], ("x-api-key", "sk-env"))]
|
||||
#[case::env_key_when_the_param_is_whitespace(
|
||||
Some(" "),
|
||||
&[("ANTHROPIC_API_KEY", "sk-env")],
|
||||
("x-api-key", "sk-env"),
|
||||
)]
|
||||
#[case::env_key_over_auth_token(
|
||||
None,
|
||||
&[("ANTHROPIC_API_KEY", "sk-env"), ("ANTHROPIC_AUTH_TOKEN", "env-token")],
|
||||
("x-api-key", "sk-env"),
|
||||
)]
|
||||
#[case::auth_token_as_a_bearer(
|
||||
None,
|
||||
&[("ANTHROPIC_AUTH_TOKEN", "env-token")],
|
||||
("authorization", "Bearer env-token"),
|
||||
)]
|
||||
#[case::auth_token_when_the_env_key_is_whitespace(
|
||||
None,
|
||||
&[("ANTHROPIC_API_KEY", " \t"), ("ANTHROPIC_AUTH_TOKEN", "env-token")],
|
||||
("authorization", "Bearer env-token"),
|
||||
)]
|
||||
#[case::oauth_env_key_as_a_plain_bearer(
|
||||
None,
|
||||
&[("ANTHROPIC_API_KEY", "sk-ant-oat01-env")],
|
||||
("authorization", "Bearer sk-ant-oat01-env"),
|
||||
)]
|
||||
fn credential_is_resolved_after_the_existing_headers(
|
||||
#[case] api_key: Option<&str>,
|
||||
#[case] env: Env,
|
||||
#[case] expected: (&str, &str),
|
||||
) {
|
||||
let forwarded = [("anthropic-beta", "web-search-2025-03-05")];
|
||||
assert_eq!(
|
||||
authenticate_with(&forwarded, api_key, env).unwrap(),
|
||||
headers(&[forwarded[0], expected])
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::no_credentials(&[], None, &[])]
|
||||
#[case::empty_api_key(&[], Some(""), &[])]
|
||||
#[case::whitespace_only_env_values(
|
||||
&[],
|
||||
None,
|
||||
&[("ANTHROPIC_API_KEY", " "), ("ANTHROPIC_AUTH_TOKEN", " \t")],
|
||||
)]
|
||||
#[case::unrelated_forwarded_headers(&[("anthropic-beta", "web-search-2025-03-05")], None, &[])]
|
||||
fn missing_credentials_are_an_auth_error(
|
||||
#[case] forwarded: &[(&str, &str)],
|
||||
#[case] api_key: Option<&str>,
|
||||
#[case] env: Env,
|
||||
) {
|
||||
assert!(matches!(
|
||||
authenticate_with(forwarded, api_key, env),
|
||||
Err(litellm_auth::Error::MissingApiKey {
|
||||
provider: "Anthropic",
|
||||
environment_variable: "ANTHROPIC_API_KEY",
|
||||
})
|
||||
));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::no_features(json!({}), &[])]
|
||||
#[case::output_format(json!({"output_format": {"type": "json_schema"}}), &[beta::STRUCTURED_OUTPUT])]
|
||||
#[case::null_output_format(json!({"output_format": null}), &[])]
|
||||
#[case::output_config_format(
|
||||
json!({"output_config": {"format": {"type": "json_schema"}, "effort": "xhigh"}}),
|
||||
&[beta::STRUCTURED_OUTPUT]
|
||||
)]
|
||||
#[case::null_output_config_format(json!({"output_config": {"format": null}}), &[])]
|
||||
#[case::top_level_output_config_without_format(json!({"output_config": {"effort": "high"}}), &[])]
|
||||
#[case::fast_speed(json!({"speed": "fast"}), &[beta::FAST_MODE_2026_02_01])]
|
||||
#[case::standard_speed(json!({"speed": "standard"}), &[])]
|
||||
#[case::compaction_param(json!({"compaction": {"enabled": true}}), &[beta::COMPACT_2026_09_04])]
|
||||
#[case::empty_compaction_param(json!({"compaction": {}}), &[beta::COMPACT_2026_09_04])]
|
||||
#[case::signed_compaction_block_in_history(
|
||||
json!({"messages": [
|
||||
{"role": "assistant", "content": [{"type": "compaction", "content": "summary", "signature": "sig"}]},
|
||||
{"role": "user", "content": "Continue"},
|
||||
]}),
|
||||
&[beta::COMPACT_2026_09_04]
|
||||
)]
|
||||
#[case::unsigned_compaction_block_in_history(
|
||||
json!({"messages": [
|
||||
{"role": "assistant", "content": [{"type": "compaction", "content": "summary", "signature": ""}]},
|
||||
{"role": "user", "content": "Continue"},
|
||||
]}),
|
||||
&[]
|
||||
)]
|
||||
#[case::advisor_tool(
|
||||
json!({"tools": [{"type": "advisor_20260301", "name": "advisor", "model": "claude-opus-4-6"}]}),
|
||||
&[beta::ADVISOR_TOOL_2026_03_01]
|
||||
)]
|
||||
#[case::no_tools(json!({"tools": []}), &[])]
|
||||
#[case::regex_tool_search(
|
||||
json!({"tools": [{"type": "tool_search_tool_regex_20251119"}]}),
|
||||
&[beta::ADVANCED_TOOL_USE_2025_11_20]
|
||||
)]
|
||||
#[case::bm25_tool_search(
|
||||
json!({"tools": [{"type": "tool_search_tool_bm25_20251119"}]}),
|
||||
&[beta::ADVANCED_TOOL_USE_2025_11_20]
|
||||
)]
|
||||
#[case::unrelated_server_tool(json!({"tools": [{"type": "web_search_20250305", "name": "web_search"}]}), &[])]
|
||||
#[case::only_compact_edits(
|
||||
json!({"context_management": {"edits": [{"type": "compact_20260112"}]}}),
|
||||
&[beta::COMPACT_2026_01_12]
|
||||
)]
|
||||
#[case::only_other_edits(
|
||||
json!({"context_management": {"edits": [{"type": "clear_tool_uses_20250919", "keep": {"type": "tool_uses", "value": 3}}]}}),
|
||||
&[beta::CONTEXT_MANAGEMENT_2025_06_27]
|
||||
)]
|
||||
#[case::compact_and_other_edits(
|
||||
json!({"context_management": {"edits": [{"type": "compact_20260112"}, {"type": "clear_tool_uses_20250919"}]}}),
|
||||
&[beta::COMPACT_2026_01_12, beta::CONTEXT_MANAGEMENT_2025_06_27]
|
||||
)]
|
||||
#[case::edit_without_a_type(json!({"context_management": {"edits": [{}]}}), &[beta::CONTEXT_MANAGEMENT_2025_06_27])]
|
||||
#[case::empty_edits(json!({"context_management": {"edits": []}}), &[])]
|
||||
#[case::context_management_without_edits(json!({"context_management": {}}), &[])]
|
||||
#[case::per_message_output_config(
|
||||
json!({"messages": [{"role": "user", "content": "hi", "output_config": {"effort": "low"}}]}),
|
||||
&[beta::PER_TURN_CONTROL_2026_07_01]
|
||||
)]
|
||||
#[case::per_message_null_output_config(
|
||||
json!({"messages": [{"role": "user", "content": "hi", "output_config": null}]}),
|
||||
&[beta::PER_TURN_CONTROL_2026_07_01]
|
||||
)]
|
||||
fn feature_betas_follow_the_request(#[case] fields: Value, #[case] expected: &[&str]) {
|
||||
assert_eq!(feature_betas(&request(fields)), expected);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::no_betas(&[("x-api-key", "k"), ("anthropic-version", "2023-06-01")], json!({}))]
|
||||
#[case::blank_beta_header(&[("Anthropic-Beta", " , "), ("x-api-key", "k")], json!({}))]
|
||||
fn headers_without_any_beta_value_are_untouched(
|
||||
#[case] input: &[(&str, &str)],
|
||||
#[case] fields: Value,
|
||||
) {
|
||||
assert_eq!(
|
||||
with_feature_betas(headers(input), &request(fields)),
|
||||
headers(input)
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::feature_beta_is_appended(
|
||||
&[("x-api-key", "k")],
|
||||
json!({"speed": "fast"}),
|
||||
&[("x-api-key", "k"), ("anthropic-beta", beta::FAST_MODE_2026_02_01)],
|
||||
)]
|
||||
#[case::existing_betas_are_normalized_without_features(
|
||||
&[("Anthropic-Beta", "web-search-2025-03-05, interleaved-thinking-2025-05-14 ,web-search-2025-03-05"), ("x-api-key", "k")],
|
||||
json!({}),
|
||||
&[("x-api-key", "k"), ("anthropic-beta", "interleaved-thinking-2025-05-14,web-search-2025-03-05")],
|
||||
)]
|
||||
#[case::existing_advisor_beta_is_kept_without_an_advisor_tool(
|
||||
&[("anthropic-beta", beta::ADVISOR_TOOL_2026_03_01)],
|
||||
json!({"tools": []}),
|
||||
&[("anthropic-beta", beta::ADVISOR_TOOL_2026_03_01)],
|
||||
)]
|
||||
#[case::feature_already_sent_is_not_duplicated(
|
||||
&[("anthropic-beta", beta::FAST_MODE_2026_02_01)],
|
||||
json!({"speed": "fast"}),
|
||||
&[("anthropic-beta", beta::FAST_MODE_2026_02_01)],
|
||||
)]
|
||||
fn feature_betas_merge_into_the_headers(
|
||||
#[case] input: &[(&str, &str)],
|
||||
#[case] fields: Value,
|
||||
#[case] expected: &[(&str, &str)],
|
||||
) {
|
||||
assert_eq!(
|
||||
with_feature_betas(headers(input), &request(fields)),
|
||||
headers(expected)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn differently_cased_beta_header_is_replaced_by_one_sorted_header() {
|
||||
let merged = with_feature_betas(
|
||||
headers(&[("Anthropic-Beta", "interleaved-thinking-2025-05-14")]),
|
||||
&request(
|
||||
json!({"messages": [{"role": "system", "content": "env", "output_config": {"effort": "low"}}]}),
|
||||
),
|
||||
);
|
||||
assert_eq!(
|
||||
merged,
|
||||
headers(&[(
|
||||
"anthropic-beta",
|
||||
&betas(&[
|
||||
"interleaved-thinking-2025-05-14",
|
||||
beta::PER_TURN_CONTROL_2026_07_01
|
||||
])
|
||||
)])
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn every_beta_header_casing_is_unioned_into_one_header() {
|
||||
let merged = with_feature_betas(
|
||||
headers(&[
|
||||
("anthropic-beta", "interleaved-thinking-2025-05-14"),
|
||||
("Anthropic-Beta", "web-search-2025-03-05"),
|
||||
]),
|
||||
&request(json!({"speed": "fast"})),
|
||||
);
|
||||
assert_eq!(
|
||||
merged,
|
||||
headers(&[(
|
||||
"anthropic-beta",
|
||||
&betas(&[
|
||||
beta::FAST_MODE_2026_02_01,
|
||||
"interleaved-thinking-2025-05-14",
|
||||
"web-search-2025-03-05"
|
||||
])
|
||||
)])
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn unknown_client_betas_survive_alongside_the_added_one() {
|
||||
let client_betas = [
|
||||
"claude-code-20250219",
|
||||
"interleaved-thinking-2025-05-14",
|
||||
beta::CONTEXT_MANAGEMENT_2025_06_27,
|
||||
beta::PER_TURN_CONTROL_2026_07_01,
|
||||
"effort-2025-11-24",
|
||||
];
|
||||
let merged = with_feature_betas(
|
||||
headers(&[("anthropic-beta", &betas(&client_betas))]),
|
||||
&request(
|
||||
json!({"messages": [{"role": "user", "content": "hi", "output_config": {"effort": "low"}}]}),
|
||||
),
|
||||
);
|
||||
assert_eq!(
|
||||
merged,
|
||||
headers(&[(
|
||||
"anthropic-beta",
|
||||
&betas(&[
|
||||
"claude-code-20250219",
|
||||
beta::CONTEXT_MANAGEMENT_2025_06_27,
|
||||
"effort-2025-11-24",
|
||||
"interleaved-thinking-2025-05-14",
|
||||
beta::PER_TURN_CONTROL_2026_07_01,
|
||||
])
|
||||
)])
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn every_feature_merges_with_the_oauth_beta_sorted_and_last() {
|
||||
let oauth_headers = authenticate_with(&[], Some(OAUTH_TOKEN), &[]).unwrap();
|
||||
let all_features = request(json!({
|
||||
"compaction": {"enabled": true},
|
||||
"output_format": {"type": "json_schema"},
|
||||
"speed": "fast",
|
||||
"tools": [{"type": "advisor_20260301"}, {"type": "tool_search_tool_bm25_20251119"}],
|
||||
"context_management": {"edits": [{"type": "compact_20260112"}, {"type": "clear_thinking_20251015"}]},
|
||||
"messages": [{"role": "user", "content": "hi", "output_config": {"effort": "low"}}],
|
||||
}));
|
||||
assert_eq!(
|
||||
with_feature_betas(oauth_headers, &all_features),
|
||||
headers(&[
|
||||
BROWSER_ACCESS,
|
||||
("authorization", OAUTH_BEARER),
|
||||
(
|
||||
"anthropic-beta",
|
||||
&betas(&[
|
||||
beta::ADVANCED_TOOL_USE_2025_11_20,
|
||||
beta::ADVISOR_TOOL_2026_03_01,
|
||||
beta::COMPACT_2026_01_12,
|
||||
beta::COMPACT_2026_09_04,
|
||||
beta::CONTEXT_MANAGEMENT_2025_06_27,
|
||||
beta::FAST_MODE_2026_02_01,
|
||||
ANTHROPIC_OAUTH_BETA_HEADER,
|
||||
beta::PER_TURN_CONTROL_2026_07_01,
|
||||
beta::STRUCTURED_OUTPUT,
|
||||
])
|
||||
),
|
||||
])
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
@ -1,5 +1,4 @@
|
|||
pub mod handler;
|
||||
pub mod headers;
|
||||
pub mod streaming_iterator;
|
||||
pub mod thinking;
|
||||
pub mod transformation;
|
||||
|
|
|
|||
|
|
@ -1,31 +1,35 @@
|
|||
use litellm_auth::CredentialPlacement;
|
||||
use litellm_core_utils::settings::{Lookup, ProcessEnvironment};
|
||||
use litellm_types::llms::anthropic_messages::anthropic_request::{
|
||||
AnthropicMessagesOptionalParams, AnthropicMessagesRequest,
|
||||
use litellm_types::{
|
||||
llms::{
|
||||
anthropic::{AnthropicBeta, BetaSet},
|
||||
anthropic_messages::anthropic_request::{
|
||||
AnthropicMessage, AnthropicMessagesOptionalParams, AnthropicMessagesRequest,
|
||||
ContextEdit, ContextManagement, Speed,
|
||||
},
|
||||
},
|
||||
recognized::Recognized,
|
||||
};
|
||||
use serde_json::{Map, Value, json};
|
||||
|
||||
use super::{
|
||||
headers::{validate_environment, with_feature_betas},
|
||||
thinking::{ThinkingBudgets, ThinkingContext, translate_thinking},
|
||||
};
|
||||
use super::thinking::{ThinkingBudgets, ThinkingContext, translate_thinking};
|
||||
use crate::{
|
||||
Error,
|
||||
anthropic::common_utils::{
|
||||
AnthropicModelCapabilities, has_advisor_tool, strip_advisor_blocks,
|
||||
strip_encrypted_reasoning_blocks,
|
||||
ANTHROPIC_API_BASE_ENV, ANTHROPIC_API_KEY_ENV, ANTHROPIC_AUTH_TOKEN_ENV,
|
||||
ANTHROPIC_BASE_URL_ENV, AnthropicModelCapabilities, OauthHandling, complete_anthropic_url,
|
||||
get_auth_header, has_advisor_tool, has_anthropic_credential, is_tool_search_used,
|
||||
merge_beta_headers, optionally_handle_anthropic_oauth, requires_native_compaction_beta,
|
||||
strip_advisor_blocks, strip_encrypted_reasoning_blocks,
|
||||
},
|
||||
base_llm::anthropic_messages::transformation::{
|
||||
BaseAnthropicMessagesConfig, Headers, MessagesTransformContext, ValidatedEnvironment,
|
||||
base_llm::{
|
||||
anthropic_messages::transformation::{
|
||||
BaseAnthropicMessagesConfig, Headers, MessagesTransformContext, ValidatedEnvironment,
|
||||
},
|
||||
auth::AuthScheme,
|
||||
},
|
||||
};
|
||||
|
||||
const ANTHROPIC_API_KEY_ENV: &str = "ANTHROPIC_API_KEY";
|
||||
const ANTHROPIC_AUTH_TOKEN_ENV: &str = "ANTHROPIC_AUTH_TOKEN";
|
||||
const ANTHROPIC_API_BASE_ENV: &str = "ANTHROPIC_API_BASE";
|
||||
const ANTHROPIC_BASE_URL_ENV: &str = "ANTHROPIC_BASE_URL";
|
||||
const DEFAULT_ANTHROPIC_API_BASE: &str = "https://api.anthropic.com";
|
||||
const MESSAGES_PATH_SUFFIX: &str = "/v1/messages";
|
||||
|
||||
pub struct AnthropicMessagesConfig;
|
||||
|
||||
pub const ANTHROPIC_MESSAGES_CONFIG: AnthropicMessagesConfig = AnthropicMessagesConfig;
|
||||
|
|
@ -66,18 +70,15 @@ impl BaseAnthropicMessagesConfig for AnthropicMessagesConfig {
|
|||
context: &MessagesTransformContext,
|
||||
) -> Result<AnthropicMessagesRequest, Error> {
|
||||
if request.params.max_tokens.is_none() {
|
||||
return Err(Error::InvalidRequest(
|
||||
"max_tokens is required for Anthropic /v1/messages API".to_string(),
|
||||
));
|
||||
return Err(Error::MissingField("max_tokens"));
|
||||
}
|
||||
let request = drop_unsupported_params(request, context)?;
|
||||
let request = translate_thinking(request, &context.thinking)?;
|
||||
let context_management = request
|
||||
.params
|
||||
.context_management
|
||||
.as_ref()
|
||||
.and_then(map_openai_context_management_to_anthropic)
|
||||
.or_else(|| request.params.context_management.clone());
|
||||
.clone()
|
||||
.map(map_openai_context_management_to_anthropic);
|
||||
let messages = if has_advisor_tool(request.params.tools.as_deref()) {
|
||||
request.messages
|
||||
} else {
|
||||
|
|
@ -102,6 +103,8 @@ impl BaseAnthropicMessagesConfig for AnthropicMessagesConfig {
|
|||
]
|
||||
}
|
||||
|
||||
/// Python's `validate_anthropic_messages_environment` up to the beta merge, which
|
||||
/// `request_headers` does once the request is final.
|
||||
fn validate_environment(
|
||||
&self,
|
||||
headers: Headers,
|
||||
|
|
@ -109,14 +112,99 @@ impl BaseAnthropicMessagesConfig for AnthropicMessagesConfig {
|
|||
_model: &str,
|
||||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> Result<ValidatedEnvironment, Error> {
|
||||
validate_environment(headers, api_key, env_lookup).map_err(Error::from)
|
||||
let headers = match optionally_handle_anthropic_oauth(headers, api_key) {
|
||||
OauthHandling::Bearer { headers, token } => {
|
||||
return Ok(ValidatedEnvironment {
|
||||
headers,
|
||||
auth: AuthScheme::Credential {
|
||||
placement: CredentialPlacement::Bearer,
|
||||
secret: token,
|
||||
},
|
||||
});
|
||||
}
|
||||
OauthHandling::Untouched(headers) => headers,
|
||||
};
|
||||
if has_anthropic_credential(&headers) {
|
||||
return Ok(ValidatedEnvironment {
|
||||
headers,
|
||||
auth: AuthScheme::Forwarded,
|
||||
});
|
||||
}
|
||||
let auth = get_auth_header(api_key, env_lookup).ok_or(Error::Auth(
|
||||
litellm_auth::Error::MissingApiKey {
|
||||
provider: "Anthropic",
|
||||
environment_variable: ANTHROPIC_API_KEY_ENV,
|
||||
},
|
||||
))?;
|
||||
Ok(ValidatedEnvironment { headers, auth })
|
||||
}
|
||||
|
||||
fn request_headers(&self, headers: Headers, request: &AnthropicMessagesRequest) -> Headers {
|
||||
with_feature_betas(headers, request)
|
||||
update_headers_with_anthropic_beta(headers, request)
|
||||
}
|
||||
}
|
||||
|
||||
fn update_headers_with_anthropic_beta(
|
||||
headers: Headers,
|
||||
request: &AnthropicMessagesRequest,
|
||||
) -> Headers {
|
||||
merge_beta_headers(headers, feature_betas(request))
|
||||
}
|
||||
|
||||
fn feature_betas(request: &AnthropicMessagesRequest) -> BetaSet {
|
||||
let params = &request.params;
|
||||
let tools = params.tools.as_deref();
|
||||
[
|
||||
requires_native_compaction_beta(params.compaction.as_ref(), &request.messages)
|
||||
.then_some(AnthropicBeta::Compact20260904),
|
||||
uses_structured_output(params).then_some(AnthropicBeta::StructuredOutputs20251113),
|
||||
(params.speed == Some(Recognized::Known(Speed::Fast)))
|
||||
.then_some(AnthropicBeta::FastMode20260201),
|
||||
messages_carry_output_config(&request.messages)
|
||||
.then_some(AnthropicBeta::PerTurnControl20260701),
|
||||
has_advisor_tool(tools).then_some(AnthropicBeta::AdvisorTool20260301),
|
||||
is_tool_search_used(tools).then_some(AnthropicBeta::AdvancedToolUse20251120),
|
||||
]
|
||||
.into_iter()
|
||||
.flatten()
|
||||
.chain(context_management_betas(params.context_management.as_ref()))
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn is_compact_edit(edit: &Recognized<ContextEdit>) -> bool {
|
||||
matches!(edit, Recognized::Known(ContextEdit::Compact { .. }))
|
||||
}
|
||||
|
||||
fn context_management_betas(
|
||||
context_management: Option<&Recognized<ContextManagement>>,
|
||||
) -> impl Iterator<Item = AnthropicBeta> {
|
||||
let edits = context_management
|
||||
.and_then(Recognized::known)
|
||||
.and_then(|context_management| context_management.edits.as_deref())
|
||||
.unwrap_or_default();
|
||||
let compact = edits.iter().any(is_compact_edit);
|
||||
let other = edits.iter().any(|edit| !is_compact_edit(edit));
|
||||
compact
|
||||
.then_some(AnthropicBeta::Compact20260112)
|
||||
.into_iter()
|
||||
.chain(other.then_some(AnthropicBeta::ContextManagement20250627))
|
||||
}
|
||||
|
||||
fn uses_structured_output(params: &AnthropicMessagesOptionalParams) -> bool {
|
||||
params.output_format.is_some()
|
||||
|| params
|
||||
.output_config
|
||||
.as_ref()
|
||||
.and_then(Recognized::known)
|
||||
.is_some_and(|config| config.format.is_some())
|
||||
}
|
||||
|
||||
fn messages_carry_output_config(messages: &[AnthropicMessage]) -> bool {
|
||||
messages
|
||||
.iter()
|
||||
.any(|message| message.extra.contains_key("output_config"))
|
||||
}
|
||||
|
||||
fn unsupported_param(model: &str, param: &str, value: &str, hint: &str) -> Error {
|
||||
Error::InvalidRequest(format!(
|
||||
"{model} does not support {param}={value}. {hint}To drop unsupported params, set `litellm.drop_params = True`."
|
||||
|
|
@ -136,9 +224,9 @@ fn drop_unsupported_params(
|
|||
Err(unsupported_param(&model, param, &value, hint))
|
||||
};
|
||||
let params = request.params;
|
||||
let speed = match params.speed.as_deref() {
|
||||
let speed = match ¶ms.speed {
|
||||
Some(speed) if !capabilities.supports_speed => {
|
||||
reject("speed", format!("'{speed}'"), "")?;
|
||||
reject("speed", format!("'{}'", speed_text(speed)), "")?;
|
||||
None
|
||||
}
|
||||
_ => params.speed.clone(),
|
||||
|
|
@ -178,101 +266,73 @@ fn drop_unsupported_params(
|
|||
})
|
||||
}
|
||||
|
||||
pub fn map_openai_context_management_to_anthropic(context_management: &Value) -> Option<Value> {
|
||||
match context_management {
|
||||
Value::Object(edits) if edits.contains_key("edits") => Some(context_management.clone()),
|
||||
Value::Array(entries) => {
|
||||
let edits: Vec<Value> = entries
|
||||
.iter()
|
||||
.filter_map(Value::as_object)
|
||||
.filter(|entry| entry.get("type").and_then(Value::as_str) == Some("compaction"))
|
||||
.map(|entry| {
|
||||
let trigger = entry.get("compact_threshold").and_then(Value::as_f64).map(
|
||||
|threshold| json!({"type": "input_tokens", "value": threshold as i64}),
|
||||
);
|
||||
let passthrough = entry
|
||||
.iter()
|
||||
.filter(|(key, _)| !matches!(key.as_str(), "type" | "compact_threshold"))
|
||||
.map(|(key, value)| (key.clone(), value.clone()));
|
||||
Value::Object(
|
||||
[("type".to_string(), json!("compact_20260112"))]
|
||||
.into_iter()
|
||||
.chain(trigger.map(|trigger| ("trigger".to_string(), trigger)))
|
||||
.chain(passthrough)
|
||||
.collect::<Map<String, Value>>(),
|
||||
)
|
||||
})
|
||||
.collect();
|
||||
(!edits.is_empty()).then(|| json!({"edits": edits}))
|
||||
}
|
||||
_ => None,
|
||||
fn speed_text(speed: &Recognized<Speed>) -> String {
|
||||
match speed {
|
||||
Recognized::Known(speed) => speed.as_str().to_string(),
|
||||
Recognized::Unrecognized(Value::String(text)) => text.clone(),
|
||||
Recognized::Unrecognized(other) => other.to_string(),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn non_empty(value: Option<&str>) -> Option<&str> {
|
||||
value.map(str::trim).filter(|value| !value.is_empty())
|
||||
}
|
||||
|
||||
pub fn resolve_anthropic_api_key(
|
||||
api_key: Option<&str>,
|
||||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> Result<String, litellm_auth::Error> {
|
||||
non_empty(api_key)
|
||||
.map(str::to_string)
|
||||
.or_else(|| env_lookup(ANTHROPIC_API_KEY_ENV).filter(|value| !value.trim().is_empty()))
|
||||
.ok_or(litellm_auth::Error::MissingApiKey {
|
||||
provider: "Anthropic",
|
||||
environment_variable: ANTHROPIC_API_KEY_ENV,
|
||||
})
|
||||
}
|
||||
|
||||
pub fn complete_anthropic_url(
|
||||
api_base: Option<&str>,
|
||||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> String {
|
||||
let api_base = resolve_anthropic_api_base(api_base, env_lookup);
|
||||
|
||||
let api_base = api_base.trim_end_matches('/');
|
||||
if api_base.ends_with(MESSAGES_PATH_SUFFIX) {
|
||||
return api_base.to_string();
|
||||
fn compact_edit_from_openai(entry: &Map<String, Value>) -> Option<ContextEdit> {
|
||||
if entry.get("type").and_then(Value::as_str) != Some("compaction") {
|
||||
return None;
|
||||
}
|
||||
format!("{api_base}{MESSAGES_PATH_SUFFIX}")
|
||||
let trigger = entry
|
||||
.get("compact_threshold")
|
||||
.and_then(Value::as_f64)
|
||||
.map(|threshold| json!({"type": "input_tokens", "value": threshold as i64}));
|
||||
let passthrough = entry
|
||||
.iter()
|
||||
.filter(|(key, _)| !matches!(key.as_str(), "type" | "compact_threshold"))
|
||||
.map(|(key, value)| (key.clone(), value.clone()));
|
||||
Some(ContextEdit::Compact {
|
||||
extra: trigger
|
||||
.map(|trigger| ("trigger".to_string(), trigger))
|
||||
.into_iter()
|
||||
.chain(passthrough)
|
||||
.collect(),
|
||||
})
|
||||
}
|
||||
|
||||
pub fn resolve_anthropic_api_base(
|
||||
api_base: Option<&str>,
|
||||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> String {
|
||||
let env = |name: &str| env_lookup(name).filter(|value| !value.trim().is_empty());
|
||||
non_empty(api_base)
|
||||
.map(str::to_string)
|
||||
.or_else(|| env(ANTHROPIC_API_BASE_ENV))
|
||||
.or_else(|| env(ANTHROPIC_BASE_URL_ENV))
|
||||
.unwrap_or_else(|| DEFAULT_ANTHROPIC_API_BASE.to_string())
|
||||
/// An OpenAI-style `context_management` list becomes Anthropic `edits` when it holds
|
||||
/// compaction entries. Anything else, native edits included, is sent as it came.
|
||||
pub fn map_openai_context_management_to_anthropic(
|
||||
context_management: Recognized<ContextManagement>,
|
||||
) -> Recognized<ContextManagement> {
|
||||
let Recognized::Unrecognized(Value::Array(entries)) = &context_management else {
|
||||
return context_management;
|
||||
};
|
||||
let edits: Vec<Recognized<ContextEdit>> = entries
|
||||
.iter()
|
||||
.filter_map(Value::as_object)
|
||||
.filter_map(compact_edit_from_openai)
|
||||
.map(Recognized::Known)
|
||||
.collect();
|
||||
if edits.is_empty() {
|
||||
return context_management;
|
||||
}
|
||||
Recognized::Known(ContextManagement {
|
||||
edits: Some(edits),
|
||||
extra: Map::new(),
|
||||
})
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::process::Command;
|
||||
|
||||
use litellm_auth::CredentialPlacement;
|
||||
use rstest::{fixture, rstest};
|
||||
|
||||
use super::*;
|
||||
use crate::{
|
||||
anthropic::common_utils::{ENCRYPTED_REASONING_SIGNATURE_PREFIX, beta},
|
||||
base_llm::auth::AuthScheme,
|
||||
};
|
||||
use crate::anthropic::common_utils::ENCRYPTED_REASONING_SIGNATURE_PREFIX;
|
||||
|
||||
type Env = &'static [(&'static str, &'static str)];
|
||||
|
||||
const BOTH_BASE_ENVS: Env = &[
|
||||
(ANTHROPIC_API_BASE_ENV, "https://api-base.example.com"),
|
||||
(ANTHROPIC_BASE_URL_ENV, "https://base-url.example.com"),
|
||||
];
|
||||
const API_KEY_ENV: Env = &[(ANTHROPIC_API_KEY_ENV, "sk-env")];
|
||||
const MISSING_API_KEY: &str =
|
||||
"Missing Anthropic API Key - Set `api_key` or the ANTHROPIC_API_KEY environment variable";
|
||||
const OAUTH_TOKEN: &str = "sk-ant-oat01-token";
|
||||
const OAUTH_BEARER: &str = "Bearer sk-ant-oat01-token";
|
||||
const OAUTH_BETA: &str = "oauth-2025-04-20";
|
||||
const BROWSER_ACCESS: (&str, &str) = ("anthropic-dangerous-direct-browser-access", "true");
|
||||
const LOW_BUDGET_ENV: &str = "DEFAULT_REASONING_EFFORT_LOW_THINKING_BUDGET";
|
||||
const PROCESS_ENV_PROBE: &str = "LITELLM_MESSAGES_TRANSFORM_CONTEXT_PROBE";
|
||||
|
||||
|
|
@ -377,7 +437,7 @@ mod tests {
|
|||
fn missing_max_tokens_is_rejected(#[case] fields: Value, unmapped: AnthropicModelCapabilities) {
|
||||
assert_eq!(
|
||||
transform(fields, unmapped, false),
|
||||
invalid("max_tokens is required for Anthropic /v1/messages API")
|
||||
Err(Error::MissingField("max_tokens"))
|
||||
);
|
||||
}
|
||||
|
||||
|
|
@ -569,17 +629,19 @@ mod tests {
|
|||
#[case::empty_list(json!([]), None)]
|
||||
#[case::anthropic_edits_pass_through(
|
||||
json!({"edits": [{"type": "compact_20260112", "trigger": {"type": "input_tokens", "value": 150000}}]}),
|
||||
Some(json!({"edits": [{"type": "compact_20260112", "trigger": {"type": "input_tokens", "value": 150000}}]}))
|
||||
None
|
||||
)]
|
||||
#[case::object_without_edits(json!({"type": "compaction"}), None)]
|
||||
#[case::scalar(json!("compaction"), None)]
|
||||
fn openai_context_management_maps_to_anthropic_edits(
|
||||
#[case] context_management: Value,
|
||||
#[case] expected: Option<Value>,
|
||||
#[case] mapped: Option<Value>,
|
||||
) {
|
||||
let parsed: Recognized<ContextManagement> =
|
||||
serde_json::from_value(context_management.clone()).unwrap();
|
||||
assert_eq!(
|
||||
map_openai_context_management_to_anthropic(&context_management),
|
||||
expected
|
||||
serde_json::to_value(map_openai_context_management_to_anthropic(parsed)).unwrap(),
|
||||
mapped.unwrap_or(context_management)
|
||||
);
|
||||
}
|
||||
|
||||
|
|
@ -723,47 +785,6 @@ mod tests {
|
|||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::public_endpoint_by_default(None, &[], "https://api.anthropic.com")]
|
||||
#[case::explicit_api_base_beats_env(
|
||||
Some("https://explicit.example.com"),
|
||||
BOTH_BASE_ENVS,
|
||||
"https://explicit.example.com"
|
||||
)]
|
||||
#[case::explicit_api_base_is_trimmed(
|
||||
Some(" https://explicit.example.com "),
|
||||
&[],
|
||||
"https://explicit.example.com"
|
||||
)]
|
||||
#[case::blank_api_base_falls_back_to_env(
|
||||
Some(" "),
|
||||
BOTH_BASE_ENVS,
|
||||
"https://api-base.example.com"
|
||||
)]
|
||||
#[case::api_base_env_beats_base_url_env(None, BOTH_BASE_ENVS, "https://api-base.example.com")]
|
||||
#[case::base_url_env_without_api_base_env(
|
||||
None,
|
||||
&[(ANTHROPIC_BASE_URL_ENV, "https://base-url.example.com")],
|
||||
"https://base-url.example.com"
|
||||
)]
|
||||
#[case::blank_api_base_env_falls_back_to_base_url_env(
|
||||
None,
|
||||
&[(ANTHROPIC_API_BASE_ENV, " \t "), (ANTHROPIC_BASE_URL_ENV, "https://base-url.example.com")],
|
||||
"https://base-url.example.com"
|
||||
)]
|
||||
#[case::blank_envs_fall_back_to_public_endpoint(
|
||||
None,
|
||||
&[(ANTHROPIC_API_BASE_ENV, ""), (ANTHROPIC_BASE_URL_ENV, " ")],
|
||||
"https://api.anthropic.com"
|
||||
)]
|
||||
fn api_base_resolution(
|
||||
#[case] api_base: Option<&str>,
|
||||
#[case] vars: Env,
|
||||
#[case] expected: &str,
|
||||
) {
|
||||
assert_eq!(resolve_anthropic_api_base(api_base, &env(vars)), expected);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::public_endpoint(None, &[], "https://api.anthropic.com/v1/messages")]
|
||||
#[case::base_url_env(
|
||||
|
|
@ -794,77 +815,274 @@ mod tests {
|
|||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::param_beats_env(Some("sk-param"), API_KEY_ENV, Ok("sk-param"))]
|
||||
#[case::param_is_trimmed(Some(" sk-param "), &[], Ok("sk-param"))]
|
||||
#[case::blank_param_falls_back_to_env(Some(" "), API_KEY_ENV, Ok("sk-env"))]
|
||||
#[case::env_without_param(None, API_KEY_ENV, Ok("sk-env"))]
|
||||
#[case::blank_env_is_missing(None, &[(ANTHROPIC_API_KEY_ENV, " ")], Err(MISSING_API_KEY))]
|
||||
#[case::nothing_is_missing(None, &[], Err(MISSING_API_KEY))]
|
||||
fn api_key_resolution(
|
||||
#[case] api_key: Option<&str>,
|
||||
#[case] vars: Env,
|
||||
#[case] expected: Result<&str, &str>,
|
||||
) {
|
||||
assert_eq!(
|
||||
resolve_anthropic_api_key(api_key, &env(vars)).map_err(|error| error.to_string()),
|
||||
expected.map(str::to_string).map_err(str::to_string)
|
||||
);
|
||||
fn betas(values: &[&str]) -> BetaSet {
|
||||
values.join(",").parse().unwrap()
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn config_reports_a_missing_key_as_an_auth_error() {
|
||||
fn validated(
|
||||
forwarded: &[(&str, &str)],
|
||||
api_key: Option<&str>,
|
||||
vars: Env,
|
||||
) -> Result<ValidatedEnvironment, Error> {
|
||||
ANTHROPIC_MESSAGES_CONFIG.validate_environment(
|
||||
headers(forwarded),
|
||||
api_key,
|
||||
"claude",
|
||||
&env(vars),
|
||||
)
|
||||
}
|
||||
|
||||
fn credential(auth: &AuthScheme) -> Option<(&'static str, &str)> {
|
||||
match auth {
|
||||
AuthScheme::Credential { placement, secret } => {
|
||||
Some((placement.header_name(), secret.expose()))
|
||||
}
|
||||
AuthScheme::Forwarded => None,
|
||||
other => panic!("unexpected auth scheme {other:?}"),
|
||||
}
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::forwarded_oauth_bearer(
|
||||
&[("anthropic-version", "2023-06-01"), ("X-Api-Key", "sk-caller"), ("Authorization", OAUTH_BEARER)],
|
||||
Some("sk-deployment"),
|
||||
&[("ANTHROPIC_API_KEY", "sk-env")],
|
||||
&[("anthropic-version", "2023-06-01"), ("anthropic-beta", OAUTH_BETA), BROWSER_ACCESS],
|
||||
Some(("Authorization", OAUTH_TOKEN)),
|
||||
)]
|
||||
#[case::oauth_api_key(
|
||||
&[("x-api-key", OAUTH_TOKEN), ("anthropic-beta", "web-search-2025-03-05")],
|
||||
Some(OAUTH_TOKEN),
|
||||
&[],
|
||||
&[("anthropic-beta", "oauth-2025-04-20,web-search-2025-03-05"), BROWSER_ACCESS],
|
||||
Some(("Authorization", OAUTH_TOKEN)),
|
||||
)]
|
||||
#[case::forwarded_x_api_key_is_kept_over_the_deployment_key(
|
||||
&[("X-API-KEY", "caller-key")],
|
||||
Some("sk-other"),
|
||||
&[("ANTHROPIC_API_KEY", "sk-env")],
|
||||
&[("X-API-KEY", "caller-key")],
|
||||
None,
|
||||
)]
|
||||
#[case::forwarded_non_oauth_bearer_is_kept(
|
||||
&[("Authorization", "Bearer some-proxy-token")],
|
||||
Some("sk-ant-api03-regular"),
|
||||
&[],
|
||||
&[("Authorization", "Bearer some-proxy-token")],
|
||||
None,
|
||||
)]
|
||||
#[case::oauth_token_without_the_bearer_scheme_is_kept(
|
||||
&[("authorization", OAUTH_TOKEN)],
|
||||
None,
|
||||
&[],
|
||||
&[("authorization", OAUTH_TOKEN)],
|
||||
None,
|
||||
)]
|
||||
#[case::api_key_param(
|
||||
&[("anthropic-beta", "web-search-2025-03-05")],
|
||||
Some("sk-param"),
|
||||
&[("ANTHROPIC_API_KEY", "sk-env"), ("ANTHROPIC_AUTH_TOKEN", "env-token")],
|
||||
&[("anthropic-beta", "web-search-2025-03-05")],
|
||||
Some(("x-api-key", "sk-param")),
|
||||
)]
|
||||
#[case::env_key_when_the_param_is_blank(
|
||||
&[],
|
||||
Some(" "),
|
||||
&[("ANTHROPIC_API_KEY", "sk-env"), ("ANTHROPIC_AUTH_TOKEN", "env-token")],
|
||||
&[],
|
||||
Some(("x-api-key", "sk-env")),
|
||||
)]
|
||||
#[case::auth_token_when_no_key_is_set(
|
||||
&[],
|
||||
None,
|
||||
&[("ANTHROPIC_API_KEY", " \t"), ("ANTHROPIC_AUTH_TOKEN", "env-token")],
|
||||
&[],
|
||||
Some(("Authorization", "env-token")),
|
||||
)]
|
||||
#[case::oauth_env_key_as_a_bearer(
|
||||
&[],
|
||||
None,
|
||||
&[("ANTHROPIC_API_KEY", "sk-ant-oat01-env")],
|
||||
&[],
|
||||
Some(("Authorization", "sk-ant-oat01-env")),
|
||||
)]
|
||||
fn validate_environment_shapes_the_headers_and_names_the_credential(
|
||||
#[case] forwarded: &[(&str, &str)],
|
||||
#[case] api_key: Option<&str>,
|
||||
#[case] vars: Env,
|
||||
#[case] expected_headers: &[(&str, &str)],
|
||||
#[case] expected_credential: Option<(&str, &str)>,
|
||||
) {
|
||||
let environment = validated(forwarded, api_key, vars).unwrap();
|
||||
assert_eq!(environment.headers, headers(expected_headers));
|
||||
assert_eq!(credential(&environment.auth), expected_credential);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::no_credentials(&[], None, &[])]
|
||||
#[case::empty_api_key(&[], Some(""), &[])]
|
||||
#[case::whitespace_only_env_values(&[], None, &[("ANTHROPIC_API_KEY", " "), ("ANTHROPIC_AUTH_TOKEN", " \t")])]
|
||||
#[case::unrelated_forwarded_headers(&[("anthropic-beta", "web-search-2025-03-05")], None, &[])]
|
||||
fn missing_credentials_are_an_auth_error(
|
||||
#[case] forwarded: &[(&str, &str)],
|
||||
#[case] api_key: Option<&str>,
|
||||
#[case] vars: Env,
|
||||
) {
|
||||
assert!(matches!(
|
||||
ANTHROPIC_MESSAGES_CONFIG.validate_environment(vec![], None, "claude", &no_env),
|
||||
validated(forwarded, api_key, vars),
|
||||
Err(Error::Auth(litellm_auth::Error::MissingApiKey {
|
||||
provider: "Anthropic",
|
||||
environment_variable: ANTHROPIC_API_KEY_ENV,
|
||||
environment_variable: "ANTHROPIC_API_KEY",
|
||||
}))
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn config_authenticates_with_the_anthropic_auth_token() {
|
||||
let validated = ANTHROPIC_MESSAGES_CONFIG
|
||||
.validate_environment(
|
||||
vec![],
|
||||
None,
|
||||
"claude",
|
||||
&env(&[("ANTHROPIC_AUTH_TOKEN", "auth-token")]),
|
||||
)
|
||||
.unwrap();
|
||||
assert!(matches!(
|
||||
validated.auth,
|
||||
AuthScheme::Credential {
|
||||
placement: CredentialPlacement::Bearer,
|
||||
ref secret
|
||||
} if secret.expose() == "auth-token"
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn config_requests_the_betas_the_request_features_need() {
|
||||
assert_eq!(
|
||||
ANTHROPIC_MESSAGES_CONFIG.request_headers(
|
||||
headers(&[("x-api-key", "sk")]),
|
||||
&request(json!({"speed": "fast"}))
|
||||
),
|
||||
headers(&[
|
||||
("x-api-key", "sk"),
|
||||
("anthropic-beta", beta::FAST_MODE_2026_02_01)
|
||||
])
|
||||
);
|
||||
#[rstest]
|
||||
#[case::no_features(json!({}), &[])]
|
||||
#[case::output_format(json!({"output_format": {"type": "json_schema"}}), &["structured-outputs-2025-11-13"])]
|
||||
#[case::null_output_format(json!({"output_format": null}), &[])]
|
||||
#[case::output_config_format(
|
||||
json!({"output_config": {"format": {"type": "json_schema"}, "effort": "xhigh"}}),
|
||||
&["structured-outputs-2025-11-13"]
|
||||
)]
|
||||
#[case::null_output_config_format(json!({"output_config": {"format": null}}), &[])]
|
||||
#[case::top_level_output_config_without_format(json!({"output_config": {"effort": "high"}}), &[])]
|
||||
#[case::fast_speed(json!({"speed": "fast"}), &["fast-mode-2026-02-01"])]
|
||||
#[case::standard_speed(json!({"speed": "standard"}), &[])]
|
||||
#[case::unknown_speed(json!({"speed": "turbo"}), &[])]
|
||||
#[case::compaction_param(json!({"compaction": {"enabled": true}}), &["compact-2026-09-04"])]
|
||||
#[case::empty_compaction_param(json!({"compaction": {}}), &["compact-2026-09-04"])]
|
||||
#[case::signed_compaction_block_in_history(
|
||||
json!({"messages": [
|
||||
{"role": "assistant", "content": [{"type": "compaction", "content": "summary", "signature": "sig"}]},
|
||||
{"role": "user", "content": "Continue"},
|
||||
]}),
|
||||
&["compact-2026-09-04"]
|
||||
)]
|
||||
#[case::unsigned_compaction_block_in_history(
|
||||
json!({"messages": [
|
||||
{"role": "assistant", "content": [{"type": "compaction", "content": "summary", "signature": ""}]},
|
||||
{"role": "user", "content": "Continue"},
|
||||
]}),
|
||||
&[]
|
||||
)]
|
||||
#[case::advisor_tool(
|
||||
json!({"tools": [{"type": "advisor_20260301", "name": "advisor", "model": "claude-opus-4-6"}]}),
|
||||
&["advisor-tool-2026-03-01"]
|
||||
)]
|
||||
#[case::no_tools(json!({"tools": []}), &[])]
|
||||
#[case::regex_tool_search(
|
||||
json!({"tools": [{"type": "tool_search_tool_regex_20251119"}]}),
|
||||
&["advanced-tool-use-2025-11-20"]
|
||||
)]
|
||||
#[case::bm25_tool_search(
|
||||
json!({"tools": [{"type": "tool_search_tool_bm25_20251119"}]}),
|
||||
&["advanced-tool-use-2025-11-20"]
|
||||
)]
|
||||
#[case::unrelated_server_tool(json!({"tools": [{"type": "web_search_20250305", "name": "web_search"}]}), &[])]
|
||||
#[case::only_compact_edits(
|
||||
json!({"context_management": {"edits": [{"type": "compact_20260112"}]}}),
|
||||
&["compact-2026-01-12"]
|
||||
)]
|
||||
#[case::only_other_edits(
|
||||
json!({"context_management": {"edits": [{"type": "clear_tool_uses_20250919", "keep": {"type": "tool_uses", "value": 3}}]}}),
|
||||
&["context-management-2025-06-27"]
|
||||
)]
|
||||
#[case::compact_and_other_edits(
|
||||
json!({"context_management": {"edits": [{"type": "compact_20260112"}, {"type": "clear_tool_uses_20250919"}]}}),
|
||||
&["compact-2026-01-12", "context-management-2025-06-27"]
|
||||
)]
|
||||
#[case::edit_without_a_type(json!({"context_management": {"edits": [{}]}}), &["context-management-2025-06-27"])]
|
||||
#[case::unknown_edit_type(json!({"context_management": {"edits": [{"type": "future"}]}}), &["context-management-2025-06-27"])]
|
||||
#[case::empty_edits(json!({"context_management": {"edits": []}}), &[])]
|
||||
#[case::context_management_without_edits(json!({"context_management": {}}), &[])]
|
||||
#[case::unmapped_openai_context_management(json!({"context_management": [{"type": "other"}]}), &[])]
|
||||
#[case::per_message_output_config(
|
||||
json!({"messages": [{"role": "user", "content": "hi", "output_config": {"effort": "low"}}]}),
|
||||
&["per-turn-control-2026-07-01"]
|
||||
)]
|
||||
#[case::per_message_null_output_config(
|
||||
json!({"messages": [{"role": "user", "content": "hi", "output_config": null}]}),
|
||||
&["per-turn-control-2026-07-01"]
|
||||
)]
|
||||
fn feature_betas_follow_the_request(#[case] fields: Value, #[case] expected: &[&str]) {
|
||||
assert_eq!(feature_betas(&request(fields)), betas(expected));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::absent(None, None)]
|
||||
#[case::blank(Some(" \t "), None)]
|
||||
#[case::padded(Some(" value "), Some("value"))]
|
||||
fn non_empty_trims_and_drops_blank_values(
|
||||
#[case] value: Option<&str>,
|
||||
#[case] expected: Option<&str>,
|
||||
#[case::no_betas(&[("x-api-key", "k"), ("anthropic-version", "2023-06-01")], json!({}), &[("x-api-key", "k"), ("anthropic-version", "2023-06-01")])]
|
||||
#[case::blank_beta_header(&[("Anthropic-Beta", " , "), ("x-api-key", "k")], json!({}), &[("Anthropic-Beta", " , "), ("x-api-key", "k")])]
|
||||
#[case::feature_beta_is_appended(
|
||||
&[("x-api-key", "k")],
|
||||
json!({"speed": "fast"}),
|
||||
&[("x-api-key", "k"), ("anthropic-beta", "fast-mode-2026-02-01")],
|
||||
)]
|
||||
#[case::existing_betas_are_normalized_without_features(
|
||||
&[("Anthropic-Beta", "web-search-2025-03-05, interleaved-thinking-2025-05-14 ,web-search-2025-03-05"), ("x-api-key", "k")],
|
||||
json!({}),
|
||||
&[("x-api-key", "k"), ("anthropic-beta", "interleaved-thinking-2025-05-14,web-search-2025-03-05")],
|
||||
)]
|
||||
#[case::existing_advisor_beta_is_kept_without_an_advisor_tool(
|
||||
&[("anthropic-beta", "advisor-tool-2026-03-01")],
|
||||
json!({"tools": []}),
|
||||
&[("anthropic-beta", "advisor-tool-2026-03-01")],
|
||||
)]
|
||||
#[case::feature_already_sent_is_not_duplicated(
|
||||
&[("anthropic-beta", "fast-mode-2026-02-01")],
|
||||
json!({"speed": "fast"}),
|
||||
&[("anthropic-beta", "fast-mode-2026-02-01")],
|
||||
)]
|
||||
#[case::differently_cased_beta_header_is_replaced_by_one_sorted_header(
|
||||
&[("Anthropic-Beta", "interleaved-thinking-2025-05-14")],
|
||||
json!({"messages": [{"role": "system", "content": "env", "output_config": {"effort": "low"}}]}),
|
||||
&[("anthropic-beta", "interleaved-thinking-2025-05-14,per-turn-control-2026-07-01")],
|
||||
)]
|
||||
#[case::every_beta_header_casing_is_unioned_into_one_header(
|
||||
&[("anthropic-beta", "interleaved-thinking-2025-05-14"), ("Anthropic-Beta", "web-search-2025-03-05")],
|
||||
json!({"speed": "fast"}),
|
||||
&[("anthropic-beta", "fast-mode-2026-02-01,interleaved-thinking-2025-05-14,web-search-2025-03-05")],
|
||||
)]
|
||||
#[case::unknown_client_betas_survive_alongside_the_added_one(
|
||||
&[("anthropic-beta", "claude-code-20250219,interleaved-thinking-2025-05-14,context-management-2025-06-27,per-turn-control-2026-07-01,effort-2025-11-24")],
|
||||
json!({"messages": [{"role": "user", "content": "hi", "output_config": {"effort": "low"}}]}),
|
||||
&[("anthropic-beta", "claude-code-20250219,context-management-2025-06-27,effort-2025-11-24,interleaved-thinking-2025-05-14,per-turn-control-2026-07-01")],
|
||||
)]
|
||||
fn request_headers_merge_the_feature_betas(
|
||||
#[case] input: &[(&str, &str)],
|
||||
#[case] fields: Value,
|
||||
#[case] expected: &[(&str, &str)],
|
||||
) {
|
||||
assert_eq!(non_empty(value), expected);
|
||||
assert_eq!(
|
||||
ANTHROPIC_MESSAGES_CONFIG.request_headers(headers(input), &request(fields)),
|
||||
headers(expected)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn every_feature_merges_with_the_oauth_beta_sorted() {
|
||||
let environment = validated(&[], Some(OAUTH_TOKEN), &[]).unwrap();
|
||||
let all_features = request(json!({
|
||||
"compaction": {"enabled": true},
|
||||
"output_format": {"type": "json_schema"},
|
||||
"speed": "fast",
|
||||
"tools": [{"type": "advisor_20260301"}, {"type": "tool_search_tool_bm25_20251119"}],
|
||||
"context_management": {"edits": [{"type": "compact_20260112"}, {"type": "clear_thinking_20251015"}]},
|
||||
"messages": [{"role": "user", "content": "hi", "output_config": {"effort": "low"}}],
|
||||
}));
|
||||
assert_eq!(
|
||||
ANTHROPIC_MESSAGES_CONFIG.request_headers(environment.headers, &all_features),
|
||||
headers(&[
|
||||
BROWSER_ACCESS,
|
||||
(
|
||||
"anthropic-beta",
|
||||
"advanced-tool-use-2025-11-20,advisor-tool-2026-03-01,compact-2026-01-12,compact-2026-09-04,context-management-2025-06-27,fast-mode-2026-02-01,oauth-2025-04-20,per-turn-control-2026-07-01,structured-outputs-2025-11-13"
|
||||
),
|
||||
])
|
||||
);
|
||||
assert_eq!(
|
||||
credential(&environment.auth),
|
||||
Some(("Authorization", OAUTH_TOKEN))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
use litellm_auth::{CredentialPlacement, SecretValue};
|
||||
use litellm_auth::SecretValue;
|
||||
use litellm_http::request::{has_bearer_auth, has_header};
|
||||
use litellm_types::llms::anthropic_messages::{
|
||||
anthropic_request::{
|
||||
|
|
@ -10,8 +10,9 @@ use litellm_types::llms::anthropic_messages::{
|
|||
|
||||
use crate::{
|
||||
Error,
|
||||
anthropic::messages::transformation::{
|
||||
ANTHROPIC_MESSAGES_CONFIG, AnthropicMessagesConfig, non_empty,
|
||||
anthropic::{
|
||||
common_utils::{API_KEY_PLACEMENT, MESSAGES_PATH_SUFFIX, non_empty},
|
||||
messages::transformation::{ANTHROPIC_MESSAGES_CONFIG, AnthropicMessagesConfig},
|
||||
},
|
||||
base_llm::{
|
||||
anthropic_messages::transformation::{
|
||||
|
|
@ -24,9 +25,7 @@ use crate::{
|
|||
const AZURE_API_KEY_ENV: &str = "AZURE_API_KEY";
|
||||
const AZURE_API_BASE_ENV: &str = "AZURE_API_BASE";
|
||||
const ANTHROPIC_PATH_SEGMENT: &str = "/anthropic";
|
||||
const MESSAGES_PATH_SUFFIX: &str = "/v1/messages";
|
||||
const SYSTEM_ROLE: &str = "system";
|
||||
const API_KEY_HEADER: &str = "x-api-key";
|
||||
|
||||
pub struct AzureAnthropicMessagesConfig {
|
||||
anthropic: AnthropicMessagesConfig,
|
||||
|
|
@ -86,14 +85,14 @@ impl BaseAnthropicMessagesConfig for AzureAnthropicMessagesConfig {
|
|||
_model: &str,
|
||||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> Result<ValidatedEnvironment, Error> {
|
||||
if has_header(&headers, API_KEY_HEADER) || has_bearer_auth(&headers) {
|
||||
if has_header(&headers, API_KEY_PLACEMENT.header_name()) || has_bearer_auth(&headers) {
|
||||
return Ok(ValidatedEnvironment {
|
||||
headers,
|
||||
auth: AuthScheme::Forwarded,
|
||||
});
|
||||
}
|
||||
let auth = AuthScheme::Credential {
|
||||
placement: CredentialPlacement::Header(API_KEY_HEADER),
|
||||
placement: API_KEY_PLACEMENT,
|
||||
secret: SecretValue::new(resolve_azure_api_key(api_key, env_lookup)?),
|
||||
};
|
||||
Ok(ValidatedEnvironment { headers, auth })
|
||||
|
|
@ -216,6 +215,8 @@ mod tests {
|
|||
use rstest::rstest;
|
||||
use serde_json::json;
|
||||
|
||||
use litellm_auth::CredentialPlacement;
|
||||
|
||||
use super::*;
|
||||
use crate::anthropic::common_utils::AnthropicModelCapabilities;
|
||||
|
||||
|
|
|
|||
|
|
@ -7,6 +7,7 @@
|
|||
|
||||
use litellm_auth::{AuthServices, CredentialPlacement, SecretValue, TokenProviderHandle};
|
||||
use litellm_auth_aws::{AwsCredentialSource, SigV4Signer};
|
||||
use litellm_http::request::without_headers;
|
||||
|
||||
pub type Headers = Vec<(String, String)>;
|
||||
|
||||
|
|
@ -107,9 +108,8 @@ fn with_credential(headers: Headers, placement: CredentialPlacement, credential:
|
|||
CredentialPlacement::Bearer => format!("Bearer {credential}"),
|
||||
CredentialPlacement::Header(_) => credential.to_string(),
|
||||
};
|
||||
headers
|
||||
without_headers(headers, &[name])
|
||||
.into_iter()
|
||||
.filter(|(header, _)| !header.eq_ignore_ascii_case(name))
|
||||
.chain([(name.to_ascii_lowercase(), value)])
|
||||
.collect()
|
||||
}
|
||||
|
|
|
|||
|
|
@ -7,6 +7,7 @@ pub mod cohere;
|
|||
mod error;
|
||||
pub mod mistral;
|
||||
pub mod openai;
|
||||
pub mod openai_like;
|
||||
pub mod reducto;
|
||||
pub mod vertex_ai;
|
||||
|
||||
|
|
|
|||
1
litellm-rust/crates/llms/src/openai_like/chat/mod.rs
Normal file
1
litellm-rust/crates/llms/src/openai_like/chat/mod.rs
Normal file
|
|
@ -0,0 +1 @@
|
|||
pub mod transformation;
|
||||
270
litellm-rust/crates/llms/src/openai_like/chat/transformation.rs
Normal file
270
litellm-rust/crates/llms/src/openai_like/chat/transformation.rs
Normal file
|
|
@ -0,0 +1,270 @@
|
|||
//! `litellm/llms/openai_like/chat/transformation.py`: the chat config every
|
||||
//! OpenAI-compatible endpoint shares. The body is already OpenAI-shaped, so
|
||||
//! parameters pass through verbatim; the port keeps Python's two deviations,
|
||||
//! the `max_completion_tokens` -> `max_tokens` rename and the usage
|
||||
//! `*_tokens` null-to-zero sanitize.
|
||||
|
||||
use litellm_auth::{CredentialPlacement, SecretValue};
|
||||
use litellm_core_utils::core_helpers::unix_now;
|
||||
use litellm_types::{
|
||||
llms::openai::ChatMessage,
|
||||
utils::{ChatCompletionsChoice, ChatCompletionsChoiceMessage, ChatCompletionsResponse},
|
||||
};
|
||||
use serde_json::{Map, Value, json};
|
||||
|
||||
use crate::{
|
||||
Error,
|
||||
base_llm::{
|
||||
auth::AuthScheme,
|
||||
chat::transformation::{
|
||||
BaseConfig, Headers, ProviderChatRequestData, ProviderChatResponseData,
|
||||
ValidatedEnvironment,
|
||||
},
|
||||
},
|
||||
openai_like::common_utils::{complete_openai_like_url, openai_compatible_provider_info},
|
||||
};
|
||||
|
||||
/// OpenAI parameter names the Rust path can place verbatim in the request body.
|
||||
/// Tool parameters are absent on purpose: the message gate already declines
|
||||
/// tool-call content, and a `tools` request that did get through would produce
|
||||
/// a tool-call response this port cannot normalize yet, so it declines before
|
||||
/// the call instead of after it.
|
||||
const SUPPORTED_PARAMS: &[(&str, &str)] = &[
|
||||
("frequency_penalty", "frequency_penalty"),
|
||||
("logit_bias", "logit_bias"),
|
||||
("logprobs", "logprobs"),
|
||||
("top_logprobs", "top_logprobs"),
|
||||
("max_tokens", "max_tokens"),
|
||||
("max_completion_tokens", "max_completion_tokens"),
|
||||
("modalities", "modalities"),
|
||||
("prediction", "prediction"),
|
||||
("n", "n"),
|
||||
("presence_penalty", "presence_penalty"),
|
||||
("seed", "seed"),
|
||||
("stop", "stop"),
|
||||
("stream_options", "stream_options"),
|
||||
("temperature", "temperature"),
|
||||
("top_p", "top_p"),
|
||||
("audio", "audio"),
|
||||
("web_search_options", "web_search_options"),
|
||||
("service_tier", "service_tier"),
|
||||
("safety_identifier", "safety_identifier"),
|
||||
("prompt_cache_key", "prompt_cache_key"),
|
||||
("prompt_cache_retention", "prompt_cache_retention"),
|
||||
("store", "store"),
|
||||
("response_format", "response_format"),
|
||||
];
|
||||
|
||||
/// Call configuration the caller may pass that never enters the request body.
|
||||
const CONFIG_PARAMS: &[&str] = &["custom_endpoint", "extra_headers", "max_retries"];
|
||||
|
||||
pub struct OpenAILikeChatConfig;
|
||||
|
||||
pub const OPENAI_LIKE_CHAT_COMPLETIONS_CONFIG: OpenAILikeChatConfig = OpenAILikeChatConfig;
|
||||
|
||||
impl BaseConfig for OpenAILikeChatConfig {
|
||||
fn supported_openai_param_mappings(&self) -> &'static [(&'static str, &'static str)] {
|
||||
SUPPORTED_PARAMS
|
||||
}
|
||||
|
||||
fn get_complete_url(
|
||||
&self,
|
||||
api_base: Option<&str>,
|
||||
_model: &str,
|
||||
optional_params: &Map<String, Value>,
|
||||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> Result<String, Error> {
|
||||
let custom_endpoint = optional_params
|
||||
.get("custom_endpoint")
|
||||
.and_then(Value::as_bool)
|
||||
.unwrap_or(false);
|
||||
complete_openai_like_url(api_base, custom_endpoint, env_lookup)
|
||||
}
|
||||
|
||||
fn transform_request(
|
||||
&self,
|
||||
model: &str,
|
||||
messages: Vec<ChatMessage>,
|
||||
optional_params: Map<String, Value>,
|
||||
) -> Result<ProviderChatRequestData, Error> {
|
||||
let mut params = Map::from_iter(
|
||||
optional_params
|
||||
.into_iter()
|
||||
.filter(|(key, _)| !CONFIG_PARAMS.contains(&key.as_str())),
|
||||
);
|
||||
// Most OpenAI-compatible endpoints take `max_tokens`, not
|
||||
// `max_completion_tokens`, so Python's `map_openai_params` renames it
|
||||
// and lets it overwrite a `max_tokens` the caller also sent.
|
||||
if let Some(limit) = params.remove("max_completion_tokens") {
|
||||
params.insert("max_tokens".to_string(), limit);
|
||||
}
|
||||
let body = Map::from_iter(
|
||||
[
|
||||
("model".to_string(), json!(model)),
|
||||
("messages".to_string(), json!(messages)),
|
||||
]
|
||||
.into_iter()
|
||||
.chain(params),
|
||||
);
|
||||
Ok(ProviderChatRequestData {
|
||||
body: Value::Object(body),
|
||||
stream_shape: Default::default(),
|
||||
})
|
||||
}
|
||||
|
||||
fn transform_response(
|
||||
&self,
|
||||
model: &str,
|
||||
response: ProviderChatResponseData,
|
||||
) -> Result<ChatCompletionsResponse, Error> {
|
||||
let mut body = response.body;
|
||||
sanitize_usage(&mut body);
|
||||
let body = body
|
||||
.as_object()
|
||||
.ok_or_else(|| Error::InvalidResponse("chat response is not an object".into()))?;
|
||||
|
||||
let choices = body
|
||||
.get("choices")
|
||||
.and_then(Value::as_array)
|
||||
.ok_or(Error::MissingField("choices"))?
|
||||
.iter()
|
||||
.enumerate()
|
||||
.map(|(position, choice)| normalize_choice(position, choice))
|
||||
.collect::<Result<Vec<_>, _>>()?;
|
||||
|
||||
let usage = body.get("usage").and_then(Value::as_object);
|
||||
let field = |name: &str| {
|
||||
usage
|
||||
.and_then(|usage| usage.get(name))
|
||||
.and_then(Value::as_u64)
|
||||
.unwrap_or(0)
|
||||
};
|
||||
let details = usage.and_then(|usage| usage.get("prompt_tokens_details"));
|
||||
|
||||
Ok(ChatCompletionsResponse {
|
||||
created: body
|
||||
.get("created")
|
||||
.and_then(Value::as_u64)
|
||||
.unwrap_or_else(unix_now),
|
||||
model: body
|
||||
.get("model")
|
||||
.and_then(Value::as_str)
|
||||
.unwrap_or(model)
|
||||
.to_string(),
|
||||
choices,
|
||||
usage: litellm_types::utils::ChatCompletionsUsage {
|
||||
prompt_tokens: field("prompt_tokens"),
|
||||
completion_tokens: field("completion_tokens"),
|
||||
total_tokens: field("total_tokens"),
|
||||
prompt_tokens_details: litellm_types::utils::PromptTokensDetails {
|
||||
cached_tokens: details
|
||||
.and_then(|d| d.get("cached_tokens"))
|
||||
.and_then(Value::as_u64)
|
||||
.unwrap_or(0),
|
||||
cache_creation_tokens: details
|
||||
.and_then(|d| d.get("cache_creation_tokens"))
|
||||
.and_then(Value::as_u64)
|
||||
.unwrap_or(0),
|
||||
text_tokens: details
|
||||
.and_then(|d| d.get("text_tokens"))
|
||||
.and_then(Value::as_u64)
|
||||
.unwrap_or(0),
|
||||
},
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
/// `OpenAILikeBase._validate_environment`: a forwarded `authorization` is
|
||||
/// the whole credential, and any other call authenticates with the
|
||||
/// resolved key as a bearer. The key resolves to `""` when neither the
|
||||
/// deployment nor `OPENAI_LIKE_API_KEY` sets one, because vllm-compatible
|
||||
/// endpoints take no key; Python still sends `Bearer ` in that case.
|
||||
fn validate_environment(
|
||||
&self,
|
||||
headers: Headers,
|
||||
api_key: Option<&str>,
|
||||
_model: &str,
|
||||
_optional_params: &Map<String, Value>,
|
||||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> Result<ValidatedEnvironment, Error> {
|
||||
if headers
|
||||
.iter()
|
||||
.any(|(name, _)| name.eq_ignore_ascii_case("authorization"))
|
||||
{
|
||||
return Ok(ValidatedEnvironment {
|
||||
headers,
|
||||
auth: AuthScheme::Forwarded,
|
||||
});
|
||||
}
|
||||
let (_, key) = openai_compatible_provider_info(None, api_key, env_lookup);
|
||||
Ok(ValidatedEnvironment {
|
||||
headers,
|
||||
auth: AuthScheme::Credential {
|
||||
placement: CredentialPlacement::Bearer,
|
||||
secret: SecretValue::new(key.unwrap_or_default()),
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
fn config_params(&self) -> &'static [&'static str] {
|
||||
CONFIG_PARAMS
|
||||
}
|
||||
}
|
||||
|
||||
/// `OpenAILikeChatConfig._sanitize_usage_obj`: a provider that reports a null
|
||||
/// `*_tokens` entry breaks OpenAI clients, so nulls become 0. Python scrubs
|
||||
/// every top-level usage key ending in `_tokens`.
|
||||
fn sanitize_usage(body: &mut Value) {
|
||||
if let Some(usage) = body.get_mut("usage").and_then(Value::as_object_mut) {
|
||||
for (key, value) in usage.iter_mut() {
|
||||
if key.ends_with("_tokens") && value.is_null() {
|
||||
*value = json!(0);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn normalize_choice(position: usize, choice: &Value) -> Result<ChatCompletionsChoice, Error> {
|
||||
let message = choice
|
||||
.get("message")
|
||||
.and_then(Value::as_object)
|
||||
.ok_or(Error::MissingField("message"))?;
|
||||
if message
|
||||
.get("tool_calls")
|
||||
.and_then(Value::as_array)
|
||||
.is_some_and(|calls| !calls.is_empty())
|
||||
{
|
||||
// Python rewrites the lone tool call into content only under
|
||||
// `json_mode`, a request flag `transform_response` cannot see, and the
|
||||
// normalized type cannot carry tool calls at all. Declining is
|
||||
// terminal at this point, but passing back an empty assistant turn
|
||||
// would fabricate the reply.
|
||||
return Err(Error::Unsupported("tool call response"));
|
||||
}
|
||||
if message.get("refusal").is_some_and(|value| !value.is_null()) {
|
||||
return Err(Error::Unsupported("refusal response"));
|
||||
}
|
||||
let content = message.get("content");
|
||||
if content.is_some_and(|value| !value.is_null() && !value.is_string()) {
|
||||
return Err(Error::Unsupported("non-text response content"));
|
||||
}
|
||||
Ok(ChatCompletionsChoice {
|
||||
index: choice
|
||||
.get("index")
|
||||
.and_then(Value::as_u64)
|
||||
.unwrap_or(position as u64),
|
||||
message: ChatCompletionsChoiceMessage {
|
||||
role: message
|
||||
.get("role")
|
||||
.and_then(Value::as_str)
|
||||
.unwrap_or("assistant")
|
||||
.to_string(),
|
||||
content: content.and_then(Value::as_str).map(str::to_string),
|
||||
},
|
||||
finish_reason: choice
|
||||
.get("finish_reason")
|
||||
.and_then(Value::as_str)
|
||||
.unwrap_or("")
|
||||
.to_string(),
|
||||
})
|
||||
}
|
||||
58
litellm-rust/crates/llms/src/openai_like/common_utils.rs
Normal file
58
litellm-rust/crates/llms/src/openai_like/common_utils.rs
Normal file
|
|
@ -0,0 +1,58 @@
|
|||
//! Shared OpenAI-like credential and endpoint resolution, mirroring
|
||||
//! `litellm/llms/openai_like/common_utils.py`.
|
||||
|
||||
use crate::Error;
|
||||
|
||||
/// `OpenAILikeChatConfig._get_openai_compatible_provider_info`: the deployment's
|
||||
/// `api_base` wins over `OPENAI_LIKE_API_BASE`, and the deployment key over
|
||||
/// `OPENAI_LIKE_API_KEY`, with an empty key allowed because vllm-compatible
|
||||
/// endpoints do not require one.
|
||||
pub fn openai_compatible_provider_info(
|
||||
api_base: Option<&str>,
|
||||
api_key: Option<&str>,
|
||||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> (Option<String>, Option<String>) {
|
||||
let api_base = api_base
|
||||
.map(str::to_string)
|
||||
.or_else(|| env_lookup("OPENAI_LIKE_API_BASE"));
|
||||
let api_key = api_key
|
||||
.map(str::to_string)
|
||||
.or_else(|| env_lookup("OPENAI_LIKE_API_KEY"))
|
||||
.or(Some(String::new()));
|
||||
(api_base, api_key)
|
||||
}
|
||||
|
||||
/// `OpenAILikeBase._validate_environment` requires an api base and, when the
|
||||
/// caller gave no `custom_endpoint`, appends the route suffix. A caller-supplied
|
||||
/// `custom_endpoint` base is used as is.
|
||||
pub fn complete_openai_like_url(
|
||||
api_base: Option<&str>,
|
||||
custom_endpoint: bool,
|
||||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> Result<String, Error> {
|
||||
let (api_base, _) = openai_compatible_provider_info(api_base, None, env_lookup);
|
||||
let api_base = api_base.ok_or_else(|| {
|
||||
Error::InvalidRequest(
|
||||
"Missing API Base - A call is being made to LLM Provider but no api base is set either in the environment variables ({LLM_PROVIDER}_API_KEY) or via params"
|
||||
.to_string(),
|
||||
)
|
||||
})?;
|
||||
if custom_endpoint {
|
||||
return Ok(api_base);
|
||||
}
|
||||
Ok(format!(
|
||||
"{}/chat/completions",
|
||||
api_base.trim_end_matches('/')
|
||||
))
|
||||
}
|
||||
|
||||
/// The api key the call resolves to. `None` means neither the deployment nor the
|
||||
/// environment supplied one, which is valid for endpoints that take no key.
|
||||
pub fn resolve_openai_like_api_key(
|
||||
api_key: Option<&str>,
|
||||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> Option<String> {
|
||||
openai_compatible_provider_info(None, api_key, env_lookup)
|
||||
.1
|
||||
.filter(|key| !key.is_empty())
|
||||
}
|
||||
2
litellm-rust/crates/llms/src/openai_like/mod.rs
Normal file
2
litellm-rust/crates/llms/src/openai_like/mod.rs
Normal file
|
|
@ -0,0 +1,2 @@
|
|||
pub mod chat;
|
||||
pub mod common_utils;
|
||||
|
|
@ -0,0 +1,343 @@
|
|||
use litellm_llms::{
|
||||
Error,
|
||||
base_llm::{
|
||||
auth::AuthScheme,
|
||||
chat::transformation::{BaseConfig, ProviderChatResponseData, Unsupported},
|
||||
},
|
||||
openai_like::chat::transformation::OPENAI_LIKE_CHAT_COMPLETIONS_CONFIG,
|
||||
};
|
||||
use litellm_types::{llms::openai::ChatMessage, utils::ChatCompletionsResponse};
|
||||
use rstest::rstest;
|
||||
use serde_json::{Map, Value, json};
|
||||
|
||||
fn messages(value: Value) -> Vec<ChatMessage> {
|
||||
serde_json::from_value(value).expect("valid messages")
|
||||
}
|
||||
|
||||
fn params(value: Value) -> Map<String, Value> {
|
||||
match value {
|
||||
Value::Object(map) => map,
|
||||
other => panic!("params must be an object, got {other}"),
|
||||
}
|
||||
}
|
||||
|
||||
fn no_env(_: &str) -> Option<String> {
|
||||
None
|
||||
}
|
||||
|
||||
fn env_with<'a>(name: &'a str, value: &'a str) -> impl Fn(&str) -> Option<String> + 'a {
|
||||
move |key| (key == name).then(|| value.to_string())
|
||||
}
|
||||
|
||||
fn transform(model: &str, msgs: Value, opts: Value) -> Value {
|
||||
OPENAI_LIKE_CHAT_COMPLETIONS_CONFIG
|
||||
.transform_request(model, messages(msgs), params(opts))
|
||||
.expect("request transforms")
|
||||
.body
|
||||
}
|
||||
|
||||
fn transform_response(body: Value) -> Result<ChatCompletionsResponse, Error> {
|
||||
OPENAI_LIKE_CHAT_COMPLETIONS_CONFIG
|
||||
.transform_response("some-model", ProviderChatResponseData { body })
|
||||
}
|
||||
|
||||
fn reason(msgs: Value, opts: Value) -> Option<Unsupported> {
|
||||
OPENAI_LIKE_CHAT_COMPLETIONS_CONFIG.unsupported_reason(&messages(msgs), ¶ms(opts))
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn builds_the_openai_shaped_body() {
|
||||
let body = transform(
|
||||
"my-model",
|
||||
json!([
|
||||
{"role": "system", "content": "be terse"},
|
||||
{"role": "user", "content": "hi"},
|
||||
]),
|
||||
json!({"temperature": 0.5, "max_tokens": 8}),
|
||||
);
|
||||
assert_eq!(body["model"], json!("my-model"));
|
||||
assert_eq!(
|
||||
body["messages"],
|
||||
json!([
|
||||
{"role": "system", "content": "be terse"},
|
||||
{"role": "user", "content": "hi"},
|
||||
])
|
||||
);
|
||||
assert_eq!(body["temperature"], json!(0.5));
|
||||
assert_eq!(body["max_tokens"], json!(8));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn renames_max_completion_tokens_to_max_tokens() {
|
||||
// `OpenAILikeChatConfig.map_openai_params`: most OpenAI-compatible providers
|
||||
// support `max_tokens`, not `max_completion_tokens`.
|
||||
let body = transform(
|
||||
"my-model",
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({"max_completion_tokens": 12}),
|
||||
);
|
||||
assert_eq!(body["max_tokens"], json!(12));
|
||||
assert!(body.get("max_completion_tokens").is_none());
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn max_completion_tokens_wins_when_both_limits_are_sent() {
|
||||
// Python assigns `max_tokens = max_completion_tokens` after copying the
|
||||
// params, so the renamed value outranks a caller-supplied `max_tokens`.
|
||||
let body = transform(
|
||||
"my-model",
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({"max_tokens": 8, "max_completion_tokens": 12}),
|
||||
);
|
||||
assert_eq!(body["max_tokens"], json!(12));
|
||||
assert!(body.get("max_completion_tokens").is_none());
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn call_configuration_never_enters_the_body() {
|
||||
let body = transform(
|
||||
"my-model",
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({"custom_endpoint": true, "extra_headers": {"x": "y"}, "max_retries": 2}),
|
||||
);
|
||||
assert_eq!(
|
||||
body.as_object().unwrap().keys().collect::<Vec<_>>(),
|
||||
vec!["model", "messages"]
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::appends_the_chat_completions_suffix("https://vllm.example.com/v1", json!({}), "https://vllm.example.com/v1/chat/completions")]
|
||||
#[case::trims_a_trailing_slash("https://vllm.example.com/v1/", json!({}), "https://vllm.example.com/v1/chat/completions")]
|
||||
#[case::a_custom_endpoint_is_used_as_is("https://vllm.example.com/v1/chat/completions", json!({"custom_endpoint": true}), "https://vllm.example.com/v1/chat/completions")]
|
||||
fn complete_url(#[case] api_base: &str, #[case] opts: Value, #[case] expected: &str) {
|
||||
assert_eq!(
|
||||
OPENAI_LIKE_CHAT_COMPLETIONS_CONFIG
|
||||
.get_complete_url(Some(api_base), "my-model", ¶ms(opts), &no_env)
|
||||
.expect("url resolves"),
|
||||
expected
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn api_base_falls_back_to_the_environment() {
|
||||
assert_eq!(
|
||||
OPENAI_LIKE_CHAT_COMPLETIONS_CONFIG
|
||||
.get_complete_url(
|
||||
None,
|
||||
"my-model",
|
||||
¶ms(json!({})),
|
||||
&env_with("OPENAI_LIKE_API_BASE", "https://env.example.com/v1"),
|
||||
)
|
||||
.expect("url resolves"),
|
||||
"https://env.example.com/v1/chat/completions"
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn a_missing_api_base_is_an_error() {
|
||||
assert!(matches!(
|
||||
OPENAI_LIKE_CHAT_COMPLETIONS_CONFIG
|
||||
.get_complete_url(None, "my-model", ¶ms(json!({})), &no_env),
|
||||
Err(Error::InvalidRequest(message)) if message.starts_with("Missing API Base")
|
||||
));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn the_resolved_key_authenticates_as_a_bearer() {
|
||||
let validated = OPENAI_LIKE_CHAT_COMPLETIONS_CONFIG
|
||||
.validate_environment(
|
||||
vec![],
|
||||
Some("sk-test"),
|
||||
"my-model",
|
||||
¶ms(json!({})),
|
||||
&no_env,
|
||||
)
|
||||
.expect("validates");
|
||||
assert!(matches!(
|
||||
validated.auth,
|
||||
AuthScheme::Credential {
|
||||
placement: litellm_auth::CredentialPlacement::Bearer,
|
||||
..
|
||||
}
|
||||
));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn the_key_falls_back_to_the_environment() {
|
||||
let validated = OPENAI_LIKE_CHAT_COMPLETIONS_CONFIG
|
||||
.validate_environment(
|
||||
vec![],
|
||||
None,
|
||||
"my-model",
|
||||
¶ms(json!({})),
|
||||
&env_with("OPENAI_LIKE_API_KEY", "sk-env"),
|
||||
)
|
||||
.expect("validates");
|
||||
let AuthScheme::Credential { secret, .. } = validated.auth else {
|
||||
panic!("expected a bearer credential");
|
||||
};
|
||||
assert_eq!(secret.expose(), "sk-env");
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn a_forwarded_authorization_is_the_whole_credential() {
|
||||
// Python adds `Bearer <key>` only when the caller did not already send
|
||||
// `Authorization`, so the forwarded header wins over the deployment key.
|
||||
let headers = vec![("Authorization".to_string(), "Bearer caller".to_string())];
|
||||
let validated = OPENAI_LIKE_CHAT_COMPLETIONS_CONFIG
|
||||
.validate_environment(
|
||||
headers,
|
||||
Some("sk-test"),
|
||||
"my-model",
|
||||
¶ms(json!({})),
|
||||
&no_env,
|
||||
)
|
||||
.expect("validates");
|
||||
assert!(matches!(validated.auth, AuthScheme::Forwarded));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn keyless_calls_still_validate_for_endpoints_that_take_no_key() {
|
||||
// vllm-compatible endpoints require no api key; Python resolves `""` and
|
||||
// sends `Bearer `, so validation must not fail on the missing key.
|
||||
let validated = OPENAI_LIKE_CHAT_COMPLETIONS_CONFIG
|
||||
.validate_environment(vec![], None, "my-model", ¶ms(json!({})), &no_env)
|
||||
.expect("validates");
|
||||
let AuthScheme::Credential { secret, .. } = validated.auth else {
|
||||
panic!("expected a bearer credential");
|
||||
};
|
||||
assert_eq!(secret.expose(), "");
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn normalizes_an_openai_response() {
|
||||
let response = transform_response(json!({
|
||||
"created": 1_700_000_000,
|
||||
"model": "served-model-name",
|
||||
"choices": [{
|
||||
"index": 0,
|
||||
"message": {"role": "assistant", "content": "hello"},
|
||||
"finish_reason": "stop",
|
||||
}],
|
||||
"usage": {"prompt_tokens": 3, "completion_tokens": 5, "total_tokens": 8},
|
||||
}))
|
||||
.expect("response normalizes");
|
||||
assert_eq!(response.created, 1_700_000_000);
|
||||
assert_eq!(response.model, "served-model-name");
|
||||
assert_eq!(
|
||||
response.choices[0].message.content.as_deref(),
|
||||
Some("hello")
|
||||
);
|
||||
assert_eq!(response.choices[0].finish_reason, "stop");
|
||||
assert_eq!(response.usage.prompt_tokens, 3);
|
||||
assert_eq!(response.usage.completion_tokens, 5);
|
||||
assert_eq!(response.usage.total_tokens, 8);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn null_token_fields_in_usage_become_zero() {
|
||||
// `_sanitize_usage_obj`: providers that return null token values break
|
||||
// OpenAI clients, so the response is scrubbed at the source.
|
||||
let response = transform_response(json!({
|
||||
"model": "m",
|
||||
"choices": [{"message": {"role": "assistant", "content": "hi"}, "finish_reason": "stop"}],
|
||||
"usage": {"prompt_tokens": 3, "completion_tokens": null, "total_tokens": null},
|
||||
}))
|
||||
.expect("response normalizes");
|
||||
assert_eq!(response.usage.completion_tokens, 0);
|
||||
assert_eq!(response.usage.total_tokens, 0);
|
||||
assert_eq!(response.usage.prompt_tokens, 3);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn a_tool_call_response_declines_instead_of_dropping_the_calls() {
|
||||
// The `json_mode` rewrite needs a request flag the route does not carry, so
|
||||
// a tool-call answer falls back to Python rather than losing the calls.
|
||||
assert_eq!(
|
||||
transform_response(json!({
|
||||
"model": "m",
|
||||
"choices": [{
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"content": null,
|
||||
"tool_calls": [{
|
||||
"id": "call_1",
|
||||
"type": "function",
|
||||
"function": {"name": "f", "arguments": "{}"},
|
||||
}],
|
||||
},
|
||||
"finish_reason": "tool_calls",
|
||||
}],
|
||||
})),
|
||||
Err(Error::Unsupported("tool call response"))
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn a_refusal_declines_instead_of_returning_an_empty_reply() {
|
||||
assert_eq!(
|
||||
transform_response(json!({
|
||||
"model": "m",
|
||||
"choices": [{
|
||||
"message": {"role": "assistant", "content": null, "refusal": "cannot help"},
|
||||
"finish_reason": "stop",
|
||||
}],
|
||||
})),
|
||||
Err(Error::Unsupported("refusal response"))
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn a_non_text_response_content_declines() {
|
||||
assert_eq!(
|
||||
transform_response(json!({
|
||||
"model": "m",
|
||||
"choices": [{
|
||||
"message": {"role": "assistant", "content": [{"type": "text", "text": "hi"}]},
|
||||
"finish_reason": "stop",
|
||||
}],
|
||||
})),
|
||||
Err(Error::Unsupported("non-text response content"))
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::streaming(json!({"stream": true}), "streaming")]
|
||||
#[case::unrecognized_param(json!({"some_provider_knob": 1}), "unrecognized request parameter")]
|
||||
fn declines(#[case] opts: Value, #[case] expected: &'static str) {
|
||||
assert_eq!(
|
||||
reason(json!([{"role": "user", "content": "hi"}]), opts),
|
||||
Some(Unsupported(expected))
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn accepts_standard_openai_params() {
|
||||
assert_eq!(
|
||||
reason(
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({
|
||||
"temperature": 0.2,
|
||||
"top_p": 0.9,
|
||||
"max_tokens": 16,
|
||||
"response_format": {"type": "json_object"},
|
||||
"custom_endpoint": true,
|
||||
}),
|
||||
),
|
||||
None
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn tool_parameters_decline_before_the_call() {
|
||||
// A `tools` request would come back with tool calls this port cannot
|
||||
// normalize, so it declines at the gate instead of after the call.
|
||||
assert_eq!(
|
||||
reason(
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({"tools": [{"type": "function", "function": {"name": "f"}}]}),
|
||||
),
|
||||
Some(Unsupported("unrecognized request parameter"))
|
||||
);
|
||||
}
|
||||
|
|
@ -104,6 +104,9 @@ pub struct ModelInfo {
|
|||
/// Rate applied once the prompt exceeds the token threshold in the field name.
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub cache_read_input_token_cost_above_512k_tokens: Option<f64>,
|
||||
/// Balanced service-tier rate for the same-named base field.
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub cache_read_input_token_cost_balanced: Option<f64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub cache_read_input_token_cost_batches: Option<f64>,
|
||||
/// Flex service-tier rate for the same-named base field.
|
||||
|
|
@ -211,6 +214,9 @@ pub struct ModelInfo {
|
|||
/// Rate applied once the prompt exceeds the token threshold in the field name.
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub input_cost_per_token_above_512k_tokens: Option<f64>,
|
||||
/// Balanced service-tier rate for the same-named base field.
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub input_cost_per_token_balanced: Option<f64>,
|
||||
/// USD per prompt token via the provider's batch API.
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub input_cost_per_token_batches: Option<f64>,
|
||||
|
|
@ -357,6 +363,9 @@ pub struct ModelInfo {
|
|||
/// Rate applied once the prompt exceeds the token threshold in the field name.
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub output_cost_per_token_above_512k_tokens: Option<f64>,
|
||||
/// Balanced service-tier rate for the same-named base field.
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub output_cost_per_token_balanced: Option<f64>,
|
||||
/// USD per generated token via the provider's batch API.
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub output_cost_per_token_batches: Option<f64>,
|
||||
|
|
|
|||
|
|
@ -2,9 +2,8 @@ use std::convert::Infallible;
|
|||
|
||||
use bytes::Bytes;
|
||||
use litellm_core::messages::{
|
||||
Error,
|
||||
route::{Messages, MessagesCall, MessagesOutput, MessagesStreamHead, messages_body},
|
||||
types::MessagesShaping,
|
||||
Error, MessagesCall, MessagesShaping, messages_body,
|
||||
route::{Messages, MessagesOutput, MessagesStreamHead},
|
||||
};
|
||||
use litellm_host_python::{InvokeError, ProtocolHost, from_py, lookup, to_py};
|
||||
use litellm_http::transport::Error as TransportError;
|
||||
|
|
@ -76,6 +75,11 @@ fn native_error(py: Python<'_>, error: Error) -> PyResult<PyErr> {
|
|||
error.value(py).setattr(REQUEST_ERROR_MARKER, true)?;
|
||||
Ok(error)
|
||||
}
|
||||
Error::MissingField(field) => {
|
||||
let error = PyValueError::new_err(format!("missing required field: {field}"));
|
||||
error.value(py).setattr(REQUEST_ERROR_MARKER, true)?;
|
||||
Ok(error)
|
||||
}
|
||||
other => Ok(route_error_to_pyerr(other)),
|
||||
}
|
||||
}
|
||||
|
|
@ -314,6 +318,7 @@ mod tests {
|
|||
|
||||
#[rstest]
|
||||
#[case::rejected_request(Error::InvalidRequest("does not support top_k=5".into()), true)]
|
||||
#[case::missing_field(Error::MissingField("max_tokens"), true)]
|
||||
#[case::unresolvable_provider(Error::InvalidProvider("openai".into()), false)]
|
||||
#[case::upstream_failure(
|
||||
Error::Transport(TransportError::Http { status: 400, body: "bad".into() }),
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
use std::time::Duration;
|
||||
|
||||
use litellm_core::messages::types::MessagesShaping;
|
||||
use litellm_core::messages::MessagesShaping;
|
||||
|
||||
#[derive(Clone, Debug, Default)]
|
||||
pub struct Deployment {
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
use std::time::Duration;
|
||||
|
||||
use litellm_config::Config;
|
||||
use litellm_core::messages::types::MessagesShaping;
|
||||
use litellm_core::messages::MessagesShaping;
|
||||
use litellm_router::{Deployment, Router};
|
||||
use rstest::rstest;
|
||||
|
||||
|
|
|
|||
|
|
@ -48,6 +48,7 @@ pub struct CyberArkSecretManager {
|
|||
token: Cache<(), SecretValue>,
|
||||
secrets: SecretCache<String, SecretValue>,
|
||||
authentication_lock: Arc<tokio::sync::Mutex<()>>,
|
||||
policy_load_lock: Arc<tokio::sync::Mutex<()>>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||||
|
|
|
|||
|
|
@ -26,6 +26,7 @@ impl CyberArkSecretManager {
|
|||
token,
|
||||
secrets,
|
||||
authentication_lock: Arc::new(tokio::sync::Mutex::new(())),
|
||||
policy_load_lock: Arc::new(tokio::sync::Mutex::new(())),
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -1,5 +1,8 @@
|
|||
use super::*;
|
||||
|
||||
const POLICY_LOAD_ATTEMPTS: u32 = 5;
|
||||
const POLICY_LOAD_RETRY_DELAY: std::time::Duration = std::time::Duration::from_millis(200);
|
||||
|
||||
impl CyberArkSecretManager {
|
||||
pub async fn async_write_secret(
|
||||
&self,
|
||||
|
|
@ -105,37 +108,43 @@ impl CyberArkSecretManager {
|
|||
"- !variable {}\n",
|
||||
serde_json::to_string(name).expect("serializing a string cannot fail")
|
||||
);
|
||||
let response = with_timeout(
|
||||
self.client
|
||||
.post(policy_url)
|
||||
.header("Authorization", authorization)
|
||||
.header("Content-Type", "application/x-yaml")
|
||||
.body(body),
|
||||
context,
|
||||
)
|
||||
.send()
|
||||
.await;
|
||||
match response {
|
||||
Ok(response) if response.status().is_success() => {}
|
||||
Ok(response)
|
||||
if matches!(
|
||||
response.status(),
|
||||
reqwest::StatusCode::CONFLICT | reqwest::StatusCode::UNPROCESSABLE_ENTITY
|
||||
) =>
|
||||
{
|
||||
litellm_tracing::debug!(
|
||||
"CyberArk variable policy already exists or conflicts: {}",
|
||||
response.status()
|
||||
);
|
||||
}
|
||||
Ok(response) => {
|
||||
litellm_tracing::warn!(
|
||||
"Could not ensure CyberArk variable exists: {}",
|
||||
response.status()
|
||||
);
|
||||
}
|
||||
Err(error) => {
|
||||
litellm_tracing::warn!("Error ensuring CyberArk variable exists: {error}");
|
||||
let _policy_load = self.policy_load_lock.lock().await;
|
||||
for attempt in 0..POLICY_LOAD_ATTEMPTS {
|
||||
let response = with_timeout(
|
||||
self.client
|
||||
.post(policy_url.clone())
|
||||
.header("Authorization", authorization.clone())
|
||||
.header("Content-Type", "application/x-yaml")
|
||||
.body(body.clone()),
|
||||
context,
|
||||
)
|
||||
.send()
|
||||
.await;
|
||||
match response {
|
||||
Ok(response)
|
||||
if response.status() == reqwest::StatusCode::CONFLICT
|
||||
&& attempt + 1 < POLICY_LOAD_ATTEMPTS =>
|
||||
{
|
||||
tokio::time::sleep(POLICY_LOAD_RETRY_DELAY * 2_u32.pow(attempt)).await;
|
||||
}
|
||||
Ok(response) if response.status().is_success() => return,
|
||||
Ok(response) if response.status() == reqwest::StatusCode::UNPROCESSABLE_ENTITY => {
|
||||
litellm_tracing::debug!(
|
||||
"CyberArk variable policy was rejected as unprocessable"
|
||||
);
|
||||
return;
|
||||
}
|
||||
Ok(response) => {
|
||||
litellm_tracing::warn!(
|
||||
"Could not ensure CyberArk variable exists: {}",
|
||||
response.status()
|
||||
);
|
||||
return;
|
||||
}
|
||||
Err(error) => {
|
||||
litellm_tracing::warn!("Error ensuring CyberArk variable exists: {error}");
|
||||
return;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -6,7 +6,7 @@ async fn rejected_write_token_is_reauthenticated_once() {
|
|||
let server = MockServer::start().await;
|
||||
mount_auth(&server, 2).await;
|
||||
Mock::given(path("/policies/acct/policy/root"))
|
||||
.respond_with(ResponseTemplate::new(409))
|
||||
.respond_with(ResponseTemplate::new(201))
|
||||
.expect(1)
|
||||
.mount(&server)
|
||||
.await;
|
||||
|
|
@ -45,7 +45,6 @@ async fn rejected_write_token_is_reauthenticated_once() {
|
|||
|
||||
#[rstest]
|
||||
#[case::created(201)]
|
||||
#[case::already_exists(409)]
|
||||
#[case::unprocessable(422)]
|
||||
#[case::server_error(500)]
|
||||
#[tokio::test]
|
||||
|
|
@ -81,13 +80,81 @@ async fn writes_tolerate_policy_status_and_cache_value(#[case] policy_status: u1
|
|||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn policy_load_conflict_is_retried_before_the_value_write() {
|
||||
let server = MockServer::start().await;
|
||||
mount_auth(&server, 1).await;
|
||||
let policy_loads = Arc::new(AtomicUsize::new(0));
|
||||
let policy_loads_for_response = Arc::clone(&policy_loads);
|
||||
Mock::given(path("/policies/acct/policy/root"))
|
||||
.respond_with(move |_: &Request| {
|
||||
if policy_loads_for_response.fetch_add(1, Ordering::SeqCst) < 2 {
|
||||
ResponseTemplate::new(409)
|
||||
} else {
|
||||
ResponseTemplate::new(201)
|
||||
}
|
||||
})
|
||||
.expect(3)
|
||||
.mount(&server)
|
||||
.await;
|
||||
let policy_loads_at_value_write = Arc::clone(&policy_loads);
|
||||
Mock::given(method("POST"))
|
||||
.and(path("/secrets/acct/variable/key"))
|
||||
.respond_with(move |_: &Request| {
|
||||
if policy_loads_at_value_write.load(Ordering::SeqCst) == 3 {
|
||||
ResponseTemplate::new(201)
|
||||
} else {
|
||||
ResponseTemplate::new(404)
|
||||
}
|
||||
})
|
||||
.expect(1)
|
||||
.mount(&server)
|
||||
.await;
|
||||
let manager = manager(&server, Duration::from_secs(60));
|
||||
|
||||
manager
|
||||
.async_write_secret("key", &SecretValue::new("v"), None)
|
||||
.await
|
||||
.unwrap();
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn concurrent_writes_load_policy_one_at_a_time() {
|
||||
let server = MockServer::start().await;
|
||||
mount_auth(&server, 1).await;
|
||||
Mock::given(path("/policies/acct/policy/root"))
|
||||
.respond_with(ResponseTemplate::new(201).set_delay(Duration::from_millis(100)))
|
||||
.expect(4)
|
||||
.mount(&server)
|
||||
.await;
|
||||
Mock::given(method("POST"))
|
||||
.respond_with(ResponseTemplate::new(201))
|
||||
.mount(&server)
|
||||
.await;
|
||||
let manager = manager(&server, Duration::from_secs(60));
|
||||
let started = std::time::Instant::now();
|
||||
|
||||
let value = SecretValue::new("v");
|
||||
let results = tokio::join!(
|
||||
manager.async_write_secret("key-0", &value, None),
|
||||
manager.async_write_secret("key-1", &value, None),
|
||||
manager.async_write_secret("key-2", &value, None),
|
||||
manager.async_write_secret("key-3", &value, None),
|
||||
);
|
||||
|
||||
assert!(results.0.is_ok() && results.1.is_ok() && results.2.is_ok() && results.3.is_ok());
|
||||
assert!(started.elapsed() >= Duration::from_millis(400));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn failed_value_write_is_not_cached() {
|
||||
let server = MockServer::start().await;
|
||||
mount_auth(&server, 1).await;
|
||||
Mock::given(path("/policies/acct/policy/root"))
|
||||
.respond_with(ResponseTemplate::new(409))
|
||||
.respond_with(ResponseTemplate::new(201))
|
||||
.mount(&server)
|
||||
.await;
|
||||
Mock::given(path("/secrets/acct/variable/key"))
|
||||
|
|
|
|||
|
|
@ -6,8 +6,8 @@ use litellm_http::{HttpClientConfig, HttpClientPool};
|
|||
use crate::{Error, KeyManagementSettings, KeyManagementSystem, SecretManager};
|
||||
|
||||
pub async fn load_native_manager(
|
||||
pool: &HttpClientPool,
|
||||
config: &HttpClientConfig,
|
||||
_pool: &HttpClientPool,
|
||||
_config: &HttpClientConfig,
|
||||
system: KeyManagementSystem,
|
||||
settings: KeyManagementSettings,
|
||||
environment: Arc<dyn Lookup + Send + Sync>,
|
||||
|
|
@ -33,14 +33,14 @@ pub async fn load_native_manager(
|
|||
#[cfg(feature = "azure")]
|
||||
(KeyManagementSystem::AzureKeyVault, _, environment, _) => Ok(
|
||||
SecretManager::AzureKeyVault(crate::azure::AzureKeyVault::new(
|
||||
pool.client(config, litellm_http::ClientVariant::Provider)?,
|
||||
_pool.client(_config, litellm_http::ClientVariant::Provider)?,
|
||||
environment,
|
||||
)?),
|
||||
),
|
||||
#[cfg(feature = "google")]
|
||||
(KeyManagementSystem::GoogleSecretManager, _, environment, enterprise_enabled) => Ok(
|
||||
SecretManager::GoogleSecretManager(crate::google::GoogleSecretManager::new(
|
||||
pool.client(config, litellm_http::ClientVariant::Provider)?,
|
||||
_pool.client(_config, litellm_http::ClientVariant::Provider)?,
|
||||
environment,
|
||||
enterprise_enabled,
|
||||
)?),
|
||||
|
|
@ -61,8 +61,8 @@ pub async fn load_native_manager(
|
|||
#[cfg(feature = "cyberark")]
|
||||
(KeyManagementSystem::Cyberark, _, environment, enterprise_enabled) => Ok(
|
||||
SecretManager::Cyberark(crate::cyberark::CyberArkSecretManager::new(
|
||||
pool,
|
||||
config,
|
||||
_pool,
|
||||
_config,
|
||||
environment,
|
||||
enterprise_enabled,
|
||||
)?),
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ license.workspace = true
|
|||
repository.workspace = true
|
||||
|
||||
[dependencies]
|
||||
base64.workspace = true
|
||||
fancy-regex.workspace = true
|
||||
percent-encoding.workspace = true
|
||||
serde_json.workspace = true
|
||||
|
|
|
|||
|
|
@ -5,6 +5,7 @@ use std::{
|
|||
pin::pin,
|
||||
};
|
||||
|
||||
use base64::{Engine, engine::general_purpose::STANDARD};
|
||||
use serde_json::{Map, Value};
|
||||
use tracing::{
|
||||
Dispatch, Event, Subscriber,
|
||||
|
|
@ -20,6 +21,31 @@ pub use processing::{DiagnosticInput, DiagnosticOutput, Policy, Processor};
|
|||
pub use redaction::{REDACTED, SecretRedactor};
|
||||
pub use tracing::{Level, Metadata, debug, error, info, trace, warn};
|
||||
|
||||
pub struct ByteChunk<'a>(&'a [u8]);
|
||||
|
||||
impl<'a> ByteChunk<'a> {
|
||||
pub fn new(data: &'a [u8]) -> Self {
|
||||
Self(data)
|
||||
}
|
||||
|
||||
pub fn encoding(&self) -> &'static str {
|
||||
if std::str::from_utf8(self.0).is_ok() {
|
||||
"utf8"
|
||||
} else {
|
||||
"base64"
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl fmt::Display for ByteChunk<'_> {
|
||||
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
match std::str::from_utf8(self.0) {
|
||||
Ok(text) => formatter.write_str(text),
|
||||
Err(_) => formatter.write_str(&STANDARD.encode(self.0)),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub trait Sink: Send + Sync + 'static {
|
||||
fn enabled(&self, metadata: &Metadata<'_>) -> bool;
|
||||
fn emit(&self, record: &Record);
|
||||
|
|
@ -44,6 +70,10 @@ impl Logger {
|
|||
}
|
||||
}
|
||||
|
||||
pub fn install_global(&self) -> Result<(), tracing::dispatcher::SetGlobalDefaultError> {
|
||||
tracing::dispatcher::set_global_default(self.dispatch.clone())
|
||||
}
|
||||
|
||||
pub fn scope<T>(&self, operation: impl FnOnce() -> T) -> T {
|
||||
if EMITTING.get() {
|
||||
return operation();
|
||||
|
|
|
|||
|
|
@ -4,7 +4,9 @@ use std::sync::{
|
|||
mpsc,
|
||||
};
|
||||
|
||||
use litellm_tracing::{Level, Logger, Metadata, Record, Sink, info, warn};
|
||||
use base64::{Engine, engine::general_purpose::STANDARD};
|
||||
use litellm_tracing::{ByteChunk, Level, Logger, Metadata, Record, Sink, info, warn};
|
||||
use rstest::rstest;
|
||||
use serde_json::{Value, json};
|
||||
|
||||
struct Output {
|
||||
|
|
@ -120,3 +122,18 @@ fn nested_scopes_restore_the_previous_sink() {
|
|||
["inside"]
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::utf8(b"event: message_stop\n\n", "utf8")]
|
||||
#[case::binary(&[0xff, 0x00, 0x80], "base64")]
|
||||
fn byte_chunk_logging_preserves_exact_bytes(#[case] bytes: &[u8], #[case] encoding: &str) {
|
||||
let chunk = ByteChunk::new(bytes);
|
||||
assert_eq!(chunk.encoding(), encoding);
|
||||
let text = chunk.to_string();
|
||||
let recovered = match encoding {
|
||||
"utf8" => text.into_bytes(),
|
||||
"base64" => STANDARD.decode(text).unwrap(),
|
||||
_ => unreachable!(),
|
||||
};
|
||||
assert_eq!(recovered, bytes);
|
||||
}
|
||||
|
|
|
|||
240
litellm-rust/crates/types/src/llms/anthropic.rs
Normal file
240
litellm-rust/crates/types/src/llms/anthropic.rs
Normal file
|
|
@ -0,0 +1,240 @@
|
|||
use std::{
|
||||
cmp::Ordering,
|
||||
collections::BTreeSet,
|
||||
convert::Infallible,
|
||||
fmt,
|
||||
hash::{Hash, Hasher},
|
||||
str::FromStr,
|
||||
};
|
||||
|
||||
/// One value of the `anthropic-beta` header. Equality, ordering and hashing follow the wire
|
||||
/// string, so a value parsed from a caller's header never disagrees with the matching variant.
|
||||
#[derive(Clone, Debug, strum::AsRefStr, strum::Display, strum::EnumString)]
|
||||
pub enum AnthropicBeta {
|
||||
#[strum(serialize = "oauth-2025-04-20")]
|
||||
Oauth20250420,
|
||||
#[strum(serialize = "web-fetch-2025-09-10")]
|
||||
WebFetch20250910,
|
||||
#[strum(serialize = "web-search-2025-03-05")]
|
||||
WebSearch20250305,
|
||||
#[strum(serialize = "context-management-2025-06-27")]
|
||||
ContextManagement20250627,
|
||||
#[strum(serialize = "compact-2026-01-12")]
|
||||
Compact20260112,
|
||||
#[strum(serialize = "compact-2026-09-04")]
|
||||
Compact20260904,
|
||||
#[strum(serialize = "structured-outputs-2025-11-13")]
|
||||
StructuredOutputs20251113,
|
||||
#[strum(serialize = "advanced-tool-use-2025-11-20")]
|
||||
AdvancedToolUse20251120,
|
||||
#[strum(serialize = "fast-mode-2026-02-01")]
|
||||
FastMode20260201,
|
||||
#[strum(serialize = "advisor-tool-2026-03-01")]
|
||||
AdvisorTool20260301,
|
||||
#[strum(serialize = "per-turn-control-2026-07-01")]
|
||||
PerTurnControl20260701,
|
||||
#[strum(serialize = "dangerous-tool-use-2026-09-03")]
|
||||
DangerousToolUse20260903,
|
||||
#[strum(default, transparent)]
|
||||
Other(String),
|
||||
}
|
||||
|
||||
impl AnthropicBeta {
|
||||
pub const KNOWN: [Self; 12] = [
|
||||
Self::Oauth20250420,
|
||||
Self::WebFetch20250910,
|
||||
Self::WebSearch20250305,
|
||||
Self::ContextManagement20250627,
|
||||
Self::Compact20260112,
|
||||
Self::Compact20260904,
|
||||
Self::StructuredOutputs20251113,
|
||||
Self::AdvancedToolUse20251120,
|
||||
Self::FastMode20260201,
|
||||
Self::AdvisorTool20260301,
|
||||
Self::PerTurnControl20260701,
|
||||
Self::DangerousToolUse20260903,
|
||||
];
|
||||
|
||||
pub fn as_str(&self) -> &str {
|
||||
self.as_ref()
|
||||
}
|
||||
}
|
||||
|
||||
impl PartialEq for AnthropicBeta {
|
||||
fn eq(&self, other: &Self) -> bool {
|
||||
self.as_str() == other.as_str()
|
||||
}
|
||||
}
|
||||
|
||||
impl Eq for AnthropicBeta {}
|
||||
|
||||
impl Hash for AnthropicBeta {
|
||||
fn hash<H: Hasher>(&self, state: &mut H) {
|
||||
self.as_str().hash(state);
|
||||
}
|
||||
}
|
||||
|
||||
impl PartialOrd for AnthropicBeta {
|
||||
fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
|
||||
Some(self.cmp(other))
|
||||
}
|
||||
}
|
||||
|
||||
impl Ord for AnthropicBeta {
|
||||
fn cmp(&self, other: &Self) -> Ordering {
|
||||
self.as_str().cmp(other.as_str())
|
||||
}
|
||||
}
|
||||
|
||||
/// The values of one `anthropic-beta` header: sorted, deduplicated, comma-joined on the wire.
|
||||
#[derive(Clone, Debug, Default, PartialEq, Eq)]
|
||||
pub struct BetaSet(BTreeSet<AnthropicBeta>);
|
||||
|
||||
impl BetaSet {
|
||||
pub fn is_empty(&self) -> bool {
|
||||
self.0.is_empty()
|
||||
}
|
||||
|
||||
pub fn contains(&self, beta: &AnthropicBeta) -> bool {
|
||||
self.0.contains(beta)
|
||||
}
|
||||
|
||||
pub fn iter(&self) -> impl Iterator<Item = &AnthropicBeta> {
|
||||
self.0.iter()
|
||||
}
|
||||
|
||||
pub fn union(self, other: Self) -> Self {
|
||||
self.0.into_iter().chain(other.0).collect()
|
||||
}
|
||||
}
|
||||
|
||||
impl FromIterator<AnthropicBeta> for BetaSet {
|
||||
fn from_iter<I: IntoIterator<Item = AnthropicBeta>>(betas: I) -> Self {
|
||||
Self(betas.into_iter().collect())
|
||||
}
|
||||
}
|
||||
|
||||
impl IntoIterator for BetaSet {
|
||||
type Item = AnthropicBeta;
|
||||
type IntoIter = std::collections::btree_set::IntoIter<AnthropicBeta>;
|
||||
|
||||
fn into_iter(self) -> Self::IntoIter {
|
||||
self.0.into_iter()
|
||||
}
|
||||
}
|
||||
|
||||
impl FromStr for BetaSet {
|
||||
type Err = Infallible;
|
||||
|
||||
fn from_str(header: &str) -> Result<Self, Infallible> {
|
||||
Ok(header
|
||||
.split(',')
|
||||
.map(str::trim)
|
||||
.filter(|piece| !piece.is_empty())
|
||||
.map(|piece| AnthropicBeta::from_str(piece).unwrap_or_else(|never| match never {}))
|
||||
.collect())
|
||||
}
|
||||
}
|
||||
|
||||
impl fmt::Display for BetaSet {
|
||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
let mut betas = self.0.iter();
|
||||
let Some(first) = betas.next() else {
|
||||
return Ok(());
|
||||
};
|
||||
f.write_str(first.as_str())?;
|
||||
betas.try_for_each(|beta| write!(f, ",{beta}"))
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use rstest::rstest;
|
||||
|
||||
use super::*;
|
||||
|
||||
fn set(header: &str) -> BetaSet {
|
||||
header.parse().unwrap_or_else(|never| match never {})
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn every_known_beta_parses_back_to_itself(
|
||||
#[values(
|
||||
AnthropicBeta::Oauth20250420,
|
||||
AnthropicBeta::WebFetch20250910,
|
||||
AnthropicBeta::WebSearch20250305,
|
||||
AnthropicBeta::ContextManagement20250627,
|
||||
AnthropicBeta::Compact20260112,
|
||||
AnthropicBeta::Compact20260904,
|
||||
AnthropicBeta::StructuredOutputs20251113,
|
||||
AnthropicBeta::AdvancedToolUse20251120,
|
||||
AnthropicBeta::FastMode20260201,
|
||||
AnthropicBeta::AdvisorTool20260301,
|
||||
AnthropicBeta::PerTurnControl20260701,
|
||||
AnthropicBeta::DangerousToolUse20260903
|
||||
)]
|
||||
beta: AnthropicBeta,
|
||||
) {
|
||||
let parsed: AnthropicBeta = beta.as_str().parse().unwrap();
|
||||
assert!(!matches!(parsed, AnthropicBeta::Other(_)));
|
||||
assert_eq!(parsed, beta);
|
||||
assert!(AnthropicBeta::KNOWN.contains(&beta));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn unknown_values_are_kept_verbatim() {
|
||||
let parsed: AnthropicBeta = "claude-code-20250219".parse().unwrap();
|
||||
assert_eq!(
|
||||
parsed,
|
||||
AnthropicBeta::Other("claude-code-20250219".to_string())
|
||||
);
|
||||
assert_eq!(parsed.to_string(), "claude-code-20250219");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn a_known_value_spelled_as_other_is_the_same_beta() {
|
||||
let spelled_out = AnthropicBeta::Other("compact-2026-01-12".to_string());
|
||||
assert_eq!(spelled_out, AnthropicBeta::Compact20260112);
|
||||
assert_eq!(
|
||||
spelled_out.cmp(&AnthropicBeta::Compact20260112),
|
||||
Ordering::Equal
|
||||
);
|
||||
assert_eq!(
|
||||
BetaSet::from_iter([spelled_out, AnthropicBeta::Compact20260112]).to_string(),
|
||||
"compact-2026-01-12"
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::empty("", "")]
|
||||
#[case::blank_pieces(" , ,", "")]
|
||||
#[case::single("b", "b")]
|
||||
#[case::sorted("c,a", "a,c")]
|
||||
#[case::trimmed_and_deduplicated("b, a ,b", "a,b")]
|
||||
#[case::blank_pieces_skipped("a,,b", "a,b")]
|
||||
#[case::known_and_unknown_sort_together(
|
||||
"web-search-2025-03-05,claude-code-20250219,fast-mode-2026-02-01",
|
||||
"claude-code-20250219,fast-mode-2026-02-01,web-search-2025-03-05"
|
||||
)]
|
||||
fn header_values_round_trip_sorted_and_deduplicated(#[case] header: &str, #[case] wire: &str) {
|
||||
assert_eq!(set(header).to_string(), wire);
|
||||
assert_eq!(set(header).is_empty(), wire.is_empty());
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::disjoint("a,c", "b", "a,b,c")]
|
||||
#[case::overlapping("a,b", "b,c", "a,b,c")]
|
||||
#[case::empty_right("a", "", "a")]
|
||||
#[case::empty_left("", "a", "a")]
|
||||
fn union_merges_both_sides(#[case] left: &str, #[case] right: &str, #[case] wire: &str) {
|
||||
assert_eq!(set(left).union(set(right)).to_string(), wire);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn contains_matches_by_wire_value() {
|
||||
let betas = set("oauth-2025-04-20,claude-code-20250219");
|
||||
assert!(betas.contains(&AnthropicBeta::Oauth20250420));
|
||||
assert!(betas.contains(&AnthropicBeta::Other("claude-code-20250219".into())));
|
||||
assert!(!betas.contains(&AnthropicBeta::FastMode20260201));
|
||||
}
|
||||
}
|
||||
|
|
@ -111,6 +111,70 @@ impl From<EffortLevel> for ReasoningEffort {
|
|||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, IntoStaticStr, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "lowercase")]
|
||||
#[strum(serialize_all = "lowercase")]
|
||||
pub enum Speed {
|
||||
Fast,
|
||||
Standard,
|
||||
}
|
||||
|
||||
impl Speed {
|
||||
pub fn as_str(self) -> &'static str {
|
||||
self.into()
|
||||
}
|
||||
}
|
||||
|
||||
/// The tools whose presence changes how the request is sent. Every other tool, custom or
|
||||
/// server, deserializes as `Recognized::Unrecognized` and passes through verbatim.
|
||||
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
|
||||
#[serde(tag = "type")]
|
||||
pub enum AnthropicTool {
|
||||
#[serde(rename = "advisor_20260301")]
|
||||
Advisor {
|
||||
#[serde(flatten)]
|
||||
extra: Map<String, Value>,
|
||||
},
|
||||
#[serde(rename = "tool_search_tool_regex_20251119")]
|
||||
ToolSearchRegex {
|
||||
#[serde(flatten)]
|
||||
extra: Map<String, Value>,
|
||||
},
|
||||
#[serde(rename = "tool_search_tool_bm25_20251119")]
|
||||
ToolSearchBm25 {
|
||||
#[serde(flatten)]
|
||||
extra: Map<String, Value>,
|
||||
},
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
|
||||
#[serde(tag = "type")]
|
||||
pub enum ContextEdit {
|
||||
#[serde(rename = "compact_20260112")]
|
||||
Compact {
|
||||
#[serde(flatten)]
|
||||
extra: Map<String, Value>,
|
||||
},
|
||||
#[serde(rename = "clear_tool_uses_20250919")]
|
||||
ClearToolUses {
|
||||
#[serde(flatten)]
|
||||
extra: Map<String, Value>,
|
||||
},
|
||||
#[serde(rename = "clear_thinking_20251015")]
|
||||
ClearThinking {
|
||||
#[serde(flatten)]
|
||||
extra: Map<String, Value>,
|
||||
},
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
|
||||
pub struct ContextManagement {
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub edits: Option<Vec<Recognized<ContextEdit>>>,
|
||||
#[serde(flatten)]
|
||||
pub extra: Map<String, Value>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
|
||||
pub struct OutputConfig {
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
|
|
@ -210,7 +274,7 @@ pub struct AnthropicMessagesOptionalParams {
|
|||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub top_k: Option<i64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub tools: Option<Vec<Value>>,
|
||||
pub tools: Option<Vec<Recognized<AnthropicTool>>>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub tool_choice: Option<Value>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
|
|
@ -222,13 +286,13 @@ pub struct AnthropicMessagesOptionalParams {
|
|||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub mcp_servers: Option<Vec<Value>>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub context_management: Option<Value>,
|
||||
pub context_management: Option<Recognized<ContextManagement>>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub output_format: Option<Value>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub output_config: Option<Recognized<OutputConfig>>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub speed: Option<String>,
|
||||
pub speed: Option<Recognized<Speed>>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub inference_geo: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
|
|
@ -397,6 +461,33 @@ mod tests {
|
|||
"thinking": {"type": "future", "budget_tokens": 1},
|
||||
"output_config": "bogus"
|
||||
}))]
|
||||
#[case::tools_speed_and_context_management(json!({
|
||||
"model": "m",
|
||||
"messages": [],
|
||||
"speed": "fast",
|
||||
"tools": [
|
||||
{"name": "get_weather", "input_schema": {"type": "object"}},
|
||||
{"type": "custom", "name": "f", "input_schema": {}},
|
||||
{"type": "web_search_20250305", "name": "web_search", "max_uses": 3},
|
||||
{"type": "advisor_20260301", "name": "advisor", "model": "claude-opus-4-6"},
|
||||
{"type": "tool_search_tool_regex_20251119", "name": "tool_search_tool_regex"},
|
||||
{"type": "tool_search_tool_bm25_20251119"}
|
||||
],
|
||||
"context_management": {"edits": [
|
||||
{"type": "compact_20260112", "trigger": {"type": "input_tokens", "value": 1000}},
|
||||
{"type": "clear_tool_uses_20250919", "keep": {"type": "tool_uses", "value": 3}},
|
||||
{"type": "clear_thinking_20251015"},
|
||||
{"type": "future_edit"},
|
||||
{}
|
||||
], "future": true}
|
||||
}))]
|
||||
#[case::unrecognized_tools_speed_and_context_management_are_kept_verbatim(json!({
|
||||
"model": "m",
|
||||
"messages": [],
|
||||
"speed": "turbo",
|
||||
"tools": ["none", 5],
|
||||
"context_management": [{"type": "compaction", "compact_threshold": 5}]
|
||||
}))]
|
||||
fn request_round_trips_unchanged(#[case] request: Value) {
|
||||
assert_eq!(round_trip::<AnthropicMessagesRequest>(&request), request);
|
||||
}
|
||||
|
|
@ -444,6 +535,114 @@ mod tests {
|
|||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::advisor(
|
||||
json!({"type": "advisor_20260301", "name": "advisor"}),
|
||||
Recognized::Known(AnthropicTool::Advisor { extra: Map::from_iter([("name".to_string(), json!("advisor"))]) })
|
||||
)]
|
||||
#[case::regex_tool_search(
|
||||
json!({"type": "tool_search_tool_regex_20251119"}),
|
||||
Recognized::Known(AnthropicTool::ToolSearchRegex { extra: Map::new() })
|
||||
)]
|
||||
#[case::bm25_tool_search(
|
||||
json!({"type": "tool_search_tool_bm25_20251119"}),
|
||||
Recognized::Known(AnthropicTool::ToolSearchBm25 { extra: Map::new() })
|
||||
)]
|
||||
#[case::custom_tool_without_a_type(
|
||||
json!({"name": "advisor", "input_schema": {}}),
|
||||
Recognized::Unrecognized(json!({"name": "advisor", "input_schema": {}}))
|
||||
)]
|
||||
#[case::other_server_tool(
|
||||
json!({"type": "web_search_20250305", "name": "web_search"}),
|
||||
Recognized::Unrecognized(json!({"type": "web_search_20250305", "name": "web_search"}))
|
||||
)]
|
||||
#[case::not_an_object(json!("advisor_20260301"), Recognized::Unrecognized(json!("advisor_20260301")))]
|
||||
fn tools_are_recognized_by_their_exact_type(
|
||||
#[case] tool: Value,
|
||||
#[case] expected: Recognized<AnthropicTool>,
|
||||
) {
|
||||
assert_eq!(
|
||||
serde_json::from_value::<Recognized<AnthropicTool>>(tool).unwrap(),
|
||||
expected
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::compact(
|
||||
json!({"type": "compact_20260112", "trigger": {"type": "input_tokens", "value": 1}}),
|
||||
Recognized::Known(ContextEdit::Compact {
|
||||
extra: Map::from_iter([("trigger".to_string(), json!({"type": "input_tokens", "value": 1}))]),
|
||||
})
|
||||
)]
|
||||
#[case::clear_tool_uses(
|
||||
json!({"type": "clear_tool_uses_20250919"}),
|
||||
Recognized::Known(ContextEdit::ClearToolUses { extra: Map::new() })
|
||||
)]
|
||||
#[case::clear_thinking(
|
||||
json!({"type": "clear_thinking_20251015"}),
|
||||
Recognized::Known(ContextEdit::ClearThinking { extra: Map::new() })
|
||||
)]
|
||||
#[case::unknown_type(json!({"type": "future"}), Recognized::Unrecognized(json!({"type": "future"})))]
|
||||
#[case::no_type(json!({}), Recognized::Unrecognized(json!({})))]
|
||||
fn context_edits_are_recognized_by_their_exact_type(
|
||||
#[case] edit: Value,
|
||||
#[case] expected: Recognized<ContextEdit>,
|
||||
) {
|
||||
assert_eq!(
|
||||
serde_json::from_value::<Recognized<ContextEdit>>(edit).unwrap(),
|
||||
expected
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::edits(
|
||||
json!({"edits": [{"type": "compact_20260112"}]}),
|
||||
Recognized::Known(ContextManagement {
|
||||
edits: Some(vec![Recognized::Known(ContextEdit::Compact { extra: Map::new() })]),
|
||||
extra: Map::new(),
|
||||
})
|
||||
)]
|
||||
#[case::object_without_edits(
|
||||
json!({"future": 1}),
|
||||
Recognized::Known(ContextManagement {
|
||||
edits: None,
|
||||
extra: Map::from_iter([("future".to_string(), json!(1))]),
|
||||
})
|
||||
)]
|
||||
#[case::openai_list(json!([{"type": "compaction"}]), Recognized::Unrecognized(json!([{"type": "compaction"}])))]
|
||||
#[case::edits_not_a_list(json!({"edits": 5}), Recognized::Unrecognized(json!({"edits": 5})))]
|
||||
#[case::scalar(json!("compaction"), Recognized::Unrecognized(json!("compaction")))]
|
||||
fn context_management_is_known_only_as_an_edits_object(
|
||||
#[case] value: Value,
|
||||
#[case] expected: Recognized<ContextManagement>,
|
||||
) {
|
||||
assert_eq!(
|
||||
serde_json::from_value::<Recognized<ContextManagement>>(value).unwrap(),
|
||||
expected
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::fast(json!("fast"), Recognized::Known(Speed::Fast))]
|
||||
#[case::standard(json!("standard"), Recognized::Known(Speed::Standard))]
|
||||
#[case::unknown(json!("turbo"), Recognized::Unrecognized(json!("turbo")))]
|
||||
#[case::wrong_case(json!("Fast"), Recognized::Unrecognized(json!("Fast")))]
|
||||
#[case::not_a_string(json!(1), Recognized::Unrecognized(json!(1)))]
|
||||
fn speed_is_known_only_as_a_documented_value(
|
||||
#[case] value: Value,
|
||||
#[case] expected: Recognized<Speed>,
|
||||
) {
|
||||
assert_eq!(
|
||||
serde_json::from_value::<Recognized<Speed>>(value).unwrap(),
|
||||
expected
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn speed_names_match_the_wire(#[values(Speed::Fast, Speed::Standard)] speed: Speed) {
|
||||
assert_eq!(serde_json::to_value(speed).unwrap(), json!(speed.as_str()));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn effort_level_names_match_the_wire(
|
||||
#[values(
|
||||
|
|
|
|||
|
|
@ -1,2 +1,3 @@
|
|||
pub mod anthropic;
|
||||
pub mod anthropic_messages;
|
||||
pub mod openai;
|
||||
|
|
|
|||
|
|
@ -172,6 +172,7 @@ _custom_logger_compatible_callbacks_literal = Literal[
|
|||
"levo",
|
||||
"compression_interception",
|
||||
"newrelic",
|
||||
"signoz",
|
||||
]
|
||||
cold_storage_custom_logger: Optional[_custom_logger_compatible_callbacks_literal] = None
|
||||
logged_real_time_event_types: Optional[Union[List[str], Literal["*"]]] = None
|
||||
|
|
@ -1691,7 +1692,7 @@ if TYPE_CHECKING:
|
|||
SagemakerNovaConfig as SagemakerNovaConfig,
|
||||
)
|
||||
from .llms.cohere.chat.transformation import CohereChatConfig as CohereChatConfig
|
||||
from .llms.anthropic.experimental_pass_through.messages.transformation import (
|
||||
from .llms.anthropic.pass_through.messages.transformation import (
|
||||
AnthropicMessagesConfig as AnthropicMessagesConfig,
|
||||
)
|
||||
from .llms.bedrock.messages.invoke_transformations.anthropic_claude3_transformation import (
|
||||
|
|
|
|||
|
|
@ -21,6 +21,22 @@ is_internal_call: Final[ContextVar[bool]] = ContextVar("is_internal_call", defau
|
|||
# moment they can land on either side of a window boundary and disagree with each other.
|
||||
_billing_time: Final[ContextVar[datetime | None]] = ContextVar("billing_time", default=None)
|
||||
|
||||
_post_response: Final[ContextVar[bool]] = ContextVar("post_response", default=False)
|
||||
|
||||
|
||||
@contextmanager
|
||||
def post_response_phase() -> Generator[None]:
|
||||
"""Work the caller no longer waits for (success callbacks, response-cache writes), including tasks it spawns."""
|
||||
token: Final = _post_response.set(True)
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
_post_response.reset(token)
|
||||
|
||||
|
||||
def in_post_response_phase() -> bool:
|
||||
return _post_response.get()
|
||||
|
||||
|
||||
@contextmanager
|
||||
def pinned_billing_time(moment: datetime) -> Generator[None]:
|
||||
|
|
|
|||
|
|
@ -742,7 +742,7 @@ _LLM_CONFIGS_IMPORT_MAP: Final = {
|
|||
),
|
||||
"CohereChatConfig": (".llms.cohere.chat.transformation", "CohereChatConfig"),
|
||||
"AnthropicMessagesConfig": (
|
||||
".llms.anthropic.experimental_pass_through.messages.transformation",
|
||||
".llms.anthropic.pass_through.messages.transformation",
|
||||
"AnthropicMessagesConfig",
|
||||
),
|
||||
"BedrockClaudePlatformMessagesConfig": (
|
||||
|
|
|
|||
|
|
@ -159,6 +159,7 @@ class ServiceLogging(CustomLogger):
|
|||
parent_otel_span: Span | None = None,
|
||||
start_time: datetime | float | None = None,
|
||||
end_time: float | datetime | None = None,
|
||||
caller: str | None = None,
|
||||
):
|
||||
"""
|
||||
Handles both sync and async monitoring by checking for existing event loop.
|
||||
|
|
@ -172,6 +173,7 @@ class ServiceLogging(CustomLogger):
|
|||
service=service,
|
||||
duration=duration,
|
||||
call_type=call_type,
|
||||
caller=caller,
|
||||
parent_otel_span=parent_otel_span,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
|
|
@ -187,6 +189,7 @@ class ServiceLogging(CustomLogger):
|
|||
parent_otel_span: Span | None = None,
|
||||
start_time: datetime | float | None = None,
|
||||
end_time: float | datetime | None = None,
|
||||
caller: str | None = None,
|
||||
):
|
||||
"""
|
||||
Handles both sync and async monitoring by checking for existing event loop.
|
||||
|
|
@ -200,6 +203,7 @@ class ServiceLogging(CustomLogger):
|
|||
duration=duration,
|
||||
error=error,
|
||||
call_type=call_type,
|
||||
caller=caller,
|
||||
parent_otel_span=parent_otel_span,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
|
|
@ -215,6 +219,7 @@ class ServiceLogging(CustomLogger):
|
|||
start_time: datetime | float | None = None,
|
||||
end_time: datetime | float | None = None,
|
||||
event_metadata: dict | None = None,
|
||||
caller: str | None = None,
|
||||
):
|
||||
"""
|
||||
- For counting if the redis, postgres call is successful
|
||||
|
|
@ -228,6 +233,7 @@ class ServiceLogging(CustomLogger):
|
|||
service=service,
|
||||
duration=duration,
|
||||
call_type=call_type,
|
||||
caller=caller,
|
||||
event_metadata=event_metadata,
|
||||
)
|
||||
|
||||
|
|
@ -313,6 +319,7 @@ class ServiceLogging(CustomLogger):
|
|||
start_time: datetime | float | None = None,
|
||||
end_time: float | datetime | None = None,
|
||||
event_metadata: dict | None = None,
|
||||
caller: str | None = None,
|
||||
):
|
||||
"""
|
||||
- For counting if the redis, postgres call is unsuccessful
|
||||
|
|
@ -332,6 +339,7 @@ class ServiceLogging(CustomLogger):
|
|||
service=service,
|
||||
duration=duration,
|
||||
call_type=call_type,
|
||||
caller=caller,
|
||||
event_metadata=event_metadata,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -43,7 +43,7 @@ def _uses_native_vertex_output(
|
|||
) -> bool:
|
||||
if custom_llm_provider != "vertex_ai":
|
||||
return False
|
||||
if model_name and getattr(litellm, "disable_vertex_batch_output_transformation", False):
|
||||
if model_name and litellm.disable_vertex_batch_output_transformation:
|
||||
return True
|
||||
return first_row is not None and is_native_vertex_batch_output_row(first_row)
|
||||
|
||||
|
|
|
|||
|
|
@ -25,7 +25,7 @@ from litellm._logging import verbose_logger
|
|||
from litellm.constants import CACHED_STREAMING_CHUNK_DELAY
|
||||
from litellm.litellm_core_utils.model_param_helper import ModelParamHelper
|
||||
from litellm.types.caching import *
|
||||
from litellm.types.utils import EmbeddingResponse, all_litellm_params
|
||||
from litellm.types.utils import EmbeddingResponse, is_litellm_owned_kwarg
|
||||
|
||||
from .azure_blob_cache import AzureBlobCache
|
||||
from .base_cache import BaseCache
|
||||
|
|
@ -377,7 +377,6 @@ class Cache:
|
|||
return preset_cache_key
|
||||
|
||||
combined_kwargs: Final = ModelParamHelper._get_all_llm_api_params()
|
||||
litellm_param_kwargs: Final = all_litellm_params
|
||||
is_semantic_cache: Final = self._is_semantic_cache()
|
||||
scope_excluded_params: Final = self._SEMANTIC_CACHE_SCOPE_EXCLUDED_PARAMS if is_semantic_cache else frozenset()
|
||||
for param in kwargs:
|
||||
|
|
@ -387,7 +386,7 @@ class Cache:
|
|||
param_value: str | None = self._get_param_value(param, kwargs)
|
||||
if param_value is not None:
|
||||
cache_key += f"{param}: {param_value}"
|
||||
elif param not in litellm_param_kwargs: # check if user passed in optional param - e.g. top_k
|
||||
elif not is_litellm_owned_kwarg(param):
|
||||
if litellm.enable_caching_on_provider_specific_optional_params is True: # feature flagged for now
|
||||
if kwargs[param] is None:
|
||||
continue # ignore None params
|
||||
|
|
|
|||
|
|
@ -24,6 +24,7 @@ from typing import TYPE_CHECKING, Any, Final, Optional, TypeVar
|
|||
from pydantic import BaseModel, ConfigDict, ValidationError
|
||||
|
||||
import litellm
|
||||
from litellm._internal_context import post_response_phase
|
||||
from litellm._logging import print_verbose, verbose_logger
|
||||
from litellm.caching import InMemoryCache
|
||||
from litellm.caching.caching import S3Cache
|
||||
|
|
@ -51,7 +52,7 @@ from litellm.types.utils import (
|
|||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages.response_cache import (
|
||||
from litellm.llms.anthropic.pass_through.messages.response_cache import (
|
||||
AnthropicMessagesStreamCacheWriter,
|
||||
)
|
||||
from litellm.types.utils import PromptTokensDetailsWrapper
|
||||
|
|
@ -126,7 +127,7 @@ def _should_defer_streaming_cache_hit_callbacks(*, cached_result: object) -> boo
|
|||
spend and callback records. A plain (non-stream) replay logs here, since nothing
|
||||
else will.
|
||||
"""
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages.response_cache import (
|
||||
from litellm.llms.anthropic.pass_through.messages.response_cache import (
|
||||
CachedAnthropicMessagesStreamIterator,
|
||||
)
|
||||
from litellm.responses.streaming_iterator import BaseResponsesAPIStreamingIterator
|
||||
|
|
@ -158,7 +159,8 @@ async def _complete_cache_write_despite_cancellation(write_factory: Callable[[],
|
|||
|
||||
|
||||
def create_cache_write_task(write_factory: Callable[[], Awaitable[None]]) -> "asyncio.Task[None]":
|
||||
task: Final = asyncio.create_task(_complete_cache_write_despite_cancellation(write_factory))
|
||||
with post_response_phase():
|
||||
task: Final = asyncio.create_task(_complete_cache_write_despite_cancellation(write_factory))
|
||||
_PENDING_CACHE_WRITES.add(task)
|
||||
task.add_done_callback(_PENDING_CACHE_WRITES.discard)
|
||||
return task
|
||||
|
|
@ -928,7 +930,7 @@ class LLMCachingHandler:
|
|||
elif (
|
||||
call_type == CallTypes.anthropic_messages.value or call_type == CallTypes.aanthropic_messages.value
|
||||
) and isinstance(cached_result, dict):
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages.response_cache import (
|
||||
from litellm.llms.anthropic.pass_through.messages.response_cache import (
|
||||
convert_cached_anthropic_messages_result,
|
||||
)
|
||||
|
||||
|
|
@ -1148,7 +1150,7 @@ class LLMCachingHandler:
|
|||
return result
|
||||
if not isinstance(result, AsyncIterator):
|
||||
return result
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages.response_cache import (
|
||||
from litellm.llms.anthropic.pass_through.messages.response_cache import (
|
||||
AnthropicMessagesStreamCacheWriter,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -839,7 +839,8 @@ class RedisCache(BaseCache):
|
|||
self.service_logger_obj.service_success_hook(
|
||||
service=ServiceTypes.REDIS,
|
||||
duration=_duration,
|
||||
call_type=f"set_cache <- {_get_call_stack_info()}",
|
||||
call_type="set_cache",
|
||||
caller=_get_call_stack_info(),
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
)
|
||||
|
|
@ -860,7 +861,8 @@ class RedisCache(BaseCache):
|
|||
self.service_logger_obj.service_success_hook(
|
||||
service=ServiceTypes.REDIS,
|
||||
duration=_duration,
|
||||
call_type=f"increment_cache <- {_get_call_stack_info()}",
|
||||
call_type="increment_cache",
|
||||
caller=_get_call_stack_info(),
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
)
|
||||
|
|
@ -874,7 +876,8 @@ class RedisCache(BaseCache):
|
|||
self.service_logger_obj.service_success_hook(
|
||||
service=ServiceTypes.REDIS,
|
||||
duration=_duration,
|
||||
call_type=f"increment_cache_ttl <- {_get_call_stack_info()}",
|
||||
call_type="increment_cache_ttl",
|
||||
caller=_get_call_stack_info(),
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
)
|
||||
|
|
@ -887,7 +890,8 @@ class RedisCache(BaseCache):
|
|||
self.service_logger_obj.service_success_hook(
|
||||
service=ServiceTypes.REDIS,
|
||||
duration=_duration,
|
||||
call_type=f"increment_cache_expire <- {_get_call_stack_info()}",
|
||||
call_type="increment_cache_expire",
|
||||
caller=_get_call_stack_info(),
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
)
|
||||
|
|
@ -963,7 +967,8 @@ class RedisCache(BaseCache):
|
|||
self.service_logger_obj.async_service_success_hook(
|
||||
service=ServiceTypes.REDIS,
|
||||
duration=_duration,
|
||||
call_type=f"async_scan_iter <- {_get_call_stack_info()}",
|
||||
call_type="async_scan_iter",
|
||||
caller=_get_call_stack_info(),
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
)
|
||||
|
|
@ -979,7 +984,8 @@ class RedisCache(BaseCache):
|
|||
service=ServiceTypes.REDIS,
|
||||
duration=_duration,
|
||||
error=e,
|
||||
call_type=f"async_scan_iter <- {_get_call_stack_info()}",
|
||||
call_type="async_scan_iter",
|
||||
caller=_get_call_stack_info(),
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
)
|
||||
|
|
@ -1100,7 +1106,8 @@ class RedisCache(BaseCache):
|
|||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
parent_otel_span=_get_parent_otel_span_from_kwargs(kwargs),
|
||||
call_type=f"async_set_cache <- {_get_call_stack_info()}",
|
||||
call_type="async_set_cache",
|
||||
caller=_get_call_stack_info(),
|
||||
)
|
||||
)
|
||||
log_redis_failure(
|
||||
|
|
@ -1129,7 +1136,8 @@ class RedisCache(BaseCache):
|
|||
self.service_logger_obj.async_service_success_hook(
|
||||
service=ServiceTypes.REDIS,
|
||||
duration=_duration,
|
||||
call_type=f"async_set_cache <- {_get_call_stack_info()}",
|
||||
call_type="async_set_cache",
|
||||
caller=_get_call_stack_info(),
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
parent_otel_span=_get_parent_otel_span_from_kwargs(kwargs),
|
||||
|
|
@ -1145,7 +1153,8 @@ class RedisCache(BaseCache):
|
|||
service=ServiceTypes.REDIS,
|
||||
duration=_duration,
|
||||
error=e,
|
||||
call_type=f"async_set_cache <- {_get_call_stack_info()}",
|
||||
call_type="async_set_cache",
|
||||
caller=_get_call_stack_info(),
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
parent_otel_span=_get_parent_otel_span_from_kwargs(kwargs),
|
||||
|
|
@ -1213,7 +1222,8 @@ class RedisCache(BaseCache):
|
|||
self.service_logger_obj.async_service_success_hook(
|
||||
service=ServiceTypes.REDIS,
|
||||
duration=_duration,
|
||||
call_type=f"async_set_cache_pipeline <- {_get_call_stack_info()}",
|
||||
call_type="async_set_cache_pipeline",
|
||||
caller=_get_call_stack_info(),
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
parent_otel_span=_get_parent_otel_span_from_kwargs(kwargs),
|
||||
|
|
@ -1229,7 +1239,8 @@ class RedisCache(BaseCache):
|
|||
service=ServiceTypes.REDIS,
|
||||
duration=_duration,
|
||||
error=e,
|
||||
call_type=f"async_set_cache_pipeline <- {_get_call_stack_info()}",
|
||||
call_type="async_set_cache_pipeline",
|
||||
caller=_get_call_stack_info(),
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
parent_otel_span=_get_parent_otel_span_from_kwargs(kwargs),
|
||||
|
|
@ -1263,7 +1274,8 @@ class RedisCache(BaseCache):
|
|||
self.service_logger_obj.async_service_success_hook(
|
||||
service=ServiceTypes.REDIS,
|
||||
duration=time.time() - start_time,
|
||||
call_type=f"async_set_cache_pipeline_with_ttls <- {_get_call_stack_info()}",
|
||||
call_type="async_set_cache_pipeline_with_ttls",
|
||||
caller=_get_call_stack_info(),
|
||||
start_time=start_time,
|
||||
end_time=time.time(),
|
||||
)
|
||||
|
|
@ -1274,7 +1286,8 @@ class RedisCache(BaseCache):
|
|||
service=ServiceTypes.REDIS,
|
||||
duration=time.time() - start_time,
|
||||
error=e,
|
||||
call_type=f"async_set_cache_pipeline_with_ttls <- {_get_call_stack_info()}",
|
||||
call_type="async_set_cache_pipeline_with_ttls",
|
||||
caller=_get_call_stack_info(),
|
||||
start_time=start_time,
|
||||
end_time=time.time(),
|
||||
)
|
||||
|
|
@ -1322,7 +1335,8 @@ class RedisCache(BaseCache):
|
|||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
parent_otel_span=_get_parent_otel_span_from_kwargs(kwargs),
|
||||
call_type=f"async_set_cache_sadd <- {_get_call_stack_info()}",
|
||||
call_type="async_set_cache_sadd",
|
||||
caller=_get_call_stack_info(),
|
||||
)
|
||||
)
|
||||
# NON blocking - notify users Redis is throwing an exception
|
||||
|
|
@ -1342,7 +1356,8 @@ class RedisCache(BaseCache):
|
|||
self.service_logger_obj.async_service_success_hook(
|
||||
service=ServiceTypes.REDIS,
|
||||
duration=_duration,
|
||||
call_type=f"async_set_cache_sadd <- {_get_call_stack_info()}",
|
||||
call_type="async_set_cache_sadd",
|
||||
caller=_get_call_stack_info(),
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
parent_otel_span=_get_parent_otel_span_from_kwargs(kwargs),
|
||||
|
|
@ -1356,7 +1371,8 @@ class RedisCache(BaseCache):
|
|||
service=ServiceTypes.REDIS,
|
||||
duration=_duration,
|
||||
error=e,
|
||||
call_type=f"async_set_cache_sadd <- {_get_call_stack_info()}",
|
||||
call_type="async_set_cache_sadd",
|
||||
caller=_get_call_stack_info(),
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
parent_otel_span=_get_parent_otel_span_from_kwargs(kwargs),
|
||||
|
|
@ -1427,7 +1443,8 @@ class RedisCache(BaseCache):
|
|||
self.service_logger_obj.async_service_success_hook(
|
||||
service=ServiceTypes.REDIS,
|
||||
duration=_duration,
|
||||
call_type=f"async_increment <- {_get_call_stack_info()}",
|
||||
call_type="async_increment",
|
||||
caller=_get_call_stack_info(),
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
parent_otel_span=parent_otel_span,
|
||||
|
|
@ -1443,7 +1460,8 @@ class RedisCache(BaseCache):
|
|||
service=ServiceTypes.REDIS,
|
||||
duration=_duration,
|
||||
error=e,
|
||||
call_type=f"async_increment <- {_get_call_stack_info()}",
|
||||
call_type="async_increment",
|
||||
caller=_get_call_stack_info(),
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
parent_otel_span=parent_otel_span,
|
||||
|
|
@ -1531,7 +1549,8 @@ class RedisCache(BaseCache):
|
|||
self.service_logger_obj.service_success_hook(
|
||||
service=ServiceTypes.REDIS,
|
||||
duration=_duration,
|
||||
call_type=f"get_cache <- {_get_call_stack_info()}",
|
||||
call_type="get_cache",
|
||||
caller=_get_call_stack_info(),
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
parent_otel_span=parent_otel_span,
|
||||
|
|
@ -1590,7 +1609,8 @@ class RedisCache(BaseCache):
|
|||
self.service_logger_obj.service_success_hook(
|
||||
service=ServiceTypes.REDIS,
|
||||
duration=_duration,
|
||||
call_type=f"batch_get_cache <- {_get_call_stack_info()}",
|
||||
call_type="batch_get_cache",
|
||||
caller=_get_call_stack_info(),
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
parent_otel_span=parent_otel_span,
|
||||
|
|
@ -1614,7 +1634,8 @@ class RedisCache(BaseCache):
|
|||
service=ServiceTypes.REDIS,
|
||||
duration=failed_at - start_time,
|
||||
error=e,
|
||||
call_type=f"batch_get_cache <- {_get_call_stack_info()}",
|
||||
call_type="batch_get_cache",
|
||||
caller=_get_call_stack_info(),
|
||||
start_time=start_time,
|
||||
end_time=failed_at,
|
||||
parent_otel_span=parent_otel_span,
|
||||
|
|
@ -1643,7 +1664,8 @@ class RedisCache(BaseCache):
|
|||
self.service_logger_obj.async_service_success_hook(
|
||||
service=ServiceTypes.REDIS,
|
||||
duration=_duration,
|
||||
call_type=f"async_get_cache <- {_get_call_stack_info()}",
|
||||
call_type="async_get_cache",
|
||||
caller=_get_call_stack_info(),
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
parent_otel_span=parent_otel_span,
|
||||
|
|
@ -1659,7 +1681,8 @@ class RedisCache(BaseCache):
|
|||
service=ServiceTypes.REDIS,
|
||||
duration=_duration,
|
||||
error=e,
|
||||
call_type=f"async_get_cache <- {_get_call_stack_info()}",
|
||||
call_type="async_get_cache",
|
||||
caller=_get_call_stack_info(),
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
parent_otel_span=parent_otel_span,
|
||||
|
|
@ -1704,7 +1727,8 @@ class RedisCache(BaseCache):
|
|||
self.service_logger_obj.async_service_success_hook(
|
||||
service=ServiceTypes.REDIS,
|
||||
duration=_duration,
|
||||
call_type=f"async_batch_get_cache <- {_get_call_stack_info()}",
|
||||
call_type="async_batch_get_cache",
|
||||
caller=_get_call_stack_info(),
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
parent_otel_span=parent_otel_span,
|
||||
|
|
@ -1732,7 +1756,8 @@ class RedisCache(BaseCache):
|
|||
service=ServiceTypes.REDIS,
|
||||
duration=_duration,
|
||||
error=e,
|
||||
call_type=f"async_batch_get_cache <- {_get_call_stack_info()}",
|
||||
call_type="async_batch_get_cache",
|
||||
caller=_get_call_stack_info(),
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
parent_otel_span=parent_otel_span,
|
||||
|
|
@ -1757,7 +1782,8 @@ class RedisCache(BaseCache):
|
|||
self.service_logger_obj.service_success_hook(
|
||||
service=ServiceTypes.REDIS,
|
||||
duration=_duration,
|
||||
call_type=f"sync_ping <- {_get_call_stack_info()}",
|
||||
call_type="sync_ping",
|
||||
caller=_get_call_stack_info(),
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
)
|
||||
|
|
@ -1771,7 +1797,8 @@ class RedisCache(BaseCache):
|
|||
service=ServiceTypes.REDIS,
|
||||
duration=_duration,
|
||||
error=e,
|
||||
call_type=f"sync_ping <- {_get_call_stack_info()}",
|
||||
call_type="sync_ping",
|
||||
caller=_get_call_stack_info(),
|
||||
)
|
||||
verbose_logger.error("LiteLLM Redis Cache PING: - Got exception from REDIS : %s", e)
|
||||
raise e
|
||||
|
|
@ -1789,7 +1816,8 @@ class RedisCache(BaseCache):
|
|||
self.service_logger_obj.async_service_success_hook(
|
||||
service=ServiceTypes.REDIS,
|
||||
duration=_duration,
|
||||
call_type=f"async_ping <- {_get_call_stack_info()}",
|
||||
call_type="async_ping",
|
||||
caller=_get_call_stack_info(),
|
||||
)
|
||||
)
|
||||
return response
|
||||
|
|
@ -1803,7 +1831,8 @@ class RedisCache(BaseCache):
|
|||
service=ServiceTypes.REDIS,
|
||||
duration=_duration,
|
||||
error=e,
|
||||
call_type=f"async_ping <- {_get_call_stack_info()}",
|
||||
call_type="async_ping",
|
||||
caller=_get_call_stack_info(),
|
||||
)
|
||||
)
|
||||
verbose_logger.error("LiteLLM Redis Cache PING: - Got exception from REDIS : %s", e)
|
||||
|
|
@ -1955,7 +1984,8 @@ class RedisCache(BaseCache):
|
|||
self.service_logger_obj.async_service_success_hook(
|
||||
service=ServiceTypes.REDIS,
|
||||
duration=_duration,
|
||||
call_type=f"async_increment_pipeline <- {_get_call_stack_info()}",
|
||||
call_type="async_increment_pipeline",
|
||||
caller=_get_call_stack_info(),
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
parent_otel_span=_get_parent_otel_span_from_kwargs(kwargs),
|
||||
|
|
@ -1971,7 +2001,8 @@ class RedisCache(BaseCache):
|
|||
service=ServiceTypes.REDIS,
|
||||
duration=_duration,
|
||||
error=e,
|
||||
call_type=f"async_increment_pipeline <- {_get_call_stack_info()}",
|
||||
call_type="async_increment_pipeline",
|
||||
caller=_get_call_stack_info(),
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
parent_otel_span=_get_parent_otel_span_from_kwargs(kwargs),
|
||||
|
|
@ -2049,7 +2080,8 @@ class RedisCache(BaseCache):
|
|||
self.service_logger_obj.async_service_success_hook(
|
||||
service=ServiceTypes.REDIS,
|
||||
duration=_duration,
|
||||
call_type=f"async_rpush <- {_get_call_stack_info()}",
|
||||
call_type="async_rpush",
|
||||
caller=_get_call_stack_info(),
|
||||
)
|
||||
)
|
||||
return response
|
||||
|
|
@ -2063,7 +2095,8 @@ class RedisCache(BaseCache):
|
|||
service=ServiceTypes.REDIS,
|
||||
duration=_duration,
|
||||
error=e,
|
||||
call_type=f"async_rpush <- {_get_call_stack_info()}",
|
||||
call_type="async_rpush",
|
||||
caller=_get_call_stack_info(),
|
||||
)
|
||||
)
|
||||
log_redis_failure(verbose_logger, logging.ERROR, "LiteLLM Redis Cache RPUSH: - Got exception from REDIS", e)
|
||||
|
|
@ -2096,7 +2129,8 @@ class RedisCache(BaseCache):
|
|||
self.service_logger_obj.async_service_success_hook(
|
||||
service=ServiceTypes.REDIS,
|
||||
duration=time.time() - start_time,
|
||||
call_type=f"async_rpush_and_trim <- {_get_call_stack_info()}",
|
||||
call_type="async_rpush_and_trim",
|
||||
caller=_get_call_stack_info(),
|
||||
)
|
||||
)
|
||||
return int(results[0])
|
||||
|
|
@ -2106,7 +2140,8 @@ class RedisCache(BaseCache):
|
|||
service=ServiceTypes.REDIS,
|
||||
duration=time.time() - start_time,
|
||||
error=e,
|
||||
call_type=f"async_rpush_and_trim <- {_get_call_stack_info()}",
|
||||
call_type="async_rpush_and_trim",
|
||||
caller=_get_call_stack_info(),
|
||||
)
|
||||
)
|
||||
log_redis_failure(
|
||||
|
|
@ -2163,7 +2198,8 @@ class RedisCache(BaseCache):
|
|||
self.service_logger_obj.async_service_success_hook(
|
||||
service=ServiceTypes.REDIS,
|
||||
duration=_duration,
|
||||
call_type=f"async_rpush_pipeline <- {_get_call_stack_info()}",
|
||||
call_type="async_rpush_pipeline",
|
||||
caller=_get_call_stack_info(),
|
||||
)
|
||||
)
|
||||
return results
|
||||
|
|
@ -2176,7 +2212,8 @@ class RedisCache(BaseCache):
|
|||
service=ServiceTypes.REDIS,
|
||||
duration=_duration,
|
||||
error=e,
|
||||
call_type=f"async_rpush_pipeline <- {_get_call_stack_info()}",
|
||||
call_type="async_rpush_pipeline",
|
||||
caller=_get_call_stack_info(),
|
||||
)
|
||||
)
|
||||
log_redis_failure(
|
||||
|
|
@ -2230,7 +2267,8 @@ class RedisCache(BaseCache):
|
|||
self.service_logger_obj.async_service_success_hook(
|
||||
service=ServiceTypes.REDIS,
|
||||
duration=_duration,
|
||||
call_type=f"async_lpop <- {_get_call_stack_info()}",
|
||||
call_type="async_lpop",
|
||||
caller=_get_call_stack_info(),
|
||||
)
|
||||
)
|
||||
|
||||
|
|
@ -2256,7 +2294,8 @@ class RedisCache(BaseCache):
|
|||
service=ServiceTypes.REDIS,
|
||||
duration=_duration,
|
||||
error=e,
|
||||
call_type=f"async_lpop <- {_get_call_stack_info()}",
|
||||
call_type="async_lpop",
|
||||
caller=_get_call_stack_info(),
|
||||
)
|
||||
)
|
||||
log_redis_failure(verbose_logger, logging.ERROR, "LiteLLM Redis Cache LPOP: - Got exception from REDIS", e)
|
||||
|
|
@ -2354,7 +2393,8 @@ class RedisCache(BaseCache):
|
|||
self.service_logger_obj.async_service_success_hook(
|
||||
service=ServiceTypes.REDIS,
|
||||
duration=_duration,
|
||||
call_type=f"async_lpop_pipeline <- {_get_call_stack_info()}",
|
||||
call_type="async_lpop_pipeline",
|
||||
caller=_get_call_stack_info(),
|
||||
)
|
||||
)
|
||||
return results
|
||||
|
|
@ -2367,7 +2407,8 @@ class RedisCache(BaseCache):
|
|||
service=ServiceTypes.REDIS,
|
||||
duration=_duration,
|
||||
error=e,
|
||||
call_type=f"async_lpop_pipeline <- {_get_call_stack_info()}",
|
||||
call_type="async_lpop_pipeline",
|
||||
caller=_get_call_stack_info(),
|
||||
)
|
||||
)
|
||||
log_redis_failure(
|
||||
|
|
|
|||
|
|
@ -49,6 +49,7 @@ DEFAULT_FLUSH_INTERVAL_SECONDS: Final = int(os.getenv("DEFAULT_FLUSH_INTERVAL_SE
|
|||
DEFAULT_S3_FLUSH_INTERVAL_SECONDS: Final = int(os.getenv("DEFAULT_S3_FLUSH_INTERVAL_SECONDS", 10))
|
||||
DEFAULT_S3_BATCH_SIZE: Final = int(os.getenv("DEFAULT_S3_BATCH_SIZE", 512))
|
||||
DEFAULT_S3_MAX_CONCURRENT_UPLOADS: Final = int(os.getenv("DEFAULT_S3_MAX_CONCURRENT_UPLOADS", "16"))
|
||||
DEFAULT_S3_MAX_ADAPTIVE_CONCURRENCY: Final = get_env_int("DEFAULT_S3_MAX_ADAPTIVE_CONCURRENCY", 200)
|
||||
# https://docs.aws.amazon.com/AmazonS3/latest/userguide/object-keys.html
|
||||
MAX_S3_OBJECT_KEY_BYTES: Final = 1024
|
||||
S3_BOUNDED_OBJECT_KEY_HEAD_BYTES: Final = 64
|
||||
|
|
@ -945,6 +946,7 @@ openai_compatible_endpoints: Final[list] = [
|
|||
"https://api.libertai.io/v1",
|
||||
"https://pinstripes.io/v1",
|
||||
"https://api.meta.ai/v1",
|
||||
"https://api.sailresearch.com/v1",
|
||||
"https://api.cognition.ai/v1",
|
||||
"https://api.scx.ai/v1",
|
||||
"https://gigachat.devices.sberbank.ru/api/v1",
|
||||
|
|
@ -1020,6 +1022,7 @@ openai_compatible_providers: Final[list] = [
|
|||
"meta", # Meta Model API (Muse Spark) - JSON-configured provider
|
||||
"cognition",
|
||||
"scx-ai",
|
||||
"sail",
|
||||
]
|
||||
|
||||
OPENAI_AUDIO_TRANSCRIPTION_PROVIDERS: Final = frozenset({"openai"} | frozenset(openai_compatible_providers))
|
||||
|
|
@ -1607,6 +1610,8 @@ ALLOWED_VERTEX_AI_PASSTHROUGH_HEADERS: Final = {
|
|||
# e.g., 'x-pass-anthropic-beta: value' becomes 'anthropic-beta: value'
|
||||
# Works for all LLM pass-through endpoints (Vertex AI, Anthropic, Bedrock, etc.)
|
||||
PASS_THROUGH_HEADER_PREFIX: Final = "x-pass-"
|
||||
INTERNAL_KWARG_PREFIX: Final = "_litellm_"
|
||||
CONTROL_OPTIONS_KEY: Final = f"{INTERNAL_KWARG_PREFIX}control"
|
||||
|
||||
AZURE_SPEECH_CUSTOM_LLM_PROVIDER: Final = "azure_speech"
|
||||
AZURE_SPEECH_PASS_THROUGH_ROUTE_PREFIX: Final = "/azure_speech"
|
||||
|
|
|
|||
|
|
@ -25,7 +25,7 @@ from litellm.llms.base_llm import BaseImageEditConfig, BaseImageGenerationConfig
|
|||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
|
||||
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
|
||||
from litellm.llms.custom_llm import CustomLLM
|
||||
from litellm.utils import exception_type, get_litellm_params
|
||||
from litellm.utils import exception_type, filter_out_litellm_params, get_litellm_params
|
||||
|
||||
#################### Initialize provider clients ####################
|
||||
llm_http_handler: BaseLLMHTTPHandler = BaseLLMHTTPHandler()
|
||||
|
|
@ -52,7 +52,6 @@ from litellm.types.router import GenericLiteLLMParams
|
|||
from litellm.types.utils import (
|
||||
LITELLM_IMAGE_VARIATION_PROVIDERS,
|
||||
LlmProviders,
|
||||
all_litellm_params,
|
||||
)
|
||||
from litellm.utils import (
|
||||
ImageResponse,
|
||||
|
|
@ -249,11 +248,7 @@ def image_generation(
|
|||
"size",
|
||||
"style",
|
||||
]
|
||||
litellm_params: Final = all_litellm_params
|
||||
default_params: Final = openai_params + litellm_params
|
||||
non_default_params: Final = {
|
||||
k: v for k, v in kwargs.items() if k not in default_params
|
||||
} # model-specific params - pass them straight to the model/provider
|
||||
non_default_params: Final = filter_out_litellm_params(kwargs, excluding=openai_params)
|
||||
|
||||
image_generation_config: BaseImageGenerationConfig | None = None
|
||||
if custom_llm_provider is not None and custom_llm_provider in LlmProviders._member_map_.values():
|
||||
|
|
@ -757,11 +752,7 @@ def image_edit(
|
|||
"style",
|
||||
"async_call",
|
||||
]
|
||||
litellm_params_list: Final = all_litellm_params
|
||||
default_params: Final = openai_params + litellm_params_list
|
||||
non_default_params: Final = {
|
||||
k: v for k, v in kwargs.items() if k not in default_params
|
||||
} # model-specific params - pass them straight to the model/provider
|
||||
non_default_params: Final = filter_out_litellm_params(kwargs, excluding=openai_params)
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj")
|
||||
litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id", None)
|
||||
model_info: Final = kwargs.get("model_info", None)
|
||||
|
|
|
|||
78
litellm/integrations/adaptive_concurrency.py
Normal file
78
litellm/integrations/adaptive_concurrency.py
Normal file
|
|
@ -0,0 +1,78 @@
|
|||
"""
|
||||
Adaptive in-flight concurrency limiter (AIMD, Vector ARC style).
|
||||
|
||||
Grows the limit additively after `limit` consecutive clean completions and
|
||||
halves it only on an explicit throttle signal (429, 503, SlowDown, or a
|
||||
transport error out of the PUT). With floor == ceiling it degenerates to a
|
||||
fixed-width semaphore.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
from collections import deque
|
||||
from contextlib import suppress
|
||||
from dataclasses import dataclass
|
||||
from typing import Final
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class PutSample:
|
||||
throttled: bool
|
||||
|
||||
|
||||
class AdaptiveConcurrencyLimiter:
|
||||
"""AIMD in-flight limiter used as `async with limiter:`."""
|
||||
|
||||
def __init__(self, initial: int, floor: int, ceiling: int) -> None:
|
||||
if not 1 <= floor <= ceiling:
|
||||
raise ValueError(f"adaptive limiter bounds must satisfy 1 <= floor <= ceiling, got {floor}..{ceiling}")
|
||||
self._limit: int = min(max(initial, floor), ceiling)
|
||||
self._floor: Final[int] = floor
|
||||
self._ceiling: Final[int] = ceiling
|
||||
self._clean_streak: int = 0
|
||||
self._in_flight: int = 0
|
||||
self._waiters: deque[asyncio.Future[None]] = deque() # mutable-ok: waiters queue up behind a full limit
|
||||
|
||||
@property
|
||||
def limit(self) -> int:
|
||||
return self._limit
|
||||
|
||||
async def __aenter__(self) -> "AdaptiveConcurrencyLimiter":
|
||||
if self._in_flight < self._limit:
|
||||
self._in_flight += 1
|
||||
return self
|
||||
waiter: Final = asyncio.get_running_loop().create_future()
|
||||
self._waiters.append(waiter)
|
||||
try:
|
||||
await waiter
|
||||
except asyncio.CancelledError:
|
||||
if waiter.done() and not waiter.cancelled():
|
||||
self._in_flight -= 1
|
||||
self._grant()
|
||||
else:
|
||||
with suppress(ValueError):
|
||||
self._waiters.remove(waiter)
|
||||
raise
|
||||
return self
|
||||
|
||||
def _grant(self) -> None:
|
||||
while self._in_flight < self._limit and self._waiters:
|
||||
waiter = self._waiters.popleft()
|
||||
if waiter.done():
|
||||
continue
|
||||
self._in_flight += 1
|
||||
waiter.set_result(None)
|
||||
|
||||
async def __aexit__(self, *_: object) -> None:
|
||||
self._in_flight -= 1
|
||||
self._grant()
|
||||
|
||||
def record(self, sample: PutSample) -> None:
|
||||
if sample.throttled:
|
||||
self._limit = max(self._floor, self._limit // 2)
|
||||
self._clean_streak = 0
|
||||
return
|
||||
self._clean_streak += 1
|
||||
if self._clean_streak >= self._limit and self._limit < self._ceiling:
|
||||
self._limit += 1
|
||||
self._clean_streak = 0
|
||||
self._grant()
|
||||
|
|
@ -695,6 +695,7 @@ class AnthropicCacheControlHook(CustomPromptManagement):
|
|||
tools: list | None = None,
|
||||
cache_control: object = None,
|
||||
request_kwargs: object = None,
|
||||
on_messages_route: bool = False,
|
||||
) -> bool:
|
||||
"""Return True if the request already carries any client-supplied cache_control.
|
||||
|
||||
|
|
@ -704,10 +705,14 @@ class AnthropicCacheControlHook(CustomPromptManagement):
|
|||
envelope. Configured injection points are an explicit instruction and are
|
||||
applied alongside the client's marks, bounded by the provider cap.
|
||||
"""
|
||||
return (
|
||||
AnthropicCacheControlHook.count_request_cache_breakpoints(messages, system)
|
||||
+ AnthropicCacheControlHook.count_external_cache_breakpoints(tools, cache_control, request_kwargs)
|
||||
) > 0
|
||||
external_breakpoints: Final = (
|
||||
AnthropicCacheControlHook.count_external_cache_breakpoints_on_messages_route(
|
||||
tools, cache_control, request_kwargs
|
||||
)
|
||||
if on_messages_route
|
||||
else AnthropicCacheControlHook.count_external_cache_breakpoints(tools, cache_control, request_kwargs)
|
||||
)
|
||||
return AnthropicCacheControlHook.count_request_cache_breakpoints(messages, system) + external_breakpoints > 0
|
||||
|
||||
@staticmethod
|
||||
def get_default_injection_points(
|
||||
|
|
@ -719,6 +724,7 @@ class AnthropicCacheControlHook(CustomPromptManagement):
|
|||
enable_prompt_caching: bool | None = None,
|
||||
cache_control: object = None,
|
||||
request_kwargs: object = None,
|
||||
on_messages_route: bool = False,
|
||||
) -> list[CacheControlInjectionPoint]:
|
||||
"""Default breakpoints when ``litellm.enable_anthropic_prompt_caching`` is on.
|
||||
|
||||
|
|
@ -739,7 +745,9 @@ class AnthropicCacheControlHook(CustomPromptManagement):
|
|||
if not supports_anthropic_cache_control(model, custom_llm_provider):
|
||||
return []
|
||||
|
||||
if AnthropicCacheControlHook._request_has_cache_control(messages, system, tools, cache_control, request_kwargs):
|
||||
if AnthropicCacheControlHook._request_has_cache_control(
|
||||
messages, system, tools, cache_control, request_kwargs, on_messages_route
|
||||
):
|
||||
return []
|
||||
|
||||
if is_claude_code_one_shot_subagent_request(
|
||||
|
|
@ -968,6 +976,7 @@ class AnthropicCacheControlHook(CustomPromptManagement):
|
|||
enable_prompt_caching=enable_prompt_caching,
|
||||
cache_control=cache_control,
|
||||
request_kwargs=kwargs,
|
||||
on_messages_route=True,
|
||||
)
|
||||
if model is not None
|
||||
else ()
|
||||
|
|
|
|||
|
|
@ -502,6 +502,27 @@
|
|||
},
|
||||
"description": "S3 Bucket (AWS) Logging Integration"
|
||||
},
|
||||
{
|
||||
"id": "signoz",
|
||||
"displayName": "SigNoz",
|
||||
"logo": "signoz.svg",
|
||||
"supports_key_team_logging": true,
|
||||
"dynamic_params": {
|
||||
"signoz_ingestion_endpoint": {
|
||||
"type": "text",
|
||||
"ui_name": "SigNoz Ingestion Endpoint",
|
||||
"description": "Ingestion endpoint for this team, e.g. https://ingest.us.signoz.cloud:443 for SigNoz Cloud or your own collector. Leave blank to use the proxy's configured endpoint. Regions: https://signoz.io/docs/ingestion/signoz-cloud/overview/",
|
||||
"required": false
|
||||
},
|
||||
"signoz_ingestion_key": {
|
||||
"type": "password",
|
||||
"ui_name": "SigNoz Ingestion Key (optional)",
|
||||
"description": "Ingestion key for this team, so its traces land in its own SigNoz account. Not needed for self-hosted SigNoz. Keys: https://signoz.io/docs/ingestion/signoz-cloud/keys/",
|
||||
"required": false
|
||||
}
|
||||
},
|
||||
"description": "SigNoz Logging Integration. Setup: https://signoz.io/docs/litellm-observability/"
|
||||
},
|
||||
{
|
||||
"id": "sqs",
|
||||
"displayName": "SQS",
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@ import copy
|
|||
import hashlib
|
||||
import os
|
||||
import secrets
|
||||
from collections.abc import Mapping
|
||||
from collections.abc import Mapping, Sequence
|
||||
from datetime import datetime
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, ClassVar, Final, Literal, Optional, get_args
|
||||
|
|
@ -37,6 +37,7 @@ from litellm.types.utils import (
|
|||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
dc: Final = DualCache()
|
||||
|
||||
|
||||
|
|
@ -106,6 +107,33 @@ def is_guardrail_intervention(e: Exception) -> bool:
|
|||
return is_fastapi_http_exception(e, _GUARDRAIL_BLOCK_STATUS_CODES)
|
||||
|
||||
|
||||
def _user_api_key_auth_from_request(request_data: Mapping[str, object]) -> "UserAPIKeyAuth":
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
metadata: Final = request_data.get(get_metadata_variable_name_from_kwargs(request_data))
|
||||
stamped: Final[Mapping[str, object]] = metadata if isinstance(metadata, dict) else {}
|
||||
|
||||
def stamped_str(field: str) -> str | None:
|
||||
value: Final = stamped.get(field)
|
||||
return value if isinstance(value, str) else None
|
||||
|
||||
return UserAPIKeyAuth(
|
||||
user_id=stamped_str("user_api_key_user_id"),
|
||||
team_id=stamped_str("user_api_key_team_id"),
|
||||
end_user_id=stamped_str("user_api_key_end_user_id"),
|
||||
api_key=stamped_str("user_api_key_hash"),
|
||||
request_route=stamped_str("user_api_key_request_route"),
|
||||
)
|
||||
|
||||
|
||||
def _unified_hook_fields(guardrail: "CustomGuardrail", request_data: Mapping[str, object]) -> Mapping[str, object]:
|
||||
metadata_bucket: Final = request_data.get(get_metadata_variable_name_from_kwargs(request_data))
|
||||
return {
|
||||
"guardrail_to_apply": guardrail,
|
||||
**({"litellm_metadata": metadata_bucket} if isinstance(metadata_bucket, dict) else {}),
|
||||
}
|
||||
|
||||
|
||||
def _strict_guardrail_modes_enabled() -> bool:
|
||||
"""Whether guardrail-mode validation raises (default) or logs a warning.
|
||||
|
||||
|
|
@ -789,8 +817,6 @@ class CustomGuardrail(CustomLogger):
|
|||
return unified_guardrail
|
||||
|
||||
async def async_pre_call_deployment_hook(self, kwargs: dict[str, Any], call_type: CallTypes | None) -> dict | None:
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
# should run guardrail
|
||||
litellm_guardrails: Final = kwargs.get("guardrails")
|
||||
if litellm_guardrails is None or not isinstance(litellm_guardrails, list):
|
||||
|
|
@ -808,13 +834,7 @@ class CustomGuardrail(CustomLogger):
|
|||
if target is not self:
|
||||
kwargs["guardrail_to_apply"] = self
|
||||
result: Final = await target.async_pre_call_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(
|
||||
user_id=kwargs.get("user_api_key_user_id"),
|
||||
team_id=kwargs.get("user_api_key_team_id"),
|
||||
end_user_id=kwargs.get("user_api_key_end_user_id"),
|
||||
api_key=kwargs.get("user_api_key_hash"),
|
||||
request_route=kwargs.get("user_api_key_request_route"),
|
||||
),
|
||||
user_api_key_dict=_user_api_key_auth_from_request(kwargs),
|
||||
cache=dc,
|
||||
data=kwargs,
|
||||
call_type="completion" if call_type == CallTypes.completion else "acompletion",
|
||||
|
|
@ -827,6 +847,52 @@ class CustomGuardrail(CustomLogger):
|
|||
|
||||
return kwargs
|
||||
|
||||
async def async_pre_call_hook_on_messages(
|
||||
self,
|
||||
request_data: Mapping[str, object],
|
||||
messages: Sequence[AllMessageValues],
|
||||
) -> tuple[AllMessageValues, ...]:
|
||||
from litellm.proxy.guardrails.exception_utils import (
|
||||
enrich_http_exception_with_guardrail_context,
|
||||
pre_call_rejection,
|
||||
)
|
||||
|
||||
target: Final = self._deployment_hook_target()
|
||||
scan_request: Final[dict[str, object]] = { # mutable-ok: async_pre_call_hook writes into the dict it is handed
|
||||
**{key: value for key, value in request_data.items() if key not in _PRE_CALL_CONTENT_KEYS},
|
||||
"messages": list(messages),
|
||||
**({} if target is self else _unified_hook_fields(self, request_data)),
|
||||
}
|
||||
try:
|
||||
result: Final = await target.async_pre_call_hook(
|
||||
user_api_key_dict=_user_api_key_auth_from_request(scan_request),
|
||||
cache=dc,
|
||||
data=scan_request,
|
||||
call_type="acompletion",
|
||||
)
|
||||
except SensitiveDataRouteException as e:
|
||||
unroutable: Final = pre_call_rejection(
|
||||
f"{e.guardrail_name or self.guardrail_name} asked to reroute the request to {e.route_to_model} "
|
||||
"over retrieved content; a request cannot be rerouted after retrieval, so it was blocked",
|
||||
self.guardrail_name,
|
||||
)
|
||||
enrich_http_exception_with_guardrail_context(unroutable, self)
|
||||
raise unroutable from e
|
||||
except Exception as e:
|
||||
enrich_http_exception_with_guardrail_context(e, self)
|
||||
raise
|
||||
if result is None:
|
||||
return tuple(messages)
|
||||
if isinstance(result, dict):
|
||||
scanned: Final = result.get("messages")
|
||||
return tuple(scanned) if isinstance(scanned, list) else tuple(messages)
|
||||
if isinstance(result, str):
|
||||
rejection: Final = pre_call_rejection(result, self.guardrail_name)
|
||||
enrich_http_exception_with_guardrail_context(rejection, self)
|
||||
raise rejection
|
||||
enrich_http_exception_with_guardrail_context(result, self)
|
||||
raise result
|
||||
|
||||
async def async_post_call_success_deployment_hook(
|
||||
self,
|
||||
request_data: dict,
|
||||
|
|
@ -836,8 +902,6 @@ class CustomGuardrail(CustomLogger):
|
|||
"""
|
||||
Allow modifying / reviewing the response just after it's received from the deployment.
|
||||
"""
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
# should run guardrail
|
||||
litellm_guardrails: Final = request_data.get("guardrails")
|
||||
if litellm_guardrails is None or not isinstance(litellm_guardrails, list):
|
||||
|
|
@ -851,13 +915,7 @@ class CustomGuardrail(CustomLogger):
|
|||
if target is not self:
|
||||
request_data["guardrail_to_apply"] = self # rebind-ok: dispatch consumes this key
|
||||
result: Final = await target.async_post_call_success_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(
|
||||
user_id=request_data.get("user_api_key_user_id"),
|
||||
team_id=request_data.get("user_api_key_team_id"),
|
||||
end_user_id=request_data.get("user_api_key_end_user_id"),
|
||||
api_key=request_data.get("user_api_key_hash"),
|
||||
request_route=request_data.get("user_api_key_request_route"),
|
||||
),
|
||||
user_api_key_dict=_user_api_key_auth_from_request(request_data),
|
||||
data=request_data,
|
||||
response=response,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -684,7 +684,7 @@ class LangfuseSpanExporter(SpanExporter):
|
|||
def _round(self, halving: _Halving) -> _Halving:
|
||||
sent: Final = tuple((batch, self._send_batch(batch)) for batch in halving.pending)
|
||||
return _Halving(
|
||||
pending=tuple(part for batch, outcome in sent if outcome == "too_large" for part in _smaller(batch)),
|
||||
pending=tuple(chain.from_iterable(_smaller(batch) for batch, outcome in sent if outcome == "too_large")),
|
||||
settled=halving.settled
|
||||
+ tuple(
|
||||
SpanExportResult.SUCCESS if outcome == "delivered" else SpanExportResult.FAILURE
|
||||
|
|
|
|||
|
|
@ -10,6 +10,14 @@ if TYPE_CHECKING:
|
|||
else:
|
||||
Span = Any
|
||||
|
||||
LANGTRACE_DEFAULT_HOST: Final = "https://app.langtrace.ai"
|
||||
LANGTRACE_TRACE_PATH: Final = "/api/trace"
|
||||
|
||||
|
||||
def langtrace_trace_endpoint(api_host: str | None) -> str:
|
||||
host: Final = (api_host or LANGTRACE_DEFAULT_HOST).rstrip("/")
|
||||
return host if host.endswith(LANGTRACE_TRACE_PATH) else host + LANGTRACE_TRACE_PATH
|
||||
|
||||
|
||||
class LangtraceAttributes:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -15,6 +15,7 @@ from litellm.integrations._types.open_inference import (
|
|||
SpanAttributes,
|
||||
)
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.integrations.langtrace import LANGTRACE_TRACE_PATH
|
||||
from litellm.integrations.opentelemetry_utils.gen_ai_semconv import (
|
||||
OTEL_SEMCONV_STABILITY_OPT_IN_ENV,
|
||||
OTELGenAISemconvMixin,
|
||||
|
|
@ -25,7 +26,7 @@ from litellm.integrations.otel.mappers.utils import drop_none
|
|||
from litellm.integrations.otel.model.baggage import promoted_metadata
|
||||
from litellm.integrations.otel.model.db_endpoint import db_span_attributes
|
||||
from litellm.integrations.otel.model.metadata import flatten_metadata
|
||||
from litellm.integrations.otel.model.semconv import Metric
|
||||
from litellm.integrations.otel.model.semconv import LiteLLM, Metric
|
||||
from litellm.integrations.otel.plumbing.otlp_tls import resolve_otlp_http_tls
|
||||
from litellm.litellm_core_utils.internal_call_metadata import is_unbilled_non_inference_call_from_params
|
||||
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
|
||||
|
|
@ -784,6 +785,8 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
|
|||
)
|
||||
for key, value in attributes.items():
|
||||
self.safe_set_attribute(span=span, key=key, value=value)
|
||||
if payload.caller is not None:
|
||||
self.safe_set_attribute(span=span, key=LiteLLM.SERVICE_CALLER, value=payload.caller)
|
||||
return span
|
||||
|
||||
async def async_service_success_hook(
|
||||
|
|
@ -3332,6 +3335,9 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
|
|||
if signal_type == "traces" and "/v2/trace/otlp" in endpoint:
|
||||
return endpoint
|
||||
|
||||
if signal_type == "traces" and self.callback_name == "langtrace" and endpoint.endswith(LANGTRACE_TRACE_PATH):
|
||||
return endpoint
|
||||
|
||||
# Check if endpoint already ends with the correct signal path
|
||||
target_path: Final = f"/v1/{signal_type}"
|
||||
if endpoint.endswith(target_path):
|
||||
|
|
|
|||
|
|
@ -60,7 +60,10 @@ traceable units of work:
|
|||
instead (see below).
|
||||
|
||||
Spans are named `"{service} {call_type}"` (e.g. `"redis set"`) so repeated calls
|
||||
to one service stay distinguishable. Like every other span they parent to the
|
||||
to one service stay distinguishable. `call_type` is the operation only; the
|
||||
litellm call chain that issued it (`async_set_cache <- async_add_cache`) travels
|
||||
as `ServiceLoggerPayload.caller` and lands on the `litellm.service.caller`
|
||||
attribute, so one operation is one span name. Like every other span they parent to the
|
||||
**ambient** context, falling back to the threaded `litellm_parent_otel_span` only
|
||||
when ambient has no live span; a background job with neither starts its own root
|
||||
trace.
|
||||
|
|
@ -69,16 +72,21 @@ trace.
|
|||
and the spend-counter increment all run after the response is on the wire, so they
|
||||
add nothing to the request's latency. Parenting them under the (already ended)
|
||||
server span stretched the request trace past the request itself, which is what a
|
||||
viewer shows as trace duration. `context.resolve_service_span_context` compares
|
||||
the call's end time with the resolved parent's end time: a call that finished
|
||||
after its parent ended starts a **new root trace** carrying a **span link** back
|
||||
to the request span (the `FollowsFrom` relationship of OpenTracing; the default
|
||||
`:link` propagation style of the OTel Ruby ActiveJob and Sidekiq
|
||||
instrumentations). Identity Baggage still rides along, so the detached span keeps
|
||||
its team / key / user attributes. Only an SDK span that has really ended detaches:
|
||||
a sampled-out or remote `NonRecordingSpan` is never recording but is still the
|
||||
right parent. A call that ended before the server span did stays a child even when
|
||||
its `asyncio.create_task`-dispatched hook runs after the response.
|
||||
viewer shows as trace duration. `context.resolve_service_span_context` detaches
|
||||
a call in two cases: it was logged from the post-response phase
|
||||
(`litellm._internal_context.post_response_phase`, entered by the success
|
||||
handlers and by the response-cache write task, inherited by every task spawned
|
||||
inside), or it finished after the resolved parent ended. Either way it starts a
|
||||
**new root trace** carrying a **span link** back to the request span (the
|
||||
`FollowsFrom` relationship of OpenTracing; the default `:link` propagation style
|
||||
of the OTel Ruby ActiveJob and Sidekiq instrumentations). The phase check matters
|
||||
for streaming: the stream-finished callbacks run before the ASGI server span
|
||||
closes, so by end time alone the cache write would look like request latency.
|
||||
Identity Baggage still rides along, so the detached span keeps its team / key /
|
||||
user attributes. Only an SDK span detaches: a sampled-out or remote
|
||||
`NonRecordingSpan` is never recording but is still the right parent. A call that
|
||||
ended before the server span did stays a child even when its
|
||||
`asyncio.create_task`-dispatched hook runs after the response.
|
||||
|
||||
Caller-supplied `event_metadata` is **sanitized** before it reaches a span
|
||||
(primitives only, no live objects, no secrets/headers, bounded) — see
|
||||
|
|
|
|||
|
|
@ -3,6 +3,7 @@
|
|||
from collections import OrderedDict
|
||||
from collections.abc import Callable, Iterator, Mapping, Sequence
|
||||
from contextlib import contextmanager
|
||||
from dataclasses import replace
|
||||
from datetime import datetime
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Final, cast
|
||||
|
|
@ -661,12 +662,7 @@ class OpenTelemetryV2(CustomLogger):
|
|||
if error_override is None and start_time is None and end_time is None and parent_otel_span is None:
|
||||
return None
|
||||
if error_override is not None and data.error is None:
|
||||
data = ServiceSpanData(
|
||||
service_name=data.service_name,
|
||||
call_type=data.call_type,
|
||||
error=SpanError(message=error_override),
|
||||
event_metadata=data.event_metadata,
|
||||
)
|
||||
data = replace(data, error=SpanError(message=error_override))
|
||||
# Parent like every other span: ambient context first (so identity Baggage
|
||||
# rides along and the call nests under whatever request phase is active —
|
||||
# e.g. a DB lookup under the live ``auth`` span), falling back to the
|
||||
|
|
|
|||
|
|
@ -148,6 +148,7 @@ class GenAIMapper:
|
|||
_SERVICE_ATTRS: dict[str, Callable[[ServiceSpanData], AttrValue | None]] = {
|
||||
LiteLLM.SERVICE_NAME: lambda d: d.service_name,
|
||||
LiteLLM.SERVICE_CALL_TYPE: lambda d: d.call_type,
|
||||
LiteLLM.SERVICE_CALLER: lambda d: d.caller,
|
||||
}
|
||||
|
||||
def __init__(self, tool_attr_budget: int = MAX_TOOL_DEFINITION_ATTRS_PER_SPAN) -> None:
|
||||
|
|
|
|||
|
|
@ -37,6 +37,7 @@ _LEGACY_PRESENCE_PENALTY: Final = "llm.presence_penalty"
|
|||
_LEGACY_STOP_SEQUENCES: Final = "llm.chat.stop_sequences"
|
||||
_LEGACY_SERVICE: Final = "service"
|
||||
_LEGACY_CALL_TYPE: Final = "call_type"
|
||||
_LEGACY_CALLER: Final = "caller"
|
||||
_LEGACY_ERROR: Final = Error.MESSAGE_LEGACY
|
||||
|
||||
|
||||
|
|
@ -66,6 +67,7 @@ class LegacyMapper:
|
|||
_SERVICE_ATTRS: dict[str, Callable[[ServiceSpanData], AttrValue | None]] = {
|
||||
_LEGACY_SERVICE: lambda d: d.service_name,
|
||||
_LEGACY_CALL_TYPE: lambda d: d.call_type,
|
||||
_LEGACY_CALLER: lambda d: d.caller,
|
||||
_LEGACY_ERROR: lambda d: d.error.message if d.error is not None and d.error.message else None,
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -41,6 +41,7 @@ class ExporterOwner(str, Enum):
|
|||
LEVO = "levo"
|
||||
AGENTOPS = "agentops"
|
||||
NEWRELIC = "newrelic"
|
||||
SIGNOZ = "signoz"
|
||||
|
||||
|
||||
class _OTelV2Flag(BaseSettings):
|
||||
|
|
|
|||
|
|
@ -309,6 +309,7 @@ class GuardrailSpanData:
|
|||
class ServiceSpanData:
|
||||
service_name: str
|
||||
call_type: str | None = None
|
||||
caller: str | None = None
|
||||
error: SpanError | None = None
|
||||
# Caller-supplied attributes to stamp on the service span, passed through
|
||||
# from ``async_service_*_hook(event_metadata=...)``. The mapper owns how
|
||||
|
|
@ -330,6 +331,7 @@ class ServiceSpanData:
|
|||
return cls(
|
||||
service_name=payload.service.value,
|
||||
call_type=payload.call_type,
|
||||
caller=payload.caller,
|
||||
error=SpanError(message=payload.error) if payload.error else None,
|
||||
event_metadata=sanitize_event_metadata(event_metadata),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -326,6 +326,7 @@ class LiteLLM:
|
|||
GUARDRAIL_COST_IN_SPEND: Final = "litellm.guardrail.cost_in_spend"
|
||||
SERVICE_NAME: Final = "litellm.service.name"
|
||||
SERVICE_CALL_TYPE: Final = "litellm.service.call_type"
|
||||
SERVICE_CALLER: Final = "litellm.service.caller"
|
||||
PREPROCESSING_MS: Final = "litellm.preprocessing.duration_ms"
|
||||
# The logical name of the MCP server a tool call was routed to. There is no
|
||||
# semconv key for an MCP server's *name* (the convention uses ``server.address``
|
||||
|
|
|
|||
|
|
@ -21,6 +21,7 @@ from opentelemetry.trace.propagation.tracecontext import (
|
|||
TraceContextTextMapPropagator,
|
||||
)
|
||||
|
||||
from litellm._internal_context import in_post_response_phase
|
||||
from litellm.integrations.otel.model.semconv import HTTP
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
|
@ -231,21 +232,28 @@ def resolve_service_span_context(
|
|||
) -> tuple[Context, tuple[Link, ...]]:
|
||||
"""Parent context + links for a service/DB span that ended at ``end_time_ns``.
|
||||
|
||||
A call that finished after its parent ended (post-response spend tracking)
|
||||
starts its own root trace with a span link back to the parent instead of
|
||||
stretching the parent's trace. Baggage stays on the returned context.
|
||||
Work the caller did not wait for starts its own root trace with a span link
|
||||
back to the parent instead of stretching the parent's trace: anything logged
|
||||
from the post-response phase (success callbacks, the response-cache write,
|
||||
see :func:`litellm._internal_context.post_response_phase`), whether or not
|
||||
the server span has closed yet, and anything that finished after its parent
|
||||
ended. Baggage stays on the returned context.
|
||||
"""
|
||||
ctx: Final = resolve_parent_context(threaded)
|
||||
parent: Final = get_current_span(ctx)
|
||||
if not _ended_before(parent, end_time_ns):
|
||||
if not _is_post_response(parent, end_time_ns):
|
||||
return ctx, ()
|
||||
return set_span_in_context(INVALID_SPAN, ctx), (Link(parent.get_span_context()),)
|
||||
|
||||
|
||||
def _ended_before(span: Span, end_time_ns: int | None) -> bool:
|
||||
if not isinstance(span, ReadableSpan) or span.end_time is None:
|
||||
def _is_post_response(parent: Span, end_time_ns: int | None) -> bool:
|
||||
if not isinstance(parent, ReadableSpan):
|
||||
return False
|
||||
return end_time_ns is None or end_time_ns > span.end_time
|
||||
if in_post_response_phase():
|
||||
return True
|
||||
if parent.end_time is None:
|
||||
return False
|
||||
return end_time_ns is None or end_time_ns > parent.end_time
|
||||
|
||||
|
||||
def resolve_request_span_context() -> Context:
|
||||
|
|
|
|||
|
|
@ -30,6 +30,11 @@ from litellm.integrations.otel.presets.phoenix import (
|
|||
phoenix_preset,
|
||||
phoenix_project_headers,
|
||||
)
|
||||
from litellm.integrations.otel.presets.signoz import (
|
||||
signoz_dynamic_endpoint,
|
||||
signoz_dynamic_headers,
|
||||
signoz_preset,
|
||||
)
|
||||
from litellm.integrations.otel.presets.weave import weave_dynamic_headers, weave_preset
|
||||
from litellm.types.utils import StandardCallbackDynamicParams
|
||||
|
||||
|
|
@ -44,6 +49,7 @@ PRESET_BY_CALLBACK: Final[Mapping[str, Preset]] = MappingProxyType(
|
|||
"langtrace": langtrace_preset,
|
||||
"levo": levo_preset,
|
||||
"newrelic": newrelic_preset,
|
||||
"signoz": signoz_preset,
|
||||
"weave_otel": weave_preset,
|
||||
}
|
||||
)
|
||||
|
|
@ -58,6 +64,7 @@ DYNAMIC_HEADERS_BY_CALLBACK: Final[Mapping[str, Callable[[StandardCallbackDynami
|
|||
"arize": arize_dynamic_headers,
|
||||
"langfuse_otel": langfuse_dynamic_headers,
|
||||
"newrelic": newrelic_dynamic_headers,
|
||||
"signoz": signoz_dynamic_headers,
|
||||
"weave_otel": weave_dynamic_headers,
|
||||
}
|
||||
)
|
||||
|
|
@ -71,6 +78,7 @@ DYNAMIC_ENDPOINT_BY_CALLBACK: Final[Mapping[str, Callable[[StandardCallbackDynam
|
|||
MappingProxyType(
|
||||
{
|
||||
"newrelic": newrelic_dynamic_endpoint,
|
||||
"signoz": signoz_dynamic_endpoint,
|
||||
}
|
||||
)
|
||||
)
|
||||
|
|
@ -153,5 +161,6 @@ __all__ = [
|
|||
"newrelic_preset",
|
||||
"phoenix_preset",
|
||||
"project_routing_headers",
|
||||
"signoz_preset",
|
||||
"weave_preset",
|
||||
]
|
||||
|
|
|
|||
95
litellm/integrations/otel/presets/signoz.py
Normal file
95
litellm/integrations/otel/presets/signoz.py
Normal file
|
|
@ -0,0 +1,95 @@
|
|||
from functools import lru_cache
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
|
||||
from pydantic import Field
|
||||
from pydantic_settings import BaseSettings, SettingsConfigDict
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.integrations.otel.model.config import (
|
||||
ExporterOwner,
|
||||
ExporterSpec,
|
||||
OpenTelemetryV2Config,
|
||||
)
|
||||
from litellm.integrations.otel.presets.utils import ensure_mappers
|
||||
from litellm.litellm_core_utils.url_utils import is_url_destination_allowed_by_host
|
||||
from litellm.types.utils import StandardCallbackDynamicParams
|
||||
|
||||
SIGNOZ_INGESTION_ENDPOINT_ENV: Final = "SIGNOZ_INGESTION_ENDPOINT"
|
||||
|
||||
|
||||
class _SigNozSettings(BaseSettings):
|
||||
model_config = SettingsConfigDict(case_sensitive=False, extra="ignore")
|
||||
|
||||
endpoint: str | None = Field(default=None, validation_alias=SIGNOZ_INGESTION_ENDPOINT_ENV)
|
||||
ingestion_key: str | None = Field(default=None, validation_alias="SIGNOZ_INGESTION_KEY")
|
||||
|
||||
|
||||
def signoz_preset(
|
||||
*,
|
||||
config_overrides: OpenTelemetryV2Config | None = None,
|
||||
allow_missing_credentials: bool = False,
|
||||
) -> OpenTelemetryV2Config:
|
||||
settings: Final = _SigNozSettings()
|
||||
base: Final = config_overrides or OpenTelemetryV2Config()
|
||||
key: Final = settings.ingestion_key
|
||||
spec: Final = ExporterSpec(
|
||||
kind="otlp_http",
|
||||
endpoint=settings.endpoint,
|
||||
headers=(f"signoz-ingestion-key={key}" if key else None),
|
||||
owner=ExporterOwner.SIGNOZ,
|
||||
requires_headers=bool(key),
|
||||
)
|
||||
return base.model_copy(
|
||||
update=MappingProxyType(
|
||||
{
|
||||
"exporters": (*base.exporters, spec),
|
||||
"mapper_names": ensure_mappers(base.mapper_names, "genai"),
|
||||
}
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
@lru_cache(maxsize=128)
|
||||
def _warn_host_not_allowlisted(endpoint: str) -> None:
|
||||
verbose_logger.warning(
|
||||
"SigNoz: not exporting to key/team endpoint '%s'. Add its host to "
|
||||
"litellm_settings.provider_url_destination_allowed_hosts to permit it",
|
||||
endpoint,
|
||||
)
|
||||
|
||||
|
||||
@lru_cache(maxsize=128)
|
||||
def _warn_endpoint_without_key(endpoint: str) -> None:
|
||||
verbose_logger.warning(
|
||||
"SigNoz: not exporting to key/team endpoint '%s'. Set signoz_ingestion_key alongside it; "
|
||||
"a keyless collector needs the global callback",
|
||||
endpoint,
|
||||
)
|
||||
|
||||
|
||||
def _tenant_endpoint_is_unusable(params: StandardCallbackDynamicParams) -> bool:
|
||||
return bool(params.get("signoz_ingestion_endpoint")) and signoz_dynamic_endpoint(params) is None
|
||||
|
||||
|
||||
def signoz_dynamic_endpoint(params: StandardCallbackDynamicParams) -> str | None:
|
||||
endpoint: Final = params.get("signoz_ingestion_endpoint")
|
||||
if not endpoint or not endpoint.startswith(("http://", "https://")):
|
||||
return None
|
||||
if not params.get("signoz_ingestion_key"):
|
||||
_warn_endpoint_without_key(endpoint)
|
||||
return None
|
||||
if not is_url_destination_allowed_by_host(endpoint, litellm.provider_url_destination_allowed_hosts):
|
||||
_warn_host_not_allowlisted(endpoint)
|
||||
return None
|
||||
return endpoint
|
||||
|
||||
|
||||
def signoz_dynamic_headers(
|
||||
params: StandardCallbackDynamicParams,
|
||||
) -> dict[str, str]: # mutable-ok: DYNAMIC_HEADERS_BY_CALLBACK returns a dict
|
||||
key: Final = params.get("signoz_ingestion_key")
|
||||
if _tenant_endpoint_is_unusable(params) or not key:
|
||||
return {} # mutable-ok: same registry contract
|
||||
return {"signoz-ingestion-key": key} # mutable-ok: same registry contract
|
||||
|
|
@ -36,22 +36,86 @@ def resolve_s3_log_prompts_only(configured: object, environ: Mapping[str, str] |
|
|||
return True
|
||||
|
||||
|
||||
def resolve_s3_max_concurrent_uploads(configured: object, fallback: int) -> int:
|
||||
def _resolve_positive_int(setting: str, configured: object, fallback: int, *, reject_bool: bool) -> int:
|
||||
if configured is None or configured == "":
|
||||
return fallback
|
||||
if reject_bool and isinstance(configured, bool):
|
||||
verbose_logger.warning(
|
||||
"s3 logging: %s=%r is a boolean, not an integer, using %s", setting, configured, fallback
|
||||
)
|
||||
return fallback
|
||||
try:
|
||||
bound: Final = _UPLOAD_BOUND.validate_python(configured.strip() if isinstance(configured, str) else configured)
|
||||
except ValidationError:
|
||||
verbose_logger.warning("s3 logging: %s=%r is not an integer, using %s", setting, configured, fallback)
|
||||
return fallback
|
||||
if bound < 1:
|
||||
verbose_logger.warning("s3 logging: %s=%r must be at least 1, using %s", setting, configured, fallback)
|
||||
return fallback
|
||||
return bound
|
||||
|
||||
|
||||
def resolve_s3_max_concurrent_uploads(configured: object, fallback: int) -> int:
|
||||
return _resolve_positive_int("s3_max_concurrent_uploads", configured, fallback, reject_bool=False)
|
||||
|
||||
|
||||
def resolve_s3_max_queue_size(configured: object, fallback: int) -> int:
|
||||
return _resolve_positive_int("s3_max_queue_size", configured, fallback, reject_bool=True)
|
||||
|
||||
|
||||
def resolve_s3_max_retry_age_seconds(configured: object, fallback: int | None) -> int | None:
|
||||
if configured is None or configured == "":
|
||||
return None
|
||||
if isinstance(configured, bool):
|
||||
verbose_logger.warning(
|
||||
"s3 logging: s3_max_retry_age_seconds=%r is a boolean, not an integer, falling back to %r",
|
||||
configured,
|
||||
fallback,
|
||||
)
|
||||
return fallback
|
||||
try:
|
||||
bound: Final = _UPLOAD_BOUND.validate_python(configured.strip() if isinstance(configured, str) else configured)
|
||||
except ValidationError:
|
||||
verbose_logger.warning(
|
||||
"s3 logging: s3_max_concurrent_uploads=%r is not an integer, using %s", configured, fallback
|
||||
"s3 logging: s3_max_retry_age_seconds=%r is not an integer, falling back to %r", configured, fallback
|
||||
)
|
||||
return fallback
|
||||
if bound < 1:
|
||||
if bound < 0:
|
||||
verbose_logger.warning(
|
||||
"s3 logging: s3_max_concurrent_uploads=%r must be at least 1, using %s", configured, fallback
|
||||
"s3 logging: s3_max_retry_age_seconds=%r must be at least 0, falling back to %r", configured, fallback
|
||||
)
|
||||
return fallback
|
||||
return bound
|
||||
return bound or None
|
||||
|
||||
|
||||
def resolve_s3_max_adaptive_concurrency(configured: object, fallback: int) -> int:
|
||||
return _resolve_positive_int("s3_max_adaptive_concurrency", configured, fallback, reject_bool=True)
|
||||
|
||||
|
||||
def resolve_s3_drop_on_terminal_error(configured: object) -> bool:
|
||||
if configured is None or configured == "":
|
||||
return True
|
||||
try:
|
||||
return _S3_BOOL.validate_python(configured.strip() if isinstance(configured, str) else configured)
|
||||
except ValidationError:
|
||||
verbose_logger.warning(
|
||||
"s3 logging: s3_drop_on_terminal_error=%r is not a boolean, dropping terminal-failed uploads",
|
||||
configured,
|
||||
)
|
||||
return True
|
||||
|
||||
|
||||
def resolve_s3_adaptive_concurrency(configured: object) -> bool:
|
||||
if configured is None or configured == "":
|
||||
return False
|
||||
try:
|
||||
return _S3_BOOL.validate_python(configured.strip() if isinstance(configured, str) else configured)
|
||||
except ValidationError:
|
||||
verbose_logger.warning(
|
||||
"s3 logging: s3_adaptive_concurrency=%r is not a boolean, keeping the fixed upload width",
|
||||
configured,
|
||||
)
|
||||
return False
|
||||
|
||||
|
||||
def resolve_s3_batch_file_upload(configured: object) -> bool:
|
||||
|
|
|
|||
|
|
@ -3,14 +3,19 @@ s3 Bucket Logging Integration
|
|||
|
||||
async_log_success_event: Processes the event, stores it in memory for DEFAULT_S3_FLUSH_INTERVAL_SECONDS seconds or until DEFAULT_S3_BATCH_SIZE and then flushes to s3
|
||||
async_log_failure_event: Processes the event, stores it in memory for DEFAULT_S3_FLUSH_INTERVAL_SECONDS seconds or until DEFAULT_S3_BATCH_SIZE and then flushes to s3
|
||||
NOTE 1: S3 does not provide a BATCH PUT API endpoint; by default each element is uploaded concurrently (bounded by s3_max_concurrent_uploads), or with s3_batch_file_upload the whole flush is written as one .jsonl file
|
||||
NOTE 1: S3 does not provide a BATCH PUT API endpoint; by default each element is uploaded concurrently with the fixed s3_max_concurrent_uploads bound (or an adaptive bound when s3_adaptive_concurrency is on, backing off only on throttling), or with s3_batch_file_upload the whole flush is written as one .jsonl file
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import contextvars
|
||||
import logging
|
||||
import re
|
||||
import time
|
||||
from collections.abc import Mapping
|
||||
from collections.abc import Awaitable, Callable, Mapping
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timezone
|
||||
from typing import TYPE_CHECKING, Final, cast
|
||||
from functools import partial
|
||||
from typing import TYPE_CHECKING, Final, Literal, cast
|
||||
from urllib.parse import quote
|
||||
from uuid import uuid4
|
||||
|
||||
|
|
@ -21,15 +26,22 @@ from litellm._logging import print_verbose, verbose_logger
|
|||
from litellm.constants import (
|
||||
DEFAULT_S3_BATCH_SIZE,
|
||||
DEFAULT_S3_FLUSH_INTERVAL_SECONDS,
|
||||
DEFAULT_S3_MAX_ADAPTIVE_CONCURRENCY,
|
||||
DEFAULT_S3_MAX_CONCURRENT_UPLOADS,
|
||||
)
|
||||
from litellm.integrations.adaptive_concurrency import AdaptiveConcurrencyLimiter, PutSample
|
||||
from litellm.integrations.s3 import (
|
||||
get_s3_object_download_filename,
|
||||
get_s3_object_key,
|
||||
prompts_only_payload,
|
||||
resolve_s3_adaptive_concurrency,
|
||||
resolve_s3_batch_file_upload,
|
||||
resolve_s3_drop_on_terminal_error,
|
||||
resolve_s3_log_prompts_only,
|
||||
resolve_s3_max_adaptive_concurrency,
|
||||
resolve_s3_max_concurrent_uploads,
|
||||
resolve_s3_max_queue_size,
|
||||
resolve_s3_max_retry_age_seconds,
|
||||
resolve_sse_params,
|
||||
)
|
||||
from litellm.litellm_core_utils.aws_partition import get_aws_dns_suffix
|
||||
|
|
@ -50,6 +62,42 @@ if TYPE_CHECKING:
|
|||
from botocore.credentials import Credentials
|
||||
|
||||
|
||||
UploadOutcome = Literal["delivered", "retry", "dropped"]
|
||||
|
||||
_TERMINAL_ERROR_CODES: Final = frozenset(
|
||||
{
|
||||
"EntityTooLarge",
|
||||
"InvalidArgument",
|
||||
"MalformedXML",
|
||||
"InvalidDigest",
|
||||
"KeyTooLongError",
|
||||
"BadDigest",
|
||||
"InvalidRequest",
|
||||
}
|
||||
)
|
||||
_BODY_CODED_STATUSES: Final = frozenset({400, 403})
|
||||
_RETRYABLE_STATUSES: Final = frozenset({403, 500, 503})
|
||||
_S3_ERROR_CODE: Final = re.compile(r"<Code>([^<]+)</Code>")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _PreparedPut:
|
||||
json_string: str
|
||||
headers: Mapping[str, str]
|
||||
|
||||
|
||||
def _s3_error_code(response: httpx.Response) -> str | None:
|
||||
text: Final = response.text
|
||||
match: Final = _S3_ERROR_CODE.search(text) if isinstance(text, str) else None
|
||||
return match.group(1) if match else None
|
||||
|
||||
|
||||
def _is_terminal(response: httpx.Response) -> bool:
|
||||
"""True only for object-specific, unrecoverable rejections (400/403 with a terminal XML code).
|
||||
Unknown codes, empty or non-XML bodies, and every other status fail safe toward retry."""
|
||||
return response.status_code in _BODY_CODED_STATUSES and _s3_error_code(response) in _TERMINAL_ERROR_CODES
|
||||
|
||||
|
||||
def _s3_key_parent(s3_object_key: str) -> str:
|
||||
return s3_object_key.rsplit("/", 1)[0] if "/" in s3_object_key else ""
|
||||
|
||||
|
|
@ -58,11 +106,19 @@ class S3BatchUploadError(Exception):
|
|||
def __init__(self, failed: int, total: int) -> None:
|
||||
self.failed = failed
|
||||
self.total = total
|
||||
super().__init__(f"{failed} of {total} S3 uploads failed; events kept in queue for the next flush")
|
||||
super().__init__(f"{failed} of {total} S3 uploads failed; transient failures kept in queue for the next flush")
|
||||
|
||||
|
||||
_in_flush: Final[contextvars.ContextVar[bool]] = contextvars.ContextVar("s3_v2_in_flush", default=False)
|
||||
|
||||
|
||||
class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
||||
preserve_events_added_during_flush = True
|
||||
_flush_retries: int = 0
|
||||
_requeued_count: int = 0
|
||||
_upload_limiter: asyncio.Semaphore | AdaptiveConcurrencyLimiter | None = None
|
||||
s3_drop_on_terminal_error: bool = True
|
||||
s3_max_retry_age_seconds: int | None = 3600
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
|
|
@ -92,6 +148,11 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
s3_sse_kms_key_id: str | None = None,
|
||||
s3_log_prompts_only: bool | None = None,
|
||||
s3_max_concurrent_uploads: int = DEFAULT_S3_MAX_CONCURRENT_UPLOADS,
|
||||
s3_max_queue_size: int | None = None,
|
||||
s3_max_retry_age_seconds: int | None = 3600,
|
||||
s3_drop_on_terminal_error: bool = True,
|
||||
s3_adaptive_concurrency: bool = False,
|
||||
s3_max_adaptive_concurrency: int | None = None,
|
||||
s3_batch_file_upload: bool = False,
|
||||
s3_callback_params_override: dict | None = None,
|
||||
**kwargs,
|
||||
|
|
@ -135,9 +196,22 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
s3_sse_kms_key_id=s3_sse_kms_key_id,
|
||||
s3_log_prompts_only=s3_log_prompts_only,
|
||||
s3_max_concurrent_uploads=s3_max_concurrent_uploads,
|
||||
s3_max_queue_size=s3_max_queue_size,
|
||||
s3_max_retry_age_seconds=s3_max_retry_age_seconds,
|
||||
s3_drop_on_terminal_error=s3_drop_on_terminal_error,
|
||||
s3_adaptive_concurrency=s3_adaptive_concurrency,
|
||||
s3_max_adaptive_concurrency=s3_max_adaptive_concurrency,
|
||||
s3_batch_file_upload=s3_batch_file_upload,
|
||||
)
|
||||
self._upload_semaphore = asyncio.Semaphore(self.s3_max_concurrent_uploads)
|
||||
self._upload_limiter = (
|
||||
AdaptiveConcurrencyLimiter(
|
||||
initial=self.s3_max_concurrent_uploads,
|
||||
floor=self.s3_max_concurrent_uploads,
|
||||
ceiling=max(self.s3_max_concurrent_uploads, self.s3_max_adaptive_concurrency),
|
||||
)
|
||||
if self.s3_adaptive_concurrency
|
||||
else asyncio.Semaphore(self.s3_max_concurrent_uploads)
|
||||
)
|
||||
verbose_logger.debug("s3 logger using endpoint url %s", s3_endpoint_url)
|
||||
|
||||
# IMPORTANT
|
||||
|
|
@ -158,8 +232,12 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
flush_lock=self.flush_lock,
|
||||
flush_interval=s3_flush_interval,
|
||||
batch_size=s3_batch_size,
|
||||
max_queue_size=self.s3_max_queue_size,
|
||||
)
|
||||
self.log_queue: list[s3BatchLoggingElement] = []
|
||||
self._requeued_count = 0
|
||||
self._flush_retries = 0
|
||||
self._flush_dropped: dict[int, s3BatchLoggingElement] = {}
|
||||
|
||||
# Call BaseAWSLLM's __init__
|
||||
BaseAWSLLM.__init__(self)
|
||||
|
|
@ -194,6 +272,11 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
s3_sse_kms_key_id: str | None = None,
|
||||
s3_log_prompts_only: bool | None = None,
|
||||
s3_max_concurrent_uploads: int = DEFAULT_S3_MAX_CONCURRENT_UPLOADS,
|
||||
s3_max_queue_size: int | None = None,
|
||||
s3_max_retry_age_seconds: int | None = 3600,
|
||||
s3_drop_on_terminal_error: bool = True,
|
||||
s3_adaptive_concurrency: bool = False,
|
||||
s3_max_adaptive_concurrency: int | None = None,
|
||||
s3_batch_file_upload: bool = False,
|
||||
params_source: dict | None = None,
|
||||
):
|
||||
|
|
@ -259,6 +342,37 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
DEFAULT_S3_MAX_CONCURRENT_UPLOADS,
|
||||
)
|
||||
|
||||
configured_queue_size: Final = params.get("s3_max_queue_size")
|
||||
constructor_queue_size: Final = resolve_s3_max_queue_size(
|
||||
s3_max_queue_size, CustomBatchLogger.DEFAULT_MAX_QUEUE_SIZE
|
||||
)
|
||||
self.s3_max_queue_size = resolve_s3_max_queue_size(configured_queue_size, constructor_queue_size)
|
||||
|
||||
configured_retry_age: Final = params.get("s3_max_retry_age_seconds")
|
||||
constructor_retry_age: Final = resolve_s3_max_retry_age_seconds(s3_max_retry_age_seconds, 3600)
|
||||
self.s3_max_retry_age_seconds = (
|
||||
constructor_retry_age
|
||||
if configured_retry_age is None or configured_retry_age == ""
|
||||
else resolve_s3_max_retry_age_seconds(configured_retry_age, constructor_retry_age)
|
||||
)
|
||||
|
||||
configured_drop: Final = params.get("s3_drop_on_terminal_error")
|
||||
self.s3_drop_on_terminal_error = resolve_s3_drop_on_terminal_error(
|
||||
configured_drop if configured_drop is not None else s3_drop_on_terminal_error
|
||||
)
|
||||
|
||||
self.s3_adaptive_concurrency = s3_adaptive_concurrency or resolve_s3_adaptive_concurrency(
|
||||
params.get("s3_adaptive_concurrency")
|
||||
)
|
||||
|
||||
configured_adaptive_ceiling: Final = params.get("s3_max_adaptive_concurrency")
|
||||
self.s3_max_adaptive_concurrency = resolve_s3_max_adaptive_concurrency(
|
||||
s3_max_adaptive_concurrency
|
||||
if configured_adaptive_ceiling is None or configured_adaptive_ceiling == ""
|
||||
else configured_adaptive_ceiling,
|
||||
DEFAULT_S3_MAX_ADAPTIVE_CONCURRENCY,
|
||||
)
|
||||
|
||||
self.s3_batch_file_upload = s3_batch_file_upload or resolve_s3_batch_file_upload(
|
||||
params.get("s3_batch_file_upload")
|
||||
)
|
||||
|
|
@ -310,6 +424,35 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
}
|
||||
return {key: value for key, value in candidates.items() if value}
|
||||
|
||||
def _prepare_put(self, batch_logging_element: s3BatchLoggingElement) -> _PreparedPut:
|
||||
try:
|
||||
import base64
|
||||
import hashlib
|
||||
except ImportError:
|
||||
raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.")
|
||||
|
||||
json_string: Final = (
|
||||
batch_logging_element.body
|
||||
if batch_logging_element.body is not None
|
||||
else safe_dumps(batch_logging_element.payload)
|
||||
)
|
||||
content_hash: Final = hashlib.sha256(json_string.encode("utf-8")).hexdigest()
|
||||
content_md5: Final = base64.b64encode(
|
||||
hashlib.md5(json_string.encode("utf-8"), usedforsecurity=False).digest()
|
||||
).decode()
|
||||
return _PreparedPut(
|
||||
json_string=json_string,
|
||||
headers={
|
||||
"Content-Type": batch_logging_element.content_type,
|
||||
"Content-MD5": content_md5,
|
||||
"x-amz-content-sha256": content_hash,
|
||||
"Content-Language": "en",
|
||||
"Content-Disposition": f'inline; filename="{batch_logging_element.s3_object_download_filename}"',
|
||||
"Cache-Control": "private, immutable, max-age=31536000, s-maxage=0",
|
||||
**self._sse_headers(),
|
||||
},
|
||||
)
|
||||
|
||||
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
await self._async_log_event_base(
|
||||
kwargs=kwargs,
|
||||
|
|
@ -384,12 +527,21 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
verbose_logger.exception("s3 Layer Error - %s", e)
|
||||
self.handle_callback_failure(callback_name="S3Logger")
|
||||
|
||||
async def async_upload_data_to_s3(self, batch_logging_element: s3BatchLoggingElement) -> bool:
|
||||
try:
|
||||
import base64
|
||||
import hashlib
|
||||
except ImportError:
|
||||
raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.")
|
||||
@property
|
||||
def _upload_semaphore(self) -> asyncio.Semaphore | AdaptiveConcurrencyLimiter:
|
||||
limiter: Final = self._upload_limiter
|
||||
if limiter is None:
|
||||
raise AttributeError("_upload_semaphore")
|
||||
return limiter
|
||||
|
||||
@_upload_semaphore.setter
|
||||
def _upload_semaphore(self, value: asyncio.Semaphore | AdaptiveConcurrencyLimiter) -> None:
|
||||
self._upload_limiter = value
|
||||
|
||||
async def async_upload_data_to_s3(
|
||||
self,
|
||||
batch_logging_element: s3BatchLoggingElement,
|
||||
) -> bool:
|
||||
try:
|
||||
from litellm.litellm_core_utils.asyncify import asyncify
|
||||
|
||||
|
|
@ -400,31 +552,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
|
||||
url: Final = self._build_object_url(batch_logging_element.s3_object_key)
|
||||
|
||||
# Convert JSON to string
|
||||
json_string: Final = (
|
||||
batch_logging_element.body
|
||||
if batch_logging_element.body is not None
|
||||
else safe_dumps(batch_logging_element.payload)
|
||||
)
|
||||
|
||||
# Calculate SHA256 hash of the content
|
||||
content_hash: Final = hashlib.sha256(json_string.encode("utf-8")).hexdigest()
|
||||
content_md5: Final = base64.b64encode(
|
||||
hashlib.md5(json_string.encode("utf-8"), usedforsecurity=False).digest()
|
||||
).decode()
|
||||
|
||||
# Prepare the request
|
||||
headers: Final = {
|
||||
"Content-Type": batch_logging_element.content_type,
|
||||
"Content-MD5": content_md5,
|
||||
"x-amz-content-sha256": content_hash,
|
||||
"Content-Language": "en",
|
||||
"Content-Disposition": f'inline; filename="{batch_logging_element.s3_object_download_filename}"',
|
||||
"Cache-Control": "private, immutable, max-age=31536000, s-maxage=0",
|
||||
**self._sse_headers(),
|
||||
}
|
||||
|
||||
async def signed_put() -> httpx.Response:
|
||||
async def signed_put(prepared: _PreparedPut) -> httpx.Response:
|
||||
credentials: Final = await asyncified_get_credentials(
|
||||
aws_access_key_id=self.s3_aws_access_key_id,
|
||||
aws_secret_access_key=self.s3_aws_secret_access_key,
|
||||
|
|
@ -436,18 +564,26 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
aws_web_identity_token=self.s3_aws_web_identity_token,
|
||||
aws_sts_endpoint=self.s3_aws_sts_endpoint,
|
||||
)
|
||||
signed_headers: Final = await run_aws_signing(self._sign_put, credentials, url, json_string, headers)
|
||||
signed_headers: Final = await run_aws_signing(
|
||||
self._sign_put, credentials, url, prepared.json_string, prepared.headers
|
||||
)
|
||||
try:
|
||||
return await self.async_httpx_client.put(url, data=json_string, headers=signed_headers)
|
||||
return await self.async_httpx_client.put(url, data=prepared.json_string, headers=signed_headers)
|
||||
except httpx.HTTPStatusError as error:
|
||||
return error.response
|
||||
|
||||
max_retries: Final = 3
|
||||
prepared: Final = self._prepare_put(batch_logging_element)
|
||||
for attempt in range(max_retries):
|
||||
response = await signed_put()
|
||||
if response.status_code in (403, 500, 503) and attempt < max_retries - 1:
|
||||
response = await self._recorded_put(partial(signed_put, prepared))
|
||||
if (
|
||||
response.status_code in _RETRYABLE_STATUSES
|
||||
and not (self.s3_drop_on_terminal_error and _is_terminal(response))
|
||||
and attempt < max_retries - 1
|
||||
):
|
||||
wait_time = 2**attempt # 1s, 2s
|
||||
verbose_logger.warning(
|
||||
verbose_logger.log(
|
||||
logging.DEBUG if _in_flush.get() else logging.WARNING,
|
||||
"S3 upload returned %s, retrying in %ss (attempt %s/%s) key=%s",
|
||||
response.status_code,
|
||||
wait_time,
|
||||
|
|
@ -455,6 +591,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
max_retries,
|
||||
batch_logging_element.s3_object_key,
|
||||
)
|
||||
self._flush_retries += 1
|
||||
await asyncio.sleep(wait_time)
|
||||
continue
|
||||
response.raise_for_status()
|
||||
|
|
@ -462,6 +599,13 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
except Exception as e:
|
||||
verbose_logger.exception("Error uploading to s3: %s", e)
|
||||
self.handle_callback_failure(callback_name="S3Logger")
|
||||
if isinstance(e, httpx.HTTPStatusError) and self.s3_drop_on_terminal_error and _is_terminal(e.response):
|
||||
verbose_logger.warning(
|
||||
"s3 logging: dropping object %s after terminal status %s",
|
||||
batch_logging_element.s3_object_key,
|
||||
e.response.status_code,
|
||||
)
|
||||
self._flush_dropped[id(batch_logging_element)] = batch_logging_element
|
||||
return False
|
||||
return True
|
||||
|
||||
|
|
@ -483,11 +627,64 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
# see custom_batch_logger.py which triggers the flush
|
||||
#########################################################
|
||||
uploads: Final = self._batch_file_elements(batch) if self._batch_file_mode_active() else batch
|
||||
results: Final = await asyncio.gather(*(self._upload_bounded(element) for element in uploads))
|
||||
failed: Final = tuple(element for element, ok in zip(uploads, results, strict=True) if not ok)
|
||||
if not failed:
|
||||
self._flush_retries = 0
|
||||
self._flush_dropped = {} # mutable-ok: per-flush drop marks read back by _upload_bounded
|
||||
stale: Final = min(self._requeued_count, len(uploads)) if len(uploads) == len(batch) else 0
|
||||
order: Final = (*range(stale, len(uploads)), *range(stale))
|
||||
ordered: Final = await asyncio.gather(*(self._upload_outcome(uploads[i]) for i in order))
|
||||
outcomes: Final = dict(zip(order, ordered, strict=True))
|
||||
results: Final = tuple(outcomes[i] for i in range(len(uploads)))
|
||||
if self._flush_retries:
|
||||
verbose_logger.warning(
|
||||
"s3 logging: %s in-call retries across %s uploads this flush",
|
||||
self._flush_retries,
|
||||
len(uploads),
|
||||
)
|
||||
delivered: Final = sum(1 for outcome in results if outcome == "delivered")
|
||||
bucket_wide: Final = delivered == 0
|
||||
failed: Final = tuple(
|
||||
(element, outcome) for element, outcome in zip(uploads, results, strict=True) if outcome != "delivered"
|
||||
)
|
||||
now: Final = time.monotonic()
|
||||
requeued: Final = (
|
||||
tuple(element for element, _ in failed)
|
||||
if bucket_wide
|
||||
else tuple(
|
||||
element
|
||||
if element.retrying_since is not None or self.s3_max_retry_age_seconds is None
|
||||
else element.model_copy(update={"retrying_since": now})
|
||||
for element, outcome in failed
|
||||
if outcome != "dropped"
|
||||
and not (
|
||||
self.s3_max_retry_age_seconds is not None
|
||||
and element.retrying_since is not None
|
||||
and now - element.retrying_since > self.s3_max_retry_age_seconds
|
||||
)
|
||||
)
|
||||
)
|
||||
dropped: Final = len(failed) - len(requeued)
|
||||
if dropped:
|
||||
verbose_logger.warning(
|
||||
"s3 logging: %s uploads dropped (terminal or retrying longer than s3_max_retry_age_seconds=%s)",
|
||||
dropped,
|
||||
self.s3_max_retry_age_seconds,
|
||||
)
|
||||
if not requeued:
|
||||
self._requeued_count = 0
|
||||
return
|
||||
self.log_queue = [*failed, *self.log_queue[len(batch) :]]
|
||||
arrivals: Final = self.log_queue[len(batch) :]
|
||||
overflow: Final = max(0, len(requeued) + len(arrivals) - self.max_queue_size)
|
||||
if overflow:
|
||||
verbose_logger.warning(
|
||||
"s3 logging: queue exceeded max_queue_size=%s after a failed flush, dropped %s oldest events",
|
||||
self.max_queue_size,
|
||||
overflow,
|
||||
)
|
||||
self.log_queue = [ # mutable-ok: log_queue is the flush buffer shared with custom_batch_logger
|
||||
*requeued,
|
||||
*arrivals,
|
||||
][overflow:]
|
||||
self._requeued_count = max(0, len(requeued) - overflow)
|
||||
raise S3BatchUploadError(failed=len(failed), total=len(uploads))
|
||||
|
||||
def _batch_file_mode_active(self) -> bool:
|
||||
|
|
@ -502,8 +699,37 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
return True
|
||||
|
||||
async def _upload_bounded(self, element: s3BatchLoggingElement) -> bool:
|
||||
async with self._upload_semaphore:
|
||||
return await self.async_upload_data_to_s3(element)
|
||||
token: Final = _in_flush.set(True)
|
||||
try:
|
||||
async with self._upload_semaphore:
|
||||
return await self.async_upload_data_to_s3(element)
|
||||
finally:
|
||||
_in_flush.reset(token)
|
||||
|
||||
async def _upload_outcome(self, element: s3BatchLoggingElement) -> UploadOutcome:
|
||||
delivered: Final = await self._upload_bounded(element)
|
||||
if delivered:
|
||||
return "delivered"
|
||||
if id(element) in self._flush_dropped:
|
||||
return "dropped"
|
||||
return "retry"
|
||||
|
||||
async def _recorded_put(self, signed_put: Callable[[], Awaitable[httpx.Response]]) -> httpx.Response:
|
||||
limiter: Final = self._upload_limiter
|
||||
adaptive: Final = limiter if isinstance(limiter, AdaptiveConcurrencyLimiter) else None
|
||||
try:
|
||||
response: Final = await signed_put()
|
||||
except Exception:
|
||||
if adaptive is not None:
|
||||
adaptive.record(PutSample(throttled=True))
|
||||
raise
|
||||
if adaptive is not None:
|
||||
adaptive.record(
|
||||
PutSample(
|
||||
throttled=response.status_code in (429, 503) or _s3_error_code(response) == "SlowDown",
|
||||
)
|
||||
)
|
||||
return response
|
||||
|
||||
def _batch_file_elements(self, batch: tuple[s3BatchLoggingElement, ...]) -> tuple[s3BatchLoggingElement, ...]:
|
||||
now: Final = datetime.now(timezone.utc)
|
||||
|
|
@ -527,6 +753,9 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
content_type="application/x-ndjson",
|
||||
s3_object_key=f"{parent}/{batch_name}.jsonl" if parent else f"{batch_name}.jsonl",
|
||||
s3_object_download_filename=f"{batch_name}.jsonl",
|
||||
retrying_since=min(
|
||||
(element.retrying_since for element in elements if element.retrying_since is not None), default=None
|
||||
),
|
||||
)
|
||||
|
||||
def create_s3_batch_logging_element(
|
||||
|
|
@ -596,58 +825,35 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
)
|
||||
|
||||
def upload_data_to_s3(self, batch_logging_element: s3BatchLoggingElement):
|
||||
try:
|
||||
import base64
|
||||
import hashlib
|
||||
except ImportError:
|
||||
raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.")
|
||||
try:
|
||||
verbose_logger.debug("s3_v2 logger - uploading data to s3 - %s", batch_logging_element.s3_object_key)
|
||||
|
||||
url: Final = self._build_object_url(batch_logging_element.s3_object_key)
|
||||
|
||||
# Convert JSON to string
|
||||
json_string: Final = (
|
||||
batch_logging_element.body
|
||||
if batch_logging_element.body is not None
|
||||
else safe_dumps(batch_logging_element.payload)
|
||||
)
|
||||
|
||||
# Calculate SHA256 hash of the content
|
||||
content_hash: Final = hashlib.sha256(json_string.encode("utf-8")).hexdigest()
|
||||
content_md5: Final = base64.b64encode(
|
||||
hashlib.md5(json_string.encode("utf-8"), usedforsecurity=False).digest()
|
||||
).decode()
|
||||
|
||||
# Prepare the request
|
||||
headers: Final = {
|
||||
"Content-Type": batch_logging_element.content_type,
|
||||
"Content-MD5": content_md5,
|
||||
"x-amz-content-sha256": content_hash,
|
||||
"Content-Language": "en",
|
||||
"Content-Disposition": f'inline; filename="{batch_logging_element.s3_object_download_filename}"',
|
||||
"Cache-Control": "private, immutable, max-age=31536000, s-maxage=0",
|
||||
**self._sse_headers(),
|
||||
}
|
||||
prepared: Final = self._prepare_put(batch_logging_element)
|
||||
|
||||
httpx_client: Final = _get_httpx_client(
|
||||
params=({"ssl_verify": self.s3_verify} if self.s3_verify is not None else None)
|
||||
)
|
||||
|
||||
def signed_put() -> httpx.Response:
|
||||
def signed_put(prepared_put: _PreparedPut) -> httpx.Response:
|
||||
credentials: Final = self.get_credentials(
|
||||
aws_access_key_id=self.s3_aws_access_key_id,
|
||||
aws_secret_access_key=self.s3_aws_secret_access_key,
|
||||
aws_session_token=self.s3_aws_session_token,
|
||||
aws_region_name=self.s3_region_name,
|
||||
)
|
||||
signed_headers: Final = self._sign_put(credentials, url, json_string, headers)
|
||||
return httpx_client.put(url, data=json_string, headers=signed_headers)
|
||||
signed_headers: Final = self._sign_put(credentials, url, prepared_put.json_string, prepared_put.headers)
|
||||
return httpx_client.put(url, data=prepared_put.json_string, headers=signed_headers)
|
||||
|
||||
max_retries: Final = 3
|
||||
for attempt in range(max_retries):
|
||||
response = signed_put()
|
||||
if response.status_code in (403, 500, 503) and attempt < max_retries - 1:
|
||||
response = signed_put(prepared)
|
||||
if (
|
||||
response.status_code in _RETRYABLE_STATUSES
|
||||
and not (self.s3_drop_on_terminal_error and _is_terminal(response))
|
||||
and attempt < max_retries - 1
|
||||
):
|
||||
wait_time = 2**attempt # 1s, 2s
|
||||
verbose_logger.warning(
|
||||
"S3 upload returned %s, retrying in %ss (attempt %s/%s) key=%s",
|
||||
|
|
@ -664,6 +870,12 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
except Exception as e:
|
||||
verbose_logger.exception("Error uploading to s3: %s", e)
|
||||
self.handle_callback_failure(callback_name="S3Logger")
|
||||
if isinstance(e, httpx.HTTPStatusError) and self.s3_drop_on_terminal_error and _is_terminal(e.response):
|
||||
verbose_logger.warning(
|
||||
"s3 logging: dropping object %s after terminal status %s",
|
||||
batch_logging_element.s3_object_key,
|
||||
e.response.status_code,
|
||||
)
|
||||
|
||||
async def _download_object_from_s3(self, s3_object_key: str) -> dict | None:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -111,7 +111,7 @@ def _chat_request_from_anthropic_messages(
|
|||
because the logged optional_params switch dialect per provider path (the bridge's
|
||||
inner completion rewrites them to chat shape mid-flight); the adapter translates
|
||||
them alongside the messages, and sampling params copy through untranslated."""
|
||||
from litellm.llms.anthropic.experimental_pass_through.adapters.transformation import (
|
||||
from litellm.llms.anthropic.pass_through.adapters.transformation import (
|
||||
LiteLLMAnthropicMessagesAdapter,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -2,11 +2,13 @@
|
|||
Vector Store Pre-Call Hook
|
||||
|
||||
This hook is called before making an LLM request when a vector store is configured.
|
||||
It searches the vector store for relevant context and appends it to the messages.
|
||||
It searches the vector store for relevant context, runs the request's pre-call guardrails
|
||||
over that context, and appends it to the messages.
|
||||
"""
|
||||
|
||||
from collections.abc import Awaitable, Callable, Mapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
from itertools import chain
|
||||
from typing import TYPE_CHECKING, Any, Final, Protocol, cast, get_args
|
||||
|
||||
from pydantic import TypeAdapter, ValidationError
|
||||
|
|
@ -16,7 +18,9 @@ import litellm
|
|||
import litellm.vector_stores
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.exceptions import VectorStoreSearchError
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
from litellm.types.llms.openai import (
|
||||
AllMessageValues,
|
||||
ChatCompletionUserMessage,
|
||||
|
|
@ -42,6 +46,24 @@ else:
|
|||
SEARCH_FAILURES_FIELD: Final = "vector_store_search_failures"
|
||||
_DEFAULT_FAILURE_MODE: Final[VectorStoreSearchFailureMode] = "annotate"
|
||||
_FAILURE_MODE_ADAPTER: Final = TypeAdapter(VectorStoreSearchFailureMode)
|
||||
_STR_KEYED_ADAPTER: Final = TypeAdapter(dict[str, object])
|
||||
_GUARDRAIL_KEYS_THE_PROXY_MERGES_INTO_METADATA: Final = frozenset(
|
||||
{"guardrails", "guardrail_config", "policies", "include_guardrail_response"}
|
||||
)
|
||||
|
||||
|
||||
def _scan_request(model: str, non_default_params: Mapping[str, object]) -> Mapping[str, object]:
|
||||
try:
|
||||
proxy_request: Final = _STR_KEYED_ADAPTER.validate_python(non_default_params.get("proxy_server_request"))
|
||||
client_body: Final = _STR_KEYED_ADAPTER.validate_python(proxy_request.get("body"))
|
||||
except ValidationError:
|
||||
return {**non_default_params, "model": model}
|
||||
proxy_request_params: Final = {**client_body, **non_default_params}
|
||||
return {
|
||||
key: value
|
||||
for key, value in proxy_request_params.items()
|
||||
if key not in _GUARDRAIL_KEYS_THE_PROXY_MERGES_INTO_METADATA
|
||||
}
|
||||
|
||||
|
||||
class ProxyRuntime(Protocol):
|
||||
|
|
@ -82,7 +104,7 @@ SearchOutcome = SearchSucceeded | SearchFailed
|
|||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class VectorStoreAugmentation:
|
||||
messages: tuple[AllMessageValues, ...]
|
||||
context_messages: tuple[AllMessageValues, ...]
|
||||
search_results: tuple[VectorStoreSearchResponse, ...]
|
||||
failures: tuple[VectorStoreSearchFailure, ...]
|
||||
|
||||
|
|
@ -95,7 +117,8 @@ class VectorStorePreCallHook(CustomLogger):
|
|||
When a vector store is configured, this hook:
|
||||
1. Extracts the query from the last user message
|
||||
2. Calls litellm.vector_stores.search() to get relevant context
|
||||
3. Appends the search results as context to the messages
|
||||
3. Runs the request's pre-call guardrails over each store's context message
|
||||
4. Appends the (possibly masked) context to the messages, or raises the guardrail's block
|
||||
"""
|
||||
|
||||
def __init__(self, proxy_runtime: ProxyRuntime | None = None):
|
||||
|
|
@ -170,7 +193,50 @@ class VectorStorePreCallHook(CustomLogger):
|
|||
case _:
|
||||
assert_never(failure_mode)
|
||||
|
||||
return model, list(augmentation.messages), non_default_params
|
||||
scanned_context: Final = await self._scanned_context_messages(
|
||||
model=model,
|
||||
non_default_params=non_default_params,
|
||||
context_messages=augmentation.context_messages,
|
||||
)
|
||||
return (
|
||||
model,
|
||||
self._messages_with_context(messages=messages, context_messages=scanned_context),
|
||||
non_default_params,
|
||||
)
|
||||
|
||||
async def _scanned_context_messages(
|
||||
self,
|
||||
model: str,
|
||||
non_default_params: Mapping[str, object],
|
||||
context_messages: Sequence[AllMessageValues],
|
||||
) -> tuple[AllMessageValues, ...]:
|
||||
request_data: Final = _scan_request(model, non_default_params)
|
||||
guardrails: Final = tuple(
|
||||
callback
|
||||
for callback in litellm.callbacks
|
||||
if isinstance(callback, CustomGuardrail)
|
||||
and callback.should_run_guardrail(data=request_data, event_type=GuardrailEventHooks.pre_call)
|
||||
)
|
||||
if not guardrails:
|
||||
return tuple(context_messages)
|
||||
scanned: Final = [
|
||||
await self._scan_through(guardrails=guardrails, request_data=request_data, messages=(context_message,))
|
||||
for context_message in context_messages
|
||||
]
|
||||
return tuple(chain.from_iterable(scanned))
|
||||
|
||||
async def _scan_through(
|
||||
self,
|
||||
guardrails: Sequence[CustomGuardrail],
|
||||
request_data: Mapping[str, object],
|
||||
messages: Sequence[AllMessageValues],
|
||||
) -> tuple[AllMessageValues, ...]:
|
||||
if not guardrails:
|
||||
return tuple(messages)
|
||||
scanned: Final = await guardrails[0].async_pre_call_hook_on_messages(
|
||||
request_data=request_data, messages=messages
|
||||
)
|
||||
return await self._scan_through(guardrails=guardrails[1:], request_data=request_data, messages=scanned)
|
||||
|
||||
async def _augment_messages(
|
||||
self,
|
||||
|
|
@ -234,7 +300,7 @@ class VectorStorePreCallHook(CustomLogger):
|
|||
failures: Final = tuple(outcome.failure for outcome in outcomes if isinstance(outcome, SearchFailed))
|
||||
|
||||
return VectorStoreAugmentation(
|
||||
messages=self._messages_with_context(messages=messages, search_results=search_results),
|
||||
context_messages=self._context_messages(search_results),
|
||||
search_results=search_results,
|
||||
failures=failures,
|
||||
)
|
||||
|
|
@ -309,19 +375,21 @@ class VectorStorePreCallHook(CustomLogger):
|
|||
|
||||
return None
|
||||
|
||||
def _messages_with_context(
|
||||
self,
|
||||
messages: Sequence[AllMessageValues],
|
||||
search_results: Sequence[VectorStoreSearchResponse],
|
||||
) -> tuple[AllMessageValues, ...]:
|
||||
context_messages: Final = tuple(
|
||||
def _context_messages(self, search_results: Sequence[VectorStoreSearchResponse]) -> tuple[AllMessageValues, ...]:
|
||||
return tuple(
|
||||
context_message
|
||||
for search_response in search_results
|
||||
if (context_message := self._context_message(search_response)) is not None
|
||||
)
|
||||
|
||||
def _messages_with_context(
|
||||
self,
|
||||
messages: Sequence[AllMessageValues],
|
||||
context_messages: Sequence[AllMessageValues],
|
||||
) -> list[AllMessageValues]:
|
||||
if not context_messages:
|
||||
return tuple(messages)
|
||||
return (*messages[:-1], *context_messages, *messages[-1:])
|
||||
return list(messages)
|
||||
return [*messages[:-1], *context_messages, *messages[-1:]]
|
||||
|
||||
def _context_message(self, search_response: VectorStoreSearchResponse) -> AllMessageValues | None:
|
||||
"""Build the context message for one vector store's results, or None when it returned nothing usable."""
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue