Merge pull request #26306 from BerriAI/litellm_internal_staging
Some checks failed
Unit Tests: Proxy DB Operations / proxy-db (auth-checks, tests/proxy_unit_tests/test_auth_checks.py tests/proxy_unit_tests/test_user_api_key_auth.py, 20, 8) (push) Has been cancelled
Unit Tests: Proxy DB Operations / proxy-db (key-generation, tests/proxy_unit_tests/test_key_generate_prisma.py, 30, 0) (push) Has been cancelled
Unit Tests: Proxy DB Operations / proxy-db (proxy-utils, tests/proxy_unit_tests/test_proxy_utils.py, 20, 8) (push) Has been cancelled
Unit Tests: Proxy DB Operations / proxy-db (remaining, 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_p… (push) Has been cancelled
Unit Tests: Security / security (push) Has been cancelled

merge main
This commit is contained in:
Sameer Kankute 2026-04-23 08:39:04 +05:30 committed by GitHub
commit ef1c6aeea6
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
139 changed files with 11819 additions and 1381 deletions

File diff suppressed because it is too large Load diff

136
.github/workflows/test-code-quality.yml vendored Normal file
View file

@ -0,0 +1,136 @@
name: Code Quality Checks
on:
pull_request:
branches:
- main
- litellm_internal_staging
- litellm_oss_branch
- "litellm_**"
permissions:
contents: read
concurrency:
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }}
cancel-in-progress: true
jobs:
code-quality:
runs-on: ubuntu-latest
timeout-minutes: 15
steps:
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
persist-credentials: false
- name: Checkout litellm-docs (for documentation_tests)
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
repository: BerriAI/litellm-docs
path: _litellm_docs_checkout
persist-credentials: false
- name: Wire up docs path expected by documentation_tests/*
run: |
# documentation_tests scripts read from docs/my-website/docs/...
# In litellm-docs the same files live at docs/... (repo root).
# Point docs/my-website -> litellm-docs checkout so the paths resolve.
rm -rf docs/my-website
ln -s ../_litellm_docs_checkout docs/my-website
- name: Set up Python
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
with:
python-version: "3.12"
- name: Set up uv
uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7
with:
version: "0.10.9"
- name: Cache uv dependencies
uses: actions/cache@0057852bfaa89a56745cba8c7296529d2fc39830 # v4.3.0
with:
path: |
~/.cache/uv
.venv
key: ${{ runner.os }}-uv-${{ hashFiles('uv.lock') }}
restore-keys: |
${{ runner.os }}-uv-
- name: Install dependencies
run: uv sync --frozen --all-groups --all-extras
- name: check_licenses
run: uv run --no-sync python ./tests/code_coverage_tests/check_licenses.py
- name: check_provider_folders_documented
run: uv run --no-sync python ./tests/code_coverage_tests/check_provider_folders_documented.py
- name: router_code_coverage
run: uv run --no-sync python ./tests/code_coverage_tests/router_code_coverage.py
- name: test_chat_completion_imports
run: uv run --no-sync python ./tests/code_coverage_tests/test_chat_completion_imports.py
- name: info_log_check
run: uv run --no-sync python ./tests/code_coverage_tests/info_log_check.py
- name: check_guardrail_apply_decorator
run: uv run --no-sync python ./tests/code_coverage_tests/check_guardrail_apply_decorator.py
- name: test_ban_set_verbose
run: uv run --no-sync python ./tests/code_coverage_tests/test_ban_set_verbose.py
- name: code_qa_check_tests
run: uv run --no-sync python ./tests/code_coverage_tests/code_qa_check_tests.py
- name: check_get_model_cost_key_performance
run: uv run --no-sync python ./tests/code_coverage_tests/check_get_model_cost_key_performance.py
- name: test_proxy_types_import
run: uv run --no-sync python ./tests/code_coverage_tests/test_proxy_types_import.py
- name: callback_manager_test
run: uv run --no-sync python ./tests/code_coverage_tests/callback_manager_test.py
- name: recursive_detector
run: uv run --no-sync python ./tests/code_coverage_tests/recursive_detector.py
- name: test_router_strategy_async
run: uv run --no-sync python ./tests/code_coverage_tests/test_router_strategy_async.py
- name: litellm_logging_code_coverage
run: uv run --no-sync python ./tests/code_coverage_tests/litellm_logging_code_coverage.py
- name: ensure_async_clients_test
run: uv run --no-sync python ./tests/code_coverage_tests/ensure_async_clients_test.py
- name: enforce_llms_folder_style
run: uv run --no-sync python ./tests/code_coverage_tests/enforce_llms_folder_style.py
- name: prevent_key_leaks_in_exceptions
run: uv run --no-sync python ./tests/code_coverage_tests/prevent_key_leaks_in_exceptions.py
- name: check_unsafe_enterprise_import
run: uv run --no-sync python ./tests/code_coverage_tests/check_unsafe_enterprise_import.py
- name: ban_copy_deepcopy_kwargs
run: uv run --no-sync python ./tests/code_coverage_tests/ban_copy_deepcopy_kwargs.py
- name: check_fastuuid_usage
run: uv run --no-sync python ./tests/code_coverage_tests/check_fastuuid_usage.py
- name: memory_test
run: uv run --no-sync python ./tests/code_coverage_tests/memory_test.py
- name: documentation_test_env_keys
run: uv run --no-sync python ./tests/documentation_tests/test_env_keys.py
- name: documentation_test_router_settings
run: uv run --no-sync python ./tests/documentation_tests/test_router_settings.py
- name: documentation_test_api_docs
run: uv run --no-sync python ./tests/documentation_tests/test_api_docs.py

39
.github/workflows/test-semgrep.yml vendored Normal file
View file

@ -0,0 +1,39 @@
name: Semgrep
on:
pull_request:
branches:
- main
- litellm_internal_staging
- litellm_oss_branch
- "litellm_**"
permissions:
contents: read
concurrency:
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }}
cancel-in-progress: true
jobs:
semgrep:
runs-on: ubuntu-latest
timeout-minutes: 10
steps:
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
persist-credentials: false
- name: Set up Python
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
with:
python-version: "3.12"
- name: Set up uv
uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7
with:
version: "0.10.9"
- name: Run Semgrep (custom rules)
run: uv tool run --from 'semgrep==1.157.0' semgrep scan --config .semgrep/rules . --error

View file

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

View file

@ -27,10 +27,8 @@ RUN apk add --no-cache \
npm \
libsndfile
ENV PRISMA_BINARY_CACHE_DIR=/app/.cache/prisma-python/binaries \
UV_PROJECT_ENVIRONMENT=/app/.venv \
ENV UV_PROJECT_ENVIRONMENT=/app/.venv \
UV_LINK_MODE=copy \
XDG_CACHE_HOME=/app/.cache \
PATH="/app/.venv/bin:${PATH}"
# Copy dependency metadata first for layer caching
@ -94,11 +92,14 @@ RUN apk add --no-cache bash openssl tzdata nodejs npm python3 libsndfile supervi
{ apk del --no-cache npm 2>/dev/null || true; }
WORKDIR /app
ENV PRISMA_BINARY_CACHE_DIR=/app/.cache/prisma-python/binaries \
XDG_CACHE_HOME=/app/.cache \
PATH="/app/.venv/bin:${PATH}"
ENV PATH="/app/.venv/bin:${PATH}"
COPY --from=builder /app /app
# Prisma binaries live in $HOME/.cache (default prisma-python location),
# which is /root/.cache here. Copy them from the builder so they survive
# deployments that volume-mount /app/.cache (e.g. readOnlyRootFilesystem
# + emptyDir) — otherwise the mount would shadow the baked-in query engine.
COPY --from=builder /root/.cache /root/.cache
RUN find /app/.venv -type f -path "*/tornado/test/*" -delete && \
find /app/.venv -type d -path "*/tornado/test" -delete

View file

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

View file

@ -26,10 +26,8 @@ RUN apk add --no-cache \
npm \
libsndfile
ENV PRISMA_BINARY_CACHE_DIR=/app/.cache/prisma-python/binaries \
UV_PROJECT_ENVIRONMENT=/app/.venv \
ENV UV_PROJECT_ENVIRONMENT=/app/.venv \
UV_LINK_MODE=copy \
XDG_CACHE_HOME=/app/.cache \
PATH="/app/.venv/bin:${PATH}"
# Copy dependency metadata first for layer caching
@ -92,11 +90,14 @@ RUN apk add --no-cache bash openssl tzdata nodejs npm python3 libsndfile supervi
{ apk del --no-cache npm 2>/dev/null || true; }
WORKDIR /app
ENV PRISMA_BINARY_CACHE_DIR=/app/.cache/prisma-python/binaries \
XDG_CACHE_HOME=/app/.cache \
PATH="/app/.venv/bin:${PATH}"
ENV PATH="/app/.venv/bin:${PATH}"
COPY --from=builder /app /app
# Prisma binaries live in $HOME/.cache (default prisma-python location),
# which is /root/.cache here. Copy them from the builder so they survive
# deployments that volume-mount /app/.cache (e.g. readOnlyRootFilesystem
# + emptyDir) — otherwise the mount would shadow the baked-in query engine.
COPY --from=builder /root/.cache /root/.cache
RUN find /app/.venv -type f -path "*/tornado/test/*" -delete && \
find /app/.venv -type d -path "*/tornado/test" -delete

View file

@ -138,7 +138,7 @@ RUN mkdir -p /nonexistent /var/lib/litellm/assets /var/lib/litellm/ui && \
[ -n "$LITELLM_PROXY_EXTRAS_PATH" ] && chmod -R g+w "$LITELLM_PROXY_EXTRAS_PATH" || true && \
chmod -R g+rX "$PRISMA_PATH" /var/lib/litellm/ui /var/lib/litellm/assets /app/.cache
USER nobody
USER 65534
RUN prisma generate --schema=./schema.prisma

View file

@ -2,7 +2,7 @@ schemaVersion: 2.0.0
metadataTest:
entrypoint: ["docker/prod_entrypoint.sh"]
user: "nobody"
user: "65534"
workdir: "/app"
fileExistenceTests:

View file

