Merge branch 'main' into litellm_otel_v2_tenant_internal_spans
Some checks are pending
LiteLLM Rust / rust-lint (push) Waiting to run
LiteLLM Rust / rust-test (push) Waiting to run
LiteLLM Rust / rust-wheel (push) Waiting to run

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
mrinal 2026-09-24 19:10:57 +00:00
commit 10cb68aa6e
395 changed files with 44971 additions and 2441 deletions

View file

@ -12,6 +12,9 @@ parameters:
migration_source_sha:
type: string
default: ""
routing_parity_base:
type: string
default: ""
orbs:
codecov: codecov/codecov@4.0.1
node: circleci/node@5.1.0 # Add this line to declare the node orb
@ -176,6 +179,9 @@ commands:
image:
type: string
default: postgres:14@sha256:6a70deda415ec296f977890e11aba04a0db9f632a362e3fce45e845e3db74f26
server_args:
type: string
default: ""
steps:
- run:
name: Start PostgreSQL
@ -186,7 +192,7 @@ commands:
-e POSTGRES_PASSWORD=postgres \
-e POSTGRES_DB=<< parameters.db_name >> \
-p 5432:5432 \
<< parameters.image >>
<< parameters.image >> << parameters.server_args >>
- wait_for_service:
url: tcp://localhost:5432
timeout: "60"
@ -3108,6 +3114,10 @@ jobs:
parameters:
suite:
type: string
mode:
type: enum
enum: [standard, replica]
default: standard
machine:
image: ubuntu-2204:2024.04.1
resource_class: large
@ -3142,18 +3152,19 @@ jobs:
command: cd ui/litellm-dashboard && NEXT_TELEMETRY_DISABLED=1 npm run build
- start_postgres:
image: postgres:16@sha256:e17e86066e5ef83e0952a9347f5c792b7ece00972e2aa787a6986f471b3dd3d5
server_args: "-c shared_preload_libraries=pg_stat_statements -c pg_stat_statements.track=all -c pg_stat_statements.max=20000"
- start_redis
- run:
name: Run owned integration contracts
command: bash .circleci/scripts/run_integration.sh << parameters.suite >>
command: bash .circleci/scripts/run_integration.sh << parameters.suite >> << parameters.mode >>
no_output_timeout: 15m
- run:
name: Stop owned database and Redis
when: always
command: |
mkdir -p test-results/integration-<< parameters.suite >>
docker logs postgres-db > test-results/integration-<< parameters.suite >>/postgres.log 2>&1 || true
docker logs redis-cache > test-results/integration-<< parameters.suite >>/redis.log 2>&1 || true
mkdir -p test-results/services-<< parameters.suite >>-<< parameters.mode >>
docker logs postgres-db > test-results/services-<< parameters.suite >>-<< parameters.mode >>/postgres.log 2>&1 || true
docker logs redis-cache > test-results/services-<< parameters.suite >>-<< parameters.mode >>/redis.log 2>&1 || true
docker rm -f postgres-db redis-cache
test -z "$(docker ps -aq --filter name=postgres-db --filter name=redis-cache)"
- store_test_results:
@ -3161,6 +3172,76 @@ jobs:
- store_artifacts:
path: test-results
routing_parity:
parameters:
suite:
type: string
machine:
image: ubuntu-2204:2024.04.1
resource_class: large
working_directory: ~/project
steps:
- setup_litellm_test_deps
- run:
name: Check out base product code
environment:
ROUTING_PARITY_BASE: << pipeline.parameters.routing_parity_base >>
command: |
[[ "$ROUTING_PARITY_BASE" =~ ^[0-9a-f]{40}$ ]] || exit 1
git fetch --depth 1 origin "$ROUTING_PARITY_BASE"
git rm -r -f --quiet litellm enterprise litellm-proxy-extras
git checkout "$ROUTING_PARITY_BASE" -- litellm enterprise litellm-proxy-extras
git reset --quiet
test -f litellm/rust_bridge/_native.abi3.so
- start_postgres:
image: postgres:16@sha256:e17e86066e5ef83e0952a9347f5c792b7ece00972e2aa787a6986f471b3dd3d5
server_args: "-c shared_preload_libraries=pg_stat_statements -c pg_stat_statements.track=all -c pg_stat_statements.max=20000"
- start_redis
- run:
name: Run base side
command: bash .circleci/scripts/run_integration.sh << parameters.suite >> parity base
no_output_timeout: 15m
- run:
name: Stop base database and Redis
when: always
command: |
mkdir -p test-results/services-<< parameters.suite >>-parity-base
docker logs postgres-db > test-results/services-<< parameters.suite >>-parity-base/postgres.log 2>&1 || true
docker logs redis-cache > test-results/services-<< parameters.suite >>-parity-base/redis.log 2>&1 || true
docker rm -f postgres-db redis-cache
test -z "$(docker ps -aq --filter name=postgres-db --filter name=redis-cache)"
- run:
name: Check out head product code
command: |
git rm -r -f --quiet litellm enterprise litellm-proxy-extras
git checkout "$CIRCLE_SHA1" -- litellm enterprise litellm-proxy-extras
git reset --quiet
test -f litellm/rust_bridge/_native.abi3.so
- start_postgres:
image: postgres:16@sha256:e17e86066e5ef83e0952a9347f5c792b7ece00972e2aa787a6986f471b3dd3d5
server_args: "-c shared_preload_libraries=pg_stat_statements -c pg_stat_statements.track=all -c pg_stat_statements.max=20000"
- start_redis
- run:
name: Run head side
command: bash .circleci/scripts/run_integration.sh << parameters.suite >> parity head
no_output_timeout: 15m
- run:
name: Stop head database and Redis
when: always
command: |
mkdir -p test-results/services-<< parameters.suite >>-parity-head
docker logs postgres-db > test-results/services-<< parameters.suite >>-parity-head/postgres.log 2>&1 || true
docker logs redis-cache > test-results/services-<< parameters.suite >>-parity-head/redis.log 2>&1 || true
docker rm -f postgres-db redis-cache
test -z "$(docker ps -aq --filter name=postgres-db --filter name=redis-cache)"
- run:
name: Compare routing parity
command: PYTHONPATH="$PWD/tests" .venv/bin/python -m integration._support.routing check test-results/parity-<< parameters.suite >>/base test-results/parity-<< parameters.suite >>/head
- store_test_results:
path: test-results
- store_artifacts:
path: test-results
unit:
machine:
image: ubuntu-2204:2024.04.1
@ -3224,8 +3305,22 @@ workflows:
branches:
only: main
jobs: *migration_jobs
routing_parity:
when:
not:
equal: ["", << pipeline.parameters.routing_parity_base >>]
jobs:
- routing_parity:
name: routing-parity-<< matrix.suite >>
matrix:
parameters:
suite: [management, accounting, database, providers, extensions, cost, mcp]
integration:
unless: << pipeline.parameters.run_migration_tests >>
unless:
or:
- << pipeline.parameters.run_migration_tests >>
- not:
equal: ["", << pipeline.parameters.routing_parity_base >>]
jobs:
- integration_contracts:
name: integration-<< matrix.suite >>
@ -3237,8 +3332,23 @@ workflows:
only:
- main
- /litellm_.*/
- integration_contracts:
name: integration-<< matrix.suite >>-replica
matrix:
parameters:
suite: [management, database]
mode: [replica]
filters:
branches:
only:
- main
- /litellm_.*/
build_and_test:
unless: << pipeline.parameters.run_migration_tests >>
unless:
or:
- << pipeline.parameters.run_migration_tests >>
- not:
equal: ["", << pipeline.parameters.routing_parity_base >>]
jobs:
- using_litellm_on_windows:
filters: &main_branches

View file

@ -11,7 +11,7 @@ run_full() {
[ -n "${CIRCLE_PULL_REQUEST:-}" ] || run_full "not a pull request"
candidate_bases="main"
candidate_bases="${PATH_FILTER_BASE_BRANCH:-main}"
merge_base=""
for base in $candidate_bases; do
git fetch --quiet origin "$base" 2>/dev/null || continue

View file

@ -0,0 +1,52 @@
from __future__ import annotations
import os
from typing import Final
from urllib.parse import urlsplit, urlunsplit
import psycopg
DATABASE_URL: Final = os.environ["DATABASE_URL"]
def postgres_url() -> str:
parsed: Final = urlsplit(DATABASE_URL)
return urlunsplit(parsed._replace(path="/postgres"))
def main() -> None:
with psycopg.connect(postgres_url(), autocommit=True) as admin:
admin.execute("CREATE EXTENSION IF NOT EXISTS pg_stat_statements")
admin.execute("CREATE ROLE litellm_writer LOGIN PASSWORD 'litellm-writer' NOSUPERUSER")
admin.execute("CREATE ROLE litellm_reader LOGIN PASSWORD 'litellm-reader' NOSUPERUSER NOINHERIT")
admin.execute("ALTER ROLE litellm_reader SET default_transaction_read_only = on")
admin.execute("ALTER DATABASE circle_test OWNER TO litellm_writer")
admin.execute("GRANT CONNECT ON DATABASE circle_test TO litellm_reader")
with psycopg.connect(DATABASE_URL, autocommit=True) as admin:
admin.execute("GRANT USAGE ON SCHEMA public TO litellm_reader")
admin.execute(
"ALTER DEFAULT PRIVILEGES FOR ROLE litellm_writer IN SCHEMA public GRANT SELECT ON TABLES TO litellm_reader"
)
admin.execute("GRANT SELECT ON ALL TABLES IN SCHEMA public TO litellm_reader")
parsed: Final = urlsplit(DATABASE_URL)
reader_url: Final = urlunsplit(
parsed._replace(netloc=f"litellm_reader:litellm-reader@{parsed.hostname}:{parsed.port}")
)
writer_url: Final = urlunsplit(
parsed._replace(netloc=f"litellm_writer:litellm-writer@{parsed.hostname}:{parsed.port}")
)
with psycopg.connect(reader_url, autocommit=True) as reader:
assert reader.execute("SHOW transaction_read_only").fetchone() == ("on",)
try:
reader.execute("CREATE TABLE integration_readonly_probe (id int)")
except psycopg.errors.ReadOnlySqlTransaction:
pass
else:
raise AssertionError("litellm_reader executed a write statement")
with psycopg.connect(writer_url, autocommit=True) as writer:
assert writer.execute("SELECT current_user").fetchone() == ("litellm_writer",)
if __name__ == "__main__":
main()

View file

@ -7,7 +7,15 @@ if [ "${GITHUB_ACTIONS:-}" = true ]; then
fi
suite="${1:?integration suite required}"
results="test-results/integration-${suite}"
mode="${2:-standard}"
side="${3:-}"
if [ "$mode" = replica ]; then
results="test-results/integration-${suite}-replica"
elif [ "$mode" = parity ]; then
results="test-results/parity-${suite}/${side:?parity side required}"
else
results="test-results/integration-${suite}"
fi
mkdir -p "$results"
integration_identity="$(.venv/bin/python -c 'import uuid; print(uuid.uuid4().hex)')"
upstream_pid=""
@ -80,6 +88,18 @@ export INTEGRATION_ORDER_SEED="$INTEGRATION_SEED"
uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma > "$results/prisma-generate.log" 2>&1
export INTEGRATION_PROXY_DATABASE_URL=""
export INTEGRATION_PROXY_READ_REPLICA_URL=""
export INTEGRATION_ROUTING=""
if [ "$mode" = replica ] || [ "$mode" = parity ]; then
.venv/bin/python .circleci/scripts/prepare_replica_roles.py > "$results/prepare-replica-roles.log" 2>&1
export INTEGRATION_PROXY_DATABASE_URL="postgresql://litellm_writer:litellm-writer@127.0.0.1:5432/circle_test"
export INTEGRATION_PROXY_READ_REPLICA_URL="postgresql://litellm_reader:litellm-reader@127.0.0.1:5432/circle_test"
fi
if [ "$mode" = parity ]; then
export INTEGRATION_ROUTING=capture
fi
sudo iptables -N integration_only
guard_created=true
sudo iptables -A integration_only -o lo -j ACCEPT
@ -137,8 +157,12 @@ start_proxy() {
else
cost_map_env=("LITELLM_LOCAL_MODEL_COST_MAP=True")
fi
local -a database_env=("DATABASE_URL=${INTEGRATION_PROXY_DATABASE_URL:-$DATABASE_URL}")
if [ -n "$INTEGRATION_PROXY_READ_REPLICA_URL" ]; then
database_env+=("DATABASE_URL_READ_REPLICA=$INTEGRATION_PROXY_READ_REPLICA_URL")
fi
setsid 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" \
"${database_env[@]}" REDIS_HOST="$REDIS_HOST" REDIS_PORT="$REDIS_PORT" \
INTEGRATION_UPSTREAM_URL="$INTEGRATION_UPSTREAM_URL" \
LITELLM_MASTER_KEY="$LITELLM_MASTER_KEY" LITELLM_SALT_KEY="$LITELLM_SALT_KEY" LITELLM_UI_PATH="$LITELLM_UI_PATH" PROXY_BASE_URL="http://127.0.0.1:$port" \
LITELLM_MODE=PRODUCTION STORE_MODEL_IN_DB=True "${cost_map_env[@]}" \
@ -195,6 +219,9 @@ env -i PATH="$PATH" HOME="$HOME" PYTHONPATH="$PYTHONPATH" \
INTEGRATION_SEED="$INTEGRATION_SEED" \
INTEGRATION_ORDER_SEED="$INTEGRATION_ORDER_SEED" \
LITELLM_LOCAL_MODEL_COST_MAP=True AWS_EC2_METADATA_DISABLED=true DO_NOT_TRACK=1 \
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"
if [ "${INTEGRATION_COVERAGE:-0}" = 1 ]; then

292
.circleci/tests.yml Normal file
View file

@ -0,0 +1,292 @@
version: 2.1
commands:
wait_for_service:
parameters:
url:
type: string
timeout:
type: string
default: "60"
steps:
- run:
name: "Wait for << parameters.url >>"
command: |
TIMEOUT=<< parameters.timeout >>
URL="<< parameters.url >>"
ELAPSED=0
echo "Waiting up to ${TIMEOUT}s for ${URL} ..."
if echo "$URL" | grep -q '^tcp://'; then
HOST=$(echo "$URL" | sed 's|tcp://||' | cut -d: -f1)
PORT=$(echo "$URL" | sed 's|tcp://||' | cut -d: -f2)
while ! bash -c "echo > /dev/tcp/$HOST/$PORT" 2>/dev/null; do
sleep 2; ELAPSED=$((ELAPSED+2))
if [ "$ELAPSED" -ge "$TIMEOUT" ]; then echo "Timed out"; exit 1; fi
done
else
while ! curl -sf --max-time 5 "$URL" > /dev/null 2>&1; do
sleep 2; ELAPSED=$((ELAPSED+2))
if [ "$ELAPSED" -ge "$TIMEOUT" ]; then echo "Timed out"; exit 1; fi
done
fi
echo "Service ready after ${ELAPSED}s"
install_uv:
steps:
- run:
name: Install uv (pinned 0.10.9)
command: |
curl -LsSf -o /tmp/uv-install.sh https://astral.sh/uv/0.10.9/install.sh
echo "7fc46e39cb97290b57169c0c813a17970585ac519139f19006453c99b5f2f45f /tmp/uv-install.sh" | sha256sum -c -
env UV_NO_MODIFY_PATH=1 sh /tmp/uv-install.sh
rm -f /tmp/uv-install.sh
echo 'export PATH="$HOME/.local/bin:$PATH"' >> "$BASH_ENV"
export PATH="$HOME/.local/bin:$PATH"
install_rust:
steps:
- run:
name: Install Rust (rustup 1.28.2, toolchain 1.98.0)
command: |
case "$(uname -m)" in
x86_64)
RUSTUP_TRIPLE=x86_64-unknown-linux-gnu
RUSTUP_SHA256=20a06e644b0d9bd2fbdbfd52d42540bdde820ea7df86e92e533c073da0cdd43c
;;
aarch64)
RUSTUP_TRIPLE=aarch64-unknown-linux-gnu
RUSTUP_SHA256=e3853c5a252fca15252d07cb23a1bdd9377a8c6f3efa01531109281ae47f841c
;;
*)
echo "install_rust: unsupported architecture $(uname -m)" >&2
exit 1
;;
esac
curl -sSLf -o /tmp/rustup-init \
"https://static.rust-lang.org/rustup/archive/1.28.2/${RUSTUP_TRIPLE}/rustup-init"
echo "${RUSTUP_SHA256} /tmp/rustup-init" | sha256sum -c -
chmod +x /tmp/rustup-init
/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"
export PATH="$HOME/.cargo/bin:$PATH"
rustc --version
cargo --version
install_codecov_cli:
steps:
- run:
name: Install Codecov CLI (pinned v11.3.1)
command: |
curl -sSLf -o /tmp/codecov https://cli.codecov.io/v11.3.1/linux/codecov
curl -sSLf -o /tmp/codecov.SHA256SUM https://cli.codecov.io/v11.3.1/linux/codecov.SHA256SUM
[ "$(cat /tmp/codecov.SHA256SUM)" = "ca1d64196d2d34771084afe76ea657d581bf628e31d993ff8e52ea09cc88a56d codecov" ]
(cd /tmp && sha256sum -c codecov.SHA256SUM)
chmod +x /tmp/codecov
mkdir -p "$HOME/.local/bin"
mv /tmp/codecov "$HOME/.local/bin/codecov"
setup_litellm_enterprise_pip:
steps:
- run:
name: "Install local version of litellm-enterprise"
command: |
uv run --no-sync python -c "import litellm_enterprise; print('litellm-enterprise OK:', litellm_enterprise.__file__)"
setup_test_deps:
steps:
- checkout
- 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" }}
- run:
name: Generate Prisma client
command: uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma
skip_unless_relevant:
parameters:
category:
type: string
default: backend
base_ref:
type: string
default: ""
pull_request_url:
type: string
default: ""
steps:
- run:
name: "Skip job when no << parameters.category >>-relevant files changed"
command: |
export CIRCLE_PULL_REQUEST="${CIRCLE_PULL_REQUEST:-<< parameters.pull_request_url >>}"
export PATH_FILTER_BASE_BRANCH="<< parameters.base_ref >>"
[ -n "$PATH_FILTER_BASE_BRANCH" ] || unset PATH_FILTER_BASE_BRANCH
bash .circleci/scripts/path_filter.sh << parameters.category >>
start_postgres:
parameters:
db_name:
type: string
default: circle_test
image:
type: string
default: postgres:14@sha256:6a70deda415ec296f977890e11aba04a0db9f632a362e3fce45e845e3db74f26
steps:
- run:
name: Start PostgreSQL
command: |
docker run -d \
--name postgres-db \
-e POSTGRES_USER=postgres \
-e POSTGRES_PASSWORD=postgres \
-e POSTGRES_DB=<< parameters.db_name >> \
-p 5432:5432 \
<< parameters.image >>
- wait_for_service:
url: tcp://localhost:5432
timeout: "60"
start_redis:
steps:
- run:
name: Start Redis
command: |
docker run -d \
--name redis-cache \
-p 6379:6379 \
redis:7-alpine@sha256:7aec734b2bb298a1d769fd8729f13b8514a41bf90fcdd1f38ec52267fbaa8ee6
- wait_for_service:
url: tcp://localhost:6379
timeout: "60"
jobs:
unit:
parameters:
tests_path:
type: string
default: tests/unit
flag:
type: string
default: unit
shards:
type: integer
default: 6
base_ref:
type: string
default: ""
pull_request_url:
type: string
default: ""
machine:
image: ubuntu-2204:2024.04.1
resource_class: large
working_directory: ~/project
parallelism: << parameters.shards >>
environment:
LITELLM_LOCAL_MODEL_COST_MAP: "True"
steps:
- setup_test_deps
- skip_unless_relevant:
base_ref: << parameters.base_ref >>
pull_request_url: << parameters.pull_request_url >>
- run:
name: "Run << parameters.tests_path >> shard"
no_output_timeout: 20m
command: |
mkdir -p test-results/<< parameters.flag >>
mapfile -t files < <(find << parameters.tests_path >> -name 'test_*.py' | sort | circleci tests split --split-by=timings --timings-type=filename)
if [ "${#files[@]}" -eq 0 ]; then echo "shard ${CIRCLE_NODE_INDEX} received no << parameters.tests_path >> files; nothing to run"; exit 0; fi
set +e
uv run --no-sync pytest "${files[@]}" -p no:rerunfailures -p no:pytest-retry --timeout=90 -n 4 --dist=loadscope --tb=short --durations=20 -o junit_family=xunit1 --junitxml=test-results/<< parameters.flag >>/junit.xml --cov=./litellm --cov=./enterprise/litellm_enterprise --cov-report=xml:coverage.xml --cov-config=pyproject.toml
status=$?
set -e
if [ "$status" -eq 5 ]; then echo "pytest collected no tests from the shard; passing"; exit 0; fi
exit "$status"
- install_codecov_cli
- run:
name: Upload coverage
when: always
command: |
[ -f coverage.xml ] || { echo "no coverage.xml produced; skipping upload"; exit 0; }
codecov upload-process --disable-search -f coverage.xml -F << parameters.flag >> -C "$CIRCLE_SHA1" -n "<< parameters.flag >>-${CIRCLE_NODE_INDEX}-${CIRCLE_BUILD_NUM}" --git-service github
- store_test_results:
path: test-results
- store_artifacts:
path: test-results
- store_artifacts:
path: coverage.xml
documentation:
machine:
image: ubuntu-2204:2024.04.1
resource_class: large
working_directory: ~/project
steps:
- setup_test_deps
- run:
name: Checkout litellm-docs
command: rm -rf docs/my-website && git clone --depth 1 https://github.com/BerriAI/litellm-docs.git docs/my-website
- run:
name: Run documentation validation
command: |
uv run --no-sync python ./tests/documentation_tests/test_env_keys.py
uv run --no-sync python ./tests/documentation_tests/test_router_settings.py
uv run --no-sync python ./tests/documentation_tests/test_api_docs.py
uv run --no-sync python ./tests/documentation_tests/test_circular_imports.py
integration:
parameters:
suite:
type: string
base_ref:
type: string
default: ""
pull_request_url:
type: string
default: ""
machine:
image: ubuntu-2204:2024.04.1
resource_class: large
working_directory: ~/project
steps:
- setup_test_deps
- skip_unless_relevant:
base_ref: << parameters.base_ref >>
pull_request_url: << parameters.pull_request_url >>
- start_postgres:
image: postgres:16@sha256:e17e86066e5ef83e0952a9347f5c792b7ece00972e2aa787a6986f471b3dd3d5
- start_redis
- run:
name: Run owned integration contracts
command: bash .circleci/scripts/run_integration.sh << parameters.suite >>
no_output_timeout: 15m
- run:
name: Stop owned database and Redis
when: always
command: |
mkdir -p test-results/integration-<< parameters.suite >>
docker logs postgres-db > test-results/integration-<< parameters.suite >>/postgres.log 2>&1 || true
docker logs redis-cache > test-results/integration-<< parameters.suite >>/redis.log 2>&1 || true
docker rm -f postgres-db redis-cache
test -z "$(docker ps -aq --filter name=postgres-db --filter name=redis-cache)"
- store_test_results:
path: test-results
- store_artifacts:
path: test-results
workflows:
tests:
when: (pipeline.event.name == "push" and pipeline.git.branch == "main") or pipeline.event.name == "api" or (pipeline.event.name == "pull_request" and (pipeline.event.github.pull_request.base.ref == "main" or pipeline.event.github.pull_request.base.ref starts-with "litellm_"))
jobs:
- unit:
base_ref: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.base.ref or "" >>
pull_request_url: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.url or "" >>
- documentation
- integration:
name: integration-<< matrix.suite >>
matrix:
parameters:
suite: [sdk]
base_ref: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.base.ref or "" >>
pull_request_url: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.url or "" >>

View file

@ -10,12 +10,13 @@ test_paths:
paths:
- tests/rust-python-harness
- reason: >-
What is left of the caching suite in tests/local_testing that runs nowhere. Every job that
globs that directory either deselects it (local_testing_part1 and part2 carry `-k "... and
not caching and not cache"`) or keeps only another keyword (langfuse, router, assistants),
and no job names these files the way redis_caching_unit_tests names test_dual_cache.py.
The gap was eight files and 118 tests when measured 2026-08-20; the five keyless ones now
run in the caching-local shard, leaving these three. Measured 2026-08-21 with no provider
Live-provider caching cases in tests/local_testing that remain outside CI. Jobs that
glob that directory either deselect them (local_testing_part1 and part2 carry `-k "... and
not caching and not cache"`) or keep only another keyword (langfuse, router, assistants).
Separately, test-redis-compat.yml selects two IAM cluster authentication tests in
test_caching.py by node ID. It does not run that file's other tests.
The gap was eight files and 118 tests when measured 2026-08-20; the five keyless files now
run in the caching-local shard, leaving live cases in these three. Measured 2026-08-21 with no provider
credentials and no Redis: test_caching.py needs both (37 of 65 fail without them),
test_disk_cache_unit_tests.py needs OPENAI_API_KEY for 2 of its 4, and
test_gcs_cache_unit_tests.py needs GCS credentials for all 4. They want the keyless/live

View file

@ -1,6 +1,8 @@
<!-- The whole description's target audience is humans, not AI agents: write it in plain, simple,
everyday engineering language, extremely parsable and readable at a glance. This goes double for
the TLDR, User Flow, and Caveats sections -->
the TLDR, User Flow, and Caveats sections
Drop every section you have nothing to put in, heading included: a bare "## Relevant issues" or
"## Affected release" with nothing under it must not appear in the final description -->
## TLDR
@ -21,6 +23,7 @@ How it solves it:
<!-- Two ordered lists, Before and After, walking the same end user through the same task, written strictly from that user's seat
Read the linked issue, ticket, or customer thread first so the flow reflects the real application and the routes its users actually hit; don't invent a generic scenario
Lead each list with one plain sentence saying where the flow fails (Before) or succeeds (After), then number the steps
Keep it tight: aim for 3 to 5 steps per list, one line each, roughly 20 words max, and never pad a shorter flow with filler steps to hit the count. Cover the one path the PR changes and fold variants (case, other field, second endpoint) into a clause on the step they belong to rather than their own steps. The example below is the target length
Every step is something the user does or observes: the HTTP method and full URL they hit, what they sent, and what visibly came back (status code, error text, the shape of an ID). UI steps name the page URL and what is on screen
No LiteLLM internals: never name functions, files, DB tables, config classes, hooks, callbacks, or code paths. "The upload hands back an ID that looks like OpenAI's own `file-abc123` instead of the scrambled one the gateway returned" is right, "no managed-file row was registered" is wrong
Keep the two lists step-for-step identical until they diverge, so the changed step is obvious
@ -45,15 +48,15 @@ After: the same request comes back with real token counts, so the dashboard show
## Relevant issues
<!-- e.g., "Fixes #000" -->
<!-- e.g., "Fixes #000". Drop the section if there is none -->
## 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. Leave the section blank 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", 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 -->
## Linear ticket
<!-- if you are an internal contributor, add "Resolves " followed by the Linear ticket e.g., "Resolves LIT-1234" to link the Linear ticket to the GitHub PR. If you don't have one, leave the section blank rather than guessing -->
<!-- if you are an internal contributor, add "Resolves " followed by the Linear ticket e.g., "Resolves LIT-1234" to link the Linear ticket to the GitHub PR. If you don't have one, drop the section rather than guessing -->
## Pre-Submission checklist
@ -134,7 +137,7 @@ If you're seeing a delay in your PR being merged, ping the LiteLLM Team on [Slac
human reader
If you assumed something instead of testing it, e.g. "only reproduces with X on" or "no
user-observable behavior difference", list it here too with what breaks if it is wrong
Leave this section empty if there are none -->
Drop this section if there are none -->
## QA runbook

View file

@ -9,6 +9,7 @@ on:
- "litellm/_redis.py"
- "litellm/_redis_credential_provider.py"
- "tests/test_litellm/test_redis.py"
- "tests/local_testing/test_caching.py"
- "tests/test_litellm/caching/test_redis_connection_pool.py"
- ".github/workflows/test-redis-compat.yml"
- "pyproject.toml"
@ -26,6 +27,9 @@ jobs:
name: "redis-py ${{ matrix.redis-version }}"
runs-on: ubuntu-latest
timeout-minutes: 15
permissions:
contents: read
id-token: write
strategy:
fail-fast: false
@ -55,7 +59,7 @@ jobs:
- name: Install dependencies
run: |
.github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router
.github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra extra_proxy --extra semantic-router
- name: Pin redis-py to the matrix version
env:
@ -64,12 +68,33 @@ jobs:
uv pip install "redis==${REDIS_VERSION:?}"
uv run --no-sync python -c "import redis; assert redis.__version__ == '${REDIS_VERSION:?}', redis.__version__; print('redis-py', redis.__version__)"
- name: Build Redis for cluster authentication tests
run: |
curl --fail --location --retry 3 https://download.redis.io/releases/redis-7.2.16.tar.gz -o "$RUNNER_TEMP/redis-7.2.16.tar.gz"
echo "960a8ec15e34ff40e57ff16837b26b33bd81f2da6d24497bb63de532a323a18e $RUNNER_TEMP/redis-7.2.16.tar.gz" | sha256sum --check
tar -xzf "$RUNNER_TEMP/redis-7.2.16.tar.gz" -C "$RUNNER_TEMP"
make -C "$RUNNER_TEMP/redis-7.2.16" -j2 MALLOC=libc OPTIMIZATION=-O1 redis-server
echo "$RUNNER_TEMP/redis-7.2.16/src" >> "$GITHUB_PATH"
- name: Run redis unit tests
run: |
redis-server --version
uv run --no-sync pytest \
tests/test_litellm/test_redis.py \
tests/test_litellm/caching/test_redis_connection_pool.py \
tests/local_testing/test_caching.py::test_sync_cluster_authenticates_with_azure_credentials \
tests/local_testing/test_caching.py::test_sync_cluster_authenticates_with_gcp_credentials \
--tb=short -vv \
--reruns 2 \
--reruns-delay 1 \
--durations=20
--durations=20 \
--cov=./litellm --cov-report=xml:coverage-redis.xml
- name: Upload Redis coverage
if: matrix.redis-version == '5.3.1'
uses: codecov/codecov-action@75cd11691c0faa626561e295848008c8a7dddffe # v5.5.4
with:
use_oidc: true
files: coverage-redis.xml
flags: redis-compat
fail_ci_if_error: false

View file

@ -111,6 +111,7 @@ jobs:
tests/test_litellm/images
tests/test_litellm/interactions
tests/test_litellm/messages
tests/test_litellm/embeddings
tests/test_litellm/ocr
tests/test_litellm/passthrough
tests/test_litellm/rag

View file

