Merge branch 'main' into fix/bedrock-converse-redacted-thinking-replay-43009

This commit is contained in:
Krrish Dholakia 2026-09-26 21:28:19 -07:00 • committed by GitHub
commit 64bc20a5d0
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
512 changed files with 30155 additions and 5291 deletions

View file

@ -141,7 +141,7 @@ commands:
node --version
npm --version
install_rust:
description: "Install pinned rustup (1.28.2) and Rust toolchain (1.98.0) with checksum verification. Adds ~/.cargo/bin to PATH. Run this before any `uv sync` or `uv build` of the workspace: the root package builds litellm-rust through maturin, and on an image without cargo maturin fetches an unpinned rustup and a floating toolchain by itself."
description: "Install pinned rustup (1.28.2) and Rust toolchain (1.98.0) with checksum verification. Adds ~/.cargo/bin to PATH. Run this before any `uv sync` or `uv build` of the workspace: the root package builds litellm-rust through maturin, and on an image without cargo maturin fetches an unpinned rustup and a floating toolchain by itself. Also restores the dev-profile cargo cache that save_cargo_target writes on main, minus the workspace crates' fingerprints so those always rebuild from the checked-out source."
steps:
- run:
name: Install Rust (rustup 1.28.2, toolchain 1.98.0)
@ -167,9 +167,29 @@ commands:
/tmp/rustup-init -y --no-modify-path --profile minimal --default-toolchain 1.98.0
rm -f /tmp/rustup-init
echo 'export PATH="$HOME/.cargo/bin:$PATH"' >> "$BASH_ENV"
echo 'export CARGO_INCREMENTAL=0' >> "$BASH_ENV"
export PATH="$HOME/.cargo/bin:$PATH"
rustc --version
cargo --version
{ rustc -vV; cc --version; cat /etc/os-release; } > /tmp/cargo-build-env
- restore_cache:
keys:
- v1-cargo-dev-{{ checksum "/tmp/cargo-build-env" }}-{{ checksum "litellm-rust/Cargo.lock" }}
- v1-cargo-dev-{{ checksum "/tmp/cargo-build-env" }}-
- run:
name: Force a rebuild of the workspace crates restored from the cargo cache
command: rm -rf litellm-rust/target/debug/.fingerprint/litellm-*
save_cargo_target:
steps:
- when:
condition:
equal: [main, << pipeline.git.branch >>]
steps:
- save_cache:
key: v1-cargo-dev-{{ checksum "/tmp/cargo-build-env" }}-{{ checksum "litellm-rust/Cargo.lock" }}
paths:
- ~/.cargo/registry
- ~/project/litellm-rust/target/debug
start_postgres:
description: "Start a postgres-db container on port 5432 and wait until it accepts connections."
parameters:
@ -281,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

View file

