Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_fix_responses_queued_id_encryption

# Conflicts:
#	type-discipline-budget.json
This commit is contained in:
mateo-berri 2026-09-03 14:45:19 -07:00
commit 2df5f4a7c8
100 changed files with 4428 additions and 906 deletions

View file

@ -19,9 +19,6 @@ jobs:
build-ui:
runs-on: ubuntu-latest
timeout-minutes: 10
defaults:
run:
working-directory: ui/litellm-dashboard
steps:
- name: Checkout repository
@ -35,18 +32,11 @@ jobs:
with:
category: ui
- name: Setup Node.js
# Built through the image stage rather than the checkout, because the
# stage copies ui/litellm-dashboard/ alone: an import reaching above the
# dashboard root resolves in a checkout and fails in every image we ship.
# Dockerfile, docker/Dockerfile.non_root and ui/Dockerfile share this
# stage verbatim, so building one covers all three.
- name: Build the dashboard as the shipped images build it
if: steps.changes.outputs.decision != 'skip'
uses: actions/setup-node@a0853c24544627f65ddf259abe73b1d18a591444 # v5.0.0
with:
node-version-file: ui/litellm-dashboard/.nvmrc
cache: "npm"
cache-dependency-path: ui/litellm-dashboard/package-lock.json
- name: Install dependencies
if: steps.changes.outputs.decision != 'skip'
run: npm ci
- name: Build
if: steps.changes.outputs.decision != 'skip'
run: npm run build
run: docker build --target ui-builder -f Dockerfile .

View file

@ -74,12 +74,19 @@ jobs:
- name: Run Clippy with Bedrock auth
run: cargo clippy -p litellm-core --all-targets --features bedrock-auth --locked -- -D warnings
- name: Run Clippy with all gateway features
run: cargo clippy -p litellm-ai-gateway --all-targets --all-features --locked -- -D warnings
- name: Run Rust tests
run: cargo test --workspace --locked
- name: Run core tests with Bedrock auth
run: cargo test -p litellm-core --features bedrock-auth --locked
# Not --all-features: python-config links libpython, which this job does not install.
- name: Run gateway tests with the server feature
run: cargo test -p litellm-ai-gateway --features server --locked
release-wheel:
name: release wheel
runs-on: ubuntu-latest

View file

@ -151,6 +151,7 @@ jobs:
tests/test_litellm/proxy/google_endpoints
tests/test_litellm/proxy/openai_files_endpoint
tests/test_litellm/proxy/batches_endpoints
tests/test_litellm/proxy/container_endpoints
tests/test_litellm/proxy/fine_tuning_endpoints
tests/test_litellm/proxy/vector_store_files_endpoints
tests/test_litellm/proxy/video_endpoints

View file

@ -3,7 +3,7 @@
"limit": 14074
},
"reportArgumentType": {
"limit": 2215
"limit": 2214
},
"reportAssignmentType": {
"limit": 319
@ -42,7 +42,7 @@
"limit": 12
},
"reportIndexIssue": {
"limit": 25
"limit": 24
},
"reportInvalidTypeForm": {
"limit": 34

View file

@ -428,3 +428,16 @@ envFrom:
{{- end }}
{{- end }}
{{- end -}}
{{/*
ingress-nginx's admission webhook rejects a dot in an Exact or Prefix path
(strict-validate-path-type) and serves ImplementationSpecific as a plain
prefix location, so a dotted path takes that type there.
*/}}
{{- define "litellm.ingress.pathType" -}}
{{- if and (eq .controller "nginx") (contains "." .path) -}}
ImplementationSpecific
{{- else -}}
{{- .pathType -}}
{{- end -}}
{{- end -}}

View file

@ -5,6 +5,10 @@
{{- $gatewayPort := .Values.gateway.service.port -}}
{{- $backendPort := .Values.backend.service.port -}}
{{- $uiPort := .Values.ui.service.port -}}
{{- $controller := .Values.ingress.controller | default "alb" -}}
{{- if not (has $controller (list "alb" "nginx")) }}
{{- fail (printf "ingress.controller: unknown controller %q, expected one of alb, nginx" $controller) }}
{{- end }}
{{/*
Backends addressable from ingress.extraPaths, keyed by the `service` field.
*/}}
@ -27,10 +31,11 @@
/litellm-asset-prefix, so without /*.txt they fall to the backend catch-all
→ 404 → client-side navigation never settles and the login flow spins in an
infinite redirect loop (/ ⇄ /ui/login). ui/nginx.conf already serves *.txt
from the export; the rule only routes the request to it. Needs an ingress
controller whose ImplementationSpecific path is a wildcard pattern
(AWS ALB: `*` = 0+ chars); this chart targets the AWS Load Balancer
Controller.
from the export; the rule only routes the request to it. It needs an
ingress controller whose ImplementationSpecific path is a wildcard pattern
(AWS ALB: `*` = 0+ chars), so it is rendered for ingress.controller=alb
only: ingress-nginx serves ImplementationSpecific as a literal prefix
location, where /*.txt can never match.
*/}}
{{- $uiPaths := list
(dict "path" "/" "pathType" "Exact")
@ -38,8 +43,10 @@
(dict "path" "/litellm-asset-prefix" "pathType" "Prefix")
(dict "path" "/_next" "pathType" "Prefix")
(dict "path" "/ui" "pathType" "Prefix")
(dict "path" "/*.txt" "pathType" "ImplementationSpecific")
-}}
{{- if eq $controller "alb" }}
{{- $uiPaths = append $uiPaths (dict "path" "/*.txt" "pathType" "ImplementationSpecific") }}
{{- end }}
{{/*
Gateway data-plane prefixes — must mirror gateway/routes/allowlist.py.
Versioned paths are listed explicitly to avoid routing management routes
@ -83,12 +90,6 @@
adding to it.
*/}}
{{- $builtinPathKeys := list "/test|Exact" "/|Prefix" -}}
{{- range $uiPaths }}
{{- $builtinPathKeys = append $builtinPathKeys (printf "%s|%s" .path .pathType) }}
{{- end }}
{{- range $gatewayPrefixes }}
{{- $builtinPathKeys = append $builtinPathKeys (printf "%s|Prefix" .) }}
{{- end }}
apiVersion: networking.k8s.io/v1
kind: Ingress
metadata:
@ -115,8 +116,10 @@ spec:
paths:
# --- UI (Next.js static export) ---
{{- range $uiPaths }}
{{- $pathType := include "litellm.ingress.pathType" (dict "controller" $controller "path" .path "pathType" .pathType) }}
{{- $builtinPathKeys = append $builtinPathKeys (printf "%s|%s" .path $pathType) }}
- path: {{ .path }}
pathType: {{ .pathType }}
pathType: {{ $pathType }}
backend:
service:
name: {{ $uiName }}
@ -134,8 +137,10 @@ spec:
port:
number: {{ $gatewayPort }}
{{- range $gatewayPrefixes }}
{{- $pathType := include "litellm.ingress.pathType" (dict "controller" $controller "path" . "pathType" "Prefix") }}
{{- $builtinPathKeys = append $builtinPathKeys (printf "%s|%s" . $pathType) }}
- path: {{ . }}
pathType: Prefix
pathType: {{ $pathType }}
backend:
service:
name: {{ $gatewayName }}
@ -147,10 +152,11 @@ spec:
Rendered after every built-in path so an entry can never take
precedence over a default, and before the backend catch-all.
Position only decides the match on controllers that honour manifest
order: the AWS Load Balancer Controller this chart targets sorts
Exact paths first and Prefix paths longest-first, but keeps
order: the AWS Load Balancer Controller (ingress.controller=alb)
sorts Exact paths first and Prefix paths longest-first, but keeps
ImplementationSpecific paths in manifest order, which is what the
/*.txt rule above already depends on.
/*.txt rule above already depends on. ingress-nginx ignores order
and serves the longest matching location.
*/}}
{{- range $idx, $extra := .Values.ingress.extraPaths }}
{{- if not (kindIs "map" $extra) }}
@ -164,10 +170,11 @@ spec:
{{- if not $target }}
{{- fail (printf "ingress.extraPaths[%d] (path %s): unknown service %q, expected one of backend, gateway, ui" $idx $extra.path $service) }}
{{- end }}
{{- $pathType := $extra.pathType | default "Prefix" }}
{{- if not (has $pathType (list "Prefix" "Exact" "ImplementationSpecific")) }}
{{- fail (printf "ingress.extraPaths[%d] (path %s): unknown pathType %q, expected one of Exact, ImplementationSpecific, Prefix" $idx $extra.path $pathType) }}
{{- $requestedPathType := $extra.pathType | default "Prefix" }}
{{- if not (has $requestedPathType (list "Prefix" "Exact" "ImplementationSpecific")) }}
{{- fail (printf "ingress.extraPaths[%d] (path %s): unknown pathType %q, expected one of Exact, ImplementationSpecific, Prefix" $idx $extra.path $requestedPathType) }}
{{- end }}
{{- $pathType := include "litellm.ingress.pathType" (dict "controller" $controller "path" $extra.path "pathType" $requestedPathType) }}
{{- if eq $extra.path "/" }}
{{- fail (printf "ingress.extraPaths[%d]: path / is already routed in both directions, Exact to ui and Prefix to backend, so no pathType leaves a request for an entry here to capture" $idx) }}
{{- end }}

View file

