mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
Merge branch 'litellm_internal_staging' into litellm_adaptive_routing
This commit is contained in:
commit
c7342bdc4f
51 changed files with 3723 additions and 397 deletions
|
|
@ -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"
|
||||
|
|
|
|||
9
.github/workflows/test-unit-proxy-db.yml
vendored
9
.github/workflows/test-unit-proxy-db.yml
vendored
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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" && \
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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 (
|
||||
|
|
|
|||
|
|
@ -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": (
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
0
litellm/llms/bedrock/chat/mantle/__init__.py
Normal file
0
litellm/llms/bedrock/chat/mantle/__init__.py
Normal file
91
litellm/llms/bedrock/chat/mantle/transformation.py
Normal file
91
litellm/llms/bedrock/chat/mantle/transformation.py
Normal 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
|
||||
|
|
@ -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":
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
69
litellm/llms/bedrock/messages/mantle_transformation.py
Normal file
69
litellm/llms/bedrock/messages/mantle_transformation.py
Normal 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
|
||||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
21
litellm/router_strategy/quality_router/__init__.py
Normal file
21
litellm/router_strategy/quality_router/__init__.py
Normal 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",
|
||||
]
|
||||
74
litellm/router_strategy/quality_router/config.py
Normal file
74
litellm/router_strategy/quality_router/config.py
Normal 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")
|
||||
446
litellm/router_strategy/quality_router/quality_router.py
Normal file
446
litellm/router_strategy/quality_router/quality_router.py
Normal 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,
|
||||
)
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
149
tests/llm_translation/test_bedrock_mantle.py
Normal file
149
tests/llm_translation/test_bedrock_mantle.py
Normal 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']}"
|
||||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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"),
|
||||
},
|
||||
},
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
105
tests/test_litellm/llms/bedrock/test_mantle.py
Normal file
105
tests/test_litellm/llms/bedrock/test_mantle.py
Normal 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"]
|
||||
|
|
@ -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"])
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
1033
tests/test_litellm/router_strategy/test_quality_router.py
Normal file
1033
tests/test_litellm/router_strategy/test_quality_router.py
Normal file
File diff suppressed because it is too large
Load diff
Loading…
Add table
Reference in a new issue