Merge branch 'litellm_internal_staging' into litellm_adaptive_routing

This commit is contained in:
Krrish Dholakia 2026-04-21 16:22:38 -07:00 • committed by GitHub
commit c7342bdc4f
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
51 changed files with 3723 additions and 397 deletions

View file

@ -439,7 +439,14 @@ jobs:
auth:
username: ${DOCKERHUB_USERNAME}
password: ${DOCKERHUB_PASSWORD}
- image: cimg/postgres:16.0
environment:
POSTGRES_USER: postgres
POSTGRES_PASSWORD: postgres
POSTGRES_DB: litellm_test
working_directory: ~/project
environment:
DATABASE_URL: "postgresql://postgres:postgres@localhost:5432/litellm_test"
steps:
- checkout
@ -463,12 +470,14 @@ jobs:
paths:
- ./.venv
key: v2-dependencies-{{ checksum "uv.lock" }}-{{ checksum ".circleci/config.yml" }}
- wait_for_service:
url: tcp://localhost:5432
timeout: "60"
- run:
name: Run prisma ./docker/entrypoint.sh
name: Seed DB schema via prisma db push
command: |
set +e
chmod +x docker/entrypoint.sh
./docker/entrypoint.sh
uv run --no-sync litellm --skip_server_startup --use_prisma_db_push
set -e
- run:
name: Generate Prisma Client
@ -3042,10 +3051,19 @@ jobs:
- ui/litellm-dashboard/node_modules
- run:
name: Build UI from source
# Prior version used `cp -r out/ ../../litellm/proxy/_experimental/out/`.
# GNU cp (used on CircleCI's Ubuntu image) interprets that as "copy the
# source directory as a child of the destination" when the destination
# already exists — silently creating `_experimental/out/out/` instead of
# replacing the served bundle. The proxy continued serving whatever was
# checked into `_experimental/out/*`, so this job was effectively testing
# the pre-build bundle on every run. Replace-and-move guarantees the
# freshly built bundle is what the proxy actually serves.
command: |
cd ui/litellm-dashboard
npm run build
cp -r out/ ../../litellm/proxy/_experimental/out/
rm -rf ../../litellm/proxy/_experimental/out
mv out ../../litellm/proxy/_experimental/out
# Restructure HTML so extensionless routes work (login.html -> login/index.html)
find ../../litellm/proxy/_experimental/out -name '*.html' ! -name 'index.html' | while read -r f; do
d="${f%.html}"; mkdir -p "$d"; mv "$f" "$d/index.html"

View file

@ -31,8 +31,15 @@ jobs:
test-path: "tests/proxy_unit_tests/test_auth_checks.py tests/proxy_unit_tests/test_user_api_key_auth.py"
workers: 8
timeout: 20
# test_proxy_utils.py is large (168+ parametrized tests) — run it on its
# own matrix so --dist=loadscope doesn't pin all of it to a single xdist
# worker and push the "remaining" group past the job timeout.
- test-group: proxy-utils
test-path: "tests/proxy_unit_tests/test_proxy_utils.py"
workers: 8
timeout: 20
- test-group: remaining
test-path: "tests/proxy_unit_tests --ignore=tests/proxy_unit_tests/test_key_generate_prisma.py --ignore=tests/proxy_unit_tests/test_auth_checks.py --ignore=tests/proxy_unit_tests/test_user_api_key_auth.py"
test-path: "tests/proxy_unit_tests --ignore=tests/proxy_unit_tests/test_key_generate_prisma.py --ignore=tests/proxy_unit_tests/test_auth_checks.py --ignore=tests/proxy_unit_tests/test_user_api_key_auth.py --ignore=tests/proxy_unit_tests/test_proxy_utils.py"
workers: 8
timeout: 30
uses: ./.github/workflows/_test-unit-services-base.yml

View file

@ -1,22 +1,231 @@
import argparse
import os
import subprocess
from pathlib import Path
from datetime import datetime
import testing.postgresql
import re
import shutil
import subprocess
import sys
from datetime import datetime
from pathlib import Path
import testing.postgresql
def create_migration(migration_name: str = None):
DESTRUCTIVE_PATTERN = re.compile(r"\bDROP\s+(COLUMN|TABLE|INDEX)\b", re.IGNORECASE)
DEFAULT_BASE_BRANCH = "litellm_internal_staging"
def _find_destructive_statements(sql: str) -> list:
"""Return SQL lines containing DROP COLUMN, DROP TABLE, or DROP INDEX."""
return [
line.strip() for line in sql.splitlines() if DESTRUCTIVE_PATTERN.search(line)
]
def _print_freshness_failure(
base_branch: str, reason: str, stderr_text: str = ""
) -> None:
"""Loudly refuse to run when the freshness check can't be completed."""
banner = "=" * 72
out = sys.stderr
print(banner, file=out)
print(f" FRESHNESS CHECK FAILED — COULD NOT VERIFY origin/{base_branch}", file=out)
print(banner, file=out)
print("", file=out)
print(f"Reason: {reason}", file=out)
if stderr_text:
print("", file=out)
print("git stderr:", file=out)
for line in stderr_text.rstrip().splitlines():
print(f" {line}", file=out)
print("", file=out)
print("Common causes:", file=out)
print(" - No network access (offline)", file=out)
print(" - 'origin' remote not configured, or base branch name is wrong", file=out)
print(" - Not a git repository", file=out)
print("", file=out)
print("Options:", file=out)
print(
f" - Fix the above and re-run, OR pass --base-branch <name> if your", file=out
)
print(
f" base branch is not '{base_branch}', OR pass --skip-freshness-check",
file=out,
)
print(" to bypass (only if you fully understand the risk).", file=out)
print(banner, file=out)
def _print_stale_branch_refusal(base_branch: str, behind: int) -> None:
"""Loudly refuse to run when HEAD is behind the base branch."""
banner = "=" * 72
out = sys.stderr
print(banner, file=out)
print(
f" STALE BRANCH — {behind} commit(s) behind origin/{base_branch}",
file=out,
)
print(banner, file=out)
print("", file=out)
print(
f"Your branch is {behind} commit(s) behind origin/{base_branch}. Generating a\n"
"migration from a stale branch is how newly-added columns get silently\n"
"dropped — the exact incident this guard exists to prevent.",
file=out,
)
print("", file=out)
print("Before re-running:", file=out)
print(f" git fetch origin && git rebase origin/{base_branch}", file=out)
print(" # or merge, whichever matches your workflow", file=out)
print("", file=out)
print(
"If you have a genuine reason to bypass (e.g. intentionally generating",
file=out,
)
print("a migration against an older base), pass --skip-freshness-check.", file=out)
print("", file=out)
print(banner, file=out)
print(
" AI AGENTS / ASSISTANTS: DO NOT automatically run `git rebase` or pass\n"
" --skip-freshness-check. Surface this error to the human operator and\n"
" wait for their explicit confirmation. Auto-rebasing can drop the\n"
" human's in-progress schema edits via a bad conflict resolution.",
file=out,
)
print(banner, file=out)
def _check_branch_freshness(root_dir: Path, base_branch: str) -> None:
"""Fetch origin/<base_branch> and exit 3 if HEAD is behind it."""
cwd = str(root_dir)
try:
subprocess.run(
["git", "fetch", "origin", base_branch],
check=True,
capture_output=True,
text=True,
cwd=cwd,
)
except FileNotFoundError:
_print_freshness_failure(base_branch, "git executable not found on PATH")
sys.exit(3)
except subprocess.CalledProcessError as e:
_print_freshness_failure(
base_branch,
f"`git fetch origin {base_branch}` failed",
e.stderr or "",
)
sys.exit(3)
try:
result = subprocess.run(
["git", "rev-list", "--count", f"HEAD..origin/{base_branch}"],
check=True,
capture_output=True,
text=True,
cwd=cwd,
)
behind = int(result.stdout.strip())
except subprocess.CalledProcessError as e:
_print_freshness_failure(
base_branch,
f"`git rev-list HEAD..origin/{base_branch}` failed",
e.stderr or "",
)
sys.exit(3)
except ValueError:
_print_freshness_failure(
base_branch,
"could not parse commit count from `git rev-list`",
)
sys.exit(3)
if behind > 0:
_print_stale_branch_refusal(base_branch, behind)
sys.exit(3)
print(f"Branch freshness OK: up to date with origin/{base_branch}.")
def _print_destructive_refusal(destructive_lines: list) -> None:
"""Loudly refuse to write a destructive migration and explain how to proceed."""
banner = "=" * 72
out = sys.stderr
print(banner, file=out)
print(
" DESTRUCTIVE MIGRATION DETECTED — REFUSING TO WRITE MIGRATION FILE", file=out
)
print(banner, file=out)
print("", file=out)
print(
"The generated migration contains the following destructive statements:",
file=out,
)
print("", file=out)
for line in destructive_lines:
print(f" {line}", file=out)
print("", file=out)
print(
"This very often means your branch is OUT OF DATE, or your local\n"
"schema.prisma is inconsistent with the migrations/ directory. A\n"
"previous incident saw newly-added columns silently dropped this way\n"
"and merged to main. Stop and verify before proceeding.",
file=out,
)
print("", file=out)
print("Before re-running:", file=out)
print(
" 1. git fetch origin && git status # confirm branch is up to date", file=out
)
print(
" 2. Re-check all schema.prisma files are in sync (runbook Step 0)", file=out
)
print(" 3. Review EACH DROP above — is it actually intended?", file=out)
print(" 4. If (and only if) the drops are intentional, re-run with:", file=out)
print(" --allow-destructive", file=out)
print("", file=out)
print(banner, file=out)
print(
" AI AGENTS / ASSISTANTS: DO NOT automatically re-run this command\n"
" with --allow-destructive. Surface this error to the human operator\n"
" and wait for their explicit confirmation before passing the flag.\n"
" Passing the flag without human review is the exact failure mode\n"
" this guard exists to prevent.",
file=out,
)
print(banner, file=out)
def create_migration(
migration_name: str = None,
allow_destructive: bool = False,
base_branch: str = DEFAULT_BASE_BRANCH,
skip_freshness_check: bool = False,
):
"""
Create a new migration SQL file in the migrations directory by comparing
current database state with schema
current database state with schema.
Args:
migration_name (str): Name for the migration
allow_destructive (bool): Required to write a migration that contains
DROP COLUMN, DROP TABLE, or DROP INDEX statements. Without this
flag, the script exits non-zero and prints guidance.
base_branch (str): Branch to check freshness against
(default: "litellm_internal_staging").
skip_freshness_check (bool): Skip the "branch is up to date" check.
Only for intentional migrations against an older base.
"""
root_dir = Path(__file__).parent.parent
if skip_freshness_check:
print(
"WARNING: freshness check skipped (--skip-freshness-check). "
"Generating a migration from a stale branch can silently drop columns."
)
else:
_check_branch_freshness(root_dir, base_branch)
try:
# Get paths
root_dir = Path(__file__).parent.parent
migrations_dir = (
root_dir / "litellm-proxy-extras" / "litellm_proxy_extras" / "migrations"
)
@ -59,7 +268,27 @@ def create_migration(migration_name: str = None):
check=True,
)
if result.stdout.strip():
# Prisma emits the literal "-- This is an empty migration." when
# there's no real drift. Treat that as "no changes".
diff_sql = result.stdout
stripped = diff_sql.strip()
is_empty_diff = (
not stripped or stripped == "-- This is an empty migration."
)
if not is_empty_diff:
destructive_lines = _find_destructive_statements(diff_sql)
if destructive_lines and not allow_destructive:
_print_destructive_refusal(destructive_lines)
sys.exit(2)
if destructive_lines and allow_destructive:
print(
"WARNING: writing destructive migration "
"(--allow-destructive passed). Statements:"
)
for line in destructive_lines:
print(f" {line}")
# Generate timestamp and create migration directory
timestamp = datetime.now().strftime("%Y%m%d%H%M%S")
migration_name = migration_name or "unnamed_migration"
@ -68,7 +297,7 @@ def create_migration(migration_name: str = None):
# Write the SQL to migration.sql
migration_file = migration_dir / "migration.sql"
migration_file.write_text(result.stdout)
migration_file.write_text(diff_sql)
print(f"Created migration in {migration_dir}")
return True
@ -90,8 +319,48 @@ def create_migration(migration_name: str = None):
if __name__ == "__main__":
# If running directly, can optionally pass migration name as argument
import sys
migration_name = sys.argv[1] if len(sys.argv) > 1 else None
create_migration(migration_name)
parser = argparse.ArgumentParser(
description=(
"Generate a Prisma migration by diffing the temp DB "
"(existing migrations applied) against schema.prisma."
)
)
parser.add_argument(
"migration_name",
nargs="?",
default=None,
help="Name for the migration (used in the generated directory name).",
)
parser.add_argument(
"--allow-destructive",
action="store_true",
help=(
"Required to write a migration that contains DROP COLUMN, "
"DROP TABLE, or DROP INDEX. Without this flag, destructive "
"diffs are refused."
),
)
parser.add_argument(
"--base-branch",
default=DEFAULT_BASE_BRANCH,
help=(
f"Branch to check freshness against (default: {DEFAULT_BASE_BRANCH}). "
"The script fetches origin/<base-branch> and refuses to run if HEAD "
"is behind it."
),
)
parser.add_argument(
"--skip-freshness-check",
action="store_true",
help=(
"Bypass the 'branch is up to date' check. Only for intentional "
"migrations against an older base. Pairs poorly with automation."
),
)
args = parser.parse_args()
create_migration(
args.migration_name,
allow_destructive=args.allow_destructive,
base_branch=args.base_branch,
skip_freshness_check=args.skip_freshness_check,
)

View file