@ -33,11 +33,11 @@ End-to-end tests belong in `tests/e2e/` and must follow the harness conventions
When creating PRs, target the repository's current default branch for both internal and external / OSS contributions. Check it with `python3 scripts/default_branch.py --branch` instead of assuming a branch name or relying on cached `origin/HEAD`
When writing a PR body, treat the comments and imperative instructions inside .github/pull_request_template.md as rules to follow, not just layout. Agent harnesses may strip HTML comments from copies of that file injected into context, so read .github/pull_request_template.md from disk before writing a PR body to make sure you see every comment rule
When writing a PR body, treat the comments and imperative instructions inside .github/pull_request_template.md as rules to follow, not just layout. Agent harnesses may strip HTML comments from copies of that file injected into context, so read .github/pull_request_template.md from disk before writing a PR body to make sure you see every comment rule. A section you have nothing to put in (Relevant issues, Affected release, Linear ticket, Caveats, QA runbook, and so on) is removed entirely, heading included, never left as an empty title
Same applies for filing bug reports and feature requests, with .github/ISSUE_TEMPLATE/bug_report.yml and .github/ISSUE_TEMPLATE/feature_request.yml, respectively
If you're resolving a linear ticket, in the "## Linear ticket" section of the PR, say "Resolves LIT-1234", replacing "LIT-1234" with the actual ticket id that you're resolving. If you don't have the ticket id, don't make one up or search for it. Just leave the section blank
If you're resolving a linear ticket, in the "## Linear ticket" section of the PR, say "Resolves LIT-1234", replacing "LIT-1234" with the actual ticket id that you're resolving. If you don't have the ticket id, don't make one up or search for it. Just drop the section
Never use `pytest` commands or the like as "Screenshots / Proof of Fix". We prefer curl'ing a live proxy instance running on localhost:4000 (I like to run it with `python litellm/proxy/proxy_cli.py --config litellm/proxy/dev_config.yaml --detailed_debug --reload --use_v2_migration_resolver 2>&1 | tee litellm.log`; the Admin UI dev server is `npm run dev` in `ui/litellm-dashboard`, served on port 3000) and showing both the command run and the output. Also, it should hit real LLM provider APIs, not mocks, and cost real $$$ because that is the most realistic test. The proof of fix should be exactly what the end user / customer would see / do. The run logs in PR #27703 is a prime example of how to do it (not a huge fan of using a python test script that future me and the team will have no visibility into; I prefer just curl commands or a short list of bash commands (e.g., using `for`)). If it's a UI thing, or the main use case runs through a headful agentic coding tool like Claude Code or Codex, drive that surface yourself and embed your own before and after screenshots of it in the PR (the Admin UI page, or what the coding tool shows), next to an ordered list of the URLs to go to (e.g., http://localhost:4000/ui/?page=logs), where to click, and what fields to fill out so a reviewer can reproduce it

View file

@ -3414,6 +3414,7 @@ dependencies = [
name = "litellm-types"
version = "0.1.0"
dependencies = [
"rstest",
"serde",
"serde_json",
]

View file

@ -0,0 +1,274 @@
use serde_json::Value;
#[derive(Clone, Debug, PartialEq, Eq)]
enum Segment {
Field(String),
Every,
Index(usize),
}
fn parse_segments(path: &str) -> Option<Vec<Segment>> {
let mut segments = Vec::new();
let mut rest = path;
while !rest.is_empty() {
if let Some(after_open) = rest.strip_prefix('[') {
let (inside, after) = after_open.split_once(']')?;
segments.push(match inside {
"*" => Segment::Every,
index => Segment::Index(index.trim().parse().ok()?),
});
rest = after.strip_prefix('.').unwrap_or(after);
continue;
}
let end = rest.find(['.', '[']).unwrap_or(rest.len());
let (field, after) = rest.split_at(end);
if !field.is_empty() {
segments.push(Segment::Field(field.to_string()));
}
rest = after.strip_prefix('.').unwrap_or(after);
}
Some(segments)
}
fn without_path(value: Value, segments: &[Segment]) -> Value {
let Some((segment, tail)) = segments.split_first() else {
return value;
};
match (segment, value) {
(Segment::Field(name), Value::Object(object)) => Value::Object(
object
.into_iter()
.filter_map(|(key, item)| {
if key != *name {
return Some((key, item));
}
(!tail.is_empty()).then(|| (key, without_path(item, tail)))
})
.collect(),
),
(Segment::Every, Value::Array(items)) => Value::Array(
items
.into_iter()
.map(|item| without_path(item, tail))
.collect(),
),
(Segment::Index(index), Value::Array(items)) => Value::Array(
items
.into_iter()
.enumerate()
.map(|(position, item)| {
if position == *index {
without_path(item, tail)
} else {
item
}
})
.collect(),
),
(_, value) => value,
}
}
pub fn delete_nested_value(value: Value, path: &str) -> Value {
match parse_segments(path) {
Some(segments) => without_path(value, &segments),
None => value,
}
}
#[cfg(test)]
mod tests {
use rstest::{fixture, rstest};
use serde_json::json;
use super::*;
#[fixture]
fn body() -> Value {
json!({
"tools": [
{"name": "t0", "examples": ["a"], "arr": [{"f": 1, "k": 1}, {"f": 2, "k": 2}]},
{"name": "t1", "examples": ["b"], "arr": [{"f": 3, "k": 3}]}
],
"meta": {"user": "u", "inner": {"drop": 1, "keep": 2}},
"top": 0.7
})
}
#[rstest]
#[case::top_level_field("top", json!({
"tools": [
{"name": "t0", "examples": ["a"], "arr": [{"f": 1, "k": 1}, {"f": 2, "k": 2}]},
{"name": "t1", "examples": ["b"], "arr": [{"f": 3, "k": 3}]}
],
"meta": {"user": "u", "inner": {"drop": 1, "keep": 2}}
}))]
#[case::whole_object("meta", json!({
"tools": [
{"name": "t0", "examples": ["a"], "arr": [{"f": 1, "k": 1}, {"f": 2, "k": 2}]},
{"name": "t1", "examples": ["b"], "arr": [{"f": 3, "k": 3}]}
],
"top": 0.7
}))]
#[case::nested_field("meta.inner.drop", json!({
"tools": [
{"name": "t0", "examples": ["a"], "arr": [{"f": 1, "k": 1}, {"f": 2, "k": 2}]},
{"name": "t1", "examples": ["b"], "arr": [{"f": 3, "k": 3}]}
],
"meta": {"user": "u", "inner": {"keep": 2}},
"top": 0.7
}))]
#[case::trailing_dot("meta.inner.drop.", json!({
"tools": [
{"name": "t0", "examples": ["a"], "arr": [{"f": 1, "k": 1}, {"f": 2, "k": 2}]},
{"name": "t1", "examples": ["b"], "arr": [{"f": 3, "k": 3}]}
],
"meta": {"user": "u", "inner": {"keep": 2}},
"top": 0.7
}))]
#[case::leading_and_doubled_dots(".meta..inner.drop", json!({
"tools": [
{"name": "t0", "examples": ["a"], "arr": [{"f": 1, "k": 1}, {"f": 2, "k": 2}]},
{"name": "t1", "examples": ["b"], "arr": [{"f": 3, "k": 3}]}
],
"meta": {"user": "u", "inner": {"keep": 2}},
"top": 0.7
}))]
#[case::field_in_every_element("tools[*].examples", json!({
"tools": [
{"name": "t0", "arr": [{"f": 1, "k": 1}, {"f": 2, "k": 2}]},
{"name": "t1", "arr": [{"f": 3, "k": 3}]}
],
"meta": {"user": "u", "inner": {"drop": 1, "keep": 2}},
"top": 0.7
}))]
#[case::whole_array_field_in_every_element("tools[*].arr", json!({
"tools": [
{"name": "t0", "examples": ["a"]},
{"name": "t1", "examples": ["b"]}
],
"meta": {"user": "u", "inner": {"drop": 1, "keep": 2}},
"top": 0.7
}))]
#[case::field_in_indexed_element("tools[1].examples", json!({
"tools": [
{"name": "t0", "examples": ["a"], "arr": [{"f": 1, "k": 1}, {"f": 2, "k": 2}]},
{"name": "t1", "arr": [{"f": 3, "k": 3}]}
],
"meta": {"user": "u", "inner": {"drop": 1, "keep": 2}},
"top": 0.7
}))]
#[case::padded_index("tools[ 1 ].examples", json!({
"tools": [
{"name": "t0", "examples": ["a"], "arr": [{"f": 1, "k": 1}, {"f": 2, "k": 2}]},
{"name": "t1", "arr": [{"f": 3, "k": 3}]}
],
"meta": {"user": "u", "inner": {"drop": 1, "keep": 2}},
"top": 0.7
}))]
#[case::field_right_after_bracket("tools[0]examples", json!({
"tools": [
{"name": "t0", "arr": [{"f": 1, "k": 1}, {"f": 2, "k": 2}]},
{"name": "t1", "examples": ["b"], "arr": [{"f": 3, "k": 3}]}
],
"meta": {"user": "u", "inner": {"drop": 1, "keep": 2}},
"top": 0.7
}))]
#[case::index_then_wildcard("tools[0].arr[*].f", json!({
"tools": [
{"name": "t0", "examples": ["a"], "arr": [{"k": 1}, {"k": 2}]},
{"name": "t1", "examples": ["b"], "arr": [{"f": 3, "k": 3}]}
],
"meta": {"user": "u", "inner": {"drop": 1, "keep": 2}},
"top": 0.7
}))]
#[case::wildcard_then_index_only_where_it_exists("tools[*].arr[1].f", json!({
"tools": [
{"name": "t0", "examples": ["a"], "arr": [{"f": 1, "k": 1}, {"k": 2}]},
{"name": "t1", "examples": ["b"], "arr": [{"f": 3, "k": 3}]}
],
"meta": {"user": "u", "inner": {"drop": 1, "keep": 2}},
"top": 0.7
}))]
#[case::nested_wildcards("tools[*].arr[*].f", json!({
"tools": [
{"name": "t0", "examples": ["a"], "arr": [{"k": 1}, {"k": 2}]},
{"name": "t1", "examples": ["b"], "arr": [{"k": 3}]}
],
"meta": {"user": "u", "inner": {"drop": 1, "keep": 2}},
"top": 0.7
}))]
fn deletes_the_addressed_field(body: Value, #[case] path: &str, #[case] expected: Value) {
assert_eq!(delete_nested_value(body, path), expected);
}
#[rstest]
#[case::empty_path("")]
#[case::missing_field("missing")]
#[case::missing_parent("missing.field")]
#[case::field_through_a_scalar("top.value")]
#[case::field_on_an_array("tools.name")]
#[case::index_on_an_object("meta[0].user")]
#[case::wildcard_on_an_object("meta[*].user")]
#[case::wildcard_over_scalars("tools[*].examples[*].name")]
#[case::index_out_of_range("tools[5].name")]
#[case::every_element_itself("tools[*]")]
#[case::indexed_element_itself("tools[0]")]
#[case::nested_element_itself("tools[*].arr[0]")]
#[case::negative_index("tools[-1].name")]
#[case::non_numeric_index("tools[x].name")]
#[case::empty_index("tools[].name")]
#[case::unclosed_bracket("top[0")]
fn leaves_the_value_untouched(body: Value, #[case] path: &str) {
assert_eq!(delete_nested_value(body.clone(), path), body);
}
#[rstest]
#[case::wildcards_indices_and_nesting(
json!({"tools": [
{"name": "t0", "configs": [{"id": "c0", "remove_me": 1, "keep": 1}, {"id": "c1", "remove_me": 2, "keep": 2}], "metadata": {"drop_this": 1, "preserve": 1}},
{"name": "t1", "configs": [{"id": "c0", "remove_me": 3, "keep": 3}, {"id": "c1", "remove_me": 4, "keep": 4}], "metadata": {"drop_this": 2, "preserve": 2}},
{"name": "t2", "configs": [{"id": "c0", "remove_me": 5, "keep": 5}], "metadata": {"drop_this": 3, "preserve": 3}}
]}),
&["tools[*].configs[1].remove_me", "tools[1].metadata.drop_this", "tools[*].configs[*].id"],
json!({"tools": [
{"name": "t0", "configs": [{"remove_me": 1, "keep": 1}, {"keep": 2}], "metadata": {"drop_this": 1, "preserve": 1}},
{"name": "t1", "configs": [{"remove_me": 3, "keep": 3}, {"keep": 4}], "metadata": {"preserve": 2}},
{"name": "t2", "configs": [{"remove_me": 5, "keep": 5}], "metadata": {"drop_this": 3, "preserve": 3}}
]}),
)]
#[case::simple_and_wildcard_nesting(
json!({
"tools": [{"name": "t1", "simple_nested": {"remove": 1, "keep": 2}, "complex": [{"nested": {"remove": 3, "keep": 4}}]}],
"top_level_remove": "should_go",
"top_level_keep": "should_stay"
}),
&["tools[*].simple_nested.remove", "tools[*].complex[*].nested.remove"],
json!({
"tools": [{"name": "t1", "simple_nested": {"keep": 2}, "complex": [{"nested": {"keep": 4}}]}],
"top_level_remove": "should_go",
"top_level_keep": "should_stay"
}),
)]
#[case::triple_nested_wildcards(
json!({"tools": [{"name": "t1", "arr1": [
{"arr2": [{"field": 1, "keep": 1}, {"field": 2, "keep": 2}]},
{"arr2": [{"field": 3, "keep": 3}]}
]}]}),
&["tools[*].arr1[*].arr2[*].field"],
json!({"tools": [{"name": "t1", "arr1": [
{"arr2": [{"keep": 1}, {"keep": 2}]},
{"arr2": [{"keep": 3}]}
]}]}),
)]
fn applies_paths_in_sequence(
#[case] value: Value,
#[case] paths: &[&str],
#[case] expected: Value,
) {
let deleted = paths
.iter()
.fold(value, |value, path| delete_nested_value(value, path));
assert_eq!(deleted, expected);
}
}

View file

@ -0,0 +1,93 @@
use litellm_types::utils::{ProviderSpecificHeader, ProviderSpecificHeaders};
use serde_json::{Map, Value};
pub fn get_provider_specific_headers(
provider_specific_header: Option<&ProviderSpecificHeaders>,
custom_llm_provider: &str,
) -> Map<String, Value> {
let entries: &[ProviderSpecificHeader] = match provider_specific_header {
None => &[],
Some(ProviderSpecificHeaders::One(entry)) => std::slice::from_ref(entry),
Some(ProviderSpecificHeaders::Many(entries)) => entries,
};
entries
.iter()
.filter(|entry| {
entry
.custom_llm_provider
.split(',')
.any(|scoped| scoped.trim() == custom_llm_provider)
})
.flat_map(|entry| entry.extra_headers.clone())
.collect()
}
#[cfg(test)]
mod tests {
use rstest::rstest;
use serde_json::json;
use super::*;
#[rstest]
#[case::single_entry_for_the_provider(
json!({"custom_llm_provider": "anthropic", "extra_headers": {"Authorization": "Bearer t", "Custom-Header": "v"}}),
json!({"Authorization": "Bearer t", "Custom-Header": "v"}),
)]
#[case::single_entry_for_another_provider(
json!({"custom_llm_provider": "openai", "extra_headers": {"Authorization": "Bearer t"}}),
json!({}),
)]
#[case::provider_in_a_comma_separated_scope(
json!({"custom_llm_provider": "bedrock,anthropic,vertex_ai", "extra_headers": {"anthropic-beta": "context-1m-2025-08-07"}}),
json!({"anthropic-beta": "context-1m-2025-08-07"}),
)]
#[case::provider_missing_from_a_comma_separated_scope(
json!({"custom_llm_provider": "bedrock,vertex_ai", "extra_headers": {"anthropic-beta": "test"}}),
json!({}),
)]
#[case::scope_with_spaces(
json!({"custom_llm_provider": "bedrock, anthropic , vertex_ai", "extra_headers": {"anthropic-beta": "test"}}),
json!({"anthropic-beta": "test"}),
)]
#[case::scope_names_must_match_exactly(
json!({"custom_llm_provider": "anthropic_text", "extra_headers": {"anthropic-beta": "test"}}),
json!({}),
)]
#[case::entries_scope_independently(
json!([
{"custom_llm_provider": "anthropic,bedrock,vertex_ai", "extra_headers": {"anthropic-beta": "context-1m-2025-08-07"}},
{"custom_llm_provider": "bedrock", "extra_headers": {"x-bedrock-only": "no"}},
{"custom_llm_provider": "anthropic", "extra_headers": {"authorization": "Bearer sk-ant-oat01-fake-token"}}
]),
json!({"anthropic-beta": "context-1m-2025-08-07", "authorization": "Bearer sk-ant-oat01-fake-token"}),
)]
#[case::later_entries_win(
json!([
{"custom_llm_provider": "anthropic", "extra_headers": {"x-scoped": "first"}},
{"custom_llm_provider": "anthropic", "extra_headers": {"x-scoped": "second"}}
]),
json!({"x-scoped": "second"}),
)]
#[case::empty_list(json!([]), json!({}))]
#[case::entry_without_scope(json!({"extra_headers": {"x-scoped": "yes"}}), json!({}))]
#[case::entry_without_headers(json!({"custom_llm_provider": "anthropic"}), json!({}))]
fn provider_specific_headers_match_the_scoped_provider(
#[case] configured: Value,
#[case] expected: Value,
) {
let configured: ProviderSpecificHeaders = serde_json::from_value(configured).unwrap();
assert_eq!(
Value::Object(get_provider_specific_headers(
Some(&configured),
"anthropic"
)),
expected
);
}
#[test]
fn no_configured_headers_match_nothing() {
assert_eq!(get_provider_specific_headers(None, "anthropic"), Map::new());
}
}

View file

@ -1,7 +1,9 @@
pub mod call_arguments;
pub mod core_helpers;
pub mod dot_notation_indexing;
pub mod exception_mapping_utils;
pub mod get_llm_provider_logic;
pub mod get_provider_specific_headers;
pub mod params;
pub mod prompt_templates;
pub mod secret_redaction;

View file

@ -1,5 +1,5 @@
use litellm_http::request::string_headers as shared_string_headers;
pub(super) use litellm_http::request::{has_bearer_auth, has_header, truncate_error_body};
pub(super) use litellm_http::request::truncate_error_body;
use litellm_llms::{
anthropic::experimental_pass_through::messages::transformation::ANTHROPIC_MESSAGES_CONFIG,
azure_ai::anthropic::messages_transformation::AZURE_ANTHROPIC_MESSAGES_CONFIG,

View file

@ -1,3 +1,5 @@
use std::sync::Arc;
use litellm_llms::base_llm::chat::transformation::Error as LlmError;
#[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)]
@ -18,8 +20,34 @@ pub enum Error {
Transport(#[from] litellm_http::transport::Error),
#[error(transparent)]
Headers(#[from] litellm_http::request::HeaderError),
#[error(transparent)]
Secret(#[from] SecretError),
}
#[derive(Clone, Debug, thiserror::Error)]
#[error(transparent)]
pub struct SecretError(Arc<litellm_secrets::Error>);
impl SecretError {
pub fn source_error(&self) -> &litellm_secrets::Error {
&self.0
}
}
impl From<litellm_secrets::Error> for Error {
fn from(error: litellm_secrets::Error) -> Self {
Self::Secret(SecretError(Arc::new(error)))
}
}
impl PartialEq for SecretError {
fn eq(&self, other: &Self) -> bool {
Arc::ptr_eq(&self.0, &other.0)
}
}
impl Eq for SecretError {}
impl From<LlmError> for Error {
fn from(error: LlmError) -> Self {
match error {

View file

@ -12,6 +12,9 @@ mod common_utils;
mod handler;
mod prepare;
pub mod route;
use std::sync::Arc;
use litellm_secrets::source::EnvironmentSecrets;
use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse;
use route::{LocalMessagesHost, MessagesCall, MessagesOutput, messages_machine};
use serde_json::Value;
@ -31,9 +34,12 @@ pub async fn messages(request: MessagesRequest<'_>) -> Result<AnthropicMessagesR
api_base: request.api_base.map(Into::into),
custom_llm_provider: request.custom_llm_provider.map(Into::into),
extra_headers: request.extra_headers,
provider_specific_header: request.provider_specific_header,
timeout: request.timeout,
shaping: request.shaping,
};
match litellm_host::run::run(messages_machine(), &LocalMessagesHost::new(call)).await? {
let secrets = Arc::new(EnvironmentSecrets::python_compatible());
match litellm_host::run::run(messages_machine(secrets), &LocalMessagesHost::new(call)).await? {
MessagesOutput::Message(message) => Ok(*message),
MessagesOutput::Streamed => Err(Error::Unsupported(
"streamed responses need a streaming host",

View file

@ -1,51 +1,102 @@
use litellm_core_utils::get_llm_provider_logic::{CustomLlmProvider, get_custom_llm_provider};
use litellm_llms::base_llm::anthropic_messages::transformation::{
BaseAnthropicMessagesConfig, MessagesAuthStrategy,
use litellm_core_utils::{
dot_notation_indexing::delete_nested_value,
get_llm_provider_logic::{CustomLlmProvider, get_custom_llm_provider},
get_provider_specific_headers::get_provider_specific_headers,
settings::Lookup,
};
use litellm_llms::{
anthropic::experimental_pass_through::messages::handler::shape_anthropic_messages_request,
base_llm::anthropic_messages::transformation::{
BaseAnthropicMessagesConfig, MessagesTransformContext,
},
};
use litellm_types::llms::anthropic_messages::anthropic_request::AnthropicMessagesRequest;
use serde_json::{Map, Value};
use super::{
Error,
common_utils::{has_bearer_auth, has_header, messages_provider_config, string_headers},
common_utils::{messages_provider_config, string_headers},
};
use crate::messages::types::{MessagesRequest, ProviderMessagesRequest};
pub(super) fn prepare_provider_request(
request: MessagesRequest<'_>,
) -> Result<ProviderMessagesRequest, Error> {
let provider_info = get_custom_llm_provider(request.model, request.custom_llm_provider)
pub(super) struct ResolvedProvider<'a> {
pub(super) model: &'a str,
pub(super) provider: &'a str,
pub(super) config: &'static dyn BaseAnthropicMessagesConfig,
}
pub(super) fn resolve_provider<'a>(
model: &'a str,
custom_llm_provider: Option<&'a str>,
) -> Result<ResolvedProvider<'a>, Error> {
let CustomLlmProvider {
model,
custom_llm_provider: provider,
} = get_custom_llm_provider(model, custom_llm_provider)
.or_else(|| {
request
.custom_llm_provider
.map(|provider| CustomLlmProvider {
model: request.model,
custom_llm_provider: provider,
})
custom_llm_provider.map(|provider| CustomLlmProvider {
model,
custom_llm_provider: provider,
})
})
.ok_or_else(|| {
Error::InvalidProvider(
"unable to resolve custom_llm_provider for messages request".to_string(),
)
})?;
let model = provider_info.model.to_string();
let provider = provider_info.custom_llm_provider;
let config = messages_provider_config(provider)
.ok_or_else(|| Error::InvalidProvider(provider.to_string()))?;
let env_lookup = |key: &str| std::env::var(key).ok();
Ok(ResolvedProvider {
model,
provider,
config,
})
}
let headers =
validate_environment(config, request.extra_headers, request.api_key, &env_lookup)?;
pub(super) fn prepare_provider_request(
request: MessagesRequest<'_>,
resolved: ResolvedProvider<'_>,
secrets: &dyn Lookup,
) -> Result<ProviderMessagesRequest, Error> {
let ResolvedProvider {
model,
provider,
config,
} = resolved;
let model = model.to_string();
let env_lookup = |key: &str| secrets.get(key);
let typed_request: AnthropicMessagesRequest =
serde_json::from_value(request.body).map_err(|err| {
Error::InvalidRequest(format!("invalid Anthropic messages request: {err}"))
})?;
let transformed = config.transform_anthropic_messages_request(AnthropicMessagesRequest {
model: model.clone(),
..typed_request
})?;
serde_json::from_value(request.body).map_err(invalid_request)?;
let sanitized = shape_anthropic_messages_request(
AnthropicMessagesRequest {
model: model.clone(),
..typed_request
},
request.shaping.reasoning_auto_summary,
)?;
let trimmed =
without_additional_drop_params(sanitized, &request.shaping.additional_drop_params)?;
let transformed = config.transform_anthropic_messages_request(
trimmed,
&MessagesTransformContext::new(request.shaping.capabilities, request.shaping.drop_params),
)?;
let scoped = get_provider_specific_headers(request.provider_specific_header.as_ref(), provider);
let forwarded = string_headers(Some(
request
.extra_headers
.into_iter()
.flatten()
.chain(scoped)
.collect(),
))?;
let authenticated = config.authenticate(forwarded, request.api_key, &env_lookup)?;
let headers = config.request_headers(
with_default_headers(authenticated, config.default_headers()),
&transformed,
);
let body = serde_json::to_value(transformed).map_err(|err| {
Error::InvalidRequest(format!(
"failed to serialize Anthropic messages request: {err}"
@ -65,33 +116,371 @@ pub(super) fn prepare_provider_request(
})
}
fn validate_environment(
config: &dyn BaseAnthropicMessagesConfig,
extra_headers: Option<Map<String, Value>>,
api_key: Option<&str>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> Result<Vec<(String, String)>, Error> {
let mut headers = string_headers(extra_headers)?;
let auth_strategy = config.auth_strategy();
let already_authorized = has_header(&headers, auth_strategy.header_name())
|| (config.accepts_bearer_auth() && has_bearer_auth(&headers));
if !already_authorized {
let api_key = config.resolve_api_key(api_key, env_lookup)?;
let auth_header = match auth_strategy {
MessagesAuthStrategy::Bearer => {
("authorization".to_string(), format!("Bearer {api_key}"))
}
MessagesAuthStrategy::Header(name) => (name.to_string(), api_key),
};
headers.push(auth_header);
}
for (name, value) in config.default_headers() {
if !has_header(&headers, name) {
headers.push((name.to_string(), value.to_string()));
}
}
Ok(headers)
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],
) -> Result<AnthropicMessagesRequest, Error> {
if paths.is_empty() {
return Ok(request);
}
let Value::Object(fields) = serde_json::to_value(request).map_err(invalid_request)? else {
return Err(Error::InvalidRequest(
"Anthropic messages request did not serialize to an object".to_string(),
));
};
let (required, optional): (Map<String, Value>, Map<String, Value>) = fields
.into_iter()
.partition(|(key, _)| matches!(key.as_str(), "model" | "messages"));
let trimmed = paths.iter().fold(Value::Object(optional), |body, path| {
delete_nested_value(body, path)
});
let merged: Map<String, Value> = required
.into_iter()
.chain(trimmed.as_object().cloned().unwrap_or_default())
.collect();
serde_json::from_value(Value::Object(merged)).map_err(invalid_request)
}
fn with_default_headers(
headers: Vec<(String, String)>,
defaults: &[(&str, &str)],
) -> Vec<(String, String)> {
let missing: Vec<(String, String)> = defaults
.iter()
.filter(|(name, _)| {
!headers
.iter()
.any(|(header, _)| header.eq_ignore_ascii_case(name))
})
.map(|(name, value)| ((*name).to_string(), (*value).to_string()))
.collect();
headers.into_iter().chain(missing).collect()
}
#[cfg(test)]
mod tests {
use litellm_types::utils::ProviderSpecificHeaders;
use rstest::{fixture, rstest};
use serde_json::json;
use super::*;
use crate::messages::types::MessagesShaping;
#[fixture]
fn shaping() -> MessagesShaping {
MessagesShaping::default()
}
fn prepare(request: MessagesRequest<'_>) -> Result<ProviderMessagesRequest, Error> {
prepare_with_secrets(request, &|_: &str| None)
}
fn prepare_with_secrets(
request: MessagesRequest<'_>,
secrets: &dyn Lookup,
) -> Result<ProviderMessagesRequest, Error> {
let resolved = resolve_provider(request.model, request.custom_llm_provider)?;
prepare_provider_request(request, resolved, secrets)
}
#[rstest]
#[case::api_key(
&[("ANTHROPIC_API_KEY", "sk-secret")],
&[("x-api-key", "sk-secret")],
"https://api.anthropic.com/v1/messages"
)]
#[case::auth_token(
&[("ANTHROPIC_AUTH_TOKEN", "token")],
&[("authorization", "Bearer token")],
"https://api.anthropic.com/v1/messages"
)]
#[case::api_base(
&[("ANTHROPIC_API_KEY", "sk-secret"), ("ANTHROPIC_API_BASE", "https://gateway.test")],
&[("x-api-key", "sk-secret")],
"https://gateway.test/v1/messages"
)]
#[case::sdk_base_url(
&[("ANTHROPIC_API_KEY", "sk-secret"), ("ANTHROPIC_BASE_URL", "https://sdk.test")],
&[("x-api-key", "sk-secret")],
"https://sdk.test/v1/messages"
)]
fn credentials_and_base_come_from_the_resolved_secrets(
shaping: MessagesShaping,
#[case] secrets: &[(&str, &str)],
#[case] expected_auth: &[(&str, &str)],
#[case] expected_url: &str,
) {
let lookup = |name: &str| {
secrets
.iter()
.find(|(key, _)| *key == name)
.map(|(_, value)| value.to_string())
};
let prepared = prepare_with_secrets(
MessagesRequest {
model: "claude-test",
body: json!({"model": "claude-test", "messages": [{"role": "user", "content": "hi"}], "max_tokens": 16}),
api_key: None,
api_base: None,
custom_llm_provider: Some("anthropic"),
extra_headers: None,
provider_specific_header: None,
timeout: None,
shaping,
},
&lookup,
)
.unwrap();
let auth: Vec<(&str, &str)> = prepared
.upstream_headers
.iter()
.filter(|(name, _)| matches!(name.as_str(), "x-api-key" | "authorization"))
.map(|(name, value)| (name.as_str(), value.as_str()))
.collect();
assert_eq!(
(auth.as_slice(), prepared.url.as_str()),
(expected_auth, expected_url)
);
}
fn prepared_body(body: Value, shaping: MessagesShaping) -> Result<Value, Error> {
prepare(MessagesRequest {
model: "anthropic/claude-test",
body,
api_key: Some("sk-test"),
api_base: Some("https://anthropic.test"),
custom_llm_provider: Some("anthropic"),
extra_headers: None,
provider_specific_header: None,
timeout: None,
shaping,
})
.map(|prepared| prepared.body)
}
#[rstest]
#[case::nothing_forwarded(
&[],
&[("x-version", "1"), ("content-type", "application/json")],
&[("x-version", "1"), ("content-type", "application/json")],
)]
#[case::forwarded_header_wins_in_any_case(
&[("X-Version", "custom"), ("x-api-key", "k")],
&[("x-version", "1"), ("content-type", "application/json")],
&[("X-Version", "custom"), ("x-api-key", "k"), ("content-type", "application/json")],
)]
#[case::no_defaults(&[("x-api-key", "k")], &[], &[("x-api-key", "k")])]
fn default_headers_fill_only_missing_names(
#[case] forwarded: &[(&str, &str)],
#[case] defaults: &[(&str, &str)],
#[case] expected: &[(&str, &str)],
) {
let owned = |headers: &[(&str, &str)]| -> Vec<(String, String)> {
headers
.iter()
.map(|(name, value)| ((*name).to_string(), (*value).to_string()))
.collect()
};
assert_eq!(
with_default_headers(owned(forwarded), defaults),
owned(expected)
);
}
#[rstest]
#[case::top_level_and_nested_paths(
json!({
"max_tokens": 1024,
"thinking": {"type": "enabled", "budget_tokens": 2048},
"context_management": {"edits": [{"type": "clear_thinking_20251015"}]},
"metadata": {"user_id": "u1"},
"tools": [{"name": "lookup", "input_schema": {"type": "object"}, "input_examples": [{"q": "x"}]}]
}),
&["thinking", "context_management", "tools[*].input_examples"],
json!({
"max_tokens": 1024,
"metadata": {"user_id": "u1"},
"tools": [{"name": "lookup", "input_schema": {"type": "object"}}]
}),
)]
#[case::no_paths(
json!({"max_tokens": 16, "safeguards": [{"type": "dangerous_tool_use"}]}),
&[],
json!({"max_tokens": 16, "safeguards": [{"type": "dangerous_tool_use"}]}),
)]
#[case::model_and_messages_are_never_dropped(
json!({"max_tokens": 16}),
&["model", "messages", "messages[0].content"],
json!({"max_tokens": 16}),
)]
fn prepared_body_drops_configured_paths(
shaping: MessagesShaping,
#[case] fields: Value,
#[case] additional_drop_params: &[&str],
#[case] expected_fields: Value,
) {
let with_messages = |fields: Value| -> Value {
let Value::Object(fields) = fields else {
unreachable!()
};
Value::Object(
[
("model".to_string(), json!("claude-test")),
(
"messages".to_string(),
json!([{"role": "user", "content": "hi"}]),
),
]
.into_iter()
.chain(fields)
.collect(),
)
};
let shaping = MessagesShaping {
additional_drop_params: additional_drop_params
.iter()
.map(ToString::to_string)
.collect(),
..shaping
};
assert_eq!(
prepared_body(with_messages(fields), shaping),
Ok(with_messages(expected_fields))
);
}
#[rstest]
#[case::model_prefix_picks_the_provider(
"azure_ai/claude-test",
None,
&[("x-priority", "extra"), ("x-scoped", "azure_ai")]
)]
#[case::explicit_provider(
"claude-test",
Some("anthropic"),
&[("x-priority", "scoped"), ("x-scoped", "anthropic")]
)]
#[case::provider_prefix_on_an_anthropic_model(
"anthropic/claude-test",
None,
&[("x-priority", "scoped"), ("x-scoped", "anthropic")]
)]
fn provider_specific_headers_follow_the_resolved_provider(
shaping: MessagesShaping,
#[case] model: &str,
#[case] custom_llm_provider: Option<&str>,
#[case] expected: &[(&str, &str)],
) {
let configured: ProviderSpecificHeaders = serde_json::from_value(json!([
{"custom_llm_provider": "azure_ai", "extra_headers": {"x-scoped": "azure_ai"}},
{"custom_llm_provider": "anthropic", "extra_headers": {"x-scoped": "anthropic", "x-priority": "scoped"}}
]))
.unwrap();
let prepared = prepare(MessagesRequest {
model,
body: json!({"model": model, "messages": [{"role": "user", "content": "hi"}], "max_tokens": 16}),
api_key: Some("sk-test"),
api_base: Some("https://resource.services.ai.azure.com"),
custom_llm_provider,
extra_headers: Some(serde_json::from_value(json!({"x-priority": "extra"})).unwrap()),
provider_specific_header: Some(configured),
timeout: None,
shaping,
})
.unwrap();
let caller_headers: Vec<(&str, &str)> = prepared
.upstream_headers
.iter()
.filter(|(name, _)| matches!(name.as_str(), "x-priority" | "x-scoped"))
.map(|(name, value)| (name.as_str(), value.as_str()))
.collect();
assert_eq!(caller_headers, expected);
}
#[rstest]
fn prepared_body_carries_the_provider_stripped_model(shaping: MessagesShaping) {
assert_eq!(
prepared_body(
json!({
"model": "anthropic/claude-test",
"messages": [{"role": "user", "content": "hi"}],
"max_tokens": 16
}),
shaping,
),
Ok(json!({
"model": "claude-test",
"messages": [{"role": "user", "content": "hi"}],
"max_tokens": 16
}))
);
}
#[rstest]
fn dropped_thinking_display_is_not_restored_by_auto_summary(shaping: MessagesShaping) {
let shaping = MessagesShaping {
reasoning_auto_summary: true,
additional_drop_params: vec!["thinking.display".to_string()],
..shaping
};
assert_eq!(
prepared_body(
json!({
"model": "claude-test",
"messages": [{"role": "user", "content": "hi"}],
"max_tokens": 4096,
"thinking": {"type": "enabled", "budget_tokens": 2048}
}),
shaping,
),
Ok(json!({
"model": "claude-test",
"messages": [{"role": "user", "content": "hi"}],
"max_tokens": 4096,
"thinking": {"type": "enabled", "budget_tokens": 2048}
}))
);
}
#[rstest]
fn dropping_an_invalid_metadata_user_id_does_not_skip_its_validation(shaping: MessagesShaping) {
let shaping = MessagesShaping {
additional_drop_params: vec!["metadata.user_id".to_string()],
..shaping
};
assert!(matches!(
prepared_body(
json!({
"model": "claude-test",
"messages": [{"role": "user", "content": "hi"}],
"max_tokens": 16,
"metadata": {"user_id": 123}
}),
shaping,
),
Err(Error::InvalidRequest(_))
));
}
#[rstest]
fn prepared_body_rejects_invalid_metadata_before_the_call(shaping: MessagesShaping) {
assert_eq!(
prepared_body(
json!({
"model": "claude-test",
"messages": [{"role": "user", "content": "hi"}],
"max_tokens": 16,
"metadata": {"user_id": 123}
}),
shaping,
),
Err(Error::InvalidRequest(
"metadata.user_id must be a string, got 123".to_string()
))
);
}
}