@ -0,0 +1,155 @@
# [BETA] Adaptive Router
:::info
Beta feature. Share feedback on [Discord](https://discord.gg/wuPM9dRgDw) or [Slack](https://join.slack.com/t/litellmossslack/shared_invite/zt-3o7nkuyfr-p_kbNJj8taRfXGgQI1~YyA).
:::
**Requirements:** LiteLLM Proxy with a Postgres database. Quality estimates are stored in Postgres and loaded on startup — without a database the router works but forgets everything learned on restart.
You have a cheap model and an expensive one. You want to use the cheap one when it's good enough, and the expensive one when it actually matters — without hardcoding rules you'll spend months tuning.
The adaptive router does this automatically. It tracks which model performs best for each type of request (code, writing, analysis, etc.) and routes accordingly, balancing quality against cost based on weights you control.
## Quick start
```yaml
model_list:
- model_name: gpt-4o
litellm_params:
model: openai/gpt-4o
model_info:
input_cost_per_token: 0.0000025
adaptive_router_preferences:
quality_tier: 3 # 1=budget, 2=mid, 3=frontier
strengths: ["code_generation", "analytical_reasoning"]
- model_name: gpt-4o-mini
litellm_params:
model: openai/gpt-4o-mini
model_info:
input_cost_per_token: 0.00000015
adaptive_router_preferences:
quality_tier: 2
strengths: ["factual_lookup"]
- model_name: my-router
litellm_params:
model: auto_router/adaptive_router
adaptive_router_config:
available_models: ["gpt-4o", "gpt-4o-mini"]
weights:
quality: 0.7 # raise this if quality complaints; lower if bill too high
cost: 0.3 # must sum to 1.0 with quality
```
Route to it by setting `model` to your adaptive router's name:
```bash
curl -X POST {{baseURL}}/v1/chat/completions \
-H "Content-Type: application/json" \
-H "Authorization: Bearer $LITELLM_API_KEY" \
-d '{
"model": "my-router",
"messages": [
{"role": "user", "content": "build me a python script that parses CSV"},
{"role": "assistant", "content": "Here is a script using csv.DictReader..."},
{"role": "user", "content": "now add error handling for missing files"},
{"role": "assistant", "content": "Wrap the open() call in a try/except FileNotFoundError..."},
{"role": "user", "content": "perfect, that worked. thanks!"}
]
}'
```
The response includes a header telling you which model was actually picked:
```
x-litellm-adaptive-router-model: gpt-4o
```
The "thanks!" turn in the example above fires a satisfaction signal — that's what moves the bandit.
## Tuning cost vs. quality
The `weights` are your main lever:
| Goal | quality | cost |
|---|---|---|
| Minimize cost, quality is secondary | 0.3 | 0.7 |
| Balanced | 0.5 | 0.5 |
| Quality-first (default) | 0.7 | 0.3 |
| Quality non-negotiable | 0.9 | 0.1 |
The router learns over time. For the first ~10 requests per model, it relies on the tiers you declared. After that, real performance data takes over.
## Force a minimum quality tier per request
If a specific request needs a frontier model regardless of cost, pass this header:
```
x-litellm-min-quality-tier: 3
```
You can also pass `min_quality_tier` via request metadata instead of a header.
## What's being learned
The router classifies each request into one of 7 types and tracks how each model performs on each independently. A model that's great at factual lookup but poor at code will win factual requests and lose code requests — even if it's cheaper overall.
| Type | Example |
|---|---|
| `code_generation` | "write me a Python sort function" |
| `code_understanding` | "explain what this function does" |
| `technical_design` | "how should I design this API?" |
| `analytical_reasoning` | "calculate the probability that..." |
| `writing` | "draft an email to my team about..." |
| `factual_lookup` | "what is the capital of France?" |
| `general` | anything else |
[**See classifier code**](https://github.com/BerriAI/litellm/blob/litellm_adaptive_routing/litellm/router_strategy/adaptive_router/classifier.py)
Learning signals are inspired by [Signals: Trajectory Sampling and Triage for Agentic Interactions](https://arxiv.org/pdf/2604.00356).
## Inspect the current state
```
GET /adaptive_router/{router_name}/state
```
Returns current quality estimates per model per request type. Useful for understanding why a model is or isn't being picked.
```json
{
"routers": [
{
"router_name": "smart-cheap-router",
"available_models": ["fast", "smart"],
"weights": { "quality": 0.7, "cost": 0.3 },
"cells": [
{
"request_type": "analytical_reasoning",
"model": "fast",
"quality_mean": 0.5,
"samples": 0
},
{
"request_type": "analytical_reasoning",
"model": "smart",
"quality_mean": 0.95,
"samples": 0
}
]
}
]
}
```
`quality_mean` is the key number — it's the router's current estimate of how well that model handles that request type. `samples` counts how many real observations have moved the prior (starts at 0; the cold-start prior mass is excluded).
## Known limitations
- Latency isn't scored — a slow model can still win on quality + cost
- Signals are regex-based and English-biased — no LLM judge
- Hard cap of 200 observations per cell; no decay yet
- Once a model is picked for a session, other models' turns in that session don't contribute to learning

View file

@ -26,6 +26,7 @@
},
"devDependencies": {
"@docusaurus/module-type-aliases": "3.8.1",
"ajv": "^8.18.0",
"dotenv": "16.6.1"
},
"engines": {

View file

@ -32,6 +32,7 @@
},
"devDependencies": {
"@docusaurus/module-type-aliases": "3.8.1",
"ajv": "^8.18.0",
"dotenv": "16.6.1"
},
"browserslist": {

View file

@ -1060,6 +1060,7 @@ const sidebars = {
},
items: [
"routing",
"adaptive_router",
"scheduler",
"proxy/auto_routing",
"proxy/load_balancing",

View file

@ -0,0 +1,39 @@
-- One row per (router, request_type, model). Hot path on every routing decision.
CREATE TABLE "LiteLLM_AdaptiveRouterState" (
router_name TEXT NOT NULL,
request_type TEXT NOT NULL,
model_name TEXT NOT NULL,
alpha DOUBLE PRECISION NOT NULL,
beta DOUBLE PRECISION NOT NULL,
total_samples INTEGER NOT NULL DEFAULT 0,
last_updated_at TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
PRIMARY KEY (router_name, request_type, model_name)
);
-- One row per (session, router, model). Updated per turn via the queue.
CREATE TABLE "LiteLLM_AdaptiveRouterSession" (
session_id TEXT NOT NULL,
router_name TEXT NOT NULL,
model_name TEXT NOT NULL,
classified_type TEXT NOT NULL,
misalignment_count INTEGER NOT NULL DEFAULT 0,
stagnation_count INTEGER NOT NULL DEFAULT 0,
disengagement_count INTEGER NOT NULL DEFAULT 0,
satisfaction_count INTEGER NOT NULL DEFAULT 0,
failure_count INTEGER NOT NULL DEFAULT 0,
loop_count INTEGER NOT NULL DEFAULT 0,
exhaustion_count INTEGER NOT NULL DEFAULT 0,
last_user_content TEXT,
last_assistant_content TEXT,
tool_call_history JSONB NOT NULL DEFAULT '[]',
pending_tool_calls JSONB NOT NULL DEFAULT '{}',
turn_count INTEGER NOT NULL DEFAULT 0,
last_processed_turn INTEGER NOT NULL DEFAULT -1,
clean_credit_awarded BOOLEAN NOT NULL DEFAULT FALSE,
terminal_status INTEGER,
last_activity_at TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
PRIMARY KEY (session_id, router_name, model_name)
);
CREATE INDEX "idx_adaptive_router_session_activity"
ON "LiteLLM_AdaptiveRouterSession" (last_activity_at);

View file

@ -0,0 +1,3 @@
-- AlterTable
ALTER TABLE "LiteLLM_TeamMembership" ADD COLUMN "total_spend" DOUBLE PRECISION NOT NULL DEFAULT 0.0;

View file

@ -616,6 +616,7 @@ model LiteLLM_TeamMembership {
user_id String
team_id String
spend Float @default(0.0)
total_spend Float @default(0.0)
budget_id String?
litellm_budget_table LiteLLM_BudgetTable? @relation(fields: [budget_id], references: [budget_id])
@@id([user_id, team_id])
@ -1223,3 +1224,46 @@ model LiteLLM_ClaudeCodePluginTable {
@@map("LiteLLM_ClaudeCodePluginTable")
}
// Per-(router, request_type, model) Beta posterior for the adaptive router.
model LiteLLM_AdaptiveRouterState {
router_name String
request_type String
model_name String
alpha Float
beta Float
total_samples Int @default(0)
last_updated_at DateTime @default(now()) @updatedAt
@@id([router_name, request_type, model_name])
}
// Per-(session, router, model) signal counters for the adaptive router.
model LiteLLM_AdaptiveRouterSession {
session_id String
router_name String
model_name String
classified_type String
misalignment_count Int @default(0)
stagnation_count Int @default(0)
disengagement_count Int @default(0)
satisfaction_count Int @default(0)
failure_count Int @default(0)
loop_count Int @default(0)
exhaustion_count Int @default(0)
last_user_content String?
last_assistant_content String?
tool_call_history Json @default("[]")
pending_tool_calls Json @default("{}")
turn_count Int @default(0)
last_processed_turn Int @default(-1)
clean_credit_awarded Boolean @default(false)
terminal_status Int?
last_activity_at DateTime @default(now()) @updatedAt
@@id([session_id, router_name, model_name])
@@index([last_activity_at], map: "idx_adaptive_router_session_activity")
}

View file

@ -30,6 +30,26 @@ def _get_prisma_env() -> dict:
return prisma_env
_MIGRATION_TS_RE = re.compile(r"^(\d{14})_")
def _migration_timestamp(name: str) -> int:
"""Extract the leading `YYYYMMDDHHMMSS` timestamp from a migration name.
Returns 0 if the name doesn't match the Prisma pattern — unexpected-format
entries sort as "oldest" and are treated as historical.
"""
m = _MIGRATION_TS_RE.match(name)
return int(m.group(1)) if m else 0
def _max_migration_timestamp(names) -> int:
"""Max timestamp in a set/list of migration names (0 if empty)."""
if not names:
return 0
return max(_migration_timestamp(n) for n in names)
def _get_prisma_command() -> str:
"""Get the Prisma command to use, bypassing Python wrapper in offline mode."""
if str_to_bool(os.getenv("PRISMA_OFFLINE_MODE")):
@ -383,18 +403,301 @@ class ProxyExtrasDBManager:
)
@staticmethod
def setup_database(use_migrate: bool = False) -> bool:
def _strip_prisma_query_params(url: str) -> str:
"""Remove Prisma-specific query params (connection_limit, pool_timeout,
schema, etc.) from DATABASE_URL so psycopg can parse it."""
from urllib.parse import urlparse, urlunparse, parse_qsl, urlencode
parsed = urlparse(url)
if not parsed.query:
return url
libpq_params = {
"sslmode",
"sslcert",
"sslkey",
"sslrootcert",
"sslpassword",
"application_name",
"connect_timeout",
"client_encoding",
"options",
"service",
"gssencmode",
"krbsrvname",
"target_session_attrs",
}
kept = [(k, v) for k, v in parse_qsl(parsed.query) if k in libpq_params]
return urlunparse(parsed._replace(query=urlencode(kept)))
@staticmethod
def _warn_if_db_ahead_of_head(migrations_dir: str) -> None:
"""
Log a warning if _prisma_migrations contains applied migrations with
timestamps newer than every migration this build ships.
This is informational only for the v2 resolver it tells the operator
the DB was likely migrated by a newer deployment, which is usually a
signal that this (older) version shouldn't run against it. We do NOT
block startup: many users have weird _prisma_migrations state from
prior thrashing bugs, and blocking them would be a breaking change.
Safe no-op if psycopg isn't installed or DB isn't reachable.
"""
database_url = os.getenv("DATABASE_URL")
if not database_url:
return
try:
import psycopg
except ImportError:
return
cleaned_url = ProxyExtrasDBManager._strip_prisma_query_params(database_url)
known = set(ProxyExtrasDBManager._get_migration_names(migrations_dir))
try:
# autocommit=True keeps the SELECT outside a transaction. Without
# it, psycopg3's `with conn` calls COMMIT on clean exit — which
# fails after `UndefinedTable` (fresh DB) leaves the transaction
# in an aborted state.
with psycopg.connect(
cleaned_url, connect_timeout=10, autocommit=True
) as conn:
try:
rows = conn.execute(
"SELECT migration_name FROM _prisma_migrations "
"WHERE finished_at IS NOT NULL AND rolled_back_at IS NULL"
).fetchall()
except psycopg.errors.UndefinedTable:
return
except (psycopg.OperationalError, psycopg.DatabaseError):
# Swallow connection failures AND any other DB-layer error
# (e.g. InsufficientPrivilege if the runtime user lacks SELECT
# on _prisma_migrations). This is an informational check —
# never block startup on it.
return
applied = {r[0] for r in rows}
unknown = applied - known
if not unknown:
return
head_newest_ts = _max_migration_timestamp(known)
hostile = {
name for name in unknown if _migration_timestamp(name) > head_newest_ts
}
if not hostile:
return
sorted_hostile = sorted(hostile)
logger.warning(
"Database has %d migration(s) applied that are NEWER than any "
"migration this LiteLLM version ships. This usually means the "
"database was migrated by a newer LiteLLM deployment. Some API "
"endpoints may fail because this proxy's Prisma client does not "
"know about those schema changes. Consider upgrading this "
"deployment. Unknown: %s",
len(hostile),
", ".join(sorted_hostile[:5]) + (" ..." if len(sorted_hostile) > 5 else ""),
)
@staticmethod
def _setup_database_v2(use_migrate: bool) -> bool:
"""
v2 migration resolver (opt-in via --use_v2_migration_resolver).
Runs `prisma migrate deploy` and handles standard recovery paths
(P3005 baseline, P3009/P3018 idempotent errors). Critically, it does
NOT call `_resolve_all_migrations` the diff-and-force recovery that
caused schema thrashing when two LiteLLM versions contended for the
same DB during rolling deploys.
Ahead-of-HEAD state (DB has migrations newer than this build ships)
is logged as a warning, not a fatal error users whose DBs got into
weird shapes from the old thrashing should still be able to start.
"""
schema_path = ProxyExtrasDBManager._get_prisma_dir() + "/schema.prisma"
migrations_dir = ProxyExtrasDBManager._get_prisma_dir()
if not use_migrate:
# Preserve `prisma db push` path unchanged.
original_dir = os.getcwd()
os.chdir(migrations_dir)
try:
subprocess.run(
[_get_prisma_command(), "db", "push", "--accept-data-loss"],
timeout=60,
check=True,
env=_get_prisma_env(),
)
return True
except (
subprocess.CalledProcessError,
subprocess.TimeoutExpired,
) as e:
# Re-raise as RuntimeError so proxy_cli.py's
# `except RuntimeError` catches it and exits cleanly.
raise RuntimeError(f"prisma db push failed.\n\nDetail: {e}") from e
finally:
os.chdir(original_dir)
# Informational — never blocks.
ProxyExtrasDBManager._warn_if_db_ahead_of_head(migrations_dir)
original_dir = os.getcwd()
os.chdir(migrations_dir)
try:
for attempt in range(4):
try:
result = subprocess.run(
[_get_prisma_command(), "migrate", "deploy"],
timeout=60,
check=True,
capture_output=True,
text=True,
env=_get_prisma_env(),
)
logger.info(f"prisma migrate deploy stdout: {result.stdout}")
return True
except subprocess.TimeoutExpired:
logger.info(
f"prisma migrate deploy attempt {attempt + 1} timed out, retrying"
)
time.sleep(random.randrange(5, 15))
continue
except subprocess.CalledProcessError as e:
stderr = e.stderr or ""
if "P3005" in stderr and "database schema is not empty" in stderr:
logger.info(
"Schema exists but no migrations ledger — creating baseline"
)
ProxyExtrasDBManager._create_baseline_migration(schema_path)
continue
if "P3009" in stderr:
migration_match = re.search(r"`(\d+_\S+?)`", stderr)
if (
migration_match
and ProxyExtrasDBManager._is_idempotent_error(stderr)
):
name = migration_match.group(1)
logger.info(
f"Migration {name} failed idempotently — marking applied and retrying"
)
try:
ProxyExtrasDBManager._roll_back_migration(name)
except (
subprocess.CalledProcessError,
subprocess.TimeoutExpired,
):
pass # may already be rolled-back
try:
ProxyExtrasDBManager._resolve_specific_migration(name)
except (
subprocess.CalledProcessError,
subprocess.TimeoutExpired,
) as resolve_err:
# We're already inside the outer
# `except CalledProcessError` handler —
# re-raising CalledProcessError from here
# would escape as itself, bypassing
# proxy_cli.py's `except RuntimeError`.
raise RuntimeError(
f"Failed to mark migration {name} as applied "
f"after idempotent recovery. Manual "
f"intervention may be required.\n\n"
f"Detail: {resolve_err}"
) from resolve_err
continue
raise RuntimeError(
"Database migration failed and cannot be auto-recovered. "
f"Manual intervention required.\n\nPrisma error:\n{stderr}"
) from e
if "P3018" in stderr:
if ProxyExtrasDBManager._is_permission_error(stderr):
raise RuntimeError(
"Database migration failed due to insufficient "
"permissions. Please grant the required privileges "
f"and retry.\n\nPrisma error:\n{stderr}"
) from e
migration_match = re.search(
r"Migration name: (\d+_\S+)", stderr
)
if (
migration_match
and ProxyExtrasDBManager._is_idempotent_error(stderr)
):
name = migration_match.group(1)
logger.info(
f"Migration {name} SQL hit idempotent error — marking applied and retrying"
)
try:
ProxyExtrasDBManager._roll_back_migration(name)
except (
subprocess.CalledProcessError,
subprocess.TimeoutExpired,
):
pass # may already be rolled-back
try:
ProxyExtrasDBManager._resolve_specific_migration(name)
except (
subprocess.CalledProcessError,
subprocess.TimeoutExpired,
) as resolve_err:
raise RuntimeError(
f"Failed to mark migration {name} as applied "
f"after idempotent recovery. Manual "
f"intervention may be required.\n\n"
f"Detail: {resolve_err}"
) from resolve_err
continue
raise RuntimeError(
"Database migration failed and cannot be auto-recovered. "
f"Manual intervention required.\n\nPrisma error:\n{stderr}"
) from e
raise RuntimeError(
"Database migration failed and cannot be auto-recovered. "
f"Manual intervention required.\n\nPrisma error:\n{stderr}"
) from e
raise RuntimeError(
"Database migration failed after 4 attempts (retry loop "
"exhausted by timeouts or repeated idempotent-recovery "
"continues). Check database connectivity, load, and "
"_prisma_migrations ledger state."
)
finally:
os.chdir(original_dir)
@staticmethod
def setup_database(
use_migrate: bool = False, use_v2_resolver: bool = False
) -> bool:
"""
Set up the database using either prisma migrate or prisma db push
Uses migrations from litellm-proxy-extras package
Args:
schema_path (str): Path to the Prisma schema file
use_migrate (bool): Whether to use prisma migrate instead of db push
use_migrate: Whether to use prisma migrate instead of db push
use_v2_resolver: Opt into the v2 migration resolver (safer during
rolling deploys; does not run the diff-and-force recovery
that causes schema thrashing). Defaults to False for
backwards compatibility.
Returns:
bool: True if setup was successful, False otherwise
"""
if use_v2_resolver:
logger.info("Using v2 migration resolver (--use_v2_migration_resolver)")
return ProxyExtrasDBManager._setup_database_v2(use_migrate=use_migrate)
schema_path = ProxyExtrasDBManager._get_prisma_dir() + "/schema.prisma"
for attempt in range(4):
original_dir = os.getcwd()

View file

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

View file

@ -1,6 +1,6 @@
[project]
name = "litellm-proxy-extras"
version = "0.4.67"
version = "0.4.68"
description = "Additional files for the LiteLLM Proxy. Reduces the size of the main litellm package."
readme = "README.md"
requires-python = ">=3.9"
@ -25,7 +25,7 @@ required-version = "==0.10.9"
module-root = ""
[tool.commitizen]
version = "0.4.67"
version = "0.4.68"
version_files = [
"pyproject.toml:^version",
"../pyproject.toml:litellm-proxy-extras==",

View file

@ -0,0 +1,242 @@
"""Regression tests for ProxyExtrasDBManager v2 migration resolver.
The v2 resolver is opt-in via `--use_v2_migration_resolver` / the
`use_v2_resolver=True` kwarg. These tests exercise the v2 path; the v1
(default) behavior is unchanged from pre-fix.
"""
import subprocess
from unittest.mock import patch
import pytest
from litellm_proxy_extras.utils import (
ProxyExtrasDBManager,
_max_migration_timestamp,
_migration_timestamp,
)
def _fake_migrate_deploy_failure(returncode: int, stderr: str):
def _run(*args, **kwargs):
raise subprocess.CalledProcessError(
returncode=returncode,
cmd=args[0],
stderr=stderr,
output="",
)
return _run
def test_v2_p3018_permission_error_raises_runtime_error(monkeypatch, tmp_path):
"""v2: a permission failure during migrate deploy raises RuntimeError."""
monkeypatch.setenv("DATABASE_URL", "postgresql://u:p@localhost:9/x")
monkeypatch.setattr(
ProxyExtrasDBManager, "_warn_if_db_ahead_of_head", lambda _: None
)
monkeypatch.setattr(ProxyExtrasDBManager, "_get_prisma_dir", lambda: str(tmp_path))
(tmp_path / "schema.prisma").write_text("// stub")
stderr = (
"Error: P3018\nMigration name: 20250326162113_baseline\n"
"Database error code: 42501\npermission denied for schema public"
)
with patch("subprocess.run", side_effect=_fake_migrate_deploy_failure(1, stderr)):
with pytest.raises(RuntimeError, match="permission"):
ProxyExtrasDBManager.setup_database(use_migrate=True, use_v2_resolver=True)
def test_v2_non_idempotent_p3009_raises_runtime_error(monkeypatch, tmp_path):
"""v2: a non-idempotent migration failure raises (no silent recovery)."""
monkeypatch.setenv("DATABASE_URL", "postgresql://u:p@localhost:9/x")
monkeypatch.setattr(
ProxyExtrasDBManager, "_warn_if_db_ahead_of_head", lambda _: None
)
monkeypatch.setattr(ProxyExtrasDBManager, "_get_prisma_dir", lambda: str(tmp_path))
(tmp_path / "schema.prisma").write_text("// stub")
stderr = (
"Error: P3009\nMigration `20260101000000_genuinely_broken` failed\n"
'Reason: syntax error at or near "BRKN" LINE 42'
)
with patch("subprocess.run", side_effect=_fake_migrate_deploy_failure(1, stderr)):
with pytest.raises(RuntimeError, match="cannot be auto-recovered"):
ProxyExtrasDBManager.setup_database(use_migrate=True, use_v2_resolver=True)
def test_strip_prisma_query_params_removes_connection_limit():
"""DATABASE_URLs with Prisma-specific params should be parseable by psycopg."""
url = "postgresql://u:p@h:5432/db?connection_limit=100&pool_timeout=60&sslmode=require"
stripped = ProxyExtrasDBManager._strip_prisma_query_params(url)
assert "connection_limit" not in stripped
assert "pool_timeout" not in stripped
assert "sslmode=require" in stripped
def test_strip_prisma_query_params_passthrough_no_query():
"""URLs without query strings are returned unchanged."""
url = "postgresql://u:p@h:5432/db"
assert ProxyExtrasDBManager._strip_prisma_query_params(url) == url
def test_migration_timestamp_extracts_leading_digits():
assert _migration_timestamp("20260101000000_add_foo") == 20260101000000
assert _migration_timestamp("20250326162113_baseline") == 20250326162113
def test_migration_timestamp_returns_zero_on_malformed():
assert _migration_timestamp("0_init") == 0
assert _migration_timestamp("not_a_migration") == 0
def test_max_migration_timestamp():
names = {"20250326000000_a", "20260415000000_b", "20251115000000_c"}
assert _max_migration_timestamp(names) == 20260415000000
def test_max_migration_timestamp_empty_set():
assert _max_migration_timestamp(set()) == 0
def test_v1_default_still_calls_resolve_all_migrations(monkeypatch, tmp_path):
"""v1 (default) continues to call _resolve_all_migrations on the happy path.
This is the existing buggy behavior we're not fixing it in v1, only
offering v2 as opt-in. This test pins the default so that a future
inadvertent default flip is caught.
"""
monkeypatch.setattr(ProxyExtrasDBManager, "_get_prisma_dir", lambda: str(tmp_path))
(tmp_path / "schema.prisma").write_text("// stub")
# Stub `prisma migrate deploy` to claim success with pending migrations
# applied, which is the code path that triggers the legacy post-migration
# sanity check (a call to _resolve_all_migrations).
class FakeResult:
stdout = "Applied migration.\n"
stderr = ""
def fake_run(cmd, *args, **kwargs):
return FakeResult()
resolve_called = {"n": 0}
def fake_resolve(*args, **kwargs):
resolve_called["n"] += 1
monkeypatch.setattr("subprocess.run", fake_run)
monkeypatch.setattr(ProxyExtrasDBManager, "_resolve_all_migrations", fake_resolve)
ok = ProxyExtrasDBManager.setup_database(use_migrate=True) # v2 flag NOT set
assert ok is True
assert resolve_called["n"] == 1, "v1 default should still invoke the legacy path"
def test_v2_db_push_wraps_subprocess_error_as_runtime_error(monkeypatch, tmp_path):
"""v2: a failing `prisma db push` must raise RuntimeError, not leak
CalledProcessError past proxy_cli.py's `except RuntimeError`."""
monkeypatch.setattr(ProxyExtrasDBManager, "_get_prisma_dir", lambda: str(tmp_path))
(tmp_path / "schema.prisma").write_text("// stub")
stderr = "db push error"
with patch("subprocess.run", side_effect=_fake_migrate_deploy_failure(1, stderr)):
with pytest.raises(RuntimeError, match="prisma db push failed"):
ProxyExtrasDBManager.setup_database(use_migrate=False, use_v2_resolver=True)
def test_v2_warn_ahead_of_head_swallows_db_errors(monkeypatch, tmp_path):
"""_warn_if_db_ahead_of_head must never raise — it's informational.
Non-connection DB errors (e.g. InsufficientPrivilege from a user
without SELECT on _prisma_migrations) must be caught, not propagated.
"""
import psycopg
monkeypatch.setenv("DATABASE_URL", "postgresql://u:p@localhost:9/x")
monkeypatch.setattr(ProxyExtrasDBManager, "_get_prisma_dir", lambda: str(tmp_path))
(tmp_path / "schema.prisma").write_text("// stub")
class _FakeConn:
def __enter__(self):
return self
def __exit__(self, *a):
return False
def execute(self, *a, **kw):
# Simulate an InsufficientPrivilege (subclass of DatabaseError).
raise psycopg.errors.InsufficientPrivilege("permission denied")
def _fake_connect(*a, **kw):
return _FakeConn()
monkeypatch.setattr("psycopg.connect", _fake_connect)
# Must not raise.
ProxyExtrasDBManager._warn_if_db_ahead_of_head(str(tmp_path))
def test_v2_resolve_specific_migration_failure_raises_runtime_error(
monkeypatch, tmp_path
):
"""If marking a migration as applied fails inside P3009 idempotent
recovery, the subprocess error must be re-raised as RuntimeError so
proxy_cli.py catches it cleanly (instead of leaking CalledProcessError)."""
monkeypatch.setattr(
ProxyExtrasDBManager, "_warn_if_db_ahead_of_head", lambda _: None
)
monkeypatch.setattr(ProxyExtrasDBManager, "_get_prisma_dir", lambda: str(tmp_path))
(tmp_path / "schema.prisma").write_text("// stub")
monkeypatch.setattr(
ProxyExtrasDBManager, "_roll_back_migration", lambda *a, **kw: None
)
# First call: migrate deploy -> P3009 idempotent error.
# Recovery path tries _resolve_specific_migration; that also raises.
def _failing_resolve(*a, **kw):
raise subprocess.CalledProcessError(
returncode=1,
cmd="prisma migrate resolve --applied",
stderr="resolve failed",
output="",
)
monkeypatch.setattr(
ProxyExtrasDBManager, "_resolve_specific_migration", _failing_resolve
)
stderr = (
"Error: P3009\nMigration `20260101000000_some_migration` failed\n"
"relation already exists"
)
with patch("subprocess.run", side_effect=_fake_migrate_deploy_failure(1, stderr)):
with pytest.raises(
RuntimeError, match="Failed to mark migration .* as applied"
):
ProxyExtrasDBManager.setup_database(use_migrate=True, use_v2_resolver=True)
def test_v2_does_not_call_resolve_all_migrations(monkeypatch, tmp_path):
"""v2 must never call _resolve_all_migrations — that's the bug it fixes."""
monkeypatch.setattr(
ProxyExtrasDBManager, "_warn_if_db_ahead_of_head", lambda _: None
)
monkeypatch.setattr(ProxyExtrasDBManager, "_get_prisma_dir", lambda: str(tmp_path))
(tmp_path / "schema.prisma").write_text("// stub")
class FakeResult:
stdout = "Applied migration.\n"
stderr = ""
monkeypatch.setattr("subprocess.run", lambda *a, **kw: FakeResult())
resolve_called = {"n": 0}
monkeypatch.setattr(
ProxyExtrasDBManager,
"_resolve_all_migrations",
lambda *a, **kw: resolve_called.__setitem__("n", resolve_called["n"] + 1),
)
ok = ProxyExtrasDBManager.setup_database(use_migrate=True, use_v2_resolver=True)
assert ok is True
assert resolve_called["n"] == 0, "v2 must not invoke the diff-and-force recovery"

View file

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

View file

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

View file

@ -164,6 +164,7 @@ MCP_STDIO_ALLOWED_COMMANDS: frozenset = frozenset(
LITELLM_UI_ALLOW_HEADERS = [
"x-litellm-semantic-filter",
"x-litellm-semantic-filter-tools",
"x-litellm-adaptive-router-model",
]
# Gemini model-specific minimal thinking budget constants

View file

@ -2,6 +2,7 @@ from datetime import datetime
from typing import (
TYPE_CHECKING,
Any,
ClassVar,
Dict,
List,
Literal,
@ -12,6 +13,7 @@ from typing import (
)
from litellm._logging import verbose_logger
from litellm.litellm_core_utils.core_helpers import redact_nested_match_and_regex_keys
from litellm.caching import DualCache
from litellm.integrations.custom_logger import CustomLogger
from litellm.types.guardrails import (
@ -81,6 +83,9 @@ class ModifyResponseException(Exception):
class CustomGuardrail(CustomLogger):
# If True, during_call runs async_moderation_hook instead of the unified apply_guardrail path.
use_native_during_call_hook: ClassVar[bool] = False
def __init__(
self,
guardrail_name: Optional[str] = None,
@ -637,6 +642,13 @@ class CustomGuardrail(CustomLogger):
if isinstance(item, dict):
item.pop("secret_fields", None)
# Default-safe behavior: never persist raw matched spans in standard
# guardrail logging payloads (single shared implementation; Bedrock hooks pass
# raw provider JSON so redaction is not duplicated upstream).
clean_guardrail_response = redact_nested_match_and_regex_keys(
clean_guardrail_response
)
slg = StandardLoggingGuardrailInformation(
guardrail_name=self.guardrail_name,
guardrail_provider=guardrail_provider,

View file

@ -1615,6 +1615,14 @@ class OpenTelemetry(CustomLogger):
value=response_id,
)
litellm_call_id = standard_logging_payload.get("litellm_call_id")
if litellm_call_id:
self.safe_set_attribute(
span=span,
key="litellm.call_id",
value=litellm_call_id,
)
# The model used to generate the response.
if response_obj and response_obj.get("model"):
self.safe_set_attribute(
@ -2281,6 +2289,10 @@ class OpenTelemetry(CustomLogger):
# Remove trailing slash
endpoint = endpoint.rstrip("/")
# Splunk Observability Cloud OTLP/HTTP uses /v2/trace/otlp (not /v1/traces). Do not rewrite.
if signal_type == "traces" and "/v2/trace/otlp" in endpoint:
return endpoint
# Check if endpoint already ends with the correct signal path
target_path = f"/v1/{signal_type}"
if endpoint.endswith(target_path):

View file

@ -1,5 +1,6 @@
# What is this?
## Helper utilities
import copy
from typing import TYPE_CHECKING, Any, Iterable, List, Literal, Optional, Union
import httpx
@ -435,3 +436,42 @@ def filter_internal_params(
# Filter out internal parameters
return {k: v for k, v in data.items() if k not in internal_params}
def redact_nested_match_and_regex_keys(
payload: Union[dict, List[Any], str, None],
) -> Union[dict, List[Any], str, None]:
"""
Deep-copy `payload` and replace every `match` / `regex` string field with
"[REDACTED]" anywhere in nested dict/list structures.
Used for guardrail spend/compliance logging so raw spans are not persisted.
"""
if payload is None or isinstance(payload, str):
return payload
try:
redacted: Union[dict, List[Any], str, None] = copy.deepcopy(payload)
except Exception:
return payload
# Iterative traversal; `seen` guards against cyclic refs preserved by deepcopy.
try:
seen: set = set()
stack: List[Any] = [redacted]
while stack:
node = stack.pop()
node_id = id(node)
if node_id in seen:
continue
seen.add(node_id)
if isinstance(node, dict):
if "match" in node:
node["match"] = "[REDACTED]"
if "regex" in node:
node["regex"] = "[REDACTED]"
stack.extend(node.values())
elif isinstance(node, list):
stack.extend(node)
except Exception:
return payload
return redacted

View file

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

View file

@ -14,6 +14,8 @@ class AzureImageEditConfig(OpenAIImageEditConfig):
headers: dict,
model: str,
api_key: Optional[str] = None,
litellm_params: Optional[dict] = None,
api_base: Optional[str] = None,
) -> dict:
api_key = (
api_key

View file

@ -65,6 +65,8 @@ class AzureFoundryFlux2ImageEditConfig(OpenAIImageEditConfig):
headers: dict,
model: str,
api_key: Optional[str] = None,
litellm_params: Optional[dict] = None,
api_base: Optional[str] = None,
) -> dict:
"""
Validate Azure AI Foundry environment and set up authentication

View file

@ -25,6 +25,8 @@ class AzureFoundryFluxImageEditConfig(OpenAIImageEditConfig):
headers: dict,
model: str,
api_key: Optional[str] = None,
litellm_params: Optional[dict] = None,
api_base: Optional[str] = None,
) -> dict:
"""
Validate Azure AI Foundry environment and set up authentication

View file

@ -67,6 +67,8 @@ class BaseImageEditConfig(ABC):
headers: dict,
model: str,
api_key: Optional[str] = None,
litellm_params: Optional[dict] = None,
api_base: Optional[str] = None,
) -> dict:
return {}

View file

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

View file

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

View file

@ -483,6 +483,8 @@ class BedrockAmazonNovaCanvasImageEditConfig(BaseImageEditConfig):
headers: dict,
model: str,
api_key: Optional[str] = None,
litellm_params: Optional[dict] = None,
api_base: Optional[str] = None,
) -> dict:
if headers is None:
headers = {}

View file

@ -372,6 +372,8 @@ class BedrockStabilityImageEditConfig(BaseImageEditConfig):
headers: dict,
model: str,
api_key: Optional[str] = None,
litellm_params: Optional[dict] = None,
api_base: Optional[str] = None,
) -> dict:
"""
Validate environment for Bedrock Stability image edit.

View file

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

View file

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

View file

@ -123,6 +123,8 @@ class BlackForestLabsImageEditConfig(BaseImageEditConfig):
headers: dict,
model: str,
api_key: Optional[str] = None,
litellm_params: Optional[dict] = None,
api_base: Optional[str] = None,
) -> dict:
"""
Validate environment and set up headers for Black Forest Labs.

View file

@ -5515,6 +5515,8 @@ class BaseLLMHTTPHandler:
api_key=litellm_params.api_key,
headers=image_edit_optional_request_params.get("extra_headers", {}) or {},
model=model,
litellm_params=dict(litellm_params),
api_base=litellm_params.api_base,
)
if extra_headers:
@ -5611,6 +5613,8 @@ class BaseLLMHTTPHandler:
api_key=litellm_params.api_key,
headers=image_edit_optional_request_params.get("extra_headers", {}) or {},
model=model,
litellm_params=dict(litellm_params),
api_base=litellm_params.api_base,
)
if extra_headers:

View file

@ -54,6 +54,8 @@ class GeminiImageEditConfig(BaseImageEditConfig):
headers: dict,
model: str,
api_key: Optional[str] = None,
litellm_params: Optional[dict] = None,
api_base: Optional[str] = None,
) -> dict:
final_api_key: Optional[str] = api_key or get_secret_str("GEMINI_API_KEY")
if not final_api_key:

View file

@ -8,7 +8,12 @@ class LiteLLMProxyImageEditConfig(OpenAIImageEditConfig):
"""Configuration for image edit requests routed through LiteLLM Proxy."""
def validate_environment(
self, headers: dict, model: str, api_key: Optional[str] = None
self,
headers: dict,
model: str,
api_key: Optional[str] = None,
litellm_params: Optional[dict] = None,
api_base: Optional[str] = None,
) -> dict:
api_key = api_key or get_secret_str("LITELLM_PROXY_API_KEY")
headers.update({"Authorization": f"Bearer {api_key}"})

View file

@ -165,6 +165,8 @@ class OpenAIImageEditConfig(BaseImageEditConfig):
headers: dict,
model: str,
api_key: Optional[str] = None,
litellm_params: Optional[dict] = None,
api_base: Optional[str] = None,
) -> dict:
api_key = (
api_key

View file

@ -116,6 +116,8 @@ class OpenRouterImageEditConfig(BaseImageEditConfig):
headers: dict,
model: str,
api_key: Optional[str] = None,
litellm_params: Optional[dict] = None,
api_base: Optional[str] = None,
) -> dict:
api_key = api_key or litellm.api_key or get_secret_str("OPENROUTER_API_KEY")
if not api_key:

View file

@ -81,6 +81,8 @@ class RecraftImageEditConfig(BaseImageEditConfig):
headers: dict,
model: str,
api_key: Optional[str] = None,
litellm_params: Optional[dict] = None,
api_base: Optional[str] = None,
) -> dict:
final_api_key: Optional[str] = api_key or get_secret_str("RECRAFT_API_KEY")
if not final_api_key:

View file

@ -149,6 +149,8 @@ class StabilityImageEditConfig(BaseImageEditConfig):
headers: dict,
model: str,
api_key: Optional[str] = None,
litellm_params: Optional[dict] = None,
api_base: Optional[str] = None,
) -> dict:
"""
Validate environment and set up headers for Stability AI.

View file

@ -103,10 +103,24 @@ class VertexAIImagenImageEditConfig(BaseImageEditConfig, VertexLLM):
headers: dict,
model: str,
api_key: Optional[str] = None,
litellm_params: Optional[dict] = None,
api_base: Optional[str] = None,
) -> dict:
headers = headers or {}
vertex_project = self._resolve_vertex_project()
vertex_credentials = self._resolve_vertex_credentials()
litellm_params = litellm_params or {}
_api_base = litellm_params.get("api_base") or api_base
if _api_base is not None:
return headers
vertex_project = (
self.safe_get_vertex_ai_project(litellm_params)
or self._resolve_vertex_project()
)
vertex_credentials = (
self.safe_get_vertex_ai_credentials(litellm_params)
or self._resolve_vertex_credentials()
)
access_token, _ = self._ensure_access_token(
credentials=vertex_credentials,
project_id=vertex_project,
@ -123,8 +137,14 @@ class VertexAIImagenImageEditConfig(BaseImageEditConfig, VertexLLM):
"""
Get the complete URL for Vertex AI Imagen predict API
"""
vertex_project = self._resolve_vertex_project()
vertex_location = self._resolve_vertex_location()
vertex_project = (
self.safe_get_vertex_ai_project(litellm_params)
or self._resolve_vertex_project()
)
vertex_location = (
self.safe_get_vertex_ai_location(litellm_params)
or self._resolve_vertex_location()
)
if not vertex_project or not vertex_location:
raise ValueError(

View file

@ -22872,6 +22872,22 @@
"supports_video_input": true,
"supports_vision": true
},
"moonshot/kimi-k2.6": {
"cache_read_input_token_cost": 1.6e-07,
"input_cost_per_token": 9.5e-07,
"litellm_provider": "moonshot",
"max_input_tokens": 262144,
"max_output_tokens": 262144,
"max_tokens": 262144,
"mode": "chat",
"output_cost_per_token": 4e-06,
"source": "https://platform.kimi.ai/docs/pricing/chat-k26",
"supports_function_calling": true,
"supports_reasoning": true,
"supports_tool_choice": true,
"supports_video_input": true,
"supports_vision": true
},
"moonshot/kimi-latest": {
"cache_read_input_token_cost": 1.5e-07,
"input_cost_per_token": 2e-06,

View file

@ -323,6 +323,14 @@ async def authorize_with_server(
)
parsed = urlparse(redirect_uri)
if parsed.scheme not in ("http", "https"):
raise HTTPException(
status_code=400,
detail={
"error": "invalid_redirect_uri",
"message": "redirect_uri must use http or https scheme",
},
)
base_url = urlunparse(parsed._replace(query=""))
request_base_url = get_request_base_url(request)
encoded_state = encode_state_with_base_url(

View file

@ -1,42 +1,83 @@
# model_list:
# - model_name: claude-sonnet-4-6
# litellm_params: {model: anthropic/claude-sonnet-4-6}
# model_info:
# litellm_routing_preferences:
# quality_tier: 1
# keywords: [tin]
# - model_name: gpt-4o-mini
# litellm_params: {model: openai/gpt-4o-mini}
# model_info:
# litellm_routing_preferences:
# quality_tier: 1
# keywords: []
# - model_name: gpt-4o
# litellm_params: {model: openai/gpt-4o}
# model_info:
# litellm_routing_preferences:
# quality_tier: 2
# keywords: [vision, function_calling]
# - model_name: opus
# litellm_params: {model: anthropic/claude-opus-4-7}
# model_info:
# litellm_routing_preferences:
# quality_tier: 3
# keywords: ["architecture", "design"]
# - model_name: my-quality-router
# litellm_params:
# model: auto_router/adaptive_router
# adaptive_router_default_model: gpt-4o-mini
# adaptive_router_config:
# available_models: [gpt-4o-mini, gpt-4o, opus, claude-sonnet-4-6]
# Example proxy config for the adaptive router (v0).
#
# Wires one logical router ("smart-cheap-router") that adaptively picks between
# two real deployments ("fast" and "smart") based on per-session feedback signals.
#
# How to use from a client:
# POST /v1/chat/completions { "model": "smart-cheap-router", ... }
# Add { "metadata": { "litellm_session_id": "<your-session-id>" } } to enable
# sticky-session routing within a conversation.
#
# Required env vars: OPENAI_API_KEY, DATABASE_URL.
model_list:
# OpenAI model for /v1/chat/completions test — 200x custom pricing
- model_name: "gpt-4.1-mini"
# ---- The adaptive router "control" deployment -------------------------
# `model_name` is what clients call. `available_models` lists the underlying
# deployments the router is allowed to pick from (must match other model_name
# entries in this list).
- model_name: smart-cheap-router
litellm_params:
model: openai/gpt-4.1-mini
api_key: os.environ/OPENAI_API_KEY
model_info:
id: gpt-4.1-mini-custom-pricing
input_cost_per_token: 0.00004 # 100x standard ($0.40/1M = $0.0000004)
output_cost_per_token: 0.00016 # 100x standard ($1.60/1M = $0.0000016)
model: auto_router/adaptive_router
adaptive_router_config:
available_models: ["fast", "smart"]
weights:
quality: 0.7
cost: 0.3
# OpenAI model for /v1/responses test — 100x custom pricing
- model_name: "gpt-5"
litellm_params:
model: openai/gpt-5
api_key: os.environ/OPENAI_API_KEY
model_info:
id: gpt-5-custom-pricing
mode: "chat"
input_cost_per_token: 125 # 100x standard ($1.25/1M = $0.00000125)
output_cost_per_token: 10 # 100x standard ($10.00/1M = $0.00001)
# Anthropic model for /v1/messages test — 100x custom pricing
- model_name: "claude-sonnet-4-6"
# ---- Underlying deployments the router picks from ---------------------
- model_name: fast
litellm_params:
model: anthropic/claude-sonnet-4-6
api_key: os.environ/ANTHROPIC_API_KEY
input_cost_per_token: 0.00000015
model_info:
id: claude-sonnet-4-custom-pricing
input_cost_per_token: 0.0003 # 100x standard ($0.000003)
output_cost_per_token: 0.0015 # 100x standard ($0.000015)
- model_name: my-auto
adaptive_router_preferences:
quality_tier: 2
strengths: []
- model_name: smart
litellm_params:
model: auto_router/complexity_router
complexity_router_config:
tiers:
SIMPLE: "gpt-4.1-mini"
COMPLEX: claude-sonnet-4-6
tier_boundaries:
simple_medium: 0.30
complexity_router_default_model: small-model
model: anthropic/claude-opus-4-7
api_key: os.environ/ANTHROPIC_API_KEY
input_cost_per_token: 0.0000050
model_info:
adaptive_router_preferences:
quality_tier: 3
strengths: ["code_generation", "technical_design", "analytical_reasoning"]
litellm_settings:
drop_params: True
general_settings:
master_key: sk-1234 # REPLACE in production

View file

@ -1997,7 +1997,12 @@ class TeamRequest(LiteLLMPydanticObjectBase):
class LiteLLM_BudgetTable(LiteLLMPydanticObjectBase):
"""Represents user-controllable params for a LiteLLM_BudgetTable record"""
"""Represents user-controllable params for a LiteLLM_BudgetTable record.
Budget-write paths use `model_fields.keys()` on this class as an allowlist
for user input. Keep server-managed fields (e.g. `budget_reset_at`) on
`LiteLLM_BudgetTableFull` so they aren't user-settable.
"""
budget_id: Optional[str] = None
soft_budget: Optional[float] = None
@ -2015,7 +2020,7 @@ class LiteLLM_BudgetTable(LiteLLMPydanticObjectBase):
class LiteLLM_BudgetTableFull(LiteLLM_BudgetTable):
"""Represents all params for a LiteLLM_BudgetTable record"""
"""LiteLLM_BudgetTable + server-managed fields returned on API responses."""
budget_reset_at: Optional[datetime] = None
created_at: datetime
@ -3695,7 +3700,11 @@ class LiteLLM_TeamMembership(LiteLLMPydanticObjectBase):
team_id: str
budget_id: Optional[str] = None
spend: Optional[float] = 0.0
litellm_budget_table: Optional[LiteLLM_BudgetTable]
total_spend: Optional[float] = 0.0
# Union so Pydantic picks Full when data has server-managed fields
# (/team/info) and Base when callers/tests construct with only
# user-settable fields.
litellm_budget_table: Optional[Union[LiteLLM_BudgetTableFull, LiteLLM_BudgetTable]]
def safe_get_team_member_rpm_limit(self) -> Optional[int]:
if self.litellm_budget_table is not None:

View file

@ -626,11 +626,17 @@ async def common_checks( # noqa: PLR0915
and user_object.max_budget is not None
):
user_budget = user_object.max_budget
if user_budget < user_object.spend:
from litellm.proxy.proxy_server import get_current_spend
user_spend = await get_current_spend(
counter_key=f"spend:user:{user_object.user_id}",
fallback_spend=user_object.spend or 0.0,
)
if user_spend >= user_budget:
raise litellm.BudgetExceededError(
current_cost=user_object.spend,
current_cost=user_spend,
max_budget=user_budget,
message=f"ExceededBudget: User={user_object.user_id} over budget. Spend={user_object.spend}, Budget={user_budget}",
message=f"ExceededBudget: User={user_object.user_id} over budget. Spend={user_spend}, Budget={user_budget}",
)
## 4.2 check team member budget, if team key
@ -3665,12 +3671,20 @@ async def _organization_max_budget_check(
if org_max_budget is None or org_max_budget <= 0:
return
# Read spend from cross-pod counter (Redis-first) or cached object (fallback)
from litellm.proxy.proxy_server import get_current_spend
org_spend = await get_current_spend(
counter_key=f"spend:org:{org_id}",
fallback_spend=org_table.spend or 0.0,
)
# Check if organization spend exceeds max budget
if org_table.spend >= org_max_budget:
if org_spend >= org_max_budget:
# Trigger budget alert
call_info = CallInfo(
token=valid_token.token,
spend=org_table.spend,
spend=org_spend,
max_budget=org_max_budget,
user_id=valid_token.user_id,
team_id=valid_token.team_id,
@ -3686,9 +3700,9 @@ async def _organization_max_budget_check(
)
raise litellm.BudgetExceededError(
current_cost=org_table.spend,
current_cost=org_spend,
max_budget=org_max_budget,
message=f"Budget has been exceeded! Organization={org_id} Current cost: {org_table.spend}, Max budget: {org_max_budget}",
message=f"Budget has been exceeded! Organization={org_id} Current cost: {org_spend}, Max budget: {org_max_budget}",
)

View file

@ -433,7 +433,8 @@ def add_guardrail_to_applied_guardrails_header(
return
_metadata = request_data.get("metadata", None) or {}
if "applied_guardrails" in _metadata:
_metadata["applied_guardrails"].append(guardrail_name)
if guardrail_name not in _metadata["applied_guardrails"]:
_metadata["applied_guardrails"].append(guardrail_name)
else:
_metadata["applied_guardrails"] = [guardrail_name]
# Ensure metadata is set back to request_data (important when metadata didn't exist)

View file

@ -1300,7 +1300,10 @@ class DBSpendUpdateWriter:
batcher.litellm_teammembership.update_many( # 'update_many' prevents error from being raised if no row exists
where={"team_id": team_id, "user_id": user_id},
data={"spend": {"increment": response_cost}},
data={
"spend": {"increment": response_cost},
"total_spend": {"increment": response_cost},
},
)
# Transaction succeeded, break out of retry loop
break

View file

@ -22,6 +22,7 @@ from litellm.constants import (
)
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.proxy._types import (
BaseDailySpendTransaction,
DailyAgentSpendTransaction,
DailyEndUserSpendTransaction,
DailyOrganizationSpendTransaction,
@ -29,6 +30,8 @@ from litellm.proxy._types import (
DailyTeamSpendTransaction,
DailyUserSpendTransaction,
DBSpendUpdateTransactions,
Litellm_EntityType,
SpendUpdateQueueItem,
)
from litellm.proxy.db.db_transaction_queue.base_update_queue import service_logger_obj
from litellm.proxy.db.db_transaction_queue.daily_spend_update_queue import (
@ -259,9 +262,36 @@ class RedisUpdateBuffer:
if len(rpush_list) == 0:
return
result_lengths = await self.redis_cache.async_rpush_pipeline(
rpush_list=rpush_list,
)
try:
result_lengths = await self.redis_cache.async_rpush_pipeline(
rpush_list=rpush_list,
)
except Exception as e:
# The in-memory queues were already drained above. If we let the
# exception propagate without restoring, the aggregated spend is
# permanently lost. Re-enqueue so the next scheduler tick retries.
verbose_proxy_logger.error(
"Spend tracking - failed to push aggregated spend updates to Redis. "
"Restoring %d transaction sets to in-memory queues for retry on next tick. "
"Error: %s",
len(rpush_list),
str(e),
)
await self._restore_spend_updates_to_in_memory_queues(
db_spend_update_transactions=db_spend_update_transactions,
daily_spend_update_transactions=daily_spend_update_transactions,
daily_team_spend_update_transactions=daily_team_spend_update_transactions,
daily_org_spend_update_transactions=daily_org_spend_update_transactions,
daily_end_user_spend_update_transactions=daily_end_user_spend_update_transactions,
daily_agent_spend_update_transactions=daily_agent_spend_update_transactions,
spend_update_queue=spend_update_queue,
daily_spend_update_queue=daily_spend_update_queue,
daily_team_spend_update_queue=daily_team_spend_update_queue,
daily_org_spend_update_queue=daily_org_spend_update_queue,
daily_end_user_spend_update_queue=daily_end_user_spend_update_queue,
daily_agent_spend_update_queue=daily_agent_spend_update_queue,
)
return
# Emit gauge events for each queue
for i, queue_size in enumerate(result_lengths):
@ -271,6 +301,101 @@ class RedisUpdateBuffer:
service=service_types[i],
)
@staticmethod
async def _restore_spend_updates_to_in_memory_queues(
db_spend_update_transactions: Optional[DBSpendUpdateTransactions],
daily_spend_update_transactions: Optional[Dict[str, BaseDailySpendTransaction]],
daily_team_spend_update_transactions: Optional[
Dict[str, BaseDailySpendTransaction]
],
daily_org_spend_update_transactions: Optional[
Dict[str, BaseDailySpendTransaction]
],
daily_end_user_spend_update_transactions: Optional[
Dict[str, BaseDailySpendTransaction]
],
daily_agent_spend_update_transactions: Optional[
Dict[str, BaseDailySpendTransaction]
],
spend_update_queue: SpendUpdateQueue,
daily_spend_update_queue: DailySpendUpdateQueue,
daily_team_spend_update_queue: DailySpendUpdateQueue,
daily_org_spend_update_queue: DailySpendUpdateQueue,
daily_end_user_spend_update_queue: DailySpendUpdateQueue,
daily_agent_spend_update_queue: DailySpendUpdateQueue,
) -> None:
"""
Put drained-but-unpushed transactions back into in-memory queues.
Called when the Redis rpush pipeline raises. Without this, all spend
data aggregated during the current scheduler tick is permanently lost
because the source queues were already drained before the rpush.
"""
if db_spend_update_transactions is not None:
entity_entries: List[
Tuple[Litellm_EntityType, Optional[Dict[str, float]]]
] = [
(
Litellm_EntityType.USER,
db_spend_update_transactions.get("user_list_transactions"),
),
(
Litellm_EntityType.END_USER,
db_spend_update_transactions.get("end_user_list_transactions"),
),
(
Litellm_EntityType.KEY,
db_spend_update_transactions.get("key_list_transactions"),
),
(
Litellm_EntityType.TEAM,
db_spend_update_transactions.get("team_list_transactions"),
),
(
Litellm_EntityType.TEAM_MEMBER,
db_spend_update_transactions.get("team_member_list_transactions"),
),
(
Litellm_EntityType.ORGANIZATION,
db_spend_update_transactions.get("org_list_transactions"),
),
(
Litellm_EntityType.TAG,
db_spend_update_transactions.get("tag_list_transactions"),
),
(
Litellm_EntityType.AGENT,
db_spend_update_transactions.get("agent_list_transactions"),
),
]
for entity_type, entities in entity_entries:
if not entities:
continue
for entity_id, cost in entities.items():
await spend_update_queue.add_update(
SpendUpdateQueueItem(
entity_type=entity_type,
entity_id=entity_id,
response_cost=cost,
)
)
daily_pairs: List[
Tuple[Optional[Dict[str, BaseDailySpendTransaction]], DailySpendUpdateQueue]
] = [
(daily_spend_update_transactions, daily_spend_update_queue),
(daily_team_spend_update_transactions, daily_team_spend_update_queue),
(daily_org_spend_update_transactions, daily_org_spend_update_queue),
(
daily_end_user_spend_update_transactions,
daily_end_user_spend_update_queue,
),
(daily_agent_spend_update_transactions, daily_agent_spend_update_queue),
]
for daily_txns, daily_queue in daily_pairs:
if daily_txns:
await daily_queue.update_queue.put(daily_txns)
@staticmethod
def _number_of_transactions_to_store_in_redis(
db_spend_update_transactions: DBSpendUpdateTransactions,

View file

@ -403,10 +403,18 @@ class PrismaManager:
return dname
@staticmethod
def setup_database(use_migrate: bool = False) -> bool:
def setup_database(
use_migrate: bool = False, use_v2_resolver: bool = False
) -> bool:
"""
Set up the database using either prisma migrate or prisma db push
Args:
use_migrate: Use `prisma migrate deploy` instead of `db push`.
use_v2_resolver: Opt into the v2 migration resolver that avoids
the diff-and-force recovery behavior (which caused schema
thrashing during rolling deploys). Defaults to False.
Returns:
bool: True if setup was successful, False otherwise
"""
@ -427,7 +435,10 @@ class PrismaManager:
prisma_dir = PrismaManager._get_prisma_dir()
return ProxyExtrasDBManager.setup_database(use_migrate=use_migrate)
return ProxyExtrasDBManager.setup_database(
use_migrate=use_migrate,
use_v2_resolver=use_v2_resolver,
)
else:
# Use prisma db push with increased timeout
subprocess.run(

View file

@ -0,0 +1,52 @@
# Example proxy config for the adaptive router (v0).
#
# Wires one logical router ("smart-cheap-router") that adaptively picks between
# two real deployments ("fast" and "smart") based on per-session feedback signals.
#
# How to use from a client:
# POST /v1/chat/completions { "model": "smart-cheap-router", ... }
# Add { "metadata": { "litellm_session_id": "<your-session-id>" } } to enable
# sticky-session routing within a conversation.
#
# Required env vars: OPENAI_API_KEY, DATABASE_URL.
model_list:
# ---- The adaptive router "control" deployment -------------------------
# `model_name` is what clients call. `available_models` lists the underlying
# deployments the router is allowed to pick from (must match other model_name
# entries in this list).
- model_name: smart-cheap-router
litellm_params:
model: auto_router/adaptive_router # required prefix -- triggers adaptive-router init
adaptive_router_config:
available_models: ["fast", "smart"]
weights:
quality: 0.7
cost: 0.3
# ---- Underlying deployments the router picks from ---------------------
- model_name: fast
litellm_params:
model: openai/gpt-4o-mini
api_key: os.environ/OPENAI_API_KEY
input_cost_per_token: 0.00000015
model_info:
adaptive_router_preferences:
quality_tier: 2
strengths: []
- model_name: smart
litellm_params:
model: openai/gpt-4o
api_key: os.environ/OPENAI_API_KEY
input_cost_per_token: 0.0000050
model_info:
adaptive_router_preferences:
quality_tier: 3
strengths: ["code_generation", "technical_design", "analytical_reasoning"]
litellm_settings:
drop_params: True
general_settings:
master_key: sk-1234 # REPLACE in production

View file

@ -5,7 +5,6 @@
# +-------------------------------------------------------------+
# Thank you users! We ❤️ you! - Krrish & Ishaan
import copy
import os
import sys
@ -18,6 +17,7 @@ from typing import (
TYPE_CHECKING,
Any,
AsyncGenerator,
ClassVar,
Dict,
List,
Literal,
@ -33,6 +33,7 @@ from fastapi import HTTPException
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.litellm_core_utils.core_helpers import redact_nested_match_and_regex_keys
from litellm.caching import DualCache
from litellm.exceptions import GuardrailInterventionNormalStringError
from litellm.integrations.custom_guardrail import CustomGuardrail
@ -62,6 +63,7 @@ from litellm.types.utils import (
CallTypesLiteral,
Choices,
GuardrailStatus,
Message,
ModelResponse,
ModelResponseStream,
StreamingChoices,
@ -78,56 +80,33 @@ class GuardrailMessageFilterResult(NamedTuple):
def _redact_pii_matches(response_json: dict) -> dict:
try:
# Create a deep copy to avoid modifying the original response
redacted_response = copy.deepcopy(response_json)
"""
Redact match-like fields from a Bedrock ApplyGuardrail JSON payload.
# Get assessments from the response
# NOTE: We use `.get("key") or []` instead of `.get("key", [])` because
# the Bedrock API can return explicit `null` for list fields (e.g. "regexes": null).
# In Python, dict.get("key", []) returns None (not []) when the key exists
# with a None/null value. The `or []` ensures we always get an iterable,
# preventing "TypeError: 'NoneType' object is not iterable".
assessments = redacted_response.get("assessments") or []
if not assessments:
return redacted_response
Delegates to :func:`redact_nested_match_and_regex_keys` (same rules as spend
logging). Kept as a Bedrock-module entry point for existing unit tests.
"""
redacted = redact_nested_match_and_regex_keys(response_json)
return redacted if isinstance(redacted, dict) else response_json
for assessment in assessments:
# Redact PII entities in sensitive information policy
sensitive_info_policy = assessment.get("sensitiveInformationPolicy")
if sensitive_info_policy:
pii_entities = sensitive_info_policy.get("piiEntities") or []
for pii_entity in pii_entities:
if "match" in pii_entity:
pii_entity["match"] = "[REDACTED]"
# Redact regex matches
regexes = sensitive_info_policy.get("regexes") or []
for regex_match in regexes:
if "match" in regex_match:
regex_match["match"] = "[REDACTED]"
def _redact_assessment_match_fields(assessments: List[dict]) -> List[dict]:
"""
Redact sensitive match-like fields from blocked assessment summaries.
# Redact custom word matches in word policy
word_policy = assessment.get("wordPolicy")
if word_policy:
custom_words = word_policy.get("customWords") or []
for custom_word in custom_words:
if "match" in custom_word:
custom_word["match"] = "[REDACTED]"
managed_words = word_policy.get("managedWordLists") or []
for managed_word in managed_words:
if "match" in managed_word:
managed_word["match"] = "[REDACTED]"
return redacted_response
except Exception as e:
# We do not want to fail in any case so this is just a warning
verbose_proxy_logger.warning("Guardrail log redaction failed: %s", str(e))
return response_json
This is used for customer-visible error payloads (HTTPException.detail) where
we want to preserve policy/type/action metadata without echoing raw matched
content.
"""
redacted = redact_nested_match_and_regex_keys(assessments)
return redacted if isinstance(redacted, list) else assessments
class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
# During-call must use async_moderation_hook (not unified apply_guardrail), otherwise
# OpenAI translation always passes input_type="request" and spend/UI show PRE-CALL.
use_native_during_call_hook: ClassVar[bool] = True
def __init__(
self,
guardrailIdentifier: Optional[str] = None,
@ -418,6 +397,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
messages: Optional[List[AllMessageValues]] = None,
response: Optional[Union[Any, litellm.ModelResponse]] = None,
request_data: Optional[dict] = None,
logging_event_type: Optional[GuardrailEventHooks] = None,
) -> BedrockGuardrailResponse:
from datetime import datetime
@ -455,11 +435,17 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
prepared_request.headers,
)
event_type = (
GuardrailEventHooks.pre_call
if source == "INPUT"
else GuardrailEventHooks.post_call
)
# UI / spend logs use event_type. Bedrock's `source` is INPUT vs OUTPUT for the API
# body, which must not be confused with the proxy hook (pre_call / during_call /
# post_call). When omitted, keep legacy mapping for backward compatibility.
if logging_event_type is not None:
event_type = logging_event_type
else:
event_type = (
GuardrailEventHooks.pre_call
if source == "INPUT"
else GuardrailEventHooks.post_call
)
try:
httpx_response = await self.async_handler.post(
@ -514,9 +500,12 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
#########################################################
# Add guardrail information to request trace
#########################################################
_json_response = httpx_response.json()
# Raw Bedrock JSON is passed here; match/regex redaction runs once inside
# CustomGuardrail.add_standard_logging_guardrail_information_to_request_data.
self.add_standard_logging_guardrail_information_to_request_data(
guardrail_provider=self.guardrail_provider,
guardrail_json_response=httpx_response.json(),
guardrail_json_response=_json_response,
request_data=request_data or {},
guardrail_status=self._get_bedrock_guardrail_response_status(
response=httpx_response
@ -529,9 +518,10 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
#########################################################
if httpx_response.status_code == 200:
# check if the response was flagged
_json_response = httpx_response.json()
redacted_response = _redact_pii_matches(_json_response)
verbose_proxy_logger.debug("Bedrock AI response : %s", redacted_response)
verbose_proxy_logger.debug(
"Bedrock AI response : %s",
redact_nested_match_and_regex_keys(_json_response),
)
bedrock_guardrail_response = BedrockGuardrailResponse(**_json_response)
if self._should_raise_guardrail_blocked_exception(
bedrock_guardrail_response
@ -808,7 +798,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
assessments = self._extract_blocked_assessments(response)
if assessments:
detail["assessments"] = assessments
detail["assessments"] = _redact_assessment_match_fields(assessments)
return HTTPException(status_code=400, detail=detail)
@ -830,8 +820,8 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
return False
# Check assessments to determine if any actions were BLOCKED (vs ANONYMIZED)
# NOTE: Use `or []` instead of default param to handle explicit null from Bedrock API.
# See _redact_pii_matches() for detailed explanation of the null safety pattern.
# NOTE: Use `.get("k") or []` not `.get("k", [])` — Bedrock can return explicit
# JSON null; dict.get("k", []) then yields None, and `for x in None` raises.
assessments = response.get("assessments") or []
if not assessments:
return False
@ -951,7 +941,10 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
)
try:
bedrock_guardrail_response = await self.make_bedrock_api_request(
source="INPUT", messages=filtered_messages, request_data=data
source="INPUT",
messages=filtered_messages,
request_data=data,
logging_event_type=GuardrailEventHooks.pre_call,
)
except GuardrailInterventionNormalStringError as e:
bedrock_guardrail_response = e.message
@ -1023,7 +1016,10 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
)
try:
bedrock_guardrail_response = await self.make_bedrock_api_request(
source="INPUT", messages=filtered_messages, request_data=data
source="INPUT",
messages=filtered_messages,
request_data=data,
logging_event_type=GuardrailEventHooks.during_call,
)
except GuardrailInterventionNormalStringError as e:
bedrock_guardrail_response = e.message
@ -1127,9 +1123,13 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
source="INPUT",
messages=input_messages,
request_data=data,
logging_event_type=GuardrailEventHooks.post_call,
)
output_task = self.make_bedrock_api_request(
source="OUTPUT", response=response, request_data=data
source="OUTPUT",
response=response,
request_data=data,
logging_event_type=GuardrailEventHooks.post_call,
)
# Execute both requests in parallel
@ -1143,7 +1143,10 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
# Only run OUTPUT validation (INPUT was already validated in pre_call or during_call)
try:
output_content_bedrock = await self.make_bedrock_api_request(
source="OUTPUT", response=response, request_data=data
source="OUTPUT",
response=response,
request_data=data,
logging_event_type=GuardrailEventHooks.post_call,
)
except GuardrailInterventionNormalStringError as e:
output_content_bedrock = e.message
@ -1270,9 +1273,13 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
source="INPUT",
messages=input_messages,
request_data=request_data,
logging_event_type=GuardrailEventHooks.post_call,
) # Only input messages
output_task = self.make_bedrock_api_request(
source="OUTPUT", response=assembled_model_response
source="OUTPUT",
response=assembled_model_response,
request_data=request_data,
logging_event_type=GuardrailEventHooks.post_call,
) # Only response
# Execute both requests in parallel
@ -1286,7 +1293,10 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
# Only run OUTPUT validation (INPUT was already validated in pre_call or during_call)
try:
output_guardrail_response = await self.make_bedrock_api_request(
source="OUTPUT", response=assembled_model_response
source="OUTPUT",
response=assembled_model_response,
request_data=request_data,
logging_event_type=GuardrailEventHooks.post_call,
)
except GuardrailInterventionNormalStringError as e:
output_guardrail_response = e.message
@ -1563,11 +1573,50 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
# Bedrock will throw an error if there is no text to process
if filtered_messages:
bedrock_response = await self.make_bedrock_api_request(
source="INPUT",
messages=filtered_messages,
request_data=request_data,
_log_hook = (
GuardrailEventHooks.pre_call
if input_type == "request"
else GuardrailEventHooks.post_call
)
# Map the abstract input_type to the Bedrock source parameter.
# "request" -> INPUT (scan user-supplied content)
# "response" -> OUTPUT (scan model-generated content)
# Bedrock guardrail policies are often configured differently
# for Input vs Output (e.g. PII blocking only on Output), so
# the source MUST match where the text originated.
bedrock_source: Literal["INPUT", "OUTPUT"] = (
"OUTPUT" if input_type == "response" else "INPUT"
)
if bedrock_source == "OUTPUT":
# Build a synthetic ModelResponse whose choices carry the
# text(s) to scan, so _create_bedrock_output_content_request
# can produce the correct Bedrock OUTPUT payload.
synthetic_response = ModelResponse(
choices=[
Choices(
index=_idx,
message=Message(
role="assistant",
content=str(_msg.get("content") or ""),
),
finish_reason="stop",
)
for _idx, _msg in enumerate(filtered_messages)
]
)
bedrock_response = await self.make_bedrock_api_request(
source="OUTPUT",
response=synthetic_response,
request_data=request_data,
logging_event_type=_log_hook,
)
else:
bedrock_response = await self.make_bedrock_api_request(
source="INPUT",
messages=filtered_messages,
request_data=request_data,
logging_event_type=_log_hook,
)
# Apply any masking that was applied by the guardrail
output_list = bedrock_response.get("output")

View file

@ -21,20 +21,30 @@ class _PROXY_MaxBudgetLimiter(CustomLogger):
):
try:
verbose_proxy_logger.debug("Inside Max Budget Limiter Pre-Call Hook")
cache_key = f"{user_api_key_dict.user_id}_user_api_key_user_id"
user_row = await cache.async_get_cache(
cache_key, parent_otel_span=user_api_key_dict.parent_otel_span
max_budget = user_api_key_dict.user_max_budget
user_id = user_api_key_dict.user_id
if max_budget is None or user_id is None:
return
# Personal budget applies only to non-team requests, matching
# the explicit team-key exemption in common_checks section 4.1.
if user_api_key_dict.team_id is not None:
return
from litellm.proxy.proxy_server import get_current_spend
curr_spend = await get_current_spend(
counter_key=f"spend:user:{user_id}",
fallback_spend=user_api_key_dict.user_spend or 0.0,
)
if user_row is None: # value not yet cached
return
max_budget = user_row["max_budget"]
curr_spend = user_row["spend"]
if max_budget is None:
return
if curr_spend is None:
return
verbose_proxy_logger.debug(
"MaxBudgetLimiter: user_id=%s, spend=%.6f, max=%.6f",
user_id,
curr_spend,
max_budget,
)
# CHECK IF REQUEST ALLOWED
if curr_spend >= max_budget:

View file

@ -213,6 +213,7 @@ class _ProxyDBLogger(CustomLogger):
team_id=team_id,
user_id=user_id,
response_cost=response_cost,
org_id=org_id,
)
# update cache (fire-and-forget for backward compat:

View file

@ -355,6 +355,7 @@ async def _upsert_budget_and_membership(
tpm_limit: Optional[int] = None,
rpm_limit: Optional[int] = None,
allowed_models: Optional[List[str]] = None,
team_default_budget_id: Optional[str] = None,
):
"""
Helper function to Create/Update or Delete the budget within the team membership
@ -368,6 +369,11 @@ async def _upsert_budget_and_membership(
tpm_limit: Tokens per minute limit for the team member
rpm_limit: Requests per minute limit for the team member
allowed_models: Per-member model scope. None = don't change. [] = remove restrictions. Non-empty list = enforce.
team_default_budget_id: The team's shared default member budget id (from
team metadata.team_member_budget_id), if any. When the membership's
existing_budget_id matches this, we clone-on-write so editing one
member's budget does not mutate the shared default (and therefore
every other member who still points at it).
If max_budget, tpm_limit, rpm_limit, and allowed_models are all None, the user's budget is removed from the team membership.
If any of these values exist, a budget is updated or created and linked to the team membership.
@ -385,7 +391,13 @@ async def _upsert_budget_and_membership(
)
return
if existing_budget_id is not None:
is_shared_default = (
existing_budget_id is not None
and team_default_budget_id is not None
and existing_budget_id == team_default_budget_id
)
if existing_budget_id is not None and not is_shared_default:
# Update the existing budget in-place to preserve fields not being changed.
# Only write fields that the caller explicitly provided (non-None).
update_data: Dict[str, Any] = {
@ -405,11 +417,40 @@ async def _upsert_budget_and_membership(
)
return
# No existing budget — create a new one and link it to the membership.
# Either there is no existing budget, OR the membership is still pointing
# at the team's shared default member budget. In both cases we create a
# NEW private budget for this user and (re)link the membership to it.
create_data: Dict[str, Any] = {
"created_by": user_api_key_dict.user_id or "",
"updated_by": user_api_key_dict.user_id or "",
}
# If we're forking off the shared default, seed the new row with the
# default's values so fields the caller did not change carry over.
if is_shared_default:
default_budget_row = await tx.litellm_budgettable.find_unique(
where={"budget_id": existing_budget_id}
)
if default_budget_row is not None:
default_budget_dict = default_budget_row.model_dump()
for field in (
"max_budget",
"soft_budget",
"max_parallel_requests",
"tpm_limit",
"rpm_limit",
"model_max_budget",
"budget_duration",
"allowed_models",
):
value = default_budget_dict.get(field)
if value is None:
continue
if isinstance(value, list) and len(value) == 0:
continue
create_data[field] = value
# Caller-provided values take precedence over the cloned defaults.
if max_budget is not None:
create_data["max_budget"] = max_budget
if tpm_limit is not None:

View file

@ -1336,7 +1336,9 @@ if MCP_AVAILABLE:
return _redact_mcp_credentials(temp_record)
def _get_cached_temporary_mcp_server_or_404(server_id: str) -> MCPServer:
def _get_cached_temporary_mcp_server_or_404(
server_id: str, request: Optional[Request] = None
) -> MCPServer:
server = get_cached_temporary_mcp_server(server_id)
if server is None:
# Fall back to real DB/config server (e.g. for the user-side OAuth flow
@ -1344,10 +1346,14 @@ if MCP_AVAILABLE:
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
from litellm.proxy.auth.ip_address_utils import IPAddressUtils
client_ip = IPAddressUtils.get_mcp_client_ip(request) if request else None
server = global_mcp_server_manager.get_mcp_server_by_id(
server_id
) or global_mcp_server_manager.get_mcp_server_by_name(server_id)
) or global_mcp_server_manager.get_mcp_server_by_name(
server_id, client_ip=client_ip
)
if server is None:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
@ -1358,10 +1364,12 @@ if MCP_AVAILABLE:
@router.get(
"/server/oauth/{server_id}/authorize",
include_in_schema=False,
dependencies=[Depends(user_api_key_auth)],
)
async def mcp_authorize(
request: Request,
server_id: str,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
client_id: Optional[str] = None,
redirect_uri: str = Query(...),
state: str = "",
@ -1370,7 +1378,7 @@ if MCP_AVAILABLE:
response_type: Optional[str] = None,
scope: Optional[str] = None,
):
mcp_server = _get_cached_temporary_mcp_server_or_404(server_id)
mcp_server = _get_cached_temporary_mcp_server_or_404(server_id, request=request)
# Use the server's stored client_id when the caller doesn't supply one
resolved_client_id = mcp_server.client_id or client_id or ""
if not resolved_client_id:
@ -1399,10 +1407,12 @@ if MCP_AVAILABLE:
@router.post(
"/server/oauth/{server_id}/token",
include_in_schema=False,
dependencies=[Depends(user_api_key_auth)],
)
async def mcp_token(
request: Request,
server_id: str,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
grant_type: str = Form(...),
code: Optional[str] = Form(None),
redirect_uri: Optional[str] = Form(None),
@ -1412,7 +1422,7 @@ if MCP_AVAILABLE:
refresh_token: Optional[str] = Form(None),
scope: Optional[str] = Form(None),
):
mcp_server = _get_cached_temporary_mcp_server_or_404(server_id)
mcp_server = _get_cached_temporary_mcp_server_or_404(server_id, request=request)
resolved_client_id = mcp_server.client_id or client_id or ""
if not resolved_client_id:
raise HTTPException(
@ -1441,9 +1451,14 @@ if MCP_AVAILABLE:
@router.post(
"/server/oauth/{server_id}/register",
include_in_schema=False,
dependencies=[Depends(user_api_key_auth)],
)
async def mcp_register(request: Request, server_id: str):
mcp_server = _get_cached_temporary_mcp_server_or_404(server_id)
async def mcp_register(
request: Request,
server_id: str,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
mcp_server = _get_cached_temporary_mcp_server_or_404(server_id, request=request)
request_data = await _read_request_body(request=request)
data: dict = {**request_data}

View file

@ -2608,6 +2608,15 @@ async def team_member_update(
identified_budget_id = tm.budget_id
break
# If this membership still points at the team's shared default member
# budget, _upsert_budget_and_membership will clone-on-write so that the
# update only touches this user (not every member sharing the default).
team_default_budget_id: Optional[str] = None
if team_table.metadata is not None:
raw_default_budget_id = team_table.metadata.get("team_member_budget_id")
if isinstance(raw_default_budget_id, str):
team_default_budget_id = raw_default_budget_id
### upsert new budget
async with prisma_client.db.tx() as tx:
await _upsert_budget_and_membership(
@ -2620,6 +2629,7 @@ async def team_member_update(
tpm_limit=data.tpm_limit,
rpm_limit=data.rpm_limit,
allowed_models=data.allowed_models,
team_default_budget_id=team_default_budget_id,
)
### update team member role

View file

@ -140,6 +140,62 @@ async def handle_budget_for_entity(
return existing_budget_id
# Fields on LiteLLM_BudgetTable that represent the budget's *configuration*
# (i.e. the values an admin sets). We copy these when cloning a team's
# default member-budget into an individual member-budget so that the new
# row starts with the same limits as the default.
_CLONABLE_BUDGET_FIELDS: Tuple[str, ...] = (
"max_budget",
"soft_budget",
"max_parallel_requests",
"tpm_limit",
"rpm_limit",
"model_max_budget",
"budget_duration",
"allowed_models",
)
async def _clone_team_default_budget_for_member(
prisma_client: PrismaClient,
default_team_budget_id: str,
user_api_key_dict: UserAPIKeyAuth,
litellm_proxy_admin_name: str,
) -> Optional[str]:
"""
Create a new budget row that copies the values from the team's default
member budget. Returns the new budget_id, or None if the default budget
no longer exists in the DB.
Used when adding a new team member without an explicit per-member budget,
so the member starts with the team default's values but gets their own
private budget row (which can be edited independently).
"""
default_budget = await prisma_client.db.litellm_budgettable.find_unique(
where={"budget_id": default_team_budget_id}
)
if default_budget is None:
return None
default_budget_dict = default_budget.model_dump()
cloned_data: dict = {
"created_by": user_api_key_dict.user_id or litellm_proxy_admin_name,
"updated_by": user_api_key_dict.user_id or litellm_proxy_admin_name,
}
for field in _CLONABLE_BUDGET_FIELDS:
value = default_budget_dict.get(field)
if value is None:
continue
# Skip empty list defaults (e.g. allowed_models = []) so the cloned
# row matches the "no value set" shape rather than carrying a default.
if isinstance(value, list) and len(value) == 0:
continue
cloned_data[field] = value
new_budget = await prisma_client.db.litellm_budgettable.create(data=cloned_data)
return new_budget.budget_id
async def add_new_member(
new_member: Member,
max_budget_in_team: Optional[float],
@ -221,8 +277,20 @@ async def add_new_member(
response = await prisma_client.db.litellm_budgettable.create(data=budget_data)
_budget_id = response.budget_id
elif default_team_budget_id is not None:
# No per-member budget was provided, but the team has a default member
# budget. Clone the default budget into a new row for this user so that
# later edits to one member's budget do not bleed into other members.
# If the default no longer exists in the DB, fall back to no budget.
_budget_id = await _clone_team_default_budget_for_member(
prisma_client=prisma_client,
default_team_budget_id=default_team_budget_id,
user_api_key_dict=user_api_key_dict,
litellm_proxy_admin_name=litellm_proxy_admin_name,
)
else:
_budget_id = default_team_budget_id
# No per-member budget and no team default → member gets no budget.
_budget_id = None
if _budget_id and returned_user is not None and returned_user.user_id is not None:
_returned_team_membership = (

View file

@ -577,6 +577,16 @@ class ProxyInitializationHelpers:
help="Exit with error if database migration fails on startup.",
envvar="ENFORCE_PRISMA_MIGRATION_CHECK",
)
@click.option(
"--use_v2_migration_resolver",
is_flag=True,
default=False,
help=(
"Opt into the v2 migration resolver. Avoids the diff-and-force recovery "
"path that can cause schema thrashing during rolling deploys where two "
"LiteLLM versions contend for the same DB. Default is the v1 resolver."
),
)
@click.option(
"--reload",
is_flag=True,
@ -624,6 +634,7 @@ def run_server( # noqa: PLR0915
keepalive_timeout,
max_requests_before_restart,
enforce_prisma_migration_check: bool,
use_v2_migration_resolver: bool,
reload: bool,
):
if setup:
@ -893,9 +904,31 @@ def run_server( # noqa: PLR0915
):
check_prisma_schema_diff(db_url=None)
else:
if not PrismaManager.setup_database(
use_migrate=not use_prisma_db_push
):
if not use_v2_migration_resolver:
print( # noqa
"\033[1;33mLiteLLM Proxy: Using default (v1) migration resolver. "
"If your deployment has seen schema thrashing during rolling "
"deploys, try --use_v2_migration_resolver (safer: avoids the "
"diff-and-force recovery that caused the thrash).\033[0m"
)
try:
setup_ok = PrismaManager.setup_database(
use_migrate=not use_prisma_db_push,
use_v2_resolver=use_v2_migration_resolver,
)
except RuntimeError as e:
# v2 resolver raises on unrecoverable migration errors
# (e.g. non-idempotent failures, permission issues).
# v1 never raises here, so this only fires when the
# operator opted into v2.
print( # noqa
"\033[1;31mLiteLLM Proxy: Database migration cannot proceed. "
f"{e}\033[0m",
file=sys.stderr,
flush=True,
)
sys.exit(2)
if not setup_ok:
if enforce_prisma_migration_check:
print( # noqa
"\033[1;31mLiteLLM Proxy: Database setup failed after multiple retries. "

View file

@ -952,6 +952,17 @@ async def proxy_startup_event(app: FastAPI): # noqa: PLR0915
_run_background_health_check()
) # start the background health check coroutine.
# Start adaptive-router queue flusher unconditionally — adaptive routers
# may be added later via `/config/reload`, and the flusher is a no-op when
# `llm_router.adaptive_routers` is empty. Per-router DB state is loaded
# lazily by the flusher on first tick (see `_state_loaded` flag) so
# hot-reloaded routers also get their persisted priors.
if llm_router is not None and getattr(llm_router, "adaptive_routers", None):
for _ar in llm_router.adaptive_routers.values():
await _ar.load_state_from_db(prisma_client)
_ar._state_loaded = True
asyncio.create_task(_adaptive_router_flusher_loop())
## [Optional] Initialize dd tracer
ProxyStartupEvent._init_dd_tracer()
@ -1795,6 +1806,7 @@ async def increment_spend_counters(
team_id: Optional[str],
user_id: Optional[str],
response_cost: Optional[float],
org_id: Optional[str] = None,
):
"""
Atomically increment spend counters for budget enforcement.
@ -1881,6 +1893,20 @@ async def increment_spend_counters(
increment=response_cost,
)
if user_id is not None:
await _init_and_increment_spend_counter(
counter_key=f"spend:user:{user_id}",
source_cache_key=user_id,
increment=response_cost,
)
if org_id is not None:
await _init_and_increment_spend_counter(
counter_key=f"spend:org:{org_id}",
source_cache_key=f"org_id:{org_id}",
increment=response_cost,
)
async def _init_and_increment_spend_counter(
counter_key: str,
@ -2427,6 +2453,38 @@ def _write_health_state_to_router_cache(
)
_ADAPTIVE_ROUTER_FLUSH_INTERVAL_SECONDS = 10
async def _adaptive_router_flusher_loop():
"""
Drain every AdaptiveRouter's in-memory state + session aggregators into
Postgres on a fixed cadence. Hot-path writes go to memory; this loop is
the only writer to the adaptive router DB tables.
"""
global llm_router, prisma_client
while True:
try:
await asyncio.sleep(_ADAPTIVE_ROUTER_FLUSH_INTERVAL_SECONDS)
adaptive_routers = getattr(llm_router, "adaptive_routers", None) or {}
if not adaptive_routers or prisma_client is None:
continue
for ar in adaptive_routers.values():
# Lazy state load: covers adaptive routers registered via
# `/config/reload` after proxy boot.
if not getattr(ar, "_state_loaded", False):
try:
await ar.load_state_from_db(prisma_client)
finally:
ar._state_loaded = True
await ar.queue.flush_state_to_db(prisma_client)
await ar.queue.flush_session_to_db(prisma_client)
except asyncio.CancelledError:
raise
except Exception:
verbose_proxy_logger.exception("adaptive_router flusher iteration failed")
async def _run_background_health_check():
"""
Periodically run health checks in the background on the endpoints.
@ -13953,6 +14011,38 @@ async def home(request: Request):
return "LiteLLM: RUNNING"
@router.get(
"/adaptive_router/state",
tags=["adaptive_router"],
dependencies=[Depends(user_api_key_auth)],
)
async def get_adaptive_router_state(
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
"""Return live bandit posteriors + queue depth for every configured adaptive router.
Admin-only. Returns 404 if no adaptive router is configured.
Response shape: `{"routers": [<snapshot>, ...]}` one snapshot per
adaptive-router deployment. Each snapshot's `router_name` field identifies
which deployment it came from.
"""
if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN:
raise HTTPException(
status_code=403,
detail={"error": CommonProxyErrors.not_allowed_access.value},
)
if llm_router is None or not llm_router.adaptive_routers:
raise HTTPException(
status_code=404,
detail={"error": "No adaptive_router is configured on this proxy."},
)
snapshots = [
await ar.get_state_snapshot() for ar in llm_router.adaptive_routers.values()
]
return {"routers": snapshots}
@router.get("/routes", dependencies=[Depends(user_api_key_auth)])
async def get_routes():
"""

View file

@ -616,6 +616,7 @@ model LiteLLM_TeamMembership {
user_id String
team_id String
spend Float @default(0.0)
total_spend Float @default(0.0)
budget_id String?
litellm_budget_table LiteLLM_BudgetTable? @relation(fields: [budget_id], references: [budget_id])
@@id([user_id, team_id])
@ -1223,3 +1224,46 @@ model LiteLLM_ClaudeCodePluginTable {
@@map("LiteLLM_ClaudeCodePluginTable")
}
// Per-(router, request_type, model) Beta posterior for the adaptive router.
model LiteLLM_AdaptiveRouterState {
router_name String
request_type String
model_name String
alpha Float
beta Float
total_samples Int @default(0)
last_updated_at DateTime @default(now()) @updatedAt
@@id([router_name, request_type, model_name])
}
// Per-(session, router, model) signal counters for the adaptive router.
model LiteLLM_AdaptiveRouterSession {
session_id String
router_name String
model_name String
classified_type String
misalignment_count Int @default(0)
stagnation_count Int @default(0)
disengagement_count Int @default(0)
satisfaction_count Int @default(0)
failure_count Int @default(0)
loop_count Int @default(0)
exhaustion_count Int @default(0)
last_user_content String?
last_assistant_content String?
tool_call_history Json @default("[]")
pending_tool_calls Json @default("{}")
turn_count Int @default(0)
last_processed_turn Int @default(-1)
clean_credit_awarded Boolean @default(false)
terminal_status Int?
last_activity_at DateTime @default(now()) @updatedAt
@@id([session_id, router_name, model_name])
@@index([last_activity_at], map: "idx_adaptive_router_session_activity")
}

View file

@ -940,7 +940,11 @@ class ProxyLogging:
Result from the guardrail execution
"""
# Use unified_guardrail if callback has apply_guardrail method
use_unified = "apply_guardrail" in type(callback).__dict__
has_apply_guardrail = "apply_guardrail" in type(callback).__dict__
use_unified = has_apply_guardrail and not (
hook_type == "during_call"
and getattr(callback, "use_native_during_call_hook", False)
)
if use_unified:
data["guardrail_to_apply"] = callback
@ -1540,6 +1544,7 @@ class ProxyLogging:
if (
"apply_guardrail" in type(callback).__dict__
and user_api_key_dict is not None
and not getattr(callback, "use_native_during_call_hook", False)
):
data["guardrail_to_apply"] = callback
guardrail_task = self._run_guardrail_task_with_enrichment(

View file

@ -200,6 +200,9 @@ if TYPE_CHECKING:
from litellm.router_strategy.complexity_router.complexity_router import (
ComplexityRouter,
)
from litellm.router_strategy.adaptive_router.adaptive_router import (
AdaptiveRouter,
)
from litellm.router_strategy.quality_router.quality_router import (
QualityRouter,
)
@ -209,6 +212,7 @@ else:
Span = Any
AutoRouter = Any
ComplexityRouter = Any
AdaptiveRouter = Any
QualityRouter = Any
PreRoutingHookResponse = Any
@ -468,6 +472,7 @@ class Router:
) # {"TEAM_ID": PatternMatchRouter}
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
@ -5369,8 +5374,13 @@ class Router:
_request_team_id: Optional[str] = (kwargs.get("metadata", {}) or {}).get(
"user_api_key_team_id"
)
all_deployments = self._get_all_deployments(
model_name=original_model_group, team_id=_request_team_id
# Use wildcard-aware lookup so order-based fallback also works for model
# groups resolved via pattern routing (e.g. `openai/*` -> `openai/gpt-4.1-mini`).
all_deployments = (
self.get_model_list(
model_name=original_model_group, team_id=_request_team_id
)
or []
)
_order_set: set = {
litellm.utils._get_deployment_order(d)
@ -6815,10 +6825,13 @@ class Router:
Check if the deployment is an auto-router deployment (semantic router).
Returns True if the litellm_params model starts with "auto_router/"
but NOT "auto_router/complexity_router" (which uses complexity routing).
but NOT "auto_router/complexity_router" or "auto_router/adaptive_router"
(which use the complexity-router and adaptive-router strategies).
"""
if litellm_params.model.startswith("auto_router/complexity_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/"):
@ -6927,6 +6940,144 @@ class Router:
)
self.complexity_routers[deployment.model_name] = complexity_router
def _is_adaptive_router_deployment(self, litellm_params: LiteLLM_Params) -> bool:
"""True when this deployment opts in via the `auto_router/adaptive_router` model prefix."""
return litellm_params.model.startswith("auto_router/adaptive_router")
def _finalize_adaptive_router_if_configured(self) -> None:
"""Locate every adaptive-router deployment in the finalized model_list and
build an AdaptiveRouter for each. Safe no-op when none are configured.
Idempotent: skips any deployment whose model_name is already initialized."""
# Drop any adaptive-router hooks left over from a previous Router
# instance (e.g. after `/config/reload` replaced `llm_router`). Without
# this, stale AdaptiveRouterPostCallHook callbacks from the old Router
# remain wired up in `litellm.callbacks` and double-fire signal
# recording for every request.
from litellm.router_strategy.adaptive_router.hooks import (
AdaptiveRouterPostCallHook,
)
for _cb_list in (
litellm.callbacks,
litellm.success_callback,
litellm.failure_callback,
litellm._async_success_callback,
litellm._async_failure_callback,
):
litellm.logging_callback_manager.remove_callbacks_by_type(
_cb_list, AdaptiveRouterPostCallHook
)
for entry in self.model_list or []:
lp = (
entry.get("litellm_params")
if isinstance(entry, dict)
else entry.litellm_params
)
lp_model = (
(lp.get("model") if isinstance(lp, dict) else lp.model) if lp else None
)
if not (lp_model and lp_model.startswith("auto_router/adaptive_router")):
continue
model_name = (
entry.get("model_name") if isinstance(entry, dict) else entry.model_name
)
if not model_name or not lp:
continue
if model_name in self.adaptive_routers:
continue
deployment = Deployment(
model_name=model_name,
litellm_params=(
lp if not isinstance(lp, dict) else LiteLLM_Params(**lp)
),
model_info=(
entry.get("model_info")
if isinstance(entry, dict)
else entry.model_info
),
)
self.init_adaptive_router_deployment(deployment=deployment)
def init_adaptive_router_deployment(self, deployment: Deployment) -> None:
"""
Build an AdaptiveRouter instance for this deployment and register its
post-call hook. Multiple adaptive routers can coexist on a single Router,
keyed by `deployment.model_name`.
`model_to_prefs` and `model_to_cost` are derived from the OTHER models
already registered in `self.model_list` whose `model_name` appears in
`available_models`. Models not yet registered fall back to defaults.
"""
# Local import: AdaptiveRouter -> hooks -> classifier all import litellm
# internals which transitively import this module. (AGENTS.md exception clause.)
from litellm.router_strategy.adaptive_router.adaptive_router import (
AdaptiveRouter,
)
from litellm.router_strategy.adaptive_router.hooks import (
AdaptiveRouterPostCallHook,
)
from litellm.types.router import (
AdaptiveRouterConfig,
AdaptiveRouterPreferences,
)
raw_config = deployment.litellm_params.adaptive_router_config
if raw_config is None:
raise ValueError(
"adaptive_router_config is required for adaptive-router deployments."
)
config = AdaptiveRouterConfig(**raw_config)
model_to_prefs: Dict[str, AdaptiveRouterPreferences] = {}
model_to_cost: Dict[str, float] = {}
# O(k) via the name→indices map: only touch deployments whose name
# is listed in `available_models`, instead of scanning model_list.
for name in config.available_models:
indices = self.model_name_to_deployment_indices.get(name, [])
if not indices:
continue
d = (self.model_list or [])[indices[0]]
mi = d.get("model_info") if isinstance(d, dict) else d.model_info
mi_dict: Dict[str, Any] = (
mi if isinstance(mi, dict) else (mi.model_dump() if mi else {})
)
prefs_raw = mi_dict.get("adaptive_router_preferences")
if prefs_raw is not None:
model_to_prefs[name] = AdaptiveRouterPreferences(**prefs_raw)
# `input_cost_per_token` is a LiteLLM_Params field per types/router.py.
lp = d.get("litellm_params") if isinstance(d, dict) else d.litellm_params
lp_dict: Dict[str, Any] = (
lp if isinstance(lp, dict) else (lp.model_dump() if lp else {})
)
cost = lp_dict.get("input_cost_per_token")
if cost is not None:
model_to_cost[name] = float(cost)
if deployment.model_name in self.adaptive_routers:
raise ValueError(
f"Adaptive-router deployment {deployment.model_name} already exists. "
"Please use a different model name."
)
adaptive_router = AdaptiveRouter(
router_name=deployment.model_name,
config=config,
model_to_prefs=model_to_prefs,
model_to_cost=model_to_cost,
)
self.adaptive_routers[deployment.model_name] = adaptive_router
litellm.logging_callback_manager.add_litellm_callback(
AdaptiveRouterPostCallHook(adaptive_router=adaptive_router)
)
verbose_router_logger.info(
"AdaptiveRouter[%s] initialized with %d models",
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.
@ -7077,6 +7228,10 @@ class Router:
# Note: model_name_to_deployment_indices is already built incrementally
# by _create_deployment -> _add_model_to_list_and_index_map
# Deferred: build the AdaptiveRouter strategy now that all underlying
# deployments have been registered.
self._finalize_adaptive_router_if_configured()
def _add_deployment(self, deployment: Deployment) -> Deployment:
import os
@ -7204,6 +7359,10 @@ class Router:
):
self.init_complexity_router_deployment(deployment=deployment)
# NOTE: adaptive-router deployments are deferred to the end of
# 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
#########################################################
@ -9841,6 +10000,19 @@ class Router:
specific_deployment=specific_deployment,
)
#########################################################
# Check if an adaptive-router should be used
#########################################################
adaptive_router = self.adaptive_routers.get(model)
if adaptive_router is not None:
return await adaptive_router.async_pre_routing_hook(
model=model,
request_kwargs=request_kwargs,
messages=messages,
input=input,
specific_deployment=specific_deployment,
)
#########################################################
# Check if any quality-router should be used
#########################################################

View file

@ -0,0 +1,95 @@
# Adaptive Router (v0)
A request-type-aware routing strategy. For each incoming request, classify the
prompt into one of seven `RequestType` buckets (code generation, writing,
analytical reasoning, …), then Thompson-sample a Beta(α, β) bandit posterior
per `(request_type, model)` cell to pick the best model. Quality estimates are
combined with a normalized cost score via a weighted linear sum.
A post-call hook reads the response and runs lightweight regex + tool-call
detectors (see `signals.py`) to award per-turn credit/blame to the model that
served the turn. Updates are batched in-memory and flushed to Postgres every
~10s by a background task in `proxy_server.py`.
## Config example
```yaml
model_list:
- model_name: gpt-4o
litellm_params:
model: openai/gpt-4o
model_info:
input_cost_per_token: 0.0000025
adaptive_router_preferences:
quality_tier: 3
strengths: ["code_generation", "analytical_reasoning"]
- model_name: gpt-4o-mini
litellm_params:
model: openai/gpt-4o-mini
model_info:
input_cost_per_token: 0.00000015
adaptive_router_preferences:
quality_tier: 2
strengths: ["general", "factual_lookup"]
- model_name: smart-router
litellm_params:
model: auto_router/adaptive_router
adaptive_router_default_model: gpt-4o-mini
adaptive_router_config:
available_models: ["gpt-4o", "gpt-4o-mini"]
weights:
quality: 0.7
cost: 0.3
```
Callers may pass header `x-litellm-min-quality-tier: 3` (or metadata key
`min_quality_tier: 3`) to force selection from tier-3-or-higher models only.
## Behavior summary
- **Cold start.** Each `(request_type, model)` cell starts with a
Beta prior whose mean = `BASE_TIER_WEIGHT[tier] (+ STRENGTH_BONUS if declared)`
and total mass = `COLD_START_MASS` (10). About ten real observations move it
meaningfully.
- **Per-request decision.** Sample once per eligible model, score with
`quality_weight·sample + cost_weight·normalized_cost`, pick the argmax.
Routing is stateless per-turn — no sticky lookup. Each call resamples.
- **Owner-cache attribution.** Post-call, the conversation's first picked
model claims an "owner slot" for `OWNER_CACHE_TTL_SECONDS` (24h). Later
turns of the same conversation only fire bandit/state updates if the
same model handled them — mismatches are dropped (no attribution) and
counted in `skipped_updates_total`. Conversation identity is the
client-supplied `litellm_session_id` if present, otherwise a sha256 over
caller identity (api key hash, team, user, end-user) + the first message.
- **Per-turn updates.** `satisfaction → +α`. `misalignment, stagnation,
disengagement, failure → +β` (each). `loop → +0.5β`. `exhaustion → 0`
(uptime, not quality). Skipped if conversation has fewer than
`SIGNAL_GATE_MIN_MESSAGES` messages.
- **Persistence.** Bandit cells: aggregated deltas, eventually consistent.
Session rows: last-write-wins snapshots.
## Known v0 limitations
- **Latency is not in the score.** Quality + cost only. A pathologically slow
model can still be picked.
- **Hard sample cap at 200.** Once `α + β > 200`, deltas are silently dropped.
No rescaling — drift is a v1 concern.
- **24h owner-cache TTL.** No explicit eviction below TTL. The in-memory map
can grow if traffic patterns produce many one-shot sessions.
- **Owner-recovery skew.** If model A "owns" a conversation but is then
dethroned in the bandit, later turns served by model B are dropped — so
bandit updates for that conversation flatline until A's TTL expires.
Tracked via `skipped_updates_total`.
- **Signals are regex + tool-call only.** No LLM-judge, no embedding similarity,
no exemplar storage. Signals are best-effort and biased toward English.
- **One AdaptiveRouter per `Router`.** Multiple `adaptive_router/*` deployments
on the same `litellm.Router` raise at init.
- **Bandit-delta mapping is unvalidated.** `_compute_bandit_delta` is a v0
guess; expect to retune after the first ~1000 sessions of real traffic.
- **`request_type` is classified per turn from the latest user message.** For
non-GENERAL turns, the current-turn type is used for bandit attribution (so
genuine mid-session topic shifts update the correct cell). For GENERAL turns
("thanks!", "ok", "sounds good"), attribution falls back to the session's
original type to avoid misattributing closing pleasantries.

View file

@ -0,0 +1,6 @@
"""Adaptive router strategy. See README.md for design overview."""
from litellm.router_strategy.adaptive_router.adaptive_router import AdaptiveRouter
from litellm.router_strategy.adaptive_router.hooks import AdaptiveRouterPostCallHook
__all__ = ["AdaptiveRouter", "AdaptiveRouterPostCallHook"]

View file

@ -0,0 +1,454 @@
"""
Main adaptive router strategy. See README.md for design overview.
One AdaptiveRouter instance per router_name. Holds in-memory caches:
- _cells: Beta(alpha, beta) bandit posteriors per (request_type, model)
- _owner_cache: session_key -> (owner_model, expires_at) the first model
picked for a conversation owns its bandit-update slot
- _session_states: (session_key, model) -> SessionState for incremental signal updates
Owns the AdaptiveRouterUpdateQueue used by the proxy's flusher to persist
state and session snapshots back to Postgres.
Routing is stateless per-turn (Thompson sample fresh on every call). The
owner cache is consulted only at post-call time to decide whether a turn's
signals should fire a bandit update turns served by a different model than
the conversation's owner are skipped to avoid cross-model misattribution.
"""
from __future__ import annotations
import asyncio
import time
from dataclasses import asdict
from typing import Any, Dict, List, Optional, Tuple, Union, cast
from litellm._logging import verbose_router_logger
from litellm.litellm_core_utils.prompt_templates.common_utils import (
get_last_user_message,
)
from litellm.router_strategy.adaptive_router.bandit import (
BanditCell,
apply_delta,
initial_cell,
pick_best,
)
from litellm.router_strategy.adaptive_router.classifier import classify_prompt
from litellm.router_strategy.adaptive_router.config import (
ADAPTIVE_ROUTER_CHOSEN_MODEL_KEY,
MIN_QUALITY_TIER_HEADER,
MIN_QUALITY_TIER_METADATA_KEY,
OWNER_CACHE_TTL_SECONDS,
)
from litellm.router_strategy.adaptive_router.signals import (
SessionState,
SignalDelta,
Turn,
apply_turn,
)
from litellm.router_strategy.adaptive_router.update_queue import (
AdaptiveRouterUpdateQueue,
)
# Sweep session-state cache when it exceeds this many live entries. Expired
# entries are dropped in bulk; amortizes to O(1) per insert.
_SESSION_STATE_SWEEP_THRESHOLD: int = 1024
# Same pattern for the owner cache.
_OWNER_CACHE_SWEEP_THRESHOLD: int = 1024
from litellm.types.llms.openai import AllMessageValues
from litellm.types.router import (
AdaptiveRouterConfig,
AdaptiveRouterPreferences,
PreRoutingHookResponse,
RequestType,
)
def _default_prefs() -> AdaptiveRouterPreferences:
"""Tier-2 prior with no declared strengths; used when a model omits prefs."""
return AdaptiveRouterPreferences(quality_tier=2, strengths=[])
class AdaptiveRouter:
"""One instance per router_name. Holds in-memory caches + the update queue."""
def __init__(
self,
router_name: str,
config: AdaptiveRouterConfig,
model_to_prefs: Dict[str, AdaptiveRouterPreferences],
model_to_cost: Dict[str, float],
) -> None:
self.router_name = router_name
self.config = config
self.model_to_prefs = model_to_prefs
self.model_to_cost = model_to_cost
self.queue = AdaptiveRouterUpdateQueue()
self._cells: Dict[Tuple[RequestType, str], BanditCell] = {}
self._owner_cache: Dict[str, Tuple[str, float]] = {}
self._session_states: Dict[Tuple[str, str], SessionState] = {}
# Parallel expiry map for _session_states, same TTL as _owner_cache.
# Evicted opportunistically in `get_or_create_session_state`.
self._session_states_expiry: Dict[Tuple[str, str], float] = {}
self._skipped_updates_total: int = 0
# Set to True once the proxy flusher has loaded persisted priors from
# Postgres. Checked to support lazy-load on hot-reloaded routers.
self._state_loaded: bool = False
self._lock = asyncio.Lock()
self._init_cold_start_cells()
# ---- Cold-start ------------------------------------------------------
def _init_cold_start_cells(self) -> None:
"""Populate _cells with cold-start priors for every (rt, model) combination."""
for rt in RequestType:
for model in self.config.available_models:
prefs = self.model_to_prefs.get(model) or _default_prefs()
self._cells[(rt, model)] = initial_cell(prefs, rt)
async def load_state_from_db(self, prisma_client: Any) -> None:
"""Override cold-start cells with persisted state. Called once at startup."""
if prisma_client is None:
return
try:
rows = await prisma_client.db.litellm_adaptiverouterstate.find_many(
where={"router_name": self.router_name}
)
loaded = 0
for row in rows:
try:
rt = RequestType(row.request_type)
except ValueError:
# Unknown taxonomy entry from an older/newer version. Skip.
continue
if row.model_name not in self.config.available_models:
continue
self._cells[(rt, row.model_name)] = BanditCell(
alpha=row.alpha, beta=row.beta
)
loaded += 1
verbose_router_logger.info(
"AdaptiveRouter[%s]: loaded %d cells from DB",
self.router_name,
loaded,
)
except Exception as e:
verbose_router_logger.exception(
"AdaptiveRouter[%s]: failed to load state from DB: %s",
self.router_name,
e,
)
# ---- Pre-routing hook ------------------------------------------------
async def async_pre_routing_hook(
self,
model: str,
request_kwargs: Dict[str, Any],
messages: Optional[List[Dict[str, Any]]] = None,
input: Optional[Union[str, List]] = None,
specific_deployment: Optional[bool] = False,
) -> Optional[PreRoutingHookResponse]:
"""
Plugin entry point invoked by `Router.async_pre_routing_hook` when the
inbound `model` matches this adaptive router's `router_name`.
Classifies the last user message, picks a logical model via the bandit,
and stashes the chosen model on `request_kwargs["metadata"]` so the
post-call hook can surface it as a response header.
Routing is stateless per-turn: every call Thompson-samples fresh,
regardless of any prior pick for the same session. Cross-turn
attribution is enforced post-call via the owner cache (see
`claim_or_check_owner`).
"""
user_text = (
get_last_user_message(cast(List[AllMessageValues], messages or [])) or ""
)
request_type = classify_prompt(user_text)
min_quality_tier = self._extract_min_quality_tier(request_kwargs)
chosen_model = await self.pick_model(
request_type=request_type, min_quality_tier=min_quality_tier
)
verbose_router_logger.debug(
"AdaptiveRouter[%s]: classified=%s -> chose %s",
self.router_name,
request_type.value,
chosen_model,
)
# Relay the chosen logical model to the post-call hook, which surfaces
# it as the `x-litellm-adaptive-router-model` response header. We use
# `metadata` (not a top-level kwarg) so the value doesn't leak into
# `litellm.acompletion(**input_kwargs)`.
kwargs_metadata = request_kwargs.setdefault("metadata", {})
if isinstance(kwargs_metadata, dict):
kwargs_metadata[ADAPTIVE_ROUTER_CHOSEN_MODEL_KEY] = chosen_model
return PreRoutingHookResponse(model=chosen_model, messages=messages)
# ---- Pick model ------------------------------------------------------
async def pick_model(
self,
request_type: RequestType,
min_quality_tier: Optional[int] = None,
) -> str:
"""Thompson-sample across eligible models. Stateless per-turn."""
eligible = self._eligible_models(min_quality_tier)
if not eligible:
raise ValueError(
f"AdaptiveRouter[{self.router_name}]: no models meet "
f"min_quality_tier={min_quality_tier}"
)
cells = {m: self._cells[(request_type, m)] for m in eligible}
costs = {m: self.model_to_cost.get(m, 0.0) for m in eligible}
return pick_best(
cells,
costs,
quality_weight=self.config.weights.quality,
cost_weight=self.config.weights.cost,
)
def claim_or_check_owner(self, session_key: str, current_model: str) -> bool:
"""Resolve attribution for a turn under stateless routing.
Returns True iff this turn should fire a bandit/state update. The
first call for a `session_key` claims ownership for `current_model`
and returns True. Subsequent calls return True only if the owner is
still live AND matches `current_model`. Mismatches (a different
model handled this turn) and expired owners both increment
`_skipped_updates_total` and return False no attribution.
"""
now = time.time()
existing = self._owner_cache.get(session_key)
if existing is not None and existing[1] > now:
owner_model, _ = existing
if owner_model == current_model:
return True
self._skipped_updates_total += 1
return False
# Opportunistic bulk sweep — sessions that never come back would
# otherwise pile up here forever. Same threshold pattern as the
# session-state cache.
if len(self._owner_cache) >= _OWNER_CACHE_SWEEP_THRESHOLD:
self._evict_expired_owner_cache(now)
# No live owner -> claim for current_model.
self._owner_cache[session_key] = (
current_model,
now + OWNER_CACHE_TTL_SECONDS,
)
return True
def _evict_expired_owner_cache(self, now: float) -> None:
expired = [k for k, (_, exp) in self._owner_cache.items() if exp <= now]
for k in expired:
self._owner_cache.pop(k, None)
async def get_state_snapshot(self) -> Dict[str, Any]:
"""In-memory snapshot for the introspection endpoint. Cheap; no DB hit."""
cells = []
for (rt, model), cell in sorted(
self._cells.items(), key=lambda kv: (kv[0][0].value, kv[0][1])
):
total = cell.alpha + cell.beta
cells.append(
{
"request_type": rt.value,
"model": model,
"alpha": cell.alpha,
"beta": cell.beta,
# Net observations that have moved the posterior, excluding
# the cold-start prior mass. `alpha + beta` would show the
# initial COLD_START_MASS (e.g. 10) before any real traffic
# arrives, which confuses operators reading the endpoint.
"samples": cell.total_samples,
"quality_mean": cell.alpha / total if total > 0 else 0.0,
}
)
queue = await self.queue.queue_size()
now = time.time()
owner_cache_live = sum(1 for _, exp in self._owner_cache.values() if exp > now)
return {
"router_name": self.router_name,
"available_models": list(self.config.available_models),
"weights": {
"quality": self.config.weights.quality,
"cost": self.config.weights.cost,
},
"model_costs": dict(self.model_to_cost),
"cells": cells,
"owner_cache_live": owner_cache_live,
"skipped_updates_total": self._skipped_updates_total,
"queue": queue,
}
@staticmethod
def _extract_min_quality_tier(
request_kwargs: Dict[str, Any],
) -> Optional[int]:
"""Pull `min_quality_tier` from request headers or metadata.
Precedence: headers (`x-litellm-min-quality-tier`) over metadata
(`min_quality_tier`). Headers arrive lowercased from the proxy but we
lookup case-insensitively to be safe. Unparseable values are ignored
(treated as "not set") rather than raising a bad header shouldn't
fail the request.
"""
headers = request_kwargs.get("headers") or {}
if isinstance(headers, dict):
for k, v in headers.items():
if isinstance(k, str) and k.lower() == MIN_QUALITY_TIER_HEADER:
try:
return int(v)
except (TypeError, ValueError):
return None
metadata = request_kwargs.get("metadata") or {}
if isinstance(metadata, dict):
raw = metadata.get(MIN_QUALITY_TIER_METADATA_KEY)
if raw is not None:
try:
return int(raw)
except (TypeError, ValueError):
return None
return None
def _eligible_models(self, min_quality_tier: Optional[int]) -> List[str]:
if min_quality_tier is None:
return list(self.config.available_models)
return [
m
for m in self.config.available_models
if (self.model_to_prefs.get(m) or _default_prefs()).quality_tier
>= min_quality_tier
]
# ---- Session state ---------------------------------------------------
def get_or_create_session_state(
self,
session_id: str,
model_name: str,
request_type: RequestType,
) -> SessionState:
key = (session_id, model_name)
now = time.time()
# Opportunistic bulk sweep when the cache grows past the threshold.
# Cheap relative to the alternative of a bounded LRU — conversations
# naturally become inactive within OWNER_CACHE_TTL_SECONDS.
if len(self._session_states) >= _SESSION_STATE_SWEEP_THRESHOLD:
self._evict_expired_session_states(now)
state = self._session_states.get(key)
if state is None:
state = SessionState(
session_id=session_id,
router_name=self.router_name,
model_name=model_name,
classified_type=request_type.value,
)
self._session_states[key] = state
self._session_states_expiry[key] = now + OWNER_CACHE_TTL_SECONDS
return state
def _evict_expired_session_states(self, now: float) -> None:
"""Drop session states whose TTL has passed. O(n) but amortized O(1)
per insert thanks to `_SESSION_STATE_SWEEP_THRESHOLD`."""
expired = [k for k, exp in self._session_states_expiry.items() if exp <= now]
for k in expired:
self._session_states.pop(k, None)
self._session_states_expiry.pop(k, None)
async def record_turn(
self,
session_id: str,
model_name: str,
request_type: RequestType,
turn: Turn,
) -> SignalDelta:
"""Apply one turn, push session snapshot + bandit deltas to the queue."""
state = self.get_or_create_session_state(session_id, model_name, request_type)
delta = apply_turn(state, turn)
verbose_router_logger.debug(
"AdaptiveRouter[%s]: record_turn delta=%s", self.router_name, delta
)
# Strip the raw conversation content before persisting. The
# last_user/assistant_content and tool_call_history fields are only
# needed in-memory for the next turn's incremental signal detection;
# writing user prompts and tool payloads to the DB would store PII
# for every adaptive-router conversation. Counts + bookkeeping is
# all the persisted row needs.
snapshot = asdict(state)
for sensitive in (
"last_user_content",
"last_assistant_content",
"tool_call_history",
"pending_tool_calls",
):
snapshot.pop(sensitive, None)
await self.queue.add_session_state(
session_id, self.router_name, model_name, snapshot
)
d_alpha, d_beta = self._compute_bandit_delta(delta)
verbose_router_logger.debug(
"AdaptiveRouter[%s]: bandit delta alpha=%.2f beta=%.2f",
self.router_name,
d_alpha,
d_beta,
)
if d_alpha != 0 or d_beta != 0:
# For non-GENERAL turns, attribute to the current-turn classification
# so genuine mid-session topic shifts (e.g. code → math) update the
# correct cell. For GENERAL turns ("thanks!", "ok", "sounds good"), fall
# back to the session's original type so closing pleasantries don't
# misattribute the reward.
attribution_type = (
request_type
if request_type != RequestType.GENERAL
else RequestType(state.classified_type)
)
cell_key = (attribution_type, model_name)
self._cells[cell_key] = apply_delta(self._cells[cell_key], d_alpha, d_beta)
await self.queue.add_state_delta(
self.router_name,
attribution_type.value,
model_name,
d_alpha,
d_beta,
)
return delta
@staticmethod
def _compute_bandit_delta(delta: SignalDelta) -> Tuple[float, float]:
"""
Translate per-turn signal deltas into bandit-cell deltas.
v0 mapping (UNVALIDATED D6):
- satisfaction -> +1 alpha
- misalignment, stagnation,
disengagement, failure -> +1 beta each
- loop -> +0.5 beta (weak; could be model OR user)
- exhaustion -> 0 (uptime issue, tracked separately later)
"""
d_alpha = float(delta.satisfaction)
d_beta = (
float(
delta.misalignment
+ delta.stagnation
+ delta.disengagement
+ delta.failure
)
+ 0.5 * delta.loop
)
return d_alpha, d_beta

View file

@ -0,0 +1,142 @@
"""
Thompson sampling and prior initialization for the adaptive router bandit.
Each (router, request_type, model) cell is a Beta(alpha, beta) posterior.
- alpha = pseudo-successes
- beta = pseudo-failures
- mean = alpha / (alpha + beta)
- total samples = alpha + beta - COLD_START_MASS (informative prior, not data)
Hot path: thompson_sample() pure function, no I/O.
"""
import random
from dataclasses import dataclass
from typing import Dict, List, Optional
from litellm.router_strategy.adaptive_router.config import (
BASE_TIER_WEIGHT,
COLD_START_MASS,
DEFAULT_COST_WEIGHT,
DEFAULT_QUALITY_WEIGHT,
SAMPLE_CAP,
STRENGTH_BONUS,
)
from litellm.types.router import AdaptiveRouterPreferences, RequestType
@dataclass(frozen=True)
class BanditCell:
"""Posterior state for a single (router, request_type, model) cell."""
alpha: float
beta: float
@property
def mean(self) -> float:
total = self.alpha + self.beta
return self.alpha / total if total > 0 else 0.5
@property
def total_samples(self) -> int:
return max(0, int(self.alpha + self.beta - COLD_START_MASS))
def initial_cell(
prefs: AdaptiveRouterPreferences, request_type: RequestType
) -> BanditCell:
"""
Cold-start prior for a (model, request_type) cell.
mean = base_tier_weight[tier] + (STRENGTH_BONUS if request_type in strengths else 0)
capped at 0.95 to avoid an over-confident prior.
Total mass = COLD_START_MASS so that ~10 real observations can move it noticeably.
"""
if prefs.quality_tier not in BASE_TIER_WEIGHT:
valid = sorted(BASE_TIER_WEIGHT)
raise ValueError(
f"quality_tier={prefs.quality_tier} is not supported; "
f"valid tiers are {valid}"
)
base = BASE_TIER_WEIGHT[prefs.quality_tier]
bonus = STRENGTH_BONUS if request_type in prefs.strengths else 0.0
mean = min(0.95, base + bonus)
alpha = mean * COLD_START_MASS
beta = (1.0 - mean) * COLD_START_MASS
return BanditCell(alpha=alpha, beta=beta)
def apply_delta(cell: BanditCell, delta_alpha: float, delta_beta: float) -> BanditCell:
"""
Apply a learning update to a cell, enforcing the sample cap.
SAMPLE_CAP is a HARD cap on (alpha + beta). When the cap would be exceeded,
we drop the update. (D5: hard cap, no rescaling keep v0 simple.)
"""
new_alpha = cell.alpha + delta_alpha
new_beta = cell.beta + delta_beta
if new_alpha + new_beta > SAMPLE_CAP:
return cell
return BanditCell(alpha=new_alpha, beta=new_beta)
def thompson_sample(cell: BanditCell, rng: Optional[random.Random] = None) -> float:
"""Draw a sample from Beta(alpha, beta). Returns a quality estimate in [0, 1]."""
r = rng if rng is not None else random
return r.betavariate(cell.alpha, cell.beta)
def normalized_cost(model_cost: float, all_costs: List[float]) -> float:
"""
Map a raw $/1k-token cost into [0, 1] where 0 = most expensive, 1 = cheapest.
Returns 0.5 when there's no spread.
"""
if not all_costs:
return 0.5
lo, hi = min(all_costs), max(all_costs)
if hi == lo:
return 0.5
return 1.0 - ((model_cost - lo) / (hi - lo))
def score(
quality_sample: float,
model_cost: float,
all_costs: List[float],
quality_weight: float = DEFAULT_QUALITY_WEIGHT,
cost_weight: float = DEFAULT_COST_WEIGHT,
) -> float:
"""
Multi-objective score. V0 is a weighted linear sum of (quality, normalized_cost).
Higher is better. Both inputs are in [0, 1].
"""
cost_score = normalized_cost(model_cost, all_costs)
return quality_weight * quality_sample + cost_weight * cost_score
def pick_best(
cells: Dict[str, BanditCell],
model_costs: Dict[str, float],
quality_weight: float = DEFAULT_QUALITY_WEIGHT,
cost_weight: float = DEFAULT_COST_WEIGHT,
rng: Optional[random.Random] = None,
) -> str:
"""
Sample once per model, score each, return the model with highest score.
cells: {model_name: BanditCell}
model_costs: {model_name: $/1k tokens}
"""
if not cells:
raise ValueError("pick_best called with no models")
all_costs = list(model_costs.values())
best_model: Optional[str] = None
best_score = float("-inf")
for model, cell in cells.items():
q = thompson_sample(cell, rng=rng)
s = score(q, model_costs[model], all_costs, quality_weight, cost_weight)
if s > best_score:
best_score = s
best_model = model
assert best_model is not None
return best_model

View file

@ -0,0 +1,140 @@
"""
Rule-based classifier mapping a user prompt to a RequestType.
V0 design choice: deterministic regex over the FIRST user message in a session.
Result is cached per session (caller's responsibility, not ours).
Order matters: we check more specific types first, falling back to GENERAL.
"""
import re
from typing import List, Pattern, Tuple
from litellm.types.router import RequestType
_RULES: List[Tuple[Pattern[str], RequestType]] = [
(
re.compile(
r"\b(write|create|generate|implement|build)\s+(?:a |an |the |me )?(?:python|javascript|typescript|java|rust|go|c\+\+|sql|bash|shell)\b",
re.IGNORECASE,
),
RequestType.CODE_GENERATION,
),
(
re.compile(
r"\b(write|create|implement|build)\b(?:\s+\w+){0,4}?\s+(function|class|method|script|program|api|endpoint|microservice)\b",
re.IGNORECASE,
),
RequestType.CODE_GENERATION,
),
(
re.compile(
r"\b(explain|describe|understand|walk me through|what does)\b.*\b(code|function|method|class|algorithm|snippet)\b",
re.IGNORECASE,
),
RequestType.CODE_UNDERSTANDING,
),
(
re.compile(
r"\b(debug|fix|why (?:is|does|isn't)|what.s wrong|trace)\b.*\b(error|bug|exception|stacktrace|stack trace|traceback)\b",
re.IGNORECASE,
),
RequestType.CODE_UNDERSTANDING,
),
(
re.compile(
r"\b(review|critique)\s+(?:this |my |the )?(?:code|pr|pull request|diff|patch)\b",
re.IGNORECASE,
),
RequestType.CODE_UNDERSTANDING,
),
(
re.compile(
r"\b(design|architect|plan|architecture)\b.*\b(system|service|api|database|schema|module|microservice)\b",
re.IGNORECASE,
),
RequestType.TECHNICAL_DESIGN,
),
(
re.compile(
r"\b(should i (?:use|choose|pick)|tradeoffs? between|compare)\b.*\b(library|framework|language|database|protocol|postgres|postgresql|mongodb|dynamodb|mysql|redis|kafka|sql|nosql)\b",
re.IGNORECASE,
),
RequestType.TECHNICAL_DESIGN,
),
(
re.compile(
r"\bhow (?:should|do) i (?:design|structure|organize|model)\b",
re.IGNORECASE,
),
RequestType.TECHNICAL_DESIGN,
),
(
re.compile(
r"\b(solve|compute|calculate|prove|derive)\b.*\b(equation|integral|derivative|theorem|proof|problem)\b",
re.IGNORECASE,
),
RequestType.ANALYTICAL_REASONING,
),
(
re.compile(r"\b(if .+ then|given .+ find|suppose|assume)\b", re.IGNORECASE),
RequestType.ANALYTICAL_REASONING,
),
(
re.compile(
r"\b(probability|statistics|combinatorics|optimization problem)\b",
re.IGNORECASE,
),
RequestType.ANALYTICAL_REASONING,
),
(
re.compile(
r"\b(write|draft|compose|rewrite|edit|proofread|polish)\b.*\b(email|essay|blog|post|article|letter|memo|copy|paragraph|sentence)\b",
re.IGNORECASE,
),
RequestType.WRITING,
),
(
re.compile(
r"\b(make (?:this|it)|help me)\s+(?:more |less )?(?:concise|formal|casual|professional|persuasive)\b",
re.IGNORECASE,
),
RequestType.WRITING,
),
(
re.compile(
r"^\s*(who|what|when|where|which)\s+(?:is|was|were|are)\b", re.IGNORECASE
),
RequestType.FACTUAL_LOOKUP,
),
(
re.compile(r"^\s*(define|definition of|meaning of)\b", re.IGNORECASE),
RequestType.FACTUAL_LOOKUP,
),
(
re.compile(
r"^\s*how (?:do you spell|to spell|many .* are there|tall is)\b",
re.IGNORECASE,
),
RequestType.FACTUAL_LOOKUP,
),
]
def classify_prompt(text: str) -> RequestType:
"""
Classify a single user prompt.
Falls back to GENERAL when no rule matches. Empty/whitespace-only also
returns GENERAL.
"""
if not text or not text.strip():
return RequestType.GENERAL
truncated = text[:2000]
for pattern, request_type in _RULES:
if pattern.search(truncated):
return request_type
return RequestType.GENERAL

View file

@ -0,0 +1,54 @@
"""
Configuration constants for the adaptive_router strategy.
All magic numbers are first-pass guesses (D3-D6 in the handoff plan).
Expect to retune after first 1000 sessions of real traffic.
"""
from typing import Dict
from litellm.types.router import RequestType # re-export for convenience # noqa: F401
# D3 — Score weights (default; user-overridable via AdaptiveRouterConfig.weights)
DEFAULT_QUALITY_WEIGHT: float = 0.7 # UNVALIDATED — calibrated against [0] sessions
DEFAULT_COST_WEIGHT: float = 0.3 # UNVALIDATED — calibrated against [0] sessions
# D4 — Cold-start prior: (alpha + beta) total mass = COLD_START_MASS
# Mean of Beta = base_tier_weight + (strength_bonus if declared)
BASE_TIER_WEIGHT: Dict[int, float] = {1: 0.3, 2: 0.5, 3: 0.7} # UNVALIDATED
STRENGTH_BONUS: float = 0.3 # UNVALIDATED
COLD_START_MASS: float = 10.0
# D5 — Sample cap. Hard cap, no rescaling (drift handling is v1).
SAMPLE_CAP: int = 200
# D6 — Clean-trace credit: minimum turns before α += 1 can fire.
MIN_TURNS_FOR_CLEAN_CREDIT: int = 3
# D2 — Owner-cache TTL (seconds). 24h.
# A conversation's first-picked model "owns" the bandit-update slot for
# this long. Subsequent turns of the same conversation only contribute a
# bandit/state update when the same model is re-sampled.
OWNER_CACHE_TTL_SECONDS: int = 24 * 3600
# Below this many messages we skip post-call signal recording. Most signals
# (misalignment, stagnation, satisfaction-in-response-to-prior-turn) need at
# least one full prior exchange to be meaningful.
SIGNAL_GATE_MIN_MESSAGES: int = 4
# Detector thresholds (from Plano/Chen 2026 paper).
MISALIGNMENT_JACCARD_THRESHOLD: float = 0.45
STAGNATION_JACCARD_NEAR_DUP: float = 0.50
LOOP_REPEAT_THRESHOLD: int = 3
TOOL_CALL_HISTORY_MAX: int = 20
# D1 — Caller filter for min quality tier.
MIN_QUALITY_TIER_HEADER: str = "x-litellm-min-quality-tier"
MIN_QUALITY_TIER_METADATA_KEY: str = "min_quality_tier"
# Pre-routing -> post-call relay: the chosen logical model is stashed on
# request_kwargs["metadata"][ADAPTIVE_ROUTER_CHOSEN_MODEL_KEY] by the
# pre-routing hook, then read by the post-call hook to surface as the
# ADAPTIVE_ROUTER_RESPONSE_HEADER response header.
ADAPTIVE_ROUTER_CHOSEN_MODEL_KEY: str = "adaptive_router_chosen_model"
ADAPTIVE_ROUTER_RESPONSE_HEADER: str = "x-litellm-adaptive-router-model"

View file

@ -0,0 +1,278 @@
"""
Post-call hook for the adaptive router.
On each successful or failed completion, build a Turn from the request/response
and push it through `AdaptiveRouter.record_turn`. The router then updates the
in-memory bandit cell + session state and queues writes for the proxy flusher.
All work happens after the response has been returned to the caller. Any
exception is swallowed signal recording must never break a request.
"""
from __future__ import annotations
import hashlib
import json
from typing import Any, Dict, List, Optional
from litellm._logging import verbose_router_logger
from litellm.integrations.custom_logger import CustomLogger
from litellm.router_strategy.adaptive_router.adaptive_router import AdaptiveRouter
from litellm.router_strategy.adaptive_router.classifier import classify_prompt
from litellm.router_strategy.adaptive_router.config import (
ADAPTIVE_ROUTER_CHOSEN_MODEL_KEY,
ADAPTIVE_ROUTER_RESPONSE_HEADER,
SIGNAL_GATE_MIN_MESSAGES,
)
from litellm.router_strategy.adaptive_router.signals import Turn
# Identity fields hashed into a derived session key so the same conversation
# from the same caller produces a stable key, while different keys/teams/users
# stay segregated even if they happen to send identical first messages.
_IDENTITY_FIELDS = (
"user_api_key_hash",
"user_api_key_team_id",
"user_api_key_user_id",
"user_api_key_end_user_id",
)
def _resolve_session_key(kwargs: Dict[str, Any]) -> Optional[str]:
"""Pick a stable per-conversation key for owner-cache attribution.
Order:
1. Honor a client-supplied session id (`litellm_session_id` on either
`litellm_params` or `litellm_params.metadata`, or `session_id` on
metadata) backward compat for callers already wired up.
2. Otherwise derive a sha256 over (identity fields, first
SIGNAL_GATE_MIN_MESSAGES messages) so the key is stable across turns
and only materialises once there is enough context for the bandit to
act on (matching the gate in the signal-processing path).
Returns None if the conversation is shorter than SIGNAL_GATE_MIN_MESSAGES.
"""
litellm_params = kwargs.get("litellm_params") or {}
sid = litellm_params.get("litellm_session_id")
if sid:
return str(sid)
metadata = litellm_params.get("metadata") or {}
if isinstance(metadata, dict):
sid = metadata.get("session_id") or metadata.get("litellm_session_id")
if sid:
return str(sid)
messages = kwargs.get("messages") or []
if len(messages) < SIGNAL_GATE_MIN_MESSAGES:
# Don't attribute until we have enough turns to match the signal gate —
# ensures the hash is stable (same N messages every time) and avoids
# crediting the bandit for conversations that are too short to signal.
return None
identity = ":".join(
str(metadata.get(f) or "") if isinstance(metadata, dict) else ""
for f in _IDENTITY_FIELDS
)
anchor = messages[:SIGNAL_GATE_MIN_MESSAGES]
payload = (
identity
+ "|"
+ json.dumps(
[{"role": m.get("role"), "content": m.get("content")} for m in anchor],
sort_keys=True,
default=str,
)
)
return hashlib.sha256(payload.encode("utf-8")).hexdigest()
def _last_user_content(messages: Optional[List[Dict[str, Any]]]) -> Optional[str]:
if not messages:
return None
for msg in reversed(messages):
if msg.get("role") == "user":
content = msg.get("content")
if isinstance(content, str):
return content
if isinstance(content, list):
# OpenAI vision-style content: pick first text part.
for part in content:
if isinstance(part, dict) and part.get("type") == "text":
return part.get("text")
return None
return None
def _recent_tool_results(
messages: Optional[List[Dict[str, Any]]]
) -> List[Dict[str, Any]]:
"""Extract the current turn's tool result payloads from the request messages.
Tool results are `role == "tool"` messages that sit at the tail of the
conversation i.e. after the most recent assistant message with
`tool_calls`, waiting for the model to produce a user-facing reply. Walk
backwards from the end and collect the contiguous run of tool messages;
stop at the first non-tool message.
Each result is normalized to `{content, is_error}` the only fields
`signals._detect_failure` / `_detect_exhaustion` actually read.
"""
if not messages:
return []
results: List[Dict[str, Any]] = []
for msg in reversed(messages):
if not isinstance(msg, dict):
break
if msg.get("role") != "tool":
break
content = msg.get("content")
# Some providers (Anthropic-style) carry an explicit error flag; OpenAI
# tool results don't, so fall back to an empty/missing content heuristic
# inside `_detect_failure`.
is_error = bool(msg.get("is_error"))
results.append({"content": content, "is_error": is_error})
results.reverse()
return results
def _assistant_content_and_tool_calls(response_obj: Any) -> tuple:
"""Return (assistant_text, tool_calls_list) extracted from a ModelResponse-ish object."""
if response_obj is None:
return None, []
try:
choices = getattr(response_obj, "choices", None) or response_obj.get("choices")
except Exception:
return None, []
if not choices:
return None, []
msg = choices[0]
msg = getattr(msg, "message", None) or (
msg.get("message") if isinstance(msg, dict) else None
)
if msg is None:
return None, []
content = getattr(msg, "content", None)
if content is None and isinstance(msg, dict):
content = msg.get("content")
raw_tool_calls = getattr(msg, "tool_calls", None)
if raw_tool_calls is None and isinstance(msg, dict):
raw_tool_calls = msg.get("tool_calls")
tool_calls: List[Dict[str, Any]] = []
for tc in raw_tool_calls or []:
if isinstance(tc, dict):
tool_calls.append(tc)
else:
try:
tool_calls.append(tc.model_dump())
except Exception:
tool_calls.append({"name": getattr(tc, "name", ""), "arguments": ""})
return content, tool_calls
class AdaptiveRouterPostCallHook(CustomLogger):
"""One hook instance per AdaptiveRouter. Registered into litellm.callbacks."""
def __init__(self, adaptive_router: AdaptiveRouter) -> None:
self.adaptive_router = adaptive_router
async def async_post_call_response_headers_hook(
self,
data: Dict[str, Any],
user_api_key_dict: Any,
response: Any,
request_headers: Optional[Dict[str, str]] = None,
litellm_call_info: Optional[Dict[str, Any]] = None,
) -> Optional[Dict[str, str]]:
"""
Surface the chosen logical model as the `x-litellm-adaptive-router-model`
response header for both streaming and non-streaming responses.
`async_post_call_success_hook` fires after the stream is fully consumed,
so writing to `_hidden_params["additional_headers"]` there is too late for
streaming the StreamingResponse headers are already frozen. This hook is
called during header construction (before StreamingResponse is built), so
the header is included for both paths.
"""
metadata = data.get("metadata") or {}
chosen = (
metadata.get(ADAPTIVE_ROUTER_CHOSEN_MODEL_KEY)
if isinstance(metadata, dict)
else None
)
if not chosen:
return None
return {ADAPTIVE_ROUTER_RESPONSE_HEADER: chosen}
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
await self._record(kwargs, response_obj, response_status=200)
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time):
status = kwargs.get("response_status")
if status is None:
exc = kwargs.get("exception")
status = getattr(exc, "status_code", 500) if exc is not None else 500
await self._record(kwargs, response_obj, response_status=int(status))
async def _record(
self,
kwargs: Dict[str, Any],
response_obj: Any,
response_status: int,
) -> None:
try:
messages = kwargs.get("messages") or []
if len(messages) < SIGNAL_GATE_MIN_MESSAGES:
# Too few turns for any signal to be meaningful — skip.
return
session_key = _resolve_session_key(kwargs)
if not session_key:
return
# The bandit cells are keyed by the *logical* model name from
# `available_models` (e.g. "smart"/"fast"). `kwargs["model"]` at
# post-call time is the physical upstream model
# (e.g. "anthropic/claude-opus-4-7"), so it cannot be used directly.
# The pre-routing hook stashes the logical pick under this key.
litellm_params = kwargs.get("litellm_params") or {}
metadata = litellm_params.get("metadata") or {}
current_model = (
metadata.get(ADAPTIVE_ROUTER_CHOSEN_MODEL_KEY)
if isinstance(metadata, dict)
else None
)
if not current_model:
return
if not self.adaptive_router.claim_or_check_owner(
session_key, current_model
):
# A different model owns this conversation — skip attribution.
return
user_text = _last_user_content(messages)
assistant_text, tool_calls = _assistant_content_and_tool_calls(response_obj)
tool_results = _recent_tool_results(messages)
request_type = classify_prompt(user_text or "")
turn = Turn(
user_content=user_text,
assistant_content=(
assistant_text if isinstance(assistant_text, str) else None
),
tool_calls=tool_calls,
tool_results=tool_results,
response_status=response_status,
)
await self.adaptive_router.record_turn(
session_id=session_key,
model_name=current_model,
request_type=request_type,
turn=turn,
)
except Exception as e:
verbose_router_logger.exception(
"AdaptiveRouterPostCallHook: failed to record turn: %s", e
)

View file

@ -0,0 +1,287 @@
"""
Incremental signal detection for the adaptive router.
Each session maintains a SessionState. On every turn, we call apply_turn(state, turn)
which mutates the state in place and returns a SignalDelta listing which signals
fired on THIS turn. The router then queues the delta to be flushed to DB.
Design constraint: O(1) work per turn. No re-scanning the full session history.
We keep small bounded windows: last_user_content, last_assistant_content, and a
bounded list of recent tool call signatures.
"""
from __future__ import annotations
import re
from dataclasses import dataclass, field
from typing import Any, Dict, List, Optional, Set
from litellm.router_strategy.adaptive_router.config import (
LOOP_REPEAT_THRESHOLD,
MIN_TURNS_FOR_CLEAN_CREDIT,
MISALIGNMENT_JACCARD_THRESHOLD,
STAGNATION_JACCARD_NEAR_DUP,
TOOL_CALL_HISTORY_MAX,
)
# ---- Public types ---------------------------------------------------------
@dataclass
class SignalDelta:
"""Which signals fired on a single turn. Counts are 0 or 1 (one delta per turn)."""
misalignment: int = 0
stagnation: int = 0
disengagement: int = 0
satisfaction: int = 0
failure: int = 0
loop: int = 0
exhaustion: int = 0
def any_fired(self) -> bool:
return any(
[
self.misalignment,
self.stagnation,
self.disengagement,
self.satisfaction,
self.failure,
self.loop,
self.exhaustion,
]
)
@dataclass
class SessionState:
"""In-memory rolling state for one session.
Mirrors the LiteLLM_AdaptiveRouterSession DB row (Wave 0 schema). The flusher
later persists this. We keep this as a plain dataclass no DB coupling.
"""
session_id: str
router_name: str
model_name: str
classified_type: str
misalignment_count: int = 0
stagnation_count: int = 0
disengagement_count: int = 0
satisfaction_count: int = 0
failure_count: int = 0
loop_count: int = 0
exhaustion_count: int = 0
last_user_content: Optional[str] = None
last_assistant_content: Optional[str] = None
tool_call_history: List[str] = field(default_factory=list)
pending_tool_calls: Dict[str, str] = field(default_factory=dict)
turn_count: int = 0
last_processed_turn: int = -1
clean_credit_awarded: bool = False
terminal_status: Optional[int] = None
@dataclass
class Turn:
"""One turn of input. Caller assembles this from the request/response."""
user_content: Optional[str] = None
assistant_content: Optional[str] = None
tool_calls: List[Dict[str, Any]] = field(default_factory=list)
tool_results: List[Dict[str, Any]] = field(default_factory=list)
response_status: Optional[int] = None
# ---- Detection helpers ----------------------------------------------------
_TOKEN_RE = re.compile(r"[A-Za-z0-9]+")
def _tokens(text: Optional[str]) -> Set[str]:
if not text:
return set()
return {t.lower() for t in _TOKEN_RE.findall(text)}
def _jaccard(a: Set[str], b: Set[str]) -> float:
union = a | b
if not union:
return 0.0
return len(a & b) / len(union)
_DISENGAGEMENT_PATTERNS = [
re.compile(
r"\b(forget it|never mind|give up|talk to (?:a )?human|cancel)\b", re.IGNORECASE
),
re.compile(r"\b(this (?:isn'?t|is not) working|stop|abort)\b", re.IGNORECASE),
re.compile(r"\bi'?ll do it (?:myself|manually)\b", re.IGNORECASE),
]
_SATISFACTION_PATTERNS = [
re.compile(
r"\b(that worked|that did it|works now|fixed it|solved it|nice)\b",
re.IGNORECASE,
),
re.compile(r"\b(thanks|thank you|thx|appreciated|appreciate it)\b", re.IGNORECASE),
re.compile(r"\b(perfect|great|excellent|exactly)\b", re.IGNORECASE),
]
def _detect_misalignment(prev_user: Optional[str], curr_user: Optional[str]) -> bool:
"""Fires when consecutive user messages share *some* topic (jaccard > 0)
but are sufficiently different (jaccard < threshold) i.e. user is
rephrasing, not changing topic, not repeating."""
if not prev_user or not curr_user:
return False
j = _jaccard(_tokens(prev_user), _tokens(curr_user))
return 0.0 < j < MISALIGNMENT_JACCARD_THRESHOLD
def _detect_stagnation(prev_asst: Optional[str], curr_asst: Optional[str]) -> bool:
"""Fires when consecutive assistant messages are near-duplicates."""
if not prev_asst or not curr_asst:
return False
j = _jaccard(_tokens(prev_asst), _tokens(curr_asst))
return j >= STAGNATION_JACCARD_NEAR_DUP
def _detect_disengagement(curr_user: Optional[str]) -> bool:
if not curr_user:
return False
return any(p.search(curr_user) for p in _DISENGAGEMENT_PATTERNS)
def _detect_satisfaction(curr_user: Optional[str]) -> bool:
if not curr_user:
return False
return any(p.search(curr_user) for p in _SATISFACTION_PATTERNS)
def _detect_failure(tool_results: List[Dict[str, Any]]) -> bool:
"""Any tool result explicitly flagged as an error.
We do NOT treat empty content as failure many tools legitimately return
empty output (zero-result searches, silent bash commands, void writes) and
penalizing the model for those would corrupt the bandit posterior.
"""
for r in tool_results:
if r.get("is_error"):
return True
return False
def _signature(call: Dict[str, Any]) -> str:
"""Stable signature for loop detection: name + sorted JSON-ish args."""
name = call.get("name") or call.get("function", {}).get("name", "")
call_args = call.get("arguments")
if call_args is None:
call_args = call.get("function", {}).get("arguments", "")
if isinstance(call_args, dict):
call_args = ",".join(f"{k}={call_args[k]}" for k in sorted(call_args.keys()))
return f"{name}({call_args})"
def _detect_loop(history: List[str], new_calls: List[Dict[str, Any]]) -> bool:
"""Fires if any new call's signature appears >= LOOP_REPEAT_THRESHOLD-1 times
in recent history (so this call would be the Nth)."""
if not new_calls:
return False
for call in new_calls:
sig = _signature(call)
recent_count = history.count(sig)
if recent_count >= LOOP_REPEAT_THRESHOLD - 1:
return True
return False
_EXHAUSTION_STATUSES = {408, 413, 429, 503, 504}
_EXHAUSTION_KEYWORDS = (
"context length",
"context window",
"token limit",
"rate limit",
"too many requests",
"timeout",
)
def _detect_exhaustion(
status: Optional[int], tool_results: List[Dict[str, Any]]
) -> bool:
if status is not None and status in _EXHAUSTION_STATUSES:
return True
for r in tool_results:
content = str(r.get("content", "")).lower()
if any(kw in content for kw in _EXHAUSTION_KEYWORDS):
return True
return False
# ---- Public entrypoint ----------------------------------------------------
def apply_turn(state: SessionState, turn: Turn) -> SignalDelta:
"""
Detect signals on this turn, mutate state, return the delta.
O(1) per turn (no full-history rescan). Only inspects last_*, recent tool history
(which is bounded at TOOL_CALL_HISTORY_MAX), and the new turn payload.
"""
delta = SignalDelta()
if _detect_misalignment(state.last_user_content, turn.user_content):
delta.misalignment = 1
if _detect_stagnation(state.last_assistant_content, turn.assistant_content):
delta.stagnation = 1
if _detect_disengagement(turn.user_content):
delta.disengagement = 1
if _detect_satisfaction(turn.user_content):
# Gate: only award satisfaction credit once per session, and only
# after MIN_TURNS_FOR_CLEAN_CREDIT turns of context. Early "thanks"
# on turn 1-2 is noise, not a validated quality signal.
current_turn_index = state.turn_count + 1
if (
not state.clean_credit_awarded
and current_turn_index >= MIN_TURNS_FOR_CLEAN_CREDIT
):
delta.satisfaction = 1
state.clean_credit_awarded = True
if _detect_failure(turn.tool_results):
delta.failure = 1
if _detect_loop(state.tool_call_history, turn.tool_calls):
delta.loop = 1
if _detect_exhaustion(turn.response_status, turn.tool_results):
delta.exhaustion = 1
state.misalignment_count += delta.misalignment
state.stagnation_count += delta.stagnation
state.disengagement_count += delta.disengagement
state.satisfaction_count += delta.satisfaction
state.failure_count += delta.failure
state.loop_count += delta.loop
state.exhaustion_count += delta.exhaustion
if turn.user_content:
state.last_user_content = turn.user_content
if turn.assistant_content:
state.last_assistant_content = turn.assistant_content
for call in turn.tool_calls:
state.tool_call_history.append(_signature(call))
if len(state.tool_call_history) > TOOL_CALL_HISTORY_MAX:
state.tool_call_history = state.tool_call_history[-TOOL_CALL_HISTORY_MAX:]
if turn.response_status is not None:
state.terminal_status = turn.response_status
state.turn_count += 1
state.last_processed_turn = state.turn_count
return delta

View file

@ -0,0 +1,213 @@
"""
In-memory queues for adaptive router state and session updates.
Pattern follows DailySpendUpdateQueue: hot path is fully in-memory; a background
flusher task drains the aggregator and writes batches to Postgres.
Two logical queues (one class):
1. STATE updates: increments to (router, request_type, model) bandit cell.
Aggregator key = (router_name, request_type, model_name)
Aggregated payload = {"delta_alpha": float, "delta_beta": float, "samples_added": int}
2. SESSION updates: full snapshot of a session row (last-write-wins per session+router+model).
Aggregator key = (session_id, router_name, model_name)
Aggregated payload = the full session state dict.
Hot-path API is non-blocking and synchronous from the caller's POV (it just appends
to the in-memory aggregator). Flush is async and batched.
"""
from __future__ import annotations
import asyncio
from typing import Any, Dict, Tuple
from litellm._logging import verbose_router_logger
StateKey = Tuple[str, str, str] # (router_name, request_type, model_name)
SessionKey = Tuple[str, str, str] # (session_id, router_name, model_name)
class AdaptiveRouterUpdateQueue:
"""
Single class managing both state-update aggregation and session-snapshot aggregation.
Held by the AdaptiveRouter strategy instance and started by the proxy on boot.
"""
def __init__(self) -> None:
self._state_agg: Dict[StateKey, Dict[str, float]] = {}
self._session_agg: Dict[SessionKey, Dict[str, Any]] = {}
self._lock = asyncio.Lock()
self._max_state_size_seen = 0
self._max_session_size_seen = 0
# ---- Hot-path: state delta -------------------------------------------
async def add_state_delta(
self,
router_name: str,
request_type: str,
model_name: str,
delta_alpha: float,
delta_beta: float,
) -> None:
"""Aggregate a bandit-cell delta. Multiple deltas to the same cell sum."""
key: StateKey = (router_name, request_type, model_name)
async with self._lock:
current = self._state_agg.get(key)
if current is None:
self._state_agg[key] = {
"delta_alpha": delta_alpha,
"delta_beta": delta_beta,
"samples_added": 1,
}
else:
current["delta_alpha"] += delta_alpha
current["delta_beta"] += delta_beta
current["samples_added"] += 1
if len(self._state_agg) > self._max_state_size_seen:
self._max_state_size_seen = len(self._state_agg)
# ---- Hot-path: session snapshot --------------------------------------
async def add_session_state(
self,
session_id: str,
router_name: str,
model_name: str,
state_dict: Dict[str, Any],
) -> None:
"""
Last-write-wins per session row. The state_dict is a snapshot of the
SessionState (signals counts + bookkeeping fields). The flusher will
upsert this into LiteLLM_AdaptiveRouterSession.
"""
key: SessionKey = (session_id, router_name, model_name)
async with self._lock:
self._session_agg[key] = state_dict
if len(self._session_agg) > self._max_session_size_seen:
self._max_session_size_seen = len(self._session_agg)
# ---- Flushers (called by background task) ----------------------------
async def flush_state_to_db(self, prisma_client: Any) -> int:
"""
Drain state aggregator and apply to LiteLLM_AdaptiveRouterState.
Returns number of cells flushed.
"""
async with self._lock:
batch = self._state_agg
self._state_agg = {}
if not batch:
return 0
# Sort keys to give deterministic write order across writers and
# reduce the chance of cross-row deadlocks when other workers race us.
for key in sorted(batch.keys()):
router, rt, model = key
payload = batch[key]
try:
# Atomic increment: push the delta directly into the DB so
# concurrent flushers from multiple pods don't overwrite each
# other. The upsert creates the row with the delta as the
# initial value on first write, then increments on subsequent
# writes — no read-modify-write race.
await prisma_client.db.litellm_adaptiverouterstate.upsert(
where={
"router_name_request_type_model_name": {
"router_name": router,
"request_type": rt,
"model_name": model,
}
},
data={
"create": {
"router_name": router,
"request_type": rt,
"model_name": model,
"alpha": payload["delta_alpha"],
"beta": payload["delta_beta"],
"total_samples": int(payload["samples_added"]),
},
"update": {
"alpha": {"increment": payload["delta_alpha"]},
"beta": {"increment": payload["delta_beta"]},
"total_samples": {
"increment": int(payload["samples_added"])
},
},
},
)
except Exception as e:
verbose_router_logger.exception(
"AdaptiveRouterUpdateQueue: failed to flush state for %s: %s",
key,
e,
)
return len(batch)
async def flush_session_to_db(self, prisma_client: Any) -> int:
"""
Drain session aggregator and upsert into LiteLLM_AdaptiveRouterSession.
Returns number of session rows flushed.
"""
async with self._lock:
batch = self._session_agg
self._session_agg = {}
if not batch:
return 0
for key in sorted(batch.keys()):
session_id, router, model = key
payload = batch[key]
try:
# NOTE: Prisma client lower-cases model names, so
# `LiteLLM_AdaptiveRouterSession` -> `litellm_adaptiveroutersession`
# (single 's', not 'litellm_adaptiverouterssession').
# Strip PK fields from the update payload — Prisma rejects
# writes to fields that are part of the @@id. asdict(state)
# always carries them, so build a separate update dict.
update_payload = {
k: v
for k, v in payload.items()
if k not in ("session_id", "router_name", "model_name")
}
await prisma_client.db.litellm_adaptiveroutersession.upsert(
where={
"session_id_router_name_model_name": {
"session_id": session_id,
"router_name": router,
"model_name": model,
}
},
data={
"create": {
"session_id": session_id,
"router_name": router,
"model_name": model,
**update_payload,
},
"update": update_payload,
},
)
except Exception as e:
verbose_router_logger.exception(
"AdaptiveRouterUpdateQueue: failed to flush session for %s: %s",
key,
e,
)
return len(batch)
# ---- Observability ---------------------------------------------------
async def queue_size(self) -> Dict[str, int]:
async with self._lock:
return {
"state_pending": len(self._state_agg),
"session_pending": len(self._session_agg),
"max_state_seen": self._max_state_size_seen,
"max_session_seen": self._max_session_size_seen,
}

View file

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

View file

@ -8,7 +8,7 @@ from dataclasses import dataclass
from typing import Any, Dict, List, Literal, Optional, Tuple, Union, get_type_hints
import httpx
from pydantic import BaseModel, ConfigDict, Field, model_validator
from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator
from typing_extensions import Required, TypedDict
from litellm._uuid import uuid
@ -221,6 +221,9 @@ class GenericLiteLLMParams(CredentialLiteLLMParams, CustomPricingLiteLLMParams):
complexity_router_config: Optional[Dict] = None
complexity_router_default_model: Optional[str] = None
# 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
@ -794,3 +797,44 @@ class PreRoutingHookResponse(BaseModel):
model: str
messages: Optional[List[Dict[str, Any]]]
class RequestType(str, enum.Enum):
"""Fixed v0 taxonomy. User-extensible types come in v1."""
CODE_GENERATION = "code_generation"
CODE_UNDERSTANDING = "code_understanding"
TECHNICAL_DESIGN = "technical_design"
ANALYTICAL_REASONING = "analytical_reasoning"
WRITING = "writing"
FACTUAL_LOOKUP = "factual_lookup"
GENERAL = "general"
class AdaptiveRouterWeights(BaseModel):
quality: float = Field(default=0.7, ge=0.0, le=1.0)
cost: float = Field(default=0.3, ge=0.0, le=1.0)
@field_validator("cost")
@classmethod
def _weights_sum_to_one(cls, v, info):
q = info.data.get("quality", 0.7)
if abs(q + v - 1.0) > 0.001:
raise ValueError(
f"weights must sum to 1.0, got quality={q} + cost={v} = {q + v}"
)
return v
class AdaptiveRouterConfig(BaseModel):
available_models: List[str]
weights: AdaptiveRouterWeights = Field(default_factory=AdaptiveRouterWeights)
class AdaptiveRouterPreferences(BaseModel):
"""model_info.adaptive_router_preferences — declared by each model."""
model_config = ConfigDict(use_enum_values=False)
quality_tier: int = Field(ge=1, le=3)
strengths: List[RequestType] = Field(default_factory=list)

View file

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

View file

@ -1148,6 +1148,20 @@
"tool_use_system_prompt_tokens": 346,
"supports_native_structured_output": true
},
"anthropic.claude-mythos-preview": {
"input_cost_per_token": 0,
"output_cost_per_token": 0,
"litellm_provider": "bedrock",
"max_input_tokens": 1000000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"supports_function_calling": true,
"supports_vision": true,
"supports_prompt_caching": false,
"supports_reasoning": true,
"supports_tool_choice": true
},
"global.anthropic.claude-opus-4-7": {
"cache_creation_input_token_cost": 6.25e-06,
"cache_read_input_token_cost": 5e-07,
@ -22872,6 +22886,22 @@
"supports_video_input": true,
"supports_vision": true
},
"moonshot/kimi-k2.6": {
"cache_read_input_token_cost": 1.6e-07,
"input_cost_per_token": 9.5e-07,
"litellm_provider": "moonshot",
"max_input_tokens": 262144,
"max_output_tokens": 262144,
"max_tokens": 262144,
"mode": "chat",
"output_cost_per_token": 4e-06,
"source": "https://platform.kimi.ai/docs/pricing/chat-k26",
"supports_function_calling": true,
"supports_reasoning": true,
"supports_tool_choice": true,
"supports_video_input": true,
"supports_vision": true
},
"moonshot/kimi-latest": {
"cache_read_input_token_cost": 1.5e-07,
"input_cost_per_token": 2e-06,
@ -25135,6 +25165,28 @@
"supports_vision": true,
"tool_use_system_prompt_tokens": 346
},
"openrouter/anthropic/claude-opus-4.7": {
"cache_creation_input_token_cost": 6.25e-06,
"cache_read_input_token_cost": 5e-07,
"input_cost_per_token": 5e-06,
"litellm_provider": "openrouter",
"max_input_tokens": 1000000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 2.5e-05,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
"supports_xhigh_reasoning_effort": true,
"tool_use_system_prompt_tokens": 346
},
"openrouter/bytedance/ui-tars-1.5-7b": {
"input_cost_per_token": 1e-07,
"litellm_provider": "openrouter",

View file

@ -1,6 +1,6 @@
[project]
name = "litellm"
version = "1.83.10"
version = "1.83.12"
description = "Library to easily interface with LLM API providers"
readme = "README.md"
requires-python = ">=3.10, <3.14"
@ -52,7 +52,7 @@ proxy = [
"azure-identity==1.25.2",
"azure-storage-blob==12.28.0",
"mcp==1.26.0",
"litellm-proxy-extras==0.4.67",
"litellm-proxy-extras==0.4.68",
"litellm-enterprise==0.1.38",
"RestrictedPython==8.1",
"rich==13.9.4",
@ -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]
@ -236,7 +236,7 @@ source-exclude = [
profile = "black"
[tool.commitizen]
version = "1.83.10"
version = "1.83.12"
version_files = [
"pyproject.toml:^version",
]

View file

@ -616,6 +616,7 @@ model LiteLLM_TeamMembership {
user_id String
team_id String
spend Float @default(0.0)
total_spend Float @default(0.0)
budget_id String?
litellm_budget_table LiteLLM_BudgetTable? @relation(fields: [budget_id], references: [budget_id])
@@id([user_id, team_id])
@ -1223,3 +1224,46 @@ model LiteLLM_ClaudeCodePluginTable {
@@map("LiteLLM_ClaudeCodePluginTable")
}
// Per-(router, request_type, model) Beta posterior for the adaptive router.
model LiteLLM_AdaptiveRouterState {
router_name String
request_type String
model_name String
alpha Float
beta Float
total_samples Int @default(0)
last_updated_at DateTime @default(now()) @updatedAt
@@id([router_name, request_type, model_name])
}
// Per-(session, router, model) signal counters for the adaptive router.
model LiteLLM_AdaptiveRouterSession {
session_id String
router_name String
model_name String
classified_type String
misalignment_count Int @default(0)
stagnation_count Int @default(0)
disengagement_count Int @default(0)
satisfaction_count Int @default(0)
failure_count Int @default(0)
loop_count Int @default(0)
exhaustion_count Int @default(0)
last_user_content String?
last_assistant_content String?
tool_call_history Json @default("[]")
pending_tool_calls Json @default("{}")
turn_count Int @default(0)
last_processed_turn Int @default(-1)
clean_credit_awarded Boolean @default(false)
terminal_status Int?
last_activity_at DateTime @default(now()) @updatedAt
@@id([session_id, router_name, model_name])
@@index([last_activity_at], map: "idx_adaptive_router_session_activity")
}

View file

@ -0,0 +1,157 @@
# Adaptive Router — Live Demo
A 5-minute demo of LiteLLM's adaptive router learning, in real time, that
the smart model wins for code while the fast model is fine for facts.
```
┌─ traffic.py ──┐ ┌─ litellm proxy ──────────┐ ┌─ dashboard.html ─┐
│ synthetic │──▶│ adaptive_router strategy │──▶│ bandit bars + │
│ chat sessions │ │ /adaptive_router/state │ │ cost meter + │
└───────────────┘ └──────────┬───────────────┘ │ activity log │
│ └───────────────────┘
┌─────────▼───────────┐
│ chat.html │
│ interactive chat │
│ with preset │
│ scenarios │
└─────────────────────┘
```
## Files
| File | What it does |
|---|---|
| `dashboard.html` | Live bandit dashboard — polls `/adaptive_router/state` every 500ms |
| `chat.html` | Interactive chat with preset scenarios — sends real requests through the router |
| `traffic.py` | Synthetic traffic generator — drives labeled sessions for automated demo |
## What you're watching
- **Bandit posteriors** — one Beta(α, β) bar per `(request_type, model)`
cell. Bars fill up as α grows from positive feedback signals.
- **Pick share** — softmax estimate of how often the router would currently
pick each model for that request type.
- **Cost meter** — total spend so far compared to "always use the most
expensive model". The savings line is the headline number.
- **Activity log** — every signal that moves the bandit, in real time.
## 1. Start the proxy
The repo ships with a working example config:
```bash
export OPENAI_API_KEY=sk-... # underlying models hit OpenAI
uv run litellm \
--config litellm/proxy/example_config_yaml/adaptive_router_example.yaml \
--port 4000
```
`DATABASE_URL` is optional — the proxy falls back to a bundled Neon dev DB.
Wait ~15s until you see `Application startup complete`.
## 2. Chat interactively with the router
Open `chat.html` in a browser (same `file://` or `python3 -m http.server` approach as the dashboard):
- Click **Connect** after filling in the proxy URL and API key.
- Pick a preset scenario:
- **🐛 Debug my code** — paste broken code and get a fix
- **💡 Brainstorm a feature** — ideate on a product capability
- **📚 Explain a concept** — get a clear technical explanation
- **✍️ Write something** — draft emails, docs, or any prose
- A starter message is pre-filled — edit it or send as-is.
- Each response shows which model the router picked and the inferred request type (from the `x-litellm-adaptive-router-model` and `x-litellm-request-type` response headers).
- A sidebar gate indicator tells you when the session has accumulated enough messages for the bandit to start updating (4+ turns).
> **Note on headers:** The model/type headers are only readable in the browser if the proxy sets `Access-Control-Expose-Headers`. LiteLLM defaults to exposing them. If the info panel shows `check dashboard`, the router still works — you can verify picks in `dashboard.html`.
## 4. Open the dashboard
The dashboard is a single static HTML file. Either:
- **Easy:** double-click `dashboard.html`. Most browsers will load it from
`file://` and the LiteLLM proxy's CORS defaults (`*`) will accept it.
- **If your browser blocks `file://` fetches:**
```bash
cd scripts/adaptive_router_demo
python3 -m http.server 8080
```
Then open <http://localhost:8080/dashboard.html>.
In the connect bar, fill in:
- **Proxy URL:** `http://localhost:4000`
- **Master Key:** the `master_key` from your config (`sk-1234` in the example).
Click **Connect**. The dashboard polls `GET /adaptive_router/state` every
500ms (admin-only endpoint, returns one snapshot per configured router).
## 5. Drive synthetic traffic
In a second terminal:
```bash
uv run python scripts/adaptive_router_demo/traffic.py \
--proxy-url http://localhost:4000 \
--api-key sk-1234 \
--router smart-cheap-router \
--rounds 100 \
--rate 0.5
```
What it does:
- Picks a random `(request_type, prompt)` per round from a small labeled corpus.
- Sends a 5-message conversation (passes the `SIGNAL_GATE_MIN_MESSAGES=4` gate
in one round-trip) so the post-call hook runs and updates the bandit.
- Reads the `x-litellm-adaptive-router-model` response header to see what
the router picked.
- Rolls Bernoulli against a hard-coded oracle:
```
code_generation : smart=0.92 fast=0.35
factual_lookup : smart=0.90 fast=0.85
writing : smart=0.85 fast=0.55
```
- On success → sends a follow-up engineered to match the satisfaction
regex (and re-classify into the same type). Bandit cell gets +α.
- On failure → sends a neutral follow-up. No signal fires.
After 5080 rounds you'll see `code_generation` decisively favor `smart`
while `factual_lookup` stays near a coin flip — the router learned the
asymmetry from the oracle.
## Tuning knobs
| Knob | Where | What changes |
|---|---|---|
| Quality vs. cost weight | `adaptive_router_config.weights` in proxy yaml | Bias toward quality or savings |
| Per-cell cold-start mass | `litellm/router_strategy/adaptive_router/config.py` `COLD_START_MASS` | How long until the prior is overwritten |
| Avg tokens per request | dashboard input box | How the cost meter estimates spend |
| Oracle | `traffic.py` `ORACLE` dict | Which model "should" win for which type |
| Sessions to drive | `--rounds` | Total learning budget |
| Throttle | `--rate` | Seconds between sessions |
## Multi-router
If your proxy has more than one `auto_router/adaptive_router` deployment,
the dashboard shows a router dropdown above the bars. Each router is
independent; the cost meter is per-router (and resets when you switch).
## Troubleshooting
- **"Disconnected" / HTTP 401 in the dashboard** — wrong master key.
- **HTTP 403** — your key isn't `proxy_admin`. The state endpoint is
admin-only. Use the master key.
- **HTTP 404 from `/adaptive_router/state`** — proxy started, but no
`auto_router/adaptive_router` deployment is in the model list.
- **Bars don't move** — check the proxy logs for `record_turn` activity.
Common cause: requests are not including 4+ messages, so the signal
gate skips them. `traffic.py` already builds 5-message conversations,
so this only happens if you've changed the script.
- **Cost meter stays at $0** — your model deployments don't have
`input_cost_per_token` set in `litellm_params`. Add it.
- **CORS error in the dashboard console** — set `LITELLM_CORS_ORIGINS=*`
on the proxy (the default), or serve `dashboard.html` from
`python3 -m http.server` instead of `file://`.

View file

@ -0,0 +1,838 @@
<!DOCTYPE html>
<html lang="en">
<head>
<meta charset="UTF-8" />
<title>Adaptive Router — Chat</title>
<style>
:root {
--bg: #0b0f17;
--panel: #131a26;
--panel-2: #1b2433;
--fg: #e7ecf3;
--muted: #8a95a8;
--accent: #5dd6a4;
--accent-2: #6fb6ff;
--warn: #f6b94d;
--bad: #ff6b6b;
--bar-bg: #233047;
--border: #25324a;
}
* { box-sizing: border-box; }
body {
margin: 0;
font-family: -apple-system, BlinkMacSystemFont, "SF Pro Display",
"Segoe UI", Roboto, Inter, sans-serif;
background: var(--bg);
color: var(--fg);
font-size: 14px;
line-height: 1.45;
height: 100vh;
display: flex;
flex-direction: column;
}
header {
padding: 14px 24px;
border-bottom: 1px solid var(--border);
display: flex;
align-items: center;
gap: 16px;
background: var(--panel);
flex-shrink: 0;
}
header h1 { margin: 0; font-size: 17px; font-weight: 600; }
.dot { width: 8px; height: 8px; border-radius: 50%; background: var(--bad); display: inline-block; margin-right: 6px; }
.dot.ok { background: var(--accent); }
.status { color: var(--muted); font-size: 12px; }
.header-link { margin-left: auto; color: var(--accent-2); font-size: 12px; text-decoration: none; }
.header-link:hover { text-decoration: underline; }
.connect {
padding: 12px 24px;
display: flex; gap: 10px; align-items: center;
background: var(--panel-2);
border-bottom: 1px solid var(--border);
flex-shrink: 0;
flex-wrap: wrap;
}
.connect label { color: var(--muted); font-size: 12px; }
.connect input, .connect select {
background: #0e1422;
color: var(--fg);
border: 1px solid var(--border);
padding: 6px 10px;
border-radius: 6px;
font-size: 13px;
font-family: inherit;
}
.connect input[type=text] { width: 210px; }
.connect input[type=password] { width: 180px; }
.connect button.btn-connect {
background: var(--accent-2);
color: #0b0f17;
border: none;
padding: 7px 14px;
border-radius: 6px;
font-weight: 600;
cursor: pointer;
font-size: 13px;
}
/* scenario strip */
.scenarios {
padding: 10px 24px;
display: flex; gap: 8px; flex-wrap: wrap;
background: var(--panel);
border-bottom: 1px solid var(--border);
flex-shrink: 0;
}
.scenarios .sc-btn {
background: var(--panel-2);
border: 1px solid var(--border);
color: var(--fg);
padding: 7px 14px;
border-radius: 20px;
cursor: pointer;
font-size: 13px;
font-family: inherit;
transition: border-color 0.15s, background 0.15s;
}
.scenarios .sc-btn:hover { border-color: var(--accent-2); background: #1e2d42; }
.scenarios .sc-btn.active { border-color: var(--accent-2); background: #1a2d45; color: var(--accent-2); }
.scenarios .sc-btn.new { border-color: var(--border); color: var(--muted); }
.scenarios .sc-btn.new:hover { border-color: var(--warn); color: var(--warn); background: #1f1a10; }
/* main layout */
.workspace {
display: grid;
grid-template-columns: 1fr 300px;
flex: 1;
min-height: 0;
}
@media (max-width: 900px) {
.workspace { grid-template-columns: 1fr; }
.info-panel { display: none; }
}
/* chat panel */
.chat-panel {
display: flex;
flex-direction: column;
min-height: 0;
border-right: 1px solid var(--border);
}
.messages {
flex: 1;
overflow-y: auto;
padding: 20px 24px;
display: flex;
flex-direction: column;
gap: 16px;
}
.messages .empty-state {
margin: auto;
text-align: center;
color: var(--muted);
}
.messages .empty-state h2 {
font-size: 18px;
font-weight: 600;
color: var(--fg);
margin: 0 0 8px;
}
.messages .empty-state p {
font-size: 13px;
margin: 0;
max-width: 360px;
}
.msg {
display: flex;
gap: 12px;
align-items: flex-start;
}
.msg.assistant { flex-direction: row; }
.msg.user { flex-direction: row-reverse; }
.avatar {
width: 28px; height: 28px;
border-radius: 50%;
display: flex; align-items: center; justify-content: center;
font-size: 13px;
flex-shrink: 0;
}
.msg.user .avatar { background: var(--accent-2); color: #0b0f17; font-weight: 700; }
.msg.assistant .avatar { background: var(--accent); color: #0b0f17; }
.bubble {
max-width: 75%;
padding: 10px 14px;
border-radius: 12px;
font-size: 13px;
line-height: 1.55;
white-space: pre-wrap;
word-break: break-word;
}
.msg.user .bubble {
background: #1a2d45;
border: 1px solid var(--border);
border-top-right-radius: 4px;
}
.msg.assistant .bubble {
background: var(--panel);
border: 1px solid var(--border);
border-top-left-radius: 4px;
}
.bubble code {
font-family: ui-monospace, "SF Mono", Menlo, monospace;
background: rgba(255,255,255,0.06);
padding: 1px 4px;
border-radius: 3px;
font-size: 12px;
}
.bubble pre {
background: #0e1422;
border: 1px solid var(--border);
border-radius: 6px;
padding: 10px 12px;
overflow-x: auto;
margin: 8px 0 0;
}
.bubble pre code {
background: none;
padding: 0;
font-size: 12px;
}
.msg-meta {
font-size: 11px;
color: var(--muted);
margin-top: 4px;
padding: 0 2px;
}
.msg.user .msg-meta { text-align: right; }
.msg-model { color: var(--accent); font-weight: 600; }
.msg-type { color: var(--accent-2); }
.thinking {
display: flex; gap: 4px; align-items: center;
padding: 8px 0;
}
.thinking span {
width: 6px; height: 6px; border-radius: 50%;
background: var(--muted);
animation: blink 1.2s infinite;
}
.thinking span:nth-child(2) { animation-delay: 0.2s; }
.thinking span:nth-child(3) { animation-delay: 0.4s; }
@keyframes blink {
0%, 80%, 100% { opacity: 0.2; }
40% { opacity: 1; }
}
/* input area */
.input-area {
padding: 14px 24px;
border-top: 1px solid var(--border);
background: var(--panel);
flex-shrink: 0;
display: flex;
flex-direction: column;
gap: 8px;
}
.input-row {
display: flex;
gap: 10px;
align-items: flex-end;
}
textarea {
flex: 1;
background: #0e1422;
color: var(--fg);
border: 1px solid var(--border);
border-radius: 8px;
padding: 10px 12px;
font-size: 13px;
font-family: inherit;
resize: none;
outline: none;
min-height: 44px;
max-height: 160px;
line-height: 1.45;
}
textarea:focus { border-color: var(--accent-2); }
textarea:disabled { opacity: 0.5; }
.btn-send {
background: var(--accent);
color: #0b0f17;
border: none;
padding: 10px 18px;
border-radius: 8px;
font-weight: 700;
font-size: 13px;
cursor: pointer;
font-family: inherit;
flex-shrink: 0;
align-self: flex-end;
}
.btn-send:disabled { opacity: 0.5; cursor: default; }
.input-hint {
font-size: 11px;
color: var(--muted);
}
/* info panel */
.info-panel {
background: var(--panel);
padding: 18px 16px;
overflow-y: auto;
display: flex;
flex-direction: column;
gap: 16px;
}
.info-block h3 {
margin: 0 0 10px;
font-size: 11px;
text-transform: uppercase;
letter-spacing: 0.6px;
color: var(--muted);
font-weight: 600;
}
.info-row {
display: flex;
justify-content: space-between;
align-items: center;
font-size: 12px;
margin-bottom: 6px;
}
.info-row .label { color: var(--muted); }
.info-row .val { font-family: ui-monospace, monospace; color: var(--fg); font-weight: 600; }
.info-row .val.accent { color: var(--accent); }
.info-row .val.accent-2 { color: var(--accent-2); }
.info-row .val.warn { color: var(--warn); }
.signal-gate {
padding: 8px 10px;
border-radius: 6px;
font-size: 12px;
line-height: 1.4;
}
.signal-gate.waiting {
background: rgba(246, 185, 77, 0.08);
border: 1px solid rgba(246, 185, 77, 0.3);
color: var(--warn);
}
.signal-gate.learning {
background: rgba(93, 214, 164, 0.08);
border: 1px solid rgba(93, 214, 164, 0.3);
color: var(--accent);
}
.scenario-desc {
font-size: 12px;
color: var(--muted);
line-height: 1.5;
}
.divider {
border: none;
border-top: 1px solid var(--border);
margin: 0;
}
.link-dash {
display: block;
text-align: center;
color: var(--accent-2);
font-size: 12px;
text-decoration: none;
padding: 8px;
border: 1px solid var(--border);
border-radius: 6px;
margin-top: auto;
}
.link-dash:hover { border-color: var(--accent-2); background: #1a2d45; }
</style>
</head>
<body>
<header>
<h1>⚡ Adaptive Router — Chat</h1>
<span class="status"><span id="conn-dot" class="dot"></span><span id="conn-label">Disconnected</span></span>
<a class="header-link" href="dashboard.html">→ Open live dashboard</a>
</header>
<div class="connect">
<label>Proxy URL <input id="proxy-url" type="text" value="http://localhost:4000" /></label>
<label>API Key <input id="api-key" type="password" placeholder="sk-1234" /></label>
<label>Router
<input id="router" type="text" value="smart-cheap-router" style="width:160px" />
</label>
<button class="btn-connect" id="connect-btn">Connect</button>
</div>
<div class="scenarios" id="scenario-bar">
<button class="sc-btn" data-id="debug_code">🐛 Debug my code</button>
<button class="sc-btn" data-id="brainstorm_feature">💡 Brainstorm a feature</button>
<button class="sc-btn" data-id="explain_concept">📚 Explain a concept</button>
<button class="sc-btn" data-id="write_something">✍️ Write something</button>
<button class="sc-btn new" id="new-chat-btn"> New chat</button>
</div>
<div class="workspace">
<div class="chat-panel">
<div class="messages" id="messages">
<div class="empty-state">
<h2>Pick a scenario to start</h2>
<p>Choose one of the presets above or connect to the proxy and type your own message. The adaptive router will pick the best model for each turn.</p>
</div>
</div>
<div class="input-area">
<div class="input-row">
<textarea id="input" rows="1" placeholder="Send a message… (Shift+Enter for new line)" disabled></textarea>
<button class="btn-send" id="send-btn" disabled>Send</button>
</div>
<div class="input-hint" id="input-hint">Connect first to start chatting.</div>
</div>
</div>
<aside class="info-panel">
<div class="info-block">
<h3>Session</h3>
<div class="info-row"><span class="label">ID</span><span class="val" id="info-session-id"></span></div>
<div class="info-row"><span class="label">Messages</span><span class="val" id="info-msg-count">0</span></div>
<div class="info-row"><span class="label">Scenario</span><span class="val accent-2" id="info-scenario">none</span></div>
</div>
<div id="gate-status" class="signal-gate waiting" style="display:none">
<b>Learning starts at 4 messages.</b> Keep chatting — the router will start updating its bandit after your next reply.
</div>
<hr class="divider" />
<div class="info-block">
<h3>Last response</h3>
<div class="info-row"><span class="label">Model picked</span><span class="val accent" id="info-model"></span></div>
<div class="info-row"><span class="label">Request type</span><span class="val accent-2" id="info-req-type"></span></div>
<div class="info-row"><span class="label">Latency</span><span class="val" id="info-latency"></span></div>
</div>
<hr class="divider" />
<div class="info-block">
<h3>How it works</h3>
<p class="scenario-desc">
Each message goes through the <b>adaptive router</b> which classifies your request type (code, writing, factual…) and uses a Thompson-sampling bandit to pick the model with the best quality for that category.<br><br>
After 4+ messages, positive feedback signals (✓ in the activity log) update the bandit. Watch the bars move in the <a href="dashboard.html" style="color:var(--accent-2)">live dashboard</a>.
</p>
</div>
<a class="link-dash" href="dashboard.html">📊 Live bandit dashboard →</a>
</aside>
</div>
<script>
// ---- scenarios -------------------------------------------------------
const SCENARIOS = {
debug_code: {
label: "Debug my code",
system: "You are an expert debugging assistant. Be concise. Identify the bug, explain why it's wrong, and provide the corrected code.",
starter: "I have a Python function that should return the sum of a list, but it always returns 0:\n\n```python\ndef sum_list(items):\n total = 0\n for item in items:\n total + item\n return total\n\nprint(sum_list([1, 2, 3])) # prints 0, expected 6\n```\n\nWhat's wrong with it?",
},
brainstorm_feature: {
label: "Brainstorm a feature",
system: "You are a product thinking partner. Help explore feature ideas with concrete examples, trade-offs, and implementation considerations. Be specific and opinionated.",
starter: "I'm building a note-taking app for developers. What are 5 differentiated features that would make it stand out from Notion or Obsidian? Focus on things that would genuinely solve developer pain points.",
},
explain_concept: {
label: "Explain a concept",
system: "You are a clear, concise technical educator. Explain concepts with simple language, good analogies, and concrete examples. Avoid unnecessary jargon.",
starter: "Can you explain how the Thompson Sampling algorithm works and why it's better than epsilon-greedy for multi-armed bandit problems? Use a concrete example if it helps.",
},
write_something: {
label: "Write something",
system: "You are a skilled writer. Produce clear, professional text tailored to the requested format and tone. Match the voice the user asks for.",
starter: "Write a short Slack message to my team letting them know our weekly standup is moving from 9am to 10am starting next Monday. Keep it brief, friendly, and include a clear ask for them to update their calendars.",
},
};
// ---- state -----------------------------------------------------------
const STATE = {
connected: false,
proxyUrl: "",
apiKey: "",
router: "",
sessionId: "",
messages: [], // [{role, content}] — sent to the API
scenario: null,
sending: false,
msgCount: 0, // turns in current session
lastModel: null,
lastReqType: null,
};
// ---- session ---------------------------------------------------------
function newSession() {
STATE.sessionId = "chat-" + Math.random().toString(36).slice(2, 10);
STATE.messages = [];
STATE.msgCount = 0;
STATE.lastModel = null;
STATE.lastReqType = null;
renderInfo();
renderGateStatus();
}
// ---- persistence -----------------------------------------------------
function ssGet(k) { try { return sessionStorage.getItem(k) || ""; } catch { return ""; } }
function ssSet(k, v) { try { sessionStorage.setItem(k, v); } catch {} }
// ---- rendering -------------------------------------------------------
function setConn(ok, label) {
document.getElementById("conn-dot").className = "dot" + (ok ? " ok" : "");
document.getElementById("conn-label").textContent = ok ? "Connected" : label || "Disconnected";
}
function renderInfo() {
const shortId = STATE.sessionId ? STATE.sessionId.slice(-8) : "—";
document.getElementById("info-session-id").textContent = shortId;
document.getElementById("info-msg-count").textContent = STATE.msgCount;
document.getElementById("info-scenario").textContent = STATE.scenario ? SCENARIOS[STATE.scenario].label : "none";
document.getElementById("info-model").textContent = STATE.lastModel || "—";
document.getElementById("info-req-type").textContent = STATE.lastReqType || "—";
}
function renderGateStatus() {
const el = document.getElementById("gate-status");
if (STATE.msgCount === 0) { el.style.display = "none"; return; }
el.style.display = "";
if (STATE.msgCount < 4) {
el.className = "signal-gate waiting";
el.innerHTML = `⏳ <b>${4 - STATE.msgCount} more message${4 - STATE.msgCount > 1 ? "s" : ""} until learning kicks in.</b> The bandit updates after 4+ turns.`;
} else {
el.className = "signal-gate learning";
el.innerHTML = `✅ <b>Bandit is learning!</b> Each reply is now updating the router's model quality estimates.`;
}
}
function appendEmptyState() {
const msgs = document.getElementById("messages");
msgs.innerHTML = `<div class="empty-state">
<h2>Pick a scenario to start</h2>
<p>Choose one of the presets above or type your own message. The adaptive router picks the best model for each turn.</p>
</div>`;
}
function clearMessages() {
document.getElementById("messages").innerHTML = "";
}
function appendMessage(role, content, meta) {
const msgs = document.getElementById("messages");
const div = document.createElement("div");
div.className = `msg ${role}`;
div.dataset.role = role;
const avatar = document.createElement("div");
avatar.className = "avatar";
avatar.textContent = role === "user" ? "U" : "AI";
const wrap = document.createElement("div");
const bubble = document.createElement("div");
bubble.className = "bubble";
bubble.textContent = content; // plain text; code blocks show as preformatted
renderBubble(bubble, content);
wrap.appendChild(bubble);
if (meta) {
const metaEl = document.createElement("div");
metaEl.className = "msg-meta";
metaEl.innerHTML = meta;
wrap.appendChild(metaEl);
}
div.appendChild(avatar);
div.appendChild(wrap);
msgs.appendChild(div);
msgs.scrollTop = msgs.scrollHeight;
return bubble;
}
function renderBubble(el, text) {
// Minimal markdown: fenced code blocks and inline code.
const escaped = text
.replace(/&/g, "&amp;")
.replace(/</g, "&lt;")
.replace(/>/g, "&gt;");
const withBlocks = escaped.replace(
/```(\w*)\n?([\s\S]*?)```/g,
(_, lang, code) => `<pre><code>${code.trimEnd()}</code></pre>`
);
const withInline = withBlocks.replace(/`([^`]+)`/g, "<code>$1</code>");
el.innerHTML = withInline;
}
function appendThinking() {
const msgs = document.getElementById("messages");
const div = document.createElement("div");
div.className = "msg assistant";
div.id = "thinking-bubble";
const avatar = document.createElement("div");
avatar.className = "avatar";
avatar.textContent = "AI";
const bubble = document.createElement("div");
bubble.className = "bubble";
bubble.innerHTML = `<div class="thinking"><span></span><span></span><span></span></div>`;
div.appendChild(avatar);
div.appendChild(bubble);
msgs.appendChild(div);
msgs.scrollTop = msgs.scrollHeight;
return bubble;
}
function removeThinking() {
const el = document.getElementById("thinking-bubble");
if (el) el.remove();
}
// ---- chat send -------------------------------------------------------
async function sendMessage(text) {
if (STATE.sending || !text.trim()) return;
STATE.sending = true;
setSendEnabled(false);
// Add user message to history and UI
STATE.messages.push({ role: "user", content: text });
STATE.msgCount++;
appendMessage("user", text);
renderInfo();
renderGateStatus();
const thinkingBubble = appendThinking();
const t0 = Date.now();
try {
const body = {
model: STATE.router,
messages: STATE.messages,
stream: true,
metadata: { litellm_session_id: STATE.sessionId },
};
const resp = await fetch(`${STATE.proxyUrl}/v1/chat/completions`, {
method: "POST",
headers: {
"Authorization": `Bearer ${STATE.apiKey}`,
"Content-Type": "application/json",
},
body: JSON.stringify(body),
});
if (!resp.ok) {
const err = await resp.text().catch(() => `HTTP ${resp.status}`);
removeThinking();
appendMessage("assistant", `Error ${resp.status}: ${err}`);
STATE.sending = false;
setSendEnabled(true);
return;
}
// Read model from response header (requires proxy to expose via CORS).
const chosenModel = resp.headers.get("x-litellm-adaptive-router-model");
const reqType = resp.headers.get("x-litellm-request-type");
STATE.lastModel = chosenModel || "check dashboard";
STATE.lastReqType = reqType || "—";
// Stream the response.
removeThinking();
const assistantBubble = appendMessage("assistant", "");
let fullContent = "";
const reader = resp.body.getReader();
const decoder = new TextDecoder();
let buffer = "";
while (true) {
const { done, value } = await reader.read();
if (done) break;
buffer += decoder.decode(value, { stream: true });
// Process complete SSE lines.
const lines = buffer.split("\n");
buffer = lines.pop(); // last fragment may be incomplete
for (const line of lines) {
if (!line.startsWith("data: ")) continue;
const raw = line.slice(6).trim();
if (raw === "[DONE]") break;
try {
const chunk = JSON.parse(raw);
const delta = chunk.choices?.[0]?.delta?.content || "";
fullContent += delta;
renderBubble(assistantBubble, fullContent);
assistantBubble.closest(".messages")
? (assistantBubble.closest(".messages").scrollTop = assistantBubble.closest(".messages").scrollHeight)
: null;
document.getElementById("messages").scrollTop = document.getElementById("messages").scrollHeight;
} catch { /* incomplete JSON chunk, fine */ }
}
}
const latency = ((Date.now() - t0) / 1000).toFixed(2) + "s";
document.getElementById("info-latency").textContent = latency;
// Add assistant turn meta
const metaParts = [];
if (chosenModel) metaParts.push(`<span class="msg-model">${chosenModel}</span>`);
if (reqType) metaParts.push(`<span class="msg-type">${reqType}</span>`);
metaParts.push(latency);
if (metaParts.length) {
const metaEl = document.createElement("div");
metaEl.className = "msg-meta";
metaEl.innerHTML = metaParts.join(" · ");
assistantBubble.parentNode.appendChild(metaEl);
}
STATE.messages.push({ role: "assistant", content: fullContent });
STATE.msgCount++;
renderInfo();
renderGateStatus();
} catch (e) {
removeThinking();
appendMessage("assistant", `Request failed: ${e.message}`);
}
STATE.sending = false;
setSendEnabled(true);
}
// ---- input controls --------------------------------------------------
function setSendEnabled(enabled) {
const ta = document.getElementById("input");
const btn = document.getElementById("send-btn");
ta.disabled = !enabled || !STATE.connected;
btn.disabled = !enabled || !STATE.connected;
}
function setHint(text) {
document.getElementById("input-hint").textContent = text;
}
// ---- scenario selection ----------------------------------------------
function activateScenario(id) {
STATE.scenario = id;
document.querySelectorAll(".sc-btn[data-id]").forEach(b => {
b.classList.toggle("active", b.dataset.id === id);
});
newSession();
clearMessages();
const s = SCENARIOS[id];
if (s.system) STATE.messages.push({ role: "system", content: s.system });
const ta = document.getElementById("input");
ta.value = s.starter;
ta.style.height = "auto";
ta.style.height = Math.min(ta.scrollHeight, 160) + "px";
ta.focus();
renderInfo();
setHint(`Scenario: "${s.label}". Edit the starter if you like, then hit Send.`);
}
// ---- connect ---------------------------------------------------------
function connect() {
const url = document.getElementById("proxy-url").value.trim().replace(/\/$/, "");
const key = document.getElementById("api-key").value.trim();
const router = document.getElementById("router").value.trim();
if (!url || !key || !router) {
alert("Please fill in Proxy URL, API Key, and Router name.");
return;
}
STATE.proxyUrl = url;
STATE.apiKey = key;
STATE.router = router;
STATE.connected = true;
ssSet("ar_proxy_url", url);
ssSet("ar_api_key", key);
ssSet("ar_router", router);
setConn(true);
setSendEnabled(true);
setHint("Pick a scenario above or type your own message.");
newSession();
appendEmptyState();
renderInfo();
}
// ---- textarea auto-resize & keyboard submit --------------------------
document.getElementById("input").addEventListener("input", function () {
this.style.height = "auto";
this.style.height = Math.min(this.scrollHeight, 160) + "px";
});
document.getElementById("input").addEventListener("keydown", function (e) {
if (e.key === "Enter" && !e.shiftKey) {
e.preventDefault();
doSend();
}
});
function doSend() {
const ta = document.getElementById("input");
const text = ta.value.trim();
if (!text) return;
ta.value = "";
ta.style.height = "auto";
sendMessage(text);
}
// ---- wiring ----------------------------------------------------------
document.getElementById("connect-btn").addEventListener("click", connect);
document.getElementById("send-btn").addEventListener("click", doSend);
document.getElementById("new-chat-btn").addEventListener("click", () => {
STATE.scenario = null;
document.querySelectorAll(".sc-btn[data-id]").forEach(b => b.classList.remove("active"));
newSession();
clearMessages();
appendEmptyState();
const ta = document.getElementById("input");
ta.value = "";
ta.style.height = "auto";
setHint("Type anything — the router will classify it and pick the best model.");
renderInfo();
});
document.querySelectorAll(".sc-btn[data-id]").forEach(btn => {
btn.addEventListener("click", () => {
if (!STATE.connected) {
alert("Connect to the proxy first (fill in the form above and click Connect).");
return;
}
activateScenario(btn.dataset.id);
});
});
// ---- restore session storage -----------------------------------------
window.addEventListener("DOMContentLoaded", () => {
const u = ssGet("ar_proxy_url"); if (u) document.getElementById("proxy-url").value = u;
const k = ssGet("ar_api_key"); if (k) document.getElementById("api-key").value = k;
const r = ssGet("ar_router"); if (r) document.getElementById("router").value = r;
newSession();
renderInfo();
});
</script>
</body>
</html>

View file

@ -0,0 +1,635 @@
<!DOCTYPE html>
<html lang="en">
<head>
<meta charset="UTF-8" />
<title>Adaptive Router — Live</title>
<style>
:root {
--bg: #0b0f17;
--panel: #131a26;
--panel-2: #1b2433;
--fg: #e7ecf3;
--muted: #8a95a8;
--accent: #5dd6a4;
--accent-2: #6fb6ff;
--warn: #f6b94d;
--bad: #ff6b6b;
--bar-bg: #233047;
--border: #25324a;
}
* { box-sizing: border-box; }
body {
margin: 0;
font-family: -apple-system, BlinkMacSystemFont, "SF Pro Display",
"Segoe UI", Roboto, Inter, sans-serif;
background: var(--bg);
color: var(--fg);
font-size: 14px;
line-height: 1.45;
}
header {
padding: 18px 28px;
border-bottom: 1px solid var(--border);
display: flex;
align-items: center;
gap: 16px;
background: var(--panel);
}
header h1 {
margin: 0;
font-size: 18px;
font-weight: 600;
letter-spacing: 0.2px;
}
header .dot {
width: 8px; height: 8px; border-radius: 50%;
background: var(--bad); display: inline-block; margin-right: 6px;
}
header .dot.ok { background: var(--accent); }
header .status { color: var(--muted); font-size: 12px; }
.connect {
padding: 16px 28px;
display: flex; gap: 12px; align-items: center;
background: var(--panel-2);
border-bottom: 1px solid var(--border);
flex-wrap: wrap;
}
.connect input, .connect select {
background: #0e1422;
color: var(--fg);
border: 1px solid var(--border);
padding: 7px 10px;
border-radius: 6px;
font-size: 13px;
font-family: inherit;
}
.connect input[type=text] { width: 240px; }
.connect input[type=password] { width: 200px; }
.connect input[type=number] { width: 70px; }
.connect button {
background: var(--accent-2);
color: #0b0f17;
border: none;
padding: 7px 14px;
border-radius: 6px;
font-weight: 600;
cursor: pointer;
}
.connect label { color: var(--muted); font-size: 12px; }
main {
display: grid;
grid-template-columns: 1fr 360px;
gap: 20px;
padding: 20px 28px 40px;
max-width: 1400px;
}
@media (max-width: 1000px) { main { grid-template-columns: 1fr; } }
.panel {
background: var(--panel);
border: 1px solid var(--border);
border-radius: 10px;
padding: 18px 20px;
}
.panel h2 {
margin: 0 0 14px;
font-size: 13px;
text-transform: uppercase;
letter-spacing: 0.6px;
color: var(--muted);
font-weight: 600;
}
.rt-group {
margin-bottom: 18px;
padding-bottom: 14px;
border-bottom: 1px dashed var(--border);
}
.rt-group:last-child { border-bottom: none; margin-bottom: 0; }
.rt-title {
font-weight: 600;
font-size: 13px;
color: var(--fg);
margin-bottom: 8px;
display: flex; justify-content: space-between;
}
.rt-title .meta { color: var(--muted); font-weight: 400; font-size: 12px; }
.row {
display: grid;
grid-template-columns: 80px 1fr 130px;
align-items: center;
gap: 10px;
margin-bottom: 4px;
font-size: 13px;
}
.row .name { color: var(--muted); font-family: ui-monospace, monospace; }
.row .name.lead { color: var(--accent); font-weight: 600; }
.row .num { color: var(--muted); font-family: ui-monospace, monospace;
text-align: right; font-size: 12px; }
.bar { height: 14px; background: var(--bar-bg); border-radius: 4px;
position: relative; overflow: hidden; }
.bar > .fill {
height: 100%;
background: linear-gradient(90deg, var(--accent), var(--accent-2));
border-radius: 4px;
transition: width 0.4s ease;
}
.bar > .conf {
position: absolute; top: 0; bottom: 0; width: 1px;
background: rgba(255,255,255,0.4);
}
.bar > .conf.lo { background: rgba(255,255,255,0.5); }
.bar > .conf.hi { background: rgba(255,255,255,0.5); }
.pick-pct {
margin-top: 4px;
font-size: 11px;
color: var(--muted);
padding-left: 90px;
}
.cost-grid { display: grid; grid-template-columns: 1fr 1fr; gap: 12px; }
.cost-card {
background: var(--panel-2);
border-radius: 8px;
padding: 14px 16px;
}
.cost-card .label { color: var(--muted); font-size: 11px;
text-transform: uppercase; letter-spacing: 0.5px; }
.cost-card .value { font-size: 22px; font-weight: 600; margin-top: 4px;
font-family: ui-monospace, monospace; }
.cost-card .value.big { font-size: 28px; }
.cost-card .sub { color: var(--muted); font-size: 12px; margin-top: 2px; }
.cost-card.good { border: 1px solid rgba(93, 214, 164, 0.35); }
.cost-card.warn { border: 1px solid rgba(246, 185, 77, 0.35); }
.cost-card.bad { border: 1px solid rgba(255, 107, 107, 0.35); }
.savings {
margin-top: 12px;
padding: 10px 14px;
background: rgba(93, 214, 164, 0.08);
border: 1px solid rgba(93, 214, 164, 0.3);
border-radius: 8px;
color: var(--accent);
font-weight: 600;
font-size: 14px;
}
.savings.warn {
background: rgba(246, 185, 77, 0.08);
border-color: rgba(246, 185, 77, 0.3);
color: var(--warn);
}
.savings.bad {
background: rgba(255, 107, 107, 0.08);
border-color: rgba(255, 107, 107, 0.3);
color: var(--bad);
}
.savings .verdict-sub {
display: block;
color: var(--muted);
font-weight: 400;
font-size: 12px;
margin-top: 4px;
}
.panel-explainer {
margin: -8px 0 14px;
color: var(--muted);
font-size: 12px;
line-height: 1.5;
padding: 10px 12px;
background: var(--panel-2);
border-radius: 6px;
border-left: 3px solid var(--accent-2);
}
.panel-explainer b { color: var(--fg); font-weight: 600; }
.activity {
max-height: 360px; overflow-y: auto;
font-family: ui-monospace, monospace;
font-size: 12px;
}
.activity-row {
padding: 6px 8px;
border-bottom: 1px solid var(--border);
color: var(--muted);
display: grid;
grid-template-columns: 70px 1fr;
gap: 8px;
}
.activity-row .ts { color: #4f5d77; }
.activity-row .alpha { color: var(--accent); }
.activity-row .beta { color: var(--bad); }
.queue {
display: flex; gap: 16px; flex-wrap: wrap;
font-size: 12px; color: var(--muted);
margin-top: 8px;
}
.queue span b { color: var(--fg); font-weight: 600; }
.empty {
color: var(--muted); font-style: italic;
text-align: center; padding: 20px;
}
.pill {
background: var(--panel-2);
color: var(--muted);
padding: 3px 8px;
border-radius: 12px;
font-size: 11px;
}
</style>
</head>
<body>
<header>
<h1>⚡ Adaptive Router — Live</h1>
<span class="status"><span id="conn-dot" class="dot"></span><span id="conn-label">Disconnected</span></span>
<span id="poll-info" class="status"></span>
</header>
<div class="connect">
<label>Proxy URL <input id="proxy-url" type="text" value="http://localhost:4000" /></label>
<label>Master Key <input id="api-key" type="password" placeholder="sk-1234" /></label>
<label>Avg tokens/req <input id="avg-tokens" type="number" value="500" min="1" /></label>
<label>Poll ms <input id="poll-ms" type="number" value="500" min="100" /></label>
<button id="connect-btn">Connect</button>
<select id="router-select" style="display:none;"></select>
</div>
<main>
<section class="panel" id="bandit-panel">
<h2>How well each model performs, by request type</h2>
<div class="panel-explainer">
Each bar shows the <b>fraction of recent feedback that was positive</b>
for that model on that kind of request. Wider = better. The number
next to it (<b>"N signals"</b>) is how much real feedback the bar is
based on — more signals means the router is more confident.
It picks higher-quality bars first, with cost as a tiebreaker.
</div>
<div id="cells" class="empty">Connect to see live bandit state.</div>
</section>
<aside style="display:flex; flex-direction:column; gap:20px;">
<section class="panel">
<h2>Are the savings worth it?</h2>
<div class="panel-explainer">
<b>Cost saved</b> is what you spent vs. always picking the most
expensive model. <b>Quality kept</b> is the average quality of
the model that was actually picked, divided by the average
quality of the best-known model for each request type.
<i>If quality kept stays high while cost saved is high, you're
winning. If quality drops fast, you're saving money but making
users mad.</i>
</div>
<div class="cost-grid">
<div class="cost-card good">
<div class="label">💰 Cost saved</div>
<div class="value big" id="metric-cost-pct"></div>
<div class="sub" id="metric-cost-sub">no traffic yet</div>
</div>
<div class="cost-card good">
<div class="label">⭐ Quality kept</div>
<div class="value big" id="metric-quality-pct"></div>
<div class="sub" id="metric-quality-sub">vs best-known model</div>
</div>
</div>
<div class="savings" id="verdict">Send some traffic to see how the router is balancing cost and quality.</div>
<div class="queue" id="queue-info"></div>
</section>
<section class="panel">
<h2>Activity (last 30)</h2>
<div id="activity" class="activity empty">Waiting for signals…</div>
</section>
</aside>
</main>
<script>
// ---- helpers ---------------------------------------------------------
const REQ_TYPE_ORDER = [
"code_generation",
"code_understanding",
"technical_design",
"analytical_reasoning",
"writing",
"factual_lookup",
"general",
];
function fmtPct(n) { return (n * 100).toFixed(0) + "%"; }
function fmtUSD(n) { return "$" + n.toFixed(4); }
function nowTS() {
const d = new Date();
return d.toTimeString().slice(0, 8);
}
function ssGet(k) { try { return sessionStorage.getItem(k) || ""; } catch { return ""; } }
function ssSet(k, v) { try { sessionStorage.setItem(k, v); } catch {} }
// ---- state -----------------------------------------------------------
const STATE = {
proxyUrl: "",
apiKey: "",
pollMs: 500,
avgTokens: 500,
timer: null,
routers: [], // last snapshot list
selectedRouter: null, // name
prevCells: new Map(), // (router, rt, model) -> {alpha,beta,samples}
costAdaptive: 0,
costBaseline: 0,
totalRequests: 0,
activity: [], // [{ts, msg, kind}]
};
// ---- rendering -------------------------------------------------------
function renderRouters(routers) {
const sel = document.getElementById("router-select");
if (routers.length <= 1) {
sel.style.display = "none";
} else {
sel.style.display = "";
if (sel.options.length !== routers.length) {
sel.innerHTML = "";
for (const r of routers) {
const opt = document.createElement("option");
opt.value = r.router_name; opt.textContent = r.router_name;
sel.appendChild(opt);
}
sel.value = STATE.selectedRouter || routers[0].router_name;
}
}
if (!STATE.selectedRouter) STATE.selectedRouter = routers[0].router_name;
}
function pickShare(rows) {
// Approximate prob each model wins a Thompson-sample draw against the others.
// Simple proxy: softmax over quality_mean with temperature=0.05.
if (rows.length === 0) return {};
const T = 0.05;
const expv = rows.map(r => Math.exp(r.quality_mean / T));
const sum = expv.reduce((a, b) => a + b, 0);
const out = {};
rows.forEach((r, i) => out[r.model] = expv[i] / sum);
return out;
}
function renderCells(router) {
const container = document.getElementById("cells");
container.classList.remove("empty");
const byType = new Map();
for (const c of router.cells) {
if (!byType.has(c.request_type)) byType.set(c.request_type, []);
byType.get(c.request_type).push(c);
}
const order = REQ_TYPE_ORDER.filter(t => byType.has(t));
for (const t of byType.keys()) if (!order.includes(t)) order.push(t);
let html = "";
for (const rt of order) {
const rows = byType.get(rt).sort((a, b) => b.quality_mean - a.quality_mean);
const shares = pickShare(rows);
const lead = rows[0];
html += `<div class="rt-group">`;
html += `<div class="rt-title"><span>${rt}</span>`;
html += `<span class="meta">${rows.reduce((s, r) => s + (r.samples - 10), 0)} learning signals</span>`;
html += `</div>`;
for (const r of rows) {
const pct = fmtPct(r.quality_mean);
const share = fmtPct(shares[r.model] || 0);
const observed = Math.max(0, r.samples - 10); // strip cold-start prior mass
const isLead = r.model === lead.model;
const sigLabel = observed === 0 ? "no signals yet" : `${observed.toFixed(0)} signals`;
// Tooltip exposes raw Beta(α,β) for power users.
const tip = `Beta(α=${r.alpha.toFixed(1)}, β=${r.beta.toFixed(1)}) — ` +
`started at α=5,β=5 (cold-start prior), so the bar reflects ` +
`${observed.toFixed(0)} real feedback signals so far.`;
html += `<div class="row" title="${tip}">`;
html += `<div class="name ${isLead ? 'lead' : ''}">${r.model}</div>`;
html += `<div class="bar"><div class="fill" style="width:${(r.quality_mean*100).toFixed(1)}%"></div></div>`;
html += `<div class="num">${pct} good · ${sigLabel}</div>`;
html += `</div>`;
html += `<div class="pick-pct">→ ${share} of picks (Thompson softmax estimate)</div>`;
}
html += `</div>`;
}
container.innerHTML = html;
}
function computeQualityKept(router) {
// For each request type: pick_count_per_cell × quality_mean_per_cell
// summed and divided by total picks gives "average quality delivered".
// Compare against best-cell quality per type (weighted by picks in that type).
const byType = new Map();
for (const c of router.cells) {
if (!byType.has(c.request_type)) byType.set(c.request_type, []);
byType.get(c.request_type).push(c);
}
let totalPicks = 0, weightedDelivered = 0, weightedBest = 0;
for (const cells of byType.values()) {
const bestQ = Math.max(...cells.map(c => c.quality_mean));
for (const c of cells) {
const picks = Math.max(0, c.samples - 10);
if (picks === 0) continue;
totalPicks += picks;
weightedDelivered += picks * c.quality_mean;
weightedBest += picks * bestQ;
}
}
if (totalPicks === 0 || weightedBest === 0) return null;
return {
delivered: weightedDelivered / totalPicks,
best: weightedBest / totalPicks,
keptPct: weightedDelivered / weightedBest,
totalPicks,
};
}
function renderTradeoff(router) {
const costEl = document.getElementById("metric-cost-pct");
const costSub = document.getElementById("metric-cost-sub");
const qualEl = document.getElementById("metric-quality-pct");
const qualSub = document.getElementById("metric-quality-sub");
const verdict = document.getElementById("verdict");
// ---- Cost side ---------------------------------------------------
let costSavedPct = null;
if (STATE.costBaseline > 0) {
costSavedPct = 1 - STATE.costAdaptive / STATE.costBaseline;
costEl.textContent = (costSavedPct * 100).toFixed(0) + "%";
costSub.textContent = `${fmtUSD(STATE.costAdaptive)} spent vs ${fmtUSD(STATE.costBaseline)} baseline`;
} else {
costEl.textContent = "—";
costSub.textContent = "no traffic yet";
}
// ---- Quality side ------------------------------------------------
const q = computeQualityKept(router);
if (q) {
qualEl.textContent = (q.keptPct * 100).toFixed(0) + "%";
qualSub.textContent =
`delivered ${(q.delivered*100).toFixed(0)}% vs best-known ${(q.best*100).toFixed(0)}%`;
} else {
qualEl.textContent = "—";
qualSub.textContent = "vs best-known model";
}
// ---- Color-code the cards ----------------------------------------
const costCard = costEl.closest(".cost-card");
const qualCard = qualEl.closest(".cost-card");
costCard.className = "cost-card " + (costSavedPct === null ? "good"
: costSavedPct >= 0.30 ? "good"
: costSavedPct >= 0.05 ? "warn" : "bad");
qualCard.className = "cost-card " + (!q ? "good"
: q.keptPct >= 0.90 ? "good"
: q.keptPct >= 0.75 ? "warn" : "bad");
// ---- Verdict line ------------------------------------------------
if (costSavedPct === null || !q) {
verdict.className = "savings";
verdict.textContent = "Send some traffic to see how the router is balancing cost and quality.";
return;
}
const savedTxt = costSavedPct >= 0
? `${(costSavedPct*100).toFixed(0)}% cheaper`
: `${((-costSavedPct)*100).toFixed(0)}% MORE expensive (still exploring)`;
const qualLost = (1 - q.keptPct) * 100;
let line, cls;
if (q.keptPct >= 0.95 && costSavedPct >= 0.30) {
cls = "savings"; line = `✅ Big win: ${savedTxt}, lost only ${qualLost.toFixed(0)}% quality.`;
} else if (q.keptPct >= 0.85 && costSavedPct >= 0.10) {
cls = "savings"; line = `✅ Good trade: ${savedTxt}, gave up ${qualLost.toFixed(0)}% quality.`;
} else if (q.keptPct >= 0.75) {
cls = "savings warn"; line = `⚠️ Mixed: ${savedTxt}, but ${qualLost.toFixed(0)}% quality lost. Consider raising the quality weight.`;
} else {
cls = "savings bad"; line = `❌ Saving money, hurting users: ${savedTxt} but ${qualLost.toFixed(0)}% quality lost. Raise quality weight in the router config.`;
}
verdict.className = cls;
verdict.innerHTML = line +
`<span class="verdict-sub">Based on ${q.totalPicks} feedback signals across ${STATE.totalRequests} routed requests.</span>`;
}
function renderQueue(router) {
const q = router.queue || {};
document.getElementById("queue-info").innerHTML =
`<span>state pending: <b>${q.state_pending ?? 0}</b></span>` +
`<span>session pending: <b>${q.session_pending ?? 0}</b></span>` +
`<span>sticky live: <b>${router.sticky_sessions_live ?? 0}</b></span>` +
`<span>weights: q=<b>${router.weights?.quality ?? "?"}</b> c=<b>${router.weights?.cost ?? "?"}</b></span>`;
}
function renderActivity() {
const el = document.getElementById("activity");
if (STATE.activity.length === 0) {
el.classList.add("empty");
el.textContent = "Waiting for signals…";
return;
}
el.classList.remove("empty");
el.innerHTML = STATE.activity.map(a =>
`<div class="activity-row"><span class="ts">${a.ts}</span><span>${a.msg}</span></div>`
).join("");
}
// ---- diff & cost accounting -----------------------------------------
function processDiff(router, costsByModel) {
const maxCost = Math.max(0, ...Object.values(costsByModel));
for (const c of router.cells) {
const key = `${router.router_name}|${c.request_type}|${c.model}`;
const prev = STATE.prevCells.get(key);
if (prev) {
const dA = c.alpha - prev.alpha;
const dB = c.beta - prev.beta;
const dPicks = (c.samples - 10) - (prev.samples - 10);
if (dA > 0.001 || dB > 0.001) {
const tag = dA > dB
? `<span class="alpha">+${dA.toFixed(0)} 👍</span>`
: `<span class="beta">+${dB.toFixed(0)} 👎</span>`;
const qNow = (c.alpha / (c.alpha + c.beta) * 100).toFixed(0);
STATE.activity.unshift({
ts: nowTS(),
msg: `${c.request_type} → <b>${c.model}</b> ${tag} (now ${qNow}% good)`,
});
STATE.activity = STATE.activity.slice(0, 30);
}
if (dPicks > 0) {
const cost = costsByModel[c.model] || 0;
STATE.costAdaptive += dPicks * cost * STATE.avgTokens;
STATE.costBaseline += dPicks * maxCost * STATE.avgTokens;
STATE.totalRequests += dPicks;
}
}
STATE.prevCells.set(key, {
alpha: c.alpha, beta: c.beta, samples: c.samples
});
}
}
// ---- polling ---------------------------------------------------------
async function pollOnce() {
try {
const r = await fetch(`${STATE.proxyUrl}/adaptive_router/state`, {
headers: { "Authorization": `Bearer ${STATE.apiKey}` },
});
if (!r.ok) {
setConn(false, `HTTP ${r.status}`);
return;
}
const data = await r.json();
setConn(true, `Polling every ${STATE.pollMs}ms`);
STATE.routers = data.routers || [];
renderRouters(STATE.routers);
const router = STATE.routers.find(r => r.router_name === STATE.selectedRouter)
|| STATE.routers[0];
if (!router) return;
processDiff(router, router.model_costs || {});
renderCells(router);
renderQueue(router);
renderTradeoff(router);
renderActivity();
} catch (e) {
setConn(false, e.message);
}
}
function setConn(ok, msg) {
document.getElementById("conn-dot").className = "dot" + (ok ? " ok" : "");
document.getElementById("conn-label").textContent = ok ? "Connected" : "Disconnected";
document.getElementById("poll-info").textContent = msg || "";
}
function startPolling() {
if (STATE.timer) clearInterval(STATE.timer);
pollOnce();
STATE.timer = setInterval(pollOnce, STATE.pollMs);
}
// ---- wiring ----------------------------------------------------------
document.getElementById("connect-btn").addEventListener("click", () => {
STATE.proxyUrl = document.getElementById("proxy-url").value.trim().replace(/\/$/, "");
STATE.apiKey = document.getElementById("api-key").value.trim();
STATE.pollMs = parseInt(document.getElementById("poll-ms").value, 10) || 500;
STATE.avgTokens = parseInt(document.getElementById("avg-tokens").value, 10) || 500;
ssSet("ar_proxy_url", STATE.proxyUrl);
ssSet("ar_api_key", STATE.apiKey);
startPolling();
});
document.getElementById("router-select").addEventListener("change", (e) => {
STATE.selectedRouter = e.target.value;
STATE.prevCells.clear();
});
window.addEventListener("DOMContentLoaded", () => {
const u = ssGet("ar_proxy_url"); if (u) document.getElementById("proxy-url").value = u;
const k = ssGet("ar_api_key"); if (k) document.getElementById("api-key").value = k;
});
</script>
</body>
</html>

View file

@ -0,0 +1,271 @@
# ruff: noqa: T201
"""
Adaptive router evaluator LLM-as-judge harness.
For each test case:
1. Sends the prompt to the adaptive router.
2. Reads which model was picked (x-litellm-adaptive-router-model header).
3. Asks the judge model whether the response meets the ideal criteria.
4. Prints PASS or FAIL with one line of reasoning.
Run:
uv run python scripts/adaptive_router_demo/eval.py \
--proxy-url http://localhost:4000 \
--api-key sk-1234 \
--router smart-cheap-router \
--judge-model smart
"""
from __future__ import annotations
import argparse
import asyncio
import sys
import uuid
from dataclasses import dataclass
from typing import Dict, List, Optional, Tuple
import httpx
# ---------------------------------------------------------------------------
# Test cases
# ---------------------------------------------------------------------------
@dataclass
class EvalCase:
category: str
prompt: str
ideal: str # criteria the judge checks the response against
EVAL_CASES: List[EvalCase] = [
# code_generation
EvalCase(
category="code_generation",
prompt="Write a Python function that flattens a nested list of arbitrary depth.",
ideal=(
"A Python function (def flatten(...)) that accepts a list which may "
"contain nested lists to arbitrary depth and returns a single flat list "
"with all elements in order. Must handle at least two levels of nesting."
),
),
EvalCase(
category="code_generation",
prompt="Write a Python decorator that retries a function up to 3 times on exception.",
ideal=(
"A Python decorator that wraps a callable, catches exceptions, and "
"retries the call up to 3 times before re-raising. Should use functools.wraps "
"or equivalent to preserve the wrapped function's metadata."
),
),
EvalCase(
category="code_generation",
prompt="Write a SQL query that returns the top 5 customers by total order value.",
ideal=(
"A valid SQL SELECT query that JOINs an orders or order_items table with a "
"customers table, groups by customer, sums order value, orders descending, "
"and limits to 5 rows."
),
),
# factual_lookup
EvalCase(
category="factual_lookup",
prompt="What is the capital of New Zealand?",
ideal="The answer must state Wellington as the capital of New Zealand.",
),
EvalCase(
category="factual_lookup",
prompt="In what year did World War II end?",
ideal="The answer must state 1945 as the year World War II ended.",
),
EvalCase(
category="factual_lookup",
prompt="What is the chemical symbol for gold?",
ideal="The answer must include 'Au' as the chemical symbol for gold.",
),
# writing
EvalCase(
category="writing",
prompt=(
"Write a short, polite email declining a meeting request because of "
"a scheduling conflict."
),
ideal=(
"A professional email that: (1) thanks the sender for the invitation, "
"(2) clearly declines, (3) mentions a scheduling conflict as the reason, "
"and (4) offers to reschedule or an alternative. Tone must be polite."
),
),
EvalCase(
category="writing",
prompt="Write a one-paragraph product description for noise-cancelling headphones.",
ideal=(
"A marketing paragraph for noise-cancelling headphones that mentions "
"noise cancellation as a feature, highlights at least one other benefit "
"(comfort, audio quality, battery life, or similar), and ends with a "
"persuasive call to action or closing statement."
),
),
]
# Matches the satisfaction regex in signals.py (_SATISFACTION_PATTERNS).
SATISFY_FOLLOWUP = "great, thanks!"
NEUTRAL_FOLLOWUP = "ok, noted"
FAB_ASSISTANT = "Got it. Working on that now."
JUDGE_SYSTEM = (
"You are a strict but fair evaluator. Your job is to decide whether a model "
"response meets the stated requirements. Reply with exactly two lines:\n"
"Line 1: PASS or FAIL\n"
"Line 2: One sentence of reasoning (≤ 25 words)."
)
def _judge_user(prompt: str, ideal: str, actual: str) -> str:
return (
f"Question sent to model:\n{prompt}\n\n"
f"Requirements the response must meet:\n{ideal}\n\n"
f"Actual model response:\n{actual}\n\n"
"Does the response meet the requirements? Reply PASS or FAIL."
)
# ---------------------------------------------------------------------------
# HTTP helpers
# ---------------------------------------------------------------------------
async def _chat(
client: httpx.AsyncClient,
proxy_url: str,
api_key: str,
model: str,
messages: List[Dict[str, str]],
session_id: Optional[str] = None,
) -> Tuple[str, str]:
"""
Returns (response_text, chosen_model_header).
chosen_model_header is empty for non-router calls.
"""
body: Dict = {"model": model, "messages": messages}
if session_id:
body["metadata"] = {"litellm_session_id": session_id}
resp = await client.post(
f"{proxy_url}/v1/chat/completions",
json=body,
headers={"Authorization": f"Bearer {api_key}"},
timeout=60.0,
)
resp.raise_for_status()
data = resp.json()
text = data["choices"][0]["message"]["content"]
chosen = resp.headers.get("x-litellm-adaptive-router-model", "")
return text, chosen
# ---------------------------------------------------------------------------
# Evaluation loop
# ---------------------------------------------------------------------------
async def evaluate(
proxy_url: str,
api_key: str,
router: str,
judge_model: str,
) -> None:
passed = 0
failed = 0
async with httpx.AsyncClient() as client:
for i, case in enumerate(EVAL_CASES, 1):
print(f"\n[{i}/{len(EVAL_CASES)}] category={case.category}")
print(f" prompt : {case.prompt[:80]}{'' if len(case.prompt) > 80 else ''}")
session_id = f"eval-{uuid.uuid4()}"
# Round 1: single-turn real request — get the actual LLM response to judge.
try:
response, chosen = await _chat(
client, proxy_url, api_key, router,
[{"role": "user", "content": case.prompt}],
session_id=session_id,
)
except Exception as exc: # noqa: BLE001
print(f" ERROR calling router: {exc}", file=sys.stderr)
failed += 1
continue
print(f" model : {chosen or router}")
print(f" response : {response[:120].replace(chr(10), ' ')}{'' if len(response) > 120 else ''}")
# Judge the real response.
judge_msgs = [
{"role": "system", "content": JUDGE_SYSTEM},
{"role": "user", "content": _judge_user(case.prompt, case.ideal, response)},
]
try:
verdict, _ = await _chat(
client, proxy_url, api_key, judge_model, judge_msgs,
)
except Exception as exc: # noqa: BLE001
print(f" ERROR calling judge: {exc}", file=sys.stderr)
failed += 1
continue
# Parse verdict — first non-empty line should be PASS or FAIL.
lines = [ln.strip() for ln in verdict.splitlines() if ln.strip()]
first = lines[0].upper() if lines else ""
reason = lines[1] if len(lines) > 1 else ""
is_pass = "PASS" in first
if is_pass:
passed += 1
print(f" verdict : \033[32mPASS\033[0m {reason}")
else:
failed += 1
print(f" verdict : \033[31mFAIL\033[0m {reason}")
# Round 2: 5-message conversation on the same session_id so the bandit fires.
# On PASS → satisfaction follow-up (+alpha). On FAIL → neutral (no signal).
follow_up = SATISFY_FOLLOWUP if is_pass else NEUTRAL_FOLLOWUP
bandit_msgs = [
{"role": "user", "content": case.prompt},
{"role": "assistant", "content": response},
{"role": "user", "content": "ok continue"},
{"role": "assistant", "content": FAB_ASSISTANT},
{"role": "user", "content": follow_up},
]
try:
await _chat(
client, proxy_url, api_key, router, bandit_msgs,
session_id=session_id,
)
except Exception as exc: # noqa: BLE001
print(f" WARNING: bandit update failed: {exc}", file=sys.stderr)
total = passed + failed
print(f"\n{'='*60}")
print(f"Results: {passed}/{total} passed ({failed} failed)")
if passed == total:
print("All test cases passed — the adaptive router is working well!")
elif passed >= total * 0.8:
print("Most test cases passed — minor issues to investigate.")
else:
print("Significant failures — check router config and model availability.")
print("=" * 60)
# ---------------------------------------------------------------------------
# Entry point
# ---------------------------------------------------------------------------
def main() -> None:
ap = argparse.ArgumentParser(description="Evaluate the adaptive router with LLM-as-judge.")
ap.add_argument("--proxy-url", default="http://localhost:4000")
ap.add_argument("--api-key", required=True, help="proxy API key")
ap.add_argument("--router", default="smart-cheap-router", help="adaptive router model name")
ap.add_argument("--judge-model", default="smart", help="model name for the judge (via proxy)")
args = ap.parse_args()
asyncio.run(evaluate(args.proxy_url, args.api_key, args.router, args.judge_model))
if __name__ == "__main__":
main()

View file

@ -0,0 +1,227 @@
"""
Synthetic traffic generator for the adaptive_router demo dashboard.
What it does:
- Sends labeled multi-turn chat requests to the proxy's adaptive router.
- For each turn, peeks at the `x-litellm-adaptive-router-model` response
header to learn which underlying model was picked.
- Draws a Bernoulli outcome from a hard-coded ORACLE table that says
"model M succeeds at request type T with probability p".
- Sends a final follow-up turn whose user message is engineered to
BOTH classify into the same RequestType AND match the
satisfaction regex on success (so the bandit's `(type, model)` cell
gets +alpha). On failure we send a neutral follow-up so no signal
fires over time, models the oracle favors accumulate alpha faster.
Why this shape:
- The post-call hook gates signal recording on len(messages) >= 4.
A single 5-message request passes the gate in one round-trip, which
keeps the demo cheap.
- Mock responses (`mock_response=...`) skip the real LLM call but still
flow through routing + post-call hooks, so no API keys / no spend.
Run:
uv run python scripts/adaptive_router_demo/traffic.py \\
--proxy-url http://localhost:4000 \\
--api-key sk-1234 \\
--router smart-cheap-router \\
--rounds 100 \\
--rate 0.5
Open `dashboard.html` in a browser alongside this and watch the bars move.
"""
from __future__ import annotations
import argparse
import asyncio
import random
import sys
import uuid
from typing import Dict, List, Tuple
import httpx
# ---- prompts (paired with the RequestType the classifier will assign) ----
# Each prompt is engineered to (a) classify into the listed type and (b) make
# sense as a user request. Keep prompts short to limit token cost.
PROMPTS: Dict[str, List[str]] = {
"code_generation": [
"Write a Python function that flattens a nested list",
"Create a TypeScript function that debounces another function",
"Build a Rust function that parses a CSV string",
"Generate a SQL function that returns running totals",
],
"factual_lookup": [
"What is the capital of New Zealand?",
"When was the Treaty of Westphalia signed?",
"Who is the current Secretary General of the UN?",
"Where is Mount Kilimanjaro located?",
],
"writing": [
"Write an email declining a meeting politely",
"Draft a paragraph introducing a product launch",
"Compose a short blog post about morning routines",
"Rewrite this sentence to be more concise: ...",
],
}
# Engineered satisfaction follow-ups — each one is designed to:
# (1) match the satisfaction regex (thanks/great/works/perfect/etc.), AND
# (2) re-classify into the SAME RequestType as the first prompt
# so that signals attribute to the right (type, model) bandit cell.
SATISFY: Dict[str, str] = {
"code_generation": "thanks, that works! now write me a python function that does the inverse",
"factual_lookup": "perfect, thanks! who is the current prime minister?",
"writing": "great, thanks! now write a follow-up email confirming attendance",
}
# Neutral follow-up — does not match any signal regex, does not move the bandit.
NEUTRAL_FOLLOWUP = "ok, noted"
# Oracle: P(success | request_type, model). Tunable.
# Defaults: smart dominates code/writing; both are fine for factual_lookup.
ORACLE: Dict[str, Dict[str, float]] = {
"code_generation": {"smart": 0.92, "fast": 0.35},
"factual_lookup": {"smart": 0.90, "fast": 0.85},
"writing": {"smart": 0.85, "fast": 0.55},
}
# Fabricated assistant turn — content doesn't matter for the hook, only the role.
FAB_ASSISTANT = "Got it. Working on that now."
def _build_messages(prompt: str, last_user: str) -> List[Dict[str, str]]:
"""5-message conversation that passes the SIGNAL_GATE_MIN_MESSAGES=4 gate."""
return [
{"role": "user", "content": prompt},
{"role": "assistant", "content": FAB_ASSISTANT},
{"role": "user", "content": "ok continue"},
{"role": "assistant", "content": FAB_ASSISTANT},
{"role": "user", "content": last_user},
]
async def _send(
client: httpx.AsyncClient,
proxy_url: str,
api_key: str,
router: str,
session_id: str,
messages: List[Dict[str, str]],
mock_response: str,
) -> Tuple[bool, str]:
"""Returns (ok, chosen_model)."""
body = {
"model": router,
"messages": messages,
"metadata": {"litellm_session_id": session_id},
"mock_response": mock_response,
}
try:
r = await client.post(
f"{proxy_url}/v1/chat/completions",
json=body,
headers={"Authorization": f"Bearer {api_key}"},
timeout=15.0,
)
r.raise_for_status()
except Exception as e: # noqa: BLE001
print(f" request failed: {e}", file=sys.stderr)
return False, ""
chosen = r.headers.get("x-litellm-adaptive-router-model", "")
return True, chosen
async def _drive_one_session(
client: httpx.AsyncClient,
proxy_url: str,
api_key: str,
router: str,
request_type: str,
prompt: str,
) -> str:
"""Run one labeled session. Returns the chosen model (for logging)."""
session_id = f"demo-{uuid.uuid4()}"
# Send the engineered 5-message conversation. The follow-up is chosen
# AFTER we observe what model the router would pick — but since the
# router is sticky-per-session, the model on this single round-trip
# IS the model we're crediting.
#
# Pre-decide success based on the oracle for whichever model gets picked.
# We can't know the pick before sending, so: send a neutral follow-up
# first to learn the pick, then send a second round with credit attached.
#
# Round 1: neutral follow-up → no signal fires, but we learn the pick.
ok, chosen = await _send(
client, proxy_url, api_key, router, session_id,
_build_messages(prompt, NEUTRAL_FOLLOWUP),
mock_response=FAB_ASSISTANT,
)
if not ok or not chosen:
return ""
# Decide outcome from oracle.
p = ORACLE.get(request_type, {}).get(chosen, 0.5)
success = random.random() < p
follow_up = SATISFY[request_type] if success else NEUTRAL_FOLLOWUP
# Round 2: include the round-1 turns + a new follow-up. On success the
# follow-up matches satisfaction → +alpha for (request_type, chosen).
history = _build_messages(prompt, NEUTRAL_FOLLOWUP) + [
{"role": "assistant", "content": FAB_ASSISTANT},
{"role": "user", "content": follow_up},
]
await _send(
client, proxy_url, api_key, router, session_id, history,
mock_response=FAB_ASSISTANT,
)
return chosen
async def main() -> None:
ap = argparse.ArgumentParser()
ap.add_argument("--proxy-url", default="http://localhost:4000")
ap.add_argument("--api-key", required=True, help="proxy key with /v1/chat/completions perms")
ap.add_argument("--router", default="smart-cheap-router")
ap.add_argument("--rounds", type=int, default=100)
ap.add_argument("--rate", type=float, default=0.5,
help="seconds between sessions; lower = faster")
ap.add_argument("--types", default="code_generation,factual_lookup,writing",
help="comma-separated subset of request types to drive")
args = ap.parse_args()
types = [t.strip() for t in args.types.split(",") if t.strip() in PROMPTS]
if not types:
print(f"ERROR: no valid types. Choose from: {list(PROMPTS)}", file=sys.stderr)
sys.exit(2)
print(f"driving {args.rounds} sessions across types: {types}")
print(f"oracle: {ORACLE}")
print(f"proxy: {args.proxy_url} router: {args.router}\n")
counts: Dict[Tuple[str, str], int] = {}
async with httpx.AsyncClient() as client:
for i in range(args.rounds):
rt = random.choice(types)
prompt = random.choice(PROMPTS[rt])
chosen = await _drive_one_session(
client, args.proxy_url, args.api_key, args.router, rt, prompt,
)
if chosen:
counts[(rt, chosen)] = counts.get((rt, chosen), 0) + 1
if (i + 1) % 10 == 0:
summary = ", ".join(
f"{rt}/{m}={n}" for (rt, m), n in sorted(counts.items())
)
print(f" round {i + 1}/{args.rounds} picks: {summary}")
await asyncio.sleep(args.rate)
print("\nfinal pick distribution:")
for (rt, m), n in sorted(counts.items()):
print(f" {rt:22s}{m:8s} {n}")
if __name__ == "__main__":
asyncio.run(main())

View file

@ -0,0 +1,216 @@
"""
End-to-end verification script for the adaptive router.
Requires:
- LiteLLM proxy running on http://localhost:4000 with adaptive_router configured
(see litellm/proxy/example_config_yaml/adaptive_router_example.yaml).
- Postgres reachable via DATABASE_URL (same one the proxy uses).
- LITELLM_PROXY_KEY env var set (a valid key with permission to send requests).
- Two model deployments configured under one adaptive_router:
* "fast" (cheap, lower quality)
* "smart" (expensive, higher quality)
Run:
uv run python scripts/verify_adaptive_router.py
Optional env:
LITELLM_PROXY_URL (default: http://localhost:4000)
ADAPTIVE_ROUTER_NAME (default: smart-cheap-router)
EXPECTED_WINNER (default: smart) -- model expected to dominate after training
TRAIN_SESSIONS (default: 20) -- training sessions in phase 1
CONVERGE_SESSIONS (default: 10) -- cold sessions in phase 2
WIN_THRESHOLD (default: 0.7) -- min share for EXPECTED_WINNER in phase 2
"""
from __future__ import annotations
import asyncio
import os
import sys
import time
import uuid
from typing import List, Optional
import httpx
PROXY_URL: str = os.environ.get("LITELLM_PROXY_URL", "http://localhost:4000")
try:
PROXY_KEY: str = os.environ["LITELLM_PROXY_KEY"]
except KeyError:
print(
"ERROR: LITELLM_PROXY_KEY env var must be set (a proxy key with /chat/completions perms).",
file=sys.stderr,
)
sys.exit(2)
ROUTER_NAME: str = os.environ.get("ADAPTIVE_ROUTER_NAME", "smart-cheap-router")
EXPECTED_WINNER: str = os.environ.get("EXPECTED_WINNER", "smart")
TRAIN_SESSIONS: int = int(os.environ.get("TRAIN_SESSIONS", "20"))
CONVERGE_SESSIONS: int = int(os.environ.get("CONVERGE_SESSIONS", "10"))
WIN_THRESHOLD: float = float(os.environ.get("WIN_THRESHOLD", "0.7"))
REQUEST_TIMEOUT_SECONDS: float = 30.0
RETRY_ATTEMPTS: int = 3
RETRY_BACKOFF_SECONDS: float = 1.0
FLUSHER_DRAIN_WAIT_SECONDS: float = 30.0 # proxy flusher loop is 10s; pad with margin
PROMPTS: List[str] = [
"Write a Python function that reverses a binary tree",
"Explain the time complexity of quicksort",
"Design an API for a chat application",
]
SATISFACTION_PROMPT: str = "thanks, that worked!"
async def _post_chat(
client: httpx.AsyncClient, session_id: str, prompt: str
) -> Optional[dict]:
"""POST a chat completion with retry + timeout. Returns response JSON or None."""
body = {
"model": ROUTER_NAME,
"messages": [{"role": "user", "content": prompt}],
"metadata": {"litellm_session_id": session_id},
}
last_exc: Optional[Exception] = None
for attempt in range(1, RETRY_ATTEMPTS + 1):
try:
r = await client.post(
f"{PROXY_URL}/v1/chat/completions",
json=body,
headers={"Authorization": f"Bearer {PROXY_KEY}"},
timeout=REQUEST_TIMEOUT_SECONDS,
)
r.raise_for_status()
return r.json()
except Exception as e: # noqa: BLE001
last_exc = e
if attempt < RETRY_ATTEMPTS:
await asyncio.sleep(RETRY_BACKOFF_SECONDS * attempt)
print(
f" request failed after {RETRY_ATTEMPTS} attempts (session={session_id}): {last_exc}",
file=sys.stderr,
)
return None
async def send_session(
client: httpx.AsyncClient,
session_id: str,
prompts: List[str],
satisfy: bool = True,
) -> Optional[str]:
"""Send a session of N turns. Returns the model that handled the last turn."""
last_model: Optional[str] = None
for prompt in prompts:
resp = await _post_chat(client, session_id, prompt)
if resp is None:
return None
last_model = resp.get("model") or last_model
if satisfy:
await _post_chat(client, session_id, SATISFACTION_PROMPT)
return last_model
async def _proxy_health_check(client: httpx.AsyncClient) -> bool:
"""Confirm the proxy is reachable before doing anything else."""
try:
r = await client.get(f"{PROXY_URL}/health/liveliness", timeout=5.0)
return r.status_code == 200
except Exception as e: # noqa: BLE001
print(f"proxy unreachable at {PROXY_URL}: {e}", file=sys.stderr)
return False
async def main() -> None:
print("=== verify_adaptive_router.py ===")
print(f"proxy: {PROXY_URL}")
print(f"router: {ROUTER_NAME}")
print(f"expected winner: {EXPECTED_WINNER}")
print(f"train sessions: {TRAIN_SESSIONS}")
print(f"converge runs: {CONVERGE_SESSIONS}\n")
async with httpx.AsyncClient() as client:
if not await _proxy_health_check(client):
print("FAIL: proxy health check did not return 200.", file=sys.stderr)
sys.exit(1)
# ---- Phase 1: training -------------------------------------------
print(
f"Phase 1: training ({TRAIN_SESSIONS} sessions of 3 turns + satisfaction)..."
)
for i in range(TRAIN_SESSIONS):
sid = f"verify-train-{uuid.uuid4()}"
await send_session(client, sid, PROMPTS, satisfy=True)
if (i + 1) % 5 == 0:
print(f" trained {i + 1}/{TRAIN_SESSIONS} sessions")
print(
f"\nWaiting {FLUSHER_DRAIN_WAIT_SECONDS:.0f}s for flusher to drain queue..."
)
await asyncio.sleep(FLUSHER_DRAIN_WAIT_SECONDS)
# ---- Phase 2: convergence ----------------------------------------
print(f"\nPhase 2: convergence test ({CONVERGE_SESSIONS} cold sessions)...")
picks: List[str] = []
for i in range(CONVERGE_SESSIONS):
sid = f"verify-test-{uuid.uuid4()}"
m = await send_session(client, sid, [PROMPTS[0]], satisfy=False)
if m:
picks.append(m)
print(f" session {i + 1}: picked {m}")
if not picks:
print("\nFAIL: no successful picks in convergence phase.", file=sys.stderr)
sys.exit(1)
winner_share = picks.count(EXPECTED_WINNER) / len(picks)
print(
f"\n{EXPECTED_WINNER} share: {winner_share:.0%} "
f"({picks.count(EXPECTED_WINNER)}/{len(picks)})"
)
# ---- Phase 3: sticky session -------------------------------------
print("\nPhase 3: sticky session test...")
sid = f"verify-sticky-{uuid.uuid4()}"
models: List[str] = []
for _ in range(3):
m = await send_session(client, sid, [PROMPTS[0]], satisfy=False)
if m:
models.append(m)
if len(models) == 3 and len(set(models)) == 1:
print(f" PASS: same model {models[0]} across 3 turns of session {sid}")
else:
print(
f" FAIL: models differed within session: {models}",
file=sys.stderr,
)
sys.exit(1)
# ---- Phase 4: latency benchmark ----------------------------------
print("\nPhase 4: routing latency (5 picks, p50)...")
latencies: List[float] = []
for _ in range(5):
t0 = time.perf_counter()
await send_session(
client, f"verify-lat-{uuid.uuid4()}", [PROMPTS[0]], satisfy=False
)
latencies.append(time.perf_counter() - t0)
latencies.sort()
p50 = latencies[len(latencies) // 2]
print(f" p50 e2e roundtrip: {p50 * 1000:.0f}ms")
# ---- Verdict -----------------------------------------------------
if winner_share >= WIN_THRESHOLD:
print(
f"\nPASS: convergence ({winner_share:.0%} >= {WIN_THRESHOLD:.0%}) + "
f"sticky + latency checks all green."
)
sys.exit(0)
print(
f"\nFAIL: convergence too weak ({winner_share:.0%} < {WIN_THRESHOLD:.0%}).",
file=sys.stderr,
)
sys.exit(1)
if __name__ == "__main__":
asyncio.run(main())

View file

@ -1107,7 +1107,12 @@ async def test_convert_to_bedrock_format_post_call_streaming_hook():
# Mock the make_bedrock_api_request method to track calls
async def mock_make_bedrock_api_request(
source, messages=None, response=None, request_data=None
source,
messages=None,
response=None,
request_data=None,
logging_event_type=None,
**kwargs,
):
bedrock_calls.append(
{
@ -1115,6 +1120,7 @@ async def test_convert_to_bedrock_format_post_call_streaming_hook():
"messages": messages,
"response": response,
"request_data": request_data,
"logging_event_type": logging_event_type,
}
)
# Return the mock bedrock response

View file

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

View file

@ -100,8 +100,12 @@ import pytest
import requests
def test_litellm_proxy_server_config_no_general_settings():
# Sync the local litellm packages into the project environment
def _run_proxy_server_smoke_test(extra_proxy_args=None):
"""Sync deps, generate Prisma client, start proxy with optional extra args,
send a health check + chat/completions request, and tear down."""
if extra_proxy_args is None:
extra_proxy_args = []
server_process = None
try:
_run_uv(
@ -144,6 +148,7 @@ def test_litellm_proxy_server_config_no_general_settings():
"litellm.proxy.proxy_cli",
"--config",
config_fp,
*extra_proxy_args,
],
cwd=PROJECT_ROOT,
)
@ -182,3 +187,17 @@ def test_litellm_proxy_server_config_no_general_settings():
# Additional assertions can be added here
assert True
def test_litellm_proxy_server_config_no_general_settings():
"""Exercises the default (v1) migration resolver."""
_run_proxy_server_smoke_test()
def test_litellm_proxy_server_config_no_general_settings_v2_resolver():
"""Exercises the opt-in v2 migration resolver.
Runs in a separate CI job against a local Postgres to avoid collisions
with the v1 variant when they share a database.
"""
_run_proxy_server_smoke_test(extra_proxy_args=["--use_v2_migration_resolver"])

View file

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

View file

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

View file

@ -244,6 +244,7 @@ class TestRouterIndexManagement:
# Methods that are allowed to iterate through self.model_list
ALLOWED_METHODS = [
"_get_deployment_by_litellm_model", # Edge case: lookup by litellm_params.model (not indexed)
"_finalize_adaptive_router_if_configured", # Init-time prefix scan for "auto_router/adaptive_router" (no index for prefix match)
]
# Get path to router.py

View file

@ -1,10 +1,9 @@
import pytest
import asyncio
import aiohttp
import json
import time
from httpx import AsyncClient
from typing import Any, Optional
import litellm
from litellm._uuid import uuid
"""
@ -12,15 +11,13 @@ Tests to run
Basic Tests:
1. Basic Spend Accuracy Test:
- Make 1 calibration request, poll for spend to derive SPEND_PER_REQUEST
- Make N-1 more requests (N total)
- Expect the spend for each of the following to be N * SPEND_PER_REQUEST
Key, Team, User, Org (call /info endpoint for each object to validate)
- Make N requests, compute expected total spend locally from each response's usage
- Poll until batch writer has flushed spend to the DB
- Expect spend for Key, Team, User, Org (/info endpoints) to equal the computed total
2. Long term spend accuracy test (with 2 bursts of requests)
- Burst 1: Make requests, derive SPEND_PER_REQUEST from first request
- Burst 2: Make more requests
- Verify total spend = (burst1 + burst2) * SPEND_PER_REQUEST
- Burst 1: compute expected from responses, verify
- Burst 2: compute expected from responses, verify total = burst1 + burst2
Additional Test Scenarios:
@ -38,6 +35,34 @@ Additional Test Scenarios:
- Verify accurate total spend calculation
"""
# Upstream model the proxy is configured with (spend_tracking_config.yaml).
# The proxy computes spend using this model's pricing; the local ground-truth
# calculation uses the same pricing table via litellm.cost_per_token.
UPSTREAM_MODEL = "gpt-3.5-turbo"
# Batch writer flush cadence in CI is ~2-7s (PROXY_BATCH_WRITE_AT=2 + up to 5s jitter).
# Poll every 2s for 60s — plenty of headroom for multiple ticks to land.
POLL_INTERVAL_SECONDS = 2
POLL_TIMEOUT_SECONDS = 60
TOLERANCE = 1e-10
def _make_test_session() -> aiohttp.ClientSession:
"""
Session tuned for CI reliability:
- force_close: avoid aiohttp reusing a TCP connection that the proxy/kernel
silently closed during the long idle window between setup POSTs and the
later poll loop (observed failure mode: ConnectionTimeoutError on the
first /key/info call after 20 chat completions).
- explicit connect timeout: surface a blocked proxy event loop quickly
instead of hanging on aiohttp's 5-minute default total timeout.
"""
return aiohttp.ClientSession(
connector=aiohttp.TCPConnector(force_close=True),
timeout=aiohttp.ClientTimeout(total=30, connect=10),
)
async def create_organization(session, organization_alias: str):
"""Helper function to create a new organization"""
@ -102,54 +127,83 @@ async def get_spend_info(session, entity_type: str, entity_id: str):
return await response.json()
async def poll_key_spend_until_nonzero(
session, key: str, timeout: int = 120, interval: int = 10
):
"""Poll key spend until it becomes non-zero or timeout is reached."""
async def get_proxy_readiness(session):
"""Fetch /health/readiness. Used both as a fail-fast gate and as a diagnostic on poll timeout."""
url = "http://0.0.0.0:4000/health/readiness"
headers = {"Authorization": "Bearer sk-1234"}
async with session.get(url, headers=headers) as response:
return response.status, await response.json()
async def assert_proxy_healthy(session):
"""Fail fast if the proxy's DB or cache is not reachable — no point running the test."""
status, body = await get_proxy_readiness(session)
if status != 200 or body.get("db") != "connected":
pytest.fail(
f"Proxy /health/readiness unhealthy (status={status}). "
f"Cannot run spend accuracy test. Response: {body}"
)
print(f"Proxy readiness OK: {body}")
def compute_expected_spend(responses) -> float:
"""
Compute the expected total spend locally from each response's usage tokens,
using the same pricing table the proxy uses. This is the independent ground
truth we compare the proxy's reported spend against.
"""
total = 0.0
for r in responses:
usage = r.usage
prompt_cost, completion_cost = litellm.cost_per_token(
model=UPSTREAM_MODEL,
prompt_tokens=usage.prompt_tokens,
completion_tokens=usage.completion_tokens,
)
total += prompt_cost + completion_cost
return total
async def poll_key_spend_until(session, key: str, expected: float) -> float:
"""
Poll key spend until it matches `expected` within TOLERANCE, or timeout.
Returns the last observed spend either way; caller decides how to report.
"""
start = time.time()
while time.time() - start < timeout:
key_info = await get_spend_info(session, "key", key)
spend = key_info["info"]["spend"]
if spend > 0:
print(
f"Key spend became non-zero ({spend}) after {time.time() - start:.1f}s"
)
return spend
print(f"Key spend still 0.0, waiting... ({time.time() - start:.1f}s elapsed)")
await asyncio.sleep(interval)
raise TimeoutError(
f"Key spend remained 0.0 after {timeout}s — batch writer may not be running"
)
async def calibrate_spend_per_request(session, key: str, max_retries: int = 5):
"""
Make a single calibration request and poll for its spend to derive SPEND_PER_REQUEST.
Fails fast with pytest.fail() if spend cannot be determined.
"""
response = await chat_completion(session, key)
print(f"Calibration request completed: {response}")
for attempt in range(1, max_retries + 1):
last_spend = 0.0
while time.time() - start < POLL_TIMEOUT_SECONDS:
try:
spend = await poll_key_spend_until_nonzero(
session, key, timeout=120, interval=10
)
key_info = await get_spend_info(session, "key", key)
except (aiohttp.ClientError, asyncio.TimeoutError) as exc:
print(
f"Calibrated SPEND_PER_REQUEST = {spend} "
f"(attempt {attempt}/{max_retries})"
f"Transient transport error during spend poll: "
f"{type(exc).__name__}: {exc}. Retrying... "
f"({time.time() - start:.1f}s elapsed)"
)
return spend
except TimeoutError:
if attempt < max_retries:
print(
f"Calibration attempt {attempt}/{max_retries} timed out, retrying..."
)
else:
pytest.fail(
f"Failed to calibrate SPEND_PER_REQUEST after {max_retries} attempts. "
"The batch writer may not be running or the model may have 0 cost."
)
await asyncio.sleep(POLL_INTERVAL_SECONDS)
continue
last_spend = key_info["info"]["spend"]
if abs(last_spend - expected) < TOLERANCE:
print(
f"Key spend reached expected {expected} after {time.time() - start:.1f}s"
)
return last_spend
print(
f"Key spend {last_spend}, expected {expected}, waiting... "
f"({time.time() - start:.1f}s elapsed)"
)
await asyncio.sleep(POLL_INTERVAL_SECONDS)
return last_spend
async def fail_with_diagnostics(session, stage: str, expected: float, observed: float):
"""Emit a failure with readiness state so CI output points at the real cause."""
_, readiness = await get_proxy_readiness(session)
pytest.fail(
f"{stage}: key spend did not match expected after {POLL_TIMEOUT_SECONDS}s poll. "
f"expected={expected}, observed={observed}, diff={expected - observed}. "
f"Proxy readiness: {readiness}"
)
@pytest.mark.asyncio
@ -157,63 +211,60 @@ async def test_basic_spend_accuracy():
"""
Test basic spend accuracy across different entities:
1. Create org, team, user, and key
2. Make 1 calibration request to derive SPEND_PER_REQUEST
3. Make remaining requests (NUM_LLM_REQUESTS total)
4. Verify spend accuracy for key, team, user, and org
2. Make N requests, keeping each response
3. Compute expected spend locally from response usage (independent ground truth)
4. Poll until proxy-reported spend matches expected
5. Verify spend is consistent across key, team, user, and org entities
"""
NUM_LLM_REQUESTS = 20
TOLERANCE = 1e-10
async with aiohttp.ClientSession() as session:
# Create organization
async with _make_test_session() as session:
await assert_proxy_healthy(session)
org_response = await create_organization(
session=session, organization_alias=f"test-org-{uuid.uuid4()}"
)
print("org_response: ", org_response)
org_id = org_response["organization_id"]
# Create team under organization
team_response = await create_team(session, org_id)
print("team_response: ", team_response)
team_id = team_response["team_id"]
# Create user
user_response = await create_user(session, org_id)
print("user_response: ", user_response)
user_id = user_response["user_id"]
# Generate key
key_response = await generate_key(session, user_id, team_id)
print("key_response: ", key_response)
key = key_response["key"]
# Calibrate: make 1 request and derive SPEND_PER_REQUEST
spend_per_request = await calibrate_spend_per_request(session, key)
expected_spend = NUM_LLM_REQUESTS * spend_per_request
print(f"SPEND_PER_REQUEST={spend_per_request}, expected_spend={expected_spend}")
# Make remaining requests (1 already made during calibration)
for i in range(NUM_LLM_REQUESTS - 1):
responses = []
for i in range(NUM_LLM_REQUESTS):
response = await chat_completion(session, key)
print(f"Request {i + 2}/{NUM_LLM_REQUESTS} completed")
responses.append(response)
print(f"Request {i + 1}/{NUM_LLM_REQUESTS} completed")
# Poll until batch writer has flushed all spend
start = time.time()
while time.time() - start < 120:
key_info = await get_spend_info(session, "key", key)
current_spend = key_info["info"]["spend"]
if abs(current_spend - expected_spend) < TOLERANCE:
print(
f"Key spend reached expected {expected_spend} after {time.time() - start:.1f}s"
)
break
print(f"Key spend {current_spend}, expected {expected_spend}, waiting...")
await asyncio.sleep(10)
expected_spend = compute_expected_spend(responses)
assert expected_spend > 0, (
f"Locally computed expected spend is {expected_spend}. Either cost calc "
f"is broken or upstream returned zero tokens. "
f"Usage: {[r.usage.model_dump() for r in responses]}"
)
print(f"Expected total spend (local ground truth): {expected_spend}")
# Allow extra time for all entity spend aggregations to complete
final_spend = await poll_key_spend_until(session, key, expected_spend)
if abs(final_spend - expected_spend) >= TOLERANCE:
await fail_with_diagnostics(
session,
stage="test_basic_spend_accuracy",
expected=expected_spend,
observed=final_spend,
)
# Allow a final scheduler tick for team/user/org aggregations to settle
await asyncio.sleep(5)
# Get spend information for each entity
key_info = await get_spend_info(session, "key", key)
print("key_info: ", key_info)
team_info = await get_spend_info(session, "team", team_id)
@ -223,7 +274,6 @@ async def test_basic_spend_accuracy():
org_info = await get_spend_info(session, "organization", org_id)
print("org_info: ", org_info)
# Verify spend for each entity
assert (
abs(key_info["info"]["spend"] - expected_spend) < TOLERANCE
), f"Key spend {key_info['info']['spend']} does not match expected {expected_spend}"
@ -246,91 +296,78 @@ async def test_long_term_spend_accuracy_with_bursts():
"""
Test long-term spend accuracy with multiple bursts of requests:
1. Create org, team, user, and key
2. Calibrate SPEND_PER_REQUEST from first request
3. Burst 1: Make remaining requests
4. Burst 2: Make more requests
5. Verify the total spend is tracked accurately across all entities
2. Burst 1: make requests, compute expected locally, verify proxy matches
3. Burst 2: make more requests, verify proxy total == burst1 + burst2
4. Verify total spend is consistent across all entities
"""
BURST_1_REQUESTS = 22
BURST_2_REQUESTS = 12
TOTAL_REQUESTS = BURST_1_REQUESTS + BURST_2_REQUESTS
TOLERANCE = 1e-10
async with aiohttp.ClientSession() as session:
# Create organization
async with _make_test_session() as session:
await assert_proxy_healthy(session)
org_response = await create_organization(
session=session, organization_alias=f"test-org-{uuid.uuid4()}"
)
print("org_response: ", org_response)
org_id = org_response["organization_id"]
# Create team under organization
team_response = await create_team(session, org_id)
print("team_response: ", team_response)
team_id = team_response["team_id"]
# Create user
user_response = await create_user(session, org_id)
print("user_response: ", user_response)
user_id = user_response["user_id"]
# Generate key
key_response = await generate_key(session, user_id, team_id)
print("key_response: ", key_response)
key = key_response["key"]
# Calibrate: make 1 request and derive SPEND_PER_REQUEST
spend_per_request = await calibrate_spend_per_request(session, key)
expected_spend = TOTAL_REQUESTS * spend_per_request
print(f"SPEND_PER_REQUEST={spend_per_request}, expected_spend={expected_spend}")
# First burst: remaining requests (1 already made during calibration)
print(f"Starting first burst ({BURST_1_REQUESTS - 1} remaining requests)...")
for i in range(BURST_1_REQUESTS - 1):
print(f"Starting first burst of {BURST_1_REQUESTS} requests...")
burst_1_responses = []
for i in range(BURST_1_REQUESTS):
response = await chat_completion(session, key)
print(f"Burst 1 - Request {i + 2}/{BURST_1_REQUESTS} completed")
burst_1_responses.append(response)
print(f"Burst 1 - Request {i + 1}/{BURST_1_REQUESTS} completed")
# Poll until batch writer has flushed burst 1 spend
burst_1_expected = BURST_1_REQUESTS * spend_per_request
start = time.time()
while time.time() - start < 120:
key_info_check = await get_spend_info(session, "key", key)
current_spend = key_info_check["info"]["spend"]
if abs(current_spend - burst_1_expected) < TOLERANCE:
print(
f"Burst 1 spend reached expected {burst_1_expected} after {time.time() - start:.1f}s"
)
break
print(f"Key spend {current_spend}, expected {burst_1_expected}, waiting...")
await asyncio.sleep(10)
burst_1_expected = compute_expected_spend(burst_1_responses)
assert burst_1_expected > 0, (
f"Burst 1 expected spend is {burst_1_expected}. "
f"Usage: {[r.usage.model_dump() for r in burst_1_responses]}"
)
print(f"Burst 1 expected spend: {burst_1_expected}")
# Check intermediate spend
intermediate_key_info = await get_spend_info(session, "key", key)
print(f"After Burst 1 - Key spend: {intermediate_key_info['info']['spend']}")
final_burst_1 = await poll_key_spend_until(session, key, burst_1_expected)
if abs(final_burst_1 - burst_1_expected) >= TOLERANCE:
await fail_with_diagnostics(
session,
stage="test_long_term_spend_accuracy burst 1",
expected=burst_1_expected,
observed=final_burst_1,
)
# Second burst
print(f"Starting second burst of {BURST_2_REQUESTS} requests...")
burst_2_responses = []
for i in range(BURST_2_REQUESTS):
response = await chat_completion(session, key)
burst_2_responses.append(response)
print(f"Burst 2 - Request {i + 1}/{BURST_2_REQUESTS} completed")
# Poll until key spend reaches expected total (burst 1 + burst 2)
start = time.time()
while time.time() - start < 120:
key_info_check = await get_spend_info(session, "key", key)
current_spend = key_info_check["info"]["spend"]
if abs(current_spend - expected_spend) < TOLERANCE:
print(
f"Total spend reached expected {expected_spend} after {time.time() - start:.1f}s"
)
break
print(f"Key spend {current_spend}, expected {expected_spend}, waiting...")
await asyncio.sleep(10)
total_expected = burst_1_expected + compute_expected_spend(burst_2_responses)
print(f"Total expected spend (burst 1 + burst 2): {total_expected}")
final_total = await poll_key_spend_until(session, key, total_expected)
if abs(final_total - total_expected) >= TOLERANCE:
await fail_with_diagnostics(
session,
stage="test_long_term_spend_accuracy total",
expected=total_expected,
observed=final_total,
)
# Allow extra time for all entity spend aggregations
await asyncio.sleep(5)
# Get final spend information for each entity
key_info = await get_spend_info(session, "key", key)
team_info = await get_spend_info(session, "team", team_id)
user_info = await get_spend_info(session, "user", user_id)
@ -341,19 +378,18 @@ async def test_long_term_spend_accuracy_with_bursts():
print(f"Final user spend: {user_info['user_info']['spend']}")
print(f"Final org spend: {org_info['spend']}")
# Verify total spend for each entity
assert (
abs(key_info["info"]["spend"] - expected_spend) < TOLERANCE
), f"Key spend {key_info['info']['spend']} does not match expected {expected_spend}"
abs(key_info["info"]["spend"] - total_expected) < TOLERANCE
), f"Key spend {key_info['info']['spend']} does not match expected {total_expected}"
assert (
abs(user_info["user_info"]["spend"] - expected_spend) < TOLERANCE
), f"User spend {user_info['user_info']['spend']} does not match expected {expected_spend}"
abs(user_info["user_info"]["spend"] - total_expected) < TOLERANCE
), f"User spend {user_info['user_info']['spend']} does not match expected {total_expected}"
assert (
abs(team_info["team_info"]["spend"] - expected_spend) < TOLERANCE
), f"Team spend {team_info['team_info']['spend']} does not match expected {expected_spend}"
abs(team_info["team_info"]["spend"] - total_expected) < TOLERANCE
), f"Team spend {team_info['team_info']['spend']} does not match expected {total_expected}"
assert (
abs(org_info["spend"] - expected_spend) < TOLERANCE
), f"Organization spend {org_info['spend']} does not match expected {expected_spend}"
abs(org_info["spend"] - total_expected) < TOLERANCE
), f"Organization spend {org_info['spend']} does not match expected {total_expected}"

View file

@ -1,4 +1,4 @@
from typing import Any, Dict, List
from typing import Any, Dict, List, Optional
from unittest.mock import MagicMock, patch
import pytest
@ -26,7 +26,12 @@ class MockImageEditConfig(BaseImageEditConfig):
return "https://example.com/api"
def validate_environment(
self, headers: dict, model: str, api_key: str = None
self,
headers: dict,
model: str,
api_key: Optional[str] = None,
litellm_params: Optional[dict] = None,
api_base: Optional[str] = None,
) -> dict:
return headers
@ -262,3 +267,141 @@ class TestImageEditCustomPricing:
def test_custom_pricing_not_detected_without_model_info(self):
litellm_params = {"litellm_call_id": "test-call-id"}
assert use_custom_pricing_for_model(litellm_params) is False
class TestImageEditHandlerCredentialsForwarding:
"""
Regression tests for Vertex AI image_edit credentials bug.
image_edit handler must forward litellm_params to validate_environment,
so that credentials passed via YAML config (vertex_ai_project,
vertex_ai_credentials, etc.) reach the auth layer instead of falling
through to Application Default Credentials.
"""
def test_vertex_gemini_image_edit_reads_credentials_from_litellm_params(self):
"""
VertexAIGeminiImageEditConfig.validate_environment should read
vertex_ai_project/vertex_ai_credentials from litellm_params first.
"""
from litellm.llms.vertex_ai.image_edit.vertex_gemini_transformation import (
VertexAIGeminiImageEditConfig,
)
config = VertexAIGeminiImageEditConfig()
litellm_params = {
"vertex_ai_project": "test-project-from-params",
"vertex_ai_credentials": "/path/to/creds.json",
}
with patch.object(
config, "_ensure_access_token", return_value=("token", "project")
) as mock_ensure:
config.validate_environment(
headers={},
model="test-model",
litellm_params=litellm_params,
)
mock_ensure.assert_called_once()
call_kwargs = mock_ensure.call_args[1]
assert call_kwargs["credentials"] == "/path/to/creds.json"
assert call_kwargs["project_id"] == "test-project-from-params"
def test_vertex_imagen_image_edit_reads_credentials_from_litellm_params(self):
"""
VertexAIImagenImageEditConfig.validate_environment should read
vertex_ai_project/vertex_ai_credentials from litellm_params first.
"""
from litellm.llms.vertex_ai.image_edit.vertex_imagen_transformation import (
VertexAIImagenImageEditConfig,
)
config = VertexAIImagenImageEditConfig()
litellm_params = {
"vertex_ai_project": "test-project-from-params",
"vertex_ai_credentials": "/path/to/creds.json",
}
with patch.object(
config, "_ensure_access_token", return_value=("token", "project")
) as mock_ensure:
config.validate_environment(
headers={},
model="test-model",
litellm_params=litellm_params,
)
mock_ensure.assert_called_once()
call_kwargs = mock_ensure.call_args[1]
assert call_kwargs["credentials"] == "/path/to/creds.json"
assert call_kwargs["project_id"] == "test-project-from-params"
def test_vertex_imagen_get_complete_url_reads_project_and_location_from_litellm_params(
self,
):
"""
VertexAIImagenImageEditConfig.get_complete_url should read
vertex_ai_project and vertex_ai_location from litellm_params,
not only from env vars / global settings.
"""
from litellm.llms.vertex_ai.image_edit.vertex_imagen_transformation import (
VertexAIImagenImageEditConfig,
)
config = VertexAIImagenImageEditConfig()
litellm_params = {
"vertex_ai_project": "param-project",
"vertex_ai_location": "us-east1",
}
url = config.get_complete_url(
model="vertex_ai/imagegeneration@002",
api_base=None,
litellm_params=litellm_params,
)
assert "param-project" in url
assert "us-east1" in url
def test_validate_environment_signature_includes_litellm_params(self):
"""
All image_edit config validate_environment methods should accept
litellm_params to allow credentials to be forwarded from the handler.
"""
import inspect
from litellm.llms.vertex_ai.image_edit.vertex_gemini_transformation import (
VertexAIGeminiImageEditConfig,
)
from litellm.llms.vertex_ai.image_edit.vertex_imagen_transformation import (
VertexAIImagenImageEditConfig,
)
from litellm.llms.openai.image_edit.transformation import (
OpenAIImageEditConfig,
)
configs = [
VertexAIGeminiImageEditConfig(),
VertexAIImagenImageEditConfig(),
OpenAIImageEditConfig(),
MockImageEditConfig(),
]
for config in configs:
sig = inspect.signature(config.validate_environment)
params = list(sig.parameters.keys())
assert "litellm_params" in params, (
f"{config.__class__.__name__}.validate_environment "
"missing litellm_params parameter"
)
assert "api_base" in params, (
f"{config.__class__.__name__}.validate_environment "
"missing api_base parameter"
)

View file

@ -1055,3 +1055,50 @@ class TestTracingFieldsPopulation:
assert slg["classification"] == classification
assert slg["detection_method"] == "llm-judge"
assert slg["confidence_score"] == 0.94
class TestCustomGuardrailSpendLogMatchRedaction:
"""Guardrail JSON persisted via standard_logging must not contain raw match spans."""
def test_add_standard_logging_redacts_nested_match(self):
cg = CustomGuardrail(guardrail_name="test-rail")
raw = {
"assessments": [
{
"sensitiveInformationPolicy": {
"piiEntities": [
{"type": "NAME", "match": "GG", "action": "BLOCKED"}
]
}
}
]
}
request_data: dict = {"metadata": {}}
cg.add_standard_logging_guardrail_information_to_request_data(
guardrail_json_response=raw,
request_data=request_data,
guardrail_status="guardrail_intervened",
)
slg = request_data["metadata"]["standard_logging_guardrail_information"][0]
assert (
slg["guardrail_response"]["assessments"][0]["sensitiveInformationPolicy"][
"piiEntities"
][0]["match"]
== "[REDACTED]"
)
assert raw["assessments"][0]["sensitiveInformationPolicy"]["piiEntities"][0][
"match"
] == "GG"
def test_add_standard_logging_redacts_regex_field(self):
cg = CustomGuardrail(guardrail_name="test-rail")
raw = {"filters": [{"regex": r"\d{3}-\d{2}-\d{4}", "action": "BLOCKED"}]}
request_data: dict = {"metadata": {}}
cg.add_standard_logging_guardrail_information_to_request_data(
guardrail_json_response=raw,
request_data=request_data,
guardrail_status="success",
)
slg = request_data["metadata"]["standard_logging_guardrail_information"][0]
assert slg["guardrail_response"]["filters"][0]["regex"] == "[REDACTED]"
assert raw["filters"][0]["regex"] == r"\d{3}-\d{2}-\d{4}"

View file

@ -1047,6 +1047,36 @@ class TestOpenTelemetryEndpointNormalization(unittest.TestCase):
result = otel._normalize_otel_endpoint("http://collector:4318/", "traces")
self.assertEqual(result, "http://collector:4318/v1/traces")
@parameterized.expand(
[
(
"https://ingest.eu1.observability.splunkcloud.com/v2/trace/otlp",
"https://ingest.eu1.observability.splunkcloud.com/v2/trace/otlp",
),
(
"https://ingest.us0.observability.splunkcloud.com/v2/trace/otlp/",
"https://ingest.us0.observability.splunkcloud.com/v2/trace/otlp",
),
(
"https://ingest.eu0.signalfx.com/v2/trace/otlp",
"https://ingest.eu0.signalfx.com/v2/trace/otlp",
),
(
"https://example.com/prefix/v2/trace/otlp",
"https://example.com/prefix/v2/trace/otlp",
),
]
)
def test_normalize_traces_nonstandard_otlp_ingest_urls_unchanged(
self, input_url: str, expected: str
) -> None:
"""Splunk-style /v2/trace/otlp endpoints must not get /v1/traces appended."""
otel = OpenTelemetry()
self.assertEqual(
otel._normalize_otel_endpoint(input_url, "traces"),
expected,
)
def test_normalize_endpoint_none(self):
"""Test that None endpoint returns None"""
otel = OpenTelemetry()
@ -1315,7 +1345,7 @@ class TestOpenTelemetryProtocolSelection(unittest.TestCase):
@patch.dict(
os.environ,
{
"OTEL_EXPORTER": "otlp_http",
"OTEL_EXPORTER_OTLP_PROTOCOL": "http/protobuf",
"OTEL_EXPORTER_OTLP_ENDPOINT": "http://collector:4318",
},
clear=False,
@ -1339,7 +1369,7 @@ class TestOpenTelemetryProtocolSelection(unittest.TestCase):
@patch.dict(
os.environ,
{
"OTEL_EXPORTER": "otlp_grpc",
"OTEL_EXPORTER_OTLP_PROTOCOL": "grpc",
"OTEL_EXPORTER_OTLP_ENDPOINT": "http://collector:4317",
},
clear=False,
@ -1360,6 +1390,60 @@ class TestOpenTelemetryProtocolSelection(unittest.TestCase):
self.assertIsInstance(processor, BatchSpanProcessor)
self.assertIsInstance(processor.span_exporter, OTLPSpanExporterGRPC)
@patch.dict(
os.environ,
{
"OTEL_EXPORTER": "otlp_http",
"OTEL_EXPORTER_OTLP_ENDPOINT": "http://collector:4318",
},
clear=False,
)
def test_protocol_selection_from_otel_exporter_fallback_http(self):
"""OTEL_EXPORTER drives protocol when OTEL_EXPORTER_OTLP_PROTOCOL is unset."""
from opentelemetry.exporter.otlp.proto.http.trace_exporter import (
OTLPSpanExporter as OTLPSpanExporterHTTP,
)
from opentelemetry.sdk.trace.export import BatchSpanProcessor
popped_protocol = os.environ.pop("OTEL_EXPORTER_OTLP_PROTOCOL", None)
try:
config = OpenTelemetryConfig.from_env()
self.assertEqual(config.exporter, "otlp_http")
otel = OpenTelemetry(config=config)
processor = otel._get_span_processor()
self.assertIsInstance(processor, BatchSpanProcessor)
self.assertIsInstance(processor.span_exporter, OTLPSpanExporterHTTP)
finally:
if popped_protocol is not None:
os.environ["OTEL_EXPORTER_OTLP_PROTOCOL"] = popped_protocol
@patch.dict(
os.environ,
{
"OTEL_EXPORTER": "otlp_grpc",
"OTEL_EXPORTER_OTLP_ENDPOINT": "http://collector:4317",
},
clear=False,
)
def test_protocol_selection_from_otel_exporter_fallback_grpc(self):
"""OTEL_EXPORTER drives protocol when OTEL_EXPORTER_OTLP_PROTOCOL is unset."""
from opentelemetry.exporter.otlp.proto.grpc.trace_exporter import (
OTLPSpanExporter as OTLPSpanExporterGRPC,
)
from opentelemetry.sdk.trace.export import BatchSpanProcessor
popped_protocol = os.environ.pop("OTEL_EXPORTER_OTLP_PROTOCOL", None)
try:
config = OpenTelemetryConfig.from_env()
self.assertEqual(config.exporter, "otlp_grpc")
otel = OpenTelemetry(config=config)
processor = otel._get_span_processor()
self.assertIsInstance(processor, BatchSpanProcessor)
self.assertIsInstance(processor.span_exporter, OTLPSpanExporterGRPC)
finally:
if popped_protocol is not None:
os.environ["OTEL_EXPORTER_OTLP_PROTOCOL"] = popped_protocol
def test_http_exporter_endpoint_normalization_for_traces(self):
"""Test that HTTP trace exporter gets properly normalized endpoint"""
config = OpenTelemetryConfig(
@ -2752,3 +2836,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)

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