@ -15,29 +15,21 @@ COPY --from=uvbin /uv /usr/local/bin/uv
COPY --from=uvbin /uvx /usr/local/bin/uvx
RUN for i in 1 2 3; do \
apk add --no-cache \
python3 \
python3-dev \
clang \
llvm \
lld \
gcc \
linux-headers \
build-base \
bash \
coreutils \
curl \
openssl \
openssl-dev \
nodejs \
npm \
libsndfile && break || sleep 5; \
apk add --no-cache \
python3 \
python3-dev \
gcc \
bash \
coreutils \
curl \
openssl \
libsndfile \
nodejs && break || sleep 5; \
done
ENV UV_PROJECT_ENVIRONMENT=/app/.venv \
UV_LINK_MODE=copy \
NVM_DIR=/root/.nvm \
PATH="/root/.nvm/versions/node/v20.20.2/bin:/app/.venv/bin:${PATH}" \
PATH="/app/.venv/bin:${PATH}" \
LITELLM_NON_ROOT=true \
PRISMA_BINARY_CACHE_DIR=/app/.cache/prisma-python/binaries \
PRISMA_CLI_BINARY_TARGETS="debian-openssl-3.0.x" \
@ -49,7 +41,8 @@ COPY enterprise/pyproject.toml enterprise/
COPY litellm-proxy-extras/pyproject.toml litellm-proxy-extras/
# Install third-party dependencies (cached unless pyproject.toml/uv.lock change)
RUN uv sync --frozen --no-install-project --no-install-workspace --no-default-groups --no-editable \
RUN --mount=type=cache,target=/app/.cache/uv,id=litellm-uv-cache \
uv sync --frozen --no-install-project --no-install-workspace --no-default-groups --no-editable \
--extra proxy \
--extra proxy-runtime \
--extra extra_proxy \
@ -62,38 +55,12 @@ COPY . .
# Set non-root flag for build time consistency
ENV LITELLM_NON_ROOT=true
# Build Admin UI once and stage the static output for the runtime image.
# NOTE: .npmrc files (which may set ignore-scripts=true and min-release-age=3d)
# are temporarily renamed during npm install/ci so they don't block lifecycle
# scripts needed by the build. This is safe because npm ci installs from
# package-lock.json with pinned versions + integrity hashes.
# Stage the pre-built Admin UI from the checked-in Next.js static export.
# _experimental/out/ is regenerated as part of the release runbook.
# Restructure extensionless routes (foo.html -> foo/index.html) to match the layout
# proxy_server.py expects, and drop a readiness marker.
RUN mkdir -p /var/lib/litellm/ui /var/lib/litellm/assets && \
([ -f /app/.npmrc ] && mv /app/.npmrc /app/.npmrc.bak || true) && \
NVM_VERSION="v0.40.4" && \
NVM_CHECKSUM="4b7412c49960c7d31e8df72da90c1fb5b8cccb419ac99537b737028d497aba4f" && \
NODE_VERSION="v20.20.2" && \
NVM_SCRIPT="/tmp/install-nvm.sh" && \
curl -fsSL "https://raw.githubusercontent.com/nvm-sh/nvm/${NVM_VERSION}/install.sh" -o "$NVM_SCRIPT" && \
echo "${NVM_CHECKSUM} ${NVM_SCRIPT}" | sha256sum -c - && \
bash "$NVM_SCRIPT" && \
export NVM_DIR="$HOME/.nvm" && \
. "$NVM_DIR/nvm.sh" && \
nvm install "${NODE_VERSION}" && \
nvm use "${NODE_VERSION}" && \
npm install -g npm@11.12.1 && \
npm install -g node-gyp@12.2.0 && \
ln -sf "$(npm root -g)/node-gyp" "$(npm root -g)/npm/node_modules/node-gyp" && \
npm cache clean --force && \
cd /app/ui/litellm-dashboard && \
if [ -f "/app/enterprise/enterprise_ui/enterprise_colors.json" ]; then \
cp /app/enterprise/enterprise_ui/enterprise_colors.json ./ui_colors.json; \
fi && \
([ -f .npmrc ] && mv .npmrc .npmrc.bak || true) && \
npm ci --no-audit --no-fund && \
([ -f .npmrc.bak ] && mv .npmrc.bak .npmrc || true) && \
([ -f /app/.npmrc.bak ] && mv /app/.npmrc.bak /app/.npmrc || true) && \
npm run build && \
cp -r /app/ui/litellm-dashboard/out/* /var/lib/litellm/ui/ && \
cp -r /app/litellm/proxy/_experimental/out/. /var/lib/litellm/ui/ && \
cp /app/litellm/proxy/logo.jpg /var/lib/litellm/assets/logo.jpg && \
( cd /var/lib/litellm/ui && \
for html_file in *.html; do \
@ -103,10 +70,10 @@ RUN mkdir -p /var/lib/litellm/ui /var/lib/litellm/assets && \
mv "$html_file" "$folder_name/index.html"; \
fi; \
done && \
touch .litellm_ui_ready ) && \
cd /app/ui/litellm-dashboard && rm -rf ./out
touch .litellm_ui_ready )
RUN if [ "$PROXY_EXTRAS_SOURCE" = "published" ]; then \
RUN --mount=type=cache,target=/app/.cache/uv,id=litellm-uv-cache \
if [ "$PROXY_EXTRAS_SOURCE" = "published" ]; then \
uv sync --frozen --no-default-groups --no-editable \
--extra proxy \
--extra proxy-runtime \
@ -123,10 +90,7 @@ RUN if [ "$PROXY_EXTRAS_SOURCE" = "published" ]; then \
--python python3; \
fi
RUN mkdir -p /app/.cache/npm && \
prisma generate --schema=./schema.prisma && \
prisma --version && \
prisma migrate diff --from-empty --to-schema-datamodel ./schema.prisma --script > /dev/null 2>&1 || true
RUN prisma generate --schema=./schema.prisma
RUN sed -i 's/\r$//' docker/entrypoint.sh && chmod +x docker/entrypoint.sh && \
sed -i 's/\r$//' docker/prod_entrypoint.sh && chmod +x docker/prod_entrypoint.sh
@ -137,33 +101,11 @@ WORKDIR /app
USER root
RUN for i in 1 2 3; do \
apk upgrade --no-cache && break || sleep 5; \
apk upgrade --no-cache && break || sleep 5; \
done && \
for i in 1 2 3; do \
apk add --no-cache python3 bash openssl tzdata nodejs npm supervisor libsndfile && break || sleep 5; \
done && \
apk upgrade --no-cache nodejs && \
npm install -g npm@11.12.1 tar@7.5.11 glob@11.1.0 @isaacs/brace-expansion@5.0.1 minimatch@10.2.4 diff@8.0.3 && \
GLOBAL="$(npm root -g)" && \
find "$GLOBAL/npm" -type d -name "tar" -path "*/node_modules/tar" | while read d; do \
rm -rf "$d" && cp -rL "$GLOBAL/tar" "$d"; \
done && \
find "$GLOBAL/npm" -type d -name "glob" -path "*/node_modules/glob" | while read d; do \
rm -rf "$d" && cp -rL "$GLOBAL/glob" "$d"; \
done && \
find "$GLOBAL/npm" -type d -name "brace-expansion" -path "*/node_modules/@isaacs/brace-expansion" | while read d; do \
rm -rf "$d" && cp -rL "$GLOBAL/@isaacs/brace-expansion" "$d"; \
done && \
find "$GLOBAL/npm" -type d -name "minimatch" -path "*/node_modules/minimatch" | while read d; do \
rm -rf "$d" && cp -rL "$GLOBAL/minimatch" "$d"; \
done && \
find "$GLOBAL/npm" -type d -name "diff" -path "*/node_modules/diff" | while read d; do \
rm -rf "$d" && cp -rL "$GLOBAL/diff" "$d"; \
done && \
find /usr/local/lib /usr/lib -path "*/node_modules/npm/package.json" -exec \
sed -i 's/"tar": "\^7\.5\.[0-9]*"/"tar": "^7.5.10"/g; s/"minimatch": "\^10\.[0-9.]*"/"minimatch": "^10.2.4"/g' {} + 2>/dev/null && \
npm cache clean --force && \
{ apk del --no-cache npm 2>/dev/null || true; }
apk add --no-cache python3 bash openssl tzdata supervisor libsndfile nodejs && break || sleep 5; \
done
COPY --from=builder /app /app
COPY --from=builder /var/lib/litellm/ui /var/lib/litellm/ui
@ -179,15 +121,10 @@ ENV PATH="/app/.venv/bin:${PATH}" \
PRISMA_SKIP_POSTINSTALL_GENERATE=1 \
PRISMA_HIDE_UPDATE_MESSAGE=1 \
PRISMA_ENGINES_CHECKSUM_IGNORE_MISSING=1 \
NPM_CONFIG_CACHE=/app/.cache/npm \
NPM_CONFIG_PREFER_OFFLINE=true \
PRISMA_OFFLINE_MODE=true
RUN sed -i 's/\r$//' docker/entrypoint.sh && \
sed -i 's/\r$//' docker/prod_entrypoint.sh && \
chmod +x docker/entrypoint.sh docker/prod_entrypoint.sh && \
mkdir -p /nonexistent /.npm /var/lib/litellm/assets /var/lib/litellm/ui /tmp/.npm && \
chown -R nobody:nogroup /app /var/lib/litellm/ui /var/lib/litellm/assets /nonexistent /.npm /tmp/.npm && \
RUN mkdir -p /nonexistent /var/lib/litellm/assets /var/lib/litellm/ui && \
chown -R nobody:nogroup /app /var/lib/litellm/ui /var/lib/litellm/assets /nonexistent && \
PRISMA_PATH=$(python -c "import os, prisma; print(os.path.dirname(prisma.__file__))") && \
chown -R nobody:nogroup "$PRISMA_PATH" && \
LITELLM_PKG_MIGRATIONS_PATH="$(python -c 'import os, litellm_proxy_extras; print(os.path.dirname(litellm_proxy_extras.__file__))' 2>/dev/null || echo '')/migrations" && \

View file