View file

@ -1,4 +1,7 @@
use std::{sync::Mutex, time::Duration};
use std::{
sync::{Arc, Mutex},
time::Duration,
};
use bytes::Bytes;
use litellm_auth::SecretValue;
@ -9,15 +12,19 @@ use litellm_host::{
machine::{HostChannel, MachineFault, RouteMachine},
route::Route,
};
use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse;
use litellm_secrets::source::SecretSource;
use litellm_types::{
llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse,
utils::ProviderSpecificHeaders,
};
use serde_json::{Map, Value};
use super::{
Error,
common_utils::messages_provider_config,
handler::{decode_response, network, provider_error, send},
prepare::prepare_provider_request,
types::MessagesRequest,
prepare::{prepare_provider_request, resolve_provider},
types::{MessagesRequest, MessagesShaping},
};
use crate::constants::ANTHROPIC_MESSAGES_PROVIDER;
@ -38,7 +45,9 @@ pub struct MessagesCall {
pub api_base: Option<String>,
pub custom_llm_provider: Option<String>,
pub extra_headers: Option<Map<String, Value>>,
pub provider_specific_header: Option<ProviderSpecificHeaders>,
pub timeout: Option<Duration>,
pub shaping: MessagesShaping,
}
impl MessagesCall {
@ -120,22 +129,33 @@ impl Host<Messages> for LocalMessagesHost {
}
}
pub fn messages_machine() -> MessagesMachine {
RouteMachine::new(|host| Box::pin(execute(host)))
pub fn messages_machine(secrets: Arc<dyn SecretSource>) -> MessagesMachine {
RouteMachine::new(move |host| Box::pin(execute(host, secrets.clone())))
}
async fn execute(host: MessagesHost) -> Result<MessagesOutput, Error> {
async fn execute(
host: MessagesHost,
secrets: Arc<dyn SecretSource>,
) -> Result<MessagesOutput, Error> {
let MessagesOpResult::Request(call) = host.route(MessagesOp::ProjectRequest).await?;
let stream = call.streams();
let request = prepare_provider_request(MessagesRequest {
model: &call.model,
body: Value::Object(call.body.clone()),
api_key: call.api_key.as_deref(),
api_base: call.api_base.as_deref(),
custom_llm_provider: call.custom_llm_provider.as_deref(),
extra_headers: call.extra_headers.clone(),
timeout: call.timeout,
})?;
let resolved = resolve_provider(&call.model, call.custom_llm_provider.as_deref())?;
let secrets = secrets.resolve(resolved.config.secret_names()).await?;
let request = prepare_provider_request(
MessagesRequest {
model: &call.model,
body: Value::Object(call.body.clone()),
api_key: call.api_key.as_deref(),
api_base: call.api_base.as_deref(),
custom_llm_provider: call.custom_llm_provider.as_deref(),
extra_headers: call.extra_headers.clone(),
provider_specific_header: call.provider_specific_header.clone(),
timeout: call.timeout,
shaping: call.shaping.clone(),
},
resolved,
secrets.as_ref(),
)?;
if stream && request.provider != ANTHROPIC_MESSAGES_PROVIDER {
return Err(Error::Unsupported("streaming messages for this provider"));
}

View file

@ -1,5 +1,8 @@
use std::time::Duration;
use std::{sync::Arc, time::Duration};
use futures_util::future::BoxFuture;
use litellm_http::request::{has_bearer_auth, has_header};
use litellm_secrets::{SecretValue, source::SecretSource};
use serde_json::{Map, Value, json};
use tokio::{
io::{AsyncReadExt, AsyncWriteExt},
@ -8,12 +11,132 @@ use tokio::{
use super::{
Error,
common_utils::{
has_bearer_auth, has_header, messages_provider_config, string_headers, truncate_error_body,
},
common_utils::{messages_provider_config, string_headers, truncate_error_body},
messages,
route::{LocalMessagesHost, MessagesCall, MessagesOutput, messages_machine},
};
use crate::messages::types::MessagesRequest;
use crate::messages::types::{MessagesRequest, MessagesShaping};
struct RecordingSecrets {
values: Vec<(&'static str, String)>,
fails: bool,
requested: std::sync::Mutex<Vec<String>>,
}
impl RecordingSecrets {
fn new(values: Vec<(&'static str, String)>, fails: bool) -> Self {
Self {
values,
fails,
requested: std::sync::Mutex::new(Vec::new()),
}
}
}
impl SecretSource for RecordingSecrets {
fn get_secret_str<'a>(
&'a self,
name: &'a str,
) -> BoxFuture<'a, Result<Option<SecretValue>, litellm_secrets::Error>> {
Box::pin(async move {
self.requested.lock().unwrap().push(name.to_string());
if self.fails {
return Err(litellm_secrets::Error::ManagedSecretMissing);
}
Ok(self
.values
.iter()
.find(|(key, _)| *key == name)
.map(|(_, value)| SecretValue::new(value.clone())))
})
}
}
fn secrets_call() -> MessagesCall {
let Value::Object(body) = json!({
"model": "claude-sonnet-4-5",
"max_tokens": 16,
"messages": [{"role": "user", "content": "hi"}]
}) else {
unreachable!("literal object")
};
MessagesCall {
model: "claude-sonnet-4-5".into(),
body,
api_key: None,
api_base: None,
custom_llm_provider: Some("anthropic".into()),
extra_headers: None,
provider_specific_header: None,
timeout: Some(Duration::from_secs(5)),
shaping: MessagesShaping::default(),
}
}
#[tokio::test]
async fn route_reads_the_provider_credential_and_base_from_the_secret_source() {
let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds");
let addr = listener.local_addr().expect("addr");
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.expect("accepts request");
let request = read_http_request(&mut socket).await;
let response_body = r#"{"id":"msg_1","type":"message","role":"assistant","content":[],"model":"claude-sonnet-4-5","stop_reason":"end_turn","usage":{"input_tokens":1,"output_tokens":1}}"#;
socket
.write_all(write_response(response_body).as_bytes())
.await
.expect("writes response");
request
});
let secrets = Arc::new(RecordingSecrets::new(
vec![
("ANTHROPIC_API_KEY", "sk-from-manager".to_string()),
("ANTHROPIC_BASE_URL", format!("http://{addr}")),
],
false,
));
let output = litellm_host::run::run(
messages_machine(secrets.clone()),
&LocalMessagesHost::new(secrets_call()),
)
.await
.expect("messages request succeeds");
assert!(matches!(output, MessagesOutput::Message(_)));
let request = server.await.expect("server task completes");
assert!(
request
.to_ascii_lowercase()
.contains("x-api-key: sk-from-manager"),
"{request}"
);
let requested = secrets.requested.lock().unwrap().clone();
assert_eq!(
requested,
messages_provider_config("anthropic")
.unwrap()
.secret_names()
.iter()
.map(ToString::to_string)
.collect::<Vec<_>>()
);
}
#[tokio::test]
async fn route_surfaces_a_secret_manager_failure_before_the_call() {
let Err(error) = litellm_host::run::run(
messages_machine(Arc::new(RecordingSecrets::new(Vec::new(), true))),
&LocalMessagesHost::new(secrets_call()),
)
.await
else {
panic!("a secret manager failure fails the call");
};
assert!(
matches!(&error, Error::Secret(source) if matches!(source.source_error(), litellm_secrets::Error::ManagedSecretMissing)),
"{error:?}"
);
}
async fn read_http_request(socket: &mut TcpStream) -> String {
let mut request = Vec::new();
@ -159,7 +282,9 @@ async fn messages_round_trip_builds_azure_request_and_passes_response_through()
api_base: Some(&format!("http://{addr}")),
custom_llm_provider: Some("azure_ai"),
extra_headers: None,
provider_specific_header: None,
timeout: Some(Duration::from_secs(5)),
shaping: MessagesShaping::default(),
})
.await
.expect("messages request succeeds");
@ -215,7 +340,9 @@ async fn messages_round_trip_builds_native_anthropic_request() {
api_base: Some(&format!("http://{addr}")),
custom_llm_provider: Some("anthropic"),
extra_headers: None,
provider_specific_header: None,
timeout: Some(Duration::from_secs(5)),
shaping: MessagesShaping::default(),
})
.await
.expect("messages request succeeds");
@ -268,7 +395,9 @@ async fn messages_does_not_duplicate_auth_when_x_api_key_supplied() {
api_base: Some(&format!("http://{addr}")),
custom_llm_provider: Some("azure_ai"),
extra_headers: Some(headers),
provider_specific_header: None,
timeout: Some(Duration::from_secs(5)),
shaping: MessagesShaping::default(),
})
.await
.expect("messages request succeeds");
@ -322,7 +451,9 @@ async fn messages_forwards_entra_id_bearer_without_requiring_api_key() {
api_base: Some(&format!("http://{addr}")),
custom_llm_provider: Some("azure_ai"),
extra_headers: Some(headers),
provider_specific_header: None,
timeout: Some(Duration::from_secs(5)),
shaping: MessagesShaping::default(),
})
.await
.expect("entra id request succeeds without api key");
@ -346,7 +477,9 @@ async fn messages_requires_auth_when_no_key_and_no_header() {
api_base: Some("http://127.0.0.1:1"),
custom_llm_provider: Some("azure_ai"),
extra_headers: None,
provider_specific_header: None,
timeout: Some(Duration::from_millis(50)),
shaping: MessagesShaping::default(),
})
.await
.expect_err("missing auth errors");
@ -384,7 +517,9 @@ async fn messages_ignores_malformed_authorization_and_uses_api_key() {
api_base: Some(&format!("http://{addr}")),
custom_llm_provider: Some("azure_ai"),
extra_headers: Some(headers),
provider_specific_header: None,
timeout: Some(Duration::from_secs(5)),
shaping: MessagesShaping::default(),
})
.await
.expect("falls back to api key");
@ -425,7 +560,9 @@ async fn messages_maps_provider_error_status_to_http_error() {
api_base: Some(&format!("http://{addr}")),
custom_llm_provider: Some("azure_ai"),
extra_headers: None,
provider_specific_header: None,
timeout: Some(Duration::from_secs(5)),
shaping: MessagesShaping::default(),
})
.await
.expect_err("provider error propagates");
@ -445,7 +582,9 @@ async fn messages_rejects_unsupported_provider() {
api_base: Some("http://127.0.0.1:1"),
custom_llm_provider: Some("openai"),
extra_headers: None,
provider_specific_header: None,
timeout: Some(Duration::from_millis(50)),
shaping: MessagesShaping::default(),
})
.await
.expect_err("unsupported provider errors");

View file

@ -1,8 +1,25 @@
use std::time::Duration;
use litellm_llms::base_llm::anthropic_messages::transformation::BaseAnthropicMessagesConfig;
use litellm_llms::{
anthropic::common_utils::AnthropicModelCapabilities,
base_llm::anthropic_messages::transformation::BaseAnthropicMessagesConfig,
};
use litellm_types::utils::ProviderSpecificHeaders;
use serde::{Deserialize, Serialize};
use serde_json::{Map, Value};
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
pub struct MessagesShaping {
#[serde(default)]
pub capabilities: AnthropicModelCapabilities,
#[serde(default)]
pub drop_params: bool,
#[serde(default)]
pub reasoning_auto_summary: bool,
#[serde(default)]
pub additional_drop_params: Vec<String>,
}
pub struct MessagesRequest<'a> {
pub model: &'a str,
pub body: Value,
@ -10,7 +27,9 @@ pub struct MessagesRequest<'a> {
pub api_base: Option<&'a str>,
pub custom_llm_provider: Option<&'a str>,
pub extra_headers: Option<Map<String, Value>>,
pub provider_specific_header: Option<ProviderSpecificHeaders>,
pub timeout: Option<Duration>,
pub shaping: MessagesShaping,
}
pub struct ProviderMessagesRequest {
@ -22,3 +41,86 @@ pub struct ProviderMessagesRequest {
pub upstream_headers: Vec<(String, String)>,
pub timeout: Option<Duration>,
}
#[cfg(test)]
mod tests {
use litellm_llms::anthropic::common_utils::SupportedEffortTiers;
use rstest::rstest;
use serde_json::json;
use super::*;
#[rstest]
#[case::nothing_projected(json!({}), MessagesShaping::default())]
#[case::only_drop_params(
json!({"drop_params": true}),
MessagesShaping { drop_params: true, ..MessagesShaping::default() },
)]
#[case::only_reasoning_auto_summary(
json!({"reasoning_auto_summary": true}),
MessagesShaping { reasoning_auto_summary: true, ..MessagesShaping::default() },
)]
#[case::only_additional_drop_params(
json!({"additional_drop_params": ["tools[*].input_examples"]}),
MessagesShaping {
additional_drop_params: vec!["tools[*].input_examples".to_string()],
..MessagesShaping::default()
},
)]
#[case::partial_capabilities(
json!({"capabilities": {"supports_reasoning": true}}),
MessagesShaping {
capabilities: AnthropicModelCapabilities {
supports_reasoning: true,
..AnthropicModelCapabilities::default()
},
..MessagesShaping::default()
},
)]
#[case::everything_the_python_host_projects(
json!({
"capabilities": {
"supports_reasoning": true,
"supports_adaptive_thinking": true,
"thinking_always_on": false,
"supports_legacy_thinking": false,
"supports_output_config": true,
"supports_sampling_params": false,
"supports_speed": true,
"effort_tiers": {"minimal": false, "low": true, "medium": true, "high": true, "xhigh": true, "max": false}
},
"drop_params": true,
"reasoning_auto_summary": true,
"additional_drop_params": ["metadata.user_id", "thinking"]
}),
MessagesShaping {
capabilities: AnthropicModelCapabilities {
supports_reasoning: true,
supports_adaptive_thinking: true,
thinking_always_on: false,
supports_legacy_thinking: false,
supports_output_config: true,
supports_sampling_params: false,
supports_speed: true,
effort_tiers: SupportedEffortTiers {
minimal: false,
low: true,
medium: true,
high: true,
xhigh: true,
max: false,
},
},
drop_params: true,
reasoning_auto_summary: true,
additional_drop_params: vec!["metadata.user_id".to_string(), "thinking".to_string()],
},
)]
fn shaping_deserializes_with_defaults_for_absent_fields(
#[case] projected: Value,
#[case] expected: MessagesShaping,
) {
let shaping: MessagesShaping = serde_json::from_value(projected).unwrap();
assert_eq!(shaping, expected);
}
}

File diff suppressed because it is too large Load diff

View file

@ -0,0 +1,270 @@
use litellm_types::llms::anthropic_messages::anthropic_request::{
AnthropicMessage, AnthropicMessagesRequest,
};
use serde_json::{Value, json};
use crate::{
anthropic::common_utils::{
flatten_unencrypted_web_search_results, sanitize_tool_use_ids, strip_empty_content_blocks,
strip_provider_specific_fields,
},
base_llm::chat::transformation::Error,
};
pub fn shape_anthropic_messages_request(
request: AnthropicMessagesRequest,
reasoning_auto_summary: bool,
) -> Result<AnthropicMessagesRequest, Error> {
Ok(AnthropicMessagesRequest {
messages: sanitize_anthropic_messages(request.messages),
metadata: request
.metadata
.as_ref()
.map(validate_anthropic_api_metadata)
.transpose()?,
thinking: with_reasoning_auto_summary(request.thinking, reasoning_auto_summary),
..request
})
}
fn sanitize_anthropic_messages(messages: Vec<AnthropicMessage>) -> Vec<AnthropicMessage> {
strip_provider_specific_fields(flatten_unencrypted_web_search_results(
sanitize_tool_use_ids(strip_empty_content_blocks(messages)),
))
}
fn validate_anthropic_api_metadata(metadata: &Value) -> Result<Value, Error> {
let Value::Object(fields) = metadata else {
return Err(Error::InvalidRequest(format!(
"metadata must be an object, got {metadata}"
)));
};
match fields.get("user_id") {
None | Some(Value::Null) => Ok(json!({})),
Some(Value::String(user_id)) => Ok(json!({"user_id": user_id})),
Some(other) => Err(Error::InvalidRequest(format!(
"metadata.user_id must be a string, got {other}"
))),
}
}
fn with_reasoning_auto_summary(thinking: Option<Value>, enabled: bool) -> Option<Value> {
let Some(Value::Object(thinking)) = thinking else {
return thinking;
};
if !enabled || thinking.get("type").and_then(Value::as_str) == Some("disabled") {
return Some(Value::Object(thinking));
}
Some(Value::Object(
thinking
.into_iter()
.filter(|(key, _)| key != "display")
.chain([("display".to_string(), json!("summarized"))])
.collect(),
))
}
#[cfg(test)]
mod tests {
use rstest::rstest;
use super::*;
fn messages(value: Value) -> Vec<AnthropicMessage> {
serde_json::from_value(value).unwrap()
}
fn request(body: Value) -> AnthropicMessagesRequest {
serde_json::from_value(body).unwrap()
}
#[rstest]
#[case::empty_text_next_to_a_tool_use(
json!([{"role": "assistant", "content": [
{"type": "text", "text": " "},
{"type": "tool_use", "id": "t", "name": "B", "input": {}}
]}]),
json!([{"role": "assistant", "content": [
{"type": "tool_use", "id": "t", "name": "B", "input": {}}
]}]),
)]
#[case::cross_provider_tool_ids(
json!([
{"role": "assistant", "content": [{"type": "tool_use", "id": "functions.Bash:0", "name": "Bash", "input": {}}]},
{"role": "user", "content": [{"type": "tool_result", "tool_use_id": "functions.Bash:0", "content": "ok"}]}
]),
json!([
{"role": "assistant", "content": [{"type": "tool_use", "id": "functions_Bash_0", "name": "Bash", "input": {}}]},
{"role": "user", "content": [{"type": "tool_result", "tool_use_id": "functions_Bash_0", "content": "ok"}]}
]),
)]
#[case::replayed_unencrypted_web_search_results(
json!([
{"role": "user", "content": "latest litellm version?"},
{"role": "assistant", "content": [
{"type": "server_tool_use", "id": "srvtoolu_1", "name": "web_search", "input": {"query": "latest litellm version"}},
{"type": "web_search_tool_result", "tool_use_id": "srvtoolu_1", "content": [{
"type": "web_search_result",
"url": "https://github.com/BerriAI/litellm/releases",
"title": "Releases",
"page_age": null,
"encrypted_content": "",
"snippet": "Latest release v1.95.0"
}]}
]},
{"role": "user", "content": "which version?"}
]),
json!([
{"role": "user", "content": "latest litellm version?"},
{"role": "assistant", "content": [{
"type": "text",
"text": "Web search results for 'latest litellm version':\n\nTitle: Releases\nURL: https://github.com/BerriAI/litellm/releases\nSnippet: Latest release v1.95.0"
}]},
{"role": "user", "content": "which version?"}
]),
)]
#[case::replayed_provider_specific_fields(
json!([
{"role": "assistant", "content": [{
"type": "tool_use", "id": "toolu_01", "name": "get_weather", "input": {"city": "Paris"},
"provider_specific_fields": {"signature": "sig_abc"}
}]},
{"role": "user", "content": [{"type": "tool_result", "tool_use_id": "toolu_01", "content": "Sunny"}]}
]),
json!([
{"role": "assistant", "content": [{"type": "tool_use", "id": "toolu_01", "name": "get_weather", "input": {"city": "Paris"}}]},
{"role": "user", "content": [{"type": "tool_result", "tool_use_id": "toolu_01", "content": "Sunny"}]}
]),
)]
#[case::ids_are_normalized_before_web_search_results_flatten(
json!([
{"role": "user", "content": "run it"},
{"role": "assistant", "content": [
{"type": "thinking", "thinking": "", "signature": "sig"},
{"type": "text", "text": ""},
{"type": "tool_use", "id": "functions.Bash:0", "name": "Bash", "input": {}, "provider_specific_fields": {"x": 1}},
{"type": "server_tool_use", "id": "srv.1", "name": "web_search", "input": {"query": "q"}, "provider_specific_fields": {"x": 2}},
{"type": "web_search_tool_result", "tool_use_id": "srv.1", "provider_specific_fields": {"x": 3}, "content": [
{"type": "web_search_result", "url": "u", "title": "", "encrypted_content": "", "provider_specific_fields": {"x": 4}}
]}
]},
{"role": "user", "content": [{"type": "tool_result", "tool_use_id": "functions.Bash:0", "content": "ok"}]},
{"role": "assistant", "content": [{"type": "text", "text": " "}]}
]),
json!([
{"role": "user", "content": "run it"},
{"role": "assistant", "content": [
{"type": "tool_use", "id": "functions_Bash_0", "name": "Bash", "input": {}},
{"type": "server_tool_use", "id": "srv_1", "name": "web_search", "input": {"query": "q"}},
{"type": "text", "text": "Web search results:\n\nURL: u"}
]},
{"role": "user", "content": [{"type": "tool_result", "tool_use_id": "functions_Bash_0", "content": "ok"}]}
]),
)]
fn sanitize_anthropic_messages_cleans_replayed_history(
#[case] history: Value,
#[case] expected: Value,
) {
assert_eq!(
serde_json::to_value(sanitize_anthropic_messages(messages(history))).unwrap(),
expected
);
}
#[rstest]
#[case::keeps_only_user_id(json!({"user_id": "u-1", "trace_id": "internal"}), Ok(json!({"user_id": "u-1"})))]
#[case::null_user_id(json!({"user_id": null, "trace_id": "internal"}), Ok(json!({})))]
#[case::no_user_id(json!({"trace_id": "internal"}), Ok(json!({})))]
#[case::empty(json!({}), Ok(json!({})))]
#[case::numeric_user_id(
json!({"user_id": 123}),
Err(Error::InvalidRequest("metadata.user_id must be a string, got 123".to_string())),
)]
#[case::boolean_user_id(
json!({"user_id": true}),
Err(Error::InvalidRequest("metadata.user_id must be a string, got true".to_string())),
)]
#[case::not_an_object(
json!(["u-1"]),
Err(Error::InvalidRequest(r#"metadata must be an object, got ["u-1"]"#.to_string())),
)]
fn validate_anthropic_api_metadata_passes_only_a_string_user_id(
#[case] metadata: Value,
#[case] expected: Result<Value, Error>,
) {
assert_eq!(validate_anthropic_api_metadata(&metadata), expected);
}
#[rstest]
#[case::adaptive(
Some(json!({"type": "adaptive", "budget_tokens": 5000})),
true,
Some(json!({"type": "adaptive", "budget_tokens": 5000, "display": "summarized"})),
)]
#[case::enabled(
Some(json!({"type": "enabled", "budget_tokens": 10000})),
true,
Some(json!({"type": "enabled", "budget_tokens": 10000, "display": "summarized"})),
)]
#[case::no_type(Some(json!({})), true, Some(json!({"display": "summarized"})))]
#[case::display_omitted_is_overridden(
Some(json!({"type": "enabled", "budget_tokens": 10000, "display": "omitted"})),
true,
Some(json!({"type": "enabled", "budget_tokens": 10000, "display": "summarized"})),
)]
#[case::display_summarized_is_kept(
Some(json!({"type": "enabled", "display": "summarized"})),
true,
Some(json!({"type": "enabled", "display": "summarized"})),
)]
#[case::disabled_thinking(Some(json!({"type": "disabled"})), true, Some(json!({"type": "disabled"})))]
#[case::flag_off(
Some(json!({"type": "enabled", "budget_tokens": 10000})),
false,
Some(json!({"type": "enabled", "budget_tokens": 10000})),
)]
#[case::flag_off_keeps_callers_display(
Some(json!({"type": "enabled", "display": "omitted"})),
false,
Some(json!({"type": "enabled", "display": "omitted"})),
)]
#[case::no_thinking(None, true, None)]
#[case::non_object_thinking(Some(json!("enabled")), true, Some(json!("enabled")))]
fn reasoning_auto_summary_marks_active_thinking_as_summarized(
#[case] thinking: Option<Value>,
#[case] enabled: bool,
#[case] expected: Option<Value>,
) {
assert_eq!(with_reasoning_auto_summary(thinking, enabled), expected);
}
#[test]
fn shaping_cleans_messages_metadata_and_thinking() {
let sanitized = shape_anthropic_messages_request(
request(json!({
"model": "m",
"messages": [{"role": "assistant", "content": [
{"type": "text", "text": ""},
{"type": "tool_use", "id": "functions.Bash:0", "name": "Bash", "input": {}}
]}],
"metadata": {"user_id": "u", "trace_id": "t"},
"thinking": {"type": "enabled", "budget_tokens": 1024},
"safeguards": [{"type": "dangerous_tool_use"}]
})),
true,
)
.unwrap();
assert_eq!(
serde_json::to_value(sanitized).unwrap(),
json!({
"model": "m",
"messages": [{"role": "assistant", "content": [
{"type": "tool_use", "id": "functions_Bash_0", "name": "Bash", "input": {}}
]}],
"metadata": {"user_id": "u"},
"thinking": {"type": "enabled", "budget_tokens": 1024, "display": "summarized"},
"safeguards": [{"type": "dangerous_tool_use"}]
})
);
}
}

View file