@ -1,7 +1,7 @@
#!/usr/bin/env bash
set -uo pipefail
category="${1:?usage: classify_changes.sh <backend|client|ui|provider-harness|cost-map-only|mcp-dependencies>}"
category="${1:?usage: classify_changes.sh <backend|client|ui|provider-harness|cost-map-only|mcp-dependencies|windows-release>}"
has_client=false
has_backend=false
@ -9,6 +9,7 @@ has_ci=false
has_provider_harness=false
has_cost_map=false
has_mcp_dependencies=false
has_windows_release=false
outside_cost_map_set=false
while IFS= read -r file || [ -n "$file" ]; do
[ -n "$file" ] || continue
@ -22,6 +23,10 @@ while IFS= read -r file || [ -n "$file" ]; do
tests/e2e/*.py | tests/code_coverage_tests/test_provider_cache.py | tests/code_coverage_tests/test_provider_replay_harness.py | tests/unit/test_circleci_path_filter.py | .circleci/* | pyproject.toml | uv.lock)
has_provider_harness=true ;;
esac
case "$file" in
litellm-rust/* | litellm/rust_bridge/* | rust-toolchain.toml | pyproject.toml | uv.lock | tests/windows_tests/* | .circleci/*)
has_windows_release=true ;;
esac
case "$file" in
ui/* | tests/e2e/ui/*) has_client=true ;;
docs/* | *.md | *.mdx) : ;;
@ -46,6 +51,9 @@ case "$category" in
provider-harness)
[ "$has_provider_harness" = true ] && echo run || echo skip
;;
windows-release)
[ "$has_windows_release" = true ] && echo run || echo skip
;;
backend)
[ "$has_backend" = true ] && echo run || echo skip
;;

View file

@ -26,6 +26,7 @@ guard_created=false
guard_installed=false
guard6_created=false
guard6_installed=false
egress_cgroup=litellm-integration
cleanup() {
original_status=$?
trap - EXIT INT TERM
@ -47,14 +48,14 @@ cleanup() {
fi
done
if [ "$guard_installed" = true ]; then
sudo iptables -D OUTPUT -m owner --uid-owner "$(id -u)" -j integration_only || original_status=1
sudo iptables -D OUTPUT -m cgroup --path "$egress_cgroup" -j integration_only || original_status=1
fi
if [ "$guard_created" = true ]; then
sudo iptables -F integration_only || original_status=1
sudo iptables -X integration_only || original_status=1
fi
if [ "$guard6_installed" = true ]; then
sudo ip6tables -D OUTPUT -m owner --uid-owner "$(id -u)" -j integration_only || original_status=1
sudo ip6tables -D OUTPUT -m cgroup --path "$egress_cgroup" -j integration_only || original_status=1
fi
if [ "$guard6_created" = true ]; then
sudo ip6tables -F integration_only || original_status=1
@ -100,6 +101,8 @@ if [ "$mode" = parity ]; then
export INTEGRATION_ROUTING=capture
fi
sudo mkdir -p "/sys/fs/cgroup/$egress_cgroup"
echo "$$" | sudo tee "/sys/fs/cgroup/$egress_cgroup/cgroup.procs" > /dev/null
sudo iptables -N integration_only
guard_created=true
sudo iptables -A integration_only -o lo -j ACCEPT
@ -109,13 +112,13 @@ for service in postgres-db redis-cache; do
sudo iptables -A integration_only -d "$address" -j ACCEPT
done
sudo iptables -A integration_only -j REJECT
sudo iptables -I OUTPUT 1 -m owner --uid-owner "$(id -u)" -j integration_only
sudo iptables -I OUTPUT 1 -m cgroup --path "$egress_cgroup" -j integration_only
guard_installed=true
sudo ip6tables -N integration_only
guard6_created=true
sudo ip6tables -A integration_only -o lo -j ACCEPT
sudo ip6tables -A integration_only -j REJECT
sudo ip6tables -I OUTPUT 1 -m owner --uid-owner "$(id -u)" -j integration_only
sudo ip6tables -I OUTPUT 1 -m cgroup --path "$egress_cgroup" -j integration_only
guard6_installed=true
if curl --noproxy '*' --connect-timeout 2 -s http://198.51.100.1 >/dev/null 2>&1; then
@ -209,6 +212,15 @@ if [ "$suite" = browser ]; then
exit 0
fi
node_files=()
if [ "${CIRCLE_NODE_TOTAL:-1}" -gt 1 ]; then
split="$(.venv/bin/python tests/integration/run.py "$suite" --list \
| circleci tests split --split-by=timings --timings-type=filename)"
read -r -a node_files <<< "$(printf '%s' "$split" | tr '\n' ' ')"
test "${#node_files[@]}" -gt 0
printf '%s\n' "${node_files[@]}" > "$results/node-files.txt"
fi
env -i PATH="$PATH" HOME="$HOME" PYTHONPATH="$PYTHONPATH" \
INTEGRATION_RUN_ID="$integration_identity" \
DATABASE_URL="$DATABASE_URL" REDIS_HOST="$REDIS_HOST" REDIS_PORT="$REDIS_PORT" \
@ -222,7 +234,7 @@ env -i PATH="$PATH" HOME="$HOME" PYTHONPATH="$PYTHONPATH" \
INTEGRATION_PROXY_DATABASE_URL="$INTEGRATION_PROXY_DATABASE_URL" \
INTEGRATION_PROXY_READ_REPLICA_URL="$INTEGRATION_PROXY_READ_REPLICA_URL" \
INTEGRATION_ROUTING="$INTEGRATION_ROUTING" \
.venv/bin/python tests/integration/run.py "$suite" --results "$results"
.venv/bin/python tests/integration/run.py "$suite" --results "$results" "${node_files[@]}"
if [ "${INTEGRATION_COVERAGE:-0}" = 1 ]; then
for covered_pid in "$proxy_pid" "$peer_pid"; do

View file

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

View file

@ -54,7 +54,7 @@ After: the same request comes back with real token counts, so the dashboard show
## Affected release
<!-- Only for a fix to a regression in a released or rc version (perf, memory, crash, or behavior): name the version it regressed in, e.g. "regression in v1.100.0" or "since v1.101.0-rc.1", and add the `backport-stable` label so the fix is cherry-picked onto the rc line before the stable is tagged. Drop the section otherwise -->
<!-- Only for a fix to a regression in a released or rc version (perf, memory, crash, or behavior): name the version it regressed in, e.g. "regression in v1.100.0" or "since v1.101.0-rc.1". Add the `backport-stable` label only when the regression is a P0, meaning its Linear ticket is Urgent (a security hole however narrow, data loss, or a crash or outage for every user on that version), because every labeled PR must be cherry-picked onto the baking rc line before the stable can be tagged; every other regression fix ships in the next rc unlabeled. Drop the section otherwise -->
## Linear ticket

View file

@ -516,7 +516,7 @@ def _integration_ownership(repo_root: pathlib.Path = REPO_ROOT) -> tuple[frozens
str(path.relative_to(repo_root))
for folders in groups.values()
for folder in folders
for path in (integration_root / folder).glob("test_*.py")
for path in (integration_root / folder).rglob("test_*.py")
)
browser_manifest: Final = repo_root / "tests/e2e/ui/tests/integrationCritical/expected.json"
browser_nodes: Final = json.loads(browser_manifest.read_text()) if browser_manifest.exists() else ()

View file

@ -146,6 +146,9 @@ jobs:
- name: check_migrations_no_data_rewrites
run: uv run --no-sync python ./tests/code_coverage_tests/check_migrations_no_data_rewrites.py
- name: check_unbounded_in_lists (fails on findings not in the baseline)
run: uv run --no-sync python ./tests/code_coverage_tests/check_unbounded_in_lists.py
- name: memory_test
run: uv run --no-sync python ./tests/code_coverage_tests/memory_test.py

View file

@ -336,7 +336,7 @@ Each translation is isolated in its own file, making it easy to test and modify
| `/v1/chat/completions` | Gemini | `llms/gemini/chat/transformation.py` |
| `/v1/chat/completions` | Vertex AI | `llms/vertex_ai/gemini/transformation.py` |
| `/v1/chat/completions` | OpenAI | `llms/openai/chat/gpt_transformation.py` |
| `/v1/messages` (passthrough) | Anthropic | `llms/anthropic/experimental_pass_through/messages/transformation.py` |
| `/v1/messages` (passthrough) | Anthropic | `llms/anthropic/pass_through/messages/transformation.py` |
| `/v1/messages` (passthrough) | Bedrock | `llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py` |
| `/v1/messages` (passthrough) | Vertex AI | `llms/vertex_ai/vertex_ai_partner_models/anthropic/experimental_pass_through/transformation.py` |
| Passthrough endpoints | All | `proxy/pass_through_endpoints/llm_provider_handlers/` |

View file

@ -362,6 +362,7 @@ For MCP OAuth, an upstream may advertise dynamic client registration but refuse
| [Recraft (`recraft`)](https://docs.litellm.ai/docs/providers/recraft) | | | | | ✅ | | | | | |
| [Replicate (`replicate`)](https://docs.litellm.ai/docs/providers/replicate) | ✅ | ✅ | ✅ | | | | | | | |
| [Sagemaker Chat (`sagemaker_chat`)](https://docs.litellm.ai/docs/providers/aws_sagemaker) | ✅ | ✅ | ✅ | | | | | | | |
| [Sail (`sail`)](https://docs.litellm.ai/docs/providers/sail) | ✅ | ✅ | ✅ | | | | | | | |
| [Sambanova (`sambanova`)](https://docs.litellm.ai/docs/providers/sambanova) | ✅ | ✅ | ✅ | | | | | | | |
| [Snowflake (`snowflake`)](https://docs.litellm.ai/docs/providers/snowflake) | ✅ | ✅ | ✅ | | | | | | | |
| [Text Completion Codestral (`text-completion-codestral`)](https://docs.litellm.ai/docs/providers/codestral) | ✅ | ✅ | ✅ | | | | | | | |

View file

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

View file

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

View file

@ -7,7 +7,7 @@ pub enum CredentialPlacement {
}
impl CredentialPlacement {
pub fn header_name(self) -> &'static str {
pub const fn header_name(self) -> &'static str {
match self {
Self::Bearer => "Authorization",
Self::Header(name) => name,

View file

@ -1,4 +1,6 @@
litellm-core is the LiteLLM SDK in Rust — it makes the LLM call. Each top-level call is a module under `src/<route>/` exposing a public entrypoint named after the route (`messages::messages()`, the Rust equivalent of `litellm.messages()`): you call it and get a typed non-streaming response back.
litellm-core is the LiteLLM SDK in Rust. Each top-level call is a module under `src/<route>/` exposing a public entrypoint named after the route. `messages::messages()` returns `MessagesResponse::Message` for a completed response or `MessagesResponse::Stream { headers, chunks }` when the request sets `stream: true`. The chunks are Anthropic SSE bytes in a `Stream<Item = Result<Bytes, Error>>`. Dropping the stream cancels the call. The Python bridge drives `messages::route::messages_machine()` instead, because Python has to answer the call's operations on its own thread; the gateway and the Rust SDK call the plain entrypoint
A route module has the same five pieces, in the order Python runs them. `types.rs` holds the call, the provider request, and the response. `prepare.rs` resolves the provider and credentials and shapes the request (Python's `validate_environment`, `get_complete_url`, `transform_request`). `handler.rs` resolves auth, offers the wire request to `litellm_host::hooks::RouteHooks::before_send`, sends it, reports the raw response through `emit`, and normalizes the response or stream (`pre_call`, `post`, `post_call`, `transform_response`). `mod.rs` exposes the entrypoint that runs prepare then handler with no hooks (`()`). `route.rs`, where a host needs it, wraps the same two calls in a `CallMachine` whose `HostChannel` is the hooks, and pumps a stream through `open` and `deliver`. A handler takes `&impl RouteHooks<Error>` and never a `HostChannel` directly, so it runs without a coroutine. Keep provider transport and transformation details out of the machine driver
## Crate layering
@ -10,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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -1,3 +1,6 @@
use std::time::Duration;
use litellm_auth::SecretValue;
use litellm_core_utils::{
dot_notation_indexing::delete_nested_value,
get_llm_provider_logic::{CustomLlmProvider, get_custom_llm_provider},
@ -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 {

View file

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

View file

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

View file

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

View file

@ -1,7 +1,5 @@
use litellm_llms::anthropic::common_utils::{
ANTHROPIC_ADVISOR_TOOL_TYPE, ANTHROPIC_OAUTH_BETA_HEADER, AnthropicModelCapabilities,
SupportedEffortTiers, beta,
};
use litellm_llms::anthropic::common_utils::{AnthropicModelCapabilities, SupportedEffortTiers};
use litellm_types::llms::anthropic::{AnthropicBeta, BetaSet};
use litellm_types::utils::{ProviderSpecificHeader, ProviderSpecificHeaders};
use rstest::rstest;
@ -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);
}

View file

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

View file

@ -1,12 +1,21 @@
use std::{convert::Infallible, sync::Mutex};
use std::{
convert::Infallible,
sync::{Mutex, mpsc},
};
use bytes::Bytes;
use litellm_core::messages::route::{Messages, MessagesStreamHead};
use futures_util::{StreamExt, TryStreamExt};
use litellm_core::messages::{
MessagesResponse, messages,
route::{Messages, MessagesStreamHead},
};
use litellm_host::host::{Demand, Host};
use litellm_tracing::{Logger, Metadata, Record, Sink};
use rstest::rstest;
use tokio::{
io::{AsyncReadExt, AsyncWriteExt},
net::TcpListener,
task::JoinHandle,
};
use super::*;
@ -23,6 +32,20 @@ enum Seen {
Deliver(Bytes),
}
struct TraceSink(mpsc::Sender<(String, Value)>);
impl Sink for TraceSink {
fn enabled(&self, metadata: &Metadata<'_>) -> bool {
metadata.target().starts_with("litellm_core::messages")
}
fn emit(&self, record: &Record) {
self.0
.send((record.message.clone(), Value::Object(record.fields.clone())))
.unwrap();
}
}
/// Projects like `LocalMessagesHost`, records every stream op in the order the route
/// performs it, and detaches after `detach_after` ops.
struct RecordingStreamHost {
@ -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]

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View 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");
}
}

View file

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

View file

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

View file

@ -1,3 +1,4 @@
[package]
name = "litellm"
version = "0.0.1"
edition.workspace = true

View file

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

View file

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

View file

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

View file

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

View file

@ -1,5 +1,4 @@
pub mod handler;
pub mod headers;
pub mod streaming_iterator;
pub mod thinking;
pub mod transformation;

View file

@ -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 &params.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]

View file

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

View file

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

View file

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

View file

@ -0,0 +1 @@
pub mod transformation;

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

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

View file

@ -0,0 +1,2 @@
pub mod chat;
pub mod common_utils;

View file

@ -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), &params(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", &params(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",
&params(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", &params(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",
&params(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",
&params(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",
&params(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", &params(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"))
);
}

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

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

View file

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

View file

@ -1,2 +1,3 @@
pub mod anthropic;
pub mod anthropic_messages;
pub mod openai;

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -41,6 +41,7 @@ class ExporterOwner(str, Enum):
LEVO = "levo"
AGENTOPS = "agentops"
NEWRELIC = "newrelic"
SIGNOZ = "signoz"
class _OTelV2Flag(BaseSettings):

View file

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

View file

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

View file

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

View file

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

View 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

View file

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

View file

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

View file

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

View file

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