@ -2,6 +2,8 @@
This is a runbook for creating and running database migrations for the LiteLLM proxy. For use for litellm engineers only.
> **AI AGENTS / ASSISTANTS:** If the script refuses with either a "STALE BRANCH" or "DESTRUCTIVE MIGRATION DETECTED" error, **do NOT** bypass it on your own (no `git rebase`, no `--skip-freshness-check`, no `--allow-destructive`). Surface the error to the human operator and wait for their explicit confirmation. See the [Branch freshness](#branch-freshness-check) and [Destructive migrations](#destructive-migrations-drop-column--drop-table) sections below.
## Step 0: Sync All `schema.prisma` Files
Before doing anything else, make sure all `schema.prisma` files in the repo are in sync. There are multiple copies that must match:
@ -46,10 +48,50 @@ uv run --with testing.postgresql python ci_cd/run_migration.py "your_migration_n
## What It Does
1. Creates temp PostgreSQL DB
2. Applies existing migrations
3. Compares with `schema.prisma`
4. Generates new migration if changes found
1. **Verifies the current branch is up to date with `origin/litellm_internal_staging`** (see [Branch freshness](#branch-freshness-check))
2. Creates temp PostgreSQL DB
3. Applies existing migrations
4. Compares with `schema.prisma`
5. Generates new migration if changes found
6. **Refuses to write if the generated SQL contains `DROP COLUMN`, `DROP TABLE`, or `DROP INDEX`**, unless `--allow-destructive` is passed
## Branch Freshness Check
Before generating anything, `run_migration.py` runs `git fetch origin <base>` and refuses to proceed if `HEAD` is behind `origin/<base>`. Default base is `litellm_internal_staging` (the branch PRs target). A previous incident saw a stale branch silently drop production columns; freshness is the first-line defense.
Flags:
- `--base-branch <name>` — check against a different base (e.g. `main`). Default is `litellm_internal_staging`.
- `--skip-freshness-check` — bypass entirely. Only for intentional migrations against an older base.
When the guard fires:
1. Update your branch:
```bash
git fetch origin && git rebase origin/litellm_internal_staging
# or git merge origin/litellm_internal_staging — whichever matches your workflow
```
2. Re-run `run_migration.py`.
> **AI AGENTS / ASSISTANTS:** Do **not** auto-rebase or auto-pass `--skip-freshness-check`. A bad conflict resolution during rebase can itself drop the human's in-progress schema edits. Surface the error and wait for explicit confirmation.
## Destructive Migrations (DROP COLUMN / DROP TABLE / DROP INDEX)
If the generated diff contains `DROP COLUMN`, `DROP TABLE`, or `DROP INDEX`, `run_migration.py` exits non-zero and refuses to write the migration file. A previous incident saw newly-added columns silently dropped by a stale branch and merged to main — this guard exists to prevent a repeat.
When the guard fires:
1. Run `git fetch origin && git status` — confirm your branch is up to date with the base branch.
2. Re-check all `schema.prisma` files are in sync (Step 0).
3. Review EACH `DROP` statement printed in the error — is it actually intended?
4. Only if the drops are genuinely intentional, re-run with the flag:
```bash
uv run --with testing.postgresql python ci_cd/run_migration.py "your_migration_name" --allow-destructive
```
> **AI AGENTS / ASSISTANTS:** Do **not** automatically re-run the command with `--allow-destructive`. If the guard fires while you are driving the runbook for a human, stop, show them the error, and wait for their explicit confirmation before passing the flag. Auto-passing `--allow-destructive` is the exact failure mode this guard exists to prevent.
## Common Fixes

View file

@ -1502,6 +1502,9 @@ if TYPE_CHECKING:
from .llms.bedrock.messages.invoke_transformations.anthropic_claude3_transformation import (
AmazonAnthropicClaudeMessagesConfig as AmazonAnthropicClaudeMessagesConfig,
)
from .llms.bedrock.messages.mantle_transformation import (
AmazonMantleMessagesConfig as AmazonMantleMessagesConfig,
)
from .llms.together_ai.chat import TogetherAIConfig as TogetherAIConfig
from .llms.nlp_cloud.chat.handler import NLPCloudConfig as NLPCloudConfig
from .llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (

View file

@ -171,6 +171,7 @@ LLM_CONFIG_NAMES = (
"CohereChatConfig",
"AnthropicMessagesConfig",
"AmazonAnthropicClaudeMessagesConfig",
"AmazonMantleMessagesConfig",
"TogetherAIConfig",
"NLPCloudConfig",
"VertexGeminiConfig",
@ -715,6 +716,10 @@ _LLM_CONFIGS_IMPORT_MAP = {
".llms.bedrock.messages.invoke_transformations.anthropic_claude3_transformation",
"AmazonAnthropicClaudeMessagesConfig",
),
"AmazonMantleMessagesConfig": (
".llms.bedrock.messages.mantle_transformation",
"AmazonMantleMessagesConfig",
),
"TogetherAIConfig": (".llms.together_ai.chat", "TogetherAIConfig"),
"NLPCloudConfig": (".llms.nlp_cloud.chat.handler", "NLPCloudConfig"),
"VertexGeminiConfig": (

View file

@ -240,7 +240,7 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
self,
model: str,
request_kwargs: Dict,
messages: Optional[List[Dict[str, str]]] = None,
messages: Optional[List[Dict[str, Any]]] = None,
input: Optional[Union[str, List]] = None,
specific_deployment: Optional[bool] = False,
) -> Optional[PreRoutingHookResponse]:

View file

@ -1615,6 +1615,14 @@ class OpenTelemetry(CustomLogger):
value=response_id,
)
litellm_call_id = standard_logging_payload.get("litellm_call_id")
if litellm_call_id:
self.safe_set_attribute(
span=span,
key="litellm.call_id",
value=litellm_call_id,
)
# The model used to generate the response.
if response_obj and response_obj.get("model"):
self.safe_set_attribute(

View file

@ -5512,6 +5512,8 @@ def get_standard_logging_object_payload(
payload: StandardLoggingPayload = StandardLoggingPayload(
id=str(id),
litellm_call_id=kwargs.get("litellm_call_id")
or litellm_params.get("litellm_call_id"),
trace_id=StandardLoggingPayloadSetup._get_standard_logging_payload_trace_id(
logging_obj=logging_obj,
litellm_params=litellm_params,

View file

@ -34,6 +34,7 @@ from litellm.types.llms.anthropic import (
)
from litellm.types.llms.openai import (
AllMessageValues,
ChatCompletionRequest,
ChatCompletionToolCallChunk,
ChatCompletionToolParam,
)
@ -67,6 +68,32 @@ class AnthropicMessagesHandler(BaseTranslation):
super().__init__()
self.adapter = LiteLLMAnthropicMessagesAdapter()
def _translate_to_openai(self, data: dict) -> ChatCompletionRequest:
"""Translate Anthropic request to OpenAI chat completion format."""
(
chat_completion_compatible_request,
_tool_name_mapping,
) = LiteLLMAnthropicMessagesAdapter().translate_anthropic_to_openai(
anthropic_message_request=cast(AnthropicMessagesRequest, data.copy())
)
return chat_completion_compatible_request
def get_structured_messages(self, data: dict) -> Optional[List[AllMessageValues]]:
"""
Convert Anthropic messages request data to OpenAI-spec structured messages.
Uses the Anthropic-to-OpenAI adapter to translate message format.
"""
messages = data.get("messages")
if messages is None:
return None
chat_completion_compatible_request = self._translate_to_openai(data)
result = cast(
List[AllMessageValues],
chat_completion_compatible_request.get("messages", []),
)
return result if result else None
async def process_input_messages(
self,
data: dict,
@ -82,13 +109,7 @@ class AnthropicMessagesHandler(BaseTranslation):
skip_system = effective_skip_system_message_for_guardrail(guardrail_to_apply)
(
chat_completion_compatible_request,
_tool_name_mapping,
) = LiteLLMAnthropicMessagesAdapter().translate_anthropic_to_openai(
# Use a shallow copy to avoid mutating request data (pop on litellm_metadata).
anthropic_message_request=cast(AnthropicMessagesRequest, data.copy())
)
chat_completion_compatible_request = self._translate_to_openai(data)
structured_messages = cast(
List[AllMessageValues],
@ -103,8 +124,6 @@ class AnthropicMessagesHandler(BaseTranslation):
chat_completion_compatible_request.get("tools", [])
)
task_mappings: List[Tuple[int, Optional[int]]] = []
# Track (message_index, content_index) for each text
# content_index is None for string content, int for list content
# Step 1: Extract all text content and images
for msg_idx, message in enumerate(messages):

View file

@ -5,6 +5,7 @@ if TYPE_CHECKING:
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.llms.openai import AllMessageValues
class BaseTranslation(ABC):
@ -101,6 +102,16 @@ class BaseTranslation(ABC):
"""
return responses_so_far
def get_structured_messages(self, data: dict) -> Optional[List["AllMessageValues"]]:
"""
Convert request data to OpenAI-spec structured messages.
Override in subclasses for format-specific conversion.
Returns None if no convertible content is found.
"""
return None
def extract_request_tool_names(self, data: dict) -> List[str]:
"""
Extract tool names from the request body for allowlist/policy checks.

View file

@ -0,0 +1,91 @@
"""
Transformation for Bedrock Mantle (Claude Mythos Preview)
https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-anthropic-claude-mythos-preview.html
The bedrock-mantle endpoint uses the Anthropic Messages API format but is served
at a different endpoint (bedrock-mantle.{region}.api.aws) with AWS SigV4 auth.
"""
from typing import TYPE_CHECKING, Any, List, Optional
from litellm.llms.bedrock.chat.invoke_transformations.anthropic_claude3_transformation import (
AmazonAnthropicClaudeConfig,
)
from litellm.types.llms.openai import AllMessageValues
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
LiteLLMLoggingObj = _LiteLLMLoggingObj
else:
LiteLLMLoggingObj = Any
MANTLE_ENDPOINT_TEMPLATE = "https://bedrock-mantle.{region}.api.aws/v1/messages"
class AmazonMantleConfig(AmazonAnthropicClaudeConfig):
"""
Config for the bedrock-mantle endpoint (Claude Mythos Preview).
Uses the Anthropic Messages API format with AWS SigV4 auth, but at a
different endpoint from bedrock-runtime. Model ID goes in the request body.
Usage: model="bedrock/mantle/anthropic.claude-mythos-preview"
"""
def get_complete_url(
self,
api_base: Optional[str],
api_key: Optional[str],
model: str,
optional_params: dict,
litellm_params: dict,
stream: Optional[bool] = None,
) -> str:
region = self._get_aws_region_name(optional_params=optional_params, model=model)
return MANTLE_ENDPOINT_TEMPLATE.format(region=region)
def transform_request(
self,
model: str,
messages: List[AllMessageValues],
optional_params: dict,
litellm_params: dict,
headers: dict,
) -> dict:
# Strip the "mantle/" routing prefix to get the real model ID
model_id = model.replace("mantle/", "", 1)
request = self._build_bedrock_anthropic_request_base(
model=model_id,
messages=messages,
optional_params=optional_params,
litellm_params=litellm_params,
headers=headers,
)
# The parent strips "model" from the body (Invoke API puts it in URL).
# The mantle endpoint (Messages API) requires "model" in the body.
request["model"] = model_id
return request
async def async_transform_request(
self,
model: str,
messages: List[AllMessageValues],
optional_params: dict,
litellm_params: dict,
headers: dict,
) -> dict:
model_id = model.replace("mantle/", "", 1)
request = self._build_bedrock_anthropic_request_base(
model=model_id,
messages=messages,
optional_params=optional_params,
litellm_params=litellm_params,
headers=headers,
)
await self._async_convert_document_url_sources_to_base64(request)
request["model"] = model_id
return request

View file

@ -696,6 +696,7 @@ class BedrockModelInfo(BaseLLMModelInfo):
"agentcore",
"async_invoke",
"openai",
"mantle",
]:
"""
Get the bedrock route for the given model.
@ -710,6 +711,7 @@ class BedrockModelInfo(BaseLLMModelInfo):
"agentcore",
"async_invoke",
"openai",
"mantle",
],
] = {
"invoke/": "invoke",
@ -719,6 +721,7 @@ class BedrockModelInfo(BaseLLMModelInfo):
"agentcore/": "agentcore",
"async_invoke/": "async_invoke",
"openai/": "openai",
"mantle/": "mantle",
}
# Check explicit routes first
@ -770,6 +773,13 @@ class BedrockModelInfo(BaseLLMModelInfo):
"""
return "agentcore/" in model
@staticmethod
def _explicit_mantle_route(model: str) -> bool:
"""
Check if the model is an explicit mantle route (bedrock-mantle endpoint).
"""
return "mantle/" in model
@staticmethod
def _explicit_converse_like_route(model: str) -> bool:
"""
@ -809,6 +819,16 @@ class BedrockModelInfo(BaseLLMModelInfo):
if BedrockModelInfo._explicit_converse_route(model):
return None
#########################################################
# Mantle route uses the bedrock-mantle endpoint (not bedrock-runtime)
#########################################################
if BedrockModelInfo._explicit_mantle_route(model):
from litellm.llms.bedrock.messages.mantle_transformation import (
AmazonMantleMessagesConfig,
)
return AmazonMantleMessagesConfig()
#########################################################
# This goes through litellm.AmazonAnthropicClaude3MessagesConfig()
# Since bedrock Invoke supports Native Anthropic Messages API
@ -855,6 +875,12 @@ def get_bedrock_chat_config(model: str):
)
return AmazonAgentCoreConfig()
elif bedrock_route == "mantle":
from litellm.llms.bedrock.chat.mantle.transformation import (
AmazonMantleConfig,
)
return AmazonMantleConfig()
# Handle provider-specific configs
if bedrock_invoke_provider == "amazon":

View file

@ -34,6 +34,7 @@ from litellm.llms.bedrock.common_utils import (
remove_custom_field_from_tools,
)
from litellm.types.llms.anthropic import ANTHROPIC_TOOL_SEARCH_BETA_HEADER
from litellm.types.llms.bedrock import BedrockInvokeAnthropicMessagesRequest
from litellm.types.llms.openai import AllMessageValues
from litellm.types.router import GenericLiteLLMParams
from litellm.types.utils import GenericStreamingChunk
@ -59,6 +60,10 @@ class AmazonAnthropicClaudeMessagesConfig(
DEFAULT_BEDROCK_ANTHROPIC_API_VERSION = "bedrock-2023-05-31"
BEDROCK_INVOKE_ALLOWED_TOP_LEVEL_FIELDS = frozenset(
BedrockInvokeAnthropicMessagesRequest.__annotations__.keys()
)
def __init__(self, **kwargs):
BaseAnthropicMessagesConfig.__init__(self, **kwargs)
AmazonInvokeConfig.__init__(self, **kwargs)
@ -500,10 +505,6 @@ class AmazonAnthropicClaudeMessagesConfig(
anthropic_messages_request=anthropic_messages_request,
)
# 5b. Strip `output_config` — Bedrock Invoke doesn't support it
# Fixes: https://github.com/BerriAI/litellm/issues/22797
anthropic_messages_request.pop("output_config", None)
# 5a. Remove `custom` field from tools (Bedrock doesn't support it)
# Claude Code sends `custom: {defer_loading: true}` on tool definitions,
# which causes Bedrock to reject the request with "Extra inputs are not permitted"
@ -550,14 +551,43 @@ class AmazonAnthropicClaudeMessagesConfig(
if "tool-search-tool-2025-10-19" in beta_set:
beta_set.add("tool-examples-2025-10-29")
filtered_auto_betas = filter_and_transform_beta_headers(
beta_headers=list(beta_set - user_beta_set),
provider="bedrock",
filtered_betas = sorted(
filter_and_transform_beta_headers(
beta_headers=list(beta_set),
provider="bedrock",
)
)
filtered_betas = sorted(user_beta_set.union(set(filtered_auto_betas)))
dropped_user_betas = sorted(
b
for b in user_beta_set
if not filter_and_transform_beta_headers([b], provider="bedrock")
)
if dropped_user_betas:
verbose_logger.warning(
"Bedrock Invoke: dropping unsupported anthropic-beta values "
"from client headers: %s. Bedrock has no mapping entry for "
"these; forwarding them would cause a 400.",
dropped_user_betas,
)
if filtered_betas:
anthropic_messages_request["anthropic_beta"] = filtered_betas
# 7. Final safety net: filter top-level fields to the Bedrock Invoke allowlist.
# Catches Anthropic-only extensions (context_management, output_config, speed,
# mcp_servers, ...) and any future additions Claude Code may start sending.
allowed = self.BEDROCK_INVOKE_ALLOWED_TOP_LEVEL_FIELDS
stripped = sorted(k for k in anthropic_messages_request if k not in allowed)
if stripped:
verbose_logger.debug(
"Bedrock Invoke: stripping unsupported top-level request fields: %s",
stripped,
)
anthropic_messages_request = {
k: v for k, v in anthropic_messages_request.items() if k in allowed
}
return anthropic_messages_request
def get_async_streaming_response_iterator(

View file

@ -0,0 +1,69 @@
"""
Transformation for Bedrock Mantle (Claude Mythos Preview) - /messages endpoint
Inherits all Messages API request/response transformations from
AmazonAnthropicClaudeMessagesConfig. Overrides only the URL and model-prefix
stripping that are specific to the bedrock-mantle endpoint.
"""
from typing import TYPE_CHECKING, Any, Dict, List, Optional
from litellm.llms.bedrock.messages.invoke_transformations.anthropic_claude3_transformation import (
AmazonAnthropicClaudeMessagesConfig,
)
from litellm.types.router import GenericLiteLLMParams
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
LiteLLMLoggingObj = _LiteLLMLoggingObj
else:
LiteLLMLoggingObj = Any
MANTLE_ENDPOINT_TEMPLATE = "https://bedrock-mantle.{region}.api.aws/v1/messages"
class AmazonMantleMessagesConfig(AmazonAnthropicClaudeMessagesConfig):
"""
Config for the bedrock-mantle /messages endpoint (Claude Mythos Preview).
The mantle endpoint uses the Anthropic Messages API format and requires the
model ID in the request body (unlike Bedrock Invoke which puts it in the URL).
"""
def get_complete_url(
self,
api_base: Optional[str],
api_key: Optional[str],
model: str,
optional_params: dict,
litellm_params: dict,
stream: Optional[bool] = None,
) -> str:
region = self._get_aws_region_name(optional_params=optional_params, model=model)
return MANTLE_ENDPOINT_TEMPLATE.format(region=region)
def transform_anthropic_messages_request(
self,
model: str,
messages: List[Dict],
anthropic_messages_optional_request_params: Dict,
litellm_params: GenericLiteLLMParams,
headers: dict,
) -> Dict:
# Strip "mantle/" routing prefix to get the real model ID
model_id = model.replace("mantle/", "", 1)
request = super().transform_anthropic_messages_request(
model=model_id,
messages=messages,
anthropic_messages_optional_request_params=anthropic_messages_optional_request_params,
litellm_params=litellm_params,
headers=headers,
)
# Parent (AmazonAnthropicClaudeMessagesConfig) removes "model" from the
# body (Bedrock Invoke puts model in the URL). The mantle endpoint
# (Messages API) requires "model" in the request body.
request["model"] = model_id
return request

View file

@ -48,6 +48,17 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
Methods can be overridden to customize behavior for different message formats.
"""
def get_structured_messages(self, data: dict) -> Optional[List[AllMessageValues]]:
"""
Convert chat completions request data to OpenAI-spec structured messages.
Messages are already in OpenAI format, so this is a simple extraction.
"""
messages = data.get("messages")
if messages is None:
return None
return cast(List[AllMessageValues], messages)
async def process_input_messages(
self,
data: dict,
@ -68,9 +79,6 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
tool_calls_to_check: List[ChatCompletionToolParam] = []
text_task_mappings: List[Tuple[int, Optional[int]]] = []
tool_call_task_mappings: List[Tuple[int, int]] = []
# text_task_mappings: Track (message_index, content_index) for each text
# content_index is None for string content, int for list content
# tool_call_task_mappings: Track (message_index, tool_call_index) for each tool call
# Step 1: Extract all text content, images, and tool calls
for msg_idx, message in enumerate(messages):
@ -92,12 +100,12 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
inputs["images"] = images_to_check
if tool_calls_to_check:
inputs["tool_calls"] = tool_calls_to_check # type: ignore
if messages:
msg_list = cast(List[AllMessageValues], messages)
structured_messages = self.get_structured_messages(data)
if structured_messages:
inputs["structured_messages"] = (
openai_messages_without_system(msg_list)
openai_messages_without_system(structured_messages)
if skip_system
else msg_list
else structured_messages
)
# Pass tools (function definitions) to the guardrail
tools = data.get("tools")

View file

@ -43,6 +43,7 @@ from litellm.responses.litellm_completion_transformation.transformation import (
LiteLLMCompletionResponsesConfig,
)
from litellm.types.llms.openai import (
AllMessageValues,
ChatCompletionToolCallChunk,
ChatCompletionToolParam,
)
@ -70,6 +71,24 @@ class OpenAIResponsesHandler(BaseTranslation):
Methods can be overridden to customize behavior for different message formats.
"""
def get_structured_messages(self, data: dict) -> Optional[List[AllMessageValues]]:
"""
Convert Responses API request data to OpenAI-spec structured messages.
Transforms `input` (string or ResponseInputParam) and optional
`instructions` into chat completion messages.
"""
input_data = data.get("input")
if input_data is None:
return None
messages = (
LiteLLMCompletionResponsesConfig.transform_responses_api_input_to_messages(
input=input_data,
responses_api_request=data,
)
)
return cast(List[AllMessageValues], messages) if messages else None
async def process_input_messages(
self,
data: dict,
@ -86,12 +105,7 @@ class OpenAIResponsesHandler(BaseTranslation):
if input_data is None:
return data
structured_messages = (
LiteLLMCompletionResponsesConfig.transform_responses_api_input_to_messages(
input=input_data,
responses_api_request=data,
)
)
structured_messages = self.get_structured_messages(data)
# Handle simple string input
if isinstance(input_data, str):

View file

@ -1121,14 +1121,14 @@ async def _db_health_readiness_check():
return db_health_cache
except Exception as e:
db_health_cache = {"status": "disconnected", "last_updated": datetime.now()}
PrismaDBExceptionHandler.handle_db_exception(e)
if PrismaDBExceptionHandler.is_database_transport_error(e):
try:
verbose_proxy_logger.warning(
"_db_health_readiness_check: health_check failed, attempting reconnect"
)
await prisma_client.disconnect()
await prisma_client.connect()
await prisma_client.attempt_db_reconnect(
reason="health_readiness_check"
)
await prisma_client.health_check()
verbose_proxy_logger.info(
"_db_health_readiness_check: reconnect succeeded"

View file

@ -202,6 +202,8 @@ if TYPE_CHECKING:
)
from litellm.router_strategy.adaptive_router.adaptive_router import (
AdaptiveRouter,
from litellm.router_strategy.quality_router.quality_router import (
QualityRouter,
)
Span = Union[_Span, Any]
@ -210,6 +212,7 @@ else:
AutoRouter = Any
ComplexityRouter = Any
AdaptiveRouter = Any
QualityRouter = Any
PreRoutingHookResponse = Any
@ -469,6 +472,7 @@ class Router:
self.auto_routers: Dict[str, "AutoRouter"] = {}
self.complexity_routers: Dict[str, "ComplexityRouter"] = {}
self.adaptive_routers: Dict[str, "AdaptiveRouter"] = {}
self.quality_routers: Dict[str, "QualityRouter"] = {}
# Initialize model_group_alias early since it's used in set_model_list
self.model_group_alias: Dict[str, Union[str, RouterModelGroupAliasItem]] = (
@ -5889,7 +5893,7 @@ class Router:
response = await response
## PROCESS RESPONSE HEADERS
response = await self.set_response_headers(
response=response, model_group=model_group
response=response, model_group=model_group, request_kwargs=kwargs
)
return response
@ -6822,6 +6826,8 @@ class Router:
return False # This is handled by complexity_router
if litellm_params.model.startswith("auto_router/adaptive_router"):
return False # This is handled by adaptive_router
if litellm_params.model.startswith("auto_router/quality_router"):
return False # This is handled by quality_router
if litellm_params.model.startswith("auto_router/"):
return True
return False
@ -7045,6 +7051,57 @@ class Router:
deployment.model_name,
len(config.available_models),
)
def _is_quality_router_deployment(self, litellm_params: LiteLLM_Params) -> bool:
"""
Check if the deployment is a quality-router deployment.
Returns True if the litellm_params model starts with "auto_router/quality_router".
"""
if litellm_params.model.startswith("auto_router/quality_router"):
return True
return False
def init_quality_router_deployment(self, deployment: Deployment):
"""
Initialize the quality-router deployment.
Resolves the default model from either `quality_router_default_model` or
`quality_router_config["default_model"]`, then instantiates the
QualityRouter and stores it in `self.quality_routers`.
"""
# Import here to mirror the AutoRouter / ComplexityRouter init pattern
# and avoid circular imports.
from litellm.router_strategy.quality_router.quality_router import (
QualityRouter,
)
quality_router_config: Optional[dict] = (
deployment.litellm_params.quality_router_config
)
default_model: Optional[str] = (
deployment.litellm_params.quality_router_default_model
)
if default_model is None and quality_router_config:
default_model = quality_router_config.get("default_model")
if default_model is None:
raise ValueError(
"quality_router_default_model is required for quality-router deployments, "
"or set default_model in quality_router_config. Please configure it in the litellm_params"
)
quality_router: QualityRouter = QualityRouter(
model_name=deployment.model_name,
default_model=default_model,
litellm_router_instance=self,
quality_router_config=quality_router_config,
)
if deployment.model_name in self.quality_routers:
raise ValueError(
f"Quality-router deployment {deployment.model_name} already exists. Please use a different model name."
)
self.quality_routers[deployment.model_name] = quality_router
def deployment_is_active_for_environment(self, deployment: Deployment) -> bool:
"""
@ -7092,6 +7149,11 @@ class Router:
self.model_id_to_deployment_index_map = {} # Reset the index
self.model_name_to_deployment_indices = {} # Reset the model_name index
self.team_model_to_deployment_indices = {} # Reset the team_model index
# Reset per-strategy router registries so hot-reload doesn't leave
# stale routers pointing at the old model_list.
self.quality_routers = {}
self.complexity_routers = {}
self.auto_routers = {}
self._invalidate_model_group_info_cache()
self._invalidate_access_groups_cache()
# we add api_base/api_key each model so load balancing between azure/gpt on api_base1 and api_base2 works
@ -7274,6 +7336,11 @@ class Router:
# set_model_list() because their init needs visibility into the OTHER
# deployments listed in `available_models` (which may not yet have
# been processed when this one is created).
#########################################################
# Check if this is a quality-router deployment
#########################################################
if self._is_quality_router_deployment(litellm_params=deployment.litellm_params):
self.init_quality_router_deployment(deployment=deployment)
return deployment
@ -8278,7 +8345,10 @@ class Router:
return returned_dict
async def set_response_headers(
self, response: Any, model_group: Optional[str] = None
self,
response: Any,
model_group: Optional[str] = None,
request_kwargs: Optional[dict] = None,
) -> Any:
"""
Add the most accurate rate limit headers for a given model response.
@ -8299,6 +8369,45 @@ class Router:
additional_headers = response._hidden_params["additional_headers"] # type: ignore
# Lift QualityRouter routing decision into response headers for
# transparency. The decision is stashed in request_kwargs.metadata
# by QualityRouter.async_pre_routing_hook.
metadata = (
(request_kwargs.get("metadata") or {})
if isinstance(request_kwargs, dict)
else {}
)
decision = (
metadata.get("quality_router_decision")
if isinstance(metadata, dict)
else None
)
if isinstance(decision, dict):
# Only emit headers for fields that have a meaningful value.
# `complexity_tier` and `matched_keyword` are mutually exclusive
# (the keyword path short-circuits classification), so each
# request emits one or the other but not both.
if decision.get("routed_model") is not None:
additional_headers["x-litellm-quality-router-model"] = str(
decision["routed_model"]
)
if decision.get("quality_tier") is not None:
additional_headers["x-litellm-quality-router-tier"] = str(
decision["quality_tier"]
)
if decision.get("routed_via") is not None:
additional_headers["x-litellm-quality-router-via"] = str(
decision["routed_via"]
)
if decision.get("matched_keyword") is not None:
additional_headers["x-litellm-quality-router-keyword"] = str(
decision["matched_keyword"]
)
if decision.get("complexity_tier") is not None:
additional_headers["x-litellm-quality-router-complexity"] = str(
decision["complexity_tier"]
)
if (
"x-ratelimit-remaining-tokens" not in additional_headers
and "x-ratelimit-remaining-requests" not in additional_headers
@ -8843,8 +8952,6 @@ class Router:
and self.routing_strategy == "latency-based-routing"
):
_settings_to_return[var] = self.lowestlatency_logger.routing_args.json()
elif var == "routing_strategy_args":
_settings_to_return[var] = None
return _settings_to_return
def update_settings(self, **kwargs):
@ -9755,7 +9862,7 @@ class Router:
self,
model: str,
request_kwargs: Dict,
messages: Optional[List[Dict[str, str]]] = None,
messages: Optional[List[Dict[str, Any]]] = None,
input: Optional[Union[str, List]] = None,
specific_deployment: Optional[bool] = False,
) -> Optional[PreRoutingHookResponse]:
@ -9794,6 +9901,10 @@ class Router:
adaptive_router = self.adaptive_routers.get(model)
if adaptive_router is not None:
return await adaptive_router.async_pre_routing_hook(
# Check if any quality-router should be used
#########################################################
if model in self.quality_routers:
return await self.quality_routers[model].async_pre_routing_hook(
model=model,
request_kwargs=request_kwargs,
messages=messages,

View file

@ -82,11 +82,34 @@ class AutoRouter(CustomLogger):
)
return auto_router_routes
@staticmethod
def _extract_text_from_messages(messages: List[Dict[str, Any]]) -> str:
"""
Extract text content from the last user message for routing.
Handles tool-call conversations (where the last message may be an
assistant or tool message with non-string content) and multimodal
messages (where content is a list of content blocks).
"""
for msg in reversed(messages):
if msg.get("role") == "user":
content = msg.get("content")
if content is None:
return ""
if isinstance(content, list):
return " ".join(
block.get("text", "")
for block in content
if isinstance(block, dict) and block.get("type") == "text"
)
return str(content)
return ""
async def async_pre_routing_hook(
self,
model: str,
request_kwargs: Dict,
messages: Optional[List[Dict[str, str]]] = None,
messages: Optional[List[Dict[str, Any]]] = None,
input: Optional[Union[str, List]] = None,
specific_deployment: Optional[bool] = False,
) -> Optional["PreRoutingHookResponse"]:
@ -120,8 +143,7 @@ class AutoRouter(CustomLogger):
auto_sync=self.auto_sync_value,
)
user_message: Dict[str, str] = messages[-1]
message_content: str = user_message.get("content", "")
message_content = self._extract_text_from_messages(messages)
route_choice: Optional[Union[RouteChoice, List[RouteChoice]]] = self.routelayer(
text=message_content
)

View file

@ -332,45 +332,68 @@ class ComplexityRouter(CustomLogger):
f"No model configured for tier {tier_key} and no default_model set"
)
async def async_pre_routing_hook(
def _resolve_messages(
self,
model: str,
messages: Optional[List[Dict[str, Any]]],
request_kwargs: Dict,
messages: Optional[List[Dict[str, Any]]] = None,
input: Optional[Union[str, List]] = None,
specific_deployment: Optional[bool] = False,
) -> Optional["PreRoutingHookResponse"]:
) -> Optional[List[Dict[str, Any]]]:
"""
Pre-routing hook called before the routing decision.
Resolve messages from the request, converting from other formats if needed.
Classifies the request by complexity and returns the appropriate model.
Args:
model: The original model name requested.
request_kwargs: The request kwargs.
messages: The messages in the request.
input: Optional input for embeddings.
specific_deployment: Whether a specific deployment was requested.
Returns:
PreRoutingHookResponse with the routed model, or None if no routing needed.
Uses the guardrail translation handler dispatch to convert Responses API
``input`` (or other non-chat-completions formats) into OpenAI-spec messages.
"""
from litellm.types.router import PreRoutingHookResponse
if messages:
return messages
if messages is None or len(messages) == 0:
verbose_router_logger.debug(
"ComplexityRouter: No messages provided, skipping routing"
)
return None
from litellm.litellm_core_utils.api_route_to_call_types import (
get_call_types_for_route,
)
from litellm.llms import load_guardrail_translation_mappings
from litellm.types.utils import CallTypes
# Extract the last user message and the last system prompt
mappings = load_guardrail_translation_mappings()
call_type: Optional[CallTypes] = None
# 1. Try route-based inference from proxy metadata
route = request_kwargs.get("litellm_metadata", {}).get(
"user_api_key_request_route"
)
if route:
call_types_list = get_call_types_for_route(route)
if call_types_list:
for ct in call_types_list:
if ct in mappings:
call_type = ct
break
# 2. Fallback: try each mapped handler until one produces messages
handlers_to_try: List[Any] = []
if call_type is not None and call_type in mappings:
handlers_to_try.append(mappings[call_type]())
else:
handlers_to_try.extend(handler_cls() for handler_cls in mappings.values())
for handler in handlers_to_try:
structured = handler.get_structured_messages(request_kwargs)
if structured:
return [
msg if isinstance(msg, dict) else msg.model_dump() # type: ignore
for msg in structured
]
return None
@staticmethod
def _extract_user_message_and_system_prompt(
messages: List[Dict[str, Any]],
) -> Tuple[Optional[str], Optional[str]]:
"""Extract the last user message text and last system prompt from messages."""
user_message: Optional[str] = None
system_prompt: Optional[str] = None
for msg in reversed(messages):
role = msg.get("role", "")
content = msg.get("content") or ""
# content may be a list of content parts (e.g. [{"type": "text", "text": "..."}])
if isinstance(content, list):
text_parts = [
part.get("text", "")
@ -383,6 +406,52 @@ class ComplexityRouter(CustomLogger):
user_message = content
elif role == "system" and system_prompt is None:
system_prompt = content
if user_message is not None and system_prompt is not None:
break
return user_message, system_prompt
async def async_pre_routing_hook(
self,
model: str,
request_kwargs: Dict,
messages: Optional[List[Dict[str, Any]]] = None,
input: Optional[Union[str, List]] = None,
specific_deployment: Optional[bool] = False,
) -> Optional["PreRoutingHookResponse"]:
"""
Pre-routing hook called before the routing decision.
Classifies the request by complexity and returns the appropriate model.
Supports chat completions (messages), Responses API (input), and other
formats via the guardrail translation handler dispatch.
Args:
model: The original model name requested.
request_kwargs: The request kwargs.
messages: The messages in the request.
input: Optional input for Responses API or embeddings.
specific_deployment: Whether a specific deployment was requested.
Returns:
PreRoutingHookResponse with the routed model, or None if no routing needed.
"""
from litellm.types.router import PreRoutingHookResponse
resolved_messages = self._resolve_messages(messages, request_kwargs)
if not resolved_messages:
verbose_router_logger.debug(
"ComplexityRouter: No messages could be resolved, skipping routing"
)
return None
# Determine whether the original request used messages directly
has_original_messages = messages is not None and len(messages) > 0
user_message, system_prompt = self._extract_user_message_and_system_prompt(
resolved_messages
)
if user_message is None:
verbose_router_logger.debug(
@ -391,13 +460,10 @@ class ComplexityRouter(CustomLogger):
return PreRoutingHookResponse(
model=self.config.default_model
or self.get_model_for_tier(ComplexityTier.MEDIUM),
messages=messages,
messages=messages if has_original_messages else None,
)
# Classify the request
tier, score, signals = self.classify(user_message, system_prompt)
# Get the model for this tier
routed_model = self.get_model_for_tier(tier)
verbose_router_logger.info(
@ -407,5 +473,5 @@ class ComplexityRouter(CustomLogger):
return PreRoutingHookResponse(
model=routed_model,
messages=messages,
messages=messages if has_original_messages else None,
)

View file

@ -0,0 +1,21 @@
"""
Quality-tier auto-router.
Re-uses the ComplexityRouter's classification to decide a request's complexity,
then maps that complexity to an admin-configured quality tier and resolves the
target model from each candidate's `model_info.litellm_routing_preferences`.
"""
from .config import (
DEFAULT_COMPLEXITY_TO_QUALITY,
QualityRouterConfig,
RoutingPreferences,
)
from .quality_router import QualityRouter
__all__ = [
"QualityRouter",
"QualityRouterConfig",
"RoutingPreferences",
"DEFAULT_COMPLEXITY_TO_QUALITY",
]

View file

@ -0,0 +1,74 @@
"""
Configuration models for the QualityRouter.
"""
from typing import Dict, List, Optional
from pydantic import BaseModel, ConfigDict, Field
# Default mapping from ComplexityTier name (string) to quality tier (int).
# Higher tier = higher capability requirement.
DEFAULT_COMPLEXITY_TO_QUALITY: Dict[str, int] = {
"SIMPLE": 1,
"MEDIUM": 2,
"COMPLEX": 3,
"REASONING": 4,
}
class QualityRouterConfig(BaseModel):
"""Configuration for the QualityRouter."""
available_models: List[str] = Field(
default_factory=list,
description=(
"List of candidate model names this router may route to. Each model "
"must declare its quality_tier in model_info.litellm_routing_preferences."
),
)
default_model: Optional[str] = Field(
default=None,
description="Fallback model when no quality tier resolves.",
)
complexity_to_quality: Dict[str, int] = Field(
default_factory=lambda: DEFAULT_COMPLEXITY_TO_QUALITY.copy(),
description="Mapping from ComplexityTier name to quality tier (int).",
)
model_config = ConfigDict(extra="allow")
class RoutingPreferences(BaseModel):
"""Per-deployment routing preferences declared on model_info."""
quality_tier: int = Field(
...,
description="The quality tier this deployment satisfies.",
)
keywords: List[str] = Field(
default_factory=list,
description=(
"Substring keywords (case-insensitive) that, when present in the "
"user message, route the request to this deployment. See `order` "
"for explicit collision handling, otherwise ties fall through to "
"(highest quality_tier, then cheapest model_info.input_cost_per_token)."
),
)
order: Optional[int] = Field(
default=None,
description=(
"Explicit priority used to break ties between deployments at the "
"same quality tier. Lower values win. Applies both to keyword "
"collisions and to picking between multiple deployments at the "
"same quality_tier. Tiebreak order is "
"(quality_tier DESC, order ASC, input_cost_per_token ASC, "
"model_name ASC) — quality always wins first, then explicit "
"order, then price."
),
)
model_config = ConfigDict(extra="allow")

View file

@ -0,0 +1,446 @@
"""
Quality-tier Auto Router.
Routes a request to a model at a target quality tier. The quality tier is
inferred by re-using the existing ComplexityRouter's classification, then
mapped through an admin-configured `complexity_to_quality` table. Each
candidate model declares its own `quality_tier` in
`model_info.litellm_routing_preferences`.
Optional keyword override: deployments may also declare `keywords` in
`litellm_routing_preferences`. If any declared keyword appears in the user
message (case-insensitive substring match), the router short-circuits the
complexity-classification flow and routes to the matching deployment. When
multiple deployments match, ties are broken by (highest quality_tier first,
then cheapest `model_info.input_cost_per_token`).
"""
import math
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union
from litellm._logging import verbose_router_logger
from litellm.integrations.custom_logger import CustomLogger
from litellm.router_strategy.complexity_router.complexity_router import (
ComplexityRouter,
)
from .config import QualityRouterConfig, RoutingPreferences
if TYPE_CHECKING:
from litellm.router import Router
from litellm.types.router import PreRoutingHookResponse
else:
Router = Any
PreRoutingHookResponse = Any
class QualityRouter(CustomLogger):
"""
Routes requests to a model at a target quality tier, with an optional
keyword override.
"""
def __init__(
self,
model_name: str,
litellm_router_instance: "Router",
default_model: Optional[str] = None,
quality_router_config: Optional[Dict[str, Any]] = None,
):
self.model_name = model_name
self.litellm_router_instance = litellm_router_instance
if quality_router_config:
self.config = QualityRouterConfig(**quality_router_config)
else:
self.config = QualityRouterConfig()
# Explicit default_model arg overrides anything in the config dict.
if default_model:
self.config.default_model = default_model
# Internal scorer — re-use the existing rule-based classifier.
self._scorer = ComplexityRouter(
model_name=f"{model_name}::scorer",
litellm_router_instance=litellm_router_instance,
)
# Per-model indices populated alongside the tier index. `_model_keywords`
# stores keywords lowercased so we can substring-match against the
# lowercased user message in O(total-keyword-count). `_model_quality`,
# `_model_cost`, and `_model_order` drive tiebreaking — `_model_order`
# is the explicit priority (lower wins, unset = +inf).
self._model_keywords: Dict[str, List[str]] = {}
self._model_quality: Dict[str, int] = {}
self._model_cost: Dict[str, Optional[float]] = {}
self._model_order: Dict[str, Optional[int]] = {}
# Tier → models index. Built lazily on first access so the QualityRouter
# deployment does NOT need to appear after all its referenced models in
# the config — when `_build_tier_index` runs eagerly in `__init__`, the
# router instance's `model_list` is still being assembled incrementally
# by `_create_deployment`, and any `available_models` defined AFTER the
# router entry in config.yaml would silently be reported as missing.
self._tier_to_models_cache: Optional[Dict[int, List[str]]] = None
verbose_router_logger.debug(
f"QualityRouter initialized for {model_name} with "
f"available_models={self.config.available_models}, "
f"default_model={self.config.default_model}"
)
@property
def _tier_to_models(self) -> Dict[int, List[str]]:
"""Lazy tier→models index; built on first access."""
if self._tier_to_models_cache is None:
self._tier_to_models_cache = self._build_tier_index()
return self._tier_to_models_cache
def _get_routing_preferences(self, deployment: Any) -> Optional[Dict[str, Any]]:
"""
Extract litellm_routing_preferences from a deployment, handling both
dict-shaped and Pydantic-object-shaped deployments.
"""
# Dict-shaped deployment.
if isinstance(deployment, dict):
model_info = deployment.get("model_info") or {}
if isinstance(model_info, dict):
return model_info.get("litellm_routing_preferences")
# Pydantic ModelInfo nested in a dict.
return getattr(model_info, "litellm_routing_preferences", None)
# Pydantic-object deployment.
model_info = getattr(deployment, "model_info", None)
if model_info is None:
return None
if isinstance(model_info, dict):
return model_info.get("litellm_routing_preferences")
return getattr(model_info, "litellm_routing_preferences", None)
def _get_deployment_input_cost(self, deployment: Any) -> Optional[float]:
"""
Extract `input_cost_per_token` from a deployment's model_info.
Returns None when not declared — None is treated as "infinite cost"
for the cheapest-tiebreak ordering, so unpriced models lose ties to
priced ones. (Admins who want a model to win on price must declare it.)
"""
if isinstance(deployment, dict):
model_info = deployment.get("model_info") or {}
else:
model_info = getattr(deployment, "model_info", None) or {}
if isinstance(model_info, dict):
cost = model_info.get("input_cost_per_token")
else:
cost = getattr(model_info, "input_cost_per_token", None)
if cost is None:
return None
try:
return float(cost)
except (TypeError, ValueError):
return None
def _get_deployment_model_name(self, deployment: Any) -> Optional[str]:
"""Extract `model_name` from a dict- or object-shaped deployment."""
if isinstance(deployment, dict):
return deployment.get("model_name")
return getattr(deployment, "model_name", None)
def _build_tier_index(self) -> Dict[int, List[str]]:
"""
Build {quality_tier: [model_name, ...]} for every model in
`available_models`, plus side indices `_model_keywords`,
`_model_quality`, and `_model_cost`. Raises if any listed model is
missing `litellm_routing_preferences`.
"""
model_list = getattr(self.litellm_router_instance, "model_list", None) or []
available = set(self.config.available_models)
# Track which available models we've matched so we can error on missing.
seen: Dict[str, bool] = {name: False for name in available}
tier_to_models: Dict[int, List[str]] = {}
for deployment in model_list:
name = self._get_deployment_model_name(deployment)
if name is None or name not in available:
continue
raw_prefs = self._get_routing_preferences(deployment)
if raw_prefs is None:
raise ValueError(
f"QualityRouter: model '{name}' is listed in available_models "
f"but has no model_info.litellm_routing_preferences"
)
# Validate via the Pydantic model so we get a clear error for
# missing quality_tier, wrong types, etc. This also means
# `RoutingPreferences` is the single source of truth for the
# accepted shape — readers relied on raw dicts before.
try:
if isinstance(raw_prefs, RoutingPreferences):
prefs = raw_prefs
elif isinstance(raw_prefs, dict):
prefs = RoutingPreferences(**raw_prefs)
else:
# A Pydantic object of some other shape — coerce via its dict.
prefs = RoutingPreferences(
**(
raw_prefs.model_dump()
if hasattr(raw_prefs, "model_dump")
else dict(raw_prefs)
)
)
except Exception as e:
raise ValueError(
f"QualityRouter: model '{name}' has invalid "
f"litellm_routing_preferences: {e}"
) from e
tier_int = int(prefs.quality_tier)
tier_to_models.setdefault(tier_int, []).append(name)
self._model_keywords[name] = [str(k).lower() for k in prefs.keywords if k]
self._model_quality[name] = tier_int
self._model_cost[name] = self._get_deployment_input_cost(deployment)
self._model_order[name] = prefs.order
seen[name] = True
missing = [name for name, found in seen.items() if not found]
if missing:
raise ValueError(
f"QualityRouter: the following available_models are not present in "
f"the router's model_list (or are missing routing preferences): {missing}"
)
# Sort each tier's model list so `_resolve_model_for_quality_tier`
# (which picks index [0]) honors (order ASC, cost ASC, name ASC).
# Quality is moot within a single tier; keep parity with the keyword
# tiebreak by ordering on (order, cost, name) here.
for models in tier_to_models.values():
models.sort(key=lambda n: (self._order_key(n), self._cost_key(n), n))
return tier_to_models
def _order_key(self, model_name: str) -> float:
"""`order` lookup as a float — unset becomes +inf so explicit wins."""
order = self._model_order.get(model_name)
return float(order) if order is not None else math.inf
def _cost_key(self, model_name: str) -> float:
"""`input_cost_per_token` as a float — unset becomes +inf."""
cost = self._model_cost.get(model_name)
return float(cost) if cost is not None else math.inf
def _keyword_override(self, user_message: str) -> Optional[Tuple[str, str]]:
"""
Find a deployment whose declared keywords appear in `user_message`.
Returns (model_name, matched_keyword) or None when no keyword matches.
When multiple deployments match, sorts by:
1. quality_tier DESC (best quality always wins first)
2. `order` ASC (explicit priority — unset = +inf so explicit wins
within the same tier)
3. input_cost_per_token ASC (unpriced = +inf so priced wins)
4. model_name ASC (deterministic stability)
"""
# Touch the lazy index so `_model_keywords` / `_model_quality` /
# `_model_cost` / `_model_order` are populated.
_ = self._tier_to_models
text = user_message.lower()
matches: List[Tuple[str, str]] = [] # (model_name, matched_keyword)
for model_name, keywords in self._model_keywords.items():
for kw in keywords:
if kw and kw in text:
matches.append((model_name, kw))
break # one match per model is enough
if not matches:
return None
def sort_key(match: Tuple[str, str]) -> Tuple[int, float, float, str]:
name = match[0]
quality = self._model_quality.get(name, 0)
order_val = self._order_key(name)
cost = self._model_cost.get(name)
cost_val = cost if cost is not None else math.inf
# Negate quality so higher tier sorts first under ASC sort.
return (-quality, order_val, cost_val, name)
matches.sort(key=sort_key)
return matches[0]
def _resolve_model_for_quality_tier(self, tier: int) -> str:
"""
Resolve a quality tier to a concrete model name.
Strategy:
1. Exact tier match → first model registered at that tier.
2. Round UP to the next higher tier that has a model (closer to a
request we might lack capacity for).
3. Round DOWN to the closest lower tier that has a model (degrade
gracefully instead of jumping straight to `default_model`,
which may be off-tier).
4. Fall back to `config.default_model`.
5. Otherwise raise.
"""
tier_index = self._tier_to_models
if tier in tier_index and tier_index[tier]:
return tier_index[tier][0]
# Round up.
higher_tiers = sorted(t for t in tier_index if t > tier)
for t in higher_tiers:
if tier_index[t]:
return tier_index[t][0]
# Round down — closest lower tier first.
lower_tiers = sorted((t for t in tier_index if t < tier), reverse=True)
for t in lower_tiers:
if tier_index[t]:
return tier_index[t][0]
if self.config.default_model:
return self.config.default_model
raise ValueError(
f"QualityRouter: no model available for quality tier {tier} and "
f"no default_model configured"
)
def _stash_decision(
self,
request_kwargs: Optional[Dict[str, Any]],
decision: Dict[str, Any],
) -> None:
"""
Stash the routing decision in request_kwargs.metadata so the Router can
lift it into response headers (`x-litellm-quality-router-*`). The same
dict object flows from here through to `make_call.set_response_headers`.
"""
if request_kwargs is None:
return
metadata = request_kwargs.setdefault("metadata", {})
if isinstance(metadata, dict):
metadata["quality_router_decision"] = decision
async def async_pre_routing_hook(
self,
model: str,
request_kwargs: Dict,
messages: Optional[List[Dict[str, Any]]] = None,
input: Optional[Union[str, List]] = None,
specific_deployment: Optional[bool] = False,
) -> Optional["PreRoutingHookResponse"]:
"""Try keyword override first; fall back to complexity-tier routing."""
from litellm.types.router import PreRoutingHookResponse
if messages is None or len(messages) == 0:
verbose_router_logger.debug(
"QualityRouter: No messages provided, skipping routing"
)
return None
# Extract last user message and last system prompt — same rules as
# ComplexityRouter.async_pre_routing_hook.
user_message: Optional[str] = None
system_prompt: Optional[str] = None
for msg in reversed(messages):
role = msg.get("role", "")
content = msg.get("content") or ""
if isinstance(content, list):
text_parts = [
part.get("text", "")
for part in content
if isinstance(part, dict) and part.get("type") == "text"
]
content = " ".join(text_parts).strip()
if isinstance(content, str) and content:
if role == "user" and user_message is None:
user_message = content
elif role == "system" and system_prompt is None:
system_prompt = content
if user_message is None:
verbose_router_logger.debug(
"QualityRouter: No user message found, routing to default model"
)
if not self.config.default_model:
raise ValueError(
"QualityRouter: no user message and no default_model configured"
)
return PreRoutingHookResponse(
model=self.config.default_model,
messages=messages,
)
# Try keyword override first — it short-circuits complexity classification.
keyword_match = self._keyword_override(user_message)
if keyword_match is not None:
routed_model, matched_keyword = keyword_match
verbose_router_logger.info(
f"QualityRouter: keyword override matched='{matched_keyword}' "
f"routed_model={routed_model} "
f"(quality_tier={self._model_quality.get(routed_model)}, "
f"input_cost_per_token={self._model_cost.get(routed_model)})"
)
self._stash_decision(
request_kwargs,
{
"router_model_name": self.model_name,
"routed_model": routed_model,
"routed_via": "keyword",
"matched_keyword": matched_keyword,
"quality_tier": self._model_quality.get(routed_model),
"complexity_tier": None,
},
)
return PreRoutingHookResponse(
model=routed_model,
messages=messages,
)
# No keyword match → complexity classification flow.
complexity_tier, score, signals = self._scorer.classify(
user_message, system_prompt
)
complexity_name = (
complexity_tier.value
if hasattr(complexity_tier, "value")
else str(complexity_tier)
)
quality_tier = self.config.complexity_to_quality.get(complexity_name)
if quality_tier is None:
raise ValueError(
f"QualityRouter: complexity tier '{complexity_name}' not present "
f"in complexity_to_quality mapping {self.config.complexity_to_quality}"
)
routed_model = self._resolve_model_for_quality_tier(int(quality_tier))
verbose_router_logger.info(
f"QualityRouter: complexity={complexity_name}, score={score:.3f}, "
f"signals={signals}, quality_tier={quality_tier}, "
f"routed_model={routed_model}"
)
self._stash_decision(
request_kwargs,
{
"router_model_name": self.model_name,
"routed_model": routed_model,
"routed_via": "quality_tier",
"matched_keyword": None,
"quality_tier": int(quality_tier),
"complexity_tier": complexity_name,
},
)
return PreRoutingHookResponse(
model=routed_model,
messages=messages,
)

View file

@ -997,3 +997,47 @@ class BedrockToolBlock(TypedDict, total=False):
toolSpec: Optional[ToolSpecBlock]
systemTool: Optional[SystemToolBlock] # For Nova grounding
cachePoint: Optional[CachePointBlock]
class BedrockInvokeAnthropicMessagesRequest(TypedDict, total=False):
"""
Top-level request body accepted by AWS Bedrock `InvokeModel` /
`InvokeModelWithResponseStream` when calling an Anthropic Claude model with
the Messages API format. The LiteLLM /v1/messages → Bedrock Invoke
transformation filters outgoing requests to the keys of this TypedDict; any
other field (Anthropic-only extension, internal metadata, future addition)
is dropped before signing so Bedrock doesn't 400 with
"Extra inputs are not permitted".
Reference:
https://docs.aws.amazon.com/bedrock/latest/userguide/model-parameters-anthropic-claude-messages.html
https://docs.aws.amazon.com/bedrock/latest/userguide/model-parameters-anthropic-claude-messages-request-response.html
Editing this type is the single source of truth — the runtime allowlist in
`AmazonAnthropicClaudeMessagesConfig.BEDROCK_INVOKE_ALLOWED_TOP_LEVEL_FIELDS`
is derived from `__annotations__`, and a test asserts the resolved set
exactly, so any edit forces a conscious review.
Value types are intentionally loose (`list`, `dict`) — this type exists to
pin the allowed field names, not to validate nested structure.
"""
# Required by Bedrock
anthropic_version: str
max_tokens: int
messages: list
# Documented optional fields
anthropic_beta: List[str]
system: object # str or list[TextBlock]
stop_sequences: List[str]
temperature: float
top_p: float
top_k: int
tools: list
tool_choice: dict
# `thinking` is required for Opus 4.5 / Sonnet 4 extended thinking,
# `metadata` is part of the common Anthropic Messages API shape.
thinking: dict
metadata: dict

View file

@ -224,6 +224,9 @@ class GenericLiteLLMParams(CredentialLiteLLMParams, CustomPricingLiteLLMParams):
# adaptive-router params
adaptive_router_default_model: Optional[str] = None
adaptive_router_config: Optional[Dict] = None
# quality-router params
quality_router_config: Optional[Dict] = None
quality_router_default_model: Optional[str] = None
# Batch/File API Params
s3_bucket_name: Optional[str] = None

View file

@ -2851,6 +2851,7 @@ class StandardAuditLogPayload(TypedDict):
class StandardLoggingPayload(TypedDict):
id: str
trace_id: str # Trace multiple LLM calls belonging to same overall request (e.g. fallbacks/retries)
litellm_call_id: Optional[str] # UUID returned in x-litellm-call-id response header
call_type: str
stream: Optional[bool]
response_cost: float

View file

@ -1148,6 +1148,20 @@
"tool_use_system_prompt_tokens": 346,
"supports_native_structured_output": true
},
"anthropic.claude-mythos-preview": {
"input_cost_per_token": 0,
"output_cost_per_token": 0,
"litellm_provider": "bedrock",
"max_input_tokens": 1000000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"supports_function_calling": true,
"supports_vision": true,
"supports_prompt_caching": false,
"supports_reasoning": true,
"supports_tool_choice": true
},
"global.anthropic.claude-opus-4-7": {
"cache_creation_input_token_cost": 6.25e-06,
"cache_read_input_token_cost": 5e-07,

View file

@ -208,7 +208,7 @@ build-backend = "uv_build"
[tool.uv]
default-groups = ["dev"]
required-version = "==0.10.9"
required-version = ">=0.10.9"
exclude-newer = "3 days"
[tool.uv.sources]

View file

@ -0,0 +1,149 @@
"""
E2E tests for Bedrock Mantle (Claude Mythos Preview) integration.
Tests use a fake/mocked HTTP layer to verify the full request pipeline:
- correct endpoint URL
- model ID in the request body
- AWS SigV4 Authorization header present
- response parsing
"""
import json
import os
import sys
from unittest.mock import MagicMock, patch
import httpx
import pytest
sys.path.insert(0, os.path.abspath("../.."))
import litellm
from litellm.llms.custom_httpx.http_handler import HTTPHandler
MODEL = "bedrock/mantle/anthropic.claude-mythos-preview"
REGION = "us-east-1"
EXPECTED_URL = f"https://bedrock-mantle.{REGION}.api.aws/v1/messages"
FAKE_ANTHROPIC_RESPONSE = {
"id": "msg_fake123",
"type": "message",
"role": "assistant",
"model": "anthropic.claude-mythos-preview",
"content": [{"type": "text", "text": "Hello from Mythos!"}],
"stop_reason": "end_turn",
"stop_sequence": None,
"usage": {"input_tokens": 10, "output_tokens": 5},
}
def _make_fake_response(body: dict) -> MagicMock:
mock_resp = MagicMock(spec=httpx.Response)
mock_resp.status_code = 200
mock_resp.headers = httpx.Headers({"content-type": "application/json"})
mock_resp.text = json.dumps(body)
mock_resp.json.return_value = body
mock_resp.is_error = False
mock_resp.raise_for_status = MagicMock()
return mock_resp
def test_mantle_request_url_and_body():
"""Verify the correct URL is called and model appears in the request body."""
client = HTTPHandler()
with patch.object(
client, "post", return_value=_make_fake_response(FAKE_ANTHROPIC_RESPONSE)
) as mock_post:
try:
litellm.completion(
model=MODEL,
messages=[{"role": "user", "content": "Hello"}],
max_tokens=50,
aws_region_name=REGION,
aws_access_key_id="AKIAIOSFODNN7EXAMPLE",
aws_secret_access_key="wJalrXUtnFEMI/K7MDENG/bPxRfiCYEXAMPLEKEY",
client=client,
)
except Exception:
pass # response parsing may fail on mock; we only care about the outgoing call
mock_post.assert_called_once()
call_kwargs = mock_post.call_args.kwargs
# Correct endpoint
assert (
call_kwargs["url"] == EXPECTED_URL
), f"Expected {EXPECTED_URL}, got {call_kwargs['url']}"
# Request body has model ID (without "mantle/" prefix)
raw_data = call_kwargs.get("data") or call_kwargs.get("json")
body = json.loads(raw_data) if isinstance(raw_data, (str, bytes)) else raw_data
assert (
body["model"] == "anthropic.claude-mythos-preview"
), f"body['model'] = {body.get('model')}"
assert "messages" in body
assert body["max_tokens"] == 50
# AWS SigV4 Authorization header must be present
headers = call_kwargs.get("headers", {})
assert "Authorization" in headers, f"No Authorization header in {headers}"
assert headers["Authorization"].startswith(
"AWS4-HMAC-SHA256"
), f"Expected SigV4 auth, got: {headers['Authorization'][:50]}"
def test_mantle_request_does_not_include_mantle_prefix_in_body():
"""Ensure 'mantle/' never leaks into the request body."""
client = HTTPHandler()
with patch.object(
client, "post", return_value=_make_fake_response(FAKE_ANTHROPIC_RESPONSE)
) as mock_post:
try:
litellm.completion(
model=MODEL,
messages=[{"role": "user", "content": "Hi"}],
max_tokens=10,
aws_region_name=REGION,
aws_access_key_id="AKIAIOSFODNN7EXAMPLE",
aws_secret_access_key="wJalrXUtnFEMI/K7MDENG/bPxRfiCYEXAMPLEKEY",
client=client,
)
except Exception:
pass
call_kwargs = mock_post.call_args.kwargs
raw_data = call_kwargs.get("data") or call_kwargs.get("json")
body = json.loads(raw_data) if isinstance(raw_data, (str, bytes)) else raw_data
body_str = json.dumps(body)
assert "mantle/" not in body_str, f"'mantle/' leaked into body: {body_str}"
def test_mantle_region_reflected_in_url():
"""The region from aws_region_name must appear in the endpoint URL."""
client = HTTPHandler()
for region in ["us-east-1", "us-west-2", "eu-west-1"]:
with patch.object(
client, "post", return_value=_make_fake_response(FAKE_ANTHROPIC_RESPONSE)
) as mock_post:
try:
litellm.completion(
model=MODEL,
messages=[{"role": "user", "content": "Hi"}],
max_tokens=10,
aws_region_name=region,
aws_access_key_id="AKIAIOSFODNN7EXAMPLE",
aws_secret_access_key="wJalrXUtnFEMI/K7MDENG/bPxRfiCYEXAMPLEKEY",
client=client,
)
except Exception:
pass
call_kwargs = mock_post.call_args.kwargs
expected = f"https://bedrock-mantle.{region}.api.aws/v1/messages"
assert (
call_kwargs["url"] == expected
), f"region={region}: expected URL {expected}, got {call_kwargs['url']}"

View file

@ -72,7 +72,7 @@ def test_batch_completions_models():
def test_batch_completion_models_all_responses():
try:
responses = batch_completion_models_all_responses(
models=["gemini/gemini-2.5-flash-lite", "claude-3-haiku-20240307"],
models=["gemini/gemini-2.5-flash-lite", "claude-haiku-4-5-20251001"],
messages=[{"role": "user", "content": "write a poem"}],
max_tokens=10,
)

View file

@ -142,7 +142,7 @@ def trade(model_name: str) -> List[Trade]: # type: ignore
@pytest.mark.parametrize(
"model", ["claude-3-haiku-20240307", "anthropic.claude-3-haiku-20240307-v1:0"]
"model", ["claude-haiku-4-5-20251001", "anthropic.claude-3-haiku-20240307-v1:0"]
)
@pytest.mark.flaky(retries=6, delay=10)
def test_function_call_parsing(model):

View file

@ -47,7 +47,7 @@ def get_current_weather(location, unit="fahrenheit"):
[
"gpt-3.5-turbo-1106",
"mistral/mistral-large-latest",
"claude-3-haiku-20240307",
"claude-haiku-4-5-20251001",
"gemini/gemini-2.5-flash-lite",
"anthropic.claude-3-sonnet-20240229-v1:0",
],
@ -275,7 +275,7 @@ from litellm.types.utils import ChatCompletionMessageToolCall, Function, Message
"anthropic.claude-3-sonnet-20240229-v1:0",
"bedrock",
),
("claude-3-haiku-20240307", "anthropic"),
("claude-haiku-4-5-20251001", "anthropic"),
],
)
@pytest.mark.parametrize(

View file

@ -1509,7 +1509,7 @@ def test_router_fallbacks_with_wildcard_model_name():
{
"model_name": "claude-3-haiku",
"litellm_params": {
"model": "claude-3-haiku-20240307",
"model": "claude-haiku-4-5-20251001",
"api_key": os.getenv("ANTHROPIC_API_KEY"),
"mock_response": "Hi this is claude!",
},
@ -1555,7 +1555,7 @@ def test_fallbacks_with_different_messages():
{
"model_name": "claude-3-haiku",
"litellm_params": {
"model": "claude-3-haiku-20240307",
"model": "claude-haiku-4-5-20251001",
"api_key": os.getenv("ANTHROPIC_API_KEY"),
},
},

View file

@ -1727,7 +1727,7 @@ def test_openai_chat_completion_complete_response_call():
"model",
[
"gpt-3.5-turbo",
"claude-3-haiku-20240307",
"claude-haiku-4-5-20251001",
"o1",
],
)
@ -2247,7 +2247,7 @@ def streaming_and_function_calling_format_tests(idx, chunk):
[
# "gpt-3.5-turbo",
# "anthropic.claude-3-sonnet-20240229-v1:0",
"claude-3-haiku-20240307",
"claude-haiku-4-5-20251001",
],
)
def test_streaming_and_function_calling(model):

View file

@ -253,6 +253,7 @@ def validate_redacted_message_span_attributes(span):
or attr.startswith("gen_ai.cost.")
or attr.startswith("gen_ai.operation.")
or attr.startswith("gen_ai.request.")
or attr.startswith("litellm.")
), f"Non-metadata attribute found: {attr}"
pass

View file

@ -77,7 +77,7 @@ class TestAnthropicDirectAPI(BaseAnthropicMessagesTest):
@property
def model_config(self) -> Dict[str, Any]:
return {
"model": "claude-3-haiku-20240307",
"model": "claude-haiku-4-5-20251001",
"api_key": os.getenv("ANTHROPIC_API_KEY"),
}
@ -86,7 +86,7 @@ class TestAnthropicDirectAPI(BaseAnthropicMessagesTest):
"""
This is the model name that is expected to be in the logging payload
"""
return "claude-3-haiku-20240307"
return "claude-haiku-4-5-20251001"
class TestAnthropicBedrockAPI(BaseAnthropicMessagesTest):
@ -140,7 +140,7 @@ async def test_anthropic_messages_streaming_with_bad_request():
response = await litellm.anthropic.messages.acreate(
messages=[{"role": "user", "content": "hi"}],
api_key=os.getenv("ANTHROPIC_API_KEY"),
model="claude-3-haiku-20240307",
model="claude-haiku-4-5-20251001",
max_tokens=100,
stream=True,
)
@ -168,7 +168,7 @@ async def test_anthropic_messages_router_streaming_with_bad_request():
{
"model_name": "claude-special-alias",
"litellm_params": {
"model": "claude-3-haiku-20240307",
"model": "claude-haiku-4-5-20251001",
"api_key": os.getenv("ANTHROPIC_API_KEY"),
},
}
@ -205,7 +205,7 @@ async def test_anthropic_messages_litellm_router_non_streaming():
{
"model_name": "claude-special-alias",
"litellm_params": {
"model": "claude-3-haiku-20240307",
"model": "claude-haiku-4-5-20251001",
"api_key": os.getenv("ANTHROPIC_API_KEY"),
},
}
@ -243,7 +243,7 @@ async def test_anthropic_messages_litellm_router_routing_strategy():
{
"model_name": "claude-special-alias",
"litellm_params": {
"model": "claude-3-haiku-20240307",
"model": "claude-haiku-4-5-20251001",
"api_key": os.getenv("ANTHROPIC_API_KEY"),
},
}
@ -341,7 +341,7 @@ async def test_anthropic_messages_litellm_router_latency_metadata_tracking():
"type": "message",
"role": "assistant",
"content": [{"type": "text", "text": "Here's a joke for you!"}],
"model": "claude-3-haiku-20240307",
"model": "claude-haiku-4-5-20251001",
"stop_reason": "end_turn",
"usage": {"input_tokens": 10, "output_tokens": 20},
}
@ -355,7 +355,7 @@ async def test_anthropic_messages_litellm_router_latency_metadata_tracking():
{
"model_name": MODEL_GROUP,
"litellm_params": {
"model": "claude-3-haiku-20240307",
"model": "claude-haiku-4-5-20251001",
"api_key": os.getenv("ANTHROPIC_API_KEY"),
},
}
@ -419,7 +419,7 @@ async def test_anthropic_messages_litellm_router_latency_metadata_tracking():
assert "model_info" in litellm_metadata
# Verify other call parameters
assert call_kwargs["model"] == "claude-3-haiku-20240307"
assert call_kwargs["model"] == "claude-haiku-4-5-20251001"
assert call_kwargs["messages"] == messages
assert call_kwargs["max_tokens"] == 100
assert call_kwargs["metadata"] == {"user_id": "hello"}
@ -459,7 +459,7 @@ async def test_anthropic_messages_litellm_router_non_streaming_with_logging():
{
"model_name": MODEL_GROUP,
"litellm_params": {
"model": "claude-3-haiku-20240307",
"model": "claude-haiku-4-5-20251001",
"api_key": os.getenv("ANTHROPIC_API_KEY"),
},
}
@ -496,7 +496,7 @@ async def test_anthropic_messages_litellm_router_non_streaming_with_logging():
assert test_custom_logger.logged_standard_logging_payload["response"] is not None
assert (
test_custom_logger.logged_standard_logging_payload["model"]
== "claude-3-haiku-20240307"
== "claude-haiku-4-5-20251001"
)
# check logged usage + spend
@ -543,7 +543,7 @@ async def test_anthropic_messages_with_extra_headers():
"text": "Why did the chicken cross the road? To get to the other side!",
}
],
"model": "claude-3-haiku-20240307",
"model": "claude-haiku-4-5-20251001",
"stop_reason": "end_turn",
"usage": {"input_tokens": 10, "output_tokens": 20},
}
@ -556,7 +556,7 @@ async def test_anthropic_messages_with_extra_headers():
response = await litellm.anthropic.messages.acreate(
messages=messages,
api_key=api_key,
model="claude-3-haiku-20240307",
model="claude-haiku-4-5-20251001",
max_tokens=100,
client=mock_client,
provider_specific_header={
@ -689,7 +689,7 @@ async def test_anthropic_messages_with_thinking():
"text": "Why did the chicken cross the road? To get to the other side!",
}
],
"model": "claude-3-haiku-20240307",
"model": "claude-haiku-4-5-20251001",
"stop_reason": "end_turn",
"usage": {"input_tokens": 10, "output_tokens": 20},
}
@ -702,7 +702,7 @@ async def test_anthropic_messages_with_thinking():
response = await litellm.anthropic.messages.acreate(
messages=messages,
api_key=api_key,
model="claude-3-haiku-20240307",
model="claude-haiku-4-5-20251001",
max_tokens=100,
client=mock_client,
thinking={"budget_tokens": 100},
@ -717,7 +717,7 @@ async def test_anthropic_messages_with_thinking():
request_body = json.loads(call_kwargs.get("data", {}))
print("REQUEST BODY", request_body)
assert request_body["max_tokens"] == 100
assert request_body["model"] == "claude-3-haiku-20240307"
assert request_body["model"] == "claude-haiku-4-5-20251001"
assert request_body["messages"] == messages
assert request_body["thinking"] == {"budget_tokens": 100}

View file

@ -1,48 +1,143 @@
# conftest.py
import importlib
import asyncio
import copy
import inspect
import os
import sys
import warnings
import pytest
sys.path.insert(
0, os.path.abspath("../..")
) # Adds the parent directory to the system path
import litellm
import litellm.proxy.proxy_server
# Top-level assignments of these types are the ones importlib.reload(litellm)
# would have effectively reset. We snapshot them at conftest import time and
# deep-copy the snapshot back before every test.
_SNAPSHOT_TYPES = (list, dict, set, tuple, str, int, float, bool, bytes)
def _snapshot_mutable_state(module):
"""Capture a per-module snapshot of primitive and collection attributes."""
snapshot = {}
for attr in list(vars(module)):
if attr.startswith("_"):
continue
try:
value = getattr(module, attr)
except Exception as exc:
warnings.warn(
f"conftest: could not read {module.__name__}.{attr} during snapshot: {exc}",
stacklevel=2,
)
continue
if value is None or isinstance(value, _SNAPSHOT_TYPES):
try:
snapshot[attr] = copy.deepcopy(value)
except Exception as exc:
warnings.warn(
f"conftest: could not snapshot {module.__name__}.{attr}: {exc}",
stacklevel=2,
)
return snapshot
def _restore_mutable_state(module, snapshot):
for attr, default in snapshot.items():
try:
setattr(module, attr, copy.deepcopy(default))
except Exception as exc:
warnings.warn(
f"conftest: could not restore {module.__name__}.{attr}: {exc}",
stacklevel=2,
)
def _collect_flushable_caches():
"""Return (module, attr) pairs whose values expose flush_cache()."""
targets = []
for module in (litellm, litellm.proxy.proxy_server):
for attr in list(vars(module)):
if attr.startswith("_"):
continue
try:
value = getattr(module, attr)
except Exception:
continue
# Only instances — a class reference has an unbound flush_cache
# that can't be called without a self argument.
if inspect.isclass(value) or inspect.ismodule(value):
continue
if callable(getattr(value, "flush_cache", None)):
targets.append((module, attr))
return targets
def _flush_caches(targets):
for module, attr in targets:
try:
value = getattr(module, attr)
except Exception:
continue
flush = getattr(value, "flush_cache", None)
if callable(flush):
try:
flush()
except Exception as exc:
warnings.warn(
f"conftest: flush_cache failed on {module.__name__}.{attr}: {exc}",
stacklevel=2,
)
# Snapshot once at conftest import — these are the "clean" module states.
_LITELLM_STATE = _snapshot_mutable_state(litellm)
_PROXY_SERVER_STATE = _snapshot_mutable_state(litellm.proxy.proxy_server)
_FLUSHABLE_CACHES = _collect_flushable_caches()
@pytest.fixture(scope="function", autouse=True)
def setup_and_teardown():
"""Reset mutable module state on litellm and proxy_server before each test.
Replaces a previous importlib.reload(litellm) approach that cost ~17s
per test (re-executing the full litellm __init__ import chain).
What IS reset:
- Top-level module attributes of type list / dict / set / tuple
/ str / int / float / bool / bytes, and None-valued attributes.
These cover callback lists, general_settings, master_key,
premium_user, prisma_client, etc. — anything the old reload() reset
by re-executing the module body.
- Any module-level object instance that exposes flush_cache() (the
DualCache and LLMClientCache family), which handles cache state
that can't round-trip through deepcopy because of internal locks.
What is NOT reset:
- Class instances without flush_cache() (e.g. ProxyLogging,
JWTHandler, FastAPI routers, loggers). If a test mutates such an
instance in-place (setattr on the instance, appending to one of
its internal lists, etc.), the mutation will leak into later tests.
Use pytest's monkeypatch.setattr() or a local fixture for those
cases — don't rely on this autouse fixture to undo them.
"""
This fixture reloads litellm before every function. To speed up testing by removing callbacks being chained.
"""
curr_dir = os.getcwd() # Get the current working directory
sys.path.insert(
0, os.path.abspath("../..")
) # Adds the project directory to the system path
import litellm
from litellm import Router
importlib.reload(litellm)
try:
if hasattr(litellm, "proxy") and hasattr(litellm.proxy, "proxy_server"):
importlib.reload(litellm.proxy.proxy_server)
except Exception as e:
print(f"Error reloading litellm.proxy.proxy_server: {e}")
import asyncio
_restore_mutable_state(litellm, _LITELLM_STATE)
_restore_mutable_state(litellm.proxy.proxy_server, _PROXY_SERVER_STATE)
_flush_caches(_FLUSHABLE_CACHES)
loop = asyncio.get_event_loop_policy().new_event_loop()
asyncio.set_event_loop(loop)
print(litellm)
# from litellm import Router, completion, aembedding, acompletion, embedding
yield
# Teardown code (executes after the yield point)
loop.close() # Close the loop created earlier
asyncio.set_event_loop(None) # Remove the reference to the loop
try:
yield
finally:
loop.close()
asyncio.set_event_loop(None)
def pytest_collection_modifyitems(config, items):

View file

@ -2752,3 +2752,26 @@ class TestResponseIdFallback(unittest.TestCase):
mock_span.set_attribute.assert_any_call(
"gen_ai.response.id", "litellm-img-call-101"
)
def test_litellm_call_id_emitted_as_span_attribute(self):
"""litellm.call_id must be set on the span from standard_logging_payload."""
otel = OpenTelemetry()
mock_span = MagicMock()
call_id = "my-litellm-call-uuid-456"
kwargs = {
"model": "gpt-4o",
"optional_params": {},
"litellm_params": {"custom_llm_provider": "openai"},
"standard_logging_object": {
"id": "chatcmpl-provider-id",
"litellm_call_id": call_id,
"call_type": "completion",
"metadata": {},
},
}
response_obj = {"id": "chatcmpl-provider-id", "model": "gpt-4o"}
otel.set_attributes(mock_span, kwargs, response_obj)
mock_span.set_attribute.assert_any_call("litellm.call_id", call_id)

View file

@ -2410,3 +2410,29 @@ def test_get_additional_headers_reset_fields_preserved():
assert result is not None
assert result["x_ratelimit_reset_requests"] == "1s" # type: ignore
assert result["x_ratelimit_reset_tokens"] == "100ms" # type: ignore
# ── litellm_call_id propagation ───────────────────────────────────────────────
def test_get_standard_logging_object_payload_includes_litellm_call_id(logging_obj):
"""litellm_call_id from kwargs must appear in the returned StandardLoggingPayload."""
import datetime
from litellm.litellm_core_utils.litellm_logging import (
get_standard_logging_object_payload,
)
call_id = "test-call-id-abc-123"
now = datetime.datetime.now()
payload = get_standard_logging_object_payload(
kwargs={"litellm_call_id": call_id, "model": "gpt-4o", "messages": []},
init_response_obj={},
start_time=now,
end_time=now,
logging_obj=logging_obj,
status="success",
)
assert payload is not None
assert payload["litellm_call_id"] == call_id

View file

@ -579,6 +579,155 @@ def test_bedrock_messages_strips_output_config_with_output_format():
assert "output_format" not in result
def test_bedrock_messages_strips_context_management():
"""
Ensure context_management is stripped from the request before sending to
Bedrock Invoke, which doesn't support this Anthropic-specific parameter.
Claude Code sends context_management on every request; leaving it in the body
causes a 400 "context_management: Extra inputs are not permitted" from Bedrock.
"""
from litellm.types.router import GenericLiteLLMParams
cfg = AmazonAnthropicClaudeMessagesConfig()
messages = [{"role": "user", "content": [{"type": "text", "text": "Hello"}]}]
optional_params = {
"max_tokens": 4096,
"context_management": {
"edits": [{"type": "clear_thinking_20251015", "keep": "all"}]
},
}
result = cfg.transform_anthropic_messages_request(
model="anthropic.claude-3-haiku-20240307-v1:0",
messages=messages,
anthropic_messages_optional_request_params=optional_params,
litellm_params=GenericLiteLLMParams(),
headers={},
)
assert (
"context_management" not in result
), "context_management should be stripped — Bedrock Invoke rejects it"
assert result.get("max_tokens") == 4096
def test_bedrock_messages_allowlist_filters_anthropic_only_fields():
"""
Bedrock Invoke rejects any top-level body field it doesn't recognize with
"Extra inputs are not permitted". Defend against that by filtering the
outgoing body to a Bedrock-supported allowlist — catches Anthropic-only
extensions (speed, mcp_servers, container, ...) and any future additions
Claude Code starts sending before we learn about them.
"""
from litellm.types.router import GenericLiteLLMParams
cfg = AmazonAnthropicClaudeMessagesConfig()
messages = [{"role": "user", "content": [{"type": "text", "text": "Hello"}]}]
optional_params = {
"max_tokens": 4096,
"temperature": 0.5,
"speed": "fast",
"mcp_servers": [{"type": "url", "url": "https://example.com"}],
"container": {"skills": []},
"inference_geo": "us",
"output_config": {"effort": "low"},
"context_management": {"edits": []},
}
result = cfg.transform_anthropic_messages_request(
model="anthropic.claude-3-haiku-20240307-v1:0",
messages=messages,
anthropic_messages_optional_request_params=optional_params,
litellm_params=GenericLiteLLMParams(),
headers={},
)
for bad in (
"speed",
"mcp_servers",
"container",
"inference_geo",
"output_config",
"context_management",
"model",
"stream",
):
assert bad not in result, f"{bad} should be stripped by the allowlist"
# Supported fields pass through.
assert result["max_tokens"] == 4096
assert result["temperature"] == 0.5
assert result["anthropic_version"] == cfg.DEFAULT_BEDROCK_ANTHROPIC_API_VERSION
# Every surviving key is in the allowlist.
assert set(result).issubset(cfg.BEDROCK_INVOKE_ALLOWED_TOP_LEVEL_FIELDS)
def test_bedrock_messages_filters_user_provided_unsupported_beta_header():
"""
In proxy deployments the client (e.g. Claude Code) doesn't know the backend
is Bedrock and may send Anthropic-direct beta headers Bedrock can't handle.
All betas must go through the provider mapping, not just auto-injected ones
— otherwise Bedrock 400s on the unsupported value.
"""
from litellm.types.router import GenericLiteLLMParams
cfg = AmazonAnthropicClaudeMessagesConfig()
messages = [{"role": "user", "content": [{"type": "text", "text": "Hello"}]}]
optional_params = {"max_tokens": 128}
# `advisor-tool-2026-03-01` has no bedrock mapping entry → must be dropped.
# `context-1m-2025-08-07` does → must pass through.
headers = {
"anthropic-beta": "advisor-tool-2026-03-01,context-1m-2025-08-07",
}
result = cfg.transform_anthropic_messages_request(
model="anthropic.claude-3-haiku-20240307-v1:0",
messages=messages,
anthropic_messages_optional_request_params=optional_params,
litellm_params=GenericLiteLLMParams(),
headers=headers,
)
betas = result.get("anthropic_beta") or []
assert (
"advisor-tool-2026-03-01" not in betas
), "user-provided beta not in the Bedrock mapping must be dropped"
assert (
"context-1m-2025-08-07" in betas
), "user-provided beta that IS in the Bedrock mapping should survive"
def test_bedrock_messages_renames_user_provided_aliased_beta_header():
"""
Bedrock's config maps `advanced-tool-use-2025-11-20` to
`tool-search-tool-2025-10-19`. User-provided betas must go through the
rename too, not be forwarded under their Anthropic-direct spelling.
"""
from litellm.types.router import GenericLiteLLMParams
cfg = AmazonAnthropicClaudeMessagesConfig()
messages = [{"role": "user", "content": [{"type": "text", "text": "Hello"}]}]
optional_params = {"max_tokens": 128}
headers = {"anthropic-beta": "advanced-tool-use-2025-11-20"}
result = cfg.transform_anthropic_messages_request(
model="anthropic.claude-3-haiku-20240307-v1:0",
messages=messages,
anthropic_messages_optional_request_params=optional_params,
litellm_params=GenericLiteLLMParams(),
headers=headers,
)
betas = result.get("anthropic_beta") or []
assert (
"advanced-tool-use-2025-11-20" not in betas
), "Anthropic-direct spelling should be rewritten, not forwarded verbatim"
assert (
"tool-search-tool-2025-10-19" in betas
), "user-provided beta should be renamed to the Bedrock-side spelling"
@pytest.mark.asyncio
async def test_promote_message_stop_usage_preserves_message_delta_output_tokens():
"""

View file

@ -95,7 +95,7 @@ class TestAnthropicBetaHeaderSupport:
def test_messages_transformation_anthropic_beta(self):
"""Test that Messages API transformation includes anthropic_beta in request."""
config = AmazonAnthropicClaudeMessagesConfig()
headers = {"anthropic-beta": "output-128k-2025-02-19"}
headers = {"anthropic-beta": "context-1m-2025-08-07"}
result = config.transform_anthropic_messages_request(
model="anthropic.claude-haiku-4-5-20251001-v1:0",
@ -107,7 +107,7 @@ class TestAnthropicBetaHeaderSupport:
assert "anthropic_beta" in result
# Sort both arrays before comparing to avoid flakiness from ordering differences
assert sorted(result["anthropic_beta"]) == sorted(["output-128k-2025-02-19"])
assert sorted(result["anthropic_beta"]) == sorted(["context-1m-2025-08-07"])
def test_converse_computer_use_compatibility(self):
"""Test that user anthropic_beta headers work with computer use tools."""

View file

@ -0,0 +1,105 @@
"""
Unit tests for the Bedrock Mantle (Claude Mythos Preview) integration.
Tests cover route detection, URL construction, and config dispatch for both
the /chat/completions and /messages endpoints.
"""
from litellm.llms.bedrock.common_utils import BedrockModelInfo, get_bedrock_chat_config
from litellm.llms.bedrock.chat.mantle.transformation import AmazonMantleConfig
from litellm.llms.bedrock.messages.mantle_transformation import (
AmazonMantleMessagesConfig,
)
def test_get_bedrock_route_mantle():
assert (
BedrockModelInfo.get_bedrock_route("mantle/anthropic.claude-mythos-preview")
== "mantle"
)
def test_get_bedrock_route_mantle_does_not_match_other_routes():
assert (
BedrockModelInfo.get_bedrock_route("anthropic.claude-3-sonnet-20240229-v1:0")
!= "mantle"
)
assert (
BedrockModelInfo.get_bedrock_route("converse/anthropic.claude-3-sonnet")
!= "mantle"
)
def test_explicit_mantle_route_flag():
assert (
BedrockModelInfo._explicit_mantle_route(
"mantle/anthropic.claude-mythos-preview"
)
is True
)
assert BedrockModelInfo._explicit_mantle_route("anthropic.claude-3-sonnet") is False
assert (
BedrockModelInfo._explicit_mantle_route("converse/anthropic.claude-3-sonnet")
is False
)
def test_mantle_url_construction():
config = AmazonMantleConfig()
url = config.get_complete_url(
api_base=None,
api_key=None,
model="mantle/anthropic.claude-mythos-preview",
optional_params={"aws_region_name": "us-east-1"},
litellm_params={},
)
assert url == "https://bedrock-mantle.us-east-1.api.aws/v1/messages"
def test_mantle_url_construction_different_region():
config = AmazonMantleConfig()
url = config.get_complete_url(
api_base=None,
api_key=None,
model="mantle/anthropic.claude-mythos-preview",
optional_params={"aws_region_name": "us-west-2"},
litellm_params={},
)
assert url == "https://bedrock-mantle.us-west-2.api.aws/v1/messages"
def test_get_bedrock_chat_config_returns_mantle_config():
config = get_bedrock_chat_config("mantle/anthropic.claude-mythos-preview")
assert isinstance(config, AmazonMantleConfig)
def test_get_bedrock_provider_config_for_messages_api_mantle():
config = BedrockModelInfo.get_bedrock_provider_config_for_messages_api(
"mantle/anthropic.claude-mythos-preview"
)
assert isinstance(config, AmazonMantleMessagesConfig)
def test_mantle_messages_url_construction():
config = AmazonMantleMessagesConfig()
url = config.get_complete_url(
api_base=None,
api_key=None,
model="mantle/anthropic.claude-mythos-preview",
optional_params={"aws_region_name": "us-east-1"},
litellm_params={},
)
assert url == "https://bedrock-mantle.us-east-1.api.aws/v1/messages"
def test_mantle_transform_request_strips_prefix_and_adds_model():
config = AmazonMantleConfig()
request = config.transform_request(
model="mantle/anthropic.claude-mythos-preview",
messages=[{"role": "user", "content": "Hello"}],
optional_params={"max_tokens": 100},
litellm_params={},
headers={},
)
assert request["model"] == "anthropic.claude-mythos-preview"
assert "mantle/" not in request["model"]

View file

@ -891,6 +891,61 @@ class TestOpenAIChatCompletionsHandlerStreamingOutput:
assert result == responses_so_far
class TestGetStructuredMessages:
"""Test the get_structured_messages method."""
def test_should_return_messages_from_chat_completions_request(self):
"""Test that messages are returned from a chat completions request."""
handler = OpenAIChatCompletionsHandler()
data = {
"messages": [
{"role": "system", "content": "You are helpful."},
{"role": "user", "content": "Hello"},
]
}
result = handler.get_structured_messages(data)
assert result is not None
assert len(result) == 2
assert result[0]["role"] == "system"
assert result[1]["role"] == "user"
def test_should_return_none_when_no_messages(self):
"""Test that None is returned when no messages key exists."""
handler = OpenAIChatCompletionsHandler()
data = {"model": "gpt-4"}
result = handler.get_structured_messages(data)
assert result is None
def test_should_return_none_for_none_messages(self):
"""Test that None is returned when messages is explicitly None."""
handler = OpenAIChatCompletionsHandler()
data = {"messages": None}
result = handler.get_structured_messages(data)
assert result is None
def test_should_handle_multimodal_content(self):
"""Test that messages with multimodal content are returned."""
handler = OpenAIChatCompletionsHandler()
data = {
"messages": [
{
"role": "user",
"content": [
{"type": "text", "text": "What's in this image?"},
{
"type": "image_url",
"image_url": {"url": "https://example.com/image.png"},
},
],
}
]
}
result = handler.get_structured_messages(data)
assert result is not None
assert len(result) == 1
assert isinstance(result[0]["content"], list)
if __name__ == "__main__":
# Run the tests
pytest.main([__file__, "-v"])

View file

@ -995,3 +995,63 @@ class TestOpenAIResponsesHandlerStreamingOutputProcessing:
# Should return the responses
assert result == responses_so_far
class TestGetStructuredMessages:
"""Test the get_structured_messages method for Responses API handler."""
def test_should_convert_string_input_to_messages(self):
"""Test that a simple string input is converted to OpenAI messages."""
handler = OpenAIResponsesHandler()
data = {"input": "What is the capital of France?"}
result = handler.get_structured_messages(data)
assert result is not None
assert len(result) >= 1
found_user = False
for msg in result:
if isinstance(msg, dict) and msg.get("role") == "user":
found_user = True
break
assert found_user, f"Expected a user message, got: {result}"
def test_should_convert_list_input_to_messages(self):
"""Test that list input (ResponseInputParam) is converted to OpenAI messages."""
handler = OpenAIResponsesHandler()
data = {
"input": [
{"role": "user", "content": "Hello"},
{"role": "assistant", "content": "Hi there!"},
{"role": "user", "content": "How are you?"},
]
}
result = handler.get_structured_messages(data)
assert result is not None
assert len(result) >= 3
def test_should_include_instructions_as_system_message(self):
"""Test that instructions are included as a system message."""
handler = OpenAIResponsesHandler()
data = {
"input": "Roll a d20",
"instructions": "You are a helpful dungeon master.",
}
result = handler.get_structured_messages(data)
assert result is not None
has_system = any(
isinstance(msg, dict) and msg.get("role") == "system" for msg in result
)
assert has_system, f"Expected system message from instructions, got: {result}"
def test_should_return_none_when_no_input(self):
"""Test that None is returned when input key is missing."""
handler = OpenAIResponsesHandler()
data = {"model": "gpt-4o"}
result = handler.get_structured_messages(data)
assert result is None
def test_should_return_none_for_none_input(self):
"""Test that None is returned when input is explicitly None."""
handler = OpenAIResponsesHandler()
data = {"input": None}
result = handler.get_structured_messages(data)
assert result is None

View file

@ -9,6 +9,7 @@ sys.path.insert(
0, os.path.abspath("../../..")
) # Adds the parent directory to the system path
import httpx
import pytest
from prisma.errors import ClientNotConnectedError, HTTPClientClosedError, PrismaError
@ -110,220 +111,127 @@ async def test_db_health_prisma_client_none():
@pytest.mark.asyncio
@pytest.mark.parametrize(
"prisma_error",
"transport_error",
[
PrismaError(),
httpx.ConnectError("All connection attempts failed"),
ClientNotConnectedError(),
HTTPClientClosedError(),
PrismaError("Can't reach database server"),
],
)
async def test_db_health_error_flag_off_raises_no_reconnect(prisma_error):
async def test_db_health_transport_error_never_raises(transport_error):
"""
When health_check raises and allow_requests_on_db_unavailable is False,
handle_db_exception re-raises immediately. The reconnect path is never
reached, so disconnect/connect are never called.
Regression test for the /health/readiness 503 loop bug.
handle_db_exception() used to re-raise inside _db_health_readiness_check,
turning any DB outage into a 503 "Service Unhealthy" response that never
recovered. Transport errors (ClientNotConnectedError, httpx.ConnectError,
etc.) must return {"status": "disconnected"} — never raise.
"""
mock_prisma = MagicMock()
mock_prisma.health_check = AsyncMock(side_effect=prisma_error)
mock_prisma.disconnect = AsyncMock()
mock_prisma.health_check = AsyncMock(side_effect=transport_error)
mock_prisma.attempt_db_reconnect = AsyncMock(return_value=False)
_health_endpoints_module.db_health_cache = {
"status": "connected",
"last_updated": datetime.now() - timedelta(seconds=20),
}
with (
patch("litellm.proxy.proxy_server.prisma_client", mock_prisma),
patch(
"litellm.proxy.proxy_server.general_settings",
{"allow_requests_on_db_unavailable": False},
),
):
with pytest.raises(Exception) as exc_info:
await _db_health_readiness_check()
with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma):
result = await _db_health_readiness_check()
assert exc_info.value is prisma_error
mock_prisma.disconnect.assert_not_called()
assert _health_endpoints_module.db_health_cache["status"] == "disconnected"
assert result["status"] == "disconnected"
mock_prisma.attempt_db_reconnect.assert_called_once_with(
reason="health_readiness_check"
)
@pytest.mark.asyncio
@pytest.mark.parametrize(
"prisma_error",
"transport_error",
[
PrismaError("Can't reach database server"),
httpx.ConnectError("All connection attempts failed"),
ClientNotConnectedError(),
HTTPClientClosedError(),
],
)
async def test_db_health_error_flag_on_reconnect_succeeds(prisma_error):
async def test_db_health_transport_error_reconnect_succeeds(transport_error):
"""
When health_check raises, allow_requests_on_db_unavailable is True,
and the reconnect cycle (disconnect -> connect -> health_check) succeeds,
return 'connected' and update the cache.
When health_check raises a transport error and attempt_db_reconnect
succeeds, the second health_check passes and we return 'connected'.
"""
mock_prisma = MagicMock()
mock_prisma.health_check = AsyncMock(side_effect=[prisma_error, None])
mock_prisma.disconnect = AsyncMock()
mock_prisma.connect = AsyncMock()
mock_prisma.health_check = AsyncMock(side_effect=[transport_error, None])
mock_prisma.attempt_db_reconnect = AsyncMock(return_value=True)
_health_endpoints_module.db_health_cache = {
"status": "connected",
"last_updated": datetime.now() - timedelta(seconds=20),
}
with (
patch("litellm.proxy.proxy_server.prisma_client", mock_prisma),
patch(
"litellm.proxy.proxy_server.general_settings",
{"allow_requests_on_db_unavailable": True},
),
):
with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma):
result = await _db_health_readiness_check()
assert result["status"] == "connected"
mock_prisma.disconnect.assert_called_once()
mock_prisma.connect.assert_called_once()
mock_prisma.attempt_db_reconnect.assert_called_once_with(
reason="health_readiness_check"
)
assert mock_prisma.health_check.call_count == 2
@pytest.mark.asyncio
@pytest.mark.parametrize(
"prisma_error",
"transport_error",
[
PrismaError("Can't reach database server"),
httpx.ConnectError("All connection attempts failed"),
ClientNotConnectedError(),
HTTPClientClosedError(),
],
)
async def test_db_health_error_flag_on_reconnect_fails(prisma_error):
async def test_db_health_transport_error_reconnect_fails(transport_error):
"""
When health_check raises, allow_requests_on_db_unavailable is True,
but the reconnect also fails, return 'disconnected' instead of raising.
This respects the flag's intent: keep serving even without a DB.
When health_check raises a transport error and attempt_db_reconnect also
fails, return 'disconnected' without raising.
"""
mock_prisma = MagicMock()
mock_prisma.health_check = AsyncMock(side_effect=prisma_error)
mock_prisma.disconnect = AsyncMock()
mock_prisma.connect = AsyncMock()
_health_endpoints_module.db_health_cache = {
"status": "connected",
"last_updated": datetime.now() - timedelta(seconds=20),
}
with (
patch("litellm.proxy.proxy_server.prisma_client", mock_prisma),
patch(
"litellm.proxy.proxy_server.general_settings",
{"allow_requests_on_db_unavailable": True},
),
):
result = await _db_health_readiness_check()
assert result["status"] == "disconnected"
mock_prisma.disconnect.assert_called_once()
mock_prisma.connect.assert_called_once()
@pytest.mark.asyncio
async def test_db_health_non_transport_error_flag_off_raises():
"""
When health_check raises a non-transport error and
allow_requests_on_db_unavailable is False, handle_db_exception
re-raises before reaching the is_database_transport_error guard.
Cache is still invalidated before the re-raise.
"""
non_transport_error = PrismaError("UniqueViolationError")
mock_prisma = MagicMock()
mock_prisma.health_check = AsyncMock(side_effect=non_transport_error)
mock_prisma.disconnect = AsyncMock()
mock_prisma.connect = AsyncMock()
_health_endpoints_module.db_health_cache = {
"status": "connected",
"last_updated": datetime.now() - timedelta(seconds=20),
}
with (
patch("litellm.proxy.proxy_server.prisma_client", mock_prisma),
patch(
"litellm.proxy.proxy_server.general_settings",
{"allow_requests_on_db_unavailable": False},
),
):
with pytest.raises(PrismaError):
await _db_health_readiness_check()
assert _health_endpoints_module.db_health_cache["status"] == "disconnected"
mock_prisma.disconnect.assert_not_called()
mock_prisma.connect.assert_not_called()
@pytest.mark.asyncio
async def test_db_health_non_transport_error_flag_on_skips_reconnect():
"""
When health_check raises a non-transport error (e.g. data-layer) and
allow_requests_on_db_unavailable is True, handle_db_exception swallows
the exception, then is_database_transport_error returns False so the
reconnect cycle is skipped. Returns 'disconnected' without calling
disconnect/connect.
"""
non_transport_error = PrismaError("UniqueViolationError")
mock_prisma = MagicMock()
mock_prisma.health_check = AsyncMock(side_effect=non_transport_error)
mock_prisma.disconnect = AsyncMock()
mock_prisma.connect = AsyncMock()
_health_endpoints_module.db_health_cache = {
"status": "connected",
"last_updated": datetime.now() - timedelta(seconds=20),
}
with (
patch("litellm.proxy.proxy_server.prisma_client", mock_prisma),
patch(
"litellm.proxy.proxy_server.general_settings",
{"allow_requests_on_db_unavailable": True},
),
):
result = await _db_health_readiness_check()
assert result["status"] == "disconnected"
mock_prisma.disconnect.assert_not_called()
mock_prisma.connect.assert_not_called()
@pytest.mark.asyncio
async def test_db_health_reconnect_disconnect_fails():
"""
When disconnect() itself raises during the reconnect cycle,
the inner except catches it and returns 'disconnected'.
connect() and the second health_check() are never called.
"""
transport_error = ClientNotConnectedError()
mock_prisma = MagicMock()
mock_prisma.health_check = AsyncMock(side_effect=transport_error)
mock_prisma.disconnect = AsyncMock(side_effect=RuntimeError("already closed"))
mock_prisma.connect = AsyncMock()
mock_prisma.attempt_db_reconnect = AsyncMock(
side_effect=RuntimeError("reconnect failed")
)
_health_endpoints_module.db_health_cache = {
"status": "connected",
"last_updated": datetime.now() - timedelta(seconds=20),
}
with (
patch("litellm.proxy.proxy_server.prisma_client", mock_prisma),
patch(
"litellm.proxy.proxy_server.general_settings",
{"allow_requests_on_db_unavailable": True},
),
):
with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma):
result = await _db_health_readiness_check()
assert result["status"] == "disconnected"
mock_prisma.disconnect.assert_called_once()
mock_prisma.connect.assert_not_called()
@pytest.mark.asyncio
async def test_db_health_non_transport_error_returns_disconnected():
"""
When health_check raises a non-transport error (e.g. data-layer error),
is_database_transport_error returns False so reconnect is skipped.
Returns 'disconnected' without raising and without calling attempt_db_reconnect.
"""
non_transport_error = PrismaError("UniqueViolationError")
mock_prisma = MagicMock()
mock_prisma.health_check = AsyncMock(side_effect=non_transport_error)
mock_prisma.attempt_db_reconnect = AsyncMock()
_health_endpoints_module.db_health_cache = {
"status": "connected",
"last_updated": datetime.now() - timedelta(seconds=20),
}
with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma):
result = await _db_health_readiness_check()
assert result["status"] == "disconnected"
mock_prisma.attempt_db_reconnect.assert_not_called()
@pytest.mark.asyncio

View file

@ -12,7 +12,148 @@ sys.path.insert(
from litellm.router_strategy.auto_router.auto_router import AutoRouter
pytestmark = pytest.mark.skip(reason="Skipping auto router tests - beta feature")
pytestmark_skip_beta = pytest.mark.skip(
reason="Skipping auto router tests - beta feature"
)
class TestExtractTextFromMessages:
"""Tests for AutoRouter._extract_text_from_messages (no semantic_router dependency)."""
def test_should_extract_content_from_simple_user_message(self):
messages = [{"role": "user", "content": "Hello world"}]
result = AutoRouter._extract_text_from_messages(messages)
assert result == "Hello world"
def test_should_extract_last_user_message_from_tool_call_conversation(self):
messages = [
{"role": "user", "content": "What's the weather in NYC?"},
{
"role": "assistant",
"content": None,
"tool_calls": [
{
"id": "call_abc123",
"type": "function",
"function": {
"name": "get_weather",
"arguments": '{"location": "NYC"}',
},
}
],
},
{
"role": "tool",
"tool_call_id": "call_abc123",
"content": "72°F and sunny",
},
{"role": "user", "content": "Now tell me about London"},
]
result = AutoRouter._extract_text_from_messages(messages)
assert result == "Now tell me about London"
def test_should_find_user_message_when_last_message_is_assistant_with_tool_calls(
self,
):
messages = [
{"role": "user", "content": "What's the weather?"},
{
"role": "assistant",
"content": None,
"tool_calls": [
{
"id": "call_abc",
"type": "function",
"function": {"name": "get_weather", "arguments": "{}"},
}
],
},
]
result = AutoRouter._extract_text_from_messages(messages)
assert result == "What's the weather?"
def test_should_find_user_message_when_last_message_is_tool_response(self):
messages = [
{"role": "user", "content": "What's the weather?"},
{
"role": "assistant",
"content": None,
"tool_calls": [
{
"id": "call_abc",
"type": "function",
"function": {"name": "get_weather", "arguments": "{}"},
}
],
},
{
"role": "tool",
"tool_call_id": "call_abc",
"content": "72°F and sunny",
},
]
result = AutoRouter._extract_text_from_messages(messages)
assert result == "What's the weather?"
def test_should_handle_multimodal_content_list(self):
messages = [
{
"role": "user",
"content": [
{"type": "text", "text": "What's in this image?"},
{
"type": "image_url",
"image_url": {"url": "https://example.com/img.png"},
},
],
}
]
result = AutoRouter._extract_text_from_messages(messages)
assert result == "What's in this image?"
def test_should_handle_multimodal_content_with_multiple_text_blocks(self):
messages = [
{
"role": "user",
"content": [
{"type": "text", "text": "First part"},
{"type": "text", "text": "Second part"},
{
"type": "image_url",
"image_url": {"url": "https://example.com/img.png"},
},
],
}
]
result = AutoRouter._extract_text_from_messages(messages)
assert result == "First part Second part"
def test_should_return_empty_string_when_user_content_is_none(self):
messages = [{"role": "user", "content": None}]
result = AutoRouter._extract_text_from_messages(messages)
assert result == ""
def test_should_return_empty_string_when_no_user_messages(self):
messages = [
{"role": "system", "content": "You are a helpful assistant"},
{
"role": "assistant",
"content": None,
"tool_calls": [
{
"id": "call_abc",
"type": "function",
"function": {"name": "get_weather", "arguments": "{}"},
}
],
},
]
result = AutoRouter._extract_text_from_messages(messages)
assert result == ""
def test_should_return_empty_string_for_empty_messages_list(self):
result = AutoRouter._extract_text_from_messages([])
assert result == ""
@pytest.fixture
@ -41,6 +182,7 @@ def mock_route_choice():
return mock_choice
@pytestmark_skip_beta
class TestAutoRouter:
"""Test class for AutoRouter methods."""

View file

@ -7,7 +7,7 @@ Tests the rule-based complexity scoring and tier assignment logic.
import os
import sys
from typing import Dict, List
from unittest.mock import MagicMock
from unittest.mock import MagicMock, patch
import pytest
@ -828,3 +828,222 @@ class TestRouterComplexityDeploymentMethods:
)
router.init_complexity_router_deployment(deployment)
assert "auto_router/complexity_router/test-router" in router.complexity_routers
class TestAsyncPreRoutingHookMultiFormat:
"""Test async_pre_routing_hook with multiple input formats."""
@pytest.mark.asyncio
async def test_should_route_with_chat_completions_messages(self, complexity_router):
"""Test routing with standard chat completions messages."""
result = await complexity_router.async_pre_routing_hook(
model="test-model",
request_kwargs={},
messages=[{"role": "user", "content": "What is 2+2?"}],
)
assert result is not None
assert result.model is not None
assert result.messages is not None
@pytest.mark.asyncio
async def test_should_route_with_responses_api_string_input(
self, complexity_router
):
"""Test routing with Responses API string input via handler dispatch."""
from litellm.llms.openai.responses.guardrail_translation.handler import (
OpenAIResponsesHandler,
)
from litellm.types.utils import CallTypes
mock_mappings = {CallTypes.responses: OpenAIResponsesHandler}
with patch(
"litellm.llms.load_guardrail_translation_mappings",
return_value=mock_mappings,
):
result = await complexity_router.async_pre_routing_hook(
model="test-model",
request_kwargs={"input": "What is the capital of France?"},
messages=None,
input="What is the capital of France?",
)
assert result is not None
assert result.model is not None
# messages should be None since the original request didn't have messages
assert result.messages is None
@pytest.mark.asyncio
async def test_should_route_with_responses_api_list_input(self, complexity_router):
"""Test routing with Responses API list input via handler dispatch."""
from litellm.llms.openai.responses.guardrail_translation.handler import (
OpenAIResponsesHandler,
)
from litellm.types.utils import CallTypes
mock_mappings = {CallTypes.responses: OpenAIResponsesHandler}
list_input = [
{"role": "user", "content": "Hello"},
{"role": "assistant", "content": "Hi there!"},
{
"role": "user",
"content": "Write a Python function to sort a list using merge sort",
},
]
with patch(
"litellm.llms.load_guardrail_translation_mappings",
return_value=mock_mappings,
):
result = await complexity_router.async_pre_routing_hook(
model="test-model",
request_kwargs={"input": list_input},
messages=None,
input=list_input,
)
assert result is not None
assert result.model is not None
assert result.messages is None
@pytest.mark.asyncio
async def test_should_use_route_based_inference(self, complexity_router):
"""Test that route-based call type inference is used when available."""
from litellm.llms.openai.responses.guardrail_translation.handler import (
OpenAIResponsesHandler,
)
from litellm.types.utils import CallTypes
mock_mappings = {CallTypes.responses: OpenAIResponsesHandler}
with patch(
"litellm.llms.load_guardrail_translation_mappings",
return_value=mock_mappings,
):
result = await complexity_router.async_pre_routing_hook(
model="test-model",
request_kwargs={
"input": "Roll 2d4+1",
"litellm_metadata": {
"user_api_key_request_route": "/v1/responses",
},
},
messages=None,
)
assert result is not None
assert result.model is not None
@pytest.mark.asyncio
async def test_should_return_none_when_no_messages_or_input(
self, complexity_router
):
"""Test that None is returned when neither messages nor input is available."""
result = await complexity_router.async_pre_routing_hook(
model="test-model",
request_kwargs={},
messages=None,
input=None,
)
assert result is None
@pytest.mark.asyncio
async def test_should_prefer_original_messages_over_conversion(
self, complexity_router
):
"""Test that original messages are used when both messages and input are available."""
messages = [{"role": "user", "content": "What is 2+2?"}]
result = await complexity_router.async_pre_routing_hook(
model="test-model",
request_kwargs={"input": "This should be ignored"},
messages=messages,
)
assert result is not None
assert result.messages == messages
@pytest.mark.asyncio
async def test_should_include_instructions_in_classification(
self, complexity_router
):
"""Test that Responses API instructions influence classification via system message."""
from litellm.llms.openai.responses.guardrail_translation.handler import (
OpenAIResponsesHandler,
)
from litellm.types.utils import CallTypes
mock_mappings = {CallTypes.responses: OpenAIResponsesHandler}
with patch(
"litellm.llms.load_guardrail_translation_mappings",
return_value=mock_mappings,
):
result = await complexity_router.async_pre_routing_hook(
model="test-model",
request_kwargs={
"input": "Write merge sort",
"instructions": "You are an expert Python developer. Use advanced algorithms and optimize for performance.",
},
messages=None,
)
assert result is not None
assert result.model is not None
class TestExtractUserMessageAndSystemPrompt:
"""Test the _extract_user_message_and_system_prompt static method."""
def test_should_extract_user_message(self):
"""Test extraction of the last user message."""
messages = [
{"role": "system", "content": "You are helpful."},
{"role": "user", "content": "Hello"},
{"role": "assistant", "content": "Hi!"},
{"role": "user", "content": "How are you?"},
]
user_msg, sys_prompt = ComplexityRouter._extract_user_message_and_system_prompt(
messages
)
assert user_msg == "How are you?"
assert sys_prompt == "You are helpful."
def test_should_handle_no_user_message(self):
"""Test when there is no user message."""
messages = [
{"role": "system", "content": "You are helpful."},
{"role": "assistant", "content": "Hi!"},
]
user_msg, sys_prompt = ComplexityRouter._extract_user_message_and_system_prompt(
messages
)
assert user_msg is None
assert sys_prompt == "You are helpful."
def test_should_handle_multipart_content(self):
"""Test extraction from multipart content messages."""
messages = [
{
"role": "user",
"content": [
{"type": "text", "text": "Describe this image"},
{
"type": "image_url",
"image_url": {"url": "https://example.com/img.png"},
},
],
}
]
user_msg, sys_prompt = ComplexityRouter._extract_user_message_and_system_prompt(
messages
)
assert user_msg == "Describe this image"
assert sys_prompt is None
def test_should_handle_empty_messages(self):
"""Test with empty messages list."""
user_msg, sys_prompt = ComplexityRouter._extract_user_message_and_system_prompt(
[]
)
assert user_msg is None
assert sys_prompt is None

File diff suppressed because it is too large Load diff