@ -0,0 +1,643 @@
use litellm_types::llms::anthropic_messages::anthropic_request::AnthropicMessagesRequest;
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,
};
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)))
}
fn with_oauth_bearer(headers: Headers, bearer: String) -> Headers {
let beta =
join_beta_values(existing_betas(&headers).chain([ANTHROPIC_OAUTH_BETA_HEADER.to_string()]));
without(headers, &[API_KEY_HEADER, AUTHORIZATION, BETA_HEADER])
.into_iter()
.chain([
(AUTHORIZATION.to_string(), bearer),
(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())
}
pub fn authenticate(
headers: Headers,
api_key: Option<&str>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> Result<Headers, litellm_auth::Error> {
if let Some(forwarded) = header_value(&headers, AUTHORIZATION)
&& forwarded
.strip_prefix("Bearer ")
.is_some_and(|token| token.starts_with(ANTHROPIC_OAUTH_TOKEN_PREFIX))
{
let bearer = forwarded.to_string();
return Ok(with_oauth_bearer(headers, bearer));
}
if let Some(key) = api_key.filter(|key| key.starts_with(ANTHROPIC_OAUTH_TOKEN_PREFIX)) {
return Ok(with_oauth_bearer(headers, format!("Bearer {key}")));
}
if header_value(&headers, API_KEY_HEADER).is_some()
|| header_value(&headers, AUTHORIZATION).is_some()
{
return Ok(headers);
}
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) => {
(AUTHORIZATION.to_string(), format!("Bearer {key}"))
}
Some(key) => (API_KEY_HEADER.to_string(), key),
None => match env_lookup(ANTHROPIC_AUTH_TOKEN_ENV).filter(|value| !value.trim().is_empty())
{
Some(token) => (AUTHORIZATION.to_string(), format!("Bearer {token}")),
None => {
return Err(litellm_auth::Error::MissingApiKey {
provider: "Anthropic",
environment_variable: ANTHROPIC_API_KEY_ENV,
});
}
},
};
Ok(headers.into_iter().chain([auth]).collect())
}
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.output_format.is_some()
|| request
.output_config
.as_ref()
.and_then(|config| config.get("format"))
.is_some_and(|format| !format.is_null())
}
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.tools.as_deref();
[
requires_native_compaction_beta(request.compaction.as_ref(), &request.messages)
.then_some(beta::COMPACT_2026_09_04),
uses_structured_output(request).then_some(beta::STRUCTURED_OUTPUT),
(request.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.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::*;
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())
};
authenticate(headers(forwarded), api_key, &lookup)
}
#[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([
("authorization", expected_bearer),
("anthropic-beta", ANTHROPIC_OAUTH_BETA_HEADER),
BROWSER_ACCESS,
])
.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(&[
("authorization", OAUTH_BEARER),
(
"anthropic-beta",
&betas(&[ANTHROPIC_OAUTH_BETA_HEADER, "web-search-2025-03-05"])
),
BROWSER_ACCESS,
])
);
}
#[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(&[
("authorization", OAUTH_BEARER),
BROWSER_ACCESS,
(
"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,2 +1,5 @@
pub mod handler;
pub mod headers;
pub mod streaming_iterator;
pub mod thinking;
pub mod transformation;

View file

@ -1,9 +1,28 @@
use crate::base_llm::{
anthropic_messages::transformation::BaseAnthropicMessagesConfig, chat::transformation::Error,
use litellm_core_utils::settings::{Lookup, ProcessEnvironment};
use litellm_types::llms::anthropic_messages::anthropic_request::AnthropicMessagesRequest;
use serde_json::{Map, Value, json};
use super::{
headers::{authenticate, with_feature_betas},
thinking::{ThinkingBudgets, ThinkingContext, translate_thinking},
};
use crate::{
anthropic::common_utils::{
AnthropicModelCapabilities, has_advisor_tool, strip_advisor_blocks,
strip_encrypted_reasoning_blocks,
},
base_llm::{
anthropic_messages::transformation::{
BaseAnthropicMessagesConfig, Headers, MessagesTransformContext,
},
chat::transformation::Error,
},
};
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";
@ -11,6 +30,26 @@ pub struct AnthropicMessagesConfig;
pub const ANTHROPIC_MESSAGES_CONFIG: AnthropicMessagesConfig = AnthropicMessagesConfig;
impl MessagesTransformContext {
pub fn new(capabilities: AnthropicModelCapabilities, drop_params: bool) -> Self {
Self::with_lookup(capabilities, drop_params, &ProcessEnvironment)
}
pub fn with_lookup(
capabilities: AnthropicModelCapabilities,
drop_params: bool,
env: &impl Lookup,
) -> Self {
Self {
thinking: ThinkingContext {
capabilities,
budgets: ThinkingBudgets::from_lookup(env),
},
drop_params,
}
}
}
impl BaseAnthropicMessagesConfig for AnthropicMessagesConfig {
fn get_complete_url(
&self,
@ -21,6 +60,35 @@ impl BaseAnthropicMessagesConfig for AnthropicMessagesConfig {
Ok(complete_anthropic_url(api_base, env_lookup))
}
fn transform_anthropic_messages_request(
&self,
request: AnthropicMessagesRequest,
context: &MessagesTransformContext,
) -> Result<AnthropicMessagesRequest, Error> {
if request.max_tokens.is_none() {
return Err(Error::InvalidRequest(
"max_tokens is required for Anthropic /v1/messages API".to_string(),
));
}
let request = drop_unsupported_params(request, context)?;
let request = translate_thinking(request, &context.thinking)?;
let context_management = request
.context_management
.as_ref()
.and_then(map_openai_context_management_to_anthropic)
.or_else(|| request.context_management.clone());
let messages = if has_advisor_tool(request.tools.as_deref()) {
request.messages
} else {
strip_advisor_blocks(request.messages)
};
Ok(AnthropicMessagesRequest {
messages: strip_encrypted_reasoning_blocks(messages),
context_management,
..request
})
}
fn resolve_api_key(
&self,
api_key: Option<&str>,
@ -28,6 +96,113 @@ impl BaseAnthropicMessagesConfig for AnthropicMessagesConfig {
) -> Result<String, Error> {
resolve_anthropic_api_key(api_key, env_lookup).map_err(Error::from)
}
fn secret_names(&self) -> &'static [&'static str] {
&[
ANTHROPIC_API_KEY_ENV,
ANTHROPIC_AUTH_TOKEN_ENV,
ANTHROPIC_API_BASE_ENV,
ANTHROPIC_BASE_URL_ENV,
]
}
fn authenticate(
&self,
headers: Headers,
api_key: Option<&str>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> Result<Headers, Error> {
authenticate(headers, api_key, env_lookup).map_err(Error::from)
}
fn request_headers(&self, headers: Headers, request: &AnthropicMessagesRequest) -> Headers {
with_feature_betas(headers, request)
}
}
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`."
))
}
fn drop_unsupported_params(
request: AnthropicMessagesRequest,
context: &MessagesTransformContext,
) -> Result<AnthropicMessagesRequest, Error> {
let capabilities = &context.thinking.capabilities;
let model = request.model.clone();
let reject = |param: &str, value: String, hint: &str| -> Result<(), Error> {
if context.drop_params {
return Ok(());
}
Err(unsupported_param(&model, param, &value, hint))
};
let speed = match request.speed.as_deref() {
Some(speed) if !capabilities.supports_speed => {
reject("speed", format!("'{speed}'"), "")?;
None
}
_ => request.speed.clone(),
};
if capabilities.supports_sampling_params {
return Ok(AnthropicMessagesRequest { speed, ..request });
}
let temperature = match request.temperature {
Some(temperature) if temperature != 1.0 => {
reject(
"temperature",
json!(temperature).to_string(),
"Only temperature=1 is supported. ",
)?;
None
}
temperature => temperature,
};
if let Some(top_p) = request.top_p {
reject("top_p", json!(top_p).to_string(), "")?;
}
if let Some(top_k) = request.top_k {
reject("top_k", json!(top_k).to_string(), "")?;
}
Ok(AnthropicMessagesRequest {
speed,
temperature,
top_p: None,
top_k: None,
..request
})
}
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,
}
}
pub fn non_empty(value: Option<&str>) -> Option<&str> {
@ -64,70 +239,619 @@ 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_lookup(ANTHROPIC_API_BASE_ENV).filter(|value| !value.trim().is_empty()))
.or_else(|| env(ANTHROPIC_API_BASE_ENV))
.or_else(|| env(ANTHROPIC_BASE_URL_ENV))
.unwrap_or_else(|| DEFAULT_ANTHROPIC_API_BASE.to_string())
}
#[cfg(test)]
mod tests {
use std::process::Command;
use rstest::{fixture, rstest};
use super::*;
use crate::anthropic::common_utils::{ENCRYPTED_REASONING_SIGNATURE_PREFIX, beta};
#[test]
fn url_defaults_to_public_anthropic_endpoint() {
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 LOW_BUDGET_ENV: &str = "DEFAULT_REASONING_EFFORT_LOW_THINKING_BUDGET";
const PROCESS_ENV_PROBE: &str = "LITELLM_MESSAGES_TRANSFORM_CONTEXT_PROBE";
fn merged(base: Value, fields: Value) -> Value {
Value::Object(
base.as_object()
.unwrap()
.clone()
.into_iter()
.chain(fields.as_object().unwrap().clone())
.collect(),
)
}
fn body(fields: Value) -> Value {
merged(
json!({
"model": "claude",
"max_tokens": 1024,
"messages": [{"role": "user", "content": "Hello"}]
}),
fields,
)
}
fn request(fields: Value) -> AnthropicMessagesRequest {
serde_json::from_value(body(fields)).unwrap()
}
fn no_env(_: &str) -> Option<String> {
None
}
fn env(vars: Env) -> impl Fn(&str) -> Option<String> {
move |name| {
vars.iter()
.find(|(key, _)| *key == name)
.map(|(_, value)| value.to_string())
}
}
fn headers(pairs: &[(&str, &str)]) -> Headers {
pairs
.iter()
.map(|(name, value)| (name.to_string(), value.to_string()))
.collect()
}
fn transform(
fields: Value,
capabilities: AnthropicModelCapabilities,
drop_params: bool,
) -> Result<Value, Error> {
ANTHROPIC_MESSAGES_CONFIG
.transform_anthropic_messages_request(
request(fields),
&MessagesTransformContext::with_lookup(capabilities, drop_params, &no_env),
)
.map(|transformed| serde_json::to_value(transformed).unwrap())
}
fn invalid(message: &str) -> Result<Value, Error> {
Err(Error::InvalidRequest(message.to_string()))
}
fn advisor_history() -> Value {
json!([
{"role": "user", "content": "Build a worker pool."},
{"role": "assistant", "content": [
{"type": "text", "text": "Let me consult the advisor."},
{"type": "server_tool_use", "id": "srvtoolu_abc123", "name": "advisor", "input": {}},
{"type": "advisor_tool_result", "tool_use_id": "srvtoolu_abc123", "content": {"type": "advisor_result", "text": "Use channels."}},
{"type": "text", "text": "Here is the implementation."}
]}
])
}
#[fixture]
fn unmapped() -> AnthropicModelCapabilities {
AnthropicModelCapabilities::default()
}
#[fixture]
fn sampling_removed() -> AnthropicModelCapabilities {
AnthropicModelCapabilities {
supports_sampling_params: false,
..Default::default()
}
}
#[fixture]
fn fast_mode() -> AnthropicModelCapabilities {
AnthropicModelCapabilities {
supports_speed: true,
..Default::default()
}
}
#[rstest]
#[case::alone(json!({"max_tokens": null}))]
#[case::ahead_of_the_param_gate(json!({"max_tokens": null, "speed": "fast"}))]
fn missing_max_tokens_is_rejected(#[case] fields: Value, unmapped: AnthropicModelCapabilities) {
assert_eq!(
complete_anthropic_url(None, &|_| None),
"https://api.anthropic.com/v1/messages"
transform(fields, unmapped, false),
invalid("max_tokens is required for Anthropic /v1/messages API")
);
}
#[rstest]
#[case::sampling_params_on_a_sampling_model(
unmapped(),
false,
json!({"temperature": 0.3, "top_p": 0.9, "top_k": 40})
)]
#[case::sampling_params_on_a_sampling_model_under_drop_params(
unmapped(),
true,
json!({"temperature": 0.3, "top_p": 0.9, "top_k": 40})
)]
#[case::unit_temperature_on_a_sampling_removed_model(
sampling_removed(),
false,
json!({"temperature": 1.0})
)]
#[case::unit_temperature_on_a_sampling_removed_model_under_drop_params(
sampling_removed(),
true,
json!({"temperature": 1.0})
)]
#[case::speed_on_a_fast_mode_model(fast_mode(), false, json!({"speed": "fast"}))]
#[case::speed_on_a_fast_mode_model_under_drop_params(fast_mode(), true, json!({"speed": "fast"}))]
#[case::native_context_management_edits(unmapped(), false, json!({"context_management": {"edits": [{
"type": "clear_tool_uses_20250919",
"trigger": {"type": "input_tokens", "value": 30000},
"keep": {"type": "tool_uses", "value": 3},
"clear_at_least": {"type": "input_tokens", "value": 5000},
"exclude_tools": ["web_search"],
"clear_tool_inputs": false
}]}}))]
#[case::first_party_billing_header_system_block(unmapped(), false, json!({"system": [
{"type": "text", "text": "x-anthropic-billing-header: cc_version=1"},
{"type": "text", "text": "real system prompt"}
]}))]
#[case::anthropic_signed_reasoning_history(unmapped(), false, json!({"messages": [
{"role": "user", "content": "Solve it."},
{"role": "assistant", "content": [
{"type": "thinking", "thinking": "plan", "signature": "EqQBCkYIAxgCIkA_anthropic_signed"},
{"type": "redacted_thinking", "data": "EmwKAhgBEgy_anthropic_minted"},
{"type": "text", "text": "The answer."}
]}
]}))]
#[case::advisor_history_alongside_the_advisor_tool(unmapped(), false, json!({
"messages": advisor_history(),
"tools": [{"type": "advisor_20260301", "name": "advisor"}]
}))]
fn request_is_forwarded_unchanged(
#[case] capabilities: AnthropicModelCapabilities,
#[case] drop_params: bool,
#[case] fields: Value,
) {
assert_eq!(
transform(fields.clone(), capabilities, drop_params),
Ok(body(fields))
);
}
#[rstest]
#[case::temperature(sampling_removed(), json!({"temperature": 0.3}), json!({}))]
#[case::top_p(sampling_removed(), json!({"top_p": 0.9}), json!({}))]
#[case::top_k(sampling_removed(), json!({"top_k": 40}), json!({}))]
#[case::every_sampling_param_keeping_the_rest(
sampling_removed(),
json!({"temperature": 0.3, "top_p": 0.9, "top_k": 40, "stream": true}),
json!({"stream": true})
)]
#[case::speed_on_a_sampling_model(
unmapped(),
json!({"speed": "fast", "temperature": 0.5}),
json!({"temperature": 0.5})
)]
#[case::speed_on_a_sampling_removed_model(
sampling_removed(),
json!({"speed": "fast", "temperature": 1.0}),
json!({"temperature": 1.0})
)]
fn removed_params_are_dropped_under_drop_params(
#[case] capabilities: AnthropicModelCapabilities,
#[case] fields: Value,
#[case] expected: Value,
) {
assert_eq!(transform(fields, capabilities, true), Ok(body(expected)));
}
#[rstest]
#[case::temperature(
sampling_removed(),
json!({"temperature": 0.3}),
"claude does not support temperature=0.3. Only temperature=1 is supported. To drop unsupported params, set `litellm.drop_params = True`."
)]
#[case::temperature_just_below_one(
sampling_removed(),
json!({"temperature": 0.99}),
"claude does not support temperature=0.99. Only temperature=1 is supported. To drop unsupported params, set `litellm.drop_params = True`."
)]
#[case::whole_number_temperature_keeps_its_decimal(
sampling_removed(),
json!({"temperature": 2.0}),
"claude does not support temperature=2.0. Only temperature=1 is supported. To drop unsupported params, set `litellm.drop_params = True`."
)]
#[case::top_p(
sampling_removed(),
json!({"top_p": 0.9}),
"claude does not support top_p=0.9. To drop unsupported params, set `litellm.drop_params = True`."
)]
#[case::top_k(
sampling_removed(),
json!({"top_k": 5}),
"claude does not support top_k=5. To drop unsupported params, set `litellm.drop_params = True`."
)]
#[case::top_k_next_to_unit_temperature(
sampling_removed(),
json!({"temperature": 1.0, "top_k": 5}),
"claude does not support top_k=5. To drop unsupported params, set `litellm.drop_params = True`."
)]
#[case::temperature_ahead_of_top_k(
sampling_removed(),
json!({"temperature": 0.5, "top_k": 5}),
"claude does not support temperature=0.5. Only temperature=1 is supported. To drop unsupported params, set `litellm.drop_params = True`."
)]
#[case::top_p_ahead_of_top_k(
sampling_removed(),
json!({"top_p": 0.9, "top_k": 5}),
"claude does not support top_p=0.9. To drop unsupported params, set `litellm.drop_params = True`."
)]
#[case::speed(
unmapped(),
json!({"speed": "fast"}),
"claude does not support speed='fast'. To drop unsupported params, set `litellm.drop_params = True`."
)]
#[case::speed_ahead_of_sampling_params(
sampling_removed(),
json!({"speed": "fast", "temperature": 0.5}),
"claude does not support speed='fast'. To drop unsupported params, set `litellm.drop_params = True`."
)]
fn removed_params_are_rejected_without_drop_params(
#[case] capabilities: AnthropicModelCapabilities,
#[case] fields: Value,
#[case] message: &str,
) {
assert_eq!(transform(fields, capabilities, false), invalid(message));
}
#[rstest]
#[case::compaction_threshold(
json!([{"type": "compaction", "compact_threshold": 200000}]),
Some(json!({"edits": [{"type": "compact_20260112", "trigger": {"type": "input_tokens", "value": 200000}}]}))
)]
#[case::other_keys_pass_through(
json!([{"type": "compaction", "compact_threshold": 150000, "instructions": "Focus on preserving code snippets"}]),
Some(json!({"edits": [{
"type": "compact_20260112",
"trigger": {"type": "input_tokens", "value": 150000},
"instructions": "Focus on preserving code snippets"
}]}))
)]
#[case::float_threshold_is_truncated(
json!([{"type": "compaction", "compact_threshold": 150000.9}]),
Some(json!({"edits": [{"type": "compact_20260112", "trigger": {"type": "input_tokens", "value": 150000}}]}))
)]
#[case::compaction_without_threshold(
json!([{"type": "compaction"}]),
Some(json!({"edits": [{"type": "compact_20260112"}]}))
)]
#[case::non_numeric_threshold_is_dropped(
json!([{"type": "compaction", "compact_threshold": "150000"}]),
Some(json!({"edits": [{"type": "compact_20260112"}]}))
)]
#[case::non_object_entries_are_skipped(
json!([42, "compaction", null, [], {"type": "compaction", "compact_threshold": 1000}]),
Some(json!({"edits": [{"type": "compact_20260112", "trigger": {"type": "input_tokens", "value": 1000}}]}))
)]
#[case::only_compaction_entries_are_mapped_in_order(
json!([
{"type": "compaction", "compact_threshold": 1000},
{"type": "other", "compact_threshold": 5},
{"type": "compaction", "instructions": "second"}
]),
Some(json!({"edits": [
{"type": "compact_20260112", "trigger": {"type": "input_tokens", "value": 1000}},
{"type": "compact_20260112", "instructions": "second"}
]}))
)]
#[case::list_without_compaction(json!([{"type": "other"}]), None)]
#[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}}]}))
)]
#[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>,
) {
assert_eq!(
map_openai_context_management_to_anthropic(&context_management),
expected
);
}
#[rstest]
#[case::openai_list_is_mapped(
json!([{"type": "compaction", "compact_threshold": 200000}]),
json!({"edits": [{"type": "compact_20260112", "trigger": {"type": "input_tokens", "value": 200000}}]})
)]
#[case::unmappable_list_is_kept(json!([{"type": "other"}]), json!([{"type": "other"}]))]
#[case::unmappable_object_is_kept(json!({"type": "other"}), json!({"type": "other"}))]
fn context_management_reaches_the_wire(
#[case] context_management: Value,
#[case] expected: Value,
unmapped: AnthropicModelCapabilities,
) {
assert_eq!(
transform(
json!({"context_management": context_management}),
unmapped,
false
),
Ok(body(json!({"context_management": expected})))
);
}
#[rstest]
#[case::without_tools(json!({}))]
#[case::with_only_other_tools(json!({"tools": [{"name": "get_weather", "input_schema": {"type": "object"}}]}))]
fn advisor_history_is_stripped_without_the_advisor_tool(
#[case] tools: Value,
unmapped: AnthropicModelCapabilities,
) {
let stripped = json!([
{"role": "user", "content": "Build a worker pool."},
{"role": "assistant", "content": [
{"type": "text", "text": "Let me consult the advisor."},
{"type": "text", "text": "Here is the implementation."}
]}
]);
assert_eq!(
transform(
merged(tools.clone(), json!({"messages": advisor_history()})),
unmapped,
false
),
Ok(body(merged(tools, json!({"messages": stripped}))))
);
}
#[rstest]
fn bridge_minted_reasoning_is_stripped_from_the_wire(unmapped: AnthropicModelCapabilities) {
let messages = json!([
{"role": "user", "content": "Solve it."},
{"role": "assistant", "content": [
{"type": "thinking", "thinking": "plan", "signature": format!("{ENCRYPTED_REASONING_SIGNATURE_PREFIX}gAAAA_1")},
{"type": "redacted_thinking", "data": format!("{ENCRYPTED_REASONING_SIGNATURE_PREFIX}gAAAA_2")},
{"type": "text", "text": "The answer."}
]},
{"role": "user", "content": "And the next one?"}
]);
assert_eq!(
transform(json!({"messages": messages}), unmapped, false),
Ok(body(json!({"messages": [
{"role": "user", "content": "Solve it."},
{"role": "assistant", "content": [{"type": "text", "text": "The answer."}]},
{"role": "user", "content": "And the next one?"}
]})))
);
}
#[test]
fn url_appends_messages_suffix_to_custom_base() {
fn thinking_is_translated_with_the_context_budgets() {
let context = MessagesTransformContext::with_lookup(
AnthropicModelCapabilities {
supports_reasoning: true,
..Default::default()
},
false,
&env(&[(LOW_BUDGET_ENV, "2000")]),
);
let transformed = ANTHROPIC_MESSAGES_CONFIG
.transform_anthropic_messages_request(
request(json!({"max_tokens": 4096, "reasoning_effort": "low"})),
&context,
)
.map(|transformed| serde_json::to_value(transformed).unwrap());
assert_eq!(
complete_anthropic_url(Some("https://proxy.internal"), &|_| None),
"https://proxy.internal/v1/messages"
transformed,
Ok(body(json!({
"max_tokens": 4096,
"thinking": {"type": "enabled", "budget_tokens": 2000}
})))
);
}
#[test]
fn url_leaves_complete_messages_endpoint_untouched() {
fn new_reads_thinking_budgets_from_the_process_environment() {
if std::env::var_os(PROCESS_ENV_PROBE).is_some() {
assert_eq!(
MessagesTransformContext::new(sampling_removed(), true),
MessagesTransformContext {
thinking: ThinkingContext {
capabilities: sampling_removed(),
budgets: ThinkingBudgets {
low: 2000,
..ThinkingBudgets::default()
},
},
drop_params: true,
}
);
return;
}
let (_, test_path) = concat!(
module_path!(),
"::new_reads_thinking_budgets_from_the_process_environment"
)
.split_once("::")
.unwrap();
let other_tiers = ["MINIMAL", "MEDIUM", "HIGH", "XHIGH", "MAX"]
.map(|tier| format!("DEFAULT_REASONING_EFFORT_{tier}_THINKING_BUDGET"));
let output = other_tiers
.iter()
.fold(
Command::new(std::env::current_exe().unwrap()),
|mut command, name| {
command.env_remove(name);
command
},
)
.args([test_path, "--exact"])
.env(PROCESS_ENV_PROBE, "1")
.env(LOW_BUDGET_ENV, "2000")
.output()
.unwrap();
let stdout = String::from_utf8_lossy(&output.stdout);
assert!(
output.status.success() && stdout.contains("1 passed"),
"{stdout}{}",
String::from_utf8_lossy(&output.stderr)
);
}
#[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(
None,
&[(ANTHROPIC_BASE_URL_ENV, "https://custom.example.com")],
"https://custom.example.com/v1/messages"
)]
#[case::custom_base(Some("https://proxy.internal"), &[], "https://proxy.internal/v1/messages")]
#[case::trailing_slash(Some("https://proxy.internal/"), &[], "https://proxy.internal/v1/messages")]
#[case::complete_endpoint(
Some("https://proxy.internal/v1/messages"),
&[],
"https://proxy.internal/v1/messages"
)]
#[case::complete_endpoint_with_trailing_slash(
Some("https://proxy.internal/v1/messages/"),
&[],
"https://proxy.internal/v1/messages"
)]
fn complete_url_ends_in_the_messages_path(
#[case] api_base: Option<&str>,
#[case] vars: Env,
#[case] expected: &str,
) {
assert_eq!(
complete_anthropic_url(Some("https://proxy.internal/v1/messages"), &|_| None),
"https://proxy.internal/v1/messages"
ANTHROPIC_MESSAGES_CONFIG.get_complete_url(api_base, "claude", &env(vars)),
Ok(expected.to_string())
);
}
#[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)
);
}
#[test]
fn url_falls_back_to_env_base() {
let with_env = |key: &str| {
(key == ANTHROPIC_API_BASE_ENV).then(|| "https://env.anthropic".to_string())
};
fn config_reports_a_missing_key_as_an_auth_error() {
assert_eq!(
complete_anthropic_url(Some(" "), &with_env),
"https://env.anthropic/v1/messages"
ANTHROPIC_MESSAGES_CONFIG.resolve_api_key(None, &no_env),
Err(Error::Auth(litellm_auth::Error::MissingApiKey {
provider: "Anthropic",
environment_variable: ANTHROPIC_API_KEY_ENV,
}))
);
}
#[test]
fn api_key_prefers_param_then_env_then_errors() {
fn config_authenticates_with_the_anthropic_auth_token() {
assert_eq!(
resolve_anthropic_api_key(Some("sk-param"), &|_| None).unwrap(),
"sk-param"
ANTHROPIC_MESSAGES_CONFIG.authenticate(
vec![],
None,
&env(&[("ANTHROPIC_AUTH_TOKEN", "auth-token")])
),
Ok(headers(&[("authorization", "Bearer auth-token")]))
);
let with_env = |key: &str| (key == ANTHROPIC_API_KEY_ENV).then(|| "sk-env".to_string());
}
#[test]
fn config_requests_the_betas_the_request_features_need() {
assert_eq!(
resolve_anthropic_api_key(Some(" "), &with_env).unwrap(),
"sk-env"
);
assert_eq!(
resolve_anthropic_api_key(None, &|_| None)
.expect_err("missing key")
.to_string(),
"Missing Anthropic API Key - Set `api_key` or the ANTHROPIC_API_KEY environment variable"
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::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);
}
#[test]
fn auth_strategy_and_default_headers_match_anthropic() {
assert_eq!(
@ -142,4 +866,26 @@ mod tests {
]
);
}
#[test]
fn secret_names_cover_every_credential_and_base_lookup() {
let requested = std::cell::RefCell::new(Vec::<String>::new());
let record = |name: &str| -> Option<String> {
requested.borrow_mut().push(name.to_string());
None
};
let _ = ANTHROPIC_MESSAGES_CONFIG.authenticate(Vec::new(), None, &record);
let _ = ANTHROPIC_MESSAGES_CONFIG.get_complete_url(None, "claude", &record);
let requested = requested.into_inner();
assert!(!requested.is_empty());
let undeclared: Vec<&String> = requested
.iter()
.filter(|name| {
!ANTHROPIC_MESSAGES_CONFIG
.secret_names()
.contains(&name.as_str())
})
.collect();
assert_eq!(undeclared, Vec::<&String>::new());
}
}

View file

@ -1,5 +1,6 @@
pub mod batches;
pub mod chat;
pub mod common_utils;
pub mod count_tokens;
pub mod experimental_pass_through;

View file

@ -4,14 +4,15 @@ use litellm_types::llms::anthropic_messages::{
},
anthropic_response::AnthropicMessagesResponse,
};
use serde_json::{Map, Value};
use crate::{
anthropic::experimental_pass_through::messages::transformation::{
ANTHROPIC_MESSAGES_CONFIG, AnthropicMessagesConfig, non_empty,
},
base_llm::{
anthropic_messages::transformation::{BaseAnthropicMessagesConfig, MessagesAuthStrategy},
anthropic_messages::transformation::{
BaseAnthropicMessagesConfig, Headers, MessagesAuthStrategy, MessagesTransformContext,
},
chat::transformation::Error,
},
};
@ -21,7 +22,6 @@ 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 TEXT_BLOCK_TYPE: &str = "text";
pub struct AzureAnthropicMessagesConfig {
anthropic: AnthropicMessagesConfig,
@ -45,6 +45,7 @@ impl BaseAnthropicMessagesConfig for AzureAnthropicMessagesConfig {
fn transform_anthropic_messages_request(
&self,
request: AnthropicMessagesRequest,
context: &MessagesTransformContext,
) -> Result<AnthropicMessagesRequest, Error> {
let mut request = fold_system_role_messages(request);
if let Some(system) = request.system.as_mut() {
@ -54,7 +55,8 @@ impl BaseAnthropicMessagesConfig for AzureAnthropicMessagesConfig {
.messages
.iter_mut()
.for_each(strip_scope_from_message);
self.anthropic.transform_anthropic_messages_request(request)
self.anthropic
.transform_anthropic_messages_request(request, context)
}
fn transform_anthropic_messages_response(
@ -74,6 +76,10 @@ impl BaseAnthropicMessagesConfig for AzureAnthropicMessagesConfig {
resolve_azure_api_key(api_key, env_lookup)
}
fn secret_names(&self) -> &'static [&'static str] {
&[AZURE_API_KEY_ENV, AZURE_API_BASE_ENV]
}
fn auth_strategy(&self) -> MessagesAuthStrategy {
self.anthropic.auth_strategy()
}
@ -85,6 +91,10 @@ impl BaseAnthropicMessagesConfig for AzureAnthropicMessagesConfig {
fn default_headers(&self) -> &'static [(&'static str, &'static str)] {
self.anthropic.default_headers()
}
fn request_headers(&self, headers: Headers, request: &AnthropicMessagesRequest) -> Headers {
self.anthropic.request_headers(headers, request)
}
}
pub fn resolve_azure_api_key(
@ -143,17 +153,7 @@ fn strip_scope_from_message(message: &mut AnthropicMessage) {
}
fn text_content_block(text: String) -> ContentBlock {
let extra = Map::from_iter([
(
"type".to_string(),
Value::String(TEXT_BLOCK_TYPE.to_string()),
),
("text".to_string(), Value::String(text)),
]);
ContentBlock {
cache_control: None,
extra,
}
ContentBlock::text(text)
}
fn content_into_blocks(content: MessageContent) -> Vec<ContentBlock> {
@ -202,6 +202,7 @@ mod tests {
use serde_json::json;
use super::*;
use crate::anthropic::common_utils::AnthropicModelCapabilities;
fn request_from(value: serde_json::Value) -> AnthropicMessagesRequest {
serde_json::from_value(value).expect("valid request")
@ -346,7 +347,7 @@ mod tests {
let transformed = to_value(
AZURE_ANTHROPIC_MESSAGES_CONFIG
.transform_anthropic_messages_request(request)
.transform_anthropic_messages_request(request, &MessagesTransformContext::default())
.expect("request transforms"),
);
@ -373,10 +374,13 @@ mod tests {
"messages": [{"role": "user", "content": "hi"}]
}));
let once = AZURE_ANTHROPIC_MESSAGES_CONFIG
.transform_anthropic_messages_request(request)
.transform_anthropic_messages_request(request, &MessagesTransformContext::default())
.expect("request transforms");
let twice = AZURE_ANTHROPIC_MESSAGES_CONFIG
.transform_anthropic_messages_request(once.clone())
.transform_anthropic_messages_request(
once.clone(),
&MessagesTransformContext::default(),
)
.expect("request transforms");
assert_eq!(once, twice);
assert_eq!(to_value(once)["system"], json!("plain string system"));
@ -408,9 +412,21 @@ mod tests {
"inference_geo": "us",
"litellm_metadata": {"trace": "abc"}
});
let context = MessagesTransformContext::with_lookup(
AnthropicModelCapabilities {
supports_reasoning: true,
supports_adaptive_thinking: true,
supports_legacy_thinking: true,
supports_output_config: true,
supports_speed: true,
..Default::default()
},
false,
&|_: &str| None,
);
let transformed = to_value(
AZURE_ANTHROPIC_MESSAGES_CONFIG
.transform_anthropic_messages_request(request_from(body.clone()))
.transform_anthropic_messages_request(request_from(body.clone()), &context)
.expect("request transforms"),
);
assert_eq!(transformed, body);
@ -430,7 +446,7 @@ mod tests {
let transformed = to_value(
AZURE_ANTHROPIC_MESSAGES_CONFIG
.transform_anthropic_messages_request(request)
.transform_anthropic_messages_request(request, &MessagesTransformContext::default())
.expect("request transforms"),
);
@ -460,7 +476,7 @@ mod tests {
let transformed = to_value(
AZURE_ANTHROPIC_MESSAGES_CONFIG
.transform_anthropic_messages_request(request)
.transform_anthropic_messages_request(request, &MessagesTransformContext::default())
.expect("request transforms"),
);
@ -485,9 +501,21 @@ mod tests {
{"role": "assistant", "content": "hello"}
]
});
let context = MessagesTransformContext::with_lookup(
AnthropicModelCapabilities {
supports_reasoning: true,
supports_adaptive_thinking: true,
supports_legacy_thinking: true,
supports_output_config: true,
supports_speed: true,
..Default::default()
},
false,
&|_: &str| None,
);
let transformed = to_value(
AZURE_ANTHROPIC_MESSAGES_CONFIG
.transform_anthropic_messages_request(request_from(body.clone()))
.transform_anthropic_messages_request(request_from(body.clone()), &context)
.expect("request transforms"),
);
assert_eq!(transformed, body);
@ -500,6 +528,57 @@ mod tests {
assert!(err.is_data());
}
#[rstest::rstest]
#[case::compact_context_management_edit(
json!({"context_management": {"edits": [{"type": "compact_20260112"}]}}),
&[],
&[("x-api-key", "k"), ("anthropic-beta", "compact-2026-01-12")]
)]
#[case::forwarded_beta_merged_with_structured_output(
json!({"output_config": {"format": {"type": "json_schema"}}}),
&[("anthropic-beta", "web-search-2025-03-05")],
&[("x-api-key", "k"), ("anthropic-beta", "structured-outputs-2025-11-13,web-search-2025-03-05")]
)]
#[case::no_feature_needs_a_beta(json!({}), &[], &[("x-api-key", "k")])]
fn request_headers_carry_the_anthropic_feature_betas(
#[case] fields: serde_json::Value,
#[case] forwarded: &[(&str, &str)],
#[case] expected: &[(&str, &str)],
) {
let pairs = |pairs: &[(&str, &str)]| -> Vec<(String, String)> {
pairs
.iter()
.map(|(name, value)| (name.to_string(), value.to_string()))
.collect()
};
let serde_json::Value::Object(fields) = fields else {
panic!("case fields are an object")
};
let request = request_from(serde_json::Value::Object(
[
("model".to_string(), json!("claude-sonnet")),
("max_tokens".to_string(), json!(16)),
(
"messages".to_string(),
json!([{"role": "user", "content": "hi"}]),
),
]
.into_iter()
.chain(fields)
.collect(),
));
assert_eq!(
AZURE_ANTHROPIC_MESSAGES_CONFIG.request_headers(
pairs(&[("x-api-key", "k")])
.into_iter()
.chain(pairs(forwarded))
.collect(),
&request
),
pairs(expected)
);
}
#[test]
fn transform_response_passes_through() {
let response: AnthropicMessagesResponse = serde_json::from_value(json!({
@ -521,4 +600,26 @@ mod tests {
assert_eq!(value["stop_sequence"], json!(null));
assert_eq!(value["content"][0]["text"], json!("hello"));
}
#[test]
fn secret_names_cover_every_credential_and_base_lookup() {
let requested = std::cell::RefCell::new(Vec::<String>::new());
let record = |name: &str| -> Option<String> {
requested.borrow_mut().push(name.to_string());
None
};
let _ = AZURE_ANTHROPIC_MESSAGES_CONFIG.authenticate(Vec::new(), None, &record);
let _ = AZURE_ANTHROPIC_MESSAGES_CONFIG.get_complete_url(None, "claude", &record);
let requested = requested.into_inner();
assert!(!requested.is_empty());
let undeclared: Vec<&String> = requested
.iter()
.filter(|name| {
!AZURE_ANTHROPIC_MESSAGES_CONFIG
.secret_names()
.contains(&name.as_str())
})
.collect();
assert_eq!(undeclared, Vec::<&String>::new());
}
}

View file

@ -1,8 +1,14 @@
use litellm_http::request::{has_bearer_auth, has_header};
use litellm_types::llms::anthropic_messages::{
anthropic_request::AnthropicMessagesRequest, anthropic_response::AnthropicMessagesResponse,
};
use crate::base_llm::chat::transformation::Error;
use crate::{
anthropic::experimental_pass_through::messages::thinking::ThinkingContext,
base_llm::chat::transformation::Error,
};
pub type Headers = Vec<(String, String)>;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum MessagesAuthStrategy {
@ -19,6 +25,12 @@ impl MessagesAuthStrategy {
}
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub struct MessagesTransformContext {
pub thinking: ThinkingContext,
pub drop_params: bool,
}
pub trait BaseAnthropicMessagesConfig: Sync {
fn get_complete_url(
&self,
@ -30,6 +42,7 @@ pub trait BaseAnthropicMessagesConfig: Sync {
fn transform_anthropic_messages_request(
&self,
request: AnthropicMessagesRequest,
_context: &MessagesTransformContext,
) -> Result<AnthropicMessagesRequest, Error> {
Ok(request)
}
@ -48,6 +61,8 @@ pub trait BaseAnthropicMessagesConfig: Sync {
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> Result<String, Error>;
fn secret_names(&self) -> &'static [&'static str];
fn auth_strategy(&self) -> MessagesAuthStrategy {
MessagesAuthStrategy::Header("x-api-key")
}
@ -56,10 +71,225 @@ pub trait BaseAnthropicMessagesConfig: Sync {
false
}
fn authenticate(
&self,
headers: Headers,
api_key: Option<&str>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> Result<Headers, Error> {
let strategy = self.auth_strategy();
if has_header(&headers, strategy.header_name())
|| (self.accepts_bearer_auth() && has_bearer_auth(&headers))
{
return Ok(headers);
}
let api_key = self.resolve_api_key(api_key, env_lookup)?;
let auth_header = match strategy {
MessagesAuthStrategy::Bearer => {
("authorization".to_string(), format!("Bearer {api_key}"))
}
MessagesAuthStrategy::Header(name) => (name.to_string(), api_key),
};
Ok(headers.into_iter().chain([auth_header]).collect())
}
fn default_headers(&self) -> &'static [(&'static str, &'static str)] {
&[
("anthropic-version", "2023-06-01"),
("content-type", "application/json"),
]
}
fn request_headers(&self, headers: Headers, _request: &AnthropicMessagesRequest) -> Headers {
headers
}
}
#[cfg(test)]
mod tests {
use rstest::rstest;
use super::*;
const X_API_KEY: MessagesAuthStrategy = MessagesAuthStrategy::Header("x-api-key");
struct StubConfig {
strategy: MessagesAuthStrategy,
accepts_bearer: bool,
}
impl BaseAnthropicMessagesConfig for StubConfig {
fn secret_names(&self) -> &'static [&'static str] {
&[]
}
fn get_complete_url(
&self,
_api_base: Option<&str>,
_model: &str,
_env_lookup: &dyn Fn(&str) -> Option<String>,
) -> Result<String, Error> {
Ok(String::new())
}
fn resolve_api_key(
&self,
api_key: Option<&str>,
_env_lookup: &dyn Fn(&str) -> Option<String>,
) -> Result<String, Error> {
api_key
.map(str::to_string)
.ok_or(Error::MissingField("api_key"))
}
fn auth_strategy(&self) -> MessagesAuthStrategy {
self.strategy
}
fn accepts_bearer_auth(&self) -> bool {
self.accepts_bearer
}
}
struct DefaultsConfig;
impl BaseAnthropicMessagesConfig for DefaultsConfig {
fn secret_names(&self) -> &'static [&'static str] {
&[]
}
fn get_complete_url(
&self,
_api_base: Option<&str>,
_model: &str,
_env_lookup: &dyn Fn(&str) -> Option<String>,
) -> Result<String, Error> {
Ok(String::new())
}
fn resolve_api_key(
&self,
api_key: Option<&str>,
_env_lookup: &dyn Fn(&str) -> Option<String>,
) -> Result<String, Error> {
api_key
.map(str::to_string)
.ok_or(Error::MissingField("api_key"))
}
}
#[test]
fn default_config_adds_its_key_next_to_a_forwarded_bearer() {
assert_eq!(
DefaultsConfig.authenticate(
headers(&[("authorization", "Bearer forwarded")]),
Some("sk"),
&|_| None
),
Ok(headers(&[
("authorization", "Bearer forwarded"),
("x-api-key", "sk")
]))
);
}
#[test]
fn default_request_headers_are_the_given_headers() {
let request: AnthropicMessagesRequest = serde_json::from_value(serde_json::json!({
"model": "claude",
"max_tokens": 16,
"speed": "fast",
"messages": [{"role": "user", "content": "hi"}]
}))
.unwrap();
assert_eq!(
DefaultsConfig.request_headers(headers(&[("x-api-key", "sk")]), &request),
headers(&[("x-api-key", "sk")])
);
}
fn headers(pairs: &[(&str, &str)]) -> Headers {
pairs
.iter()
.map(|(name, value)| (name.to_string(), value.to_string()))
.collect()
}
#[rstest]
#[case::own_header_is_kept(
X_API_KEY,
false,
headers(&[("x-api-key", "forwarded")]),
None,
Ok(headers(&[("x-api-key", "forwarded")]))
)]
#[case::own_header_in_any_casing_is_kept(
X_API_KEY,
false,
headers(&[("X-Api-Key", "forwarded")]),
None,
Ok(headers(&[("X-Api-Key", "forwarded")]))
)]
#[case::accepted_bearer_is_kept(
X_API_KEY,
true,
headers(&[("authorization", "Bearer forwarded")]),
None,
Ok(headers(&[("authorization", "Bearer forwarded")]))
)]
#[case::bearer_the_provider_does_not_accept_gets_the_key_too(
X_API_KEY,
false,
headers(&[("authorization", "Bearer forwarded")]),
Some("sk"),
Ok(headers(&[("authorization", "Bearer forwarded"), ("x-api-key", "sk")]))
)]
#[case::blank_bearer_gets_the_key(
X_API_KEY,
true,
headers(&[("authorization", "Bearer ")]),
Some("sk"),
Ok(headers(&[("authorization", "Bearer "), ("x-api-key", "sk")]))
)]
#[case::key_goes_in_the_provider_header(
X_API_KEY,
false,
headers(&[("content-type", "application/json")]),
Some("sk"),
Ok(headers(&[("content-type", "application/json"), ("x-api-key", "sk")]))
)]
#[case::key_goes_in_a_bearer(
MessagesAuthStrategy::Bearer,
false,
headers(&[]),
Some("sk"),
Ok(headers(&[("authorization", "Bearer sk")]))
)]
#[case::bearer_strategy_keeps_a_forwarded_authorization(
MessagesAuthStrategy::Bearer,
false,
headers(&[("authorization", "Bearer forwarded")]),
None,
Ok(headers(&[("authorization", "Bearer forwarded")]))
)]
#[case::missing_key_is_an_error(
X_API_KEY,
false,
headers(&[]),
None,
Err(Error::MissingField("api_key"))
)]
fn default_authenticate_applies_the_key_unless_a_credential_is_forwarded(
#[case] strategy: MessagesAuthStrategy,
#[case] accepts_bearer: bool,
#[case] forwarded: Headers,
#[case] api_key: Option<&str>,
#[case] expected: Result<Headers, Error>,
) {
let config = StubConfig {
strategy,
accepts_bearer,
};
assert_eq!(config.authenticate(forwarded, api_key, &|_| None), expected);
}
}

View file

@ -245,6 +245,8 @@ pub struct ModelInfo {
#[serde(default, skip_serializing_if = "Option::is_none")]
pub cache_creation_input_token_cost_above_272k_tokens_priority: Option<f64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub cache_creation_input_token_cost_above_32k_tokens: Option<f64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub cache_creation_input_token_cost_batches: Option<f64>,
/// Flex service-tier rate for the same-named base field.
#[serde(default, skip_serializing_if = "Option::is_none")]
@ -283,6 +285,8 @@ pub struct ModelInfo {
/// Priority service-tier rate for the same-named base field.
#[serde(default, skip_serializing_if = "Option::is_none")]
pub cache_read_input_token_cost_above_272k_tokens_priority: Option<f64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub cache_read_input_token_cost_above_32k_tokens: Option<f64>,
/// Rate applied once the prompt exceeds the token threshold in the field name.
#[serde(default, skip_serializing_if = "Option::is_none")]
pub cache_read_input_token_cost_above_512k_tokens: Option<f64>,
@ -377,6 +381,8 @@ pub struct ModelInfo {
/// Priority service-tier rate for the same-named base field.
#[serde(default, skip_serializing_if = "Option::is_none")]
pub input_cost_per_token_above_272k_tokens_priority: Option<f64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub input_cost_per_token_above_32k_tokens: Option<f64>,
/// Rate applied once the prompt exceeds the token threshold in the field name.
#[serde(default, skip_serializing_if = "Option::is_none")]
pub input_cost_per_token_above_512k_tokens: Option<f64>,
@ -498,6 +504,8 @@ pub struct ModelInfo {
/// Priority service-tier rate for the same-named base field.
#[serde(default, skip_serializing_if = "Option::is_none")]
pub output_cost_per_token_above_272k_tokens_priority: Option<f64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub output_cost_per_token_above_32k_tokens: Option<f64>,
/// Rate applied once the prompt exceeds the token threshold in the field name.
#[serde(default, skip_serializing_if = "Option::is_none")]
pub output_cost_per_token_above_512k_tokens: Option<f64>,

View file

@ -29,6 +29,7 @@ pub(super) enum CacheBinding {
#[pyclass(frozen, name = "_ResponseCacheRuntime")]
pub(crate) struct ResolvedCache {
binding: CacheBinding,
guard: Option<super::facade::FacadeGuard>,
pid: u32,
}
@ -36,10 +37,24 @@ impl ResolvedCache {
pub(super) fn new(binding: CacheBinding) -> Self {
Self {
binding,
guard: None,
pid: std::process::id(),
}
}
pub(super) fn with_guard(mut self, guard: super::facade::FacadeGuard) -> Self {
self.guard = Some(guard);
self
}
pub(super) fn native_service(&self) -> PyResult<Option<NativeResponseCache>> {
self.check_process()?;
Ok(match &self.binding {
CacheBinding::Native(service) => Some(service.clone()),
_ => None,
})
}
fn check_process(&self) -> PyResult<()> {
if matches!(self.binding, CacheBinding::Native(_)) && self.pid != std::process::id() {
return Err(PyRuntimeError::new_err(
@ -70,6 +85,43 @@ impl ResolvedCache {
#[pymethods]
impl ResolvedCache {
#[staticmethod]
pub(crate) fn from_selected(cache: &Bound<'_, PyAny>) -> PyResult<Self> {
let py = cache.py();
let binding = if cache.is_none() {
CacheBinding::Disabled
} else if let Ok(handle) = cache.extract::<PyRef<'_, super::handle::CacheTestHandle>>() {
CacheBinding::Native(handle.service()?)
} else if let Some(service) = super::facade::resolve(py, cache)? {
CacheBinding::Native(service)
} else if let Some(runtime) = cache
.getattr_opt("_native_cache")?
.filter(|value| !value.is_none())
{
let resolved = runtime
.getattr("native")?
.extract::<PyRef<'_, ResolvedCache>>()?;
match resolved.native_service()? {
Some(service) => {
if !resolved
.guard
.as_ref()
.is_some_and(|guard| guard.matches(py, cache).unwrap_or(false))
{
return Err(RustBridgeDeclined::new_err(
"native cache runtime no longer matches its facade",
));
}
CacheBinding::Native(service)
}
None => CacheBinding::PythonCallback(PythonCallback::new(cache.clone().unbind())),
}
} else {
CacheBinding::PythonCallback(PythonCallback::new(cache.clone().unbind()))
};
Ok(Self::new(binding))
}
#[staticmethod]
fn from_cache(cache: &Bound<'_, PyAny>) -> PyResult<Self> {
let config = match NativeCacheConfig::project(cache)? {
@ -80,7 +132,13 @@ impl ResolvedCache {
};
let backend = cache.getattr("cache")?;
let service = activate(cache.py(), &backend, config)?;
Ok(Self::new(CacheBinding::Native(service)))
let resolved = Self::new(CacheBinding::Native(service.clone()));
Ok(
match super::facade::FacadeGuard::capture(cache.py(), cache, &service) {
Ok(guard) => resolved.with_guard(guard),
Err(_) => resolved,
},
)
}
#[getter]
@ -323,6 +381,9 @@ impl ResolvedCache {
if let CacheBinding::PythonCallback(callback) = &self.binding {
callback.traverse(&visit)?;
}
if let Some(guard) = &self.guard {
guard.traverse(visit)?;
}
Ok(())
}
}

View file

@ -472,7 +472,7 @@ impl FacadeGuard {
})
}
fn matches(&self, py: Python<'_>, facade: &Bound<'_, PyAny>) -> PyResult<bool> {
pub(super) fn matches(&self, py: Python<'_>, facade: &Bound<'_, PyAny>) -> PyResult<bool> {
if !self.outer.matches(py, facade)? {
return Ok(false);
}

View file

@ -20,9 +20,7 @@ use pyo3::{
types::PyDict,
};
pub(crate) use self::{
binding::ResolvedCache, handle::CacheTestHandle, resolver::CacheTestResolver,
};
pub(crate) use self::{binding::ResolvedCache, handle::CacheTestHandle, resolver::CacheResolver};
fn cache_error(error: Error) -> PyErr {
match error {

View file

@ -1,19 +1,14 @@
use pyo3::{PyTraverseError, PyVisit, prelude::*};
use super::{
binding::{CacheBinding, ResolvedCache},
callback::PythonCallback,
facade,
handle::CacheTestHandle,
};
use super::binding::ResolvedCache;
#[pyclass(frozen, name = "_CacheTestResolver")]
pub(crate) struct CacheTestResolver {
#[pyclass(frozen, name = "_CacheResolver")]
pub(crate) struct CacheResolver {
namespace: Py<PyAny>,
}
#[pymethods]
impl CacheTestResolver {
impl CacheResolver {
#[new]
fn new(namespace: Py<PyAny>) -> Self {
Self { namespace }
@ -21,16 +16,7 @@ impl CacheTestResolver {
pub(crate) fn resolve(&self, py: Python<'_>) -> PyResult<ResolvedCache> {
let object = self.namespace.bind(py).getattr("cache")?;
let binding = if object.is_none() {
CacheBinding::Disabled
} else if let Ok(handle) = object.extract::<PyRef<'_, CacheTestHandle>>() {
CacheBinding::Native(handle.service()?)
} else if let Some(service) = facade::resolve(py, &object)? {
CacheBinding::Native(service)
} else {
CacheBinding::PythonCallback(PythonCallback::new(object.unbind()))
};
Ok(ResolvedCache::new(binding))
ResolvedCache::from_selected(&object)
}
fn __traverse__(&self, visit: PyVisit<'_>) -> Result<(), PyTraverseError> {

View file

@ -13,7 +13,7 @@ mod tokenizer;
#[pymodule(gil_used = true)]
mod _native {
use crate::cache::{CacheTestHandle, CacheTestResolver, ResolvedCache};
use crate::cache::{CacheResolver, CacheTestHandle, ResolvedCache};
#[cfg(feature = "panic-test")]
#[pymodule_export]
use crate::diagnostics::_panic_for_test;
@ -27,14 +27,16 @@ mod _native {
use crate::routes::audio_transcription::{atranscription, transcription};
#[pymodule_export]
use crate::routes::chat_completions::{
achat_completions, chat_completions, chat_completions_decline,
achat_completions, acompletion, chat_completions, chat_completions_decline, completion,
};
#[pymodule_export]
use crate::routes::embeddings::{aembedding, embedding};
#[pymodule_export]
use crate::routes::messages::{amessages, messages};
#[pymodule_export]
use crate::routes::ocr::{aocr, ocr};
#[pymodule_export]
use crate::routes::responses::ResponsesWebSocketConnection;
use crate::routes::responses::{ResponsesWebSocketConnection, aresponses, responses};
#[pymodule_export]
use crate::routes::token_counter::TokenCounter;
#[cfg(feature = "huggingface")]
@ -51,7 +53,8 @@ mod _native {
let py = module.py();
let dict = module.dict();
dict.set_item("_CacheTestHandle", py.get_type::<CacheTestHandle>())?;
dict.set_item("_CacheTestResolver", py.get_type::<CacheTestResolver>())?;
dict.set_item("_CacheResolver", py.get_type::<CacheResolver>())?;
dict.set_item("_CacheTestResolver", py.get_type::<CacheResolver>())?;
dict.set_item("_ResponseCacheRuntime", py.get_type::<ResolvedCache>())?;
dict.set_item(
"_SecretManagerRuntime",
@ -82,6 +85,8 @@ mod tests {
"ProcessReservedForForking",
"ocr",
"aocr",
"embedding",
"aembedding",
"transcription",
"atranscription",
"messages",
@ -89,6 +94,10 @@ mod tests {
"chat_completions_decline",
"chat_completions",
"achat_completions",
"completion",
"acompletion",
"responses",
"aresponses",
"ResponsesWebSocketConnection",
"NativeDiagnosticProcessor",
"TokenCounter",

View file

@ -1,3 +1,6 @@
use pyo3::types::{PyDict, PyTuple};
use crate::errors::RustBridgeDeclined;
use crate::logger::{run_async, run_sync};
use litellm_core::chat_completions::{
Error, chat_completions as run_chat_completions, chat_completions_decline_reason,
@ -123,9 +126,58 @@ pub(crate) fn achat_completions<'py>(
)
}
#[pyfunction]
#[pyo3(signature = (request, args, kwargs))]
pub(crate) fn completion(
request: Bound<'_, PyAny>,
args: Bound<'_, PyTuple>,
kwargs: Bound<'_, PyDict>,
) -> PyResult<Py<PyAny>> {
drop((request, args, kwargs));
Err(RustBridgeDeclined::new_err(
"native chat completions route is not implemented",
))
}
#[pyfunction]
#[pyo3(signature = (request, args, kwargs))]
pub(crate) fn acompletion(
request: Bound<'_, PyAny>,
args: Bound<'_, PyTuple>,
kwargs: Bound<'_, PyDict>,
) -> PyResult<Py<PyAny>> {
drop((request, args, kwargs));
Err(RustBridgeDeclined::new_err(
"native chat completions route is not implemented",
))
}
#[cfg(test)]
mod tests {
use pyo3::{prelude::*, types::PyList};
use pyo3::{
prelude::*,
types::{PyDict, PyList, PyTuple},
};
use crate::errors::RustBridgeDeclined;
#[test]
fn both_entrypoints_decline_before_provider_execution() {
Python::initialize();
Python::attach(|py| {
let request = PyDict::new(py);
let args = PyTuple::empty(py);
let kwargs = PyDict::new(py);
for entrypoint in [super::completion, super::acompletion] {
let error = entrypoint(request.clone().into_any(), args.clone(), kwargs.clone())
.expect_err(
"native chat completions must decline until a route machine exists",
);
assert!(error.is_instance_of::<RustBridgeDeclined>(py));
}
});
}
#[test]
fn chat_completions_decline_keeps_existing_reasons() {

View file

@ -0,0 +1,58 @@
use pyo3::{
prelude::*,
types::{PyDict, PyTuple},
};
use crate::errors::RustBridgeDeclined;
#[pyfunction]
#[pyo3(signature = (request, args, kwargs))]
pub(crate) fn embedding(
request: Bound<'_, PyAny>,
args: Bound<'_, PyTuple>,
kwargs: Bound<'_, PyDict>,
) -> PyResult<Py<PyAny>> {
drop((request, args, kwargs));
Err(RustBridgeDeclined::new_err(
"native embeddings route is not implemented",
))
}
#[pyfunction]
#[pyo3(signature = (request, args, kwargs))]
pub(crate) fn aembedding(
request: Bound<'_, PyAny>,
args: Bound<'_, PyTuple>,
kwargs: Bound<'_, PyDict>,
) -> PyResult<Py<PyAny>> {
drop((request, args, kwargs));
Err(RustBridgeDeclined::new_err(
"native embeddings route is not implemented",
))
}
#[cfg(test)]
mod tests {
use pyo3::{
prelude::*,
types::{PyDict, PyTuple},
};
use crate::errors::RustBridgeDeclined;
#[test]
fn both_entrypoints_decline_before_provider_execution() {
Python::initialize();
Python::attach(|py| {
let request = PyDict::new(py);
let args = PyTuple::empty(py);
let kwargs = PyDict::new(py);
for entrypoint in [super::embedding, super::aembedding] {
let error = entrypoint(request.clone().into_any(), args.clone(), kwargs.clone())
.expect_err("native embeddings must decline until a route machine exists");
assert!(error.is_instance_of::<RustBridgeDeclined>(py));
}
});
}
}

View file

@ -2,9 +2,11 @@ use bytes::Bytes;
use litellm_core::messages::{
Error,
route::{Messages, MessagesCall, MessagesOp, MessagesOpResult, MessagesOutput},
types::MessagesShaping,
};
use litellm_host_python::{InvokeError, RouteHost, from_py, lookup, to_py};
use litellm_http::transport::Error as TransportError;
use litellm_types::utils::ProviderSpecificHeaders;
use pyo3::{
exceptions::{PyException, PyValueError},
gc::{PyTraverseError, PyVisit},
@ -18,9 +20,10 @@ use crate::{
marshal::{optional_timeout, python_timeout_seconds},
};
/// The Anthropic Messages body fields a caller may pass besides `model` and `messages`,
/// as `AnthropicMessagesRequestOptionalParams` declares them.
const BODY_FIELDS: [&str; 20] = [
const ROUTE_HOST_MODULE: &str = "litellm.rust_bridge.messages.route_host";
const REQUEST_ERROR_MARKER: &str = "messages_request_error";
const BODY_FIELDS: [&str; 22] = [
"max_tokens",
"metadata",
"stop_sequences",
@ -35,14 +38,46 @@ const BODY_FIELDS: [&str; 20] = [
"top_p",
"mcp_servers",
"context_management",
"compaction",
"container",
"output_format",
"speed",
"output_config",
"cache_control",
"reasoning_effort",
"safeguards",
];
fn merge_headers(
forwarded: Option<Map<String, Value>>,
extra_headers: Option<Map<String, Value>>,
) -> Option<Map<String, Value>> {
let merged: Map<String, Value> = forwarded
.into_iter()
.flatten()
.chain(extra_headers.into_iter().flatten())
.collect();
(!merged.is_empty()).then_some(merged)
}
fn native_error(py: Python<'_>, error: Error) -> PyResult<PyErr> {
match error {
Error::Transport(TransportError::Http { status, body }) => {
let error = RustUpstreamError::new_err((status, body));
error
.value(py)
.setattr("headers", Vec::<(String, String)>::new())?;
Ok(error)
}
Error::InvalidRequest(message) => {
let error = PyValueError::new_err(message);
error.value(py).setattr(REQUEST_ERROR_MARKER, true)?;
Ok(error)
}
other => Ok(messages_error_to_pyerr(other)),
}
}
/// The Python side of the Messages route: projects the prepared arguments and builds the
/// public response, chunks and exceptions.
pub(super) struct MessagesRouteHost {
@ -84,19 +119,65 @@ impl MessagesRouteHost {
.map(|value| python_timeout_seconds(py, value.unbind()))
.transpose()?
.flatten();
let custom_llm_provider = string("custom_llm_provider")?;
let shaping = self.shaping(py, &model, custom_llm_provider.as_deref(), arguments)?;
Ok(MessagesCall {
model,
body,
api_key: string("api_key")?,
api_base: string("api_base")?,
custom_llm_provider: string("custom_llm_provider")?,
extra_headers: argument("extra_headers")?
.map(|value| from_py(&value))
.transpose()?,
extra_headers: self.merged_headers(py, arguments)?,
provider_specific_header: self.provider_specific_header(py, arguments)?,
custom_llm_provider,
timeout: optional_timeout(timeout),
shaping,
})
}
fn merged_headers(
&self,
py: Python<'_>,
arguments: &Bound<'_, PyDict>,
) -> PyResult<Option<Map<String, Value>>> {
let request = self.request.bind(py);
let mapping = |name: &str| -> PyResult<Option<Map<String, Value>>> {
lookup(arguments, request, name)?
.filter(|value| !value.is_none())
.map(|value| from_py(&value))
.transpose()
};
Ok(merge_headers(
mapping("headers")?,
mapping("extra_headers")?,
))
}
fn provider_specific_header(
&self,
py: Python<'_>,
arguments: &Bound<'_, PyDict>,
) -> PyResult<Option<ProviderSpecificHeaders>> {
lookup(arguments, self.request.bind(py), "provider_specific_header")?
.filter(|value| !value.is_none())
.map(|value| from_py(&value))
.transpose()
}
fn shaping(
&self,
py: Python<'_>,
model: &str,
custom_llm_provider: Option<&str>,
arguments: &Bound<'_, PyDict>,
) -> PyResult<MessagesShaping> {
let projected = py.import(ROUTE_HOST_MODULE)?.getattr("shaping")?.call1((
model,
custom_llm_provider,
arguments,
))?;
from_py(&projected)
}
fn provider(&self, py: Python<'_>) -> String {
self.request
.bind(py)
@ -112,7 +193,7 @@ impl MessagesRouteHost {
return error;
}
let mapped = py
.import("litellm.rust_bridge.messages.route_host")
.import(ROUTE_HOST_MODULE)
.and_then(|module| module.getattr("map_failure"))
.and_then(|map| map.call1((error.value(py), self.request.bind(py), self.provider(py))))
.and_then(|mapped| {
@ -148,7 +229,7 @@ impl RouteHost for MessagesRouteHost {
fn complete(&mut self, py: Python<'_>, response: MessagesOutput) -> PyResult<Py<PyAny>> {
match response {
MessagesOutput::Message(message) => py
.import("litellm.rust_bridge.messages.route_host")?
.import(ROUTE_HOST_MODULE)?
.getattr("response")?
.call1((to_py(py, message.as_ref())?,))
.map(Bound::unbind),
@ -161,17 +242,12 @@ impl RouteHost for MessagesRouteHost {
}
fn classify(&self, py: Python<'_>, error: Error) -> PyResult<PyErr> {
let native = match error {
Error::Transport(TransportError::Http { status, body }) => {
let error = RustUpstreamError::new_err((status, body));
error
.value(py)
.setattr("headers", Vec::<(String, String)>::new())?;
error
}
other => messages_error_to_pyerr(other),
};
Ok(self.map_failure(py, native))
if let Error::Secret(source) = &error
&& let Some(original) = crate::secrets::python_error(py, source.source_error())
{
return Ok(original);
}
Ok(self.map_failure(py, native_error(py, error)?))
}
fn host_error(error: &PyErr) -> Error {
@ -184,3 +260,62 @@ impl RouteHost for MessagesRouteHost {
visit.call(&self.request)
}
}
#[cfg(test)]
mod tests {
use rstest::rstest;
use serde_json::json;
use super::*;
fn map(value: Value) -> Map<String, Value> {
serde_json::from_value(value).unwrap()
}
#[rstest]
#[case::extra_over_forwarded(
Some(json!({"X-Priority": "forwarded", "X-Forwarded-Only": "keep"})),
Some(json!({"X-Priority": "extra", "X-Extra-Only": "also-keep"})),
Some(json!({"X-Priority": "extra", "X-Forwarded-Only": "keep", "X-Extra-Only": "also-keep"})),
)]
#[case::only_forwarded(Some(json!({"X-Forwarded": "yes"})), None, Some(json!({"X-Forwarded": "yes"})))]
#[case::only_extra_headers(
None,
Some(json!({"X-Custom-Header": "from-kwargs", "X-Auth-Token": "token123"})),
Some(json!({"X-Custom-Header": "from-kwargs", "X-Auth-Token": "token123"})),
)]
#[case::nothing(None, Some(json!({})), None)]
fn headers_merge_forwarded_then_extra(
#[case] forwarded: Option<Value>,
#[case] extra_headers: Option<Value>,
#[case] expected: Option<Value>,
) {
assert_eq!(
merge_headers(forwarded.map(map), extra_headers.map(map)),
expected.map(map)
);
}
#[rstest]
#[case::rejected_request(Error::InvalidRequest("does not support top_k=5".into()), true)]
#[case::unresolvable_provider(Error::InvalidProvider("openai".into()), false)]
#[case::upstream_failure(
Error::Transport(TransportError::Http { status: 400, body: "bad".into() }),
false,
)]
fn only_request_rejections_carry_the_request_error_marker(
#[case] error: Error,
#[case] marked: bool,
) {
Python::initialize();
Python::attach(|py| {
let native = native_error(py, error).unwrap();
let marker = native
.value(py)
.getattr_opt(REQUEST_ERROR_MARKER)
.unwrap()
.map(|value| value.extract::<bool>().unwrap());
assert_eq!(marker.unwrap_or(false), marked);
});
}
}

View file

@ -39,11 +39,12 @@ fn run_messages(
"the Rust Messages route does not serve this provider",
));
}
let secrets = crate::secrets::source(py)?;
run_legacy_call(
py,
SURFACE,
PublicCall::capture(&request, &args, &kwargs)?,
crate::logger::LoggedMachine::new(messages_machine()),
crate::logger::LoggedMachine::new(messages_machine(secrets)),
MessagesRouteHost::new(request.unbind()),
asynchronous,
)

View file

@ -1,5 +1,6 @@
pub(crate) mod audio_transcription;
pub(crate) mod chat_completions;
pub(crate) mod embeddings;
pub(crate) mod messages;
pub(crate) mod ocr;
pub(crate) mod responses;

View file

@ -1,12 +1,41 @@
use litellm_core::responses::websocket::ResponsesWebSocketConnection as RustResponsesWebSocketConnection;
use pyo3::prelude::*;
use pyo3::{
prelude::*,
types::{PyDict, PyTuple},
};
use serde_json::Value;
use crate::{
errors::responses_error_to_pyerr,
errors::{RustBridgeDeclined, responses_error_to_pyerr},
marshal::{marshal_headers, optional_timeout},
};
#[pyfunction]
#[pyo3(signature = (request, args, kwargs))]
pub(crate) fn responses(
request: Bound<'_, PyAny>,
args: Bound<'_, PyTuple>,
kwargs: Bound<'_, PyDict>,
) -> PyResult<Py<PyAny>> {
drop((request, args, kwargs));
Err(RustBridgeDeclined::new_err(
"native responses route is not implemented",
))
}
#[pyfunction]
#[pyo3(signature = (request, args, kwargs))]
pub(crate) fn aresponses(
request: Bound<'_, PyAny>,
args: Bound<'_, PyTuple>,
kwargs: Bound<'_, PyDict>,
) -> PyResult<Py<PyAny>> {
drop((request, args, kwargs));
Err(RustBridgeDeclined::new_err(
"native responses route is not implemented",
))
}
#[pyclass]
pub(crate) struct ResponsesWebSocketConnection {
inner: RustResponsesWebSocketConnection,
@ -63,7 +92,28 @@ mod tests {
use std::{ffi::CString, time::Duration};
use futures_util::{SinkExt, StreamExt};
use pyo3::{prelude::*, types::PyDict};
use pyo3::{
prelude::*,
types::{PyDict, PyTuple},
};
use crate::errors::RustBridgeDeclined;
#[test]
fn both_entrypoints_decline_before_provider_execution() {
Python::initialize();
Python::attach(|py| {
let request = PyDict::new(py);
let args = PyTuple::empty(py);
let kwargs = PyDict::new(py);
for entrypoint in [super::responses, super::aresponses] {
let error = entrypoint(request.clone().into_any(), args.clone(), kwargs.clone())
.expect_err("native responses must decline until a route machine exists");
assert!(error.is_instance_of::<RustBridgeDeclined>(py));
}
});
}
use tokio::net::TcpListener;
use tokio_tungstenite::{accept_async, tungstenite::Message};

View file

@ -8,3 +8,6 @@ repository.workspace = true
[dependencies]
serde.workspace = true
serde_json.workspace = true
[dev-dependencies]
rstest.workspace = true

View file

@ -17,12 +17,48 @@ pub enum MessageContent {
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
pub struct ContentBlock {
#[serde(rename = "type", default, skip_serializing_if = "Option::is_none")]
pub block_type: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub text: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub thinking: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub signature: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub data: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub id: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub name: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub tool_use_id: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub input: Option<Value>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub content: Option<Value>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub provider_specific_fields: Option<Value>,
#[serde(skip_serializing_if = "Option::is_none")]
pub cache_control: Option<CacheControl>,
#[serde(flatten)]
pub extra: Map<String, Value>,
}
impl ContentBlock {
pub fn text(text: impl Into<String>) -> Self {
Self {
block_type: Some("text".to_string()),
text: Some(text.into()),
..Self::default()
}
}
pub fn is_type(&self, block_type: &str) -> bool {
self.block_type.as_deref() == Some(block_type)
}
}
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
pub struct CacheControl {
#[serde(rename = "type", skip_serializing_if = "Option::is_none")]
@ -85,6 +121,126 @@ pub struct AnthropicMessagesRequest {
pub speed: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub inference_geo: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub reasoning_effort: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub compaction: Option<Value>,
#[serde(flatten)]
pub extra: Map<String, Value>,
}
impl AnthropicMessage {
pub fn blocks(&self) -> &[ContentBlock] {
match &self.content {
MessageContent::Blocks(blocks) => blocks,
MessageContent::Text(_) => &[],
}
}
pub fn with_blocks(self, blocks: Vec<ContentBlock>) -> Self {
Self {
content: MessageContent::Blocks(blocks),
..self
}
}
}
#[cfg(test)]
mod tests {
use rstest::rstest;
use serde_json::json;
use super::*;
fn round_trip<T: serde::de::DeserializeOwned + Serialize>(value: &Value) -> Value {
let parsed: T = serde_json::from_value(value.clone()).unwrap();
serde_json::to_value(parsed).unwrap()
}
#[rstest]
#[case::text(json!({"type": "text", "text": "hi"}))]
#[case::text_with_citations_and_cache_control(json!({
"type": "text",
"text": "hi",
"citations": [{"type": "char_location", "cited_text": "x"}],
"cache_control": {"type": "ephemeral", "ttl": "1h", "scope": "global", "future": 1}
}))]
#[case::image(json!({"type": "image", "source": {"type": "base64", "media_type": "image/png", "data": "AA=="}}))]
#[case::thinking(json!({"type": "thinking", "thinking": "hmm", "signature": "sig"}))]
#[case::redacted_thinking(json!({"type": "redacted_thinking", "data": "opaque"}))]
#[case::tool_use(json!({"type": "tool_use", "id": "toolu_1", "name": "f", "input": {"q": [1, null]}}))]
#[case::tool_result_with_text(json!({"type": "tool_result", "tool_use_id": "toolu_1", "content": "ok", "is_error": false}))]
#[case::tool_result_with_blocks(json!({"type": "tool_result", "tool_use_id": "toolu_1", "content": [{"type": "text", "text": "ok"}]}))]
#[case::web_search_result_with_nulls(json!({
"type": "web_search_tool_result",
"tool_use_id": "srvtoolu_1",
"content": [{"type": "web_search_result", "url": "u", "page_age": null, "encrypted_content": ""}]
}))]
#[case::provider_specific_fields(json!({"type": "tool_use", "id": "t", "name": "f", "input": {}, "provider_specific_fields": {"x": 1}}))]
#[case::untyped(json!({"unknown": {"nested": true}}))]
fn content_block_round_trips_unchanged(#[case] block: Value) {
assert_eq!(round_trip::<ContentBlock>(&block), block);
}
#[test]
fn text_constructor_serializes_as_a_text_block() {
assert_eq!(
serde_json::to_value(ContentBlock::text("hello")).unwrap(),
json!({"type": "text", "text": "hello"})
);
}
#[rstest]
#[case::same_type(json!({"type": "tool_use"}), "tool_use", true)]
#[case::other_type(json!({"type": "tool_result"}), "tool_use", false)]
#[case::prefix_of_type(json!({"type": "tool_use"}), "tool", false)]
#[case::no_type(json!({"text": "x"}), "text", false)]
fn is_type_matches_the_exact_block_type(
#[case] block: Value,
#[case] block_type: &str,
#[case] expected: bool,
) {
let block: ContentBlock = serde_json::from_value(block).unwrap();
assert_eq!(block.is_type(block_type), expected);
}
#[rstest]
#[case::string_content(json!({"role": "user", "content": "hi"}), vec![])]
#[case::block_content(
json!({"role": "user", "content": [{"type": "text", "text": "a"}, {"type": "text", "text": "b"}]}),
vec![ContentBlock::text("a"), ContentBlock::text("b")],
)]
fn message_blocks_list_only_block_content(
#[case] message: Value,
#[case] expected: Vec<ContentBlock>,
) {
let message: AnthropicMessage = serde_json::from_value(message).unwrap();
assert_eq!(message.blocks(), expected.as_slice());
}
#[rstest]
#[case::replaces_string_content(json!({"role": "assistant", "content": "old", "name": "kept"}))]
#[case::replaces_block_content(json!({"role": "assistant", "content": [{"type": "text", "text": "old"}], "name": "kept"}))]
fn with_blocks_replaces_content_and_keeps_the_rest(#[case] message: Value) {
let message: AnthropicMessage = serde_json::from_value(message).unwrap();
assert_eq!(
serde_json::to_value(message.with_blocks(vec![ContentBlock::text("new")])).unwrap(),
json!({"role": "assistant", "content": [{"type": "text", "text": "new"}], "name": "kept"})
);
}
#[rstest]
#[case::minimal(json!({"model": "m", "messages": [{"role": "user", "content": "hi"}]}))]
#[case::reasoning_effort_compaction_and_unknown_fields(json!({
"model": "m",
"messages": [{"role": "user", "content": [{"type": "text", "text": "hi"}]}],
"max_tokens": 8,
"reasoning_effort": "high",
"compaction": {"type": "auto"},
"safeguards": [{"type": "dangerous_tool_use", "classifier_context": {"v": 1}}],
"metadata": {"user_id": "u"}
}))]
fn request_round_trips_unchanged(#[case] request: Value) {
assert_eq!(round_trip::<AnthropicMessagesRequest>(&request), request);
}
}

View file

@ -9,8 +9,6 @@ pub struct AnthropicMessagesResponse {
pub role: String,
pub model: String,
pub content: Vec<Value>,
// Anthropic always includes stop_reason / stop_sequence, null until the turn
// ends; serialize them even when None so callers see the same shape as Python.
pub stop_reason: Option<String>,
pub stop_sequence: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
@ -20,3 +18,61 @@ pub struct AnthropicMessagesResponse {
#[serde(flatten)]
pub extra: Map<String, Value>,
}
#[cfg(test)]
mod tests {
use rstest::rstest;
use serde_json::json;
use super::*;
fn response(
stop_reason: Option<&str>,
stop_sequence: Option<&str>,
usage: Option<Value>,
container: Option<Value>,
) -> AnthropicMessagesResponse {
AnthropicMessagesResponse {
id: "msg_1".to_string(),
message_type: "message".to_string(),
role: "assistant".to_string(),
model: "claude".to_string(),
content: vec![],
stop_reason: stop_reason.map(str::to_string),
stop_sequence: stop_sequence.map(str::to_string),
usage,
container,
extra: Map::new(),
}
}
#[rstest]
#[case::turn_in_progress(None, None, json!(null), json!(null))]
#[case::ended_on_end_turn(Some("end_turn"), None, json!("end_turn"), json!(null))]
#[case::ended_on_stop_sequence(Some("stop_sequence"), Some("###"), json!("stop_sequence"), json!("###"))]
fn stop_fields_are_always_serialized(
#[case] stop_reason: Option<&str>,
#[case] stop_sequence: Option<&str>,
#[case] expected_reason: Value,
#[case] expected_sequence: Value,
) {
let body: Value = serde_json::to_value(response(stop_reason, stop_sequence, None, None))
.expect("serializable");
assert_eq!(body.get("stop_reason"), Some(&expected_reason));
assert_eq!(body.get("stop_sequence"), Some(&expected_sequence));
}
#[rstest]
#[case::absent(None, None)]
#[case::present(Some(json!({"input_tokens": 1})), Some(json!({"id": "c_1"})))]
fn usage_and_container_are_omitted_only_when_none(
#[case] usage: Option<Value>,
#[case] container: Option<Value>,
) {
let body: Value =
serde_json::to_value(response(None, None, usage.clone(), container.clone()))
.expect("serializable");
assert_eq!(body.get("usage").cloned(), usage);
assert_eq!(body.get("container").cloned(), container);
}
}

View file

@ -3,6 +3,21 @@ use serde_json::{Map, Value};
use crate::llms::openai::{ChatCompletionThinkingBlock, ChatCompletionToolCallChunk};
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
pub struct ProviderSpecificHeader {
#[serde(default)]
pub custom_llm_provider: String,
#[serde(default)]
pub extra_headers: Map<String, Value>,
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
#[serde(untagged)]
pub enum ProviderSpecificHeaders {
One(ProviderSpecificHeader),
Many(Vec<ProviderSpecificHeader>),
}
/// OpenAI `usage`, including the `prompt_tokens_details` split LiteLLM's Python
/// path reports so cost tracking sees the same numbers on either path.
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]

View file

@ -1461,6 +1461,7 @@ from .skills.main import (
from .containers.main import *
from .ocr.dispatch import *
from .chat_completions.dispatch import *
from .embeddings.dispatch import *
from .rust_bridge import rust
from .rag.main import *
from .sandbox.main import *
@ -1874,6 +1875,9 @@ if TYPE_CHECKING:
from .llms.openrouter.responses.transformation import (
OpenRouterResponsesAPIConfig as OpenRouterResponsesAPIConfig,
)
from .llms.bedrock.responses.transformation import (
BedrockOpenAIResponsesConfig as BedrockOpenAIResponsesConfig,
)
from .llms.bedrock_mantle.responses.transformation import (
BedrockMantleResponsesAPIConfig as BedrockMantleResponsesAPIConfig,
)

View file

@ -245,6 +245,7 @@ LLM_CONFIG_NAMES: Final = (
"PerplexityResponsesConfig",
"DatabricksResponsesAPIConfig",
"OpenRouterResponsesAPIConfig",
"BedrockOpenAIResponsesConfig",
"BedrockMantleResponsesAPIConfig",
"GoogleAIStudioInteractionsConfig",
"VertexAIInteractionsConfig",
@ -921,6 +922,10 @@ _LLM_CONFIGS_IMPORT_MAP: Final = {
"OpenAITextCompletionConfig",
),
"GroqChatConfig": (".llms.groq.chat.transformation", "GroqChatConfig"),
"BedrockOpenAIResponsesConfig": (
".llms.bedrock.responses.transformation",
"BedrockOpenAIResponsesConfig",
),
"BedrockMantleChatConfig": (
".llms.bedrock_mantle.chat.transformation",
"BedrockMantleChatConfig",

View file

@ -631,9 +631,9 @@ class LevelRoutingStreamHandler(logging.StreamHandler):
)
preferred: Final = sys.stdout if is_stdout_record else sys.stderr
if preferred is None or getattr(preferred, "closed", False):
self.stream = sys.stderr # rebind-ok: fall back to the pre-fix stream rather than raising per record
self.stream = sys.stderr
else:
self.stream = preferred # rebind-ok: StreamHandler.emit writes self.stream under the handler lock
self.stream = preferred
super().emit(record)

View file

@ -689,11 +689,12 @@ def init_redis_cluster(redis_kwargs) -> redis.RedisCluster:
verbose_logger.debug("init_redis_cluster: startup nodes are being initialized.")
from redis.cluster import ClusterNode
auth_kwargs: Final = _credential_provider_auth_kwargs(redis_kwargs)
args: Final = _get_redis_cluster_kwargs()
cluster_kwargs: Final = {}
for arg in redis_kwargs:
for arg in auth_kwargs:
if arg in args:
cluster_kwargs[arg] = redis_kwargs[arg]
cluster_kwargs[arg] = auth_kwargs[arg]
new_startup_nodes: Final[list[ClusterNode]] = []
@ -771,13 +772,13 @@ def _init_async_redis_sentinel(redis_kwargs) -> async_redis.Redis:
return sentinel.master_for(service_name, **connection_kwargs)
def _async_credential_provider(redis_connect_func: object | None) -> CredentialProvider | None:
"""The Azure AD and GCP IAM connect funcs run their AUTH exchange with the blocking client
API, so on an async connection their ``send_command``/``read_response`` calls return
coroutines nobody awaits and every connect fails. Async paths authenticate through a
``CredentialProvider`` instead, which redis-py consults per connection so the token stays
fresh. Any other ``redis_connect_func`` is left where it is, since redis-py awaits it
itself when it is a coroutine function."""
def _credential_provider_from_connect_func(redis_connect_func: object | None) -> CredentialProvider | None:
"""Translate IAM callbacks for paths that need credentials during the standard handshake.
Async connections cannot run blocking AUTH callbacks. Sync clusters authenticate before
invoking the callback, so they also need the provider during the initial handshake.
redis-py consults the provider for each connection, keeping token refresh intact.
"""
gcp_service_account: Final = getattr(redis_connect_func, "_gcp_service_account", None)
if gcp_service_account is not None:
return GCPIAMCredentialProvider(gcp_service_account)
@ -789,14 +790,13 @@ def _async_credential_provider(redis_connect_func: object | None) -> CredentialP
return None
def _async_auth_kwargs(redis_kwargs: dict) -> dict:
"""Swaps a connect func an async path cannot run for the equivalent credential provider,
which supersedes any static username or password redis-py would otherwise reject it with."""
def _credential_provider_auth_kwargs(redis_kwargs: dict) -> dict:
"""Use a credential provider instead of an IAM callback and conflicting static credentials."""
explicit_provider: Final = redis_kwargs.get("credential_provider")
credential_provider: Final = (
explicit_provider
if explicit_provider is not None
else _async_credential_provider(redis_kwargs.get("redis_connect_func"))
else _credential_provider_from_connect_func(redis_kwargs.get("redis_connect_func"))
)
if credential_provider is None:
return redis_kwargs
@ -834,7 +834,7 @@ def get_redis_async_client(
connection_pool: async_redis.BlockingConnectionPool | None = None,
**env_overrides,
) -> async_redis.Redis | async_redis.RedisCluster:
redis_kwargs: Final = _async_auth_kwargs(_get_redis_client_logic(**env_overrides))
redis_kwargs: Final = _credential_provider_auth_kwargs(_get_redis_client_logic(**env_overrides))
if "startup_nodes" in redis_kwargs:
from redis.cluster import ClusterNode
@ -906,7 +906,7 @@ def get_redis_async_client(
def get_redis_connection_pool(
**env_overrides,
) -> async_redis.BlockingConnectionPool | None:
redis_kwargs: Final = _async_auth_kwargs(_get_redis_client_logic(**env_overrides))
redis_kwargs: Final = _credential_provider_auth_kwargs(_get_redis_client_logic(**env_overrides))
verbose_logger.debug("get_redis_connection_pool: redis_kwargs", redis_kwargs)
if "startup_nodes" in redis_kwargs:

View file

@ -429,6 +429,7 @@ class LLMCachingHandler:
kwargs=kwargs,
cached_result=cached_result,
is_async=False,
custom_llm_provider=custom_llm_provider,
)
if not _should_defer_streaming_cache_hit_callbacks(cached_result=cached_result):

View file

@ -191,8 +191,6 @@ def _as_chat_reasoning_items(
) -> list[ChatCompletionReasoningItem] | None:
if not reasoning_items:
return None
# cast-ok: _BuiltReasoningItem is the structural shape ChatCompletionReasoningItem
# describes, and TypedDict invariance is what stops the two from unifying here.
return cast(list[ChatCompletionReasoningItem], list(reasoning_items))
@ -1370,7 +1368,7 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator):
if tool_call_index_map is None:
return output_index
if output_index not in tool_call_index_map:
tool_call_index_map[output_index] = len(tool_call_index_map) # mutable-ok: per-stream accumulator state
tool_call_index_map[output_index] = len(tool_call_index_map)
return tool_call_index_map[output_index]
@staticmethod

View file

@ -48,6 +48,7 @@ DEFAULT_BATCH_SIZE: Final = int(os.getenv("DEFAULT_BATCH_SIZE", 512))
DEFAULT_FLUSH_INTERVAL_SECONDS: Final = int(os.getenv("DEFAULT_FLUSH_INTERVAL_SECONDS", 5))
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"))
# 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
@ -1518,6 +1519,7 @@ PROMETHEUS_BUDGET_METRICS_REFRESH_INTERVAL_MINUTES: Final = int(
CLOUDZERO_EXPORT_INTERVAL_MINUTES: Final = int(os.getenv("CLOUDZERO_EXPORT_INTERVAL_MINUTES", 60))
MCP_TOOL_NAME_PREFIX: Final = "mcp_tool"
MAXIMUM_TRACEBACK_LINES_TO_LOG: Final = int(os.getenv("MAXIMUM_TRACEBACK_LINES_TO_LOG", 100))
PASSTHROUGH_UPSTREAM_ERROR_BODY_MAX_LOG_CHARS: Final = 4096
# Headers to control callbacks
X_LITELLM_DISABLE_CALLBACKS: Final = "x-litellm-disable-callbacks"
@ -1772,6 +1774,8 @@ RESPONSES_SESSION_LOOKUP_MAX_ATTEMPTS: Final = max(1, int(os.getenv("RESPONSES_S
RESPONSES_SESSION_LOOKUP_RETRY_INTERVAL: Final = float(os.getenv("RESPONSES_SESSION_LOOKUP_RETRY_INTERVAL", "0.2"))
SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE: Final = int(os.getenv("SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE", 10000))
PROXY_DB_LOOKUP_MAX_CONCURRENCY: Final = max(1, int(os.getenv("PROXY_DB_LOOKUP_MAX_CONCURRENCY", "25")))
PROXY_DB_LOOKUP_DEADLINE_SECONDS: Final = max(0.1, float(os.getenv("PROXY_DB_LOOKUP_DEADLINE_SECONDS", "10")))
PROXY_DB_LOOKUP_STALL_WINDOW_SECONDS: Final = max(0.0, float(os.getenv("PROXY_DB_LOOKUP_STALL_WINDOW_SECONDS", "30")))
DEFAULT_CRON_JOB_LOCK_TTL_SECONDS: Final = int(os.getenv("DEFAULT_CRON_JOB_LOCK_TTL_SECONDS", 60)) # 1 minute
PROXY_BUDGET_RESCHEDULER_MIN_TIME: Final = int(os.getenv("PROXY_BUDGET_RESCHEDULER_MIN_TIME", 597))
RESET_BUDGET_JOB_BATCH_SIZE: Final = max(1, int(os.getenv("RESET_BUDGET_JOB_BATCH_SIZE", "500")))

View file

@ -2017,8 +2017,8 @@ def _deployment_model_info(
return cast(ModelInfo, registered_deployment_info) # cast-ok: router registers deployment prices under its id
if litellm_logging_obj is None:
return None
litellm_params: Final = getattr(litellm_logging_obj, "litellm_params", None)
if litellm_params is None:
litellm_params: Final = litellm_logging_obj.litellm_params
if not litellm_params:
return None
return next(
(
@ -2036,7 +2036,9 @@ def _ocr_model_info(
router_model_id: str | None,
) -> OCRPricing | None:
deployment_info: Final = _deployment_model_info(litellm_logging_obj, custom_pricing, router_model_id)
litellm_params: Final = getattr(litellm_logging_obj, "litellm_params", None) if custom_pricing else None
litellm_params: Final = (
litellm_logging_obj.litellm_params if custom_pricing and litellm_logging_obj is not None else None
)
if litellm_params is None:
return deployment_info
return _layered_ocr_pricing(litellm_params, deployment_info)

View file

View file

@ -0,0 +1,95 @@
import inspect
from collections.abc import Awaitable, Callable, Coroutine, Mapping
from types import MappingProxyType
from typing import Final, TypeAlias, cast # noqa: TID251 # native binding selects a sync result or an async awaitable
from litellm import main
from litellm.rust_bridge.catalog import Route, RouteContext
from litellm.rust_bridge.dispatch import PublicDispatch, call_hook
from litellm.rust_bridge.embeddings.entrypoints import (
NATIVE_AEMBEDDING,
NATIVE_EMBEDDING,
LiteLLMEmbeddingRequest,
)
from litellm.rust_bridge.public_call import bind, optional_mapping, optional_str, signature
from litellm.types.utils import EmbeddingResponse
__all__ = ("aembedding", "embedding")
PythonEmbedding: TypeAlias = Callable[..., EmbeddingResponse | Coroutine[object, object, EmbeddingResponse]]
PythonAembedding: TypeAlias = Callable[..., Awaitable[EmbeddingResponse]]
_PYTHON_EMBEDDING: Final = cast( # cast-ok: [LIT006] preserve the legacy public callable contract
PythonEmbedding, main.embedding
)
_PYTHON_AEMBEDDING: Final = cast( # cast-ok: [LIT006] preserve the legacy public callable contract
PythonAembedding, main.aembedding
)
_EMBEDDING_SIGNATURE: Final = signature(_PYTHON_EMBEDDING)
def _public_request(
legacy: inspect.Signature, args: tuple[object, ...], kwargs: Mapping[str, object]
) -> LiteLLMEmbeddingRequest | None:
fields: Final = bind(legacy, args, kwargs)
if fields is None:
return None
model: Final = fields.get("model")
if not isinstance(model, str):
return None
extra: Final = optional_mapping(fields.get("kwargs")) or MappingProxyType({})
return LiteLLMEmbeddingRequest(
model=model,
input=fields.get("input"),
api_key=optional_str(fields.get("api_key")),
api_base=optional_str(fields.get("api_base")),
custom_llm_provider=optional_str(fields.get("custom_llm_provider")),
kwargs=extra,
)
def _context(request: LiteLLMEmbeddingRequest) -> RouteContext:
return RouteContext(Route.EMBEDDINGS, provider=request.custom_llm_provider, model=request.model)
_DISPATCH: Final = PublicDispatch(
route=Route.EMBEDDINGS,
request=lambda args, kwargs: _public_request(_EMBEDDING_SIGNATURE, args, kwargs),
context=_context,
bypass=lambda request: request.kwargs.get("aembedding") is True,
)
_ADISPATCH: Final = PublicDispatch(
route=Route.EMBEDDINGS,
request=lambda args, kwargs: _public_request(_EMBEDDING_SIGNATURE, args, kwargs),
context=_context,
)
def embedding(
*args: object,
**kwargs: object, # kwargs-ok: preserve the public embedding call shape
) -> EmbeddingResponse | Coroutine[object, object, EmbeddingResponse]:
return _DISPATCH.run(
args,
kwargs,
python=_PYTHON_EMBEDDING,
binding=NATIVE_EMBEDDING,
native=call_hook,
)
async def aembedding(*args: object, **kwargs: object) -> EmbeddingResponse: # kwargs-ok: preserve the public call shape
return await _ADISPATCH.arun(
args,
kwargs,
python=_PYTHON_AEMBEDDING,
binding=NATIVE_AEMBEDDING,
native=call_hook,
)
embedding.__doc__ = _PYTHON_EMBEDDING.__doc__
embedding.__wrapped__ = _PYTHON_EMBEDDING # pyright: ignore[reportFunctionMemberAccess] # preserve the legacy signature
aembedding.__doc__ = _PYTHON_AEMBEDDING.__doc__
aembedding.__wrapped__ = _PYTHON_AEMBEDDING # pyright: ignore[reportFunctionMemberAccess] # preserve the legacy signature

View file

@ -764,7 +764,7 @@ class MCPClient:
follow_redirects=True,
event_hooks=MappingProxyType(
{"response": [capture_upstream_error_response], "request": [guard] if guard else []}
), # mutable-ok: httpx types require lists of hooks
),
)
return factory
@ -921,9 +921,7 @@ class MCPClient:
with anyio.fail_after(max(self.timeout, MCP_TOOL_LISTING_TIMEOUT)):
for page_index in range(MCP_TOOL_LISTING_MAX_PAGES):
try:
page = await fetch_page( # rebind-ok: each SDK page replaces the previous one
None if cursor is None else PaginatedRequestParams(cursor=cursor)
)
page = await fetch_page(None if cursor is None else PaginatedRequestParams(cursor=cursor))
except MCPError as error:
if page_index > 0 and error.error.code == METHOD_NOT_FOUND:
raise RuntimeError("MCP list operation became unavailable during pagination") from error

View file

@ -129,7 +129,7 @@ async def list_tools_with_pagination(
)
tools.extend(result.tools)
next_cursor = getattr(result, "next_cursor", None)
next_cursor = result.next_cursor
if not isinstance(next_cursor, str) or not next_cursor:
return tools
if next_cursor in seen_cursors:

View file

@ -112,7 +112,7 @@ class ArizeLogger(OpenTelemetry):
if value is None or value in ("", "None"):
return None
try:
rate = float(value)
rate: Final = float(value)
except (TypeError, ValueError):
verbose_logger.warning(
"ArizeLogger: %s value %r is not a number; exporting the request",

View file

@ -1641,5 +1641,5 @@ def log_guardrail_information(func):
return async_wrapper(*args, **kwargs)
return sync_wrapper(*args, **kwargs)
vars(wrapper)[LOGS_GUARDRAIL_INFORMATION_MARKER] = True # rebind-ok: stamps the wrapper this call just built
vars(wrapper)[LOGS_GUARDRAIL_INFORMATION_MARKER] = True
return wrapper

View file

@ -21,7 +21,7 @@ from __future__ import annotations
import os
from datetime import datetime, timedelta, timezone
from typing import TYPE_CHECKING, Any, Final, Protocol
from typing import TYPE_CHECKING, Any, Final
import litellm
from litellm._logging import verbose_proxy_logger
@ -35,17 +35,6 @@ else:
AsyncIOScheduler = Any
class _PodLockManager(Protocol):
"""The subset of PodLockManager this logger drives to serialize the export across pods."""
@property
def redis_cache(self) -> object: ...
async def acquire_lock(self, cronjob_id: str) -> bool | None: ...
async def release_lock(self, cronjob_id: str) -> None: ...
def _parse_metrics_marker(
marker: object | None,
) -> datetime | None:
@ -237,13 +226,10 @@ class MavvrikFocusLogger(FocusLogger):
"""Scheduler entry point — uses Mavvrik-specific pod-lock key."""
from litellm.proxy.proxy_server import proxy_logging_obj # noqa: PLC0415
pod_lock_manager: _PodLockManager | None = None
if proxy_logging_obj is not None:
writer: Final[object] = getattr(proxy_logging_obj, "db_spend_update_writer", None)
if writer is not None:
pod_lock_manager = getattr(writer, "pod_lock_manager", None)
if pod_lock_manager and pod_lock_manager.redis_cache:
pod_lock_manager: Final = (
proxy_logging_obj.db_spend_update_writer.pod_lock_manager if proxy_logging_obj is not None else None
)
if pod_lock_manager is not None and pod_lock_manager.redis_cache:
acquired: Final = await pod_lock_manager.acquire_lock(cronjob_id=MAVVRIK_FOCUS_EXPORT_JOB_NAME)
if not acquired:
verbose_proxy_logger.debug("Mavvrik FOCUS export: unable to acquire pod lock")

View file

@ -366,9 +366,7 @@ class NewRelicMetricsLogger(CustomBatchLogger):
error to keep the client-error path (drop) distinct from 5xx (retry)."""
payload: Final = build_metric_payload(records=batch, window_start=window_start, now=time.time())
try:
status = (
await self.async_send_compressed_data(payload)
).status_code # rebind-ok: reassigned from the raised HTTPStatusError below
status = (await self.async_send_compressed_data(payload)).status_code
except HTTPStatusError as e:
status = e.response.status_code
except Exception as e: # noqa: BLE001 # transport/network failure re-queues the batch

View file

@ -62,10 +62,10 @@ if TYPE_CHECKING:
from litellm.proxy.proxy_server import UserAPIKeyAuth as _UserAPIKeyAuth
Span = _Span | Any
Tracer = _Tracer | Any
Context = _Context | Any
SpanExporter = _SpanExporter | Any
UserAPIKeyAuth = _UserAPIKeyAuth | Any
Tracer = _Tracer
Context = _Context
SpanExporter = _SpanExporter
UserAPIKeyAuth = _UserAPIKeyAuth
ManagementEndpointLoggingPayload = _ManagementEndpointLoggingPayload | Any
else:
Span = Any
@ -2730,7 +2730,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
self.handle_callback_failure(callback_name=self.callback_name or "opentelemetry")
verbose_logger.exception("OpenTelemetry logging error in set_attributes %s", str(e))
def _cast_as_primitive_value_type(self, value) -> str | bool | int | float:
def _cast_as_primitive_value_type(self, value: object) -> str | bool | int | float:
"""
Casts the value to a primitive OTEL type if it is not already a primitive type.

View file

@ -63,7 +63,24 @@ 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
**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. Caller-supplied `event_metadata` is **sanitized** before it reaches a span
trace.
**Post-response work is its own trace.** Spend tracking, the response cache write
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.
Caller-supplied `event_metadata` is **sanitized** before it reaches a span
(primitives only, no live objects, no secrets/headers, bounded) — see
`payloads.sanitize_event_metadata`.

View file

@ -56,8 +56,8 @@ from litellm.integrations.otel.plumbing.context import (
request_root_http_route,
request_root_span,
resolve_mcp_span_context,
resolve_parent_context,
resolve_request_span_context,
resolve_service_span_context,
set_request_baggage,
set_request_root_span,
)
@ -671,14 +671,17 @@ class OpenTelemetryV2(CustomLogger):
# 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
# server span the proxy threaded as ``parent_otel_span``. A background
# service call has neither, so it starts its own root trace.
parent_context: Final = resolve_parent_context(threaded=parent_otel_span)
# service call has neither, so it starts its own root trace, as does one
# that finished after the request span ended (linked back to it).
end_time_ns: Final = to_ns(end_time)
parent_context, links = resolve_service_span_context(threaded=parent_otel_span, end_time_ns=end_time_ns)
return self._emitter.emit(
role,
data,
parent_context=parent_context,
start_time_ns=to_ns(start_time),
end_time_ns=to_ns(end_time),
end_time_ns=end_time_ns,
links=links,
)
# ====================================================================== #

View file

@ -9,6 +9,7 @@ from opentelemetry import baggage
from opentelemetry.context import Context, get_current
from opentelemetry.sdk.trace import ReadableSpan
from opentelemetry.trace import (
INVALID_SPAN,
Link,
NonRecordingSpan,
Span,
@ -227,6 +228,28 @@ def resolve_parent_context(threaded: Span | None = None) -> Context:
return ctx
def resolve_service_span_context(
threaded: Span | None = None, end_time_ns: int | None = None
) -> 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.
"""
ctx: Final = resolve_parent_context(threaded)
parent: Final = get_current_span(ctx)
if not _ended_before(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:
return False
return end_time_ns is None or end_time_ns > span.end_time
def resolve_request_span_context() -> Context:
"""The parent context for a request-level span (the LLM call, a guardrail).

View file

@ -354,7 +354,7 @@ class _DrainPool:
def _drain_until_closed(self) -> None:
while True:
processor: SpanProcessor | None = self._pending.get() # rebind-ok: loop variable
processor: SpanProcessor | None = self._pending.get()
if processor is None:
return
_shutdown_quietly(processor)
@ -579,7 +579,7 @@ class TenantFanOutSpanProcessor(SpanProcessor):
span, destination.span_scope
):
continue
processor = self._acquire(destination) # rebind-ok: loop variable; pyright forbids Final in a loop
processor = self._acquire(destination)
if processor is None:
continue
try:

View file

@ -153,7 +153,7 @@ def destination_for(
endpoint, protocol = resolved
return OtelDestination(
endpoint=endpoint,
headers=MappingProxyType(dict(headers)), # mutable-ok: MappingProxyType needs a concrete mapping to wrap
headers=MappingProxyType(dict(headers)),
resource_attributes=MappingProxyType({"service.name": service_name}) if service_name else _NO_ATTRS,
callback_name=callback_name,
protocol=protocol,

View file

@ -131,7 +131,7 @@ def _paginated_table(repository: BaseRepository[_TableRowT]) -> _PaginatedPrisma
"""View a repository's prisma table through the pagination surface budget metrics need."""
return cast(
_PaginatedPrismaTable[_TableRowT],
repository.table, # cast-ok: prisma rows carry the budget columns the domain model declares
repository.table,
)

View file

@ -20,7 +20,8 @@ from litellm.constants import (
)
from litellm.types.utils import StandardLoggingPayload
_S3_LOG_PROMPTS_ONLY: Final = TypeAdapter(bool)
_S3_BOOL: Final = TypeAdapter(bool)
_UPLOAD_BOUND: Final = TypeAdapter(int)
def resolve_s3_log_prompts_only(configured: object, environ: Mapping[str, str] | None = None) -> bool:
@ -29,12 +30,42 @@ def resolve_s3_log_prompts_only(configured: object, environ: Mapping[str, str] |
if raw is None or raw == "":
return False
try:
return _S3_LOG_PROMPTS_ONLY.validate_python(raw.strip() if isinstance(raw, str) else raw)
return _S3_BOOL.validate_python(raw.strip() if isinstance(raw, str) else raw)
except ValidationError:
verbose_logger.warning("s3 logging: s3_log_prompts_only=%r is not a boolean, logging prompts only", raw)
return True
def resolve_s3_max_concurrent_uploads(configured: object, fallback: int) -> int:
if configured is None or configured == "":
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
)
return fallback
if bound < 1:
verbose_logger.warning(
"s3 logging: s3_max_concurrent_uploads=%r must be at least 1, using %s", configured, fallback
)
return fallback
return bound
def resolve_s3_batch_file_upload(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_batch_file_upload=%r is not a boolean, keeping per-request objects", configured
)
return False
def prompts_only_payload(payload: StandardLoggingPayload) -> StandardLoggingPayload:
return {**payload, "response": None}

View file

@ -3,26 +3,33 @@ 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, so we create tasks to upload each element individually
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
"""
import asyncio
import time
from collections.abc import Mapping
from datetime import datetime
from datetime import datetime, timezone
from typing import TYPE_CHECKING, Final, cast
from urllib.parse import quote
from uuid import uuid4
import httpx
import litellm
from litellm._logging import print_verbose, verbose_logger
from litellm.constants import DEFAULT_S3_BATCH_SIZE, DEFAULT_S3_FLUSH_INTERVAL_SECONDS
from litellm.constants import (
DEFAULT_S3_BATCH_SIZE,
DEFAULT_S3_FLUSH_INTERVAL_SECONDS,
DEFAULT_S3_MAX_CONCURRENT_UPLOADS,
)
from litellm.integrations.s3 import (
get_s3_object_download_filename,
get_s3_object_key,
prompts_only_payload,
resolve_s3_batch_file_upload,
resolve_s3_log_prompts_only,
resolve_s3_max_concurrent_uploads,
resolve_sse_params,
)
from litellm.litellm_core_utils.aws_partition import get_aws_dns_suffix
@ -43,7 +50,20 @@ if TYPE_CHECKING:
from botocore.credentials import Credentials
def _s3_key_parent(s3_object_key: str) -> str:
return s3_object_key.rsplit("/", 1)[0] if "/" in s3_object_key else ""
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")
class S3Logger(CustomBatchLogger, BaseAWSLLM):
preserve_events_added_during_flush = True
def __init__(
self,
s3_bucket_name: str | None = None,
@ -71,6 +91,8 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
s3_server_side_encryption: str | None = None,
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_batch_file_upload: bool = False,
s3_callback_params_override: dict | None = None,
**kwargs,
):
@ -112,7 +134,10 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
s3_server_side_encryption=s3_server_side_encryption,
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_batch_file_upload=s3_batch_file_upload,
)
self._upload_semaphore = asyncio.Semaphore(self.s3_max_concurrent_uploads)
verbose_logger.debug("s3 logger using endpoint url %s", s3_endpoint_url)
# IMPORTANT
@ -168,6 +193,8 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
s3_server_side_encryption: str | None = None,
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_batch_file_upload: bool = False,
params_source: dict | None = None,
):
"""
@ -226,6 +253,16 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
params.get("s3_sse_kms_key_id") or s3_sse_kms_key_id,
)
configured_bound: Final = params.get("s3_max_concurrent_uploads")
self.s3_max_concurrent_uploads = resolve_s3_max_concurrent_uploads(
s3_max_concurrent_uploads if configured_bound is None or configured_bound == "" else configured_bound,
DEFAULT_S3_MAX_CONCURRENT_UPLOADS,
)
self.s3_batch_file_upload = s3_batch_file_upload or resolve_s3_batch_file_upload(
params.get("s3_batch_file_upload")
)
def _build_object_url(self, s3_object_key: str) -> str:
"""
Build the exact URL that is both signed and sent, with the key percent-encoded once.
@ -347,7 +384,7 @@ 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):
async def async_upload_data_to_s3(self, batch_logging_element: s3BatchLoggingElement) -> bool:
try:
import base64
import hashlib
@ -364,7 +401,11 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
url: Final = self._build_object_url(batch_logging_element.s3_object_key)
# Convert JSON to string
json_string: Final = safe_dumps(batch_logging_element.payload)
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()
@ -374,7 +415,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
# Prepare the request
headers: Final = {
"Content-Type": "application/json",
"Content-Type": batch_logging_element.content_type,
"Content-MD5": content_md5,
"x-amz-content-sha256": content_hash,
"Content-Language": "en",
@ -421,27 +462,72 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
except Exception as e:
verbose_logger.exception("Error uploading to s3: %s", e)
self.handle_callback_failure(callback_name="S3Logger")
return False
return True
async def async_send_batch(self):
async def async_send_batch(self) -> None:
"""
Sends runs from self.log_queue.
Sends runs from self.log_queue
Returns: None
Raises: Does not raise an exception, will only verbose_logger.exception()
Raises S3BatchUploadError when any upload failed; CustomBatchLogger.flush_queue
keeps the surviving queue entries for the next flush.
"""
verbose_logger.debug("s3_v2 logger - sending batch of %s", len(self.log_queue))
if not self.log_queue:
batch: Final = tuple(self.log_queue)
if not batch:
return
verbose_logger.debug("s3_v2 logger - sending batch of %s", len(batch))
#########################################################
# Flush the log queue to s3
# the log queue can be bounded by DEFAULT_S3_BATCH_SIZE
# see custom_batch_logger.py which triggers the flush
#########################################################
for payload in self.log_queue:
asyncio.create_task(self.async_upload_data_to_s3(payload))
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:
return
self.log_queue = [*failed, *self.log_queue[len(batch) :]]
raise S3BatchUploadError(failed=len(failed), total=len(uploads))
def _batch_file_mode_active(self) -> bool:
if not self.s3_batch_file_upload:
return False
if litellm.cold_storage_custom_logger == "s3_v2":
verbose_logger.warning(
"s3 logging: s3_batch_file_upload is ignored because s3_v2 is the cold storage logger; "
"per-request objects are required for spend log lookups"
)
return False
return True
async def _upload_bounded(self, element: s3BatchLoggingElement) -> bool:
async with self._upload_semaphore:
return await self.async_upload_data_to_s3(element)
def _batch_file_elements(self, batch: tuple[s3BatchLoggingElement, ...]) -> tuple[s3BatchLoggingElement, ...]:
now: Final = datetime.now(timezone.utc)
groups: Final = {
parent: tuple(
element for element in batch if element.body is None and _s3_key_parent(element.s3_object_key) == parent
)
for parent in sorted({_s3_key_parent(element.s3_object_key) for element in batch if element.body is None})
}
return tuple(element for element in batch if element.body is not None) + tuple(
self._build_batch_file_element(elements, parent, now) for parent, elements in groups.items()
)
def _build_batch_file_element(
self, elements: tuple[s3BatchLoggingElement, ...], parent: str, now: datetime
) -> s3BatchLoggingElement:
batch_name: Final = f"batch_{now.strftime('%H-%M-%S')}_{uuid4().hex}"
return s3BatchLoggingElement(
payload={},
body="\n".join(safe_dumps(element.payload) for element in elements),
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",
)
def create_s3_batch_logging_element(
self,
@ -521,7 +607,11 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
url: Final = self._build_object_url(batch_logging_element.s3_object_key)
# Convert JSON to string
json_string: Final = safe_dumps(batch_logging_element.payload)
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()
@ -531,7 +621,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
# Prepare the request
headers: Final = {
"Content-Type": "application/json",
"Content-Type": batch_logging_element.content_type,
"Content-MD5": content_md5,
"x-amz-content-sha256": content_hash,
"Content-Language": "en",

View file

@ -11,6 +11,7 @@ across pods or stop races; the hook reads active jobs through a short-TTL cache.
import asyncio
import hashlib
import json
import random
import traceback
from collections.abc import Awaitable, Callable, Mapping, Sequence
@ -28,7 +29,7 @@ from litellm.caching.in_memory_cache import InMemoryCache
from litellm.constants import INTERNAL_CALL_ORIGIN_METADATA_KEY
from litellm.integrations.custom_logger import CustomLogger
from litellm.integrations.websearch_interception.tools import is_web_search_tool_responses
from litellm.litellm_core_utils.core_helpers import get_litellm_metadata_from_kwargs
from litellm.litellm_core_utils.core_helpers import get_litellm_metadata_from_kwargs, independent_snapshot
from litellm.litellm_core_utils.internal_call_metadata import sanitized_forwardable_call_metadata
from litellm.litellm_core_utils.llm_judge import (
default_router_provider,
@ -281,10 +282,8 @@ class _SurfaceOps:
request (messages plus translated generation params) and how its response yields
the judgeable final text. Membership in this table IS the sampling allowlist;
unknown call types fail closed. ``wire_params`` marks the surfaces whose params
come from the proxy's wire-body snapshot, which is taken before the guardrail
pre-call hook: those rows must not sample a request a pre-call guardrail rewrote,
or the shadow call would replay content (tools, unmasked entities) the guardrail
removed."""
come from the proxy's native request snapshot. Requests rewritten by guardrails
require a post-hook snapshot whose guardrail history is still current."""
__slots__ = ("chat_request", "final_text", "wire_params")
@ -311,19 +310,85 @@ _NON_MUTATING_GUARDRAIL_MODES: Final = frozenset(
)
def _guardrail_is_non_mutating(entry: Mapping[str, object], allowed_modes: frozenset[str]) -> bool:
modes: Final = entry.get("guardrail_mode")
return all(
isinstance(mode, str) and mode in allowed_modes
for mode in (modes if isinstance(modes, list | tuple) else (modes,))
)
def _request_mutating_guardrail_ran(request_metadata: Mapping[str, object]) -> bool:
"""Whether a guardrail that can rewrite the outbound request ran on this one, read
from the same guardrail-information entries spend logging uses. str-enum modes
compare equal to their plain-string values, and an entry whose mode is missing or
unrecognized counts as mutating."""
raw: Final = request_metadata.get("standard_logging_guardrail_information")
entries: Final = raw if isinstance(raw, Sequence) else ()
modes_per_entry: Final = tuple(entry.get("guardrail_mode") for entry in entries if isinstance(entry, Mapping))
return any(
not all(
mode in _NON_MUTATING_GUARDRAIL_MODES for mode in (modes if isinstance(modes, list | tuple) else (modes,))
not _guardrail_is_non_mutating(entry, _NON_MUTATING_GUARDRAIL_MODES)
for entry in entries
if isinstance(entry, Mapping)
)
def request_guardrail_fingerprint(request_metadata: Mapping[str, object]) -> str | None:
raw: Final = request_metadata.get("standard_logging_guardrail_information")
entries: Final = raw if isinstance(raw, Sequence) else ()
replay_safe_modes: Final = _NON_MUTATING_GUARDRAIL_MODES - frozenset(("logging_only",))
relevant: Final = tuple(
entry
for entry in entries
if isinstance(entry, Mapping) and not _guardrail_is_non_mutating(entry, replay_safe_modes)
)
try:
serialized: Final = json.dumps(relevant, sort_keys=True, default=str)
except (TypeError, ValueError):
return None
return hashlib.sha256(serialized.encode()).hexdigest()
@dataclass(frozen=True, slots=True)
class GuardrailRequestSnapshot:
body: Mapping[str, object]
fingerprint: str
@staticmethod
def capture(body: Mapping[str, object], metadata: Mapping[str, object]) -> "GuardrailRequestSnapshot | None":
if not _request_mutating_guardrail_ran(metadata):
return None
fingerprint: Final = request_guardrail_fingerprint(metadata)
if fingerprint is None:
return None
return GuardrailRequestSnapshot(
body=MappingProxyType(
_CHAT_REQUEST_ADAPTER.validate_python(
independent_snapshot(dict(body)) # mutable-ok: snapshot helper requires a plain dictionary
)
),
fingerprint=fingerprint,
)
for modes in modes_per_entry
def _post_guardrail_kwargs(
kwargs: Mapping[str, object],
request_metadata: Mapping[str, object],
ops: _SurfaceOps,
guardrail_snapshot: GuardrailRequestSnapshot | None,
) -> Mapping[str, object] | None:
if guardrail_snapshot is None or guardrail_snapshot.fingerprint != request_guardrail_fingerprint(request_metadata):
return None
raw_params: Final = kwargs.get("litellm_params")
litellm_params: Final = raw_params if isinstance(raw_params, Mapping) else _EMPTY_METADATA
raw_request: Final = litellm_params.get("proxy_server_request")
request: Final = raw_request if isinstance(raw_request, Mapping) else _EMPTY_METADATA
body: Final = guardrail_snapshot.body
return MappingProxyType(
{
**kwargs,
"messages": body.get("input" if ops is _RESPONSES_OPS else "messages"),
"system": body.get("system"),
"instructions": body.get("instructions"),
"litellm_params": MappingProxyType(
{**litellm_params, "proxy_server_request": MappingProxyType({**request, "body": body})}
),
}
)
@ -808,7 +873,6 @@ class ShadowEvalLogger(CustomLogger):
await prisma.db.litellm_shadowevalattempt.group_by(
by=["job_id"],
count=True,
# mutable-ok: Prisma aggregate spec
sum={"judge_cost": True, "shadow_cost": True, "shadow_classifier_cost": True},
where={"job_id": {"in": [str(record.id) for record in records]}}, # mutable-ok: Prisma filter
)
@ -836,7 +900,7 @@ class ShadowEvalLogger(CustomLogger):
{target: tuple(job for _, job in group) for target, group in groupby(by_target, key=itemgetter(0))}
)
await self._jobs_cache.async_set_cache(_JOBS_CACHE_KEY, jobs)
self._job_starts = {} # rebind-ok: new generation, counts absorbed into the fill
self._job_starts = {}
return jobs
except Exception as e: # noqa: BLE001 # a DB blip must never break request logging
verbose_logger.debug("shadow_eval: active-job read failed: %s", e)
@ -881,6 +945,8 @@ class ShadowEvalLogger(CustomLogger):
response_obj: object,
start_time: object,
end_time: object,
*,
guardrail_snapshot: GuardrailRequestSnapshot | None = None,
) -> None:
try:
payload: Final[StandardLoggingPayload | None] = kwargs.get("standard_logging_object") # pyright: ignore[reportAssignmentType] # untyped callback kwargs
@ -914,8 +980,13 @@ class ShadowEvalLogger(CustomLogger):
ops: Final = _SURFACE_OPS.get(str(payload.get("call_type") or ""))
if ops is None:
return # only surfaces this table can normalize are comparable; unknown types fail closed
if ops.wire_params and _request_mutating_guardrail_ran(request_metadata):
return # the wire-body snapshot predates the rewrite; replaying it would resurrect stripped content
sample_kwargs: Final = (
_post_guardrail_kwargs(kwargs, request_metadata, ops, guardrail_snapshot)
if ops.wire_params and _request_mutating_guardrail_ran(request_metadata)
else kwargs
)
if sample_kwargs is None:
return
active_jobs: Final = await self._active_jobs()
eligible: Final = self._sampled_jobs(
tuple(job for target in targets for job in active_jobs.get(target, ())),
@ -927,7 +998,7 @@ class ShadowEvalLogger(CustomLogger):
return
sample: Final = _judgeable_sample(
ops,
kwargs,
sample_kwargs,
MappingProxyType(dict(payload.get("model_parameters") or {})), # mutable-ok: frozen snapshot
response_obj,
)
@ -961,7 +1032,7 @@ class ShadowEvalLogger(CustomLogger):
real_cache_hit=real_cache_hit,
control_tier=control_tier,
shadow_params=shadow_params,
parent_metadata=MappingProxyType(dict(request_metadata)), # mutable-ok: frozen snapshot
parent_metadata=MappingProxyType(dict(request_metadata)),
)
).add_done_callback(self._release_shadow_slot)
except Exception as e: # noqa: BLE001 # logging hooks must never fail the request
@ -1275,7 +1346,7 @@ class ShadowEvalLogger(CustomLogger):
{
"role": "user",
"content": _judge_user_prompt(conversation, response_a, response_b, _tool_definitions_text(tools)),
}, # mutable-ok: SDK message
},
]
try:
response: Final = await judge_acompletion(

View file

@ -1849,7 +1849,7 @@ class WebSearchInterceptionLogger(CustomLogger):
for tool_call in tool_calls:
# Handle both Anthropic-style input and OpenAI-style function.arguments
query = None
tool_args: dict | None = None # mutable-ok: the tool call's own arguments dict
tool_args: dict[str, object] | None = None # mutable-ok: the tool call's own arguments dict
if "input" in tool_call and isinstance(tool_call["input"], dict):
tool_args = tool_call["input"]
query = tool_args.get("query")

View file

@ -79,9 +79,7 @@ async def _fetch_interaction(context: BackgroundInteractionPollContext) -> Inter
custom_llm_provider=context.custom_llm_provider,
api_key=context.api_key,
api_base=context.api_base,
**{
"no-log": True
}, # mutable-ok: "no-log" is not a valid identifier, so it can only be passed through a mapping
**{"no-log": True},
)

View file

@ -365,7 +365,7 @@ def _budget_reservation_on_auth_object(user_api_key_auth: object) -> object:
return getattr(user_api_key_auth, "budget_reservation", None)
def budget_reservation_from_metadata(metadata: Mapping[str, object]) -> dict | None:
def budget_reservation_from_metadata(metadata: Mapping[str, object]) -> dict[str, object] | None:
stamped: Final = metadata.get("user_api_key_budget_reservation")
if isinstance(stamped, dict):
return stamped
@ -764,4 +764,4 @@ def set_response_cost_in_hidden_params(response: _CarriesHiddenParams, cost: flo
**(additional_headers if isinstance(additional_headers, Mapping) else _NO_HEADERS),
RESPONSE_COST_HEADER: cost,
}
hidden_params["additional_headers"] = merged # rebind-ok: the caller's record is the point
hidden_params["additional_headers"] = merged

View file

@ -60,13 +60,13 @@ class _HasProxyErrorType(Protocol):
_MESSAGE_PATTERNS: Final[tuple[tuple[re.Pattern[str], str], ...]] = (
(re.compile(r"upstream passthrough request failed", re.IGNORECASE), UPSTREAM_PASSTHROUGH),
(
re.compile(r"budget has been exceeded|max budget|exceeded.*budget|crossed budget", re.IGNORECASE),
re.compile(r"budget has been exceeded|max budget|crossed budget", re.IGNORECASE),
BUDGET_EXCEEDED,
),
(re.compile(r"no healthy deployments?|no deployments available", re.IGNORECASE), NO_HEALTHY_DEPLOYMENTS),
(re.compile(r"not allowed to access model due to tags configuration", re.IGNORECASE), MODEL_ACCESS_DENIED),
(re.compile(r"upstream passthrough request failed", re.IGNORECASE), UPSTREAM_PASSTHROUGH),
(re.compile(r"is not supported for provider|not implemented", re.IGNORECASE), UNSUPPORTED_OPERATION),
(
re.compile(r"context window|context length|(prompt|input) is too long|tokens? ?> ?\d+ ?maximum", re.IGNORECASE),
@ -155,6 +155,14 @@ _CLASS_CODE_TABLE: Final[tuple[tuple[tuple[type[BaseException], ...], str], ...]
)
def _exceeded_before_budget(message: str) -> bool:
"""Linear-time equivalent of ``re.search(r"exceeded.*budget", message, re.IGNORECASE)``."""
return any(
(start := line.find("exceeded")) != -1 and line.find("budget", start + len("exceeded")) != -1
for line in message.lower().split("\n")
)
def _classify_by_message(message: str, patterns: tuple[tuple[re.Pattern[str], str], ...]) -> str | None:
return next((code for pattern, code in patterns if pattern.search(message)), None)
@ -183,7 +191,9 @@ def normalize_error(exc: Exception | None, status_code: str, message: str) -> st
by_proxy_type: Final = _PROXY_ERROR_TYPE_MAP.get(proxy_type) if isinstance(proxy_type, str) else None
if by_proxy_type is not None:
return by_proxy_type
by_message: Final = _classify_by_message(message, _MESSAGE_PATTERNS)
by_message: Final = (
BUDGET_EXCEEDED if _exceeded_before_budget(message) else _classify_by_message(message, _MESSAGE_PATTERNS)
)
if by_message is not None:
return by_message
by_class: Final = _classify_by_class(exc)

View file

@ -32,9 +32,7 @@ def get_supported_openai_params(
- None if unmapped
"""
if not custom_llm_provider:
custom_llm_provider = declared_authenticating_provider(
model
) # rebind-ok: resolving would run the provider's OAuth flow
custom_llm_provider = declared_authenticating_provider(model)
if not custom_llm_provider:
try:
custom_llm_provider = litellm.get_llm_provider(model=model)[1]

View file

@ -21,20 +21,18 @@ class JSONFragmentAccumulator:
def __init__(self) -> None:
self._chunks: list[str] = [] # mutable-ok: O(1) append; string concat would copy the buffer each time
self._buffer: str = (
"" # mutable-ok: lazily materialized join of _chunks, rebuilt only when _chunks is non-empty
)
self._offset: int = 0 # mutable-ok: cursor past already-consumed values; avoids re-slicing on every pop
self._could_close: bool = False # mutable-ok: cached heuristic; rescanning past fragments was itself O(n^2)
self._buffer: str = ""
self._offset: int = 0
self._could_close: bool = False
def __bool__(self) -> bool:
return bool(self._chunks) or self._offset < len(self._buffer)
def append(self, fragment: str) -> None:
self._chunks.append(fragment) # mutable-ok: see __init__
self._chunks.append(fragment)
stripped: Final = fragment.rstrip()
if stripped:
self._could_close = stripped[-1] in ("}", "]") # mutable-ok: see __init__
self._could_close = stripped[-1] in ("}", "]")
def could_close_json(self) -> bool:
"""
@ -50,8 +48,8 @@ class JSONFragmentAccumulator:
if not self._chunks:
return
unconsumed: Final = self._buffer[self._offset :]
self._buffer = unconsumed + "".join(self._chunks) # mutable-ok: merge pending fragments, once per append batch
self._offset = 0 # mutable-ok: see __init__
self._buffer = unconsumed + "".join(self._chunks)
self._offset = 0
self._chunks = [] # mutable-ok: see __init__
def pop_next_value(self) -> tuple[bool, object]:
@ -69,7 +67,7 @@ class JSONFragmentAccumulator:
while start < length and self._buffer[start].isspace():
start += 1
if start >= length:
self._offset = start # mutable-ok: see __init__
self._offset = start
return False, None
decoder: Final = json.JSONDecoder()
try:
@ -77,11 +75,11 @@ class JSONFragmentAccumulator:
except json.JSONDecodeError:
return False, None
decoded, end_index = cast("tuple[object, int]", raw_value) # cast-ok: raw_decode returns tuple[Any, int]
self._offset = end_index # mutable-ok: see __init__
self._offset = end_index
if self._offset >= len(self._buffer):
self._buffer = "" # mutable-ok: see __init__
self._offset = 0 # mutable-ok: see __init__
self._could_close = False # mutable-ok: buffer is empty, nothing can close
self._buffer = ""
self._offset = 0
self._could_close = False
return True, decoded
def snapshot(self) -> str:
@ -91,7 +89,7 @@ class JSONFragmentAccumulator:
def set(self, value: str) -> None:
"""Replace the buffer's contents with a single fragment."""
self._chunks = [] # mutable-ok: see __init__
self._buffer = value # mutable-ok: see __init__
self._offset = 0 # mutable-ok: see __init__
self._buffer = value
self._offset = 0
stripped: Final = value.rstrip()
self._could_close = bool(stripped) and stripped[-1] in ("}", "]") # mutable-ok: see __init__
self._could_close = bool(stripped) and stripped[-1] in ("}", "]")

View file

@ -226,6 +226,7 @@ if TYPE_CHECKING:
from litellm.integrations.otel.logger import OpenTelemetryV2
from litellm.integrations.otel.model.config import ExporterSpec, OpenTelemetryV2Config
from litellm.integrations.shadow_eval_logger import GuardrailRequestSnapshot
from litellm.litellm_core_utils.llm_cost_calc.utils import BilledTokenRates
from litellm.llms.base_llm.passthrough.transformation import PassthroughStreamCollector
from litellm.proxy.hooks.autorouter_baseline_cache import BaselineCacheContext, CapturedBaselineObservation
@ -714,6 +715,7 @@ class Logging(LiteLLMLoggingBaseClass):
self._defer_async_logging: bool = False
self._enqueue_deferred_logging: Callable[[], None] | None = None
self._on_detached_stream_failure: Callable[[Exception], Awaitable[None]] | None = None
self.shadow_eval_request_snapshot: GuardrailRequestSnapshot | None = None
def set_response_timing_metrics(self, timing_metrics: Mapping[str, float]) -> None:
"""Keep ``_response_ms`` / ``litellm_overhead_time_ms`` for a result that has no ``_hidden_params``."""
@ -1564,14 +1566,14 @@ class Logging(LiteLLMLoggingBaseClass):
attr = "debug"
if json_logs:
callattr = getattr(verbose_logger, attr)
callattr = verbose_logger.warning if attr == "warning" else verbose_logger.debug
callattr(
"RAW RESPONSE:\n{}\n\n".format(
self.model_call_details.get("original_response", self.model_call_details)
),
)
else:
callattr = getattr(verbose_logger, attr)
callattr = verbose_logger.warning if attr == "warning" else verbose_logger.debug
callattr(
"RAW RESPONSE:\n{}\n\n".format(
self.model_call_details.get("original_response", self.model_call_details)
@ -2825,6 +2827,7 @@ class Logging(LiteLLMLoggingBaseClass):
):
continue
self.shadow_eval_request_snapshot = None
self.model_call_details, result = callback.logging_hook(
kwargs=self.model_call_details,
result=result,
@ -3391,6 +3394,7 @@ class Logging(LiteLLMLoggingBaseClass):
):
continue
self.shadow_eval_request_snapshot = None
self.model_call_details, result = await callback.async_logging_hook(
kwargs=self.model_call_details,
result=result,
@ -3450,6 +3454,8 @@ class Logging(LiteLLMLoggingBaseClass):
)
if isinstance(callback, CustomLogger): # custom logger class
from litellm.integrations.shadow_eval_logger import ShadowEvalLogger
model_call_details: dict = self.model_call_details
##################################
# call redaction hook for custom logger
@ -3460,7 +3466,19 @@ class Logging(LiteLLMLoggingBaseClass):
model_call_details=model_call_details, custom_logger=callback
)
##################################
if self.stream is True:
if isinstance(callback, ShadowEvalLogger) and (
not self.stream or "async_complete_streaming_response" in model_call_details
):
await callback.async_log_success_event(
kwargs=model_call_details,
response_obj=model_call_details["async_complete_streaming_response"]
if self.stream
else result,
start_time=start_time,
end_time=end_time,
guardrail_snapshot=self.shadow_eval_request_snapshot,
)
elif self.stream is True:
if "async_complete_streaming_response" in model_call_details:
await callback.async_log_success_event(
kwargs=model_call_details,
@ -5173,7 +5191,7 @@ def _maybe_construct_otel_v2(callback_name: str, _in_memory_loggers: list[Custom
for callback in _in_memory_loggers:
if (
isinstance(callback, OpenTelemetryV2)
and getattr(callback, "callback_name", None) == callback_name
and callback.callback_name == callback_name
and (serves_a_destination or not _exports_nowhere(callback.config))
):
return callback
@ -5864,7 +5882,7 @@ class StandardLoggingPayloadSetup:
base_model: str | None,
custom_pricing: bool | None,
custom_llm_provider: str | None,
init_response_obj: Any | BaseModel | dict,
init_response_obj: object,
api_base: str | None = None,
) -> StandardLoggingModelInformation:
model_cost_name: Final = _select_model_name_for_cost_calc(
@ -5897,9 +5915,7 @@ class StandardLoggingPayloadSetup:
return model_cost_information
@staticmethod
def get_final_response_obj(
response_obj: dict, init_response_obj: Any | BaseModel | dict, kwargs: dict
) -> dict | str | list | None:
def get_final_response_obj(response_obj: dict, init_response_obj: object, kwargs: dict) -> dict | str | list | None:
"""
Get final response object after redacting the message input/output from logging
"""
@ -6342,7 +6358,7 @@ def _get_status_fields(
def _extract_response_obj_and_hidden_params(
init_response_obj: Any | BaseModel | dict,
init_response_obj: object,
original_exception: Exception | None,
) -> tuple[dict, dict | None]:
"""Extract response_obj and hidden_params from init_response_obj."""
@ -6645,11 +6661,11 @@ def get_standard_logging_object_payload(
cost_breakdown=request_cost_breakdown,
autorouter_savings=autorouter_savings,
autorouter_savings_estimate=(
{
{ # mutable-ok: spend-log JSON serialization requires plain mappings
"version": 3,
"status": "unknown",
"reason": "pending_projection",
} # mutable-ok: spend-log JSON serialization requires plain mappings
}
if captured_baseline is not None
else (
{ # mutable-ok: spend-log JSON serialization requires plain mappings

View file

@ -5,7 +5,12 @@ from typing import Final
from pydantic import TypeAdapter, ValidationError
from typing_extensions import assert_never
from litellm.types.utils import StandardLoggingZeroCostDiagnostic, Usage
from litellm.types.utils import (
CompletionTokensDetailsWrapper,
PromptTokensDetailsWrapper,
StandardLoggingZeroCostDiagnostic,
Usage,
)
ZERO_COST_COUNTER_NAME: Final = "litellm_zero_cost_requests_total"
@ -18,8 +23,8 @@ _NESTED_PRICING: Final = TypeAdapter(Mapping[str, object] | tuple[object, ...])
_MAX_PRICING_DEPTH: Final = 4
def _audio_tokens(details: object) -> int:
audio_tokens: Final = getattr(details, "audio_tokens", None)
def _audio_tokens(details: PromptTokensDetailsWrapper | CompletionTokensDetailsWrapper | None) -> int:
audio_tokens: Final = details.audio_tokens if details is not None else None
return audio_tokens if isinstance(audio_tokens, int) and audio_tokens > 0 else 0

View file

@ -372,9 +372,7 @@ from collections import defaultdict
def _handle_invalid_parallel_tool_calls(
tool_calls: list[
ChatCompletionMessageToolCall | ChatCompletionMessageCustomToolCall
], # mutable-ok: patched in place via slice assignment
tool_calls: list[ChatCompletionMessageToolCall | ChatCompletionMessageCustomToolCall],
):
"""
Handle hallucinated parallel tool call from openai - https://community.openai.com/t/model-tries-to-call-unknown-function-multi-tool-use-parallel/490653

View file

@ -208,7 +208,7 @@ def _content_parts_contain_image(parts: Sequence[object]) -> bool:
for _ in range(_IMAGE_SCAN_MAX_DEPTH):
if any(isinstance(part, Mapping) and part.get("type") in _IMAGE_CONTENT_PART_TYPES for part in frontier):
return True
frontier = tuple( # rebind-ok: depth-bounded frontier walk
frontier = tuple(
nested
for part in frontier
if isinstance(part, Mapping)
@ -2003,11 +2003,11 @@ def strip_encrypted_reasoning_from_messages(messages: object) -> None:
"""
if not isinstance(messages, list):
return
for content in _anthropic_content_lists(cast(list[object], messages)): # cast-ok: untyped client json
for content in anthropic_content_lists(cast(list[object], messages)): # cast-ok: untyped client json
_strip_encrypted_reasoning_from_blocks(content)
def _anthropic_content_lists(messages: Sequence[object]) -> Iterator[object]:
def anthropic_content_lists(messages: Sequence[object]) -> Iterator[object]:
return (
cast(list[object], content) # cast-ok: narrowed by isinstance
for message in messages
@ -2020,7 +2020,7 @@ def _anthropic_content_lists(messages: Sequence[object]) -> Iterator[object]:
def _strip_encrypted_reasoning_from_blocks(content: object) -> None:
blocks: Final = cast(list[object], content) # cast-ok: narrowed by the caller's isinstance
kept: Final = tuple(block for block in blocks if not is_encrypted_reasoning_block(block))
blocks[:] = kept # rebind-ok: shared with fallback snapshot
blocks[:] = kept
def _reasoning_replay_group_key(indexed_block: tuple[int, Mapping[str, object]]) -> str:

View file

@ -446,7 +446,7 @@ def _render_chat_template(env, chat_template: str, bos_token: str, eos_token: st
async def _afetch_and_extract_template(
model: str, chat_template: Any | None, get_config_fn, get_template_fn
model: str, chat_template: str | None, get_config_fn, get_template_fn
) -> tuple[str, str, str]:
"""
Async version: Fetch template and tokens from HuggingFace.
@ -500,7 +500,7 @@ async def _afetch_and_extract_template(
def _fetch_and_extract_template(
model: str, chat_template: Any | None, get_config_fn, get_template_fn
model: str, chat_template: str | None, get_config_fn, get_template_fn
) -> tuple[str, str, str]:
"""
Sync version: Fetch template and tokens from HuggingFace.

View file

@ -83,7 +83,7 @@ def get_stable_session_id(litellm_params: object | None) -> str | None:
return None
def add_provider_affinity_header( # mutable-ok: downstream handlers add auth and signing headers
def add_provider_affinity_header(
headers: Mapping[str, object], litellm_params: object | None
) -> dict[str, object]: # mutable-ok: downstream handlers add auth and signing headers
header_name: Final = _get_provider_affinity_header_name(litellm_params)

View file

@ -475,9 +475,7 @@ class ChunkProcessor:
def get_combined_tool_content(
self, tool_call_chunks: Sequence["_ToolCallChunk"]
) -> list[
ChatCompletionMessageToolCall | ChatCompletionMessageCustomToolCall
]: # mutable-ok: assigned verbatim to Message.tool_calls, a list field
) -> list[ChatCompletionMessageToolCall | ChatCompletionMessageCustomToolCall]:
tool_calls_list: list[
ChatCompletionMessageToolCall | ChatCompletionMessageCustomToolCall
] = [] # mutable-ok: see return type

View file

@ -1329,7 +1329,7 @@ class CustomStreamWrapper:
"is_finished": chunk_finish_reason is not None,
"finish_reason": chunk_finish_reason,
"original_chunk": cached_chunk,
"tool_calls": (getattr(cached_choice.delta, "tool_calls", None) if cached_choice is not None else None),
"tool_calls": cached_choice.delta.tool_calls if cached_choice is not None else None,
}
completion_obj["content"] = response_obj["text"]

View file

@ -48,7 +48,7 @@ def _registry_api_key(agent_litellm_params: Mapping[str, object]) -> str | None:
return configured_api_key if isinstance(configured_api_key, str) else None
def _registry_headers(agent_litellm_params: Mapping[str, object]) -> dict[str, Any] | None:
def _registry_headers(agent_litellm_params: Mapping[str, object]) -> dict[str, object] | None:
stored_headers: Final = agent_litellm_params.get("headers")
if not isinstance(stored_headers, Mapping):
return None

View file

@ -199,9 +199,7 @@ def _write_back_system_block(system: object, block_idx: int, response: str) -> N
return
text_blocks: Final = tuple(block for block in system if isinstance(block, dict) and block.get("type") == "text")
if block_idx < len(text_blocks):
text_blocks[block_idx]["text"] = (
response # mutable-ok: guardrails rewrite the caller's request payload in place
)
text_blocks[block_idx]["text"] = response
def _write_back_message_text(message: _WritableMessage, target: MessageTextTarget, response: str) -> None:
@ -211,22 +209,16 @@ def _write_back_message_text(message: _WritableMessage, target: MessageTextTarge
match target:
case MessageContentTarget():
if isinstance(content, str):
message["content"] = response # mutable-ok: guardrails rewrite the caller's request payload in place
message["content"] = response
case ContentBlockTextTarget(content_idx=content_idx):
if isinstance(content, list):
content[content_idx]["text"] = (
response # mutable-ok: guardrails rewrite the caller's request payload in place
)
content[content_idx]["text"] = response
case ToolResultStringTarget(content_idx=content_idx):
if isinstance(content, list):
content[content_idx]["content"] = (
response # mutable-ok: guardrails rewrite the caller's request payload in place
)
content[content_idx]["content"] = response
case ToolResultBlockTextTarget(content_idx=content_idx, block_idx=block_idx):
if isinstance(content, list):
content[content_idx]["content"][block_idx]["text"] = (
response # mutable-ok: guardrails rewrite the caller's request payload in place
)
content[content_idx]["content"][block_idx]["text"] = response
case _:
assert_never(target)
@ -248,9 +240,9 @@ def _write_back_tool_use(
block: Final = content[target.content_idx] if isinstance(content, list) else None
if not isinstance(block, dict):
return
block["input"] = rewritten_input # mutable-ok: guardrails rewrite the caller's request payload in place
block["input"] = rewritten_input
if shape.name is not None and shape.name != block.get("name"):
block["name"] = shape.name # mutable-ok: guardrails rewrite the caller's request payload in place
block["name"] = shape.name
@dataclass(frozen=True, slots=True)
@ -603,13 +595,9 @@ class AnthropicMessagesHandler(BaseTranslation):
*(item for one_message in extracted for item in one_message.scanned),
)
texts_to_check: Final = [item.text for item in scanned] # mutable-ok: GenericGuardrailAPIInputs takes list[str]
images_to_check: Final = [
image for one_message in extracted for image in one_message.images
] # mutable-ok: GenericGuardrailAPIInputs takes list[str]
images_to_check: Final = [image for one_message in extracted for image in one_message.images]
scanned_tool_calls: Final = tuple(item for one_message in extracted for item in one_message.tool_calls)
tool_calls_to_check: Final = [
item.tool_call for item in scanned_tool_calls
] # mutable-ok: GenericGuardrailAPIInputs takes list[ChatCompletionToolCallChunk]
tool_calls_to_check: Final = [item.tool_call for item in scanned_tool_calls]
pre_guardrail_tool_calls: Final = _tool_call_shapes(tool_calls_to_check)
# Step 2: Apply guardrail to all texts and tool calls in batch
@ -697,9 +685,7 @@ class AnthropicMessagesHandler(BaseTranslation):
return data
def _hoisted_top_level_system_message(
self, data: dict
) -> AllMessageValues | None: # mutable-ok: API message payload
def _hoisted_top_level_system_message(self, data: Mapping[str, object]) -> AllMessageValues | None:
"""Return the system message produced by translating the top-level prompt."""
system: Final = data.get("system")
if not system:
@ -736,7 +722,7 @@ class AnthropicMessagesHandler(BaseTranslation):
if isinstance(content, str):
return (
{"role": "system", "content": content} if content else None # mutable-ok: API message payload
) # mutable-ok: API message payload
)
if not isinstance(content, list):
return None
blocks: Final[list[dict[str, object]]] = [] # mutable-ok: API message payload
@ -749,14 +735,14 @@ class AnthropicMessagesHandler(BaseTranslation):
anthropic_block: dict[str, object] = { # mutable-ok: API message payload
"type": "text",
"text": text,
} # mutable-ok: API message payload
}
cache_control = block.get("cache_control")
if cache_control:
anthropic_block["cache_control"] = deepcopy(cache_control)
blocks.append(anthropic_block)
return (
{"role": "system", "content": blocks} if blocks else None # mutable-ok: API message payload
) # mutable-ok: API message payload
)
@staticmethod
def _fold_leading_systems_into_top_level(
@ -1098,9 +1084,7 @@ class AnthropicMessagesHandler(BaseTranslation):
match item.target:
case SystemStringTarget():
if isinstance(data.get("system"), str):
data["system"] = (
guardrail_response # mutable-ok: guardrails rewrite the caller's request payload in place
)
data["system"] = guardrail_response
case SystemBlockTextTarget(block_idx=block_idx):
_write_back_system_block(data.get("system"), block_idx, guardrail_response)
case (

View file

@ -1591,7 +1591,7 @@ def _flatten_web_search_results_in_message(message: object) -> object:
return {**message, "content": [b for b in rewritten if b is not None]} # mutable-ok: JSON wire format
def flatten_unencrypted_web_search_results_in_anthropic_messages( # mutable-ok: as sibling sanitizers
def flatten_unencrypted_web_search_results_in_anthropic_messages(
messages: list[Any],
) -> list[Any]:
"""

View file

@ -11,7 +11,6 @@ from typing import (
Final,
Literal,
Protocol,
cast, # noqa: TID251 # rebuilt message_delta dict spans the ContentBlockDelta/MessageBlockDelta union
get_args,
)
@ -27,6 +26,7 @@ from litellm.types.llms.anthropic import (
ContentBlockDelta,
ContextManagementResponse,
MessageBlockDelta,
MessageDelta,
StreamingContentBlockDeltaType,
UsageDelta,
UsageIteration,
@ -1028,26 +1028,22 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
self,
processed_chunk: ContentBlockDelta | MessageBlockDelta,
) -> ContentBlockDelta | MessageBlockDelta:
if processed_chunk.get("type") != "message_delta" or not self._refusal_text:
if processed_chunk["type"] != "message_delta" or not self._refusal_text:
return processed_chunk
delta: Final = cast(Mapping[str, object], processed_chunk["delta"]) # cast-ok: keys checked before use
delta: Final = processed_chunk["delta"]
if delta.get("stop_reason") == "max_tokens":
return processed_chunk
from litellm.llms.anthropic.experimental_pass_through.messages.utils import (
refusal_stop_details,
)
return cast( # cast-ok: rebuilt dict matches the message_delta TypedDict shape for this branch
ContentBlockDelta | MessageBlockDelta,
{ # mutable-ok: fresh translation payload; never mutated after construction
**processed_chunk,
"delta": { # mutable-ok: fresh message_delta payload; never mutated after construction
**delta,
"stop_reason": "refusal",
"stop_details": refusal_stop_details(self._refusal_text),
},
},
)
refusal_delta: Final[MessageDelta] = {
**delta,
"stop_reason": "refusal",
"stop_details": refusal_stop_details(self._refusal_text),
}
refusal_chunk: Final[MessageBlockDelta] = {**processed_chunk, "delta": refusal_delta}
return refusal_chunk
@staticmethod
def _delta_has_content(processed_chunk: Mapping[str, object]) -> bool:

View file

@ -50,14 +50,14 @@ def _build_tool_result_message(tool_results: Sequence[Mapping[str, object]]) ->
"""Turn executed tool results into the user message Anthropic expects."""
return AnthropicMessagesUserMessageParam(
role="user",
content=tuple(
content=[
AnthropicMessagesToolResultParam(
type="tool_result",
tool_use_id=str(result.get("tool_call_id") or ""),
content=str(result.get("result") or ""),
)
for result in tool_results
),
],
)

View file

@ -88,9 +88,7 @@ class AnthropicMessagesStreamCacheWriter:
try:
events: Final = _split_sse_events(collected_stream.decode("utf-8"))
cached_payload: Final = {
CACHED_STREAM_EVENTS_KEY: events
} # mutable-ok: cache backends serialize plain dicts
cached_payload: Final = {CACHED_STREAM_EVENTS_KEY: events}
await litellm.cache.async_add_cache(
cached_payload,
dynamic_cache_object=self.caching_handler.dual_cache,

View file

@ -186,7 +186,7 @@ def _sse_event(event_type: str, payload: Mapping[str, object]) -> bytes:
def _incomplete_stream_error_sse_event() -> bytes:
return _sse_event( # mutable-ok: one-shot JSON payload, never mutated after construction
return _sse_event(
"error",
{"type": "error", "error": {"type": "api_error", "message": INCOMPLETE_STREAM_ERROR_MESSAGE}},
)

View file

@ -37,7 +37,7 @@ def _mapping_field(container: object, key: str) -> object | None:
"""One key of a raw provider payload, or None when the payload is not a mapping."""
if not isinstance(container, Mapping):
return None
return cast(Mapping[str, object], container).get(key) # cast-ok: raw payload, callers re-check every value
return container.get(key)
def _mapping_str_field(container: object, key: str) -> str | None:

View file

@ -148,13 +148,11 @@ class LiteLLMAnthropicToResponsesAPIAdapter:
if isinstance(content, str):
return (
[{"type": "input_text", "text": content}] if content else [] # mutable-ok: API message payload
) # mutable-ok: API message payload
)
if not isinstance(content, list):
return [] # mutable-ok: API message payload
return [ # mutable-ok: API message payload
with_prompt_cache_breakpoint(
{"type": "input_text", "text": text}, block.get("prompt_cache_breakpoint")
) # mutable-ok: API message payload
with_prompt_cache_breakpoint({"type": "input_text", "text": text}, block.get("prompt_cache_breakpoint"))
for block in content
if isinstance(block, dict) and block.get("type") == "text" and (text := block.get("text")) # pyright: ignore[reportUnnecessaryIsInstance] # untrusted client payload
]
@ -171,7 +169,7 @@ class LiteLLMAnthropicToResponsesAPIAdapter:
cls,
summary: Iterable[object],
encrypted_content: object,
) -> dict[str, Any] | None: # mutable-ok: API message payload
) -> dict[str, object] | None: # mutable-ok: API message payload
"""The one Anthropic block for a Responses reasoning item.
The item's encrypted reasoning rides the block's opaque field (`signature`, or
@ -200,7 +198,7 @@ class LiteLLMAnthropicToResponsesAPIAdapter:
@classmethod
def _assistant_group_to_input_items(
cls, group: tuple[Mapping[str, object], ...]
) -> tuple[dict[str, Any], ...]: # mutable-ok: API message payload
) -> tuple[dict[str, object], ...]: # mutable-ok: API message payload
first: Final = group[0]
btype: Final = first.get("type")
if btype in ("thinking", "redacted_thinking"):

View file

@ -59,9 +59,7 @@ def logged_responses_stream(all_chunks: Sequence[str], logging_obj: Logging) ->
terminal_event: Final = OpenAIResponsesAPIConfig.parse_terminal_event_from_stream_chunks(all_chunks=all_chunks)
if terminal_event is None:
return None
logging_obj.call_type = (
RESPONSES_RELAY_SHAPE.call_type.value
) # rebind-ok: routes cost calculation to the relayed shape's pricing path
logging_obj.call_type = RESPONSES_RELAY_SHAPE.call_type.value
return terminal_event

View file

@ -73,9 +73,7 @@ class AzureFoundryFluxImageGenerationConfig(GPTImageGenerationConfig):
normalized_model: Final = model.lower().replace(".", "-").replace("_", "-")
return "flux-2-flex" if "flux-2-flex" in normalized_model else "flux-2-pro"
def get_supported_openai_params( # mutable-ok: inherited config contract returns a list
self, model: str
) -> list[OpenAIImageGenerationOptionalParams]:
def get_supported_openai_params(self, model: str) -> list[OpenAIImageGenerationOptionalParams]:
if not self.is_flux2_model(model):
return super().get_supported_openai_params(model)
return [ # mutable-ok: BaseImageGenerationConfig requires a list

View file

@ -95,9 +95,7 @@ def logged_relay_shape(
parsed: Final = shape.parse(body)
except ValidationError:
return None
logging_obj.call_type = (
shape.call_type.value
) # rebind-ok: routes cost calculation to the relayed shape's pricing path
logging_obj.call_type = shape.call_type.value
return parsed

View file

@ -0,0 +1,154 @@
"""Codex CLI wire-format quirks shared by the Responses API providers that need them.
Codex sends history item types that api.openai.com accepts but other Responses
backends reject with ``400 Invalid 'input': value did not match any expected
variant``. Both Amazon Bedrock endpoints reject them:
- ``bedrock-mantle.{region}.api.aws`` (verified against ``openai.gpt-5.6-sol``)
- ``bedrock-runtime.{region}.amazonaws.com/openai/v1`` (same, verified separately)
They are *history* items, so they only appear from the second turn of a session
onward -- a first-turn request succeeds and hides the problem entirely.
Codex also sends a ``web_search`` tool on every turn. api.openai.com runs that tool
itself; a backend with no server-side tools rejects the whole request over it, so
the same providers drop the tool types their backend does not accept.
Both helpers are pure transforms that report what they rewrote or dropped; callers
do their own logging, so each provider keeps its own wording.
"""
import json
from collections.abc import Mapping, Sequence
from typing import Final
from typing_extensions import ReadOnly, TypedDict
from litellm.types.llms.openai import ResponseInputParam
AGENT_MESSAGE_INPUT_ITEM_TYPE: Final = "agent_message"
CONTEXT_COMPACTION_INPUT_ITEM_TYPE: Final = "context_compaction"
LOCAL_SHELL_CALL_INPUT_ITEM_TYPE: Final = "local_shell_call"
class _RewrittenOutputTextBlock(TypedDict):
type: ReadOnly[str]
text: ReadOnly[str]
class _RewrittenAssistantMessageItem(TypedDict):
type: ReadOnly[str]
role: ReadOnly[str]
content: ReadOnly[tuple[_RewrittenOutputTextBlock, ...]]
class _RewrittenCompactionItem(TypedDict):
type: ReadOnly[str]
encrypted_content: ReadOnly[str]
class _RewrittenFunctionCallItem(TypedDict):
type: ReadOnly[str]
call_id: ReadOnly[str]
name: ReadOnly[str]
arguments: ReadOnly[str]
def _agent_message_text(item: "Mapping[str, object]") -> str:
content: Final = item.get("content")
if not isinstance(content, list):
return ""
return "".join(
str(block.get("text") or block.get("encrypted_content") or "") for block in content if isinstance(block, dict)
)
def _normalize_agent_message_item(item: "Mapping[str, object]") -> "_RewrittenAssistantMessageItem | None":
text: Final = _agent_message_text(item)
if not text:
return None
rewritten: Final[_RewrittenAssistantMessageItem] = {
"type": "message",
"role": "assistant",
"content": ({"type": "output_text", "text": text},),
}
return rewritten
def _normalize_context_compaction_item(item: "Mapping[str, object]") -> "_RewrittenCompactionItem | None":
encrypted_content: Final = item.get("encrypted_content")
if not isinstance(encrypted_content, str) or not encrypted_content:
return None
rewritten: Final[_RewrittenCompactionItem] = {"type": "compaction", "encrypted_content": encrypted_content}
return rewritten
def _normalize_local_shell_call_item(item: "Mapping[str, object]") -> "_RewrittenFunctionCallItem | None":
call_id: Final = item.get("call_id")
if not isinstance(call_id, str) or not call_id:
return None
action: Final = item.get("action")
rewritten: Final[_RewrittenFunctionCallItem] = {
"type": "function_call",
"call_id": call_id,
"name": "local_shell",
"arguments": json.dumps(action) if isinstance(action, dict) else "{}",
}
return rewritten
def _normalize_input_item(item: object) -> "tuple[object, str | None]":
"""Returns (normalized item, or None to drop it; original type when rewritten)."""
if not isinstance(item, dict):
return item, None
item_type: Final = item.get("type")
if item_type == AGENT_MESSAGE_INPUT_ITEM_TYPE:
return _normalize_agent_message_item(item), item_type
if item_type == CONTEXT_COMPACTION_INPUT_ITEM_TYPE:
return _normalize_context_compaction_item(item), item_type
if item_type == LOCAL_SHELL_CALL_INPUT_ITEM_TYPE:
return _normalize_local_shell_call_item(item), item_type
return item, None
def normalize_codex_input_items(
input: "str | ResponseInputParam",
) -> "tuple[str | ResponseInputParam, tuple[str, ...]]":
"""Rewrite the Codex history item types a Responses backend rejects.
``agent_message`` (Codex multi-agent traffic; its ``encrypted_content`` slot
carries the plaintext payload when the model never issued encrypted args)
becomes an assistant message, ``context_compaction`` becomes the ``compaction``
spelling these backends accept, and ``local_shell_call`` becomes the
``function_call`` its recorded ``function_call_output`` already pairs with.
Returns the normalized input and the sorted set of types that were rewritten,
so the caller can log in its own words. Non-list input is returned untouched.
"""
if not isinstance(input, list):
return input, ()
normalized: Final = tuple(_normalize_input_item(item) for item in input)
rewritten_types: Final = tuple(sorted(frozenset(item_type for _, item_type in normalized if item_type is not None)))
kept: Final = [i for i, _ in normalized if i is not None] # mutable-ok: downstream narrows on isinstance(list)
# Codex passthrough items sit outside the OpenAI input union.
return kept, rewritten_types # pyright: ignore[reportReturnType] # see above
def drop_unsupported_tools(
tools: "Sequence[object]", supported_types: "frozenset[str]"
) -> "tuple[tuple[object, ...], tuple[str, ...]]":
"""Keep the tools whose ``type`` the backend accepts; non-dict tools pass through.
Returns the kept tools and the sorted set of dropped types.
"""
kept: Final = tuple(tool for tool in tools if not isinstance(tool, dict) or tool.get("type") in supported_types)
dropped_types: Final = tuple(
sorted(
frozenset(
str(tool.get("type"))
for tool in tools
if isinstance(tool, dict) and tool.get("type") not in supported_types
)
)
)
return kept, dropped_types

View file

@ -130,6 +130,22 @@ class BaseResponsesAPIConfig(ABC):
) -> dict:
pass
async def async_transform_responses_api_request(
self,
model: str,
input: str | ResponseInputParam,
response_api_optional_request_params: dict,
litellm_params: GenericLiteLLMParams,
headers: dict,
) -> dict:
return self.transform_responses_api_request(
model=model,
input=input,
response_api_optional_request_params=response_api_optional_request_params,
litellm_params=litellm_params,
headers=headers,
)
@abstractmethod
def transform_response_api_response(
self,

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