@ -0,0 +1,205 @@
suite: test ingress.controller
templates:
- ingress.yaml
values:
- ./values/required.yaml
tests:
- it: keeps the AWS Load Balancer Controller path types by default
set:
ingress.enabled: true
asserts:
- contains:
path: spec.rules[0].http.paths
content:
path: /favicon.ico
pathType: Exact
backend:
service:
name: RELEASE-NAME-litellm-ui
port:
number: 3000
- contains:
path: spec.rules[0].http.paths
content:
path: /eu.assemblyai
pathType: Prefix
backend:
service:
name: RELEASE-NAME-litellm-gateway
port:
number: 4000
- contains:
path: spec.rules[0].http.paths
content:
path: /*.txt
pathType: ImplementationSpecific
backend:
service:
name: RELEASE-NAME-litellm-ui
port:
number: 3000
- it: renders no dotted Exact or Prefix path for ingress-nginx, whose admission webhook rejects them
set:
ingress.enabled: true
ingress.controller: nginx
asserts:
- notMatchRegexRaw:
pattern: 'path: /\S*\.\S*\n\s+pathType: (Exact|Prefix)\n'
- contains:
path: spec.rules[0].http.paths
content:
path: /favicon.ico
pathType: ImplementationSpecific
backend:
service:
name: RELEASE-NAME-litellm-ui
port:
number: 3000
- contains:
path: spec.rules[0].http.paths
content:
path: /eu.assemblyai
pathType: ImplementationSpecific
backend:
service:
name: RELEASE-NAME-litellm-gateway
port:
number: 4000
- it: drops the /*.txt wildcard for ingress-nginx and keeps every other route as is
set:
ingress.enabled: true
ingress.controller: nginx
asserts:
- notContains:
path: spec.rules[0].http.paths
content:
path: /*.txt
any: true
- contains:
path: spec.rules[0].http.paths
content:
path: /ui
pathType: Prefix
backend:
service:
name: RELEASE-NAME-litellm-ui
port:
number: 3000
- contains:
path: spec.rules[0].http.paths
content:
path: /test
pathType: Exact
backend:
service:
name: RELEASE-NAME-litellm-gateway
port:
number: 4000
- equal:
path: spec.rules[0].http.paths[-1]
value:
path: /
pathType: Prefix
backend:
service:
name: RELEASE-NAME-litellm-backend
port:
number: 4001
- it: rejects an extraPaths entry that repeats a built-in path at the pathType ingress-nginx renders it with
set:
ingress.enabled: true
ingress.controller: nginx
ingress.extraPaths:
- path: /favicon.ico
service: ui
pathType: ImplementationSpecific
asserts:
- failedTemplate:
errorMessage: "ingress.extraPaths[0]: path /favicon.ico with pathType ImplementationSpecific is already routed by this chart, and a duplicate would take it over rather than add to it"
- it: rejects a controller it has no path types for
set:
ingress.enabled: true
ingress.controller: traefik
asserts:
- failedTemplate:
errorMessage: 'ingress.controller: unknown controller "traefik", expected one of alb, nginx'
- it: rejects an extraPaths entry that repeats a built-in path once ingress-nginx normalizes its pathType
set:
ingress.enabled: true
ingress.controller: nginx
ingress.extraPaths:
- path: /favicon.ico
service: ui
asserts:
- failedTemplate:
errorMessage: "ingress.extraPaths[0]: path /favicon.ico with pathType ImplementationSpecific is already routed by this chart, and a duplicate would take it over rather than add to it"
- it: renders a dotted extraPaths entry as ImplementationSpecific for ingress-nginx
set:
ingress.enabled: true
ingress.controller: nginx
ingress.extraPaths:
- path: /eu.assemblyai.custom
service: gateway
- path: /robots.txt
service: ui
pathType: Exact
asserts:
- notMatchRegexRaw:
pattern: 'path: "?/\S*\.\S*"?\n\s+pathType: (Exact|Prefix)\n'
- contains:
path: spec.rules[0].http.paths
content:
path: /eu.assemblyai.custom
pathType: ImplementationSpecific
backend:
service:
name: RELEASE-NAME-litellm-gateway
port:
number: 4000
- contains:
path: spec.rules[0].http.paths
content:
path: /robots.txt
pathType: ImplementationSpecific
backend:
service:
name: RELEASE-NAME-litellm-ui
port:
number: 3000
- it: keeps the requested pathType of a dotted extraPaths entry for the AWS Load Balancer Controller
set:
ingress.enabled: true
ingress.extraPaths:
- path: /eu.assemblyai.custom
service: gateway
- path: /robots.txt
service: ui
pathType: Exact
asserts:
- contains:
path: spec.rules[0].http.paths
content:
path: /eu.assemblyai.custom
pathType: Prefix
backend:
service:
name: RELEASE-NAME-litellm-gateway
port:
number: 4000
- contains:
path: spec.rules[0].http.paths
content:
path: /robots.txt
pathType: Exact
backend:
service:
name: RELEASE-NAME-litellm-ui
port:
number: 3000

View file

@ -10,6 +10,18 @@ imagePullSecrets: []
ingress:
enabled: false
className: ""
# Which ingress controller serves this Ingress. Controllers disagree on the
# pathTypes they accept, so this picks the pathType of the dotted paths, the
# built-in ones and any dotted extraPaths entry alike:
# alb AWS Load Balancer Controller (default): Exact and Prefix paths plus
# the /*.txt wildcard that routes the UI's RSC payloads.
# nginx ingress-nginx: its admission webhook rejects a dot in an Exact or
# Prefix path (strict-validate-path-type, on by default from v1.12.0
# until v1.12.6 / v1.13.2 allowed dots again), so /favicon.ico and
# /eu.assemblyai render as ImplementationSpecific, which nginx serves
# as a plain prefix location. /*.txt is dropped: nginx has no
# wildcard pathType, so that rule could never match there.
controller: alb
annotations: {}
host: "" # optional; if set, becomes the rule's host
tls: []
@ -26,7 +38,8 @@ ingress:
#
# path required; the HTTP path to route
# service which component serves it: gateway (default), backend, or ui
# pathType Prefix (default), Exact, or ImplementationSpecific
# pathType Prefix (default), Exact, or ImplementationSpecific; a dotted
# path renders as ImplementationSpecific when controller is nginx
#
# The target component only answers paths its own route allowlist keeps, so
# a path here still has to be one that component serves.

View file

@ -18,10 +18,15 @@ recoverable one.
constant: it grows with the number of pending migrations, so a fresh database
that has to replay every migration this package ships overruns a per-command
budget sized for the short bookkeeping commands, on a laptop as much as on a
slow CI runner. The Python ``prisma`` wrapper spawns Node and the schema engine
as separate children, so killing the wrapper on timeout leaves them running:
the retry then contends with that orphan for Prisma's advisory lock and cannot
finish any sooner. Migrate deploy therefore runs under its own budget.
slow CI runner. Migrate deploy therefore runs under its own budget.
The Python ``prisma`` wrapper spawns Node, which spawns the Rust schema
engine, so killing only the wrapper on timeout leaves the engine running with
no parent: it keeps mutating the database after the proxy has given up, holds
Prisma's advisory lock so every retry and every later boot queues behind it,
and dies mid-migration once its pipes close, leaving a half-applied ledger row.
Every Prisma command therefore runs in a process group of its own, and a
timeout kills the whole group.
All three budgets are overridable so an operator can widen them without a
release: ``LITELLM_PRISMA_BOOTSTRAP_TIMEOUT`` for the toolchain install,
@ -35,10 +40,12 @@ the deploy override says otherwise.
import math
import os
import shutil
import signal
import subprocess
from collections.abc import Mapping, Sequence
from dataclasses import dataclass
from pathlib import Path
from typing import Optional
from typing import IO, Optional, Union
from litellm_proxy_extras._logging import logger
@ -167,6 +174,49 @@ def heal_incomplete_nodeenv_cache() -> bool:
return True
def _kill_process_group(process: "subprocess.Popen[str]") -> None:
if os.name == "nt":
process.kill()
return
try:
os.killpg(process.pid, signal.SIGKILL)
except ProcessLookupError:
return
def run_prisma(
argv: Sequence[str],
*,
timeout: float,
env: Mapping[str, str],
stdout: Union[IO[str], int, None] = subprocess.PIPE,
stderr: Optional[int] = subprocess.PIPE,
) -> "subprocess.CompletedProcess[str]":
"""Run one Prisma CLI command in its own process group, bounded by ``timeout``.
Raises ``subprocess.TimeoutExpired`` once the budget is spent, after killing
the command together with every process it spawned, and
``subprocess.CalledProcessError`` on a non-zero exit. Output is captured as
text unless ``stdout``/``stderr`` say otherwise.
"""
with subprocess.Popen(
argv,
env=env,
stdout=stdout,
stderr=stderr,
text=True,
start_new_session=True,
) as process:
try:
out, err = process.communicate(timeout=timeout)
except BaseException:
_kill_process_group(process)
raise
if process.returncode:
raise subprocess.CalledProcessError(process.returncode, process.args, out, err)
return subprocess.CompletedProcess(process.args, process.returncode, out, err)
def ensure_prisma_toolchain(
prisma_command: str, prisma_env: dict[str, str]
) -> ToolchainBootstrap:
@ -179,14 +229,7 @@ def ensure_prisma_toolchain(
timeout = prisma_bootstrap_timeout()
logger.info("Preparing the Prisma CLI toolchain (timeout %ss)", timeout)
try:
subprocess.run(
[prisma_command, BOOTSTRAP_ARG],
timeout=timeout,
check=True,
capture_output=True,
text=True,
env=prisma_env,
)
run_prisma([prisma_command, BOOTSTRAP_ARG], timeout=timeout, env=prisma_env)
except subprocess.TimeoutExpired:
logger.warning(
"Preparing the Prisma CLI toolchain timed out after %ss. Raise %s "

View file

@ -16,7 +16,7 @@ import tempfile
from pathlib import Path
from litellm_proxy_extras._logging import logger
from litellm_proxy_extras.prisma_toolchain import prisma_command_timeout
from litellm_proxy_extras.prisma_toolchain import prisma_command_timeout, run_prisma
REPLICA_IDENTITY_FULL_ENV_VAR = "LITELLM_SET_REPLICA_IDENTITY_FULL"
@ -66,7 +66,7 @@ def apply_replica_identity_full(
with tempfile.TemporaryDirectory(prefix="litellm_replica_identity_") as tmp_dir:
sql_path = Path(tmp_dir) / "replica_identity_full.sql"
sql_path.write_text(REPLICA_IDENTITY_FULL_SQL)
subprocess.run(
run_prisma(
[
prisma_command,
"db",
@ -77,9 +77,6 @@ def apply_replica_identity_full(
schema_path,
],
timeout=prisma_command_timeout(),
check=True,
capture_output=True,
text=True,
env=prisma_env,
)
except subprocess.CalledProcessError as e:

View file

@ -10,6 +10,7 @@ from dataclasses import dataclass, replace
from pathlib import Path
from typing import Optional
from litellm_proxy_extras import prisma_toolchain
from litellm_proxy_extras._logging import logger
from litellm_proxy_extras.replica_identity import (
REPLICA_IDENTITY_FULL_ENV_VAR,
@ -231,7 +232,7 @@ class ProxyExtrasDBManager:
# 1. Generate migration SQL file by comparing empty state to current db state
logger.info("Generating baseline migration...")
migration_file = init_dir / "migration.sql"
subprocess.run(
prisma_toolchain.run_prisma(
[
_get_prisma_command(),
"migrate",
@ -242,14 +243,13 @@ class ProxyExtrasDBManager:
"--script",
],
stdout=open(migration_file, "w"),
check=True,
timeout=prisma_command_timeout(),
env=prisma_env,
)
# 3. Mark the migration as applied since it represents current state
logger.info("Marking baseline migration as applied...")
subprocess.run(
prisma_toolchain.run_prisma(
[
_get_prisma_command(),
"migrate",
@ -257,7 +257,6 @@ class ProxyExtrasDBManager:
"--applied",
"0_init",
],
check=True,
timeout=prisma_command_timeout(),
env=prisma_env,
)
@ -286,7 +285,7 @@ class ProxyExtrasDBManager:
"""Mark a specific migration as rolled back"""
# Set up environment for offline mode if configured
prisma_env = _get_prisma_env()
subprocess.run(
prisma_toolchain.run_prisma(
[
_get_prisma_command(),
"migrate",
@ -295,8 +294,6 @@ class ProxyExtrasDBManager:
migration_name,
],
timeout=prisma_command_timeout(),
check=True,
capture_output=True,
env=prisma_env,
)
@ -348,11 +345,9 @@ class ProxyExtrasDBManager:
def _resolve_specific_migration(migration_name: str):
"""Mark a specific migration as applied"""
prisma_env = _get_prisma_env()
subprocess.run(
prisma_toolchain.run_prisma(
[_get_prisma_command(), "migrate", "resolve", "--applied", migration_name],
timeout=prisma_command_timeout(),
check=True,
capture_output=True,
env=prisma_env,
)
@ -436,7 +431,7 @@ class ProxyExtrasDBManager:
try:
logger.info("Generating migration diff between DB and schema.prisma...")
with open(diff_sql_path, "w") as f:
subprocess.run(
prisma_toolchain.run_prisma(
[
_get_prisma_command(),
"migrate",
@ -447,7 +442,6 @@ class ProxyExtrasDBManager:
schema_path,
"--script",
],
check=True,
timeout=prisma_command_timeout(),
stdout=f,
env=_get_prisma_env(),
@ -470,7 +464,7 @@ class ProxyExtrasDBManager:
migration_files = sorted(Path(migrations_dir).glob("*/migration.sql"))
for mig_file in migration_files:
try:
subprocess.run(
prisma_toolchain.run_prisma(
[
_get_prisma_command(),
"db",
@ -481,9 +475,6 @@ class ProxyExtrasDBManager:
schema_path,
],
timeout=prisma_command_timeout(),
check=True,
capture_output=True,
text=True,
env=_get_prisma_env(),
)
logger.info(f"Applied migration: {mig_file.parent.name}")
@ -516,7 +507,7 @@ class ProxyExtrasDBManager:
applied_ok = False
try:
logger.info("Running prisma db execute to apply the migration diff...")
result = subprocess.run(
result = prisma_toolchain.run_prisma(
[
_get_prisma_command(),
"db",
@ -527,9 +518,6 @@ class ProxyExtrasDBManager:
schema_path,
],
timeout=prisma_command_timeout(),
check=True,
capture_output=True,
text=True,
env=_get_prisma_env(),
)
logger.info(f"prisma db execute stdout: {result.stdout}")
@ -558,7 +546,7 @@ class ProxyExtrasDBManager:
for migration_name in migration_names:
try:
logger.info(f"Resolving migration: {migration_name}")
subprocess.run(
prisma_toolchain.run_prisma(
[
_get_prisma_command(),
"migrate",
@ -567,9 +555,6 @@ class ProxyExtrasDBManager:
migration_name,
],
timeout=prisma_command_timeout(),
check=True,
capture_output=True,
text=True,
env=_get_prisma_env(),
)
logger.debug(f"Resolved migration: {migration_name}")
@ -762,11 +747,12 @@ class ProxyExtrasDBManager:
original_dir = os.getcwd()
os.chdir(migrations_dir)
try:
subprocess.run(
prisma_toolchain.run_prisma(
[_get_prisma_command(), "db", "push", "--accept-data-loss"],
timeout=prisma_command_timeout(),
check=True,
env=_get_prisma_env(),
stdout=None,
stderr=None,
)
return True
except (
@ -789,12 +775,9 @@ class ProxyExtrasDBManager:
try:
while not budget.exhausted:
try:
result = subprocess.run(
result = prisma_toolchain.run_prisma(
[_get_prisma_command(), "migrate", "deploy"],
timeout=deploy_timeout,
check=True,
capture_output=True,
text=True,
env=_get_prisma_env(),
)
logger.info(f"prisma migrate deploy stdout: {result.stdout}")
@ -1031,12 +1014,9 @@ class ProxyExtrasDBManager:
logger.info("Running prisma migrate deploy")
try:
# Set migrations directory for Prisma
result = subprocess.run(
result = prisma_toolchain.run_prisma(
[_get_prisma_command(), "migrate", "deploy"],
timeout=prisma_migrate_deploy_timeout(),
check=True,
capture_output=True,
text=True,
env=_get_prisma_env(),
)
logger.info(f"prisma migrate deploy stdout: {result.stdout}")
@ -1108,7 +1088,7 @@ class ProxyExtrasDBManager:
f"Found failed migration: {failed_migration}, marking as rolled back"
)
# Mark the failed migration as rolled back
subprocess.run(
prisma_toolchain.run_prisma(
[
_get_prisma_command(),
"migrate",
@ -1117,9 +1097,6 @@ class ProxyExtrasDBManager:
failed_migration,
],
timeout=prisma_command_timeout(),
check=True,
capture_output=True,
text=True,
env=_get_prisma_env(),
)
logger.info(
@ -1244,10 +1221,12 @@ class ProxyExtrasDBManager:
if ProxyExtrasDBManager.spend_logs_is_partitioned():
raise RuntimeError(PARTITIONED_SPEND_LOGS_PUSH_ERROR)
# Use prisma db push with increased timeout
subprocess.run(
prisma_toolchain.run_prisma(
[_get_prisma_command(), "db", "push", "--accept-data-loss"],
timeout=prisma_command_timeout(),
check=True,
stdout=None,
stderr=None,
env=_get_prisma_env(),
)
return True
except subprocess.TimeoutExpired:

View file

@ -42,7 +42,7 @@ def test_v2_p3018_permission_error_raises_runtime_error(monkeypatch, tmp_path):
"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 patch("litellm_proxy_extras.prisma_toolchain.run_prisma", side_effect=_fake_migrate_deploy_failure(1, stderr)):
with pytest.raises(RuntimeError, match="permission"):
ProxyExtrasDBManager.setup_database(use_migrate=True, use_v2_resolver=True)
@ -60,7 +60,7 @@ def test_v2_non_idempotent_p3009_raises_runtime_error(monkeypatch, tmp_path):
"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 patch("litellm_proxy_extras.prisma_toolchain.run_prisma", 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)
@ -124,7 +124,7 @@ def test_v1_default_still_calls_resolve_all_migrations(monkeypatch, tmp_path):
def fake_resolve(*args, **kwargs):
resolve_called["n"] += 1
monkeypatch.setattr("subprocess.run", fake_run)
monkeypatch.setattr("litellm_proxy_extras.prisma_toolchain.run_prisma", fake_run)
monkeypatch.setattr(ProxyExtrasDBManager, "_resolve_all_migrations", fake_resolve)
ok = ProxyExtrasDBManager.setup_database(use_migrate=True) # v2 flag NOT set
@ -139,7 +139,7 @@ def test_v2_db_push_wraps_subprocess_error_as_runtime_error(monkeypatch, tmp_pat
(tmp_path / "schema.prisma").write_text("// stub")
stderr = "db push error"
with patch("subprocess.run", side_effect=_fake_migrate_deploy_failure(1, stderr)):
with patch("litellm_proxy_extras.prisma_toolchain.run_prisma", 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)
@ -209,7 +209,7 @@ def test_v2_resolve_specific_migration_failure_raises_runtime_error(
"Error: P3009\nMigration `20260101000000_some_migration` failed\n"
"relation already exists"
)
with patch("subprocess.run", side_effect=_fake_migrate_deploy_failure(1, stderr)):
with patch("litellm_proxy_extras.prisma_toolchain.run_prisma", side_effect=_fake_migrate_deploy_failure(1, stderr)):
with pytest.raises(
RuntimeError, match="Failed to mark migration .* as applied"
):
@ -228,7 +228,7 @@ def test_v2_does_not_call_resolve_all_migrations(monkeypatch, tmp_path):
stdout = "Applied migration.\n"
stderr = ""
monkeypatch.setattr("subprocess.run", lambda *a, **kw: FakeResult())
monkeypatch.setattr("litellm_proxy_extras.prisma_toolchain.run_prisma", lambda *a, **kw: FakeResult())
resolve_called = {"n": 0}
monkeypatch.setattr(
@ -296,7 +296,7 @@ def test_v2_p3018_deadlock_rolls_back_and_retries(monkeypatch, tmp_path):
"_resolve_specific_migration",
lambda name: pytest.fail("a deadlocked migration must never be marked applied"),
)
monkeypatch.setattr("subprocess.run", _succeed_after(1, _DEADLOCK_P3018_STDERR))
monkeypatch.setattr("litellm_proxy_extras.prisma_toolchain.run_prisma", _succeed_after(1, _DEADLOCK_P3018_STDERR))
ok = ProxyExtrasDBManager.setup_database(use_migrate=True, use_v2_resolver=True)
assert ok is True
@ -309,7 +309,7 @@ def test_v2_p3018_persistent_deadlock_exhausts_attempts(monkeypatch, tmp_path):
monkeypatch.setattr(ProxyExtrasDBManager, "_roll_back_migration", lambda name: None)
with patch(
"subprocess.run",
"litellm_proxy_extras.prisma_toolchain.run_prisma",
side_effect=_fake_migrate_deploy_failure(1, _DEADLOCK_P3018_STDERR),
):
with pytest.raises(RuntimeError, match="after 4 attempts"):
@ -343,7 +343,7 @@ def test_v2_p3009_deadlocked_ledger_row_rolls_back_and_retries(monkeypatch, tmp_
"_resolve_specific_migration",
lambda name: pytest.fail("a deadlocked migration must never be marked applied"),
)
monkeypatch.setattr("subprocess.run", _succeed_after(1, stderr))
monkeypatch.setattr("litellm_proxy_extras.prisma_toolchain.run_prisma", _succeed_after(1, stderr))
ok = ProxyExtrasDBManager.setup_database(use_migrate=True, use_v2_resolver=True)
assert ok is True
@ -372,7 +372,7 @@ def test_v2_p3009_empty_ledger_logs_rolls_back_and_retries(monkeypatch, tmp_path
"_resolve_specific_migration",
lambda name: pytest.fail("a deadlocked migration must never be marked applied"),
)
monkeypatch.setattr("subprocess.run", _succeed_after(1, stderr))
monkeypatch.setattr("litellm_proxy_extras.prisma_toolchain.run_prisma", _succeed_after(1, stderr))
ok = ProxyExtrasDBManager.setup_database(use_migrate=True, use_v2_resolver=True)
assert ok is True
@ -395,7 +395,7 @@ def test_v2_p3009_unreadable_ledger_still_raises(monkeypatch, tmp_path):
"_roll_back_migration",
lambda name: pytest.fail("an unreadable ledger must not trigger a retry"),
)
monkeypatch.setattr("subprocess.run", _succeed_after(1, stderr))
monkeypatch.setattr("litellm_proxy_extras.prisma_toolchain.run_prisma", _succeed_after(1, stderr))
with pytest.raises(RuntimeError, match="cannot be auto-recovered"):
ProxyExtrasDBManager.setup_database(use_migrate=True, use_v2_resolver=True)
@ -417,7 +417,7 @@ def test_v2_p3009_non_deadlock_ledger_row_still_raises(monkeypatch, tmp_path):
lambda name: 'ERROR: syntax error at or near "BRKN"',
)
with patch("subprocess.run", side_effect=_fake_migrate_deploy_failure(1, stderr)):
with patch("litellm_proxy_extras.prisma_toolchain.run_prisma", 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)
@ -427,7 +427,7 @@ def test_v2_bare_deadlock_stderr_retries(monkeypatch, tmp_path):
waiter as victim) is retried, not fatal."""
_stub_v2_env(monkeypatch, tmp_path)
monkeypatch.setattr(
"subprocess.run", _succeed_after(1, "Database error: deadlock detected")
"litellm_proxy_extras.prisma_toolchain.run_prisma", _succeed_after(1, "Database error: deadlock detected")
)
ok = ProxyExtrasDBManager.setup_database(use_migrate=True, use_v2_resolver=True)
@ -446,7 +446,10 @@ def test_v2_advisory_lock_timeout_retries(monkeypatch, tmp_path):
"""v2: the advisory-lock waiter that times out while a peer's retry holds
the lock retries instead of dying."""
_stub_v2_env(monkeypatch, tmp_path)
monkeypatch.setattr("subprocess.run", _succeed_after(2, _P1002_ADVISORY_LOCK_STDERR))
monkeypatch.setattr(
"litellm_proxy_extras.prisma_toolchain.run_prisma",
_succeed_after(2, _P1002_ADVISORY_LOCK_STDERR),
)
ok = ProxyExtrasDBManager.setup_database(use_migrate=True, use_v2_resolver=True)
assert ok is True
@ -456,7 +459,7 @@ def test_v2_p1002_without_advisory_lock_context_still_raises(monkeypatch, tmp_pa
"""v2: a plain P1002 (database unreachable) stays fatal."""
_stub_v2_env(monkeypatch, tmp_path)
stderr = "Error: P1002\n\nThe database server at `db`:`5432` was reached but timed out."
monkeypatch.setattr("subprocess.run", _succeed_after(1, stderr))
monkeypatch.setattr("litellm_proxy_extras.prisma_toolchain.run_prisma", _succeed_after(1, stderr))
with pytest.raises(RuntimeError, match="cannot be auto-recovered"):
ProxyExtrasDBManager.setup_database(use_migrate=True, use_v2_resolver=True)

View file

@ -26,4 +26,4 @@ variants of it. The test for a good abstraction is that adding the next provider
is a few declarative lines, not a new file of duplicated flow. Only diverge from
the base when behavior is genuinely different, and say so explicitly in the PR.
**Calling:** hosts invoke the core entrypoint — the Python bridge and the `ai-gateway` route service both call `litellm_core::messages::messages`. Never add a provider handler to `ai-gateway`. Register new modules in `lib.rs` / `mod.rs`, then run `cargo fmt && cargo clippy --workspace -- -D warnings && cargo test --workspace`.
**Calling:** hosts invoke the core entrypoint — the Python bridge and the `ai-gateway` route service both call `litellm_core::messages::messages`. Never add a provider handler to `ai-gateway`. Register new modules in `lib.rs` / `mod.rs`, then run the commands under "Checks" in [CLAUDE.md](CLAUDE.md).

View file

@ -174,10 +174,14 @@ for changes under `litellm-rust/`.
```bash
cd litellm-rust
cargo fmt --check
cargo clippy --workspace --all-targets -- -D warnings
cargo clippy -p litellm-core --all-targets --features bedrock-auth -- -D warnings
# the ai-gateway binary + server code is behind the `server` feature
cargo clippy -p litellm-ai-gateway --all-targets --features server -- -D warnings
cargo clippy -p litellm-core -p litellm-python-interop -p litellm-python-bridge --all-targets -- -D warnings
cargo clippy -p litellm-ai-gateway --all-targets --all-features -- -D warnings
cargo test --workspace
cargo test -p litellm-core --features bedrock-auth
# the `auth`, `routes`, `state` and `realtime` tests only exist under `server`
cargo test -p litellm-ai-gateway --features server
```
When a Rust path is exposed through Python, add Python parity tests that compare

View file

@ -49,11 +49,6 @@ function per top-level route, mirroring the core entrypoints.
## Checks
Run these before pushing Rust changes. GitHub Actions runs the same checks for
changes under `litellm-rust/`.
```bash
cargo fmt --check
cargo clippy --workspace --all-targets -- -D warnings
cargo test --workspace
```
Run the commands under "Checks" in [CLAUDE.md](CLAUDE.md) before pushing Rust
changes. That list is the single source of truth and matches what GitHub Actions
runs for changes under `litellm-rust/`.

View file

@ -49,11 +49,5 @@ Rules for adding or changing an LLM provider/route in `litellm-rust`. `messages`
## Checks before push
25. Run, and keep green:
```bash
cd litellm-rust
cargo fmt --check
cargo clippy -p litellm-ai-gateway --all-targets --features server -- -D warnings
cargo clippy -p litellm-core -p litellm-python-interop -p litellm-python-bridge --all-targets -- -D warnings
cargo test --workspace
```
25. Run, and keep green, the commands under "Checks" in `litellm-rust/CLAUDE.md`.
That list is the single source of truth and matches what GitHub Actions runs.

View file

@ -36,12 +36,12 @@ FROM chef AS builder
# whenever only gateway source changes.
COPY --from=planner /build/litellm-rust/recipe.json recipe.json
RUN cargo chef cook --locked --release \
-p litellm-ai-gateway --features python-config \
-p litellm-ai-gateway --features server,python-config \
--recipe-path recipe.json
# Now copy the real sources and build the gateway binary. Deps are already cooked
# above, so this step only recompiles the gateway crate.
COPY litellm-rust/ .
RUN cargo build --locked --release -p litellm-ai-gateway --features python-config
RUN cargo build --locked --release -p litellm-ai-gateway --bin litellm-ai-gateway --features server,python-config
# ---- Runtime ----------------------------------------------------------------
# python:3.11-slim-bookworm ships libpython3.11, matching the builder's PyO3

View file

@ -100,7 +100,7 @@ Worker tuning, rarely needed: `LITELLM_LOG_CHANNEL_CAPACITY` (4096),
## Build & run with Docker
The image is built `--features python-config` and installs litellm **from this
The image is built `--features server,python-config` and installs litellm **from this
repo's source** (the config reader is newer than any PyPI release), so the build
**context is the repo root**:
@ -135,10 +135,10 @@ docker run --rm -p 4001:4001 \
```bash
# config.yaml mode — needs litellm importable in the active python env
LITELLM_CONFIG_PATH=./crates/ai-gateway/config.yaml \
cargo run --release -p litellm-ai-gateway --features python-config
cargo run --release -p litellm-ai-gateway --features server,python-config
# env stand-in mode — no python, no config
cargo run --release -p litellm-ai-gateway
cargo run --release -p litellm-ai-gateway --features server
```
## Deploy on Render

View file

@ -106,16 +106,16 @@ impl RealTimeStreaming {
/// `litellm_call_id`, replacing the gateway-generated fallback.
fn on_session(&mut self, event: &RealtimeEvent) {
let session = event.data.get("session").and_then(Value::as_object);
if let Some(id) = session.and_then(|s| s.get("id")).and_then(Value::as_str) {
if !id.is_empty() {
self.id = id.to_string();
self.litellm_call_id = id.to_string();
}
if let Some(id) = session.and_then(|s| s.get("id")).and_then(Value::as_str)
&& !id.is_empty()
{
self.id = id.to_string();
self.litellm_call_id = id.to_string();
}
if let Some(model) = session.and_then(|s| s.get("model")).and_then(Value::as_str) {
if !model.is_empty() {
self.model = model.to_string();
}
if let Some(model) = session.and_then(|s| s.get("model")).and_then(Value::as_str)
&& !model.is_empty()
{
self.model = model.to_string();
}
}
@ -323,6 +323,32 @@ mod tests {
assert_eq!(streaming.dropped(), 0);
}
#[test]
fn blank_session_id_and_model_keep_the_gateway_fallbacks() {
let mut streaming = RealTimeStreaming::new(
Vec::new(),
"call_fallback".to_string(),
"gpt-realtime".to_string(),
RequestMetadata::default(),
);
streaming.observe(&event(
r#"{"type":"session.created","session":{"id":"","model":""}}"#,
));
let payload = streaming.build_payload();
assert_eq!(payload.id, "call_fallback");
assert_eq!(payload.litellm_call_id, "call_fallback");
assert_eq!(payload.model, "gpt-realtime");
streaming.observe(&event(
r#"{"type":"session.updated","session":{"id":"sess_002","model":""}}"#,
));
let payload = streaming.build_payload();
assert_eq!(payload.id, "sess_002");
assert_eq!(payload.litellm_call_id, "sess_002");
assert_eq!(payload.model, "gpt-realtime");
}
#[test]
fn payload_serializes_with_camelcase_times_and_realtime_call_type() {
let mut streaming = RealTimeStreaming::new(

View file

@ -50,18 +50,17 @@ where
provider_model,
params.api_key.as_deref(),
params.api_base.as_deref(),
) {
if let Some(handoff) = pool.take(&key) {
return crate::io::realtime::realtime_warm(
provider_model,
handoff,
idle_timeout,
observe,
client_in,
client_out,
)
.await;
}
) && let Some(handoff) = pool.take(&key)
{
return crate::io::realtime::realtime_warm(
provider_model,
handoff,
idle_timeout,
observe,
client_in,
client_out,
)
.await;
}
// Cold path: fresh dial (the original behavior).

View file

@ -1,4 +1,5 @@
import asyncio
from collections.abc import Callable, Coroutine
from datetime import datetime, timedelta
from typing import TYPE_CHECKING, Any, Final
@ -83,6 +84,47 @@ class ServiceLogging(CustomLogger):
return open_telemetry_logger
return None
@staticmethod
def _sync_dispatch_loop() -> asyncio.AbstractEventLoop | None:
"""The event loop a blocking caller can dispatch on, or ``None`` if it has none."""
try:
loop: Final = asyncio.get_event_loop()
except RuntimeError:
return None
return None if loop.is_closed() else loop
@staticmethod
async def _emit_guarded(hook: Callable[[], Coroutine[object, object, None]]) -> None:
"""Emit one service event, absorbing anything the callbacks raise.
Monitoring must not break the call it monitors. Sync callers are the ones that
swallow their own service failures (a Redis batch read returns an empty dict),
so an exception from a misconfigured callback would replace a Redis outage with
a callback error and skip the caller's fallback handling.
"""
try:
await hook()
except Exception as e:
verbose_logger.exception("Error emitting service event - %s", e)
@staticmethod
def _dispatch_from_sync(hook: Callable[[], Coroutine[object, object, None]]) -> None:
"""Run an async service hook from a blocking caller, whatever event loop it holds.
Takes a factory rather than a coroutine so the hook is built on the path that
runs it, and only ever once.
"""
loop: Final = ServiceLogging._sync_dispatch_loop()
try:
if loop is None:
asyncio.run(ServiceLogging._emit_guarded(hook))
elif loop.is_running():
loop.create_task(ServiceLogging._emit_guarded(hook))
else:
loop.run_until_complete(ServiceLogging._emit_guarded(hook))
except Exception as e:
verbose_logger.exception("Error dispatching service event - %s", e)
def service_success_hook(
self,
service: ServiceTypes,
@ -99,54 +141,45 @@ class ServiceLogging(CustomLogger):
if self.mock_testing:
self.mock_testing_sync_success_hook += 1
try:
# Try to get the current event loop
loop: Final = asyncio.get_event_loop()
# Check if the loop is running
if loop.is_running():
# If we're in a running loop, create a task
loop.create_task(
self.async_service_success_hook(
service=service,
duration=duration,
call_type=call_type,
parent_otel_span=parent_otel_span,
start_time=start_time,
end_time=end_time,
)
)
else:
# Loop exists but not running, we can use run_until_complete
loop.run_until_complete(
self.async_service_success_hook(
service=service,
duration=duration,
call_type=call_type,
parent_otel_span=parent_otel_span,
start_time=start_time,
end_time=end_time,
)
)
except RuntimeError:
# No event loop exists, create a new one and run
asyncio.run(
self.async_service_success_hook(
service=service,
duration=duration,
call_type=call_type,
parent_otel_span=parent_otel_span,
start_time=start_time,
end_time=end_time,
)
self._dispatch_from_sync(
lambda: self.async_service_success_hook(
service=service,
duration=duration,
call_type=call_type,
parent_otel_span=parent_otel_span,
start_time=start_time,
end_time=end_time,
)
)
def service_failure_hook(self, service: ServiceTypes, duration: float, error: Exception, call_type: str):
def service_failure_hook(
self,
service: ServiceTypes,
duration: float,
error: Exception,
call_type: str,
parent_otel_span: Span | None = None,
start_time: datetime | float | None = None,
end_time: float | datetime | None = None,
):
"""
[TODO] Not implemented for sync calls yet. V0 is focused on async monitoring (used by proxy).
Handles both sync and async monitoring by checking for existing event loop.
"""
if self.mock_testing:
self.mock_testing_sync_failure_hook += 1
self._dispatch_from_sync(
lambda: self.async_service_failure_hook(
service=service,
duration=duration,
error=error,
call_type=call_type,
parent_otel_span=parent_otel_span,
start_time=start_time,
end_time=end_time,
)
)
async def async_service_success_hook(
self,
service: ServiceTypes,

View file

@ -8,10 +8,9 @@ Has 4 primary methods:
- async_get_cache
"""
import asyncio
import time
import traceback
from concurrent.futures import ThreadPoolExecutor
from collections.abc import Sequence
from threading import Lock
from typing import TYPE_CHECKING, Any, Final
@ -188,31 +187,38 @@ class DualCache(BaseCache):
local_only: bool = False,
**kwargs,
):
received_args: Final = locals()
received_args.pop("self")
def run_in_new_loop():
"""Run the coroutine in a new event loop within this thread."""
new_loop: Final = asyncio.new_event_loop()
try:
asyncio.set_event_loop(new_loop)
return new_loop.run_until_complete(self.async_batch_get_cache(**received_args))
finally:
new_loop.close()
asyncio.set_event_loop(None)
try:
# First, try to get the current event loop
_ = asyncio.get_running_loop()
# If we're already in an event loop, run in a separate thread
# to avoid nested event loop issues
with ThreadPoolExecutor(max_workers=1) as executor:
future: Final = executor.submit(run_in_new_loop)
return future.result()
in_memory_result: Final = (
self.in_memory_cache.batch_get_cache(keys, **kwargs) if self.in_memory_cache is not None else None
)
result: Final = in_memory_result if in_memory_result is not None else tuple(None for _ in keys)
except RuntimeError:
# No running event loop, we can safely run in this thread
return run_in_new_loop()
if None not in result or self.redis_cache is None or local_only:
return result
sublist_keys, previous_access_times = self._reserve_redis_batch_keys(time.time(), keys, result)
if len(sublist_keys) == 0:
return result
try:
redis_result: Final = self.redis_cache.batch_get_cache(
key_list=sublist_keys, parent_otel_span=parent_otel_span
)
except Exception:
# Do not throttle subsequent callers if the Redis read fails.
self._rollback_redis_batch_key_reservations(previous_access_times)
raise
if self.in_memory_cache is not None:
for key, value in redis_result.items():
if value is not None:
self.in_memory_cache.set_cache(key, value, **self._backfill_kwargs(kwargs))
return list( # mutable-ok: public list contract
redis_result.get(key) if value is None else value for key, value in zip(keys, result)
)
except Exception:
verbose_logger.error(traceback.format_exc())
async def async_get_cache(
self,
@ -251,7 +257,7 @@ class DualCache(BaseCache):
self,
current_time: float,
keys: list[str],
result: list[Any],
result: Sequence[Any],
) -> tuple[list[str], dict[str, float | None]]:
"""
Atomically choose keys to fetch from Redis and reserve their access time.

View file

@ -78,10 +78,18 @@ class _AsyncRedisCommands(Protocol):
def pipeline(self, transaction: bool = True) -> "Pipeline[bytes]": ...
_BREAKER_GUARD_FRAME_NAMES: Final = frozenset(
{"<lambda>", "wrapper", "_run_under_circuit_breaker", "_run_under_circuit_breaker_sync"}
)
def _get_call_stack_info(num_frames: int = 2) -> str:
"""
Get the function names from the previous 1-2 functions in the call stack.
Frames belonging to this module's circuit-breaker guards are skipped so the
reported callers stay the real ones even on guarded methods.
Args:
num_frames: Number of previous frames to include (default: 2)
@ -102,11 +110,11 @@ def _get_call_stack_info(num_frames: int = 2) -> str:
return "unknown"
function_names: Final = []
for _ in range(num_frames):
if frame is None:
break
func_name = frame.f_code.co_name
function_names.append(func_name)
while frame is not None and len(function_names) < num_frames:
if frame.f_code.co_name in _BREAKER_GUARD_FRAME_NAMES and frame.f_globals.get("__name__") == __name__:
frame = frame.f_back
continue
function_names.append(frame.f_code.co_name)
frame = frame.f_back
if not function_names:
@ -241,6 +249,23 @@ def _record_swallowed_redis_failure(breaker: RedisCircuitBreaker, exc: BaseExcep
_swallowed_redis_failures.set(_swallowed_redis_failures.get() + 1)
def _enter_circuit_breaker(breaker: RedisCircuitBreaker, name: str) -> int:
"""Reject the call if the breaker is open, else return the swallowed-failure count to compare against."""
if breaker.is_open():
raise Exception(f"Redis circuit breaker is open — skipping {name}")
return _swallowed_redis_failures.get()
def _exit_circuit_breaker(breaker: RedisCircuitBreaker, swallowed_before: int) -> None:
"""Record success only when nothing failed while the call ran.
Several Redis methods catch their own connection errors and return a default, so a
method that returned is not on its own proof of a healthy Redis.
"""
if _swallowed_redis_failures.get() == swallowed_before:
breaker.record_success()
async def _run_under_circuit_breaker(
breaker: RedisCircuitBreaker,
name: str,
@ -249,20 +274,33 @@ async def _run_under_circuit_breaker(
"""Run one Redis coroutine under a circuit breaker.
Shared by the method decorator and the Lua script executor so both feed the same
health signal. Success is recorded only when nothing failed while ``call`` ran,
because several Redis methods catch their own connection errors and return a default.
health signal.
"""
if breaker.is_open():
raise Exception(f"Redis circuit breaker is open — skipping {name}")
swallowed_before: Final = _swallowed_redis_failures.get()
swallowed_before: Final = _enter_circuit_breaker(breaker, name)
try:
result: Final = await call()
except Exception as e:
if _is_redis_health_failure(e):
breaker.record_failure()
raise
if _swallowed_redis_failures.get() == swallowed_before:
breaker.record_success()
_exit_circuit_breaker(breaker, swallowed_before)
return result
def _run_under_circuit_breaker_sync(
breaker: RedisCircuitBreaker,
name: str,
call: Callable[[], _RedisCallResult],
) -> _RedisCallResult:
"""Run one blocking Redis call under a circuit breaker, feeding the same health signal as the async path."""
swallowed_before: Final = _enter_circuit_breaker(breaker, name)
try:
result: Final = call()
except Exception as e:
if _is_redis_health_failure(e):
breaker.record_failure()
raise
_exit_circuit_breaker(breaker, swallowed_before)
return result
@ -288,6 +326,14 @@ def _redis_circuit_breaker_guard(method):
return wrapper
def _redis_circuit_breaker_guard_sync(method: Callable[..., _RedisCallResult]) -> Callable[..., _RedisCallResult]:
return functools.wraps(method)(
lambda self, *args, **kwargs: _run_under_circuit_breaker_sync(
self._circuit_breaker, method.__name__, lambda: method(self, *args, **kwargs)
)
)
class RedisCache(BaseCache):
# if users don't provider one, use the default litellm cache
@ -1146,14 +1192,13 @@ class RedisCache(BaseCache):
"""
key_value_dict = {}
_key_list: Final = [key for key in key_list if key is not None]
start_time: Final = time.time()
try:
_keys: Final = []
for cache_key in _key_list:
cache_key = self.check_and_fix_namespace(key=cache_key or "")
_keys.append(cache_key)
start_time: Final = time.time()
swallowed_before: Final = _enter_circuit_breaker(self._circuit_breaker, "batch_get_cache")
_keys: Final = [self.check_and_fix_namespace(key=cache_key or "") for cache_key in _key_list]
results: Final = self._run_redis_mget_operation(keys=_keys)
_exit_circuit_breaker(self._circuit_breaker, swallowed_before)
end_time: Final = time.time()
_duration: Final = end_time - start_time
self.service_logger_obj.service_success_hook(
@ -1178,7 +1223,18 @@ class RedisCache(BaseCache):
return decoded_results
except Exception as e:
failed_at: Final = time.time()
self.service_logger_obj.service_failure_hook(
service=ServiceTypes.REDIS,
duration=failed_at - start_time,
error=e,
call_type=f"batch_get_cache <- {_get_call_stack_info()}",
start_time=start_time,
end_time=failed_at,
parent_otel_span=parent_otel_span,
)
verbose_logger.error("Error occurred in batch get cache - %s", e)
_record_swallowed_redis_failure(self._circuit_breaker, e)
return key_value_dict
@_redis_circuit_breaker_guard

View file

@ -831,10 +831,10 @@ class CustomGuardrail(CustomLogger):
# should run guardrail
litellm_guardrails: Final = request_data.get("guardrails")
if litellm_guardrails is None or not isinstance(litellm_guardrails, list):
return response
return None
if self.should_run_guardrail(data=request_data, event_type=GuardrailEventHooks.post_call) is not True:
return response
return None
# CHECK IF GUARDRAIL REJECTS THE REQUEST
result: Final = await self.async_post_call_success_hook(
@ -850,7 +850,7 @@ class CustomGuardrail(CustomLogger):
)
if not self._is_valid_response_type(result):
return response
return None
return result

View file

@ -6,7 +6,8 @@ It searches the vector store for relevant context and appends it to the messages
"""
from collections.abc import Awaitable, Callable
from typing import TYPE_CHECKING, Any, Final, cast
from dataclasses import dataclass
from typing import TYPE_CHECKING, Any, Final, Protocol, cast
import litellm
import litellm.vector_stores
@ -24,10 +25,35 @@ from litellm.types.vector_stores import (
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.proxy.utils import PrismaClient
from litellm.router import Router
else:
LiteLLMLoggingObj = Any
class ProxyRuntime(Protocol):
def llm_router(self) -> "Router | None": ...
def prisma_client(self) -> "PrismaClient | None": ...
@dataclass(frozen=True, slots=True)
class ProxyServerRuntime:
def llm_router(self) -> "Router | None":
try:
from litellm.proxy.proxy_server import llm_router
except ImportError:
return None
return llm_router
def prisma_client(self) -> "PrismaClient | None":
try:
from litellm.proxy.proxy_server import prisma_client
except ImportError:
return None
return prisma_client
class VectorStorePreCallHook(CustomLogger):
CONTENT_PREFIX_STRING = "Context:\n\n"
"""
@ -39,8 +65,9 @@ class VectorStorePreCallHook(CustomLogger):
3. Appends the search results as context to the messages
"""
def __init__(self):
def __init__(self, proxy_runtime: ProxyRuntime | None = None):
super().__init__()
self.proxy_runtime: Final[ProxyRuntime] = proxy_runtime or ProxyServerRuntime()
async def async_get_chat_completion_prompt(
self,
@ -79,21 +106,8 @@ class VectorStorePreCallHook(CustomLogger):
if litellm.vector_store_registry is None:
return model, messages, non_default_params
# Get prisma_client for database fallback
prisma_client = None
llm_router = None
try:
from litellm.proxy.proxy_server import (
llm_router as _llm_router,
)
from litellm.proxy.proxy_server import (
prisma_client as _prisma_client,
)
prisma_client = _prisma_client
llm_router = _llm_router
except ImportError:
pass
prisma_client: Final = self.proxy_runtime.prisma_client()
llm_router: Final = self.proxy_runtime.llm_router()
# Use database fallback to ensure synchronization across instances
vector_stores_to_run: list[
@ -136,15 +150,23 @@ class VectorStorePreCallHook(CustomLogger):
Callable[..., Awaitable[VectorStoreSearchResponse]],
litellm.vector_stores.asearch,
)
search_response = await search_function(
**{
"vector_store_id": vector_store_id,
"query": query,
"custom_llm_provider": custom_llm_provider,
"metadata": request_metadata,
**litellm_params_for_vector_store,
},
)
try:
search_response = await search_function(
**{
"vector_store_id": vector_store_id,
"query": query,
"custom_llm_provider": custom_llm_provider,
"metadata": request_metadata,
**litellm_params_for_vector_store,
},
)
except Exception as search_error:
verbose_logger.warning(
"Vector store search failed for vector_store_id=%s, continuing without its context: %s",
vector_store_id,
search_error,
)
continue
verbose_logger.debug("search_response: %s", search_response)
@ -153,7 +175,7 @@ class VectorStorePreCallHook(CustomLogger):
# Process search results and append as context
modified_messages = self._append_search_results_to_messages(
messages=messages, search_response=search_response
messages=modified_messages, search_response=search_response
)
# Get the number of results for logging

View file

@ -414,6 +414,11 @@ def _resolve_vertex_location_for_cost(
return VertexBase.get_vertex_region(configured_location, model)
def _provider_response_id(source: object) -> str | None:
candidate: Final = source.get("id") if isinstance(source, dict) else getattr(source, "id", None)
return candidate if isinstance(candidate, str) and candidate else None
class Logging(LiteLLMLoggingBaseClass):
global \
supabaseClient, \
@ -429,6 +434,7 @@ class Logging(LiteLLMLoggingBaseClass):
custom_pricing: bool = False
stream_options = None
litellm_request_debug: bool = False
streamed_anthropic_message_id: str | None = None
def __init__(
self,
@ -2136,7 +2142,7 @@ class Logging(LiteLLMLoggingBaseClass):
self.model_call_details["cache_hit"] = cache_hit
if self.call_type == CallTypes.anthropic_messages.value:
result = self._handle_anthropic_messages_response_logging(result=result)
result = self._anthropic_messages_logged_response(result=result)
elif (
self.call_type == CallTypes.generate_content.value
or self.call_type == CallTypes.agenerate_content.value
@ -3806,6 +3812,23 @@ class Logging(LiteLLMLoggingBaseClass):
)
return None
def record_streamed_anthropic_message_id(self, message_id: str) -> None:
self.streamed_anthropic_message_id = message_id
def _anthropic_messages_logged_response(self, result: Any) -> ModelResponse:
"""
The ModelResponse a /v1/messages spend_logs row is built from.
A streaming call bridged onto the Responses API is the one case where the `msg_` id the
caller was served is minted locally rather than issued upstream, so it is absent from the
response the row would otherwise be keyed on and has to be carried over here.
"""
logged: Final = self._handle_anthropic_messages_response_logging(result=result)
streamed_message_id: Final = self.streamed_anthropic_message_id
if streamed_message_id is None:
return logged
return logged.model_copy(update={"id": streamed_message_id})
def _handle_anthropic_messages_response_logging(self, result: Any) -> ModelResponse:
"""
Handles logging for Anthropic messages responses.
@ -3832,11 +3855,12 @@ class Logging(LiteLLMLoggingBaseClass):
if isinstance(result, ResponsesAPIResponse):
return self._translate_responses_api_response_to_model_response(result)
provider_response_id: Final = _provider_response_id(result)
httpx_response: Final = self.model_call_details.get("httpx_response", None)
if httpx_response and isinstance(httpx_response, httpx.Response):
result = litellm.AnthropicConfig().transform_response(
raw_response=httpx_response,
model_response=litellm.ModelResponse(),
model_response=litellm.ModelResponse(id=provider_response_id),
model=self.model,
messages=[],
logging_obj=self,
@ -3859,7 +3883,7 @@ class Logging(LiteLLMLoggingBaseClass):
status_code=200,
headers={},
),
model_response=litellm.ModelResponse(),
model_response=litellm.ModelResponse(id=provider_response_id),
json_mode=None,
speed=self.optional_params.get("speed") if self.optional_params else None,
)
@ -3882,7 +3906,7 @@ class Logging(LiteLLMLoggingBaseClass):
return LiteLLMResponsesTransformationHandler().transform_response(
model=self.model,
raw_response=result,
model_response=litellm.ModelResponse(),
model_response=litellm.ModelResponse(id=_provider_response_id(result)),
logging_obj=self,
request_data={},
messages=[],
@ -3897,7 +3921,7 @@ class Logging(LiteLLMLoggingBaseClass):
"usage-only ModelResponse to keep the spend_logs row.",
str(e),
)
model_response: Final = litellm.ModelResponse()
model_response: Final = litellm.ModelResponse(id=_provider_response_id(result))
model_response.model = self.model
usage: Final = getattr(result, "usage", None)
if usage is not None and ResponseAPILoggingUtils._is_response_api_usage(usage):

View file

@ -11,7 +11,7 @@ from typing import Final
from litellm.constants import MINIMUM_CUSTOM_KEY_LENGTH
_REDACTED: Final = "REDACTED"
REDACTED: Final = "REDACTED"
def _build_secret_patterns() -> "re.Pattern[str]":
@ -89,7 +89,7 @@ _SECRET_RE: Final = _build_secret_patterns()
def redact_string(value: str) -> str:
"""Scrub known secret/credential patterns from *value* and return the result."""
return _SECRET_RE.sub(_REDACTED, value)
return _SECRET_RE.sub(REDACTED, value)
_UNIX_SYSTEM_PATH: Final = r"/(?:etc|var|opt|usr|home|root|private|Users|tmp|mnt|srv)/[^\s'\"\)\]}>,]+"
@ -110,7 +110,7 @@ def redact_internal_details(value: str) -> str:
on top of redact_string(). For client-facing messages only: server logs keep this detail."""
marker_index: Final = value.find(_TRACEBACK_MARKER)
without_traceback: Final = value[:marker_index].rstrip() if marker_index != -1 else value
return _INTERNAL_DETAIL_RE.sub(_REDACTED, redact_string(without_traceback))
return _INTERNAL_DETAIL_RE.sub(REDACTED, redact_string(without_traceback))
def redact_structured_value(key: str | None, value: str) -> str:
@ -126,4 +126,4 @@ def redact_structured_value(key: str | None, value: str) -> str:
if scrubbed != value or key is None:
return scrubbed
rendered: Final = f"'{key}': '{value}'"
return _REDACTED if redact_string(rendered) != rendered else value
return REDACTED if redact_string(rendered) != rendered else value

View file

@ -1,9 +1,10 @@
from collections.abc import Mapping
from collections.abc import Mapping, Sequence
from typing import Any, Final
from pydantic import BaseModel
from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH_SENSITIVE_DATA_MASKER
from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH, DEFAULT_MAX_RECURSE_DEPTH_SENSITIVE_DATA_MASKER
from litellm.litellm_core_utils.secret_redaction import REDACTED
class SensitiveDataMasker:
@ -214,6 +215,46 @@ def mask_sensitive_keys(data: dict[str, Any], sensitive_fields: set[str]) -> dic
return masked
def redact_credentials_in_payload(data: Mapping[str, object]) -> Mapping[str, object]:
"""Return a copy of ``data`` where every value under a credential-named key is
replaced by the shared ``REDACTED`` marker, nested mappings are recursed into,
and every other value is preserved by identity.
Sensitive-key detection is delegated to the shared :class:`SensitiveDataMasker`,
so the credential names stay in one place. Unlike
:func:`mask_credentials_in_payload`, no prefix or suffix of the secret survives
and non-string secrets are covered too, which is what a payload rendered
straight to stdout needs. ``None`` is preserved so an unset credential still
reads as unset, and lists and tuples are rebuilt element by element so a
credential nested inside one is caught as well. The walk is bounded only to stop
runaway recursion, and a container sitting at that bound is replaced wholesale
rather than passed through, so burying a credential deeper than the walk goes
hides it instead of exposing it.
"""
return _redact_mapping(data, 0)
def _redact_mapping(data: Mapping[str, object], depth: int) -> Mapping[str, object]:
return {key: _redact_entry(key, value, depth) for key, value in data.items()}
def _redact_entry(key: str, value: object, depth: int) -> object:
if value is not None and _default_masker.is_sensitive_key(key):
return REDACTED
if not isinstance(value, (Mapping, list, tuple)):
return value
if depth >= DEFAULT_MAX_RECURSE_DEPTH:
return REDACTED
if isinstance(value, Mapping):
return _redact_mapping(value, depth + 1)
return _redact_sequence(value, depth + 1)
def _redact_sequence(values: Sequence[object], depth: int) -> Sequence[object]:
redacted: Final = tuple(_redact_entry("", item, depth) for item in values)
return redacted if isinstance(values, tuple) else list(redacted)
# Usage example:
"""
masker = SensitiveDataMasker()

View file

@ -21,6 +21,7 @@ from litellm.llms.anthropic.experimental_pass_through.context_management import
)
from litellm.llms.anthropic.experimental_pass_through.utils import (
is_reasoning_auto_summary_enabled,
litellm_logging_obj_from_kwargs,
local_model_name,
)
from litellm.types.llms.anthropic_messages.anthropic_response import (
@ -621,6 +622,7 @@ class LiteLLMMessagesToCompletionTransformationHandler:
tool_name_mapping=tool_name_mapping,
polyfill_result=polyfill_result,
is_async=True,
litellm_logging_obj=litellm_logging_obj_from_kwargs(kwargs),
)
if transformed_stream is not None:
return transformed_stream
@ -755,6 +757,7 @@ class LiteLLMMessagesToCompletionTransformationHandler:
tool_name_mapping=tool_name_mapping,
polyfill_result=polyfill_result,
is_async=False,
litellm_logging_obj=litellm_logging_obj_from_kwargs(kwargs),
)
if transformed_stream is not None:
return transformed_stream

View file

@ -31,6 +31,7 @@ from litellm.types.llms.anthropic import (
from litellm.types.utils import AdapterCompletionStreamWrapper, Delta
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObject
from litellm.types.utils import ModelResponseStream
@ -287,12 +288,16 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
applied_edits: list[AppliedEdit] | None = None,
compaction_block: CompactionBlock | None = None,
iterations_usage: list[UsageIteration] | None = None,
litellm_logging_obj: "LiteLLMLoggingObject | None" = None,
):
# Wrap the upstream stream so chunks that carry both content and a
# finish_reason (fake-streamed providers) are split into two — see
# _CombinedChunkSplitter.
super().__init__(_CombinedChunkSplitter(completion_stream))
self.model = model
self._message_id: str = f"msg_{uuid.uuid4()}"
if litellm_logging_obj is not None:
litellm_logging_obj.record_streamed_anthropic_message_id(self._message_id)
# Mapping of truncated tool names to original names (for OpenAI's 64-char limit)
self.tool_name_mapping = tool_name_mapping or {}
# Polyfill applied_edits on final message_delta.
@ -507,7 +512,7 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
{
"type": "message_start",
"message": {
"id": f"msg_{uuid.uuid4()}",
"id": self._message_id,
"type": "message",
"role": "assistant",
"content": [],
@ -741,7 +746,7 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
{
"type": "message_start",
"message": {
"id": f"msg_{uuid.uuid4()}",
"id": self._message_id,
"type": "message",
"role": "assistant",
"content": [],

View file

@ -174,6 +174,7 @@ from litellm.types.utils import Choices, ModelResponse, StreamingChoices, Usage
from .streaming_iterator import AnthropicStreamWrapper
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObject
from litellm.types.llms.anthropic import ContentBlockContentBlockDict
ToolResultContent: TypeAlias = str | list[ToolMessageContentPart]
@ -264,6 +265,7 @@ class AnthropicAdapter:
tool_name_mapping: dict[str, str] | None = None,
polyfill_result: PolyfillResult | None = None,
is_async: bool = True,
litellm_logging_obj: "LiteLLMLoggingObject | None" = None,
) -> AsyncIterator[bytes] | Iterator[bytes] | None:
"""
Translate OpenAI streaming response to Anthropic format.
@ -290,6 +292,7 @@ class AnthropicAdapter:
applied_edits=applied_edits,
compaction_block=compaction_block,
iterations_usage=iterations_usage,
litellm_logging_obj=litellm_logging_obj,
)
# Return the SSE-wrapped version for proper event formatting.
if is_async:

View file

@ -20,7 +20,7 @@ from litellm.types.llms.anthropic_messages.anthropic_response import (
)
from litellm.types.llms.openai import ResponsesAPIResponse
from ..utils import local_model_name
from ..utils import litellm_logging_obj_from_kwargs, local_model_name
from .streaming_iterator import AnthropicResponsesStreamWrapper
from .transformation import LiteLLMAnthropicToResponsesAPIAdapter
@ -186,7 +186,9 @@ class LiteLLMMessagesToResponsesAPIHandler:
if stream:
wrapper: Final = AnthropicResponsesStreamWrapper(
responses_stream=result, model=local_model_name(model, kwargs.get("custom_llm_provider"))
responses_stream=result,
model=local_model_name(model, kwargs.get("custom_llm_provider")),
litellm_logging_obj=litellm_logging_obj_from_kwargs(responses_kwargs),
)
return wrapper.async_anthropic_sse_wrapper()
@ -266,7 +268,9 @@ class LiteLLMMessagesToResponsesAPIHandler:
if stream:
wrapper: Final = AnthropicResponsesStreamWrapper(
responses_stream=result, model=local_model_name(model, kwargs.get("custom_llm_provider"))
responses_stream=result,
model=local_model_name(model, kwargs.get("custom_llm_provider")),
litellm_logging_obj=litellm_logging_obj_from_kwargs(responses_kwargs),
)
return wrapper.async_anthropic_sse_wrapper()

View file

@ -4,7 +4,7 @@ import json
import traceback
from collections import deque
from collections.abc import AsyncIterator, Mapping
from typing import Any, Final
from typing import TYPE_CHECKING, Any, Final
from litellm import verbose_logger
from litellm._uuid import uuid
@ -12,6 +12,9 @@ from litellm.types.llms.anthropic_messages.anthropic_response import AnthropicUs
from .transformation import LiteLLMAnthropicToResponsesAPIAdapter
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObject
class AnthropicResponsesStreamWrapper:
"""
@ -31,10 +34,13 @@ class AnthropicResponsesStreamWrapper:
self,
responses_stream: Any,
model: str,
litellm_logging_obj: "LiteLLMLoggingObject | None" = None,
) -> None:
self.responses_stream = responses_stream
self.model = model
self._message_id: str = f"msg_{uuid.uuid4()}"
if litellm_logging_obj is not None:
litellm_logging_obj.record_streamed_anthropic_message_id(self._message_id)
self._current_block_index: int = -1
# Map item_id -> content_block_index so we can stop the right block later
self._item_id_to_block_index: dict[str, int] = {}

View file

@ -1,11 +1,14 @@
import os
from collections.abc import Mapping
from types import MappingProxyType
from typing import Final
from typing import TYPE_CHECKING, Final
import litellm
from litellm.types.utils import ModelInfo
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObject
OPENAI_MAX_PROMPT_CACHE_KEY_LENGTH: Final = 64
_EFFORT_DEGRADATION_CHAIN: Final[Mapping[str, tuple[str, ...]]] = MappingProxyType(
@ -24,6 +27,14 @@ def prompt_cache_key_from_user_id(user_id: object) -> str | None:
return str(user_id)[:OPENAI_MAX_PROMPT_CACHE_KEY_LENGTH] or None
def litellm_logging_obj_from_kwargs(kwargs: Mapping[str, object]) -> "LiteLLMLoggingObject | None":
"""The logging object the bridged call logs through, when the caller supplied one."""
from litellm.litellm_core_utils.litellm_logging import Logging
candidate: Final = kwargs.get("litellm_logging_obj")
return candidate if isinstance(candidate, Logging) else None
def local_model_name(model: str, custom_llm_provider: object) -> str:
"""The id the provider itself knows, for reporting back to the caller in ``message_start``."""
return model.removeprefix(f"{custom_llm_provider}/") if isinstance(custom_llm_provider, str) else model

View file

@ -7,7 +7,7 @@ import urllib.parse
from collections.abc import Callable
from datetime import datetime
from threading import Lock
from typing import TYPE_CHECKING, Any, ClassVar, Final, Literal, cast, get_args
from typing import TYPE_CHECKING, Any, ClassVar, Final, Literal, cast, get_args, overload
import httpx
from pydantic import BaseModel, ValidationError
@ -48,12 +48,24 @@ _STS_REGION_FROM_ENDPOINT_PATTERN: Final = re.compile(
SIGV4_COMPUTED_HEADERS: Final = frozenset({"authorization", "x-amz-date", "x-amz-security-token", "date"})
class Boto3CredentialsInfo(BaseModel):
credentials: Credentials
class BedrockRequestTarget(BaseModel):
aws_region_name: str
aws_bedrock_runtime_endpoint: str | None
class Boto3CredentialsInfo(BedrockRequestTarget):
credentials: Credentials
class BearerRequestTarget(BedrockRequestTarget):
credentials: None = None
def bedrock_bearer_token(api_key: str | None) -> str | None:
token: Final = api_key if api_key is not None else get_secret_str("AWS_BEARER_TOKEN_BEDROCK")
return token or None
class _WebIdentityTokenClaims(BaseModel):
aud: str | list[str] | None = None
iss: str | None = None
@ -1387,9 +1399,26 @@ class BaseAWSLLM:
else:
return f"https://bedrock-runtime.{aws_region_name}.{dns_suffix}"
@overload
def _get_boto_credentials_from_optional_params(
self, optional_params: dict, model: str | None = None
) -> Boto3CredentialsInfo:
self,
optional_params: dict, # mutable-ok: the implementation pops the aws_* keys out of the caller's dict in place
model: str | None = None,
bearer_token: None = None,
) -> Boto3CredentialsInfo: ...
@overload
def _get_boto_credentials_from_optional_params(
self,
optional_params: dict, # mutable-ok: the implementation pops the aws_* keys out of the caller's dict in place
model: str | None = None,
*,
bearer_token: str,
) -> BearerRequestTarget: ...
def _get_boto_credentials_from_optional_params(
self, optional_params: dict, model: str | None = None, bearer_token: str | None = None
) -> Boto3CredentialsInfo | BearerRequestTarget:
"""
Get boto3 credentials from optional params
@ -1420,6 +1449,12 @@ class BaseAWSLLM:
) # https://bedrock-runtime.{region_name}.amazonaws.com
aws_external_id: Final = optional_params.pop("aws_external_id", None)
if bearer_token is not None:
return BearerRequestTarget(
aws_region_name=aws_region_name,
aws_bedrock_runtime_endpoint=aws_bedrock_runtime_endpoint,
)
credentials: Final[Credentials] = self.get_credentials(
aws_access_key_id=aws_access_key_id,
aws_secret_access_key=aws_secret_access_key,
@ -1432,7 +1467,6 @@ class BaseAWSLLM:
aws_sts_endpoint=aws_sts_endpoint,
aws_external_id=aws_external_id,
)
return Boto3CredentialsInfo(
credentials=credentials,
aws_region_name=aws_region_name,
@ -1451,14 +1485,9 @@ class BaseAWSLLM:
api_key: str | None = None,
supports_bearer_token: bool = True,
) -> AWSPreparedRequest:
if not supports_bearer_token:
aws_bearer_token: str | None = None
elif api_key is not None:
aws_bearer_token = api_key
else:
aws_bearer_token = get_secret_str("AWS_BEARER_TOKEN_BEDROCK")
aws_bearer_token: Final = bedrock_bearer_token(api_key) if supports_bearer_token else None
if aws_bearer_token:
if aws_bearer_token is not None:
try:
from botocore.awsrequest import AWSRequest
except ImportError:
@ -1555,13 +1584,9 @@ class BaseAWSLLM:
Returns:
Tuple[dict, Optional[str]]: A tuple containing the headers and the json str body of the request
"""
if api_key is not None:
aws_bearer_token: str | None = api_key
else:
aws_bearer_token = get_secret_str("AWS_BEARER_TOKEN_BEDROCK")
aws_bearer_token: Final = bedrock_bearer_token(api_key)
# If aws bearer token is set, use it directly in the header
if aws_bearer_token:
if aws_bearer_token is not None:
headers = headers or {}
headers["Content-Type"] = "application/json"
headers["Authorization"] = f"Bearer {aws_bearer_token}"

View file

@ -21,7 +21,7 @@ from litellm.rust_bridge.chat_completions import rust_chat_completions_accepts
from litellm.types.utils import ModelResponse
from litellm.utils import CustomStreamWrapper
from ..base_aws_llm import BaseAWSLLM, Credentials
from ..base_aws_llm import BaseAWSLLM, Credentials, bedrock_bearer_token
from ..common_utils import BedrockError, _get_all_bedrock_regions
from .invoke_handler import AWSEventStreamDecoder, MockResponseIterator, make_call
@ -349,17 +349,21 @@ class BedrockConverseLLM(BaseAWSLLM):
litellm_params["aws_region_name"] = aws_region_name # [DO NOT DELETE] important for async calls
credentials: Final[Credentials | None] = self.get_credentials(
aws_access_key_id=aws_access_key_id,
aws_secret_access_key=aws_secret_access_key,
aws_session_token=aws_session_token,
aws_region_name=aws_region_name,
aws_session_name=aws_session_name,
aws_profile_name=aws_profile_name,
aws_role_name=aws_role_name,
aws_web_identity_token=aws_web_identity_token,
aws_sts_endpoint=aws_sts_endpoint,
aws_external_id=aws_external_id,
credentials: Final[Credentials | None] = (
None
if bedrock_bearer_token(api_key) is not None
else self.get_credentials(
aws_access_key_id=aws_access_key_id,
aws_secret_access_key=aws_secret_access_key,
aws_session_token=aws_session_token,
aws_region_name=aws_region_name,
aws_session_name=aws_session_name,
aws_profile_name=aws_profile_name,
aws_role_name=aws_role_name,
aws_web_identity_token=aws_web_identity_token,
aws_sts_endpoint=aws_sts_endpoint,
aws_external_id=aws_external_id,
)
)
### SET RUNTIME ENDPOINT ###

View file

@ -149,19 +149,15 @@ class AmazonMoonshotConfig(AmazonInvokeConfig, MoonshotChatConfig):
- Temperature and parameter validation
"""
# Filter out AWS credentials using the existing method from BaseAWSLLM
self._get_boto_credentials_from_optional_params(optional_params, model)
inference_params: Final = {k: v for k, v in optional_params.items() if k not in self.aws_authentication_params}
# Strip routing prefixes to get the actual model ID
clean_model_id: Final = self._get_model_id(model)
# Use Moonshot's transform_request which handles message transformation
# and tool_choice="required" workaround
return MoonshotChatConfig.transform_request(
self,
model=clean_model_id,
messages=messages,
optional_params=optional_params,
optional_params=inference_params,
litellm_params=litellm_params,
headers=headers,
)

View file

@ -6,7 +6,7 @@ import copy
import json
import urllib.parse
from collections.abc import Callable
from typing import TYPE_CHECKING, Any, Final, get_args
from typing import TYPE_CHECKING, Final, get_args, overload
import httpx
@ -26,7 +26,7 @@ from litellm.types.llms.bedrock import (
)
from litellm.types.utils import EmbeddingResponse, LlmProviders
from ..base_aws_llm import BaseAWSLLM
from ..base_aws_llm import BaseAWSLLM, Credentials, bedrock_bearer_token
from ..common_utils import BedrockError
from .amazon_nova_transformation import AmazonNovaEmbeddingConfig
from .amazon_titan_g1_transformation import AmazonTitanG1Config
@ -42,14 +42,25 @@ if TYPE_CHECKING:
class BedrockEmbedding(BaseAWSLLM):
@overload
def _load_credentials(
self,
optional_params: dict, # mutable-ok: the implementation pops the aws_* keys out of the caller's dict in place
bearer_token: None = None,
) -> tuple[Credentials, str]: ...
@overload
def _load_credentials(
self,
optional_params: dict, # mutable-ok: the implementation pops the aws_* keys out of the caller's dict in place
bearer_token: str,
) -> tuple[None, str]: ...
def _load_credentials(
self,
optional_params: dict,
) -> tuple[Any, str]:
try:
from botocore.credentials import Credentials
except ImportError:
raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.")
bearer_token: str | None = None,
) -> tuple[Credentials | None, str]:
## CREDENTIALS ##
# pop aws_secret_access_key, aws_access_key_id, aws_session_token, aws_region_name from kwargs, since completion calls fail with them
aws_secret_access_key: Final = optional_params.pop("aws_secret_access_key", None)
@ -78,17 +89,21 @@ class BedrockEmbedding(BaseAWSLLM):
if aws_region_name is None:
aws_region_name = "us-west-2"
credentials: Final[Credentials] = self.get_credentials(
aws_access_key_id=aws_access_key_id,
aws_secret_access_key=aws_secret_access_key,
aws_session_token=aws_session_token,
aws_region_name=aws_region_name,
aws_session_name=aws_session_name,
aws_profile_name=aws_profile_name,
aws_role_name=aws_role_name,
aws_web_identity_token=aws_web_identity_token,
aws_sts_endpoint=aws_sts_endpoint,
aws_external_id=aws_external_id,
credentials: Final[Credentials | None] = (
None
if bearer_token is not None
else self.get_credentials(
aws_access_key_id=aws_access_key_id,
aws_secret_access_key=aws_secret_access_key,
aws_session_token=aws_session_token,
aws_region_name=aws_region_name,
aws_session_name=aws_session_name,
aws_profile_name=aws_profile_name,
aws_role_name=aws_role_name,
aws_web_identity_token=aws_web_identity_token,
aws_sts_endpoint=aws_sts_endpoint,
aws_external_id=aws_external_id,
)
)
return credentials, aws_region_name
@ -233,7 +248,7 @@ class BedrockEmbedding(BaseAWSLLM):
client: HTTPHandler | None,
timeout: float | httpx.Timeout | None,
batch_data: list[dict],
credentials: Any,
credentials: Credentials | None,
extra_headers: dict | None,
endpoint_url: str,
aws_region_name: str,
@ -301,7 +316,7 @@ class BedrockEmbedding(BaseAWSLLM):
client: AsyncHTTPHandler | None,
timeout: float | httpx.Timeout | None,
batch_data: list[dict],
credentials: Any,
credentials: Credentials | None,
extra_headers: dict | None,
endpoint_url: str,
aws_region_name: str,
@ -383,7 +398,9 @@ class BedrockEmbedding(BaseAWSLLM):
litellm_params: dict,
api_key: str | None = None,
) -> EmbeddingResponse:
credentials, aws_region_name = self._load_credentials(optional_params)
credentials, aws_region_name = self._load_credentials(
optional_params, bearer_token=bedrock_bearer_token(api_key)
)
### TRANSFORMATION ###
unencoded_model_id: Final = optional_params.pop("model_id", None) or model # default to model if not passed

View file

@ -29,7 +29,7 @@ from litellm.llms.custom_httpx.http_handler import (
)
from litellm.types.utils import ImageResponse
from ..base_aws_llm import BaseAWSLLM
from ..base_aws_llm import BaseAWSLLM, bedrock_bearer_token
from ..common_utils import BedrockError
if TYPE_CHECKING:
@ -198,7 +198,9 @@ class BedrockImageEdit(BaseAWSLLM):
Returns:
BedrockImageEditPreparedRequest: The prepared request object
"""
boto3_credentials_info: Final = self._get_boto_credentials_from_optional_params(optional_params, model)
boto3_credentials_info: Final = self._get_boto_credentials_from_optional_params(
optional_params, model, bearer_token=bedrock_bearer_token(api_key)
)
# Use the existing ARN-aware provider detection method
bedrock_provider: Final = self.get_bedrock_invoke_provider(model)

View file

@ -29,7 +29,7 @@ from litellm.llms.custom_httpx.http_handler import (
)
from litellm.types.utils import ImageResponse
from ..base_aws_llm import BaseAWSLLM
from ..base_aws_llm import BaseAWSLLM, bedrock_bearer_token
from ..common_utils import BedrockError
if TYPE_CHECKING:
@ -220,7 +220,9 @@ class BedrockImageGeneration(BaseAWSLLM):
prepped (httpx.Request): The prepared request object
body (bytes): The request body
"""
boto3_credentials_info: Final = self._get_boto_credentials_from_optional_params(optional_params, model)
boto3_credentials_info: Final = self._get_boto_credentials_from_optional_params(
optional_params, model, bearer_token=bedrock_bearer_token(api_key)
)
# Use the existing ARN-aware provider detection method
bedrock_provider: Final = self.get_bedrock_invoke_provider(model)

View file

@ -139,7 +139,7 @@ def _build_query_params(
return {name: value if isinstance(value, str) else str(value) for name, value in supplied if value is not None}
def _error_message_from_response(response: httpx.Response) -> str:
def error_message_from_response(response: httpx.Response) -> str:
try:
body: Final = response.json()
except ValueError:
@ -153,6 +153,16 @@ def _error_message_from_response(response: httpx.Response) -> str:
return response.text
def raise_for_error_status(response: httpx.Response, container_provider_config: "BaseContainerConfig") -> None:
if not httpx.codes.is_error(response.status_code):
return
raise container_provider_config.get_error_class(
error_message=error_message_from_response(response),
status_code=response.status_code,
headers=response.headers,
)
def _transform_response(
response: httpx.Response,
returns_binary: bool,
@ -163,7 +173,7 @@ def _transform_response(
if httpx.codes.is_error(response.status_code):
raise BaseLLMException(
status_code=response.status_code,
message=_error_message_from_response(response),
message=error_message_from_response(response),
headers=dict(response.headers),
)

View file

@ -77,6 +77,7 @@ from litellm.llms.base_llm.vector_store_files.transformation import (
BaseVectorStoreFilesConfig,
)
from litellm.llms.base_llm.videos.transformation import BaseVideoConfig
from litellm.llms.custom_httpx.container_handler import raise_for_error_status
from litellm.llms.custom_httpx.http_handler import (
AsyncHTTPHandler,
HTTPHandler,
@ -8763,17 +8764,19 @@ class BaseLLMHTTPHandler:
json=data,
timeout=timeout,
)
return container_provider_config.transform_container_create_response(
raw_response=response,
logging_obj=logging_obj,
)
except Exception as e:
raise self._handle_error(
e=e,
provider_config=container_provider_config,
)
raise_for_error_status(
response=response,
container_provider_config=container_provider_config,
)
return container_provider_config.transform_container_create_response(
raw_response=response,
logging_obj=logging_obj,
)
async def async_container_create_handler(
self,
@ -8839,17 +8842,19 @@ class BaseLLMHTTPHandler:
json=data,
timeout=timeout,
)
return container_provider_config.transform_container_create_response(
raw_response=response,
logging_obj=logging_obj,
)
except Exception as e:
raise self._handle_error(
e=e,
provider_config=container_provider_config,
)
raise_for_error_status(
response=response,
container_provider_config=container_provider_config,
)
return container_provider_config.transform_container_create_response(
raw_response=response,
logging_obj=logging_obj,
)
def container_list_handler(
self,
@ -8929,17 +8934,19 @@ class BaseLLMHTTPHandler:
headers=headers,
params=params or None,
)
return container_provider_config.transform_container_list_response(
raw_response=response,
logging_obj=logging_obj,
)
except Exception as e:
raise self._handle_error(
e=e,
provider_config=container_provider_config,
)
raise_for_error_status(
response=response,
container_provider_config=container_provider_config,
)
return container_provider_config.transform_container_list_response(
raw_response=response,
logging_obj=logging_obj,
)
async def async_container_list_handler(
self,
@ -9006,17 +9013,19 @@ class BaseLLMHTTPHandler:
headers=headers,
params=params or None,
)
return container_provider_config.transform_container_list_response(
raw_response=response,
logging_obj=logging_obj,
)
except Exception as e:
raise self._handle_error(
e=e,
provider_config=container_provider_config,
)
raise_for_error_status(
response=response,
container_provider_config=container_provider_config,
)
return container_provider_config.transform_container_list_response(
raw_response=response,
logging_obj=logging_obj,
)
def container_retrieve_handler(
self,
@ -9094,17 +9103,19 @@ class BaseLLMHTTPHandler:
headers=headers,
params=params or None,
)
return container_provider_config.transform_container_retrieve_response(
raw_response=response,
logging_obj=logging_obj,
)
except Exception as e:
raise self._handle_error(
e=e,
provider_config=container_provider_config,
)
raise_for_error_status(
response=response,
container_provider_config=container_provider_config,
)
return container_provider_config.transform_container_retrieve_response(
raw_response=response,
logging_obj=logging_obj,
)
async def async_container_retrieve_handler(
self,
@ -9171,17 +9182,19 @@ class BaseLLMHTTPHandler:
headers=headers,
params=params or None,
)
return container_provider_config.transform_container_retrieve_response(
raw_response=response,
logging_obj=logging_obj,
)
except Exception as e:
raise self._handle_error(
e=e,
provider_config=container_provider_config,
)
raise_for_error_status(
response=response,
container_provider_config=container_provider_config,
)
return container_provider_config.transform_container_retrieve_response(
raw_response=response,
logging_obj=logging_obj,
)
def container_delete_handler(
self,
@ -9259,17 +9272,19 @@ class BaseLLMHTTPHandler:
headers=headers,
params=params or None,
)
return container_provider_config.transform_container_delete_response(
raw_response=response,
logging_obj=logging_obj,
)
except Exception as e:
raise self._handle_error(
e=e,
provider_config=container_provider_config,
)
raise_for_error_status(
response=response,
container_provider_config=container_provider_config,
)
return container_provider_config.transform_container_delete_response(
raw_response=response,
logging_obj=logging_obj,
)
async def async_container_delete_handler(
self,
@ -9336,17 +9351,19 @@ class BaseLLMHTTPHandler:
headers=headers,
params=params or None,
)
return container_provider_config.transform_container_delete_response(
raw_response=response,
logging_obj=logging_obj,
)
except Exception as e:
raise self._handle_error(
e=e,
provider_config=container_provider_config,
)
raise_for_error_status(
response=response,
container_provider_config=container_provider_config,
)
return container_provider_config.transform_container_delete_response(
raw_response=response,
logging_obj=logging_obj,
)
def container_file_list_handler(
self,
@ -9428,17 +9445,19 @@ class BaseLLMHTTPHandler:
headers=headers,
params=params or None,
)
return container_provider_config.transform_container_file_list_response(
raw_response=response,
logging_obj=logging_obj,
)
except Exception as e:
raise self._handle_error(
e=e,
provider_config=container_provider_config,
)
raise_for_error_status(
response=response,
container_provider_config=container_provider_config,
)
return container_provider_config.transform_container_file_list_response(
raw_response=response,
logging_obj=logging_obj,
)
async def async_container_file_list_handler(
self,
@ -9507,17 +9526,19 @@ class BaseLLMHTTPHandler:
headers=headers,
params=params or None,
)
return container_provider_config.transform_container_file_list_response(
raw_response=response,
logging_obj=logging_obj,
)
except Exception as e:
raise self._handle_error(
e=e,
provider_config=container_provider_config,
)
raise_for_error_status(
response=response,
container_provider_config=container_provider_config,
)
return container_provider_config.transform_container_file_list_response(
raw_response=response,
logging_obj=logging_obj,
)
def container_file_content_handler(
self,
@ -9593,17 +9614,19 @@ class BaseLLMHTTPHandler:
headers=headers,
params=params or None,
)
return container_provider_config.transform_container_file_content_response(
raw_response=response,
logging_obj=logging_obj,
)
except Exception as e:
raise self._handle_error(
e=e,
provider_config=container_provider_config,
)
raise_for_error_status(
response=response,
container_provider_config=container_provider_config,
)
return container_provider_config.transform_container_file_content_response(
raw_response=response,
logging_obj=logging_obj,
)
async def async_container_file_content_handler(
self,
@ -9669,17 +9692,19 @@ class BaseLLMHTTPHandler:
headers=headers,
params=params or None,
)
return container_provider_config.transform_container_file_content_response(
raw_response=response,
logging_obj=logging_obj,
)
except Exception as e:
raise self._handle_error(
e=e,
provider_config=container_provider_config,
)
raise_for_error_status(
response=response,
container_provider_config=container_provider_config,
)
return container_provider_config.transform_container_file_content_response(
raw_response=response,
logging_obj=logging_obj,
)
###### VECTOR STORE HANDLER ######
@staticmethod

View file

@ -10305,6 +10305,24 @@
"supports_vision": true,
"supports_web_search": true
},
"azure_ai/grok-4.6": {
"cache_read_input_token_cost": 5e-07,
"input_cost_per_token": 2e-06,
"litellm_provider": "azure_ai",
"max_input_tokens": 200000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 6e-06,
"source": "https://techcommunity.microsoft.com/blog/azure-ai-foundry-blog/grok-4-6-comes-to-microsoft-foundry-models-built-for-long-horizon-reasoning-and-/4547578",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
"supports_web_search": true
},
"azure_ai/grok-4-fast-non-reasoning": {
"deprecation_date": "2026-05-01",
"input_cost_per_token": 2e-07,

View file

@ -467,6 +467,42 @@ def _getattr_object(value: object, name: str, default: object = None) -> object:
return getattr(value, name, default)
_OPENAI_ERROR_TYPE_BY_STATUS: Final[Mapping[int, str]] = MappingProxyType(
{
status.HTTP_401_UNAUTHORIZED: "authentication_error",
status.HTTP_403_FORBIDDEN: "permission_error",
status.HTTP_429_TOO_MANY_REQUESTS: "rate_limit_error",
}
)
def _error_status_code(exc: object, default: int) -> int:
"""The HTTP status an exception carries, or ``default`` when it carries none."""
carried: Final = _getattr_object(exc, "status_code")
return carried if isinstance(carried, int) and not isinstance(carried, bool) else default
def _openai_error_type(exc: object, status_code: int) -> str:
"""OpenAI types ``error.type`` as a required string, so an exception carrying none
falls back to the type its status code stands for."""
carried: Final = _getattr_object(exc, "type")
if isinstance(carried, str):
return carried
mapped: Final = _OPENAI_ERROR_TYPE_BY_STATUS.get(status_code)
if mapped is not None:
return mapped
if status_code < status.HTTP_500_INTERNAL_SERVER_ERROR:
return "invalid_request_error"
return "internal_server_error"
def _openai_error_param(exc: object) -> str | None:
"""OpenAI types ``error.param`` as nullable, so an exception carrying none
serializes as JSON ``null``."""
carried: Final = _getattr_object(exc, "param")
return carried if isinstance(carried, str) else None
class _UpstreamHttpResponse(Protocol):
@property
def status_code(self) -> int: ...
@ -540,11 +576,12 @@ def proxy_exception_from_http_exception(exc: HTTPException, headers: dict[str, s
message, structured_fields = serialize_http_exception_detail(raw_detail)
existing_fields: Final = getattr(exc, "provider_specific_fields", None) or {}
merged_fields: Final = {**existing_fields, **structured_fields} if structured_fields else (existing_fields or None)
error_status: Final = _error_status_code(exc, status.HTTP_400_BAD_REQUEST)
return ProxyException(
message=message,
type=getattr(exc, "type", "None"),
param=getattr(exc, "param", "None"),
code=getattr(exc, "status_code", status.HTTP_400_BAD_REQUEST),
type=_openai_error_type(exc, error_status),
param=_openai_error_param(exc),
code=error_status,
provider_specific_fields=merged_fields,
headers=headers,
)
@ -827,25 +864,22 @@ def sse_error_payload(exc: BaseException) -> tuple[int, Mapping[str, object]]:
are byte-identical.
"""
# Preserve status code from HTTPException (e.g. guardrail blocks)
error_status: Final = getattr(exc, "status_code", status.HTTP_500_INTERNAL_SERVER_ERROR)
error_status: Final = _error_status_code(exc, status.HTTP_500_INTERNAL_SERVER_ERROR)
raw_detail: Final = _getattr_object(exc, "detail", "Error processing stream start")
message, structured_fields = serialize_http_exception_detail(raw_detail)
existing_fields: Final = getattr(exc, "provider_specific_fields", None) or {}
merged_fields: Final = {**existing_fields, **structured_fields} if structured_fields else (existing_fields or None)
# Built in one statement then given its one optional key, rather than spread
# conditionally: the spread form costs two extra dict constructions, which
# type-discipline-budget.json's LIT002 ceiling has no room for.
error_obj: Final = {
"message": message,
"type": getattr(exc, "type", "None"),
"param": getattr(exc, "param", "None"),
"type": _openai_error_type(exc, error_status),
"param": _openai_error_param(exc),
"code": str(error_status),
}
if merged_fields:
error_obj["provider_specific_fields"] = merged_fields
return error_status, error_obj
if not merged_fields:
return error_status, error_obj
return error_status, {**error_obj, "provider_specific_fields": merged_fields}
def _sse_error_frames(error_obj: Mapping[str, object]) -> tuple[str, str]:
@ -922,7 +956,7 @@ async def create_response(
"error": {
"message": _CLIENT_DISCONNECT_DETAIL,
"type": "client_disconnect",
"param": "None",
"param": None,
"code": str(LITELLM_HTTP_STATUS_CLIENT_DISCONNECTED),
}
},
@ -3417,8 +3451,8 @@ class ProxyBaseLLMRequestProcessing:
_code = status.HTTP_500_INTERNAL_SERVER_ERROR
raise ProxyException(
message=redact_internal_details_from_client_message(getattr(e, "message", error_msg)),
type=getattr(e, "type", "None"),
param=getattr(e, "param", "None"),
type=_openai_error_type(e, _code),
param=_openai_error_param(e),
openai_code=getattr(e, "code", None),
code=_code,
provider_specific_fields=getattr(e, "provider_specific_fields", None),
@ -3628,11 +3662,12 @@ class ProxyBaseLLMRequestProcessing:
if isinstance(e, HTTPException):
raise e
stream_error_status: Final = _error_status_code(e, status.HTTP_500_INTERNAL_SERVER_ERROR)
proxy_exception: Final = ProxyException(
message=redact_internal_details_from_client_message(getattr(e, "message", str(e))),
type=getattr(e, "type", "None"),
param=getattr(e, "param", "None"),
code=getattr(e, "status_code", 500),
type=_openai_error_type(e, stream_error_status),
param=_openai_error_param(e),
code=stream_error_status,
)
stream_completed = True
yield serialize_error(proxy_exception)

View file

@ -1,10 +1,12 @@
import json
import re
from collections.abc import Collection
from typing import Any, Final
from collections.abc import Collection, Mapping
from types import MappingProxyType, UnionType
from typing import Any, Final, Union, get_args, get_origin
import orjson
from fastapi import Request, UploadFile, status
from typing_extensions import ReadOnly
from litellm._logging import verbose_proxy_logger
from litellm.constants import MAX_REQUEST_BODY_SIZE_TO_REPAIR_MB
@ -40,6 +42,65 @@ def _is_json_content_type(content_type: str) -> bool:
return _normalize_media_type(content_type) == "application/json"
def _numeric_form_type(annotation: object) -> type[int] | type[float] | None:
"""The scalar to parse an ``int``/``float``-typed field as, else ``None``."""
unwrapped: Final = get_args(annotation)[0] if get_origin(annotation) is ReadOnly else annotation
candidates: Final = (
tuple(arg for arg in get_args(unwrapped) if arg is not type(None))
if get_origin(unwrapped) in (Union, UnionType)
else (unwrapped,)
)
if len(candidates) != 1:
return None
if candidates[0] is int:
return int
if candidates[0] is float:
return float
return None
def numeric_form_fields(annotations: Mapping[str, object]) -> Mapping[str, type[int] | type[float]]:
"""
The numeric fields of a request schema, mapped to the scalar to parse them as.
Only a bare ``int``/``float`` or an optional one qualifies, so container and
literal fields are left alone and ``bool`` is excluded on purpose.
"""
return MappingProxyType(
{
name: scalar
for name, annotation in annotations.items()
if (scalar := _numeric_form_type(annotation)) is not None
}
)
def _numeric_form_value(value: object, scalar: type[int] | type[float]) -> object:
if not isinstance(value, str):
return value
try:
return scalar(value)
except ValueError:
return value
def coerce_numeric_form_fields(
parsed_body: Mapping[str, object],
numeric_fields: Mapping[str, type[int] | type[float]],
) -> Mapping[str, object]:
"""
Parse the numeric fields of a form-encoded body back into numbers.
``request.form()`` yields every field as a string, so a provider that puts the
value in a JSON body would send a string where its API requires a number. A
value that will not parse is left as-is for the provider to reject as before.
"""
return {
name: _numeric_form_value(value, numeric_fields[name]) if name in numeric_fields else value
for name, value in parsed_body.items()
}
async def _read_request_body(request: Request | None) -> dict:
"""
Safely read the request body and parse it as JSON.

View file

@ -15,10 +15,11 @@ from litellm.proxy.common_utils.openai_endpoint_utils import (
get_custom_llm_provider_from_request_headers,
get_custom_llm_provider_from_request_query,
)
from litellm.proxy.common_utils.resource_ownership import is_proxy_admin
from litellm.proxy.container_endpoints.ownership import (
assert_user_can_access_container,
filter_container_list_response,
get_container_forwarding_params,
list_owned_containers,
record_container_owner,
)
@ -173,6 +174,9 @@ async def list_containers(
request: Request,
fastapi_response: Response,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
after: str | None = None,
limit: int | None = None,
order: str | None = None,
):
"""
Container list endpoint for retrieving a list of containers.
@ -206,55 +210,54 @@ async def list_containers(
version,
)
# Read query parameters
query_params: Final = dict(request.query_params)
data: Final[dict[str, Any]] = {"query_params": query_params, "model": query_params.get("model")}
# Extract custom_llm_provider using priority chain
custom_llm_provider: Final = (
get_custom_llm_provider_from_request_headers(request=request)
or get_custom_llm_provider_from_request_query(request=request)
or "openai"
)
data: Final[dict[str, Any]] = {
"query_params": query_params,
"model": query_params.get("model"),
"order": order,
"custom_llm_provider": custom_llm_provider,
}
# Add custom_llm_provider to data
data["custom_llm_provider"] = custom_llm_provider
async def fetch_page(page_after: str | None, page_limit: int | None) -> object:
processor: Final = ProxyBaseLLMRequestProcessing(data={**data, "after": page_after, "limit": page_limit})
try:
return await processor.base_process_llm_request(
request=request,
fastapi_response=fastapi_response,
user_api_key_dict=user_api_key_dict,
route_type="alist_containers",
proxy_logging_obj=proxy_logging_obj,
llm_router=llm_router,
general_settings=general_settings,
proxy_config=proxy_config,
select_data_generator=select_data_generator,
model=None,
user_model=user_model,
user_temperature=user_temperature,
user_request_timeout=user_request_timeout,
user_max_tokens=user_max_tokens,
user_api_base=user_api_base,
version=version,
)
except Exception as e:
raise await processor._handle_llm_api_exception(
e=e,
user_api_key_dict=user_api_key_dict,
proxy_logging_obj=proxy_logging_obj,
version=version,
)
# Process request using ProxyBaseLLMRequestProcessing
processor: Final = ProxyBaseLLMRequestProcessing(data=data)
try:
response: Final = await processor.base_process_llm_request(
request=request,
fastapi_response=fastapi_response,
user_api_key_dict=user_api_key_dict,
route_type="alist_containers",
proxy_logging_obj=proxy_logging_obj,
llm_router=llm_router,
general_settings=general_settings,
proxy_config=proxy_config,
select_data_generator=select_data_generator,
model=None,
user_model=user_model,
user_temperature=user_temperature,
user_request_timeout=user_request_timeout,
user_max_tokens=user_max_tokens,
user_api_base=user_api_base,
version=version,
)
except Exception as e:
raise await processor._handle_llm_api_exception(
e=e,
user_api_key_dict=user_api_key_dict,
proxy_logging_obj=proxy_logging_obj,
version=version,
)
# Ownership filtering runs OUTSIDE the LLM-exception scope: a DB error
# in the ownership lookup is not an LLM-API error and shouldn't be
# translated to a provider-shaped failure (which would also fire the
# post_call_failure_hook for what is in fact a successful upstream call).
return await filter_container_list_response(
response=response,
if is_proxy_admin(user_api_key_dict):
return await fetch_page(after, limit)
return await list_owned_containers(
fetch_page=fetch_page,
after=after,
limit=limit,
user_api_key_dict=user_api_key_dict,
custom_llm_provider=custom_llm_provider,
)

View file

@ -6,7 +6,9 @@ FastAPI route handlers for ALL container file endpoints.
"""
import json
from collections.abc import Mapping, Sequence
from pathlib import Path
from types import MappingProxyType
from typing import Any, Final
from fastapi import APIRouter, Depends, Request, Response
@ -56,6 +58,7 @@ def _create_handler_for_path_params(
route_type: str,
returns_binary: bool = False,
is_multipart: bool = False,
query_param_names: Sequence[str] = (),
):
"""
Dynamically create a handler with the correct path parameter signature.
@ -114,6 +117,7 @@ def _create_handler_for_path_params(
user_api_key_dict=user_api_key_dict,
route_type=route_type,
path_params={"container_id": container_id},
query_param_names=query_param_names,
)
return handler_container_id
@ -133,6 +137,7 @@ def _create_handler_for_path_params(
user_api_key_dict=user_api_key_dict,
route_type=route_type,
path_params={"container_id": container_id, "file_id": file_id},
query_param_names=query_param_names,
)
return handler_container_file
@ -150,6 +155,7 @@ def _create_handler_for_path_params(
user_api_key_dict=user_api_key_dict,
route_type=route_type,
path_params={},
query_param_names=query_param_names,
)
return handler_no_params
@ -351,12 +357,17 @@ async def _process_multipart_upload_request(
)
def _declared_query_params(query_params: Mapping[str, str], query_param_names: Sequence[str]) -> Mapping[str, str]:
return MappingProxyType({name: query_params[name] for name in query_param_names if name in query_params})
async def _process_request(
request: Request,
fastapi_response: Response,
user_api_key_dict: UserAPIKeyAuth,
route_type: str,
path_params: dict[str, str],
query_param_names: Sequence[str] = (),
):
"""Common request processing logic."""
from litellm.proxy.proxy_server import (
@ -376,6 +387,7 @@ async def _process_request(
query_params: Final = dict(request.query_params)
data: Final[dict[str, Any]] = {
"query_params": query_params,
**_declared_query_params(query_params, query_param_names),
**path_params,
}
@ -452,7 +464,13 @@ def register_container_file_endpoints(router: APIRouter) -> None:
is_multipart = endpoint_config.get("is_multipart", False)
# Create handler with correct signature for path params
handler = _create_handler_for_path_params(path_params, route_type, returns_binary, is_multipart)
handler = _create_handler_for_path_params(
path_params,
route_type,
returns_binary,
is_multipart,
query_param_names=endpoint_config.get("query_params", ()),
)
# Register routes
route_method = getattr(router, method)

View file

@ -1,9 +1,10 @@
import json
from collections.abc import Mapping, Sequence
from collections.abc import Awaitable, Callable, Mapping, Sequence
from collections.abc import Set as AbstractSet
from typing import TYPE_CHECKING, Any, Final
from typing import TYPE_CHECKING, Any, Final, TypeAlias
from fastapi import HTTPException
from pydantic import BaseModel
from litellm._logging import verbose_proxy_logger
from litellm.caching.in_memory_cache import InMemoryCache
@ -46,6 +47,12 @@ _CONTAINER_STORED_ID_CACHE: Final = InMemoryCache(max_size_in_memory=10000, defa
# different users with different scopes get disjoint cache entries.
_ALLOWED_CONTAINER_IDS_CACHE: Final = InMemoryCache(max_size_in_memory=2048, default_ttl=60)
DEFAULT_CONTAINER_LIST_LIMIT: Final = 20
OWNED_CONTAINER_LIST_PAGE_SIZE: Final = 100
OWNED_CONTAINER_LIST_MAX_PAGES: Final = 5
FetchContainerListPage: TypeAlias = Callable[[str | None, int | None], Awaitable[object]]
def _allowed_container_ids_cache_key(owner_scopes: Sequence[str]) -> str:
"""JSON-encode the sorted scope list — using a separator like ``|``
@ -337,27 +344,23 @@ def _get_container_list_data(response: object) -> Sequence[object] | None:
return data if isinstance(data, list) else None
def _set_container_list_data(response: Any, data: list[object], removed_filtered_items: bool = False) -> object:
def _get_has_more(response: object) -> bool:
if isinstance(response, dict):
response["data"] = data
if data:
response["first_id"] = _get_response_id(data[0])
response["last_id"] = _get_response_id(data[-1])
else:
response["first_id"] = None
response["last_id"] = None
response["has_more"] = False
if removed_filtered_items:
response["has_more"] = False
return response
return response.get("has_more") is True
return getattr(response, "has_more", None) is True
response.data = data
response.first_id = _get_response_id(data[0]) if data else None
response.last_id = _get_response_id(data[-1]) if data else None
if not data and hasattr(response, "has_more"):
response.has_more = False
if removed_filtered_items and hasattr(response, "has_more"):
response.has_more = False
def _with_container_list_page(response: object, data: Sequence[object], has_more: bool) -> object:
page: Final = {
"data": list(data),
"first_id": _get_response_id(data[0]) if data else None,
"last_id": _get_response_id(data[-1]) if data else None,
"has_more": has_more,
}
if isinstance(response, dict):
return {**response, **page}
if isinstance(response, BaseModel):
return response.model_copy(update=page)
return response
@ -366,16 +369,16 @@ async def _get_allowed_container_ids(
) -> AbstractSet[str]:
owner_scopes: Final = get_resource_owner_scopes(user_api_key_dict)
if not owner_scopes:
return set()
return frozenset()
cache_key: Final = _allowed_container_ids_cache_key(owner_scopes)
cached: Final = _ALLOWED_CONTAINER_IDS_CACHE.get_cache(cache_key)
if cached is not None:
return set(cached)
return frozenset(cached)
prisma_client: Final = await _get_prisma_client()
if prisma_client is None:
return set()
return frozenset()
table: Final = ManagedObjectRepository(prisma_client).table
rows: Final[Sequence[prisma_models.LiteLLM_ManagedObjectTable]] = await table.find_many(
@ -384,34 +387,69 @@ async def _get_allowed_container_ids(
"created_by": {"in": owner_scopes},
}
)
allowed_ids: Final = {row.model_object_id for row in rows if getattr(row, "model_object_id", None) is not None}
# ``InMemoryCache.get_cache`` attempts ``json.loads`` on the stored
# value; passing a set would round-trip through that path
# unnecessarily. Store as a list and rehydrate above.
_ALLOWED_CONTAINER_IDS_CACHE.set_cache(cache_key, list(allowed_ids))
allowed_ids: Final = frozenset(
row.model_object_id for row in rows if getattr(row, "model_object_id", None) is not None
)
_ALLOWED_CONTAINER_IDS_CACHE.set_cache(cache_key, tuple(allowed_ids))
return allowed_ids
async def filter_container_list_response(
response: object,
def _is_owned_container(item: object, allowed_container_ids: AbstractSet[str], custom_llm_provider: str) -> bool:
container_id: Final = _get_response_id(item)
if container_id is None:
return False
original_container_id, resolved_provider = decode_container_id_for_ownership(container_id, custom_llm_provider)
return _container_model_object_id(original_container_id, resolved_provider) in allowed_container_ids
async def _collect_owned_containers(
fetch_page: FetchContainerListPage,
after: str | None,
needed: int,
allowed_container_ids: AbstractSet[str],
custom_llm_provider: str,
pages_left: int,
collected: tuple[object, ...],
) -> tuple[object, tuple[object, ...]]:
page: Final = await fetch_page(after, OWNED_CONTAINER_LIST_PAGE_SIZE)
page_data: Final = _get_container_list_data(page) or ()
owned: Final = collected + tuple(
item for item in page_data if _is_owned_container(item, allowed_container_ids, custom_llm_provider)
)
upstream_last_id: Final = _get_response_id(page_data[-1]) if page_data else None
if len(owned) >= needed or upstream_last_id is None or pages_left <= 1 or not _get_has_more(page):
return page, owned
return await _collect_owned_containers(
fetch_page=fetch_page,
after=upstream_last_id,
needed=needed,
allowed_container_ids=allowed_container_ids,
custom_llm_provider=custom_llm_provider,
pages_left=pages_left - 1,
collected=owned,
)
async def list_owned_containers(
fetch_page: FetchContainerListPage,
after: str | None,
limit: int | None,
user_api_key_dict: UserAPIKeyAuth,
custom_llm_provider: str,
) -> object:
if is_proxy_admin(user_api_key_dict):
return response
data: Final = _get_container_list_data(response)
if data is None:
return response
allowed_container_ids: Final = await _get_allowed_container_ids(user_api_key_dict)
filtered: Final[list[object]] = []
for item in data:
container_id = _get_response_id(item)
if container_id is None:
continue
original_container_id, resolved_provider = decode_container_id_for_ownership(container_id, custom_llm_provider)
if _container_model_object_id(original_container_id, resolved_provider) in allowed_container_ids:
filtered.append(item)
return _set_container_list_data(response, filtered, removed_filtered_items=len(filtered) != len(data))
page_limit: Final = limit if limit is not None else DEFAULT_CONTAINER_LIST_LIMIT
last_page, owned = await _collect_owned_containers(
fetch_page=fetch_page,
after=after,
needed=page_limit + 1,
allowed_container_ids=allowed_container_ids,
custom_llm_provider=custom_llm_provider,
pages_left=OWNED_CONTAINER_LIST_MAX_PAGES,
collected=(),
)
return _with_container_list_page(
last_page,
owned[:page_limit],
has_more=len(owned) > page_limit or _get_has_more(last_page),
)

View file

@ -41,7 +41,7 @@ from litellm.llms.anthropic.chat.guardrail_translation.handler import AnthropicM
from litellm.llms.base_llm.guardrail_translation.utils import (
effective_scan_only_tool_results_for_guardrail,
)
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM, bedrock_bearer_token
from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
httpxSpecialProvider,
@ -56,7 +56,6 @@ from litellm.proxy.guardrails.anthropic_sse import (
is_raw_sse_stream,
model_response_text,
)
from litellm.secret_managers.main import get_secret_str
from litellm.types.guardrails import (
BedrockChecksConfigModel,
BedrockGuardrailStreamingParams,
@ -713,9 +712,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
# logic becomes shared across providers.
#### CALL HOOKS - proxy only ####
def _load_credentials(
self,
):
def _load_credentials(self, bearer_token: str | None = None):
try:
from botocore.credentials import Credentials
except ImportError:
@ -737,17 +734,21 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
aws_region_name=aws_region_name,
)
credentials: Final[Credentials] = self.get_credentials(
aws_access_key_id=aws_access_key_id,
aws_secret_access_key=aws_secret_access_key,
aws_session_token=aws_session_token,
aws_region_name=aws_region_name,
aws_session_name=aws_session_name,
aws_profile_name=aws_profile_name,
aws_role_name=aws_role_name,
aws_web_identity_token=aws_web_identity_token,
aws_sts_endpoint=aws_sts_endpoint,
aws_external_id=aws_external_id,
credentials: Final[Credentials | None] = (
None
if bearer_token is not None
else self.get_credentials(
aws_access_key_id=aws_access_key_id,
aws_secret_access_key=aws_secret_access_key,
aws_session_token=aws_session_token,
aws_region_name=aws_region_name,
aws_session_name=aws_session_name,
aws_profile_name=aws_profile_name,
aws_role_name=aws_role_name,
aws_web_identity_token=aws_web_identity_token,
aws_sts_endpoint=aws_sts_endpoint,
aws_external_id=aws_external_id,
)
)
return credentials, aws_region_name
@ -779,13 +780,9 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
proxy_endpoint_url = f"{proxy_endpoint_url}{request_path}"
encoded_data: Final = json.dumps(data).encode("utf-8")
# first check api-key, if none, fall back to sigV4
if api_key is not None:
aws_bearer_token: str | None = api_key
else:
aws_bearer_token = get_secret_str("AWS_BEARER_TOKEN_BEDROCK")
aws_bearer_token: Final = bedrock_bearer_token(api_key)
if aws_bearer_token:
if aws_bearer_token is not None:
try:
from botocore.awsrequest import AWSRequest
except ImportError:
@ -916,7 +913,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
source,
)
return BedrockGuardrailResponse()
credentials, aws_region_name = self._load_credentials()
credentials, aws_region_name = self._load_credentials(bearer_token=bedrock_bearer_token(api_key))
allow_chunking: Final = not self._content_uses_contextual_grounding(content)
completed_chunk_usages: Final[list[BedrockGuardrailUsage]] = [] # mutable-ok: billed-chunk usage accumulator
@ -958,7 +955,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
self,
content: Sequence[BedrockContentItem],
base_request_data: Mapping[str, object],
credentials: "Credentials",
credentials: "Credentials | None",
aws_region_name: str,
api_key: str | None,
request_data: dict | None, # mutable-ok: proxy request body dict, mutated by the logging helper
@ -1096,7 +1093,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
self,
content: Sequence[BedrockContentItem],
base_request_data: Mapping[str, object],
credentials: "Credentials",
credentials: "Credentials | None",
aws_region_name: str,
api_key: str | None,
request_data: dict | None, # mutable-ok: proxy request body dict, mutated by the logging helper
@ -1146,7 +1143,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
self,
content: Sequence[BedrockContentItem],
base_request_data: Mapping[str, object],
credentials: "Credentials",
credentials: "Credentials | None",
aws_region_name: str,
api_key: str | None,
request_data: dict | None, # mutable-ok: proxy request body dict, mutated by the logging helper
@ -1873,9 +1870,9 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
# Nothing to scan (e.g. tool-only turn) -> allow, like ApplyGuardrail does.
return BedrockGuardrailResponse()
credentials, aws_region_name = self._load_credentials()
body: Final[dict[str, object]] = {"messages": checks_messages, "checks": self.checks}
api_key: Final[str | None] = request_data.get("api_key") if request_data else None
credentials, aws_region_name = self._load_credentials(bearer_token=bedrock_bearer_token(api_key))
body: Final[dict[str, object]] = {"messages": checks_messages, "checks": self.checks}
prepared_request: Final = self._prepare_request(
credentials=credentials,

View file

@ -121,9 +121,6 @@ def _a2a_jsonrpc_error_chunk(exc: HTTPException, request_id: str | None) -> Mapp
}
endpoint_guardrail_translation_mappings = None
def _ensure_litellm_metadata(data: dict, user_api_key_dict: UserAPIKeyAuth) -> None:
"""Populate data['litellm_metadata'] from user_api_key_dict if absent."""
if "litellm_metadata" not in data:
@ -164,7 +161,6 @@ class UnifiedLLMGuardrails(CustomLogger):
Use this if you want to MODIFY the input
"""
global endpoint_guardrail_translation_mappings
from litellm.proxy.common_utils.callback_utils import (
add_guardrail_to_applied_guardrails_header,
)
@ -186,18 +182,15 @@ class UnifiedLLMGuardrails(CustomLogger):
)
return data
if endpoint_guardrail_translation_mappings is None:
endpoint_guardrail_translation_mappings = load_guardrail_translation_mappings()
mappings: Final = load_guardrail_translation_mappings()
try:
if CallTypes(call_type) not in endpoint_guardrail_translation_mappings:
if CallTypes(call_type) not in mappings:
return data
except ValueError:
return data # handle unmapped call types
endpoint_translation: Final = _as_endpoint_translation(
endpoint_guardrail_translation_mappings[CallTypes(call_type)]()
)
endpoint_translation: Final = _as_endpoint_translation(mappings[CallTypes(call_type)]())
_ensure_litellm_metadata(data, user_api_key_dict)
@ -222,8 +215,6 @@ class UnifiedLLMGuardrails(CustomLogger):
This can NOT modify the input, only used to reject or accept a call before going to LLM API
"""
global endpoint_guardrail_translation_mappings
verbose_proxy_logger.debug("Running UnifiedLLMGuardrails moderation hook")
guardrail_to_apply: Final[CustomGuardrail] = data.pop("guardrail_to_apply", None)
@ -241,14 +232,11 @@ class UnifiedLLMGuardrails(CustomLogger):
)
return data
if endpoint_guardrail_translation_mappings is None:
endpoint_guardrail_translation_mappings = load_guardrail_translation_mappings()
if call_type is not None and CallTypes(call_type) not in endpoint_guardrail_translation_mappings:
mappings: Final = load_guardrail_translation_mappings()
if call_type is not None and CallTypes(call_type) not in mappings:
return data
endpoint_translation: Final = _as_endpoint_translation(
endpoint_guardrail_translation_mappings[CallTypes(call_type)]()
)
endpoint_translation: Final = _as_endpoint_translation(mappings[CallTypes(call_type)]())
_ensure_litellm_metadata(data, user_api_key_dict)
@ -271,7 +259,6 @@ class UnifiedLLMGuardrails(CustomLogger):
Uses Enkrypt AI guardrails to check the response for policy violations, PII, and injection attacks
"""
global endpoint_guardrail_translation_mappings
# Local import avoids a module-level cyclic import with
# litellm.integrations.custom_guardrail.
from litellm.integrations.custom_guardrail import ModifyResponseException
@ -319,10 +306,9 @@ class UnifiedLLMGuardrails(CustomLogger):
)
return response
if endpoint_guardrail_translation_mappings is None:
endpoint_guardrail_translation_mappings = load_guardrail_translation_mappings()
mappings: Final = load_guardrail_translation_mappings()
if CallTypes(call_type) not in endpoint_guardrail_translation_mappings:
if CallTypes(call_type) not in mappings:
verbose_proxy_logger.warning(
"Guardrail '%s' selected for route '%s' but call type '%s' has no guardrail translation handler; "
"skipping post-call scanning.",
@ -332,9 +318,7 @@ class UnifiedLLMGuardrails(CustomLogger):
)
return response
endpoint_translation: Final = _as_endpoint_translation(
endpoint_guardrail_translation_mappings[CallTypes(call_type)]()
)
endpoint_translation: Final = _as_endpoint_translation(mappings[CallTypes(call_type)]())
try:
response = await endpoint_translation.process_output_response(
@ -906,8 +890,6 @@ class UnifiedLLMGuardrails(CustomLogger):
sampling_rate=1 means every chunk, sampling_rate=5 means every 5th chunk, etc.
"""
global endpoint_guardrail_translation_mappings
# Local import avoids a module-level cyclic import with
# litellm.integrations.custom_guardrail.
from litellm.integrations.custom_guardrail import ModifyResponseException
@ -978,9 +960,7 @@ class UnifiedLLMGuardrails(CustomLogger):
yield item
return
# Initialize translation mappings if needed
if endpoint_guardrail_translation_mappings is None:
endpoint_guardrail_translation_mappings = load_guardrail_translation_mappings()
mappings: Final = load_guardrail_translation_mappings()
# Streaming text transformation (incremental_diff) diverges enough from the
# block_only path that it runs as its own iterator. It requires a route we
@ -989,7 +969,7 @@ class UnifiedLLMGuardrails(CustomLogger):
if streaming_transform_mode == "incremental_diff":
transform_call_type: Final = self._resolve_transform_call_type(
user_api_key_dict=user_api_key_dict,
mappings=endpoint_guardrail_translation_mappings,
mappings=mappings,
)
if transform_call_type is not None:
async for transformed_item in self._run_incremental_transform_stream(
@ -1000,7 +980,7 @@ class UnifiedLLMGuardrails(CustomLogger):
call_type=transform_call_type,
sampling_rate=sampling_rate,
end_of_stream_only=end_of_stream_only,
mappings=endpoint_guardrail_translation_mappings,
mappings=mappings,
):
yield transformed_item
return
@ -1037,7 +1017,7 @@ class UnifiedLLMGuardrails(CustomLogger):
call_type = _infer_call_type(call_type=None, completion_response=item)
# If call type not supported, just pass through all chunks
if call_type is None or CallTypes(call_type) not in endpoint_guardrail_translation_mappings:
if call_type is None or CallTypes(call_type) not in mappings:
yield item
async for remaining_item in response:
yield remaining_item
@ -1049,7 +1029,7 @@ class UnifiedLLMGuardrails(CustomLogger):
# moderation runs below.
if end_of_stream_only:
if not buffer_until_moderated:
endpoint_translation = endpoint_guardrail_translation_mappings[CallTypes(call_type)]()
endpoint_translation = mappings[CallTypes(call_type)]()
stream_has_ended = hasattr(
endpoint_translation, "_check_streaming_has_ended"
) and endpoint_translation._check_streaming_has_ended(responses_so_far)
@ -1063,7 +1043,7 @@ class UnifiedLLMGuardrails(CustomLogger):
# Process chunk based on sampling rate
if chunk_counter % sampling_rate == 0:
endpoint_translation = endpoint_guardrail_translation_mappings[CallTypes(call_type)]()
endpoint_translation = mappings[CallTypes(call_type)]()
scan_key = endpoint_translation.get_streaming_scan_key(responses_so_far)
if _is_redundant_scan(scan_key, last_scan_key):
verbose_proxy_logger.debug(
@ -1143,14 +1123,14 @@ class UnifiedLLMGuardrails(CustomLogger):
yield item
# Stream has ended - do final processing with all collected chunks
if call_type is not None and CallTypes(call_type) in endpoint_guardrail_translation_mappings:
if call_type is not None and CallTypes(call_type) in mappings:
verbose_proxy_logger.debug(
"Processing final streaming response with all %s chunks for guardrail %s",
len(responses_so_far),
guardrail_to_apply.guardrail_name,
)
endpoint_translation = endpoint_guardrail_translation_mappings[CallTypes(call_type)]()
endpoint_translation = mappings[CallTypes(call_type)]()
# When buffering, snapshot the original chunks before moderation.
# A shallow copy suffices: end-of-stream

View file

@ -2,7 +2,7 @@ import asyncio
import io
import traceback
from collections.abc import Sequence
from typing import Final
from typing import Final, get_type_hints
import orjson
from fastapi import APIRouter, Depends, File, HTTPException, Request, Response, UploadFile, status
@ -16,11 +16,18 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import (
from litellm.proxy._types import *
from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth, user_api_key_auth
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
from litellm.proxy.common_utils.http_parsing_utils import (
coerce_numeric_form_fields,
numeric_form_fields,
)
from litellm.proxy.route_llm_request import route_request
from litellm.types.images.main import ImageEditRequestParams
from litellm.types.llms.openai import ChatCompletionUserMessage
router: Final = APIRouter()
IMAGE_EDIT_NUMERIC_FORM_FIELDS: Final = numeric_form_fields(get_type_hints(ImageEditRequestParams))
async def uploadfile_to_bytesio(upload: UploadFile) -> io.BytesIO:
"""
@ -279,7 +286,12 @@ async def image_edit_api(
#########################################################
# Read request body and convert UploadFiles to BytesIO
#########################################################
data: Final = await _read_request_body(request=request)
data: Final = dict(
coerce_numeric_form_fields(
parsed_body=await _read_request_body(request=request),
numeric_fields=IMAGE_EDIT_NUMERIC_FORM_FIELDS,
)
)
image_files: Final = await batch_to_bytesio(image)
mask_files: Final = await batch_to_bytesio(mask)
if image_files:

View file

@ -107,6 +107,7 @@ class AnthropicPassthroughLoggingHandler:
start_time=start_time,
end_time=end_time,
logging_obj=logging_obj,
response_id=optional_str(response_body.get("id")),
)
return {
@ -148,8 +149,9 @@ class AnthropicPassthroughLoggingHandler:
return model
@staticmethod
def _extract_model_from_anthropic_chunks(
def _extract_message_start_field(
all_chunks: Sequence[str | bytes],
field: str,
) -> str | None:
for raw in all_chunks:
text = raw.decode("utf-8") if isinstance(raw, bytes) else raw
@ -163,11 +165,23 @@ class AnthropicPassthroughLoggingHandler:
if not isinstance(data, dict):
continue
if data.get("type") == "message_start":
model = (data.get("message") or {}).get("model")
if model:
return model
value = (data.get("message") or {}).get(field)
if isinstance(value, str) and value:
return value
return None
@staticmethod
def _extract_model_from_anthropic_chunks(
all_chunks: Sequence[str | bytes],
) -> str | None:
return AnthropicPassthroughLoggingHandler._extract_message_start_field(all_chunks, "model")
@staticmethod
def _extract_response_id_from_anthropic_chunks(
all_chunks: Sequence[str | bytes],
) -> str | None:
return AnthropicPassthroughLoggingHandler._extract_message_start_field(all_chunks, "id")
@staticmethod
def _stream_was_interrupted(
all_chunks: Sequence[str | bytes],
@ -251,6 +265,7 @@ class AnthropicPassthroughLoggingHandler:
start_time: datetime,
end_time: datetime,
logging_obj: LiteLLMLoggingObj,
response_id: str | None = None,
):
"""
Create the standard logging object for Anthropic passthrough
@ -312,8 +327,7 @@ class AnthropicPassthroughLoggingHandler:
json.dumps(kwargs, indent=4, default=str),
)
# set litellm_call_id to logging response object
litellm_model_response.id = logging_obj.litellm_call_id
litellm_model_response.id = response_id or logging_obj.litellm_call_id
litellm_model_response.model = model
logging_obj.model_call_details["model"] = model
if not logging_obj.model_call_details.get("custom_llm_provider"):
@ -413,6 +427,7 @@ class AnthropicPassthroughLoggingHandler:
start_time=start_time,
end_time=end_time,
logging_obj=litellm_logging_obj,
response_id=AnthropicPassthroughLoggingHandler._extract_response_id_from_anthropic_chunks(all_chunks),
)
return {

View file

@ -761,6 +761,7 @@ async def rag_query(
model=model,
messages=messages,
retrieval_config=merged_retrieval_config,
vector_store_params=store_data,
rerank=rerank,
stream=stream,
router=llm_router,

View file

@ -11,7 +11,7 @@ __all__ = ["aingest", "aquery", "ingest", "query"]
import asyncio
import contextvars
from collections.abc import Coroutine, Iterator
from collections.abc import Coroutine, Iterator, Mapping
from contextlib import contextmanager
from functools import partial
from types import MappingProxyType
@ -66,6 +66,10 @@ _FORWARDABLE_RETRIEVAL_CONFIG_KEYS: Final = frozenset(
}
)
_SEARCH_ARGS_SET_BY_PIPELINE: Final = frozenset(
{"vector_store_id", "query", "max_num_results", "custom_llm_provider", "router"}
)
def get_ingestion_class(provider: str) -> type[BaseRAGIngestion]:
"""
@ -225,6 +229,7 @@ async def _execute_query_pipeline(
retrieval_config: dict[str, Any],
rerank: dict[str, Any] | None = None,
stream: bool = False,
vector_store_params: Mapping[str, object] | None = None,
**kwargs,
) -> ModelResponse:
"""
@ -241,11 +246,19 @@ async def _execute_query_pipeline(
# 2. Search vector store
# Forward allowlisted provider retrieval_config extras (region, embedding
# model, bucket, credential refs) to the search call; kwargs win on conflict.
# model, bucket, credential refs) to the search call; the managed store's
# params win on conflict.
provider_search_params: Final = MappingProxyType(
{k: v for k, v in retrieval_config.items() if k in _FORWARDABLE_RETRIEVAL_CONFIG_KEYS}
)
forwarded_search_params: Final = MappingProxyType({**provider_search_params, **kwargs})
store_search_params: Final = MappingProxyType(
{
k: v
for k, v in (vector_store_params.items() if vector_store_params else ())
if k not in _SEARCH_ARGS_SET_BY_PIPELINE
}
)
forwarded_search_params: Final = MappingProxyType({**provider_search_params, **kwargs, **store_search_params})
with _suppressed_sub_call_billing():
search_response: Final = await litellm.vector_stores.asearch(
vector_store_id=retrieval_config["vector_store_id"],
@ -339,6 +352,7 @@ async def aquery(
retrieval_config: dict[str, Any],
rerank: dict[str, Any] | None = None,
stream: bool = False,
vector_store_params: Mapping[str, object] | None = None,
**kwargs,
) -> ModelResponse:
"""
@ -356,6 +370,7 @@ async def aquery(
retrieval_config=retrieval_config,
rerank=rerank,
stream=stream,
vector_store_params=vector_store_params,
**kwargs,
)
@ -386,6 +401,7 @@ def query(
retrieval_config: dict[str, Any],
rerank: dict[str, Any] | None = None,
stream: bool = False,
vector_store_params: Mapping[str, object] | None = None,
**kwargs,
) -> ModelResponse | Coroutine[None, None, ModelResponse]:
"""
@ -402,6 +418,7 @@ def query(
retrieval_config=retrieval_config,
rerank=rerank,
stream=stream,
vector_store_params=vector_store_params,
**kwargs,
)
else:
@ -412,6 +429,7 @@ def query(
retrieval_config=retrieval_config,
rerank=rerank,
stream=stream,
vector_store_params=vector_store_params,
**kwargs,
)
)

View file

@ -594,16 +594,20 @@ set_live_deployment_replay(_replay_live_router_model_cost)
# Kwargs that carry no signal about the failed attempt, so log_retry drops them from a
# breadcrumb entirely: the request payload and the router-internal walk state. Credentials are
# handled separately by mask_credentials_in_payload, which scrubs credential-named values from
# whatever kwargs remain rather than trying to enumerate every credential-bearing key here.
# breadcrumb entirely: the request payload, the proxy's snapshot of the inbound request (its body
# aliases the live request metadata, earlier breadcrumbs included, so copying it would nest every
# breadcrumb inside the next one), and the router-internal walk state. Credentials are handled
# separately by mask_credentials_in_payload, which scrubs credential-named values from whatever
# kwargs remain rather than trying to enumerate every credential-bearing key here.
RETRY_BREADCRUMB_EXCLUDED_KWARGS: Final = frozenset(
(
"messages",
"original_function",
"attempted_targets",
"proxy_server_request",
)
)
RETRY_BREADCRUMB_LIMIT: Final = 4
class Router:
@ -964,7 +968,6 @@ class Router:
self.total_calls: defaultdict = defaultdict(int) # dict to store total calls made to each model
self.fail_calls: defaultdict = defaultdict(int) # dict to store fail_calls made to each model
self.success_calls: defaultdict = defaultdict(int) # dict to store success_calls made to each model
self.previous_models: list = [] # list to store failed calls (passed in as metadata to next call)
# make Router.chat.completions.create compatible for openai.chat.completions.create
default_litellm_params = default_litellm_params or {}
@ -8143,35 +8146,31 @@ class Router:
"""
When a retry or fallback happens, log the details of the just failed model call - similar to Sentry breadcrumbing
"""
try:
_metadata_var: Final = "litellm_metadata" if "litellm_metadata" in kwargs else "metadata"
# Log failed model as the previous model
previous_model: Final = {
_metadata_var: Final = "litellm_metadata" if "litellm_metadata" in kwargs else "metadata"
request_metadata: Final[Mapping[str, object]] = kwargs[_metadata_var]
attempt_kwargs: Final = MappingProxyType(
{k: v for k, v in kwargs.items() if k != _metadata_var and k not in RETRY_BREADCRUMB_EXCLUDED_KWARGS}
)
attempt_metadata: Final = MappingProxyType(
{k: v for k, v in request_metadata.items() if k != "previous_models"}
)
previous_model: Final = MappingProxyType(
{
"exception_type": type(e).__name__,
"exception_string": str(e),
**attempt_kwargs,
_metadata_var: attempt_metadata,
}
for (
k,
v,
) in kwargs.items(): # log everything in kwargs except the old previous_models value - prevent nesting
if k != _metadata_var and k not in RETRY_BREADCRUMB_EXCLUDED_KWARGS:
previous_model[k] = v
elif k == _metadata_var and isinstance(v, dict):
previous_model[_metadata_var] = {}
for metadata_k, metadata_v in kwargs[_metadata_var].items():
if metadata_k != "previous_models":
previous_model[k][metadata_k] = metadata_v
# check current size of self.previous_models, if it's larger than 3, remove the first element
if len(self.previous_models) > 3:
self.previous_models.pop(0)
scrubbed_previous_model: Final = mask_credentials_in_payload(previous_model)
self.previous_models.append(scrubbed_previous_model)
kwargs[_metadata_var]["previous_models"] = self.previous_models
return kwargs
except Exception as e:
raise e
)
earlier_breadcrumbs: Final = request_metadata.get("previous_models")
kept_breadcrumbs: Final[tuple[object, ...]] = (
tuple(earlier_breadcrumbs)[-(RETRY_BREADCRUMB_LIMIT - 1) :]
if isinstance(earlier_breadcrumbs, (list, tuple))
else ()
)
breadcrumbs: Final = (*kept_breadcrumbs, mask_credentials_in_payload(previous_model))
kwargs[_metadata_var]["previous_models"] = breadcrumbs # rebind-ok: the logging object already holds this dict
return kwargs
def _update_usage(self, deployment_id: str, parent_otel_span: Span | None) -> int:
"""

View file

@ -1,6 +1,7 @@
#### What this does ####
# identifies lowest tpm deployment
import random
from collections.abc import Sequence
from typing import TYPE_CHECKING, Any, Final
import httpx
@ -350,9 +351,9 @@ class LowestTPMLoggingHandler_v2(BaseRoutingStrategy, CustomLogger):
model_group: str,
healthy_deployments: list,
tpm_keys: list,
tpm_values: list | None,
tpm_values: Sequence | None,
rpm_keys: list,
rpm_values: list | None,
rpm_values: Sequence | None,
messages: list[dict[str, str]] | None = None,
input: str | list | None = None,
) -> dict | None:

View file

@ -83,6 +83,7 @@ from litellm.constants import (
from litellm.litellm_core_utils.fallback_generalizations import (
match_capability_generalizations,
)
from litellm.litellm_core_utils.sensitive_data_masker import redact_credentials_in_payload
_CachingHandlerResponse = None
_LLMCachingHandler = None
@ -543,6 +544,14 @@ def print_verbose(
pass
def _print_verbose_is_active() -> bool:
"""Whether print_verbose would reach either of its two consumers, so a call site can skip
building a payload nothing would read. _is_debugging_on() is not the same predicate: it reads
litellm._logging.set_verbose, while print_verbose's print reads litellm.set_verbose, and
assigning the documented litellm.set_verbose = True rebinds only the latter."""
return litellm.set_verbose is True or verbose_logger.isEnabledFor(logging.DEBUG)
####### CLIENT ###################
# make it easy to log if completion/embedding runs succeeded or failed + see what happened | Non-Blocking
def custom_llm_setup():
@ -1284,16 +1293,18 @@ async def async_post_call_success_deployment_hook(
except ValueError:
typed_call_type = None # unknown call type
modified_response = response
CustomLogger: Final = _get_cached_custom_logger()
for callback in litellm.callbacks:
if isinstance(callback, CustomLogger):
result = await callback.async_post_call_success_deployment_hook(
request_data, cast(LLMResponseTypes, response), typed_call_type
request_data, cast(LLMResponseTypes, modified_response), typed_call_type
)
if result is not None:
return result
modified_response = result
return response
return modified_response
async def async_post_call_failure_deployment_hook(
@ -4707,7 +4718,8 @@ def get_optional_params(
openai_params=list(DEFAULT_CHAT_COMPLETION_PARAM_VALUES.keys()),
additional_drop_params=additional_drop_params,
)
print_verbose(f"Final returned optional params: {optional_params}")
if _print_verbose_is_active():
print_verbose(f"Final returned optional params: {redact_credentials_in_payload(optional_params)}")
optional_params = _apply_openai_param_overrides(
optional_params=optional_params,
non_default_params=non_default_params,
@ -7462,7 +7474,8 @@ def print_args_passed_to_litellm(original_function, args, kwargs):
return
args_str: Final = ", ".join(map(repr, args))
kwargs_str: Final = ", ".join(f"{key}={value!r}" for key, value in kwargs.items())
redacted_kwargs: Final = redact_credentials_in_payload(kwargs)
kwargs_str: Final = ", ".join(f"{key}={value!r}" for key, value in redacted_kwargs.items())
print_verbose(
"\n",
) # new line before

View file

@ -10305,6 +10305,24 @@
"supports_vision": true,
"supports_web_search": true
},
"azure_ai/grok-4.6": {
"cache_read_input_token_cost": 5e-07,
"input_cost_per_token": 2e-06,
"litellm_provider": "azure_ai",
"max_input_tokens": 200000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 6e-06,
"source": "https://techcommunity.microsoft.com/blog/azure-ai-foundry-blog/grok-4-6-comes-to-microsoft-foundry-models-built-for-long-horizon-reasoning-and-/4547578",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
"supports_web_search": true
},
"azure_ai/grok-4-fast-non-reasoning": {
"deprecation_date": "2026-05-01",
"input_cost_per_token": 2e-07,

View file

@ -42,7 +42,7 @@
"limit": 52
},
"B010": {
"limit": 188
"limit": 187
},
"B018": {
"limit": 2
@ -144,7 +144,7 @@
"limit": 1
},
"PLR1704": {
"limit": 3
"limit": 1
},
"PLR1714": {
"limit": 253
@ -240,13 +240,13 @@
"limit": 96
},
"TRY201": {
"limit": 403
"limit": 401
},
"TRY203": {
"limit": 111
"limit": 109
},
"TRY300": {
"limit": 854
"limit": 852
},
"UP028": {
"limit": 2

View file

@ -168,6 +168,35 @@ def completed_responses_object(result: StreamingResponse) -> ResponsesObject | N
return completed[-1] if completed else None
class AnthropicMessageObject(BaseModel):
id: str
class AnthropicStreamEvent(BaseModel):
"""One SSE frame of a native Anthropic stream. Only `message_start` carries the
message, so it stays optional and the deltas validate as themselves."""
type: str
message: AnthropicMessageObject | None = None
def anthropic_message_id(result: StreamingResponse) -> str | None:
"""The `msg_...` id the caller was served, which is what the spend row is keyed by
on this route: off the `message_start` frame when streaming, off the body when not."""
if not result.is_streaming:
return AnthropicMessageObject.model_validate_json(result.body).id
events = (
AnthropicStreamEvent.model_validate_json(payload)
for payload in result.stream_events
)
started = tuple(
event.message
for event in events
if event.type == "message_start" and event.message is not None
)
return started[0].id if started else None
class OpenAIResponsesBody(BaseModel):
model: str
input: str

View file

@ -2,7 +2,8 @@
Each test sends a NATIVE provider request through the proxy's passthrough route
and verifies the proxy still logged a costed SpendLogs row
(call_type="pass_through_endpoint"), correlated by the x-litellm-call-id header.
(call_type="pass_through_endpoint"), correlated by the id the caller was served:
the x-litellm-call-id header on gemini, the `msg_...` message id on anthropic.
Covered: gemini ("gemini-2.5-flash") + anthropic ("claude-haiku-4-5"), streaming +
non-streaming, plus native tool calls. See LLM_TRANSLATION_COVERAGE_MATRIX.md.
@ -14,7 +15,7 @@ A passthrough call returning non-2xx fails hard (never a skip); once it returns
import pytest
from e2e_config import CHEAP_OPENAI_MODEL, unique_marker
from e2e_http import StreamingResponse, require_successful_call, unwrap
from e2e_http import require_successful_call, unwrap
from lifecycle import ResourceManager
from models import KeyGenerateBody, SpendLogRow
from passthrough_client import (
@ -24,6 +25,7 @@ from passthrough_client import (
JsonSchema,
JsonSchemaProperty,
PassthroughClient,
anthropic_message_id,
completed_responses_object,
)
@ -33,18 +35,18 @@ REALTIME_MODEL = "gpt-realtime-2"
pytestmark = pytest.mark.e2e
def _fetch_cost_breakdown(client: PassthroughClient, result: StreamingResponse) -> SpendLogRow:
def _fetch_cost_breakdown(client: PassthroughClient, request_id: str | None) -> SpendLogRow:
"""The passthrough call's logged row, polled until it carries a cost.
Asserts (not skips) that a 2xx passthrough call produced a costed row - the
whole point of passthrough spend tracking.
"""
assert result.call_id, "passthrough response had no x-litellm-call-id header"
assert request_id, "passthrough response carried no id to correlate its spend row by"
rows = client.proxy.poll_logs_for_request_id(
result.call_id,
request_id,
predicate=lambda rs: (rs[0].spend or 0) > 0,
)
assert rows, f"no SpendLogs row for passthrough call_id {result.call_id}"
assert rows, f"no SpendLogs row for passthrough request_id {request_id}"
row = rows[0]
assert row.call_type == "pass_through_endpoint"
assert (row.spend or 0) > 0, f"passthrough call was not costed: {row}"
@ -64,7 +66,7 @@ def test_gemini_passthrough_nonstreaming_logs_cost(
)
require_successful_call(result)
row = _fetch_cost_breakdown(client, result)
row = _fetch_cost_breakdown(client, result.call_id)
assert row.custom_llm_provider == "gemini"
assert "gemini" in (row.model or "")
assert tag in (row.request_tags or []), f"tags not logged: {row.request_tags}"
@ -107,7 +109,7 @@ def test_gemini_passthrough_streaming_logs_cost(
require_successful_call(result)
assert result.chunks > 0, "streaming passthrough produced no events"
row = _fetch_cost_breakdown(client, result)
row = _fetch_cost_breakdown(client, result.call_id)
assert row.custom_llm_provider == "gemini"
@ -137,7 +139,7 @@ def test_gemini_passthrough_tool_call_logs_cost(
require_successful_call(result)
assert "functionCall" in result.body, "gemini did not emit a tool call"
row = _fetch_cost_breakdown(client, result)
row = _fetch_cost_breakdown(client, result.call_id)
assert row.custom_llm_provider == "gemini"
@ -150,7 +152,7 @@ def test_anthropic_passthrough_nonstreaming_logs_cost(
result = client.anthropic_message(scoped_key, "claude-haiku-4-5", "Say hello")
require_successful_call(result)
row = _fetch_cost_breakdown(client, result)
row = _fetch_cost_breakdown(client, anthropic_message_id(result))
assert row.custom_llm_provider == "anthropic"
assert "claude" in (row.model or "")
@ -164,7 +166,7 @@ def test_anthropic_passthrough_streaming_logs_cost(
require_successful_call(result)
assert result.chunks > 0, "streaming passthrough produced no events"
row = _fetch_cost_breakdown(client, result)
row = _fetch_cost_breakdown(client, anthropic_message_id(result))
assert row.custom_llm_provider == "anthropic"
@ -190,7 +192,7 @@ def test_anthropic_passthrough_tool_call_logs_cost(
require_successful_call(result)
assert "tool_use" in result.body, "anthropic did not emit a tool call"
row = _fetch_cost_breakdown(client, result)
row = _fetch_cost_breakdown(client, anthropic_message_id(result))
assert row.custom_llm_provider == "anthropic"

View file

@ -567,7 +567,7 @@ class TestResolveAllMigrationsLedger:
return _FakeCompleted()
return _FakeCompleted()
monkeypatch.setattr(utils_module.subprocess, "run", fake_run)
monkeypatch.setattr(utils_module.prisma_toolchain, "run_prisma", fake_run)
ProxyExtrasDBManager._resolve_all_migrations(str(tmp_path), "schema.prisma")
return calls
@ -604,9 +604,9 @@ class TestPartitionedSpendLogsPushGuard:
import litellm_proxy_extras.utils as utils_module
def fail_run(cmd, **kwargs):
raise AssertionError(f"subprocess.run should not be called, got: {cmd}")
raise AssertionError(f"run_prisma should not be called, got: {cmd}")
monkeypatch.setattr(utils_module.subprocess, "run", fail_run)
monkeypatch.setattr(utils_module.prisma_toolchain, "run_prisma", fail_run)
def test_v1_db_push_fails_fast_with_guidance(self, monkeypatch):
monkeypatch.setattr(

View file

@ -240,3 +240,35 @@ async def test_dual_cache_delete(is_async):
result = dual_cache.get_cache(test_key)
assert result is None
@pytest.mark.asyncio
async def test_dual_cache_concurrent_sync_and_async_redis_reads():
"""Sync and async batch reads share one Redis backend in one process, and sync reads never open an async connection"""
redis_cache = RedisCache(host=os.getenv("REDIS_HOST"), port=os.getenv("REDIS_PORT"))
dual_cache = DualCache(redis_cache=redis_cache)
run_id = str(uuid.uuid4())
sync_keys = [f"sync_{run_id}_{index}" for index in range(5)]
async_keys = [f"async_{run_id}_{index}" for index in range(5)]
in_loop_keys = [f"in_loop_{run_id}_{index}" for index in range(3)]
survivor_key = f"survivor_{run_id}"
expected = {key: {"key": key} for key in [*sync_keys, *async_keys, *in_loop_keys, survivor_key]}
for key, value in expected.items():
await redis_cache.async_set_cache(key, value, ttl=60)
concurrent_results = await asyncio.gather(
*(asyncio.to_thread(dual_cache.batch_get_cache, keys=[key]) for key in sync_keys),
*(dual_cache.async_batch_get_cache(keys=[key]) for key in async_keys),
)
assert list(concurrent_results) == [[expected[key]] for key in [*sync_keys, *async_keys]]
with patch.object(
redis_cache,
"async_batch_get_cache",
side_effect=AssertionError("sync batch reads must not call async Redis"),
):
in_loop_results = [dual_cache.batch_get_cache(keys=[key]) for key in in_loop_keys]
assert in_loop_results == [[expected[key]] for key in in_loop_keys]
assert await dual_cache.async_batch_get_cache(keys=[survivor_key]) == [expected[survivor_key]]

View file

@ -50,9 +50,9 @@ async def test_anthropic_basic_completion_with_headers():
anthropic_api_output_tokens = (
reported_usage.get("output_tokens", None) if reported_usage else None
)
litellm_call_id = response_headers.get("x-litellm-call-id")
anthropic_message_id = response_json.get("id")
print(f"LiteLLM Call ID: {litellm_call_id}")
print(f"Anthropic message ID: {anthropic_message_id}")
# Wait for spend to be logged
await asyncio.sleep(15)
@ -64,7 +64,7 @@ async def test_anthropic_basic_completion_with_headers():
print(f"Attempt {attempt + 1}/{max_retries} to check spend logs")
async with session.get(
f"http://0.0.0.0:4000/spend/logs?request_id={litellm_call_id}",
f"http://0.0.0.0:4000/spend/logs?request_id={anthropic_message_id}",
headers={"Authorization": "Bearer sk-1234"},
) as spend_response:
print("text spend response")
@ -84,25 +84,25 @@ async def test_anthropic_basic_completion_with_headers():
print("Waiting 10 seconds before retry...")
await asyncio.sleep(10)
# Spend data might be unavailable (auth error, slow DB write, etc.)
if (
spend_data is None
or not isinstance(spend_data, list)
or len(spend_data) == 0
or not isinstance(spend_data[0], dict)
or "request_id" not in spend_data[0]
):
print(f"Spend data not available or is error response: {spend_data}")
print("Skipping spend assertions (DB write may be slow in CI)")
if not isinstance(spend_data, list):
print(f"Spend endpoint answered with an error response: {spend_data}")
print("Skipping spend assertions (spend logs unreachable in CI)")
return
assert spend_data, (
f"GET /spend/logs?request_id={anthropic_message_id} found no row for the id "
"the caller received"
)
log_entry = spend_data[0]
# Basic existence checks
assert isinstance(log_entry, dict), "Log entry should be a dictionary"
# Request metadata assertions
assert log_entry["request_id"] == litellm_call_id, "Request ID should match"
assert (
log_entry["request_id"] == anthropic_message_id
), "Request ID should be the message id the caller received"
assert (
log_entry["call_type"] == "pass_through_endpoint"
), "Call type should be pass_through_endpoint"
@ -182,8 +182,6 @@ async def test_anthropic_streaming_with_headers():
assert response.status == 200, "Response should be successful"
response_headers = response.headers
print(f"Response headers: {response_headers}")
litellm_call_id = response_headers.get("x-litellm-call-id")
print(f"LiteLLM Call ID: {litellm_call_id}")
collected_output = []
async for line in response.content:
@ -194,13 +192,18 @@ async def test_anthropic_streaming_with_headers():
print("Collected output:", "".join(collected_output))
anthropic_api_usage_chunks = []
anthropic_message_id = None
for chunk in collected_output:
chunk_json = json.loads(chunk)
if chunk_json.get("type") == "message_start":
anthropic_message_id = chunk_json.get("message", {}).get("id")
if "usage" in chunk_json:
anthropic_api_usage_chunks.append(chunk_json["usage"])
elif "message" in chunk_json and "usage" in chunk_json["message"]:
anthropic_api_usage_chunks.append(chunk_json["message"]["usage"])
print(f"Anthropic message ID: {anthropic_message_id}")
print(
"anthropic_api_usage_chunks",
json.dumps(anthropic_api_usage_chunks, indent=4, default=str),
@ -232,7 +235,7 @@ async def test_anthropic_streaming_with_headers():
print(f"Attempt {attempt + 1}/{max_retries} to check spend logs")
async with session.get(
f"http://0.0.0.0:4000/spend/logs?request_id={litellm_call_id}",
f"http://0.0.0.0:4000/spend/logs?request_id={anthropic_message_id}",
headers={"Authorization": "Bearer sk-1234"},
) as spend_response:
spend_data = await spend_response.json()
@ -250,25 +253,25 @@ async def test_anthropic_streaming_with_headers():
print("Waiting 10 seconds before retry...")
await asyncio.sleep(10)
# Spend data might be unavailable (auth error, slow DB write, etc.)
if (
spend_data is None
or not isinstance(spend_data, list)
or len(spend_data) == 0
or not isinstance(spend_data[0], dict)
or "request_id" not in spend_data[0]
):
print(f"Spend data not available or is error response: {spend_data}")
print("Skipping spend assertions (DB write may be slow in CI)")
if not isinstance(spend_data, list):
print(f"Spend endpoint answered with an error response: {spend_data}")
print("Skipping spend assertions (spend logs unreachable in CI)")
return
assert spend_data, (
f"GET /spend/logs?request_id={anthropic_message_id} found no row for the id "
"the caller received"
)
log_entry = spend_data[0]
# Basic existence checks
assert isinstance(log_entry, dict), "Log entry should be a dictionary"
# Request metadata assertions
assert log_entry["request_id"] == litellm_call_id, "Request ID should match"
assert (
log_entry["request_id"] == anthropic_message_id
), "Request ID should be the message id the caller received"
assert (
log_entry["call_type"] == "pass_through_endpoint"
), "Call type should be pass_through_endpoint"

View file

@ -18,6 +18,7 @@ import ast
import json
import logging
import os
import signal
import sys
import time
from collections.abc import Callable
@ -47,6 +48,7 @@ FAKE_PRISMA = """#!{python}
import json
import os
import pathlib
import subprocess
import sys
import time
@ -66,6 +68,9 @@ with log_path.open("a") as log:
time.sleep(float(os.environ.get("FAKE_PRISMA_SLEEP", "0")))
if args[:2] == ["migrate", "deploy"]:
if earlier_same_command == 0:
if os.environ.get("FAKE_PRISMA_GRANDCHILD_PIDFILE"):
grandchild = subprocess.Popen([sys.executable, "-c", "import time; time.sleep(600)"])
pathlib.Path(os.environ["FAKE_PRISMA_GRANDCHILD_PIDFILE"]).write_text(str(grandchild.pid))
time.sleep(float(os.environ.get("FAKE_PRISMA_FIRST_DEPLOY_SLEEP", "0")))
elif os.environ.get("FAKE_PRISMA_LATER_DEPLOY_STDERR"):
print(os.environ["FAKE_PRISMA_LATER_DEPLOY_STDERR"], file=sys.stderr)
@ -272,6 +277,43 @@ def test_migrate_deploy_stops_at_its_own_timeout(
assert elapsed < 30
def _process_is_gone(pid: int, within_seconds: float) -> bool:
deadline = time.monotonic() + within_seconds
while time.monotonic() < deadline:
try:
os.kill(pid, 0)
except ProcessLookupError:
return True
time.sleep(0.05)
return False
def test_a_timed_out_migrate_deploy_takes_its_process_tree_with_it(
toolchain_env: tuple[Path, Path], monkeypatch: pytest.MonkeyPatch, tmp_path: Path
) -> None:
"""The real CLI forks Node and a schema engine; a timeout must not leave them running."""
_, log_path = toolchain_env
pidfile = tmp_path / "grandchild.pid"
monkeypatch.setenv("DATABASE_URL", "postgresql://u:p@localhost:9/x")
monkeypatch.setenv(PRISMA_MIGRATE_DEPLOY_TIMEOUT_ENV_VAR, "1")
monkeypatch.setenv("FAKE_PRISMA_FIRST_DEPLOY_SLEEP", "60")
monkeypatch.setenv("FAKE_PRISMA_LATER_DEPLOY_STDERR", "Error: P3018 permission denied for schema public")
monkeypatch.setenv("FAKE_PRISMA_GRANDCHILD_PIDFILE", str(pidfile))
with pytest.raises(RuntimeError, match="insufficient permissions"):
ProxyExtrasDBManager.setup_database(use_migrate=True, use_v2_resolver=True)
grandchild_pid = int(pidfile.read_text())
try:
assert len(_deploy_calls(log_path)) == 2
assert _process_is_gone(grandchild_pid, within_seconds=5)
finally:
try:
os.kill(grandchild_pid, signal.SIGKILL)
except ProcessLookupError:
pass
def test_db_push_timeout_hint_names_the_per_command_budget(
toolchain_env: tuple[Path, Path], monkeypatch: pytest.MonkeyPatch, caplog: pytest.LogCaptureFixture
) -> None:

View file

@ -61,6 +61,137 @@ async def test_dual_cache_async_batch_get_cache_rolls_back_redis_reservation_on_
assert "shared_b" not in dual_cache.last_redis_batch_access_time
def _redis_mock_for_sync_batch(redis_result: dict) -> MagicMock:
mock_redis = MagicMock(spec=RedisCache)
mock_redis.batch_get_cache.return_value = redis_result
return mock_redis
def _assert_sync_batch_used_blocking_client(dual_cache: DualCache, mock_redis: MagicMock) -> None:
with patch("asyncio.new_event_loop", side_effect=AssertionError("sync path must not create an event loop")):
result = dual_cache.batch_get_cache(keys=["lit6729_key"])
assert result == ["redis_value"]
mock_redis.batch_get_cache.assert_called_once_with(key_list=["lit6729_key"], parent_otel_span=None)
mock_redis.async_batch_get_cache.assert_not_called()
mock_redis.init_async_client.assert_not_called()
assert dual_cache.in_memory_cache.get_cache("lit6729_key") == "redis_value"
@pytest.mark.asyncio
async def test_dual_cache_batch_get_cache_uses_sync_redis_client_inside_running_loop():
"""
Regression test for LIT-6729: sync batch_get_cache ran async_batch_get_cache on a
throwaway event loop, reusing an async Redis client created on another loop and
corrupting its connection pool. The sync path must use the blocking client, never
the async one, and never create an event loop, even when called from a coroutine
(e.g. async_raise_no_deployment_exception -> get_min_cooldown).
"""
mock_redis = _redis_mock_for_sync_batch({"lit6729_key": "redis_value"})
dual_cache = DualCache(in_memory_cache=InMemoryCache(), redis_cache=mock_redis)
_assert_sync_batch_used_blocking_client(dual_cache, mock_redis)
def test_dual_cache_batch_get_cache_uses_sync_redis_client_without_running_loop():
mock_redis = _redis_mock_for_sync_batch({"lit6729_key": "redis_value"})
dual_cache = DualCache(in_memory_cache=InMemoryCache(), redis_cache=mock_redis)
_assert_sync_batch_used_blocking_client(dual_cache, mock_redis)
def test_dual_cache_batch_get_cache_only_reads_missing_keys_from_redis():
mock_redis = _redis_mock_for_sync_batch({"miss_key": "from_redis"})
dual_cache = DualCache(in_memory_cache=InMemoryCache(), redis_cache=mock_redis)
dual_cache.in_memory_cache.set_cache("hit_key", "from_memory")
result = dual_cache.batch_get_cache(keys=["hit_key", "miss_key"])
assert result == ["from_memory", "from_redis"]
mock_redis.batch_get_cache.assert_called_once_with(key_list=["miss_key"], parent_otel_span=None)
def test_dual_cache_batch_get_cache_throttles_repeat_redis_reads():
mock_redis = _redis_mock_for_sync_batch({"absent_key": None})
dual_cache = DualCache(
in_memory_cache=InMemoryCache(), redis_cache=mock_redis, default_redis_batch_cache_expiry=10
)
first = dual_cache.batch_get_cache(keys=["absent_key"])
second = dual_cache.batch_get_cache(keys=["absent_key"])
assert first == [None]
assert second == [None]
mock_redis.batch_get_cache.assert_called_once()
def test_dual_cache_batch_get_cache_rolls_back_redis_reservation_on_error():
mock_redis = MagicMock(spec=RedisCache)
mock_redis.batch_get_cache.side_effect = RuntimeError("redis unavailable")
dual_cache = DualCache(
in_memory_cache=InMemoryCache(), redis_cache=mock_redis, default_redis_batch_cache_expiry=10
)
first_result = dual_cache.batch_get_cache(keys=["shared_a"])
second_result = dual_cache.batch_get_cache(keys=["shared_a"])
assert first_result is None
assert second_result is None
assert mock_redis.batch_get_cache.call_count == 2
assert "shared_a" not in dual_cache.last_redis_batch_access_time
def test_dual_cache_batch_get_cache_returns_memory_only_when_redis_read_is_throttled():
mock_redis = _redis_mock_for_sync_batch({"throttled_key": "redis_value"})
dual_cache = DualCache(
in_memory_cache=InMemoryCache(), redis_cache=mock_redis, default_redis_batch_cache_expiry=10
)
dual_cache.last_redis_batch_access_time["throttled_key"] = time.time()
result = dual_cache.batch_get_cache(keys=["throttled_key"])
assert result == [None]
mock_redis.batch_get_cache.assert_not_called()
def test_dual_cache_sync_batch_redis_backfill_injects_default_in_memory_ttl():
"""Sync batch_get_cache's Redis-to-memory backfill must honor
default_in_memory_ttl, same as the async path."""
in_memory_cache = InMemoryCache(default_ttl=600)
mock_redis = _redis_mock_for_sync_batch({"batch_backfill_key": "redis_value"})
dual_cache = DualCache(
in_memory_cache=in_memory_cache,
redis_cache=mock_redis,
default_in_memory_ttl=60,
)
before = time.time()
result = dual_cache.batch_get_cache(keys=["batch_backfill_key"])
after = time.time()
assert result == ["redis_value"]
expiry = in_memory_cache.ttl_dict["batch_backfill_key"]
assert expiry >= before + 60
assert expiry <= after + 60
def test_dual_cache_batch_get_cache_forwards_explicit_ttl_to_backfill():
"""An explicit ttl kwarg must reach the in-memory backfill flat, not nested
under a 'kwargs' key the way the old locals()-forwarding path sent it."""
in_memory_cache = InMemoryCache(default_ttl=600)
mock_redis = _redis_mock_for_sync_batch({"explicit_ttl_key": "redis_value"})
dual_cache = DualCache(in_memory_cache=in_memory_cache, redis_cache=mock_redis)
before = time.time()
result = dual_cache.batch_get_cache(keys=["explicit_ttl_key"], ttl=5)
after = time.time()
assert result == ["redis_value"]
expiry = in_memory_cache.ttl_dict["explicit_ttl_key"]
assert expiry >= before + 5
assert expiry <= after + 5
@pytest.mark.asyncio
async def test_dual_cache_async_set_cache_injects_default_in_memory_ttl():
"""

View file

@ -1,10 +1,10 @@
import asyncio
from unittest.mock import MagicMock, patch
from collections.abc import Iterator
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from unittest.mock import AsyncMock
from litellm._service_logger import ServiceLogging
from litellm.caching.redis_cache import RedisCache
@ -17,6 +17,17 @@ def redis_no_ping():
yield
@pytest.fixture
def sync_batch_redis_cache(redis_no_ping):
with patch( # test-quality-ok: RedisCache.__init__ builds its client eagerly, with no injection point
"litellm._redis.get_redis_client", return_value=MagicMock()
) as get_client:
cache = RedisCache(host="127.0.0.1", port=6379)
cache.redis_client.mget.side_effect = OSError("redis unavailable")
get_client.assert_called_once()
yield cache
@pytest.mark.parametrize(
("namespace", "key", "expected"),
[
@ -504,6 +515,173 @@ async def test_circuit_breaker_opens_when_method_swallows_redis_failure(redis_no
await call_method(cache)
def test_circuit_breaker_open_keeps_sync_batch_get_cache_as_a_miss(sync_batch_redis_cache):
"""An open breaker must preserve the sync batch read's dictionary fallback."""
from litellm.constants import REDIS_CIRCUIT_BREAKER_FAILURE_THRESHOLD
for _ in range(REDIS_CIRCUIT_BREAKER_FAILURE_THRESHOLD):
assert sync_batch_redis_cache.batch_get_cache(key_list=["lit6729"]) == {}
assert sync_batch_redis_cache.batch_get_cache(key_list=["lit6729"]) == {}
@pytest.fixture
def sync_batch_cache_with_service_logger(redis_no_ping: None) -> Iterator[tuple[RedisCache, ServiceLogging]]:
service_logger = ServiceLogging(mock_testing=True)
failing_client = MagicMock()
failing_client.mget.side_effect = OSError("redis unavailable")
with patch( # test-quality-ok: RedisCache.__init__ builds its client eagerly, with no injection point
"litellm._redis.get_redis_client", return_value=failing_client
):
cache = RedisCache(host="127.0.0.1", port=6379, service_logger_obj=service_logger)
yield cache, service_logger
@pytest.mark.asyncio
async def test_sync_batch_get_cache_reports_a_failed_read_from_a_running_loop(
sync_batch_cache_with_service_logger: tuple[RedisCache, ServiceLogging],
):
"""A swallowed Redis failure must still be reported as a service failure event.
The routing strategies call this blocking read from inside the request's event loop,
and the read hides the Redis error by returning an empty dict. Without an emitted
failure event, litellm_redis_failed_requests_total stops moving during a Redis
outage while the success path keeps reporting, so the dashboards read healthy.
"""
cache, service_logger = sync_batch_cache_with_service_logger
assert cache.batch_get_cache(key_list=["lit6729"]) == {}
await asyncio.sleep(0.05)
assert service_logger.mock_testing_sync_failure_hook == 1
assert service_logger.mock_testing_async_failure_hook == 1
def test_sync_batch_get_cache_reports_a_failed_read_from_a_worker_thread(
sync_batch_cache_with_service_logger: tuple[RedisCache, ServiceLogging],
):
"""The same report must reach the async hook when the caller has no event loop at all."""
from concurrent.futures import ThreadPoolExecutor
cache, service_logger = sync_batch_cache_with_service_logger
with ThreadPoolExecutor(max_workers=1) as pool:
assert pool.submit(cache.batch_get_cache, key_list=["lit6729"]).result() == {}
assert service_logger.mock_testing_async_failure_hook == 1
def test_sync_batch_get_cache_reports_a_failed_read_on_an_idle_event_loop(
sync_batch_cache_with_service_logger: tuple[RedisCache, ServiceLogging],
):
"""The report must also go out when the caller holds an open loop that is not running."""
cache, service_logger = sync_batch_cache_with_service_logger
loop = asyncio.new_event_loop()
try:
asyncio.set_event_loop(loop)
assert cache.batch_get_cache(key_list=["lit6729"]) == {}
finally:
asyncio.set_event_loop(None)
loop.close()
assert service_logger.mock_testing_async_failure_hook == 1
def test_sync_batch_get_cache_survives_a_service_callback_that_raises(
sync_batch_cache_with_service_logger: tuple[RedisCache, ServiceLogging],
monkeypatch: pytest.MonkeyPatch,
):
"""A failing service callback must not replace the swallowed Redis failure.
A misconfigured callback raises while emitting (a datadog callback with no
DD_API_KEY raises at construction), and the failure event is emitted from inside
the except block that swallows the Redis error. If that exception escapes, a Redis
outage surfaces to routing as a callback error and the circuit breaker never
records the failed read.
"""
from concurrent.futures import ThreadPoolExecutor
import litellm
from litellm.constants import REDIS_CIRCUIT_BREAKER_FAILURE_THRESHOLD
cache, service_logger = sync_batch_cache_with_service_logger
monkeypatch.setattr(litellm, "service_callback", ["prometheus_system"])
monkeypatch.setattr(
service_logger,
"init_prometheus_services_logger_if_none",
AsyncMock(side_effect=Exception("callback is misconfigured")),
)
for _ in range(REDIS_CIRCUIT_BREAKER_FAILURE_THRESHOLD):
with ThreadPoolExecutor(max_workers=1) as pool:
assert pool.submit(cache.batch_get_cache, key_list=["lit6729"]).result() == {}
assert cache.batch_get_cache(key_list=["lit6729"]) == {}
def test_call_stack_info_skips_breaker_guard_frames():
"""Guarded methods must still report their real callers in service-log call_type.
The breaker guards put their own frames between a method body and its caller, so
without skipping them every guarded method logged the guard machinery instead of
who actually issued the Redis call.
"""
from litellm.caching.redis_cache import (
RedisCircuitBreaker,
_get_call_stack_info,
_redis_circuit_breaker_guard_sync,
)
class Guarded:
_circuit_breaker = RedisCircuitBreaker(failure_threshold=3, recovery_timeout=60)
@_redis_circuit_breaker_guard_sync
def probe(self):
return _get_call_stack_info()
def caller_one():
return Guarded().probe()
def caller_two():
return caller_one()
assert caller_two() == "caller_one <- caller_two"
def test_call_stack_info_skips_guard_frames_when_deployed_without_sources(monkeypatch):
"""Guard-frame skipping must survive a bytecode-only deployment.
Shipping `.pyc` files without their `.py` sources leaves the module's `__file__` pointing
at the compiled file while every frame still carries the compile-time source path, so a
check comparing those two paths stops skipping and the service log then names the guard
machinery instead of the real caller.
"""
from litellm.caching import redis_cache as redis_cache_module
from litellm.caching.redis_cache import (
RedisCircuitBreaker,
_get_call_stack_info,
_redis_circuit_breaker_guard_sync,
)
monkeypatch.setattr(redis_cache_module, "__file__", redis_cache_module.__file__ + "c")
class Guarded:
_circuit_breaker = RedisCircuitBreaker(failure_threshold=3, recovery_timeout=60)
@_redis_circuit_breaker_guard_sync
def probe(self):
return _get_call_stack_info()
def caller_one():
return Guarded().probe()
def caller_two():
return caller_one()
assert caller_two() == "caller_one <- caller_two"
@pytest.mark.asyncio
async def test_circuit_breaker_success_still_resets_the_failure_streak(redis_no_ping):
"""A reachable Redis must keep the breaker closed, however many earlier calls failed.
@ -580,7 +758,6 @@ async def test_concurrent_success_is_not_cancelled_by_another_calls_failure():
async def swallows_a_failure():
await asyncio.sleep(0.02)
_record_swallowed_redis_failure(breaker, RedisConnectionError("redis unreachable"))
return None
async def succeeds_while_the_other_fails():
await asyncio.sleep(0.05)

View file

@ -1,7 +1,7 @@
import json
import sys
from types import SimpleNamespace
from unittest.mock import AsyncMock
from unittest.mock import AsyncMock, MagicMock
import pytest
from fastapi import HTTPException
@ -242,110 +242,156 @@ async def test_should_not_reassign_existing_container_to_different_owner(monkeyp
table.update.assert_not_awaited()
@pytest.mark.asyncio
async def test_should_filter_container_list_to_owned_records(monkeypatch):
def _owned_containers_in_db(monkeypatch, *model_object_ids: str) -> AsyncMock:
table = AsyncMock()
table.find_many.return_value = [
SimpleNamespace(model_object_id="container:openai:cntr_owned"),
]
prisma_client = SimpleNamespace(
db=SimpleNamespace(litellm_managedobjecttable=table)
)
table.find_many.return_value = [SimpleNamespace(model_object_id=object_id) for object_id in model_object_ids]
monkeypatch.setattr(
ownership,
"_get_prisma_client",
AsyncMock(return_value=prisma_client),
AsyncMock(return_value=SimpleNamespace(db=SimpleNamespace(litellm_managedobjecttable=table))),
)
auth = UserAPIKeyAuth(user_id="user-1")
response = ContainerListResponse(
return table
def _upstream(pages_by_after):
calls = []
async def fetch_page(after, limit):
calls.append((after, limit))
return pages_by_after[after]
return fetch_page, calls
def _page(*container_ids: str, has_more: bool) -> ContainerListResponse:
return ContainerListResponse(
object="list",
data=[_container("cntr_owned"), _container("cntr_other")],
has_more=True,
data=[_container(container_id) for container_id in container_ids],
has_more=has_more,
)
filtered = await ownership.filter_container_list_response(
response=response,
user_api_key_dict=auth,
async def _list_owned(fetch_page, after=None, limit=None):
return await ownership.list_owned_containers(
fetch_page=fetch_page,
after=after,
limit=limit,
user_api_key_dict=UserAPIKeyAuth(user_id="user-1"),
custom_llm_provider="openai",
)
assert [item.id for item in filtered.data] == ["cntr_owned"]
assert filtered.first_id == "cntr_owned"
assert filtered.last_id == "cntr_owned"
assert filtered.has_more is False
@pytest.mark.asyncio
async def test_should_page_upstream_until_owned_containers_fill_the_limit(monkeypatch):
table = _owned_containers_in_db(monkeypatch, "container:openai:cntr_owned")
fetch_page, calls = _upstream(
{
None: _page("cntr_other_1", "cntr_other_2", has_more=True),
"cntr_other_2": _page("cntr_owned", has_more=False),
}
)
listed = await _list_owned(fetch_page, limit=1)
assert [item.id for item in listed.data] == ["cntr_owned"]
assert listed.first_id == "cntr_owned"
assert listed.last_id == "cntr_owned"
assert listed.has_more is False
assert calls == [(None, 100), ("cntr_other_2", 100)]
where = table.find_many.await_args.kwargs["where"]
assert where["file_purpose"] == ownership.CONTAINER_OBJECT_PURPOSE
assert where["created_by"]["in"] == ["user-1", "user:user-1"]
@pytest.mark.asyncio
async def test_should_clear_has_more_when_filtered_container_list_is_empty(
monkeypatch,
):
table = AsyncMock()
table.find_many.return_value = [
SimpleNamespace(model_object_id="container:openai:cntr_owned"),
]
prisma_client = SimpleNamespace(
db=SimpleNamespace(litellm_managedobjecttable=table)
)
monkeypatch.setattr(
ownership,
"_get_prisma_client",
AsyncMock(return_value=prisma_client),
)
auth = UserAPIKeyAuth(user_id="user-1")
response = ContainerListResponse(
object="list",
data=[_container("cntr_other")],
has_more=True,
)
async def test_should_trim_owned_containers_to_the_limit_without_mutating_the_upstream_page(monkeypatch):
_owned_containers_in_db(monkeypatch, "container:openai:cntr_owned_1", "container:openai:cntr_owned_2")
upstream_page = _page("cntr_owned_1", "cntr_other", "cntr_owned_2", has_more=False)
fetch_page, calls = _upstream({None: upstream_page})
filtered = await ownership.filter_container_list_response(
response=response,
user_api_key_dict=auth,
custom_llm_provider="openai",
)
listed = await _list_owned(fetch_page, limit=1)
assert filtered.data == []
assert filtered.first_id is None
assert filtered.last_id is None
assert filtered.has_more is False
assert [item.id for item in listed.data] == ["cntr_owned_1"]
assert listed.first_id == "cntr_owned_1"
assert listed.last_id == "cntr_owned_1"
assert listed.has_more is True
assert calls == [(None, 100)]
assert [item.id for item in upstream_page.data] == ["cntr_owned_1", "cntr_other", "cntr_owned_2"]
assert upstream_page.has_more is False
@pytest.mark.asyncio
async def test_should_clear_dict_has_more_when_filtered_container_list_is_empty(
monkeypatch,
):
table = AsyncMock()
table.find_many.return_value = [
SimpleNamespace(model_object_id="container:openai:cntr_owned"),
]
prisma_client = SimpleNamespace(
db=SimpleNamespace(litellm_managedobjecttable=table)
async def test_should_start_paging_from_the_requested_cursor(monkeypatch):
_owned_containers_in_db(monkeypatch, "container:openai:cntr_owned_2")
fetch_page, calls = _upstream({"cntr_owned_1": _page("cntr_other", "cntr_owned_2", has_more=False)})
listed = await _list_owned(fetch_page, after="cntr_owned_1", limit=1)
assert [item.id for item in listed.data] == ["cntr_owned_2"]
assert listed.has_more is False
assert calls == [("cntr_owned_1", 100)]
@pytest.mark.asyncio
async def test_should_default_to_twenty_owned_containers_per_page(monkeypatch):
owned_ids = tuple(f"cntr_owned_{index}" for index in range(21))
_owned_containers_in_db(monkeypatch, *(f"container:openai:{container_id}" for container_id in owned_ids))
fetch_page, _ = _upstream({None: _page(*owned_ids, has_more=False)})
listed = await _list_owned(fetch_page)
assert [item.id for item in listed.data] == list(owned_ids[:20])
assert listed.last_id == "cntr_owned_19"
assert listed.has_more is True
@pytest.mark.asyncio
async def test_should_stop_after_five_upstream_pages_and_keep_has_more(monkeypatch):
_owned_containers_in_db(monkeypatch, "container:openai:cntr_owned")
fetch_page, calls = _upstream(
{
None: _page("cntr_other_0", has_more=True),
**{f"cntr_other_{index}": _page(f"cntr_other_{index + 1}", has_more=True) for index in range(6)},
}
)
monkeypatch.setattr(
ownership,
"_get_prisma_client",
AsyncMock(return_value=prisma_client),
)
auth = UserAPIKeyAuth(user_id="user-1")
response = {
listed = await _list_owned(fetch_page, limit=1)
assert listed.data == []
assert listed.first_id is None
assert listed.last_id is None
assert listed.has_more is True
assert len(calls) == 5
@pytest.mark.asyncio
async def test_should_stop_when_upstream_has_no_more_pages(monkeypatch):
_owned_containers_in_db(monkeypatch, "container:openai:cntr_owned")
fetch_page, calls = _upstream({None: _page("cntr_other", has_more=False)})
listed = await _list_owned(fetch_page, limit=1)
assert listed.data == []
assert listed.has_more is False
assert calls == [(None, 100)]
@pytest.mark.asyncio
async def test_should_build_dict_pages_without_mutating_the_upstream_page(monkeypatch):
_owned_containers_in_db(monkeypatch, "container:openai:cntr_owned")
upstream_page = {"object": "list", "data": [{"id": "cntr_other"}, {"id": "cntr_owned"}], "has_more": False}
fetch_page, _ = _upstream({None: upstream_page})
listed = await _list_owned(fetch_page, limit=1)
assert listed == {
"object": "list",
"data": [{"id": "cntr_other"}],
"has_more": True,
"data": [{"id": "cntr_owned"}],
"first_id": "cntr_owned",
"last_id": "cntr_owned",
"has_more": False,
}
filtered = await ownership.filter_container_list_response(
response=response,
user_api_key_dict=auth,
custom_llm_provider="openai",
)
assert filtered["data"] == []
assert filtered["first_id"] is None
assert filtered["last_id"] is None
assert filtered["has_more"] is False
assert [item["id"] for item in upstream_page["data"]] == ["cntr_other", "cntr_owned"]
@pytest.mark.asyncio
@ -647,7 +693,7 @@ async def test_should_return_response_when_owner_recording_raises_unexpected(
@pytest.mark.asyncio
async def test_should_filter_container_list_inside_list_endpoint(monkeypatch):
async def test_should_list_owned_containers_inside_list_endpoint(monkeypatch):
from litellm.proxy.container_endpoints import endpoints
proxy_server_stub = SimpleNamespace(
@ -665,42 +711,37 @@ async def test_should_filter_container_list_inside_list_endpoint(monkeypatch):
)
monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", proxy_server_stub)
response = ContainerListResponse(
object="list",
data=[_container("cntr_provider")],
has_more=False,
)
class FakeProcessor:
def __init__(self, data):
pass
async def base_process_llm_request(self, **kwargs):
return response
async def _handle_llm_api_exception(self, **kwargs):
raise kwargs["e"]
filter_response = AsyncMock(return_value=response)
monkeypatch.setattr(endpoints, "ProxyBaseLLMRequestProcessing", FakeProcessor)
monkeypatch.setattr(
endpoints,
"filter_container_list_response",
filter_response,
upstream_page = _page("cntr_provider", has_more=False)
processor_cls = MagicMock(
side_effect=lambda data: SimpleNamespace(base_process_llm_request=AsyncMock(return_value=upstream_page))
)
monkeypatch.setattr(endpoints, "ProxyBaseLLMRequestProcessing", processor_cls)
list_owned = AsyncMock(return_value=upstream_page)
monkeypatch.setattr(endpoints, "list_owned_containers", list_owned)
result = await endpoints.list_containers(
request=SimpleNamespace(query_params={}, headers={}),
fastapi_response=SimpleNamespace(),
user_api_key_dict=UserAPIKeyAuth(user_id="user-1"),
after="cntr_prev",
limit=2,
order="desc",
)
assert result == response
filter_response.assert_awaited_once_with(
response=response,
user_api_key_dict=UserAPIKeyAuth(user_id="user-1"),
custom_llm_provider="openai",
)
assert result == upstream_page
kwargs = list_owned.await_args.kwargs
assert kwargs["after"] == "cntr_prev"
assert kwargs["limit"] == 2
assert kwargs["user_api_key_dict"] == UserAPIKeyAuth(user_id="user-1")
assert kwargs["custom_llm_provider"] == "openai"
processor_cls.assert_not_called()
assert await kwargs["fetch_page"]("cntr_page_cursor", 100) == upstream_page
forwarded = processor_cls.call_args.kwargs["data"]
assert forwarded["after"] == "cntr_page_cursor"
assert forwarded["limit"] == 100
assert forwarded["order"] == "desc"
assert forwarded["custom_llm_provider"] == "openai"
@pytest.mark.asyncio

View file

@ -1091,8 +1091,8 @@ class TestCustomGuardrailPassthroughSupport:
call_type=CallTypes.allm_passthrough_route,
)
# When result is None, should return the original response
assert result == mock_response
# None means the guardrail did not modify the response (LIT-5863 contract)
assert result is None
@pytest.mark.asyncio
async def test_async_post_call_success_deployment_hook_with_none_call_type(self):
@ -1120,8 +1120,8 @@ class TestCustomGuardrailPassthroughSupport:
call_type=None,
)
# Should return the original response when result is None
assert result == mock_response
# None means the guardrail did not modify the response (LIT-5863 contract)
assert result is None
def test_is_valid_response_type_with_none(self):
"""
@ -2436,3 +2436,73 @@ class TestLoggingOnlyApplyGuardrail:
assert guardrail.calls == [("request", ["hello there"]), ("response", ["general kenobi"])]
entries = logging_obj.model_call_details["standard_logging_object"]["guardrail_information"]
assert [e["guardrail_status"] for e in entries] == ["success", "success"]
class TestCustomGuardrailPostCallSuccessDeploymentHook:
"""Regression tests for LIT-5863: this hook answering the unmodified response instead of
None made the utils.py dispatcher treat the guardrail as having modified the response,
which starved every later callback in litellm.callbacks (notably the lazily-appended
VectorStorePreCallHook that attaches provider_specific_fields["search_results"])."""
@pytest.mark.asyncio
async def test_returns_none_when_request_has_no_guardrails(self):
from litellm.types.utils import ModelResponse
guardrail = CustomGuardrail(guardrail_name="test-guardrail")
response = ModelResponse()
assert (
await guardrail.async_post_call_success_deployment_hook(
request_data={}, response=response, call_type=CallTypes.acompletion
)
is None
)
assert (
await guardrail.async_post_call_success_deployment_hook(
request_data={"guardrails": "not-a-list"}, response=response, call_type=CallTypes.acompletion
)
is None
)
@pytest.mark.asyncio
async def test_returns_none_when_guardrail_should_not_run(self):
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.utils import ModelResponse
guardrail = CustomGuardrail(
guardrail_name="test-guardrail",
event_hook=GuardrailEventHooks.pre_call,
)
response = ModelResponse()
result = await guardrail.async_post_call_success_deployment_hook(
request_data={"guardrails": ["test-guardrail"]},
response=response,
call_type=CallTypes.acompletion,
)
assert result is None
@pytest.mark.asyncio
async def test_returns_modified_response_when_guardrail_runs(self):
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.utils import ModelResponse
replacement = ModelResponse()
class ReplacingGuardrail(CustomGuardrail):
async def async_post_call_success_hook(self, data, user_api_key_dict, response):
return replacement
guardrail = ReplacingGuardrail(
guardrail_name="test-guardrail",
event_hook=GuardrailEventHooks.post_call,
)
result = await guardrail.async_post_call_success_deployment_hook(
request_data={"guardrails": ["test-guardrail"]},
response=ModelResponse(),
call_type=CallTypes.acompletion,
)
assert result is replacement

View file

@ -0,0 +1,287 @@
import logging
from collections.abc import Iterator
from dataclasses import dataclass, field
from typing import Protocol
import pytest
import litellm
from litellm._logging import verbose_logger
from litellm.integrations.vector_store_integrations.vector_store_pre_call_hook import (
ProxyServerRuntime,
VectorStorePreCallHook,
)
from litellm.types.llms.openai import AllMessageValues
from litellm.types.vector_stores import (
VectorStoreResultContent,
VectorStoreSearchResponse,
VectorStoreSearchResult,
)
from litellm.vector_stores.vector_store_registry import (
LiteLLM_ManagedVectorStore,
VectorStoreRegistry,
)
def _search_response(text: str) -> VectorStoreSearchResponse:
return VectorStoreSearchResponse(
object="vector_store.search_results.page",
search_query="what is litellm?",
data=[
VectorStoreSearchResult(
score=1.0,
content=[VectorStoreResultContent(text=text, type="text")],
)
],
)
@dataclass
class RecordingRouter:
failing_vector_store_ids: frozenset[str] = frozenset()
calls: list[dict[str, object]] = field(default_factory=list)
async def avector_store_search(self, **kwargs: object) -> VectorStoreSearchResponse:
self.calls.append(kwargs)
vector_store_id = str(kwargs["vector_store_id"])
if vector_store_id in self.failing_vector_store_ids:
raise litellm.BadRequestError(
message=f"no healthy deployments for {vector_store_id}",
model="text-embedding-3-small",
llm_provider="openai",
)
return _search_response(f"context from {vector_store_id}")
@dataclass(frozen=True)
class FakeProxyRuntime:
router: RecordingRouter | None
def llm_router(self) -> RecordingRouter | None:
return self.router
def prisma_client(self) -> None:
return None
class RecordingHandler(logging.Handler):
def __init__(self) -> None:
super().__init__(level=logging.WARNING)
self.records: list[logging.LogRecord] = []
def emit(self, record: logging.LogRecord) -> None:
self.records.append(record)
class RegisterStores(Protocol):
def __call__(self, *vector_store_ids: str, custom_llm_provider: str = "bedrock") -> None: ...
@pytest.fixture
def registry_with(monkeypatch: pytest.MonkeyPatch) -> RegisterStores:
def _register(*vector_store_ids: str, custom_llm_provider: str = "bedrock") -> None:
monkeypatch.setattr(
litellm,
"vector_store_registry",
VectorStoreRegistry(
vector_stores=[
LiteLLM_ManagedVectorStore(vector_store_id=vector_store_id, custom_llm_provider=custom_llm_provider)
for vector_store_id in vector_store_ids
],
),
)
return _register
@pytest.fixture
def warnings() -> Iterator[list[logging.LogRecord]]:
handler = RecordingHandler()
verbose_logger.addHandler(handler)
yield handler.records
verbose_logger.removeHandler(handler)
class FakeLoggingObj:
def __init__(self, metadata: dict[str, str]) -> None:
self.model_call_details: dict[str, object] = {"litellm_params": {"metadata": metadata}}
async def _run_hook(
hook: VectorStorePreCallHook,
vector_store_ids: list[str],
logging_obj: FakeLoggingObj,
) -> tuple[str, list[AllMessageValues], dict[str, object]]:
return await hook.async_get_chat_completion_prompt(
model="chat-model",
messages=[{"role": "user", "content": "what is litellm?"}],
non_default_params={"vector_store_ids": vector_store_ids},
prompt_id=None,
prompt_variables=None,
dynamic_callback_params={},
litellm_logging_obj=logging_obj,
)
@pytest.mark.asyncio
async def test_hook_searches_through_the_injected_router_with_the_request_metadata(
registry_with: RegisterStores,
) -> None:
"""Regression (LIT-6752): the hook must reach the Router through its injected runtime, not a proxy_server import."""
registry_with("vs-router")
router = RecordingRouter()
logging_obj = FakeLoggingObj({"user_api_key_team_id": "team-a"})
_, messages, _ = await _run_hook(
VectorStorePreCallHook(proxy_runtime=FakeProxyRuntime(router=router)),
["vs-router"],
logging_obj,
)
assert router.calls == [
{
"vector_store_id": "vs-router",
"query": "what is litellm?",
"custom_llm_provider": "bedrock",
"metadata": {"user_api_key_team_id": "team-a"},
}
]
assert messages[0]["content"] == "Context:\n\ncontext from vs-router\n\n"
@pytest.mark.asyncio
async def test_hook_falls_back_to_the_sdk_when_the_runtime_has_no_router(
registry_with: RegisterStores,
warnings: list[logging.LogRecord],
) -> None:
registry_with("vs-sdk", custom_llm_provider="lit6752-not-a-provider")
_, messages, _ = await _run_hook(
VectorStorePreCallHook(proxy_runtime=FakeProxyRuntime(router=None)),
["vs-sdk"],
FakeLoggingObj({"user_api_key_team_id": "team-a"}),
)
assert messages == [{"role": "user", "content": "what is litellm?"}]
assert len(warnings) == 1
assert (
warnings[0]
.getMessage()
.startswith("Vector store search failed for vector_store_id=vs-sdk, continuing without its context: ")
)
assert "is not a valid LlmProviders" in warnings[0].getMessage()
@pytest.mark.asyncio
async def test_every_healthy_vector_store_contributes_its_own_context(registry_with: RegisterStores) -> None:
"""Regression (LIT-6752): each store appended its context to the original messages, so only the last one survived."""
registry_with("vs-one", "vs-two")
router = RecordingRouter()
_, messages, _ = await _run_hook(
VectorStorePreCallHook(proxy_runtime=FakeProxyRuntime(router=router)),
["vs-one", "vs-two"],
FakeLoggingObj({}),
)
assert [message["content"] for message in messages] == [
"Context:\n\ncontext from vs-one\n\n",
"Context:\n\ncontext from vs-two\n\n",
"what is litellm?",
]
@pytest.mark.asyncio
async def test_a_failing_vector_store_warns_with_its_id_and_the_other_stores_still_answer(
registry_with: RegisterStores,
warnings: list[logging.LogRecord],
) -> None:
"""Regression (LIT-6752): one unreachable store must not silently drop every other store's context."""
registry_with("vs-broken", "vs-healthy")
router = RecordingRouter(failing_vector_store_ids=frozenset({"vs-broken"}))
logging_obj = FakeLoggingObj({"user_api_key_team_id": "team-a"})
_, messages, _ = await _run_hook(
VectorStorePreCallHook(proxy_runtime=FakeProxyRuntime(router=router)),
["vs-broken", "vs-healthy"],
logging_obj,
)
search_results = logging_obj.model_call_details["search_results"]
assert [call["vector_store_id"] for call in router.calls] == ["vs-broken", "vs-healthy"]
assert messages[0]["content"] == "Context:\n\ncontext from vs-healthy\n\n"
assert isinstance(search_results, list)
assert len(search_results) == 1
assert [record.getMessage() for record in warnings] == [
"Vector store search failed for vector_store_id=vs-broken, continuing without its context: "
"litellm.BadRequestError: no healthy deployments for vs-broken"
]
@pytest.mark.asyncio
async def test_the_only_vector_store_failing_leaves_the_messages_untouched(
registry_with: RegisterStores,
warnings: list[logging.LogRecord],
) -> None:
registry_with("vs-broken")
original_messages = [{"role": "user", "content": "what is litellm?"}]
_, messages, _ = await _run_hook(
VectorStorePreCallHook(
proxy_runtime=FakeProxyRuntime(router=RecordingRouter(failing_vector_store_ids=frozenset({"vs-broken"})))
),
["vs-broken"],
FakeLoggingObj({}),
)
assert messages == original_messages
assert [(record.levelname, record.getMessage()) for record in warnings] == [
(
"WARNING",
"Vector store search failed for vector_store_id=vs-broken, continuing without its context: "
"litellm.BadRequestError: no healthy deployments for vs-broken",
)
]
@pytest.mark.asyncio
async def test_the_default_hook_reaches_the_proxy_router_through_its_runtime(
registry_with: RegisterStores,
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""Regression (LIT-6752): a hook built with no arguments must still search through the proxy's own Router."""
from litellm.proxy import proxy_server
registry_with("vs-default")
router = RecordingRouter()
monkeypatch.setattr(proxy_server, "llm_router", router)
monkeypatch.setattr(proxy_server, "prisma_client", None)
_, messages, _ = await _run_hook(
VectorStorePreCallHook(),
["vs-default"],
FakeLoggingObj({"user_api_key_team_id": "team-a"}),
)
assert [call["vector_store_id"] for call in router.calls] == ["vs-default"]
assert messages[0]["content"] == "Context:\n\ncontext from vs-default\n\n"
def test_the_default_runtime_follows_the_proxy_globals(monkeypatch: pytest.MonkeyPatch) -> None:
from litellm.proxy import proxy_server
runtime = ProxyServerRuntime()
monkeypatch.setattr(proxy_server, "llm_router", None)
monkeypatch.setattr(proxy_server, "prisma_client", None)
assert runtime.llm_router() is None
assert runtime.prisma_client() is None
router = RecordingRouter()
prisma = object()
monkeypatch.setattr(proxy_server, "llm_router", router)
monkeypatch.setattr(proxy_server, "prisma_client", prisma)
assert runtime.llm_router() is router
assert runtime.prisma_client() is prisma

View file

@ -312,3 +312,121 @@ def test_mask_credentials_in_payload_masks_only_sensitive_string_leaves():
assert masked != plaintext
assert masked.startswith(plaintext[:4])
assert masked.endswith(plaintext[-4:])
def test_redact_credentials_in_payload_leaves_no_fragment_of_the_secret():
"""A payload rendered straight to stdout cannot afford the partial reveal
mask_credentials_in_payload leaves, so every credential-named value is replaced
whole, nested header dicts included, while ordinary params survive verbatim."""
from litellm.litellm_core_utils.sensitive_data_masker import redact_credentials_in_payload
fake_key = "sk-fake-lit6823-0000000000000000"
fake_token = "fake-azure-ad-token-0000"
result = redact_credentials_in_payload(
{
"api_key": fake_key,
"azure_ad_token": fake_token,
"aws_secret_access_key": "fake-aws-secret-0000",
"vertex_credentials": {"private_key": "fake-pem"},
"extra_headers": {"Authorization": "Bearer fake-bearer-0000", "x-request-id": "abc123"},
"model": "gpt-4o-mini",
"max_tokens": 17,
"temperature": 0.25,
"api_base": None,
}
)
assert fake_key not in str(result)
assert fake_token not in str(result)
assert "fake-aws-secret-0000" not in str(result)
assert "fake-pem" not in str(result)
assert "fake-bearer-0000" not in str(result)
assert result["api_key"] == "REDACTED"
assert result["extra_headers"]["Authorization"] == "REDACTED"
assert result["extra_headers"]["x-request-id"] == "abc123"
assert result["model"] == "gpt-4o-mini"
assert result["max_tokens"] == 17
assert result["temperature"] == 0.25
assert result["api_base"] is None
def test_redact_credentials_in_payload_reaches_credentials_nested_in_sequences():
"""Free-form kwargs like extra_body and metadata routinely carry lists of dicts, so a
credential hiding one level inside a list or tuple must be replaced too, while the
surrounding container keeps its type and every ordinary element stays verbatim."""
from litellm.litellm_core_utils.sensitive_data_masker import redact_credentials_in_payload
result = redact_credentials_in_payload(
{
"extra_body": {"providers": [{"name": "openai", "api_key": "sk-fake-lit6823-in-a-list"}]},
"metadata": {"upstreams": ({"aws_secret_access_key": "fake-aws-in-a-tuple"},)},
"messages": [{"role": "user", "content": "hello"}],
}
)
assert "sk-fake-lit6823-in-a-list" not in str(result)
assert "fake-aws-in-a-tuple" not in str(result)
assert result["extra_body"]["providers"][0]["api_key"] == "REDACTED"
assert result["extra_body"]["providers"][0]["name"] == "openai"
assert isinstance(result["extra_body"]["providers"], list)
assert result["metadata"]["upstreams"][0]["aws_secret_access_key"] == "REDACTED"
assert isinstance(result["metadata"]["upstreams"], tuple)
assert result["messages"] == [{"role": "user", "content": "hello"}]
@pytest.mark.parametrize("wrap", ["mapping", "sequence"])
def test_redact_credentials_in_payload_hides_containers_at_the_recursion_limit(wrap):
"""The recursion limit exists to bound the walk, not to grant an exemption, so a caller who
buries a credential deeper than the limit must get the container hidden rather than handed
back verbatim. Nesting through lists costs depth twice as fast as nesting through mappings,
so both shapes are pushed well past the limit here."""
from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH
from litellm.litellm_core_utils.sensitive_data_masker import redact_credentials_in_payload
fake_key = "sk-fake-lit6835-past-the-limit"
node = {"api_key": fake_key}
for _ in range(2 * DEFAULT_MAX_RECURSE_DEPTH + 1):
node = {"extra_body": node} if wrap == "mapping" else {"providers": [node]}
result = redact_credentials_in_payload({**node, "max_tokens": 17})
assert fake_key not in str(result)
assert "REDACTED" in str(result)
assert result["max_tokens"] == 17
def test_redact_credentials_in_payload_leaves_a_realistic_tool_schema_intact():
"""The bound must not eat ordinary payloads: a tool whose JSON schema nests an array of
objects inside a nested object is what agent traffic looks like, and the verbose line is
useless if those leaves come back as REDACTED."""
from litellm.litellm_core_utils.sensitive_data_masker import redact_credentials_in_payload
tool = {
"type": "function",
"function": {
"name": "search_orders",
"parameters": {
"type": "object",
"properties": {
"filters": {
"type": "object",
"properties": {
"items": {
"type": "array",
"items": {
"type": "object",
"properties": {"sku": {"type": "string"}, "qty": {"type": "integer"}},
},
}
},
}
},
},
},
}
result = redact_credentials_in_payload({"model": "gpt-4o-mini", "tools": [tool], "api_key": "sk-fake-lit6835"})
assert "REDACTED" not in str(result["tools"])
assert result["tools"][0] == tool
assert result["api_key"] == "REDACTED"

View file

@ -0,0 +1,111 @@
"""
Streaming ``/v1/messages`` against a model that is neither Anthropic nor OpenAI is served by
translating the call onto ``/v1/chat/completions``, and the ``msg_`` id the caller is streamed
is minted right here. It is the only request id such a caller ever sees, so the spend row has
to be keyed on that same value rather than on the provider's own completion id.
"""
import datetime
import json
import pytest
import respx
import litellm
from litellm.llms.anthropic.experimental_pass_through.adapters.handler import (
LiteLLMMessagesToCompletionTransformationHandler,
)
from litellm.llms.anthropic.experimental_pass_through.adapters.streaming_iterator import (
AnthropicStreamWrapper,
)
MESSAGES = [{"role": "user", "content": "hello"}]
GROQ_CHAT_URL = "https://api.groq.com/openai/v1/chat/completions"
CHAT_SSE_BODY = (
b'data: {"id":"chatcmpl-lit6825","object":"chat.completion.chunk","created":1,'
b'"model":"kimi-k2","choices":[{"index":0,"delta":{"role":"assistant","content":"hi"},'
b'"finish_reason":null}]}\n\n'
b'data: {"id":"chatcmpl-lit6825","object":"chat.completion.chunk","created":1,'
b'"model":"kimi-k2","choices":[{"index":0,"delta":{},"finish_reason":"stop"}],'
b'"usage":{"prompt_tokens":3,"completion_tokens":4,"total_tokens":7}}\n\n'
b"data: [DONE]\n\n"
)
def _logging_obj(call_id: str):
from litellm.litellm_core_utils.litellm_logging import Logging
return Logging(
model="kimi-k2",
messages=MESSAGES,
stream=True,
call_type="anthropic_messages",
start_time=datetime.datetime.now(datetime.timezone.utc),
litellm_call_id=call_id,
function_id="1234",
)
def _streamed_message_id(raw_events: list[bytes]) -> str:
events = [json.loads(chunk.decode().split("data: ", 1)[1]) for chunk in raw_events]
message_start = next(e for e in events if e["type"] == "message_start")
return message_start["message"]["id"]
@pytest.fixture(autouse=True)
def _intercept_groq(respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch):
monkeypatch.setenv("GROQ_API_KEY", "gsk-lit6825-test")
monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True")
litellm.in_memory_llm_clients_cache.flush_cache()
respx_mock.post(GROQ_CHAT_URL).respond(
status_code=200,
headers={"Content-Type": "text/event-stream"},
content=CHAT_SSE_BODY,
)
@pytest.mark.asyncio
async def test_async_streaming_hands_the_logging_object_the_message_id_the_caller_is_streamed():
logging_obj = _logging_obj("6825beef-0000-4000-8000-000000000010")
sse = await LiteLLMMessagesToCompletionTransformationHandler.async_anthropic_messages_handler(
max_tokens=1024,
messages=MESSAGES,
model="groq/kimi-k2",
stream=True,
custom_llm_provider="groq",
litellm_logging_obj=logging_obj,
)
streamed_id = _streamed_message_id([chunk async for chunk in sse])
assert streamed_id.startswith("msg_")
assert logging_obj.streamed_anthropic_message_id == streamed_id
def test_sync_streaming_hands_the_logging_object_the_message_id_the_caller_is_streamed():
logging_obj = _logging_obj("6825beef-0000-4000-8000-000000000011")
sse = LiteLLMMessagesToCompletionTransformationHandler.anthropic_messages_handler(
max_tokens=1024,
messages=MESSAGES,
model="groq/kimi-k2",
stream=True,
custom_llm_provider="groq",
litellm_logging_obj=logging_obj,
)
streamed_id = _streamed_message_id(list(sse))
assert streamed_id.startswith("msg_")
assert logging_obj.streamed_anthropic_message_id == streamed_id
def test_concurrent_streams_are_keyed_on_their_own_message_id():
"""Two callers streaming at once must not be handed, or logged under, one another's id."""
first = AnthropicStreamWrapper(completion_stream=iter([]), model="kimi-k2")
second = AnthropicStreamWrapper(completion_stream=iter([]), model="kimi-k2")
assert first._message_id != second._message_id
assert _streamed_message_id(list(first.anthropic_sse_wrapper())) == first._message_id
assert _streamed_message_id(list(second.anthropic_sse_wrapper())) == second._message_id

View file

@ -1,9 +1,11 @@
import datetime
import json
import os
import sys
from unittest.mock import AsyncMock, patch
import pytest
import respx
sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "../../../../../..")))
@ -15,6 +17,18 @@ from litellm.llms.anthropic.experimental_pass_through.responses_adapters.handler
MESSAGES = [{"role": "user", "content": "hello"}]
RESPONSES_SSE_BODY = (
b"event: response.created\n"
b'data: {"type":"response.created","sequence_number":0,"response":{"id":"resp_lit6825",'
b'"object":"response","created_at":1,"status":"in_progress","model":"gpt-5.6-luna","output":[],'
b'"parallel_tool_calls":true,"tool_choice":"auto","tools":[]}}\n\n'
b"event: response.completed\n"
b'data: {"type":"response.completed","sequence_number":1,"response":{"id":"resp_lit6825",'
b'"object":"response","created_at":1,"status":"completed","model":"gpt-5.6-luna","output":[],'
b'"parallel_tool_calls":true,"tool_choice":"auto","tools":[],'
b'"usage":{"input_tokens":3,"output_tokens":4,"total_tokens":7}}}\n\n'
)
def test_build_responses_kwargs_derives_prompt_cache_key_from_user_id():
responses_kwargs = _build_responses_kwargs(
@ -82,3 +96,47 @@ async def test_streaming_message_start_reports_the_provider_local_model(requeste
message_start = next(e for e in events if e["type"] == "message_start")
assert message_start["message"]["model"] == expected_reported_model
@pytest.mark.asyncio
async def test_streaming_hands_the_logging_object_the_message_id_the_caller_is_streamed(
respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch
):
"""
The bridge mints the ``msg_`` id itself, and it is the only request id a streaming
/v1/messages caller ever sees, so the spend row has to be keyed on that same value.
"""
from litellm.litellm_core_utils.litellm_logging import Logging
monkeypatch.setenv("OPENAI_API_KEY", "sk-lit6825-test")
monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True")
litellm.in_memory_llm_clients_cache.flush_cache()
respx_mock.post("https://api.openai.com/v1/responses").respond(
status_code=200,
headers={"Content-Type": "text/event-stream"},
content=RESPONSES_SSE_BODY,
)
logging_obj = Logging(
model="gpt-5.6-luna",
messages=MESSAGES,
stream=True,
call_type="anthropic_messages",
start_time=datetime.datetime.now(datetime.timezone.utc),
litellm_call_id="6825beef-0000-4000-8000-000000000003",
function_id="1234",
)
sse = await LiteLLMMessagesToResponsesAPIHandler.async_anthropic_messages_handler(
max_tokens=1024,
messages=MESSAGES,
model="openai/gpt-5.6-luna",
stream=True,
custom_llm_provider="openai",
litellm_logging_obj=logging_obj,
)
events = [json.loads(chunk.decode().split("data: ", 1)[1]) async for chunk in sse]
message_start = next(e for e in events if e["type"] == "message_start")
assert message_start["message"]["id"].startswith("msg_")
assert logging_obj.streamed_anthropic_message_id == message_start["message"]["id"]

View file

@ -0,0 +1,66 @@
import pytest
from litellm.llms.bedrock.chat.invoke_transformations.amazon_moonshot_transformation import (
AmazonMoonshotConfig,
)
AWS_AUTH_PARAMS = {
"aws_access_key_id": "AKIAEXAMPLE",
"aws_secret_access_key": "secret",
"aws_session_token": "token",
"aws_region_name": "us-west-2",
"aws_session_name": "session",
"aws_role_name": "arn:aws:iam::000000000000:role/example",
"aws_web_identity_token": "web-identity",
"aws_sts_endpoint": "https://sts.us-west-2.amazonaws.com",
"aws_bedrock_runtime_endpoint": "https://bedrock-runtime.us-west-2.amazonaws.com",
"aws_external_id": "external",
}
def test_transform_request_never_resolves_aws_credentials():
"""A broken credential chain must not stop the request body from being built."""
config = AmazonMoonshotConfig()
transformed = config.transform_request(
model="bedrock/invoke/moonshot.kimi-k2-thinking",
messages=[{"role": "user", "content": "Hello"}],
optional_params={"aws_profile_name": "litellm-profile-that-does-not-exist", "max_tokens": 16},
litellm_params={},
headers={},
)
assert transformed["model"] == "moonshot.kimi-k2-thinking"
assert transformed["max_tokens"] == 16
assert "aws_profile_name" not in transformed
@pytest.mark.parametrize("aws_param", sorted(AWS_AUTH_PARAMS))
def test_transform_request_keeps_aws_params_out_of_the_body(aws_param: str):
config = AmazonMoonshotConfig()
transformed = config.transform_request(
model="bedrock/invoke/moonshot.kimi-k2-thinking",
messages=[{"role": "user", "content": "Hello"}],
optional_params={aws_param: AWS_AUTH_PARAMS[aws_param]},
litellm_params={},
headers={},
)
assert aws_param not in transformed
def test_transform_request_leaves_the_caller_aws_params_in_place_for_signing():
"""sign_request reads the aws_* keys off optional_params after transform_request runs."""
config = AmazonMoonshotConfig()
optional_params = dict(AWS_AUTH_PARAMS)
config.transform_request(
model="bedrock/invoke/moonshot.kimi-k2-thinking",
messages=[{"role": "user", "content": "Hello"}],
optional_params=optional_params,
litellm_params={},
headers={},
)
assert optional_params == AWS_AUTH_PARAMS

View file

@ -513,3 +513,27 @@ def test_the_rust_opt_in_needs_no_sigv4_principal():
assert not {"aws_access_key_id", "aws_secret_access_key", "aws_session_token"} & params.keys()
assert params["aws_region_name"] == "us-east-1"
assert seen["call"][0]["api_key"] == "bedrock-bearer-token"
@pytest.mark.parametrize("configured_through", ["env_var", "api_key"])
def test_bearer_token_auth_never_runs_the_sigv4_credential_chain(monkeypatch, configured_through):
"""The deployment's AWS profile does not exist, so resolving SigV4 credentials
raises; a bearer-token deployment must still serve the request, since the
bearer token alone signs it."""
if configured_through == "env_var":
monkeypatch.setenv("AWS_BEARER_TOKEN_BEDROCK", "bedrock-bearer-token")
else:
monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False)
client = _sync_client_returning_converse_response()
response = BedrockConverseLLM().completion(
**_completion_kwargs(
optional_params={"maxTokens": 16, "aws_profile_name": "litellm-no-such-aws-profile"},
litellm_params={},
client=client,
api_key="bedrock-bearer-token" if configured_through == "api_key" else None,
)
)
assert response.choices[0].message.content == "hi"
assert client.post.call_args.kwargs["headers"]["Authorization"] == "Bearer bedrock-bearer-token"

View file

@ -1033,3 +1033,29 @@ def test_load_credentials_assumes_role_with_external_id(monkeypatch):
assert credentials.token == "assumed-session-token"
assert aws_region_name == "us-east-1"
assert "aws_external_id" not in optional_params
def test_bedrock_embedding_bearer_token_never_runs_the_sigv4_credential_chain(monkeypatch):
"""The deployment's AWS profile does not exist, so resolving SigV4 credentials
raises; a bearer-token deployment must still serve the request, since the
bearer token alone signs it."""
monkeypatch.setenv("AWS_BEARER_TOKEN_BEDROCK", "env-bearer-token-12345")
client = HTTPHandler()
with patch.object(client, "post") as mock_post:
mock_response = Mock()
mock_response.status_code = 200
mock_response.text = json.dumps(titan_embedding_response)
mock_response.json = lambda: json.loads(mock_response.text)
mock_post.return_value = mock_response
response = litellm.embedding(
model="bedrock/amazon.titan-embed-text-v1",
input=test_input,
client=client,
aws_region_name="us-west-2",
aws_profile_name="litellm-no-such-aws-profile",
)
assert response.data[0]["embedding"] == titan_embedding_response["embedding"]
assert mock_post.call_args.kwargs["headers"]["Authorization"] == "Bearer env-bearer-token-12345"

View file

@ -135,3 +135,24 @@ class TestBedrockImageGeneration:
assert response is not None
assert len(response.data) > 0
mock_bedrock_image_gen.assert_called_once()
def test_image_generation_bearer_token_never_runs_the_sigv4_credential_chain(monkeypatch):
"""The deployment's AWS profile does not exist, so resolving SigV4 credentials
raises; a bearer-token deployment must still sign the request with the
bearer token alone."""
from litellm.llms.bedrock.image_generation.image_handler import BedrockImageGeneration
monkeypatch.setenv("AWS_BEARER_TOKEN_BEDROCK", "env-bearer-token-12345")
request = BedrockImageGeneration()._prepare_request(
model="amazon.nova-canvas-v1:0",
prompt="A cute baby sea otter",
optional_params={"aws_region_name": "us-west-2", "aws_profile_name": "litellm-no-such-aws-profile"},
api_base=None,
extra_headers=None,
api_key=None,
logging_obj=Mock(),
)
assert request.prepped.headers["Authorization"] == "Bearer env-bearer-token-12345"

View file

@ -3,6 +3,7 @@
import base64
import io
from typing import cast
from unittest.mock import Mock, patch
import httpx
import pytest
@ -655,3 +656,23 @@ def test_transform_response_empty_images_without_error_raises():
raw_response=resp,
logging_obj=None, # type: ignore[arg-type]
)
def test_prepare_request_bearer_token_never_runs_the_sigv4_credential_chain(monkeypatch):
"""The deployment's AWS profile does not exist, so resolving SigV4 credentials
raises; a bearer-token deployment must still sign the request with the
bearer token alone."""
monkeypatch.setenv("AWS_BEARER_TOKEN_BEDROCK", "env-bearer-token-12345")
request = BedrockImageEdit()._prepare_request(
model="amazon.nova-canvas-v1:0",
image=[io.BytesIO(b"fake-png")],
prompt="make it warmer",
optional_params={"aws_region_name": "us-west-2", "aws_profile_name": "litellm-no-such-aws-profile"},
api_base=None,
extra_headers=None,
logging_obj=Mock(),
api_key=None,
)
assert request.prepped.headers["Authorization"] == "Bearer env-bearer-token-12345"

View file

@ -3098,3 +3098,91 @@ async def test_a_provider_that_keeps_rejecting_is_not_retried_forever_on_the_asy
)
assert len(recorder.bodies) == 2
CONTAINER_NOT_FOUND_BODY = {
"error": {
"message": "Container with id 'cntr_gone' not found.",
"type": "invalid_request_error",
"param": None,
"code": None,
}
}
INVALID_API_KEY_BODY = {
"error": {
"message": "Incorrect API key provided: sk-proj-***. You can find your API key at https://platform.openai.com/account/api-keys.",
"type": "invalid_request_error",
"param": None,
"code": "invalid_api_key",
},
"status": 401,
}
CONTAINER_LIST_BODY = {
"object": "list",
"data": [{"id": "cntr_a", "object": "container", "created_at": 1, "status": "running", "name": "a"}],
"first_id": "cntr_a",
"last_id": "cntr_a",
"has_more": True,
}
def _container_sync_client(response: httpx.Response) -> HTTPHandler:
client = HTTPHandler()
client.client = httpx.Client(transport=httpx.MockTransport(lambda _request: response))
return client
def _container_async_client(response: httpx.Response) -> AsyncHTTPHandler:
client = AsyncHTTPHandler()
client.client = httpx.AsyncClient(transport=httpx.MockTransport(lambda _request: response))
return client
def test_container_retrieve_handler_raises_upstream_error_status_and_message():
from litellm.llms.openai.containers.transformation import OpenAIContainerConfig
with pytest.raises(BaseLLMException) as exc_info:
BaseLLMHTTPHandler().container_retrieve_handler(
container_id="cntr_gone",
container_provider_config=OpenAIContainerConfig(),
litellm_params=GenericLiteLLMParams(api_key="sk-test"),
logging_obj=Mock(),
client=_container_sync_client(httpx.Response(404, json=CONTAINER_NOT_FOUND_BODY)),
)
assert exc_info.value.status_code == 404
assert exc_info.value.message == "Container with id 'cntr_gone' not found."
@pytest.mark.asyncio
async def test_async_container_list_handler_raises_upstream_error_status_and_message():
from litellm.llms.openai.containers.transformation import OpenAIContainerConfig
with pytest.raises(BaseLLMException) as exc_info:
await BaseLLMHTTPHandler().async_container_list_handler(
container_provider_config=OpenAIContainerConfig(),
litellm_params=GenericLiteLLMParams(api_key="sk-rejected"),
logging_obj=Mock(),
client=_container_async_client(httpx.Response(401, json=INVALID_API_KEY_BODY)),
)
assert exc_info.value.status_code == 401
assert exc_info.value.message == INVALID_API_KEY_BODY["error"]["message"]
@pytest.mark.asyncio
async def test_async_container_list_handler_transforms_success_response():
from litellm.llms.openai.containers.transformation import OpenAIContainerConfig
response = await BaseLLMHTTPHandler().async_container_list_handler(
container_provider_config=OpenAIContainerConfig(),
litellm_params=GenericLiteLLMParams(api_key="sk-test"),
logging_obj=Mock(),
limit=1,
client=_container_async_client(httpx.Response(200, json=CONTAINER_LIST_BODY)),
)
assert [container.id for container in response.data] == ["cntr_a"]
assert response.has_more is True

View file

@ -1,4 +1,6 @@
import io
import json
from typing import get_type_hints
from unittest.mock import AsyncMock, MagicMock, patch
import orjson
@ -18,9 +20,11 @@ from litellm.proxy.common_utils.http_parsing_utils import (
_safe_get_request_parsed_body,
_safe_get_request_query_params,
_safe_set_request_parsed_body,
coerce_numeric_form_fields,
get_form_data,
get_request_body,
get_tags_from_request_body,
numeric_form_fields,
populate_request_with_path_params,
)
@ -1029,3 +1033,79 @@ class TestGetRequestBody:
mock_request = MagicMock()
mock_request.method = "GET"
assert await get_request_body(mock_request) == {}
class TestNumericFormFields:
def test_image_edit_schema_yields_only_n(self):
from litellm.types.images.main import ImageEditRequestParams
assert dict(numeric_form_fields(get_type_hints(ImageEditRequestParams))) == {"n": int}
def test_qualifiers_and_optionality_are_unwrapped(self):
from typing import Optional
from typing_extensions import Annotated, NotRequired, ReadOnly, Required, TypedDict
class Schema(TypedDict, total=False):
plain: int
optional: Optional[int]
piped: int | None
read_only: ReadOnly[int | None]
not_required: NotRequired[ReadOnly[int]]
required: Required[ReadOnly[Annotated[float, "meta"]]]
assert dict(numeric_form_fields(get_type_hints(Schema))) == {
"plain": int,
"optional": int,
"piped": int,
"read_only": int,
"not_required": int,
"required": float,
}
def test_non_scalar_and_bool_fields_are_skipped(self):
from typing import Any, Literal, Optional, Union
from typing_extensions import TypedDict
class Schema(TypedDict, total=False):
flag: bool
optional_flag: Optional[bool]
text: str
choice: Optional[Literal["high", "low"]]
numbers: list[int]
mapping: Optional[dict[str, Any]]
ambiguous: Union[int, str]
assert dict(numeric_form_fields(get_type_hints(Schema))) == {}
class TestCoerceNumericFormFields:
numeric_fields = {"n": int, "temperature": float}
def test_numeric_strings_are_parsed(self):
assert coerce_numeric_form_fields(
parsed_body={"n": "2", "temperature": "0.5"},
numeric_fields=self.numeric_fields,
) == {"n": 2, "temperature": 0.5}
def test_other_fields_keep_their_string_values(self):
result = coerce_numeric_form_fields(
parsed_body={"size": "1024x1024", "prompt": "2", "quality": "high"},
numeric_fields=self.numeric_fields,
)
assert result == {"size": "1024x1024", "prompt": "2", "quality": "high"}
def test_unparseable_value_is_left_for_the_provider_to_reject(self):
assert coerce_numeric_form_fields(
parsed_body={"n": "two", "temperature": ""},
numeric_fields=self.numeric_fields,
) == {"n": "two", "temperature": ""}
def test_already_typed_and_non_string_values_pass_through(self):
buffer = io.BytesIO(b"png")
result = coerce_numeric_form_fields(
parsed_body={"n": 3, "temperature": None, "image": buffer},
numeric_fields=self.numeric_fields,
)
assert result == {"n": 3, "temperature": None, "image": buffer}

View file

@ -0,0 +1,134 @@
import sys
from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock
import pytest
from fastapi import FastAPI
from fastapi.testclient import TestClient
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.container_endpoints import endpoints, ownership
from litellm.types.containers.main import ContainerListResponse, ContainerObject
PROXY_SERVER_STUB = SimpleNamespace(
general_settings={},
prisma_client=None,
llm_router=None,
proxy_config=None,
proxy_logging_obj=None,
select_data_generator=None,
user_api_base=None,
user_max_tokens=None,
user_model=None,
user_request_timeout=None,
user_temperature=None,
version="test",
)
ADMIN = UserAPIKeyAuth(user_id="admin-1", user_role=LitellmUserRoles.PROXY_ADMIN)
NON_ADMIN = UserAPIKeyAuth(user_id="user-1")
@pytest.fixture(autouse=True)
def clear_allowed_container_ids_cache():
ownership._ALLOWED_CONTAINER_IDS_CACHE.cache_dict.clear()
ownership._ALLOWED_CONTAINER_IDS_CACHE.ttl_dict.clear()
yield
ownership._ALLOWED_CONTAINER_IDS_CACHE.cache_dict.clear()
ownership._ALLOWED_CONTAINER_IDS_CACHE.ttl_dict.clear()
def _client(auth: UserAPIKeyAuth) -> TestClient:
app = FastAPI()
app.include_router(endpoints.router)
app.dependency_overrides[user_api_key_auth] = lambda: auth
return TestClient(app)
def _container(container_id: str) -> ContainerObject:
return ContainerObject(id=container_id, object="container", created_at=1, status="active")
def _page(*container_ids: str, has_more: bool) -> ContainerListResponse:
return ContainerListResponse(
object="list",
data=[_container(container_id) for container_id in container_ids],
has_more=has_more,
)
def _upstream_pages(monkeypatch, pages_by_after) -> MagicMock:
processor_cls = MagicMock(
side_effect=lambda data: SimpleNamespace(
base_process_llm_request=AsyncMock(return_value=pages_by_after[data["after"]])
)
)
monkeypatch.setattr(endpoints, "ProxyBaseLLMRequestProcessing", processor_cls)
return processor_cls
def _forwarded_pages(processor_cls: MagicMock):
return [(call.kwargs["data"]["after"], call.kwargs["data"]["limit"]) for call in processor_cls.call_args_list]
def test_list_containers_forwards_typed_pagination_params_for_admins(monkeypatch):
monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", PROXY_SERVER_STUB)
processor_cls = _upstream_pages(monkeypatch, {"cntr_prev": _page("cntr_next", has_more=True)})
response = _client(ADMIN).get(
"/v1/containers",
params={"limit": "1", "order": "desc", "after": "cntr_prev"},
headers={"Authorization": "Bearer sk-test"},
)
assert response.status_code == 200
assert [item["id"] for item in response.json()["data"]] == ["cntr_next"]
assert response.json()["has_more"] is True
assert _forwarded_pages(processor_cls) == [("cntr_prev", 1)]
assert processor_cls.call_args.kwargs["data"]["order"] == "desc"
def test_list_containers_rejects_a_non_integer_limit(monkeypatch):
monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", PROXY_SERVER_STUB)
processor_cls = _upstream_pages(monkeypatch, {})
response = _client(ADMIN).get(
"/v1/containers",
params={"limit": "abc"},
headers={"Authorization": "Bearer sk-test"},
)
assert response.status_code == 422
processor_cls.assert_not_called()
def test_list_containers_pages_upstream_until_non_admin_keys_see_their_containers(monkeypatch):
monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", PROXY_SERVER_STUB)
table = AsyncMock()
table.find_many.return_value = [SimpleNamespace(model_object_id="container:openai:cntr_owned")]
monkeypatch.setattr(
ownership,
"_get_prisma_client",
AsyncMock(return_value=SimpleNamespace(db=SimpleNamespace(litellm_managedobjecttable=table))),
)
processor_cls = _upstream_pages(
monkeypatch,
{
None: _page("cntr_other", has_more=True),
"cntr_other": _page("cntr_owned", has_more=False),
},
)
response = _client(NON_ADMIN).get(
"/v1/containers",
params={"limit": "1"},
headers={"Authorization": "Bearer sk-test"},
)
assert response.status_code == 200
body = response.json()
assert [item["id"] for item in body["data"]] == ["cntr_owned"]
assert body["first_id"] == "cntr_owned"
assert body["last_id"] == "cntr_owned"
assert body["has_more"] is False
assert _forwarded_pages(processor_cls) == [(None, 100), ("cntr_other", 100)]

View file

@ -0,0 +1,62 @@
import sys
from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock
from fastapi import FastAPI
from fastapi.testclient import TestClient
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.container_endpoints import endpoints, handler_factory
PROXY_SERVER_STUB = SimpleNamespace(
general_settings={},
prisma_client=None,
llm_router=None,
proxy_config=None,
proxy_logging_obj=None,
select_data_generator=None,
user_api_base=None,
user_max_tokens=None,
user_model=None,
user_request_timeout=None,
user_temperature=None,
version="test",
)
def _client() -> TestClient:
app = FastAPI()
app.include_router(endpoints.router)
app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(user_id="user-1")
return TestClient(app)
def test_list_container_files_forwards_declared_query_params(monkeypatch):
monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", PROXY_SERVER_STUB)
monkeypatch.setattr(
handler_factory,
"assert_user_can_access_container",
AsyncMock(return_value=("cntr_123", "openai")),
)
processor_cls = MagicMock()
processor_cls.return_value.base_process_llm_request = AsyncMock(
return_value={"object": "list", "data": [], "has_more": True}
)
monkeypatch.setattr(handler_factory, "ProxyBaseLLMRequestProcessing", processor_cls)
response = _client().get(
"/v1/containers/cntr_123/files",
params={"limit": "1", "order": "desc", "after": "cfile_prev", "unknown": "x"},
headers={"Authorization": "Bearer sk-test"},
)
assert response.status_code == 200
assert response.json()["has_more"] is True
assert processor_cls.return_value.base_process_llm_request.await_args.kwargs["route_type"] == "alist_container_files"
forwarded = processor_cls.call_args.kwargs["data"]
assert forwarded["container_id"] == "cntr_123"
assert forwarded["limit"] == "1"
assert forwarded["order"] == "desc"
assert forwarded["after"] == "cfile_prev"
assert "unknown" not in forwarded

View file

@ -29,7 +29,7 @@ def test_hands_the_alter_statement_to_the_prisma_cli():
return subprocess.CompletedProcess(cmd, 0)
with patch(
"litellm_proxy_extras.replica_identity.subprocess.run", side_effect=capture
"litellm_proxy_extras.replica_identity.run_prisma", side_effect=capture
):
applied = apply_replica_identity_full(
schema_path="/somewhere/schema.prisma",
@ -60,7 +60,7 @@ def test_hands_the_alter_statement_to_the_prisma_cli():
)
def test_every_failure_is_reported_instead_of_raised(failure):
with patch(
"litellm_proxy_extras.replica_identity.subprocess.run", side_effect=failure
"litellm_proxy_extras.replica_identity.run_prisma", side_effect=failure
):
assert (
apply_replica_identity_full(

View file

@ -5527,10 +5527,6 @@ async def test_streaming_end_of_stream_block_emits_error_frame_instead_of_trunca
error frame instead. The finish chunk is withheld while the end-of-stream
scan runs, so on a block it is dropped rather than relayed before the
frame."""
from litellm.llms import load_guardrail_translation_mappings
from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail import (
unified_guardrail as unified_module,
)
from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import (
UnifiedLLMGuardrails,
)
@ -5569,20 +5565,16 @@ async def test_streaming_end_of_stream_block_emits_error_frame_instead_of_trunca
yield _chunk("the forbidden ")
yield _chunk("topic answer", finish_reason="stop")
unified_module.endpoint_guardrail_translation_mappings = load_guardrail_translation_mappings()
try:
with patch.object(guardrail, "make_bedrock_api_request", new_callable=AsyncMock) as mock_api:
mock_api.side_effect = guardrail._get_http_exception_for_blocked_guardrail(blocked_response)
with patch.object(guardrail, "make_bedrock_api_request", new_callable=AsyncMock) as mock_api:
mock_api.side_effect = guardrail._get_http_exception_for_blocked_guardrail(blocked_response)
out = []
async for item in UnifiedLLMGuardrails().async_post_call_streaming_iterator_hook(
user_api_key_dict=UserAPIKeyAuth(api_key="test", request_route="/v1/chat/completions"),
response=_mock_stream(),
request_data={"guardrail_to_apply": guardrail, "model": "gpt-4"},
):
out.append(item)
finally:
unified_module.endpoint_guardrail_translation_mappings = None
out = []
async for item in UnifiedLLMGuardrails().async_post_call_streaming_iterator_hook(
user_api_key_dict=UserAPIKeyAuth(api_key="test", request_route="/v1/chat/completions"),
response=_mock_stream(),
request_data={"guardrail_to_apply": guardrail, "model": "gpt-4"},
):
out.append(item)
assert len(out) == 2
assert isinstance(out[0], ModelResponseStream)
@ -5792,3 +5784,25 @@ async def test_apply_guardrail_debug_log_masks_signed_request_headers():
assert header_lines, "expected the signed-request debug line to be logged"
assert any("X-Amz-Security-Token" in message for message in header_lines)
assert all(session_token not in message for message in rendered_messages)
@pytest.mark.asyncio
async def test_bearer_token_never_runs_the_sigv4_credential_chain(monkeypatch):
"""The guardrail's AWS profile does not exist, so resolving SigV4 credentials
raises; with a bearer token configured the guardrail must still run, since
the bearer token alone signs the request."""
monkeypatch.setenv("AWS_BEARER_TOKEN_BEDROCK", "env-bearer-token-12345")
guardrail = BedrockGuardrail(
guardrailIdentifier="test-guardrail",
guardrailVersion="DRAFT",
aws_profile_name="litellm-no-such-aws-profile",
)
mock_response = MagicMock()
mock_response.status_code = 200
mock_response.json.return_value = {"action": "NONE", "assessments": []}
with patch.object(guardrail.async_handler, "post", new_callable=AsyncMock, return_value=mock_response) as mock_post:
response = await guardrail.make_bedrock_api_request(source="INPUT", messages=[{"role": "user", "content": "hello"}])
assert response["action"] == "NONE"
assert mock_post.call_args.kwargs["headers"]["Authorization"] == "Bearer env-bearer-token-12345"

View file

@ -833,3 +833,31 @@ async def test_many_blocks_scanned_at_request_level_and_can_block():
sent_texts = [c["text"] for m in body_messages for c in m["content"]]
assert sent_texts == [f"b{i}" for i in range(25)]
assert all(len(m["content"]) <= 10 for m in body_messages)
@pytest.mark.asyncio
async def test_checks_bearer_token_never_runs_the_sigv4_credential_chain(monkeypatch):
"""Same bearer-token rule as ApplyGuardrail: the guardrail's AWS profile does
not exist, yet the InvokeGuardrailChecks call still goes out on the bearer
token and its verdict is enforced."""
monkeypatch.setenv("AWS_BEARER_TOKEN_BEDROCK", "env-bearer-token-12345")
g = BedrockGuardrail(
checks=CONTENT_FILTER_CHECKS,
content_filter_threshold=0.5,
aws_profile_name="litellm-no-such-aws-profile",
)
payload = {"results": {"contentFilter": {"results": [{"category": "VIOLENCE", "severityScore": 0.8}]}}}
post = AsyncMock(return_value=_mock_http_response(200, payload))
with patch.object(g.async_handler, "post", new=post):
with pytest.raises(HTTPException) as exc:
await g.make_bedrock_api_request(
source="INPUT",
messages=[{"role": "user", "content": "hi"}],
request_data={"messages": []},
)
assert exc.value.detail["bedrock_guardrail_checks"] == [
{"check": "contentFilter", "category": "VIOLENCE", "severityScore": 0.8}
]
assert post.call_args.kwargs["headers"]["Authorization"] == "Bearer env-bearer-token-12345"

View file

@ -75,19 +75,29 @@ class _NoopTranslation(BaseTranslation):
return response
def _patch_translation_mappings(monkeypatch, mappings):
"""Point the unified guardrail at ``mappings`` for one test, restored by pytest.
Every override goes through this one seam: competing writers to the same state
are what leaked a stale handler map into unrelated test files (LIT-6834).
"""
monkeypatch.setattr(unified_module, "load_guardrail_translation_mappings", lambda: mappings)
@pytest.fixture(autouse=True)
def _inject_mcp_handler_mapping():
def _inject_mcp_handler_mapping(monkeypatch):
"""Inject MCP handler mapping so the unified guardrail can run inside tests."""
unified_module.endpoint_guardrail_translation_mappings = {
CallTypes.call_mcp_tool: MCPGuardrailTranslationHandler,
CallTypes.anthropic_messages: _NoopTranslation,
CallTypes.ocr: OCRHandler,
CallTypes.aocr: OCRHandler,
CallTypes.responses: OpenAIResponsesHandler,
CallTypes.aresponses: OpenAIResponsesHandler,
}
yield
unified_module.endpoint_guardrail_translation_mappings = None
_patch_translation_mappings(
monkeypatch,
{
CallTypes.call_mcp_tool: MCPGuardrailTranslationHandler,
CallTypes.anthropic_messages: _NoopTranslation,
CallTypes.ocr: OCRHandler,
CallTypes.aocr: OCRHandler,
CallTypes.responses: OpenAIResponsesHandler,
CallTypes.aresponses: OpenAIResponsesHandler,
},
)
class TestUnifiedLLMGuardrails:
@ -396,7 +406,7 @@ class TestUnifiedLLMGuardrails:
class TestAsyncPostCallStreamingIteratorHook:
@pytest.mark.asyncio
async def test_streaming_content_not_lost_on_sampled_chunks(self):
async def test_streaming_content_not_lost_on_sampled_chunks(self, monkeypatch):
"""
Verify that every chunk's content is preserved in the output stream.
@ -442,10 +452,7 @@ class TestUnifiedLLMGuardrails:
return responses_so_far
# Override the mapping to use our content-clearing translation
unified_module.endpoint_guardrail_translation_mappings = {
CallTypes.acompletion: _ContentClearingTranslation,
}
_patch_translation_mappings(monkeypatch, {CallTypes.acompletion: _ContentClearingTranslation})
handler = UnifiedLLMGuardrails()
guardrail = RecordingGuardrail()
@ -885,12 +892,8 @@ class TestStreamingTransform:
completions streaming surface."""
@pytest.fixture(autouse=True)
def _use_openai_handler_mapping(self):
unified_module.endpoint_guardrail_translation_mappings = {
CallTypes.acompletion: OpenAIChatCompletionsHandler,
}
yield
unified_module.endpoint_guardrail_translation_mappings = None
def _use_openai_handler_mapping(self, monkeypatch):
_patch_translation_mappings(monkeypatch, {CallTypes.acompletion: OpenAIChatCompletionsHandler})
@pytest.mark.asyncio
async def test_block_only_drops_text_rewrites(self):
@ -1719,6 +1722,10 @@ class TestAppliedGuardrailsReflectsExecution:
decision and marks itself only when it actually ran (LIT-4650). Ordinary
guardrails are still auto-marked by the hook after dispatch."""
@pytest.fixture(autouse=True)
def _use_texts_only_mapping(self, monkeypatch):
_patch_translation_mappings(monkeypatch, {CallTypes.pass_through: _TextsOnlyTranslation})
@staticmethod
def _data(guardrail):
return {
@ -1728,7 +1735,6 @@ class TestAppliedGuardrailsReflectsExecution:
}
async def _run(self, guardrail):
unified_module.endpoint_guardrail_translation_mappings = {CallTypes.pass_through: _TextsOnlyTranslation}
data = self._data(guardrail)
await UnifiedLLMGuardrails().async_pre_call_hook(
user_api_key_dict=None,
@ -1830,10 +1836,8 @@ class TestStreamingHttpErrorFrames:
silently truncates the SSE stream (PR #38722 defect 1)."""
@pytest.fixture(autouse=True)
def _use_real_mappings(self):
unified_module.endpoint_guardrail_translation_mappings = load_guardrail_translation_mappings()
yield
unified_module.endpoint_guardrail_translation_mappings = None
def _use_real_mappings(self, monkeypatch):
_patch_translation_mappings(monkeypatch, load_guardrail_translation_mappings())
@pytest.mark.asyncio
async def test_chat_eos_block_emits_data_error_frame(self):
@ -1938,10 +1942,8 @@ class TestStreamingGuardrailInformationBucket:
guardrail_information write was diverted and /spend/logs showed null."""
@pytest.fixture(autouse=True)
def _use_real_mappings(self):
unified_module.endpoint_guardrail_translation_mappings = load_guardrail_translation_mappings()
yield
unified_module.endpoint_guardrail_translation_mappings = None
def _use_real_mappings(self, monkeypatch):
_patch_translation_mappings(monkeypatch, load_guardrail_translation_mappings())
@pytest.mark.asyncio
async def test_chat_eos_scan_writes_guardrail_information_to_metadata(self):
@ -2038,11 +2040,7 @@ class TestStreamingScanDedup:
@pytest.fixture(autouse=True)
def _use_real_mappings(self, monkeypatch):
monkeypatch.setattr(
unified_module,
"endpoint_guardrail_translation_mappings",
load_guardrail_translation_mappings(),
)
_patch_translation_mappings(monkeypatch, load_guardrail_translation_mappings())
@pytest.mark.asyncio
async def test_chat_terminal_chunk_on_sampled_index_is_scanned_once(self):
@ -2239,3 +2237,55 @@ class TestStreamingScanDedup:
assert out == chunks
assert [scan["texts"] for scan in guardrail.scans] == [["abc"]]
class TestTranslationMappingsAreReadLive:
"""The hooks must read the handler map on every call, never memoize it on the module.
A second module-level cache is what let one test's handler map outlive its own
teardown and decide how unrelated files translated their streams (LIT-6834).
"""
@staticmethod
def _ocr_request(guardrail):
return {
"guardrail_to_apply": guardrail,
"model": "mistral/mistral-ocr-latest",
"document": {
"type": "document_url",
"document_url": "https://arxiv.org/pdf/2201.04234",
},
}
async def _run_pre_call(self, guardrail):
await UnifiedLLMGuardrails().async_pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(api_key="test-key"),
cache=DualCache(),
data=self._ocr_request(guardrail),
call_type=CallTypes.aocr.value,
)
@pytest.mark.asyncio
async def test_remapping_between_calls_changes_which_handler_runs(self, monkeypatch):
_patch_translation_mappings(monkeypatch, {CallTypes.completion: _NoopTranslation})
unmapped = RecordingGuardrail()
await self._run_pre_call(unmapped)
assert unmapped.apply_calls == []
_patch_translation_mappings(monkeypatch, {CallTypes.aocr: OCRHandler})
mapped = RecordingGuardrail()
await self._run_pre_call(mapped)
assert [call["input_type"] for call in mapped.apply_calls] == ["request"]
@pytest.mark.asyncio
async def test_module_exposes_no_second_assignable_handler_map(self, monkeypatch):
_patch_translation_mappings(monkeypatch, {CallTypes.aocr: OCRHandler})
guardrail = RecordingGuardrail()
await self._run_pre_call(guardrail)
assert len(guardrail.apply_calls) == 1
assert not [
name
for name, value in vars(unified_module).items()
if isinstance(value, dict) and CallTypes.aocr in value
]

View file

@ -802,7 +802,7 @@ async def test_bedrock_guardrail_prepare_request_with_api_key():
@pytest.mark.asyncio
async def test_bedrock_guardrail_prepare_request_without_api_key():
async def test_bedrock_guardrail_prepare_request_without_api_key(monkeypatch):
"""Test _prepare_request method falls back to SigV4 when no api_key is provided"""
from unittest.mock import Mock, patch
@ -820,18 +820,13 @@ async def test_bedrock_guardrail_prepare_request_without_api_key():
# Test data without api_key
test_data = {"source": "INPUT", "content": [{"text": {"text": "test content"}}]}
monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False)
with (
patch(
"litellm.proxy.guardrails.guardrail_hooks.bedrock_guardrails.get_secret_str"
) as mock_get_secret,
patch("botocore.auth.SigV4Auth") as mock_sigv4_auth,
patch("botocore.awsrequest.AWSRequest") as mock_aws_request,
):
# Mock no AWS_BEARER_TOKEN_BEDROCK
mock_get_secret.return_value = None
# Mock SigV4Auth
mock_sigv4_instance = Mock()
mock_sigv4_auth.return_value = mock_sigv4_instance
@ -857,7 +852,7 @@ async def test_bedrock_guardrail_prepare_request_without_api_key():
@pytest.mark.asyncio
async def test_bedrock_guardrail_prepare_request_with_bearer_token_env():
async def test_bedrock_guardrail_prepare_request_with_bearer_token_env(monkeypatch):
"""Test _prepare_request method uses Bearer token from environment when available"""
from unittest.mock import Mock, patch
@ -875,15 +870,9 @@ async def test_bedrock_guardrail_prepare_request_with_bearer_token_env():
# Test data without api_key
test_data = {"source": "INPUT", "content": [{"text": {"text": "test content"}}]}
monkeypatch.setenv("AWS_BEARER_TOKEN_BEDROCK", "env-bearer-token-456")
with (
patch(
"litellm.proxy.guardrails.guardrail_hooks.bedrock_guardrails.get_secret_str"
) as mock_get_secret,
patch("botocore.awsrequest.AWSRequest") as mock_aws_request,
):
mock_get_secret.return_value = "env-bearer-token-456"
with patch("botocore.awsrequest.AWSRequest") as mock_aws_request:
mock_request_instance = Mock()
mock_request_instance.prepare.return_value = Mock()
mock_aws_request.return_value = mock_request_instance

View file

@ -5,10 +5,13 @@ from typing import Any, Dict
import orjson
import pytest
from fastapi import FastAPI
from fastapi.testclient import TestClient
from starlette.requests import Request
from starlette.responses import Response
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.image_endpoints import endpoints
@ -115,3 +118,52 @@ async def test_image_generation_prompt_rerouting(monkeypatch):
assert captured_route_request_data["prompt"] == "sanitized prompt"
assert "messages" not in captured_route_request_data
assert response.headers.get("x-callback-test") == "value"
def _image_edit_client(monkeypatch, captured: Dict[str, Any]) -> TestClient:
class CaptureProcessing:
def __init__(self, data: Dict[str, Any]) -> None:
captured.update(data)
async def base_process_llm_request(self, **_: Any) -> Dict[str, Any]:
return {"data": [{"b64_json": "aGk="}]}
monkeypatch.setattr(endpoints, "ProxyBaseLLMRequestProcessing", CaptureProcessing)
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {})
monkeypatch.setattr("litellm.proxy.proxy_server.user_model", None)
app = FastAPI()
app.include_router(endpoints.router)
app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth()
return TestClient(app)
def test_image_edit_multipart_n_reaches_the_provider_as_an_int(monkeypatch):
"""A multipart `n` must not arrive as the string Starlette parsed it into."""
captured: Dict[str, Any] = {}
response = _image_edit_client(monkeypatch, captured).post(
"/v1/images/edits",
files={"image": ("tree.png", b"\x89PNG\r\n\x1a\n", "image/png")},
data={"model": "nova-canvas", "prompt": "add a hat", "n": "2", "size": "1024x1024"},
)
assert response.status_code == 200
assert captured["n"] == 2
assert isinstance(captured["n"], int)
assert captured["size"] == "1024x1024"
assert captured["prompt"] == "add a hat"
def test_image_edit_multipart_n_that_is_not_a_number_is_left_alone(monkeypatch):
"""An unparseable `n` still reaches the provider, which rejects it as before."""
captured: Dict[str, Any] = {}
response = _image_edit_client(monkeypatch, captured).post(
"/v1/images/edits",
files={"image": ("tree.png", b"\x89PNG\r\n\x1a\n", "image/png")},
data={"model": "nova-canvas", "prompt": "add a hat", "n": "two"},
)
assert response.status_code == 200
assert captured["n"] == "two"

View file

@ -292,8 +292,8 @@ class TestUnifiedGuardrailCallTypeResolution:
with patch.object(
unified_guardrail_module,
"endpoint_guardrail_translation_mappings",
{CallTypes.pass_through: mock_handler_class},
"load_guardrail_translation_mappings",
lambda: {CallTypes.pass_through: mock_handler_class},
):
result = await unified.async_post_call_success_hook(
data=data,

View file

@ -421,6 +421,76 @@ def test_rag_query_store_params_win_over_user_retrieval_config(client_internal_u
assert forwarded_config["aws_region_name"] == "eu-west-1"
def test_rag_query_forwards_managed_store_credentials_to_search(client_internal_user):
"""
Regression for LIT-6773: the registry store's api_key / api_base and its
provider extras (Milvus outputFields, milvus_text_field) must reach the
vector store search the way the direct /v1/vector_stores/{id}/search
endpoint forwards them. Pre-fix the RAG path allowlisted them away and a
managed Milvus store 500'd with "MILVUS_API_KEY is not set".
"""
import litellm
from litellm import Router
from litellm.types.vector_stores import VectorStoreSearchResponse
mock_vector_store = {
"vector_store_id": "customer_kb",
"custom_llm_provider": "milvus",
"litellm_params": {
"vector_store_id": "customer_kb",
"custom_llm_provider": "milvus",
"api_base": "http://127.0.0.1:19530",
"api_key": "root:Milvus",
"litellm_embedding_model": "multilingual-e5-large",
"milvus_text_field": "book_intro_text",
"outputFields": ["book_intro_text"],
},
}
mock_registry = MagicMock()
mock_registry.get_litellm_managed_vector_store_from_registry.return_value = mock_vector_store
fake_search = AsyncMock(
return_value=VectorStoreSearchResponse(object="vector_store.search_results.page", search_query="q", data=[])
)
router = Router(
model_list=[
{
"model_name": "gpt-4o-mini",
"litellm_params": {"model": "openai/gpt-4o-mini", "api_key": "sk-test", "mock_response": "hi"},
}
]
)
with (
patch("litellm.vector_stores.asearch", new=fake_search), # test-quality-ok: the search boundary under test
patch.object(litellm, "vector_store_registry", mock_registry), # test-quality-ok: seeds the store under test
patch("litellm.proxy.proxy_server.llm_router", router), # test-quality-ok: mock-response router for completion
patch( # test-quality-ok: store access is not under test, so the request reaches the search boundary
"litellm.proxy.vector_store_endpoints.utils.can_user_access_vector_store",
new=AsyncMock(return_value=True),
),
):
response = client_internal_user.post(
"/v1/rag/query",
json={
"model": "gpt-4o-mini",
"messages": [{"role": "user", "content": "which database is built for similarity search?"}],
"retrieval_config": {"vector_store_id": "customer_kb", "custom_llm_provider": "milvus", "top_k": 2},
},
)
assert response.status_code == 200, response.json()
fake_search.assert_awaited_once()
search_kwargs = fake_search.await_args.kwargs
assert search_kwargs["vector_store_id"] == "customer_kb"
assert search_kwargs["custom_llm_provider"] == "milvus"
assert search_kwargs["max_num_results"] == 2
assert search_kwargs["api_base"] == "http://127.0.0.1:19530"
assert search_kwargs["api_key"] == "root:Milvus"
assert search_kwargs["litellm_embedding_model"] == "multilingual-e5-large"
assert search_kwargs["milvus_text_field"] == "book_intro_text"
assert search_kwargs["outputFields"] == ["book_intro_text"]
@pytest.mark.parametrize(
"blocked_key",
["embedding_model", "litellm_embedding_model", "litellm_embedding_config", "litellm_credential_name"],

View file

@ -3399,6 +3399,138 @@ def test_get_spend_logs_id_prefers_the_response_id_over_the_standard_logging_id(
)
@pytest.mark.asyncio
async def test_spend_log_request_id_is_the_message_id_a_bridged_streaming_caller_was_streamed():
"""A streaming /v1/messages call against a non-Anthropic model is served a msg_ id the
adapter mints itself, and it is the only request id that call ever shows the caller, so
GET /spend/logs?request_id=msg_... has to land on the row."""
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.llms.anthropic.experimental_pass_through.responses_adapters.streaming_iterator import (
AnthropicResponsesStreamWrapper,
)
from litellm.types.llms.openai import (
ResponseAPIUsage,
ResponseCompletedEvent,
ResponsesAPIResponse,
)
logging_obj = Logging(
model="gpt-5.6",
messages=[{"role": "user", "content": "hi"}],
stream=True,
call_type="anthropic_messages",
start_time=datetime.datetime.now(timezone.utc),
litellm_call_id="6825cafe-0000-4000-8000-000000000001",
function_id="1234",
)
logging_obj.optional_params = {}
completed_response = ResponsesAPIResponse(
id="resp_01Lit6825Bridged",
object="response",
created_at=1767225600,
model="gpt-5.6",
status="completed",
output=[
{
"id": "msg_bridged_output",
"type": "message",
"role": "assistant",
"status": "completed",
"content": [{"type": "output_text", "text": "epsilon", "annotations": []}],
}
],
usage=ResponseAPIUsage(input_tokens=12, output_tokens=5, total_tokens=17),
)
async def _responses_stream():
yield {"type": "response.created"}
yield {"type": "response.output_text.delta", "item_id": "msg_bridged_output", "delta": "epsilon"}
yield ResponseCompletedEvent(type="response.completed", response=completed_response)
wrapper = AnthropicResponsesStreamWrapper(
responses_stream=_responses_stream(),
model="gpt-5.6",
litellm_logging_obj=logging_obj,
)
sse_frames = [frame.decode() async for frame in wrapper.async_anthropic_sse_wrapper()]
message_start_frames = [f for f in sse_frames if f.startswith("event: message_start\n")]
assert len(message_start_frames) == 1
streamed_message_id = json.loads(message_start_frames[0].split("data: ", 1)[1])["message"]["id"]
assert streamed_message_id.startswith("msg_")
_, _, logged_response = logging_obj._success_handler_helper_fn(
result=ResponseCompletedEvent(type="response.completed", response=completed_response),
start_time=datetime.datetime.now(timezone.utc),
end_time=datetime.datetime.now(timezone.utc),
)
assert logged_response.id == streamed_message_id
payload = get_logging_payload(
kwargs={
"call_type": "anthropic_messages",
"model": "gpt-5.6",
"litellm_call_id": "6825cafe-0000-4000-8000-000000000001",
"litellm_params": {"metadata": {"user_api_key": "test-key"}},
},
response_obj=logged_response,
start_time=datetime.datetime.now(timezone.utc),
end_time=datetime.datetime.now(timezone.utc),
)
assert payload["request_id"] == streamed_message_id
@pytest.mark.asyncio
async def test_spend_log_request_id_is_untouched_when_no_message_id_was_streamed():
"""Only the bridged streaming adapter mints a msg_ id of its own, so every other
/v1/messages call must keep the id its own response carried."""
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.types.llms.openai import (
ResponseAPIUsage,
ResponseCompletedEvent,
ResponsesAPIResponse,
)
logging_obj = Logging(
model="gpt-5.6",
messages=[{"role": "user", "content": "hi"}],
stream=False,
call_type="anthropic_messages",
start_time=datetime.datetime.now(timezone.utc),
litellm_call_id="6825cafe-0000-4000-8000-000000000002",
function_id="1234",
)
logging_obj.optional_params = {}
completed_response = ResponsesAPIResponse(
id="resp_01Lit6825Unbridged",
object="response",
created_at=1767225600,
model="gpt-5.6",
status="completed",
output=[
{
"id": "msg_unbridged_output",
"type": "message",
"role": "assistant",
"status": "completed",
"content": [{"type": "output_text", "text": "epsilon", "annotations": []}],
}
],
usage=ResponseAPIUsage(input_tokens=12, output_tokens=5, total_tokens=17),
)
_, _, logged_response = logging_obj._success_handler_helper_fn(
result=ResponseCompletedEvent(type="response.completed", response=completed_response),
start_time=datetime.datetime.now(timezone.utc),
end_time=datetime.datetime.now(timezone.utc),
)
assert logged_response.id
assert not logged_response.id.startswith("msg_")
def test_batch_cost_row_does_not_collide_with_the_batch_creation_row():
"""Creating a batch writes a row keyed by the batch's own id, so keying the cost row
the same way makes the insert a duplicate of it. request_id is the primary key and the
@ -3956,3 +4088,207 @@ def test_caller_forged_router_metadata_is_discarded(bucket):
)
metadata = json.loads(payload["metadata"])
assert metadata["router_metadata"] is None
ANTHROPIC_MESSAGES_RESPONSE: Final = {
"id": "msg_01Lit6806NonStreaming",
"type": "message",
"role": "assistant",
"model": "claude-haiku-4-5",
"content": [{"type": "text", "text": "epsilon"}],
"stop_reason": "end_turn",
"stop_sequence": None,
"usage": {"input_tokens": 14, "output_tokens": 4},
}
ANTHROPIC_MESSAGES_SSE_CHUNKS: Final = (
'event: message_start\ndata: {"type":"message_start","message":{"id":"msg_01Lit6806Streaming",'
'"type":"message","role":"assistant","model":"claude-haiku-4-5","content":[],'
'"usage":{"input_tokens":14,"output_tokens":1}}}\n\n',
'event: content_block_start\ndata: {"type":"content_block_start","index":0,'
'"content_block":{"type":"text","text":""}}\n\n',
'event: content_block_delta\ndata: {"type":"content_block_delta","index":0,'
'"delta":{"type":"text_delta","text":"epsilon"}}\n\n',
'event: content_block_stop\ndata: {"type":"content_block_stop","index":0}\n\n',
'event: message_delta\ndata: {"type":"message_delta","delta":{"stop_reason":"end_turn"},'
'"usage":{"output_tokens":4}}\n\n',
"event: message_stop\ndata: {\"type\":\"message_stop\"}\n\n",
)
def _anthropic_messages_logging_obj(*, stream: bool) -> Any:
from litellm.litellm_core_utils.litellm_logging import Logging
logging_obj = Logging(
model="claude-haiku-4-5",
messages=[{"role": "user", "content": "hi"}],
stream=stream,
call_type="anthropic_messages",
start_time=datetime.datetime.now(timezone.utc),
litellm_call_id="6806cafe-0000-4000-8000-000000000001",
function_id="1234",
)
logging_obj.optional_params = {}
logging_obj.model_call_details["custom_llm_provider"] = "anthropic"
return logging_obj
def _spend_log_request_id(response_obj: Any, kwargs: dict) -> str:
payload = get_logging_payload(
kwargs=kwargs,
response_obj=response_obj,
start_time=datetime.datetime.now(timezone.utc),
end_time=datetime.datetime.now(timezone.utc),
)
return payload["request_id"]
def test_spend_log_request_id_is_the_message_id_a_non_streaming_messages_caller_received():
"""
POST /v1/messages hands the caller `id: msg_...`, the only request id they ever see, so
GET /spend/logs?request_id=msg_... has to find the row.
"""
logging_obj = _anthropic_messages_logging_obj(stream=False)
logged_response = logging_obj._handle_anthropic_messages_response_logging(
result=ANTHROPIC_MESSAGES_RESPONSE
)
assert logged_response.id == "msg_01Lit6806NonStreaming"
assert (
_spend_log_request_id(
response_obj=logged_response,
kwargs={
"call_type": "anthropic_messages",
"model": "claude-haiku-4-5",
"litellm_call_id": "6806cafe-0000-4000-8000-000000000001",
"litellm_params": {"metadata": {"user_api_key": "test-key"}},
},
)
== "msg_01Lit6806NonStreaming"
)
def test_spend_log_request_id_is_the_message_id_a_streaming_messages_caller_received():
"""
The streaming leg of /v1/messages logs through the Anthropic passthrough handler, which used
to stamp litellm_call_id over the msg_ id carried by the message_start event.
"""
from litellm.proxy.pass_through_endpoints.llm_provider_handlers.anthropic_passthrough_logging_handler import (
AnthropicPassthroughLoggingHandler,
)
from litellm.types.passthrough_endpoints.pass_through_endpoints import EndpointType
logging_obj = _anthropic_messages_logging_obj(stream=True)
logging_obj.model_call_details["stream"] = True
logged = AnthropicPassthroughLoggingHandler._handle_logging_anthropic_collected_chunks(
litellm_logging_obj=logging_obj,
passthrough_success_handler_obj=MagicMock(),
url_route="/v1/messages",
request_body={"model": "claude-haiku-4-5"},
endpoint_type=EndpointType.ANTHROPIC,
start_time=datetime.datetime.now(timezone.utc),
all_chunks=list(ANTHROPIC_MESSAGES_SSE_CHUNKS),
end_time=datetime.datetime.now(timezone.utc),
)
assert logged["result"].id == "msg_01Lit6806Streaming"
assert (
_spend_log_request_id(
response_obj=logged["result"],
kwargs={
**logged["kwargs"],
"call_type": "anthropic_messages",
"litellm_call_id": "6806cafe-0000-4000-8000-000000000001",
"litellm_params": {"metadata": {"user_api_key": "test-key"}},
},
)
== "msg_01Lit6806Streaming"
)
def test_spend_log_request_id_still_falls_back_to_litellm_call_id_without_a_provider_id():
"""
Anthropic-compatible upstreams that omit `id` must keep landing on litellm_call_id rather
than on a fresh chatcmpl- uuid nobody can look up.
"""
logging_obj = _anthropic_messages_logging_obj(stream=True)
logging_obj.model_call_details["stream"] = True
from litellm.proxy.pass_through_endpoints.llm_provider_handlers.anthropic_passthrough_logging_handler import (
AnthropicPassthroughLoggingHandler,
)
AnthropicPassthroughLoggingHandler._create_anthropic_response_logging_payload(
litellm_model_response=litellm.ModelResponse(id="chatcmpl-generated"),
model="claude-haiku-4-5",
kwargs={},
start_time=datetime.datetime.now(timezone.utc),
end_time=datetime.datetime.now(timezone.utc),
logging_obj=logging_obj,
)
assert logging_obj.model_call_details["complete_streaming_response"].id == (
"6806cafe-0000-4000-8000-000000000001"
)
def test_spend_log_request_id_for_chat_completions_is_untouched():
"""
/v1/chat/completions callers look their rows up by the chatcmpl- id in the response body.
"""
assert (
_spend_log_request_id(
response_obj=litellm.ModelResponse(id="chatcmpl-EJvWIw3DAhuKYuwp3jJI4Pnhp2vjv", choices=[]),
kwargs={
"call_type": "acompletion",
"model": "gpt-5.6",
"litellm_call_id": "6806cafe-0000-4000-8000-000000000002",
"litellm_params": {"metadata": {"user_api_key": "test-key"}},
},
)
== "chatcmpl-EJvWIw3DAhuKYuwp3jJI4Pnhp2vjv"
)
def test_spend_log_request_id_is_the_response_id_a_bridged_messages_caller_received():
"""
/v1/messages against a non-Anthropic model answers with the Responses id the caller then
looks their row up by, so the row must not fall back to a fresh chatcmpl- uuid.
"""
from litellm.types.llms.openai import ResponseAPIUsage, ResponsesAPIResponse
logging_obj = _anthropic_messages_logging_obj(stream=False)
bridged_response = ResponsesAPIResponse(
id="resp_01Lit6806Bridged",
object="response",
created_at=1767225600,
model="gpt-5.6",
status="completed",
output=[
{
"id": "msg_bridged_output",
"type": "message",
"role": "assistant",
"status": "completed",
"content": [{"type": "output_text", "text": "delta", "annotations": []}],
}
],
usage=ResponseAPIUsage(input_tokens=13, output_tokens=5, total_tokens=18),
)
logged_response = logging_obj._handle_anthropic_messages_response_logging(result=bridged_response)
assert logged_response.id == "resp_01Lit6806Bridged"
assert (
_spend_log_request_id(
response_obj=logged_response,
kwargs={
"call_type": "anthropic_messages",
"model": "gpt-5.6",
"litellm_call_id": "6806cafe-0000-4000-8000-000000000003",
"litellm_params": {"metadata": {"user_api_key": "test-key"}},
},
)
== "resp_01Lit6806Bridged"
)

View file

@ -68,12 +68,10 @@ async def test_success_hook_attaches_original_response_on_block():
user_api_key_dict = UserAPIKeyAuth(api_key="test", request_route="/chat/completions")
data = {"guardrail_to_apply": guardrail, "model": "gpt-4o"}
# Inject our translation for the inferred call type (the module global is
# cached across tests, so patch it directly rather than the loader).
with patch.object(
ug,
"endpoint_guardrail_translation_mappings",
{
"load_guardrail_translation_mappings",
lambda: {
CallTypes.acompletion: lambda: translation,
CallTypes.completion: lambda: translation,
},

View file

@ -1540,8 +1540,8 @@ class TestCommonRequestProcessingHelpers:
expected_error_data = {
"error": {
"message": "Error processing stream start",
"type": "None",
"param": "None",
"type": "internal_server_error",
"param": None,
"code": str(status.HTTP_500_INTERNAL_SERVER_ERROR),
}
}
@ -1569,8 +1569,8 @@ class TestCommonRequestProcessingHelpers:
expected_error_data = {
"error": {
"message": "Content blocked by guardrail",
"type": "None",
"param": "None",
"type": "invalid_request_error",
"param": None,
"code": "400",
}
}
@ -1934,6 +1934,104 @@ class TestCommonRequestProcessingHelpers:
assert mock_tracer.trace.call_count == 0
def _stringified_none_paths(node: object, path: str = "error") -> tuple[str, ...]:
if isinstance(node, dict):
return tuple(
found
for key, value in node.items()
for found in _stringified_none_paths(value, f"{path}.{key}")
)
if isinstance(node, (list, tuple)):
return tuple(
found
for index, value in enumerate(node)
for found in _stringified_none_paths(value, f"{path}[{index}]")
)
return (path,) if node == "None" else ()
def _blocked_guardrail_exception() -> HTTPException:
return HTTPException(
status_code=400,
detail={
"error": "Violated guardrail policy",
"bedrock_guardrail_response": {"action": "GUARDRAIL_INTERVENED"},
"guardrailIdentifier": "gf3sc1mzinjw",
"guardrailVersion": "DRAFT",
},
)
class TestGuardrailBlockErrorPayloadNeverStringifiesNone:
"""Regression for LIT-6808: a blocked-guardrail error body carried the literal string
"None" for type and param instead of a real error type and JSON null."""
def test_non_streaming_block_payload_carries_a_real_type_and_null_param(self):
from litellm.proxy.common_request_processing import (
proxy_exception_from_http_exception,
)
payload = json.loads(
json.dumps(proxy_exception_from_http_exception(_blocked_guardrail_exception(), {}).to_dict())
)
assert _stringified_none_paths(payload) == ()
assert payload["type"] == "invalid_request_error"
assert payload["param"] is None
assert payload["code"] == "400"
assert payload["message"] == "Violated guardrail policy"
def test_streaming_block_frame_carries_a_real_type_and_null_param(self):
from litellm.proxy.common_request_processing import sse_error_payload
error_status, error_obj = sse_error_payload(_blocked_guardrail_exception())
frame = json.loads(json.dumps({"error": dict(error_obj)}))
assert error_status == 400
assert _stringified_none_paths(frame["error"]) == ()
assert frame["error"]["type"] == "invalid_request_error"
assert frame["error"]["param"] is None
assert frame["error"]["code"] == "400"
@pytest.mark.parametrize(
"status_code, expected_type",
[
(400, "invalid_request_error"),
(401, "authentication_error"),
(403, "permission_error"),
(404, "invalid_request_error"),
(429, "rate_limit_error"),
(500, "internal_server_error"),
(503, "internal_server_error"),
],
)
def test_status_code_decides_the_type_when_the_exception_carries_none(self, status_code, expected_type):
from litellm.proxy.common_request_processing import (
proxy_exception_from_http_exception,
)
payload = proxy_exception_from_http_exception(
HTTPException(status_code=status_code, detail="blocked"), {}
).to_dict()
assert payload["type"] == expected_type
assert payload["param"] is None
def test_a_type_and_param_the_exception_carries_win_over_the_fallback(self):
from litellm.proxy.common_request_processing import (
proxy_exception_from_http_exception,
)
exc = HTTPException(status_code=400, detail="unknown model")
exc.type = "authentication_error"
exc.param = "model"
payload = proxy_exception_from_http_exception(exc, {}).to_dict()
assert payload["type"] == "authentication_error"
assert payload["param"] == "model"
class TestExtractErrorFromSSEChunk:
"""Tests for _extract_error_from_sse_chunk function"""
@ -2999,6 +3097,25 @@ class TestHandleLLMApiExceptionDictDetail:
assert proxy_exc.message == "Content blocked by guardrail"
assert proxy_exc.provider_specific_fields is None
async def test_blocked_guardrail_error_body_never_carries_the_string_none(self):
"""Regression for LIT-6808: the error body a blocked request returns must carry a real
error type and JSON null rather than the literal string "None"."""
proxy_exc = await self._invoke(_blocked_guardrail_exception())
payload = json.loads(json.dumps(proxy_exc.to_dict()))
assert _stringified_none_paths(payload) == ()
assert payload["type"] == "invalid_request_error"
assert payload["param"] is None
async def test_unclassified_exception_error_body_never_carries_the_string_none(self):
"""The same holds on the generic fallback, where nothing carries a type at all."""
proxy_exc = await self._invoke(ValueError("Something broke"))
payload = json.loads(json.dumps(proxy_exc.to_dict()))
assert _stringified_none_paths(payload) == ()
assert payload["type"] == "internal_server_error"
assert payload["param"] is None
async def test_not_found_error_preserves_404(self):
"""NotFoundError with status_code=404 should map to ProxyException code=404."""
from litellm.exceptions import NotFoundError

View file

@ -388,6 +388,72 @@ async def test_aquery_does_not_forward_connection_override_keys_to_search():
assert not (blocked & set(search_kwargs.keys()))
@pytest.mark.asyncio
async def test_aquery_forwards_vector_store_params_to_search_but_not_completion():
"""
Regression for LIT-6773: the server-trusted vector_store_params (a managed
store's litellm_params) must reach the search call wholesale, including the
connection keys the caller allowlist blocks, while the caller's own
retrieval_config overrides stay blocked, the caller's top-level api_key and
api_base stay on the completion only, and the completion never inherits the
store's connection params.
"""
from unittest.mock import AsyncMock
from litellm.types.vector_stores import VectorStoreSearchResponse
fake_search = AsyncMock(
return_value=VectorStoreSearchResponse(
object="vector_store.search_results.page", search_query="q", data=[]
)
)
fake_completion = AsyncMock(
return_value=ModelResponse(
id="chatcmpl-test",
choices=[{"index": 0, "message": {"role": "assistant", "content": "hi"}, "finish_reason": "stop"}],
model="gpt-4o-mini",
)
)
with (
patch("litellm.vector_stores.asearch", new=fake_search), # test-quality-ok: the search boundary under test
patch("litellm.acompletion", new=fake_completion), # test-quality-ok: the completion boundary under test
):
await litellm.aquery(
model="gpt-4o-mini",
messages=[{"role": "user", "content": "hello"}],
api_key="sk-llm-key",
api_base="https://llm.example.com",
retrieval_config={
"vector_store_id": "customer_kb",
"custom_llm_provider": "milvus",
"api_base": "https://attacker.example.com",
"api_key": "attacker-key",
},
vector_store_params={
"vector_store_id": "customer_kb",
"custom_llm_provider": "milvus",
"api_base": "http://127.0.0.1:19530",
"api_key": "root:Milvus",
"milvus_text_field": "book_intro_text",
"outputFields": ["book_intro_text"],
},
)
fake_search.assert_awaited_once()
search_kwargs = fake_search.await_args.kwargs
assert search_kwargs["vector_store_id"] == "customer_kb"
assert search_kwargs["custom_llm_provider"] == "milvus"
assert search_kwargs["api_base"] == "http://127.0.0.1:19530"
assert search_kwargs["api_key"] == "root:Milvus"
assert search_kwargs["milvus_text_field"] == "book_intro_text"
assert search_kwargs["outputFields"] == ["book_intro_text"]
fake_completion.assert_awaited_once()
completion_kwargs = fake_completion.await_args.kwargs
assert completion_kwargs["api_key"] == "sk-llm-key"
assert completion_kwargs["api_base"] == "https://llm.example.com"
assert not ({"milvus_text_field", "outputFields"} & set(completion_kwargs))
def test_rag_call_types_are_registered():
"""
query/aquery/ingest/aingest are @client-decorated entry points, so their

View file

@ -0,0 +1,55 @@
from pathlib import Path
from typing import Final
import pytest
from pydantic import TypeAdapter
from litellm import cost_per_token, get_model_info
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
REPO_ROOT: Final = Path(__file__).parents[2]
MODEL: Final = "azure_ai/grok-4.6"
SOURCE: Final = (
"https://techcommunity.microsoft.com/blog/azure-ai-foundry-blog/"
"grok-4-6-comes-to-microsoft-foundry-models-built-for-long-horizon-reasoning-and-/4547578"
)
COST_MAP_ADAPTER: Final = TypeAdapter(dict[str, dict[str, object]])
def _cost_map_entry(path: Path) -> dict[str, object]:
return COST_MAP_ADAPTER.validate_json(path.read_bytes())[MODEL]
@pytest.mark.usefixtures("local_model_cost_map")
def test_azure_ai_grok_4_6_is_priced_and_routed() -> None:
routed_model, provider, _, _ = get_llm_provider(model=MODEL)
assert (routed_model, provider) == ("grok-4.6", "azure_ai")
info = get_model_info(model=routed_model, custom_llm_provider=provider)
assert info["litellm_provider"] == "azure_ai"
assert info["mode"] == "chat"
assert info["input_cost_per_token"] == 2e-06
assert info["output_cost_per_token"] == 6e-06
assert info["cache_read_input_token_cost"] == 5e-07
assert info["max_input_tokens"] == 200000
assert info["max_output_tokens"] == 128000
assert info["max_tokens"] == 128000
assert info["supports_function_calling"] is True
assert info["supports_prompt_caching"] is True
assert info["supports_reasoning"] is True
assert info["supports_response_schema"] is True
assert info["supports_tool_choice"] is True
assert info["supports_vision"] is True
assert info["supports_web_search"] is True
prompt_cost, completion_cost = cost_per_token(model=MODEL, prompt_tokens=1_000_000, completion_tokens=1_000_000)
assert prompt_cost == pytest.approx(2.0)
assert completion_cost == pytest.approx(6.0)
def test_azure_ai_grok_4_6_entry_source_and_backup_match() -> None:
main_entry = _cost_map_entry(REPO_ROOT / "model_prices_and_context_window.json")
backup_entry = _cost_map_entry(REPO_ROOT / "litellm" / "model_prices_and_context_window_backup.json")
assert main_entry["source"] == SOURCE
assert backup_entry == main_entry

View file

@ -9275,9 +9275,11 @@ class _FallbackAttemptRecorder(CustomLogger):
def __init__(self):
super().__init__()
self.failed_targets = []
self.breadcrumbs_per_target = []
async def log_failure_fallback_event(self, original_model_group, kwargs, original_exception):
self.failed_targets.append(kwargs.get("model"))
self.breadcrumbs_per_target.append(kwargs.get("metadata", {}).get("previous_models", ()))
def _cyclic_fallback_router(num_retries=0):
@ -9348,14 +9350,16 @@ async def test_retry_breadcrumbs_do_not_carry_the_walk_state():
A retry has to be configured for the walk state to reach log_retry at all."""
router = _cyclic_fallback_router(num_retries=1)
capture = _LogCapture(logging.ERROR)
recorder = _FallbackAttemptRecorder()
await _drive_cyclic_fallback(router, capture)
await _drive_cyclic_fallback(router, capture, recorder)
assert router.previous_models, "no retry breadcrumbs were recorded"
breadcrumbs = [breadcrumb for hop in recorder.breadcrumbs_per_target for breadcrumb in hop]
assert breadcrumbs, "no retry breadcrumbs were recorded"
assert any(
"fallback_depth" in breadcrumb for breadcrumb in router.previous_models
"fallback_depth" in breadcrumb for breadcrumb in breadcrumbs
), "no breadcrumb carried router walk state, so this test cannot see the leak"
for breadcrumb in router.previous_models:
for breadcrumb in breadcrumbs:
assert "attempted_targets" not in breadcrumb
@ -9393,15 +9397,94 @@ async def test_retry_breadcrumbs_never_carry_a_forwarded_credential(container_ke
container still reaches the breadcrumb, but the raw secret never does, whatever key holds it."""
router = _cyclic_fallback_router(num_retries=1)
capture = _LogCapture(logging.ERROR)
metadata = {}
await _drive_cyclic_fallback(router, capture, **request_kwargs)
await _drive_cyclic_fallback(router, capture, metadata=metadata, **request_kwargs)
assert router.previous_models, "no retry breadcrumbs were recorded"
dumped = json.dumps(router.previous_models, default=str)
breadcrumbs = metadata["previous_models"]
assert breadcrumbs, "no retry breadcrumbs were recorded"
dumped = json.dumps(breadcrumbs, default=str)
assert container_key in dumped, "the credential-bearing kwarg never reached the breadcrumb, so this test cannot see the leak"
assert _BREADCRUMB_CREDENTIAL_CANARY not in dumped
def _always_failing_router(num_retries):
return litellm.Router(
model_list=[
{
"model_name": "broken-group",
"litellm_params": {
"model": "openai/gpt-4o-mini",
"api_key": "sk-fake",
"mock_response": "litellm.InternalServerError",
},
}
],
num_retries=num_retries,
)
async def _fail_one_proxy_shaped_request(router, request_marker):
"""The proxy hands the router a metadata dict and a proxy_server_request whose body is a
shallow copy of the request, so body["metadata"] is the very same dict the router later
stamps previous_models onto."""
metadata = {"request_marker": request_marker}
with pytest.raises(litellm.InternalServerError):
await router.acompletion(
model="broken-group",
messages=[{"role": "user", "content": "hi"}],
metadata=metadata,
proxy_server_request={
"url": "http://localhost:4000/v1/chat/completions",
"method": "POST",
"headers": {},
"body": {"model": "broken-group", "metadata": metadata},
},
)
return metadata["previous_models"]
def _nested_breadcrumb_lists(node):
if isinstance(node, dict):
return [v for k, v in node.items() if k == "previous_models"] + [
found for v in node.values() for found in _nested_breadcrumb_lists(v)
]
if isinstance(node, (list, tuple)):
return [found for item in node for found in _nested_breadcrumb_lists(item)]
return []
@pytest.mark.asyncio
async def test_retry_breadcrumbs_stay_per_request_and_flat_across_failing_requests():
"""Every failed attempt appends a breadcrumb to metadata["previous_models"], and the proxy's
request snapshot aliases that same metadata dict. Kept on the Router and copied wholesale,
each breadcrumb embedded every earlier one from every earlier request, so the breadcrumb
tree, and with it the debug repr of the kwargs, roughly doubled on each failed attempt until
a single-worker proxy spent minutes in the redaction regex and stopped answering."""
router = _always_failing_router(num_retries=2)
breadcrumbs_per_request = [
await _fail_one_proxy_shaped_request(router, f"request-{request_number}") for request_number in range(1, 7)
]
for request_number, breadcrumbs in enumerate(breadcrumbs_per_request, start=1):
assert len(breadcrumbs) == 3, "one initial attempt plus two retries failed, each leaving one breadcrumb"
assert {breadcrumb["metadata"]["request_marker"] for breadcrumb in breadcrumbs} == {f"request-{request_number}"}
for breadcrumb in breadcrumbs:
assert _nested_breadcrumb_lists(breadcrumb) == []
assert len({len(repr(breadcrumbs)) for breadcrumbs in breadcrumbs_per_request}) == 1
@pytest.mark.asyncio
async def test_retry_breadcrumbs_keep_only_the_last_four_attempts():
router = _always_failing_router(num_retries=6)
breadcrumbs = await _fail_one_proxy_shaped_request(router, "request-1")
assert len(breadcrumbs) == 4
assert [breadcrumb["metadata"]["attempted_retries"] for breadcrumb in breadcrumbs] == [3, 4, 5, 6]
@pytest.mark.asyncio
async def test_fallback_traceback_stays_available_at_debug_level():
"""Dropping the stack from the ERROR line is only safe because the fallback path still

View file

@ -40,6 +40,7 @@ from litellm.utils import (
_is_streaming_request,
_snapshot_exception_for_hook,
async_post_call_failure_deployment_hook,
async_post_call_success_deployment_hook,
client,
get_llm_provider,
get_non_default_completion_params,
@ -5808,6 +5809,90 @@ class TestHuggingFaceConfigFetch:
assert request_timeout["read"] == HF_CONFIG_FETCH_TIMEOUT_SECONDS
@pytest.mark.asyncio
async def test_success_deployment_hook_chains_past_callback_returning_response(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""Regression (LIT-5863): the dispatcher must run every callback, chaining each non-None
result into the next call, instead of returning at the first callback answering non-None.
A guardrail answering with the unmodified response used to starve every callback after it."""
from litellm.types.utils import ModelResponse
original = ModelResponse()
replacement = ModelResponse()
class PassthroughLogger(CustomLogger):
async def async_post_call_success_deployment_hook(self, request_data, response, call_type):
return response
class ReplacingLogger(CustomLogger):
def __init__(self) -> None:
super().__init__()
self.seen: list = []
async def async_post_call_success_deployment_hook(self, request_data, response, call_type):
self.seen.append(response)
return replacement
class ObservingLogger(CustomLogger):
def __init__(self) -> None:
super().__init__()
self.seen: list = []
async def async_post_call_success_deployment_hook(self, request_data, response, call_type):
self.seen.append(response)
return None
replacer = ReplacingLogger()
observer = ObservingLogger()
monkeypatch.setattr(litellm, "callbacks", [PassthroughLogger(), replacer, observer])
result = await async_post_call_success_deployment_hook(
request_data={}, response=original, call_type=CallTypes.acompletion
)
assert replacer.seen == [original]
assert observer.seen == [replacement]
assert result is replacement
@pytest.mark.asyncio
async def test_registered_guardrail_does_not_starve_vector_store_search_results(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""Regression (LIT-5863): with any guardrail registered ahead of the lazily-appended
VectorStorePreCallHook, /v1/chat/completions responses lost
provider_specific_fields["search_results"] because the guardrail answered the unmodified
response and the dispatcher stopped there."""
from types import SimpleNamespace
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.integrations.vector_store_integrations.vector_store_pre_call_hook import (
VectorStorePreCallHook,
)
from litellm.types.utils import ModelResponse
search_results: Final = [{"search_query": "coolant", "data": [{"content": [{"text": "Cryoline-9", "type": "text"}]}]}]
logging_obj = SimpleNamespace(model_call_details={"search_results": search_results})
response = ModelResponse(choices=[{"message": {"role": "assistant", "content": "Cryoline-9"}}])
monkeypatch.setattr(
litellm,
"callbacks",
[CustomGuardrail(guardrail_name="dummy-guardrail"), VectorStorePreCallHook()],
)
result = await async_post_call_success_deployment_hook(
request_data={"litellm_logging_obj": logging_obj},
response=response,
call_type=CallTypes.acompletion,
)
provider_fields = result.choices[0].message.provider_specific_fields
assert provider_fields is not None
assert provider_fields["search_results"] == search_results
class TestIsVisionExplicitlyDisabled:
"""github_copilot and chatgpt run an OAuth device flow inside get_llm_provider; the
explicit-disable lookup must adopt the declared prefix instead of resolving it, exactly
@ -5836,3 +5921,135 @@ class TestIsVisionExplicitlyDisabled:
is_vision_explicitly_disabled("fireworks_ai/accounts/fireworks/models/deepseek-v4-flash-0731") is True
)
assert is_vision_explicitly_disabled("anthropic/claude-sonnet-4-5") is False
class TestVerboseRequestLineRedaction:
"""`litellm.set_verbose = True` echoes the caller's kwargs back as a `litellm.completion(...)`
line on stdout, so a credential kwarg lands in whatever collects stdout: a terminal, a
container log drain, a CI job log. Credential-named kwargs must not survive that echo,
at any nesting depth, while ordinary params still must, or the line stops telling the
developer what they called."""
FAKE_API_KEY: Final = "sk-fake-lit6823-0000000000000000"
def _verbose_request_line(self, capsys, monkeypatch, **kwargs) -> str:
monkeypatch.setattr(litellm, "set_verbose", True)
monkeypatch.setattr("litellm._logging.set_verbose", True)
capsys.readouterr()
litellm.completion(
model="gpt-3.5-turbo",
messages=[{"role": "user", "content": "hello"}],
mock_response="hi",
**kwargs,
)
captured: Final = capsys.readouterr()
return "\n".join(line for line in (captured.out + captured.err).splitlines() if "litellm.completion(" in line)
def test_api_key_never_reaches_the_request_line(self, capsys, monkeypatch):
printed: Final = self._verbose_request_line(capsys, monkeypatch, api_key=self.FAKE_API_KEY)
assert "litellm.completion(" in printed
assert self.FAKE_API_KEY not in printed
assert "api_key='REDACTED'" in printed
def test_credential_headers_never_reach_the_request_line(self, capsys, monkeypatch):
printed: Final = self._verbose_request_line(
capsys,
monkeypatch,
api_key=self.FAKE_API_KEY,
extra_headers={"Authorization": "Bearer fake-lit6823-header", "x-request-id": "abc123"},
)
assert "fake-lit6823-header" not in printed
assert "'Authorization': 'REDACTED'" in printed
assert "'x-request-id': 'abc123'" in printed
def test_credentials_nested_in_a_list_never_reach_the_request_line(self, capsys, monkeypatch):
printed: Final = self._verbose_request_line(
capsys,
monkeypatch,
api_key=self.FAKE_API_KEY,
extra_body={"providers": [{"name": "openai", "api_key": "sk-fake-lit6823-nested"}]},
)
assert "sk-fake-lit6823-nested" not in printed
assert "'name': 'openai'" in printed
def test_ordinary_params_still_printed(self, capsys, monkeypatch):
printed: Final = self._verbose_request_line(
capsys, monkeypatch, api_key=self.FAKE_API_KEY, max_tokens=17, temperature=0.25
)
assert "model='gpt-3.5-turbo'" in printed
assert "max_tokens=17" in printed
assert "temperature=0.25" in printed
class TestFinalOptionalParamsLineRedaction:
"""A verbose run echoes the fully built optional params too, and `extra_body` carries whatever the
caller nested inside it straight onto that line, so a credential tucked in there lands in a terminal
or a log drain in plaintext. It has to be redacted on both surfaces `print_verbose` writes to, and the
line has to keep printing on both, because `litellm.set_verbose` and the DEBUG logger are independent
switches and neither implies the other."""
FAKE_NESTED_KEY: Final = "sk-fake-lit6835-nested-0000000000"
def _complete(self, **kwargs) -> None:
litellm.completion(
model="gpt-3.5-turbo",
messages=[{"role": "user", "content": "hello"}],
mock_response="hi",
**kwargs,
)
def _printed_line(self, capsys) -> str:
captured: Final = capsys.readouterr()
return "\n".join(
line for line in (captured.out + captured.err).splitlines() if "Final returned optional params" in line
)
def test_nested_credential_is_redacted_when_only_set_verbose_is_on(self, capsys, caplog, monkeypatch):
monkeypatch.setattr(litellm, "set_verbose", True)
with caplog.at_level(logging.WARNING, logger=verbose_logger.name):
capsys.readouterr()
self._complete(extra_body={"providers": [{"name": "openai", "api_key": self.FAKE_NESTED_KEY}]})
printed: Final = self._printed_line(capsys)
assert printed
assert self.FAKE_NESTED_KEY not in printed
assert "'api_key': 'REDACTED'" in printed
assert "'name': 'openai'" in printed
def test_line_still_reaches_the_logger_when_only_the_debug_logger_is_on(self, capsys, caplog, monkeypatch):
monkeypatch.setattr(litellm, "set_verbose", False)
with caplog.at_level(logging.DEBUG, logger=verbose_logger.name):
self._complete(extra_body={"providers": [{"name": "openai", "api_key": self.FAKE_NESTED_KEY}]})
logged: Final = "\n".join(
record.getMessage()
for record in caplog.records
if "Final returned optional params" in record.getMessage()
)
assert logged
assert self.FAKE_NESTED_KEY not in logged
assert "'name': 'openai'" in logged
def test_nothing_is_emitted_when_neither_verbose_switch_is_on(self, capsys, caplog, monkeypatch):
monkeypatch.setattr(litellm, "set_verbose", False)
with caplog.at_level(logging.WARNING, logger=verbose_logger.name):
capsys.readouterr()
self._complete(extra_body={"providers": [{"name": "openai", "api_key": self.FAKE_NESTED_KEY}]})
captured: Final = capsys.readouterr()
assert "Final returned optional params" not in captured.out + captured.err
assert self.FAKE_NESTED_KEY not in captured.out + captured.err
def test_ordinary_optional_params_still_reach_the_line(self, capsys, caplog, monkeypatch):
monkeypatch.setattr(litellm, "set_verbose", True)
with caplog.at_level(logging.WARNING, logger=verbose_logger.name):
capsys.readouterr()
self._complete(max_tokens=17, temperature=0.25)
printed: Final = self._printed_line(capsys)
assert "'max_tokens': 17" in printed
assert "'temperature': 0.25" in printed

View file

@ -3,7 +3,7 @@
"limit": 22328
},
"LIT002": {
"limit": 26763
"limit": 26760
},
"LIT003": {
"limit": 261
@ -30,7 +30,7 @@
"limit": 16470
},
"LIT011": {
"limit": 5517
"limit": 5516
},
"LIT012": {
"limit": 4489

View file

@ -44070,7 +44070,11 @@ export interface operations {
};
list_containers_containers_get: {
parameters: {
query?: never;
query?: {
after?: string | null;
limit?: number | null;
order?: string | null;
};
header?: never;
path?: never;
cookie?: never;
@ -44086,6 +44090,15 @@ export interface operations {
"application/json": unknown;
};
};
/** @description Validation Error */
422: {
headers: {
[name: string]: unknown;
};
content: {
"application/json": components["schemas"]["HTTPValidationError"];
};
};
};
};
create_container_containers_post: {
@ -61114,7 +61127,11 @@ export interface operations {
};
list_containers_v1_containers_get: {
parameters: {
query?: never;
query?: {
after?: string | null;
limit?: number | null;
order?: string | null;
};
header?: never;
path?: never;
cookie?: never;
@ -61130,6 +61147,15 @@ export interface operations {
"application/json": unknown;
};
};
/** @description Validation Error */
422: {
headers: {
[name: string]: unknown;
};
content: {
"application/json": components["schemas"]["HTTPValidationError"];
};
};
};
};
create_container_v1_containers_post: {