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

# Conflicts:
#	tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py
This commit is contained in:
mateo-berri 2026-09-03 15:50:55 -07:00
commit 055eaee1f9
152 changed files with 5930 additions and 1654 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

@ -18,9 +18,11 @@ caller's identity metadata, minus two things that must never be forwarded as-is:
from __future__ import annotations
from collections.abc import Mapping
from types import MappingProxyType
from typing import Final
from litellm.constants import INTERNAL_CALL_ORIGIN_METADATA_KEY, NON_INFERENCE_CALL_TYPES
from litellm.litellm_core_utils.initialize_dynamic_callback_params import initialize_standard_callback_dynamic_params
from litellm.types.utils import BACKGROUND_RESPONSE_COST_POLL_CALL_ORIGIN, InternalCallOrigin
BUDGET_RESERVATION_METADATA_KEYS: Final = frozenset({"user_api_key_budget_reservation"})
@ -142,6 +144,19 @@ def forwarded_internal_call_metadata(
}
def parent_session_kwargs(request_kwargs: Mapping[str, object] | None) -> Mapping[str, str]:
kwargs: Final = request_kwargs or MappingProxyType({})
return MappingProxyType(
{k: v for k in ("litellm_session_id", "litellm_trace_id") if isinstance(v := kwargs.get(k), str)}
)
def effective_turn_off_message_logging(request_kwargs: Mapping[str, object] | None) -> bool | None:
return initialize_standard_callback_dynamic_params(dict(request_kwargs) if request_kwargs else None).get(
"turn_off_message_logging"
)
def sanitized_forwardable_call_metadata(
parent_metadata: Mapping[str, object],
call_origin: InternalCallOrigin,

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,6 +7,7 @@ from litellm.exceptions import UnsupportedParamsError
from litellm.llms.openai.chat.gpt_5_transformation import (
OpenAIGPT5Config,
_get_effort_level,
is_gpt_reasoning_series_name,
)
from litellm.types.llms.openai import AllMessageValues
@ -35,26 +36,7 @@ class AzureOpenAIGPT5Config(AzureOpenAIConfig, OpenAIGPT5Config):
@classmethod
def is_model_gpt_5_model(cls, model: str) -> bool:
"""Check if the Azure model string refers to a gpt-5 variant.
Accepts both explicit gpt-5 model names and the ``gpt5_series/`` prefix
used for manual routing.
"""
# The gpt-5-chat* family (gpt-5-chat, gpt-5-chat-latest, gpt-5-chat-2025-08-07,
# …) are regular chat models: they support temperature and tool_choice but NOT
# reasoning_effort. They must NOT be routed through the GPT-5 reasoning path.
#
# Versioned chat models such as gpt-5.3-chat and gpt-5.1-chat ARE reasoning
# models and must stay on the GPT-5 path. The distinguishing feature is that
# the gpt-5-chat family has a literal "-chat" immediately after "gpt-5"
# (i.e. "gpt-5-chat…"), while versioned chat models interpose a minor version
# number (i.e. "gpt-5.<digit>-chat").
#
# Using a startswith("gpt-5-chat") prefix check on the normalized name (rather
# than a substring check) makes this boundary explicit and avoids any ambiguity
# if future model names coincidentally contain "gpt-5-chat" as an interior run.
_normalized: Final = model.split("/")[-1] # strip provider prefix, e.g. "azure/"
return ("gpt-5" in model and not _normalized.startswith("gpt-5-chat")) or "gpt5_series" in model
return is_gpt_reasoning_series_name(model) or "gpt5_series" in model
def get_supported_openai_params(self, model: str) -> list[str]:
"""Get supported parameters for Azure OpenAI GPT-5 models.

View file

@ -14,6 +14,7 @@ from litellm.litellm_core_utils.prompt_templates.factory import (
convert_to_azure_openai_messages,
)
from litellm.llms.base_llm.chat.transformation import BaseLLMException
from litellm.llms.openai.chat.gpt_5_transformation import GPT_REASONING_SERIES_MARKERS
from litellm.types.llms.azure import (
API_VERSION_MONTH_SUPPORTED_RESPONSE_FORMAT,
API_VERSION_YEAR_SUPPORTED_RESPONSE_FORMAT,
@ -139,7 +140,7 @@ class AzureOpenAIConfig(BaseConfig):
name family needs the rename, including the ``gpt-5-chat*`` models that are excluded from
the reasoning path by https://github.com/BerriAI/litellm/issues/13781.
"""
return "gpt-5" in model or "gpt5_series" in model
return any(marker in model for marker in GPT_REASONING_SERIES_MARKERS) or "gpt5_series" in model
def _is_response_format_supported_model(self, model: str) -> bool:
"""

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

@ -91,6 +91,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,
@ -8926,17 +8927,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,
@ -9002,17 +9005,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,
@ -9092,17 +9097,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,
@ -9169,17 +9176,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,
@ -9257,17 +9266,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,
@ -9334,17 +9345,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,
@ -9422,17 +9435,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,
@ -9499,17 +9514,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,
@ -9591,17 +9608,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,
@ -9670,17 +9689,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,
@ -9756,17 +9777,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,
@ -9832,17 +9855,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

@ -61,6 +61,14 @@ def _get_effort_level(value: str | dict | None) -> str | None:
return None
GPT_REASONING_SERIES_MARKERS: Final = ("gpt-5", "gpt-6")
def is_gpt_reasoning_series_name(model: str) -> bool:
normalized: Final = model.split("/")[-1]
return any(marker in model for marker in GPT_REASONING_SERIES_MARKERS) and not normalized.startswith("gpt-5-chat")
class OpenAIGPT5Config(OpenAIGPTConfig):
"""Configuration for gpt-5 models including GPT-5-Codex variants.
@ -73,21 +81,7 @@ class OpenAIGPT5Config(OpenAIGPTConfig):
@classmethod
def is_model_gpt_5_model(cls, model: str) -> bool:
# The gpt-5-chat* family (gpt-5-chat, gpt-5-chat-latest, gpt-5-chat-2025-08-07,
# …) are regular chat models: they support temperature and tool_choice but NOT
# reasoning_effort. They must NOT be routed through the GPT-5 reasoning path.
#
# Versioned chat models such as gpt-5.3-chat and gpt-5.1-chat ARE reasoning
# models and must stay on the GPT-5 path. The distinguishing feature is that
# the gpt-5-chat family has a literal "-chat" immediately after "gpt-5"
# (i.e. "gpt-5-chat…"), while versioned chat models interpose a minor version
# number (i.e. "gpt-5.<digit>-chat").
#
# Using a startswith("gpt-5-chat") prefix check on the normalized name (rather
# than a substring check) makes this boundary explicit and avoids any ambiguity
# if future model names coincidentally contain "gpt-5-chat" as an interior run.
_normalized: Final = model.split("/")[-1] # strip provider prefix, e.g. "openai/"
return "gpt-5" in model and not _normalized.startswith("gpt-5-chat")
return is_gpt_reasoning_series_name(model)
@classmethod
def is_model_gpt_5_search_model(cls, model: str) -> bool:
@ -122,6 +116,8 @@ class OpenAIGPT5Config(OpenAIGPTConfig):
def is_model_gpt_5_4_plus_model(cls, model: str) -> bool:
"""Check if the model is gpt-5.4 or newer (5.4, 5.5, 5.6, etc., including pro)."""
model_name: Final = model.split("/")[-1]
if model_name.startswith("gpt-6"):
return True
if not model_name.startswith("gpt-5."):
return False
try:

View file

@ -16,6 +16,7 @@ from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response impo
)
from litellm.litellm_core_utils.url_utils import encode_url_path_segment
from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig
from litellm.llms.openai.chat.gpt_5_transformation import is_gpt_reasoning_series_name
from litellm.responses.litellm_completion_transformation.custom_tools import TOOL_CALL_ITEM_ID_PREFIX_BY_TYPE
from litellm.secret_managers.main import get_secret_str
from litellm.types.llms.openai import *
@ -89,7 +90,7 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
parts: Final = model.split("/")
if len(parts) > 1 and parts[0] not in ("openai",):
return False
return "gpt-5" in model and "gpt-5-chat" not in model
return is_gpt_reasoning_series_name(model)
@staticmethod
def _supports_reasoning_effort_none(model: str) -> bool:

View file

@ -998,6 +998,16 @@ def replace_project_and_location_in_route(requested_route: str, vertex_project:
return modified_route
def _api_version_for_route(requested_route: str) -> Literal["v1", "v1beta1"]:
return "v1beta1" if "cachedContent" in requested_route else "v1"
def _with_api_version(requested_route: str) -> str:
if not requested_route.startswith("/projects/"):
return requested_route
return f"/{_api_version_for_route(requested_route)}{requested_route}"
def construct_target_url(
base_url: str,
requested_route: str,
@ -1017,18 +1027,19 @@ def construct_target_url(
new_base_url: Final = httpx.URL(base_url)
if "locations" in requested_route: # contains the target project id + location
if vertex_project and vertex_location:
requested_route = replace_project_and_location_in_route(requested_route, vertex_project, vertex_location)
return new_base_url.copy_with(path=requested_route)
targeted_route: Final = (
replace_project_and_location_in_route(requested_route, vertex_project, vertex_location)
if vertex_project and vertex_location
else requested_route
)
return new_base_url.copy_with(path=_with_api_version(targeted_route))
"""
- Add endpoint version (e.g. v1beta for cachedContent, v1 for rest)
- Add default project id
- Add default location
"""
vertex_version: Literal["v1", "v1beta1"] = "v1"
if "cachedContent" in requested_route:
vertex_version = "v1beta1"
vertex_version: Literal["v1", "v1beta1"] = _api_version_for_route(requested_route)
# Check if the requested route starts with a version
# e.g. /v1beta1/publishers/google/models/gemini-3-pro-preview:streamGenerateContent

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

@ -654,23 +654,16 @@ class InMemoryGuardrailHandler:
source: Literal["db", "config"] = "db",
) -> None:
"""
Update a guardrail in memory
- updates the guardrail in memory
- updates the guardrail params in litellm.callback_manager
Update a guardrail in memory: a changed name or litellm_params rebuilds the
live callback from the new row (fail-closed: an invalid row keeps the
previous instance and raises), anything else only refreshes the stored row
"""
self.IN_MEMORY_GUARDRAILS[guardrail_id] = guardrail
self._sources[guardrail_id] = source
tracked_callbacks: Final = self._tracked_callbacks(guardrail_id)
if not tracked_callbacks:
updated_guardrail: Final = cast(Guardrail, {**guardrail, "guardrail_id": guardrail_id})
if self._has_guardrail_params_changed(guardrail_id, updated_guardrail):
self.reinitialize_guardrail(guardrail=updated_guardrail, source=source)
return
updated_litellm_params: Final = cast(LitellmParams, guardrail.get("litellm_params", {}))
tracked_callbacks[0].update_in_memory_litellm_params(litellm_params=updated_litellm_params)
for sibling_callback in tracked_callbacks[1:]:
sibling_stage = sibling_callback.event_hook
sibling_callback.update_in_memory_litellm_params(litellm_params=updated_litellm_params)
sibling_callback.event_hook = sibling_stage
self.IN_MEMORY_GUARDRAILS[guardrail_id] = updated_guardrail
self._sources[guardrail_id] = source
def delete_in_memory_guardrail(self, guardrail_id: str) -> None:
"""

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

@ -17,6 +17,7 @@ import json
import traceback
from collections.abc import Awaitable, Mapping, Sequence
from datetime import datetime, timezone
from types import MappingProxyType
from typing import Any, Final, Literal, Protocol, cast, overload
import fastapi
@ -77,6 +78,7 @@ from litellm.types.proxy.management_endpoints.internal_user_endpoints import (
BulkUpdateUserRequest,
BulkUpdateUserResponse,
UserListResponse,
UserSearchWhere,
UserUpdateResult,
)
from litellm.types.proxy.management_endpoints.scim_v2 import (
@ -2080,6 +2082,22 @@ async def _authorize_user_list_request(
return ",".join(allowed_org_ids)
_NO_SEARCH_WHERE: Final[Mapping[str, object]] = MappingProxyType({})
def _user_search_where(search: str | None) -> Mapping[str, object]:
"""Prisma predicate for `/user/list?search=`: user_id or user_email contains it, case-insensitive."""
if not search:
return _NO_SEARCH_WHERE
search_where: Final[UserSearchWhere] = {
"OR": (
{"user_id": {"contains": search, "mode": "insensitive"}},
{"user_email": {"contains": search, "mode": "insensitive"}},
)
}
return search_where
@router.get(
"/user/list",
tags=["Internal User management"],
@ -2091,6 +2109,10 @@ async def get_users(
user_ids: str | None = fastapi.Query(default=None, description="Get list of users by user_ids"),
sso_user_ids: str | None = fastapi.Query(default=None, description="Get list of users by sso_user_id"),
user_email: str | None = fastapi.Query(default=None, description="Filter users by partial email match"),
search: str | None = fastapi.Query(
default=None,
description="Combined search: matches users whose 'user_id' or 'user_email' contains the value (case-insensitive).",
),
team: str | None = fastapi.Query(default=None, description="Filter users by team id"),
page: int = fastapi.Query(default=1, ge=1, description="Page number"),
page_size: int = fastapi.Query(default=25, ge=1, le=100, description="Number of items per page"),
@ -2121,6 +2143,8 @@ async def get_users(
Get list of users by sso_ids. Comma separated list of sso_ids.
user_email: Optional[str]
Filter users by partial email match
search: Optional[str]
Combined search: matches users whose user_id or user_email contains the value (case-insensitive)
team: Optional[str]
Filter users by team id. Will match if user has this team in their teams array.
page: int
@ -2197,7 +2221,11 @@ async def get_users(
where_conditions["organization_memberships"] = {"some": {"organization_id": {"in": org_id_list}}}
## Filter any none fastapi.Query params - e.g. where_conditions: {'user_email': {'contains': Query(None), 'mode': 'insensitive'}, 'teams': {'has': Query(None)}}
where_conditions = {k: v for k, v in where_conditions.items() if v is not None}
where: Final[Mapping[str, object]] = {
key: value
for key, value in (*where_conditions.items(), *_user_search_where(search).items())
if value is not None
}
# Build order_by conditions
@ -2206,14 +2234,14 @@ async def get_users(
)
users: Final[Sequence[prisma_models.LiteLLM_UserTable]] = await UserRepository(prisma_client).table.find_many(
where=where_conditions,
where=where,
skip=skip,
take=page_size,
order=(order_by if order_by else {"created_at": "desc"}), # Default to created_at desc if no sort specified
)
# Get total count of user rows
total_count: Final[int] = await UserRepository(prisma_client).table.count(where=where_conditions)
total_count: Final[int] = await UserRepository(prisma_client).table.count(where=where)
# Get key count for each user
user_key_counts: Final = await get_user_key_counts(prisma_client, [user.user_id for user in users])

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

@ -29,6 +29,12 @@ from litellm.proxy._types import (
from litellm.proxy.auth.auth_utils import is_request_body_safe
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_utils.path_utils import safe_filename
from litellm.proxy.prompts.prompt_registry import (
DEFAULT_PROMPT_ENVIRONMENT,
get_base_prompt_id,
get_version_number,
prompt_environment_or_default,
)
from litellm.repositories.table_repositories import PromptRepository
from litellm.types.prompts.init_prompts import (
ListPromptsResponse,
@ -102,165 +108,20 @@ def _prompt_table(prisma_client: "PrismaClient") -> _PromptTableActions:
return PromptRepository(prisma_client).table
def get_base_prompt_id(prompt_id: str) -> str:
"""
Extract the base prompt ID by stripping the version suffix if present.
Args:
prompt_id: Prompt ID that may include version suffix (e.g., "jack_success.v1" or "jack_success_v1")
Returns:
Base prompt ID without version suffix (e.g., "jack_success")
Examples:
>>> get_base_prompt_id("jack_success.v1")
"jack_success"
>>> get_base_prompt_id("jack_success_v1")
"jack_success"
>>> get_base_prompt_id("jack_success")
"jack_success"
"""
# Try dot separator first (.v)
if ".v" in prompt_id:
return prompt_id.split(".v")[0]
# Try underscore separator (_v)
if "_v" in prompt_id:
return prompt_id.split("_v")[0]
return prompt_id
def get_version_number(prompt_id: str) -> int:
"""
Extract the version number from a versioned prompt ID.
Args:
prompt_id: Prompt ID that may include version suffix (e.g., "jack_success.v2" or "jack_success_v2")
Returns:
Version number (defaults to 1 if no version suffix or invalid format)
Examples:
>>> get_version_number("jack_success.v2")
2
>>> get_version_number("jack_success_v2")
2
>>> get_version_number("jack_success")
1
"""
# Try dot separator first (.v)
if ".v" in prompt_id:
version_str = prompt_id.split(".v")[1]
try:
return int(version_str)
except ValueError:
pass
# Try underscore separator (_v)
if "_v" in prompt_id:
version_str = prompt_id.split("_v")[1]
try:
return int(version_str)
except ValueError:
pass
return 1
def construct_versioned_prompt_id(prompt_id: str, version: int | None = None) -> str:
"""
Construct a versioned prompt ID from a base prompt_id and version number.
Args:
prompt_id: Base prompt ID (e.g., "jack_success")
version: Version number (if None, returns the base prompt_id unchanged)
Returns:
Versioned prompt ID (e.g., "jack_success.v4")
Examples:
>>> construct_versioned_prompt_id("jack_success", 4)
"jack_success.v4"
>>> construct_versioned_prompt_id("jack_success", None)
"jack_success"
>>> construct_versioned_prompt_id("jack_success.v2", 4)
"jack_success.v4"
"""
if version is None:
return prompt_id
# Strip any existing version suffix first
base_id: Final = get_base_prompt_id(prompt_id)
return f"{base_id}.v{version}"
def get_latest_version_prompt_id(prompt_id: str, all_prompt_ids: Mapping[str, object]) -> str:
"""
Find the latest version of a prompt from available prompt IDs.
Args:
prompt_id: Base prompt ID or versioned prompt ID (e.g., "jack_success" or "jack_success.v2")
all_prompt_ids: Dictionary of all available prompt IDs (keys are prompt IDs)
Returns:
The prompt ID with the highest version number, or the original prompt_id if no versions exist
Examples:
>>> all_ids = {"jack.v1": {}, "jack.v2": {}, "jack.v3": {}}
>>> get_latest_version_prompt_id("jack", all_ids)
"jack.v3"
>>> get_latest_version_prompt_id("jack.v1", all_ids)
"jack.v3"
>>> all_ids = {"simple": {}}
>>> get_latest_version_prompt_id("simple", all_ids)
"simple"
"""
base_id: Final = get_base_prompt_id(prompt_id=prompt_id)
# Find all versions of this prompt
matching_versions: Final = []
for stored_prompt_id in all_prompt_ids:
if get_base_prompt_id(prompt_id=stored_prompt_id) == base_id:
version_num = get_version_number(prompt_id=stored_prompt_id)
matching_versions.append((version_num, stored_prompt_id))
# Use the highest version number
if matching_versions:
matching_versions.sort(reverse=True)
return matching_versions[0][1]
else:
# No versioned prompts found, use the base ID as-is
return prompt_id
def get_latest_prompt_versions(prompts: list[PromptSpec]) -> list[PromptSpec]:
"""
Filter a list of prompts to return only the latest version of each unique prompt.
Args:
prompts: List of PromptSpec objects
Returns:
List of PromptSpec objects with only the latest version of each prompt
Filter prompts down to the latest version per (base prompt id, environment).
"""
latest_prompts: Final[dict[str, PromptSpec]] = {}
for prompt in prompts:
base_id = get_base_prompt_id(prompt_id=prompt.prompt_id)
version = get_version_number(prompt_id=prompt.prompt_id)
# Keep the prompt with the highest version number
if base_id not in latest_prompts:
latest_prompts[base_id] = prompt
else:
existing_version = get_version_number(prompt_id=latest_prompts[base_id].prompt_id)
if version > existing_version:
latest_prompts[base_id] = prompt
sorted_prompts: Final = sorted(prompts, key=lambda prompt: get_version_number(prompt_id=prompt.prompt_id))
latest_prompts: Final = {
(get_base_prompt_id(prompt_id=prompt.prompt_id), prompt_environment_or_default(prompt.environment)): prompt
for prompt in sorted_prompts
}
return list(latest_prompts.values())
async def get_next_version_for_prompt(
prisma_client: "PrismaClient", prompt_id: str, environment: str = "development"
prisma_client: "PrismaClient", prompt_id: str, environment: str = DEFAULT_PROMPT_ENVIRONMENT
) -> int:
"""
Get the next version number for a prompt in a specific environment.
@ -403,11 +264,14 @@ async def list_prompts(
if key_metadata is not None:
prompts: Final = cast(list[str] | None, key_metadata.get("prompts", None))
if prompts is not None:
all_prompts = [
IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS[prompt_id]
for prompt_id in prompts
if prompt_id in IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS
allowed_prompt_ids: Final = frozenset(prompts)
allowed_prompts: Final = [
spec
for spec in IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS.values()
if spec.prompt_id in allowed_prompt_ids
or get_base_prompt_id(prompt_id=spec.prompt_id) in allowed_prompt_ids
]
all_prompts = get_latest_prompt_versions(prompts=allowed_prompts)
if environment:
all_prompts = [p for p in all_prompts if p.environment == environment]
prompt_list: Final = []
@ -576,7 +440,7 @@ def _get_prompt_template(prompt_spec: PromptSpec, base_prompt_id: str) -> Prompt
metadata=parsed.get("metadata"),
)
else:
prompt_callback: Final = IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_by_id(prompt_spec.prompt_id)
prompt_callback: Final = IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_for_prompt(prompt=prompt_spec)
if prompt_callback is not None:
integration_name: Final = prompt_callback.integration_name
if integration_name == "dotprompt":
@ -690,15 +554,10 @@ async def get_prompt_info(
if env_prompts:
prompt_spec = create_versioned_prompt_spec(db_prompt=env_prompts[0])
# Fallback: use in-memory registry (no environment filter)
if prompt_spec is None and environment is None:
prompt_spec = IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id(prompt_id)
if prompt_spec is None:
latest_prompt_id: Final = get_latest_version_prompt_id(
prompt_id=prompt_id,
all_prompt_ids=IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS,
)
prompt_spec = IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id(latest_prompt_id)
if prompt_spec is None:
prompt_spec = IN_MEMORY_PROMPT_REGISTRY.resolve_prompt_spec(
prompt_id, version=requested_version, environment=environment
)
if prompt_spec is None:
raise HTTPException(
@ -785,7 +644,7 @@ async def create_prompt(
environment: Final = (
request.prompt_info.environment
if request.prompt_info and request.prompt_info.environment
else "development"
else DEFAULT_PROMPT_ENVIRONMENT
)
# Get next version number
@ -885,7 +744,7 @@ async def update_prompt(
environment: Final = (
request.prompt_info.environment
if request.prompt_info and request.prompt_info.environment
else "development"
else DEFAULT_PROMPT_ENVIRONMENT
)
# Check if any version of this prompt exists (in any environment)
@ -897,9 +756,7 @@ async def update_prompt(
detail=f"Prompt with ID {base_prompt_id} not found",
)
# Check if it's a config prompt
existing_in_memory: Final = IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id(prompt_id)
if existing_in_memory and existing_in_memory.prompt_info.prompt_type == "config":
if IN_MEMORY_PROMPT_REGISTRY.has_config_prompt(base_prompt_id=base_prompt_id):
raise HTTPException(
status_code=400,
detail="Cannot update config prompts.",
@ -988,40 +845,26 @@ async def delete_prompt(
raise HTTPException(status_code=500, detail=CommonProxyErrors.db_not_connected_error.value)
try:
# Try to get prompt directly first
existing_prompt = IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id(prompt_id)
# If not found, try to find the latest version
if existing_prompt is None:
latest_prompt_id: Final = get_latest_version_prompt_id(
prompt_id=prompt_id,
all_prompt_ids=IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS,
)
existing_prompt = IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id(latest_prompt_id)
# Use the resolved prompt_id for deletion
prompt_id = latest_prompt_id
base_prompt_id: Final = get_base_prompt_id(prompt_id=prompt_id)
existing_prompt: Final = IN_MEMORY_PROMPT_REGISTRY.resolve_prompt_spec(prompt_id, environment=environment)
if existing_prompt is None:
raise HTTPException(status_code=404, detail=f"Prompt with ID {prompt_id} not found")
if existing_prompt.prompt_info.prompt_type == "config":
if IN_MEMORY_PROMPT_REGISTRY.has_config_prompt(base_prompt_id=base_prompt_id):
raise HTTPException(
status_code=400,
detail="Cannot delete config prompts.",
)
# Get the base prompt ID (without version suffix) for database deletion
base_prompt_id: Final = get_base_prompt_id(prompt_id=prompt_id)
# Build delete filter; scope to environment if provided
delete_where: Final[dict[str, str]] = {"prompt_id": base_prompt_id}
if environment:
delete_where["environment"] = environment
# Delete versions from the database (scoped to environment if provided)
delete_where: Final[dict[str, str]] = {
"prompt_id": base_prompt_id,
**({"environment": environment} if environment else {}),
}
await _prompt_table(prisma_client).delete_many(where=delete_where)
IN_MEMORY_PROMPT_REGISTRY.delete_prompts_by_base_id(base_prompt_id, environment=environment or None)
IN_MEMORY_PROMPT_REGISTRY.delete_prompts_by_base_id(
base_prompt_id=base_prompt_id, environment=environment or None
)
env_msg: Final = f" from {environment}" if environment else ""
return {"message": f"Prompt {base_prompt_id} deleted successfully{env_msg}"}
@ -1093,7 +936,7 @@ async def patch_prompt(
try:
# Resolve the target row: find the latest version in the given environment
base_prompt_id: Final = get_base_prompt_id(prompt_id=prompt_id)
env: Final = environment or "development"
env: Final = prompt_environment_or_default(environment)
requested_version: Final = get_version_number(prompt_id=prompt_id) if prompt_id != base_prompt_id else None
# Build query to find the exact row by composite unique key
@ -1117,11 +960,7 @@ async def patch_prompt(
target_row: Final = db_rows[0]
# Check if prompt exists in memory
versioned_id: Final = f"{base_prompt_id}.v{target_row.version}"
existing_prompt: Final = IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id(versioned_id)
if existing_prompt and existing_prompt.prompt_info.prompt_type == "config":
if IN_MEMORY_PROMPT_REGISTRY.has_config_prompt(base_prompt_id=base_prompt_id):
raise HTTPException(
status_code=400,
detail="Cannot update config prompts.",

View file

@ -1,6 +1,6 @@
import importlib
import os
from collections.abc import Callable
from collections.abc import Callable, Sequence
from pathlib import Path
from typing import Final
@ -14,6 +14,87 @@ from litellm.types.prompts.init_prompts import (
prompt_initializer_registry = {}
DEFAULT_PROMPT_ENVIRONMENT: Final = "development"
PROMPT_ENVIRONMENT_SERVE_PRECEDENCE: Final = ("production", "staging", "development")
def get_base_prompt_id(prompt_id: str) -> str:
"""
Extract the base prompt ID by stripping the version suffix if present.
Examples:
>>> get_base_prompt_id("jack_success.v1")
"jack_success"
>>> get_base_prompt_id("jack_success_v1")
"jack_success"
>>> get_base_prompt_id("jack_success")
"jack_success"
"""
if ".v" in prompt_id:
return prompt_id.split(".v")[0]
if "_v" in prompt_id:
return prompt_id.split("_v")[0]
return prompt_id
def get_version_number(prompt_id: str) -> int:
"""
Extract the version number from a versioned prompt ID (defaults to 1).
Examples:
>>> get_version_number("jack_success.v2")
2
>>> get_version_number("jack_success_v2")
2
>>> get_version_number("jack_success")
1
"""
if ".v" in prompt_id:
version_str = prompt_id.split(".v")[1]
try:
return int(version_str)
except ValueError:
pass
if "_v" in prompt_id:
version_str = prompt_id.split("_v")[1]
try:
return int(version_str)
except ValueError:
pass
return 1
def prompt_environment_or_default(environment: str | None) -> str:
return environment or DEFAULT_PROMPT_ENVIRONMENT
def registry_key_for_prompt(prompt: PromptSpec) -> str:
return f"{prompt.prompt_id}::{prompt_environment_or_default(prompt.environment)}"
def parse_prompt_version(raw_version: object) -> int | None:
if isinstance(raw_version, bool):
return None
if isinstance(raw_version, int):
return raw_version
if isinstance(raw_version, str) and raw_version.isdigit():
return int(raw_version)
return None
def _spec_version(prompt: PromptSpec) -> int:
return prompt.version if prompt.version is not None else get_version_number(prompt_id=prompt.prompt_id)
def _default_serve_environment(prompts: Sequence[PromptSpec]) -> str:
present: Final = frozenset(prompt_environment_or_default(prompt.environment) for prompt in prompts)
ladder_pick: Final = next((env for env in PROMPT_ENVIRONMENT_SERVE_PRECEDENCE if env in present), None)
if ladder_pick is not None:
return ladder_pick
return min(present) if present else DEFAULT_PROMPT_ENVIRONMENT
def get_prompt_initializer_from_integrations():
"""
@ -113,17 +194,16 @@ class InMemoryPromptRegistry:
"""
import litellm
prompt_id: Final = prompt.prompt_id
if prompt_id in self.IN_MEMORY_PROMPTS:
verbose_proxy_logger.debug("prompt_id already exists in IN_MEMORY_PROMPTS")
return self.IN_MEMORY_PROMPTS[prompt_id]
registry_key: Final = registry_key_for_prompt(prompt)
if registry_key in self.IN_MEMORY_PROMPTS:
verbose_proxy_logger.debug("prompt already exists in IN_MEMORY_PROMPTS")
return self.IN_MEMORY_PROMPTS[registry_key]
parsed_prompt, custom_prompt_callback = self._build_prompt_callback(prompt=prompt)
litellm.logging_callback_manager.add_litellm_callback(custom_prompt_callback)
# store references to the prompt in memory
self.IN_MEMORY_PROMPTS[prompt_id] = parsed_prompt
self.prompt_id_to_custom_prompt[prompt_id] = custom_prompt_callback
self.IN_MEMORY_PROMPTS[registry_key] = parsed_prompt
self.prompt_id_to_custom_prompt[registry_key] = custom_prompt_callback
return parsed_prompt
@ -166,68 +246,93 @@ class InMemoryPromptRegistry:
import litellm
parsed_prompt, new_callback = self._build_prompt_callback(prompt=prompt)
stale_callback: Final = self.prompt_id_to_custom_prompt.pop(prompt.prompt_id, None)
self.IN_MEMORY_PROMPTS.pop(prompt.prompt_id, None)
registry_key: Final = registry_key_for_prompt(parsed_prompt)
stale_callback: Final = self.prompt_id_to_custom_prompt.pop(registry_key, None)
self.IN_MEMORY_PROMPTS.pop(registry_key, None)
if stale_callback is not None:
litellm.logging_callback_manager.remove_callback_from_all_lists(stale_callback)
litellm.logging_callback_manager.add_litellm_callback(new_callback)
self.IN_MEMORY_PROMPTS[prompt.prompt_id] = parsed_prompt
self.prompt_id_to_custom_prompt[prompt.prompt_id] = new_callback
self.IN_MEMORY_PROMPTS[registry_key] = parsed_prompt
self.prompt_id_to_custom_prompt[registry_key] = new_callback
return parsed_prompt
def sync_prompt_from_db(self, prompt: PromptSpec) -> PromptSpec | None:
existing: Final = self.IN_MEMORY_PROMPTS.get(prompt.prompt_id)
existing: Final = self.IN_MEMORY_PROMPTS.get(registry_key_for_prompt(prompt))
if existing is None:
return self.initialize_prompt(prompt=prompt)
if existing.litellm_params == prompt.litellm_params and existing.prompt_info == prompt.prompt_info:
return existing
return self.reload_prompt(prompt=prompt)
def get_prompt_by_id(self, prompt_id: str) -> PromptSpec | None:
def resolve_prompt_spec(
self,
prompt_id: str,
version: int | None = None,
environment: str | None = None,
) -> PromptSpec | None:
"""
Get a prompt by its ID from memory
"""
return self.IN_MEMORY_PROMPTS.get(prompt_id)
Resolve a prompt spec by base prompt id, optional version, and optional environment.
def get_prompt_callback_by_id(self, prompt_id: str) -> CustomPromptManagement | None:
With no environment, resolves within the default serve environment
(production > staging > development > alphabetical first present).
With no version, resolves to the highest version in the chosen environment.
"""
Get a prompt callback by its ID from memory
"""
return self.prompt_id_to_custom_prompt.get(prompt_id)
base_prompt_id: Final = get_base_prompt_id(prompt_id=prompt_id)
base_matches: Final = tuple(
spec
for spec in self.IN_MEMORY_PROMPTS.values()
if get_base_prompt_id(prompt_id=spec.prompt_id) == base_prompt_id
)
if not base_matches:
return None
resolved_environment: Final = (
environment if environment is not None else _default_serve_environment(base_matches)
)
env_matches: Final = tuple(
spec for spec in base_matches if prompt_environment_or_default(spec.environment) == resolved_environment
)
if not env_matches:
return None
if version is not None:
return next((spec for spec in env_matches if _spec_version(spec) == version), None)
return max(env_matches, key=_spec_version)
def remove_prompt(self, prompt_id: str) -> None:
def get_prompt_callback_for_prompt(self, prompt: PromptSpec) -> CustomPromptManagement | None:
return self.prompt_id_to_custom_prompt.get(registry_key_for_prompt(prompt))
def has_config_prompt(self, base_prompt_id: str) -> bool:
return any(
spec.prompt_info.prompt_type == "config"
for spec in self.IN_MEMORY_PROMPTS.values()
if get_base_prompt_id(prompt_id=spec.prompt_id) == base_prompt_id
)
def remove_prompt(self, registry_key: str) -> None:
import litellm
self.IN_MEMORY_PROMPTS.pop(prompt_id, None)
stale_callback: Final = self.prompt_id_to_custom_prompt.pop(prompt_id, None)
self.IN_MEMORY_PROMPTS.pop(registry_key, None)
stale_callback: Final = self.prompt_id_to_custom_prompt.pop(registry_key, None)
if stale_callback is not None:
litellm.logging_callback_manager.remove_callback_from_all_lists(stale_callback)
def delete_prompts_by_base_id(self, base_prompt_id: str, environment: str | None = None) -> list[str]:
"""
Delete all prompts matching the given base prompt ID from memory, along with their
registered callbacks; scoped to one environment when given.
Delete matching prompts from memory, along with their registered callbacks,
scoped to one environment when given.
Args:
base_prompt_id: The base prompt ID (without version suffix)
environment: When set, only delete prompts deployed to this environment
Returns:
List of prompt IDs that were deleted
Returns the registry keys that were deleted.
"""
from litellm.proxy.prompts.prompt_endpoints import get_base_prompt_id
prompts_to_delete: Final = [
pid
for pid, prompt in self.IN_MEMORY_PROMPTS.items()
if get_base_prompt_id(prompt_id=pid) == base_prompt_id
and (environment is None or prompt.environment == environment)
keys_to_delete: Final = [
key
for key, spec in self.IN_MEMORY_PROMPTS.items()
if get_base_prompt_id(prompt_id=spec.prompt_id) == base_prompt_id
and (environment is None or prompt_environment_or_default(spec.environment) == environment)
]
for pid in prompts_to_delete:
self.remove_prompt(prompt_id=pid)
for key in keys_to_delete:
self.remove_prompt(registry_key=key)
return prompts_to_delete
return keys_to_delete
IN_MEMORY_PROMPT_REGISTRY: Final = InMemoryPromptRegistry()

View file

@ -7594,7 +7594,7 @@ class ProxyConfig:
return create_versioned_prompt_spec(db_prompt=db_prompt)
async def _init_prompts_in_db(self, prisma_client: PrismaClient):
from litellm.proxy.prompts.prompt_registry import IN_MEMORY_PROMPT_REGISTRY
from litellm.proxy.prompts.prompt_registry import IN_MEMORY_PROMPT_REGISTRY, registry_key_for_prompt
from litellm.types.prompts.init_prompts import PromptSpec
def parse_row(db_prompt: object) -> PromptSpec | None:
@ -7609,21 +7609,12 @@ class ProxyConfig:
return None
try:
prompt_ids_loaded_before_db_read: Final = frozenset(IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS)
registry_keys_loaded_before_db_read: Final = frozenset(IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS)
prompts_in_db: Final[Sequence[object]] = await PromptRepository(prisma_client).table.find_many()
parsed_specs: Final[tuple[PromptSpec, ...]] = tuple(
spec for row in prompts_in_db if (spec := parse_row(row)) is not None
)
newest_spec_per_id: Final[Mapping[str, PromptSpec]] = MappingProxyType(
{
spec.prompt_id: spec
for spec in sorted(
parsed_specs,
key=lambda s: s.updated_at.timestamp() if s.updated_at else float("-inf"),
)
}
)
for prompt_spec in newest_spec_per_id.values():
for prompt_spec in parsed_specs:
try:
IN_MEMORY_PROMPT_REGISTRY.sync_prompt_from_db(prompt=prompt_spec)
except Exception as prompt_sync_error: # noqa: BLE001 # one poisoned row must not block syncing the remaining prompts
@ -7635,15 +7626,16 @@ class ProxyConfig:
# An unparsable row still exists in the DB, so skip the sweep rather than unload its in-memory copy
every_row_parsed: Final = len(parsed_specs) == len(prompts_in_db)
if every_row_parsed:
deleted_db_prompt_ids: Final = tuple(
prompt_id
for prompt_id in prompt_ids_loaded_before_db_read
if (loaded_spec := IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS.get(prompt_id)) is not None
db_registry_keys: Final = frozenset(registry_key_for_prompt(spec) for spec in parsed_specs)
deleted_db_registry_keys: Final = tuple(
registry_key
for registry_key in registry_keys_loaded_before_db_read
if (loaded_spec := IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS.get(registry_key)) is not None
and loaded_spec.prompt_info.prompt_type == "db"
and prompt_id not in newest_spec_per_id
and registry_key not in db_registry_keys
)
for deleted_prompt_id in deleted_db_prompt_ids:
IN_MEMORY_PROMPT_REGISTRY.remove_prompt(prompt_id=deleted_prompt_id)
for deleted_registry_key in deleted_db_registry_keys:
IN_MEMORY_PROMPT_REGISTRY.remove_prompt(registry_key=deleted_registry_key)
except Exception as e:
verbose_proxy_logger.debug("litellm.proxy.proxy_server.py::ProxyConfig:_init_prompts_in_db - %s", e)

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

@ -1478,28 +1478,27 @@ class ProxyLogging:
) -> None:
"""Process prompt template if applicable."""
from litellm.proxy.prompts.prompt_endpoints import (
construct_versioned_prompt_id,
get_latest_version_prompt_id,
)
from litellm.proxy.prompts.prompt_registry import IN_MEMORY_PROMPT_REGISTRY
from litellm.responses.utils import ResponsesAPIRequestUtils
from litellm.utils import get_non_default_completion_params
if prompt_version is None:
lookup_prompt_id = get_latest_version_prompt_id(
prompt_id=prompt_id,
all_prompt_ids=IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS,
)
else:
lookup_prompt_id = construct_versioned_prompt_id(prompt_id=prompt_id, version=prompt_version)
custom_logger: Final = IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_by_id(lookup_prompt_id)
prompt_spec: Final = IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id(lookup_prompt_id)
raw_prompt_environment: Final = data.get("prompt_environment", None)
prompt_environment: Final = raw_prompt_environment if isinstance(raw_prompt_environment, str) else None
prompt_spec: Final = IN_MEMORY_PROMPT_REGISTRY.resolve_prompt_spec(
prompt_id,
version=prompt_version,
environment=prompt_environment,
)
custom_logger: Final = (
IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_for_prompt(prompt=prompt_spec)
if prompt_spec is not None
else None
)
litellm_prompt_id: str | None = None
if prompt_spec is not None:
litellm_prompt_id = prompt_spec.litellm_params.prompt_id
data.pop("prompt_id", None)
data.pop("prompt_environment", None)
if custom_logger and prompt_spec is not None:
is_responses_call: Final = call_type == "aresponses"
@ -1542,6 +1541,7 @@ class ProxyLogging:
data.pop("prompt_variables", None)
data.pop("prompt_label", None)
data.pop("prompt_version", None)
data.pop("prompt_environment", None)
def _process_guardrail_metadata(self, data: dict) -> None:
"""Process guardrails from metadata and add to applied_guardrails."""
@ -1750,7 +1750,6 @@ class ProxyLogging:
litellm_logging_obj: Final = cast(Optional["LiteLLMLoggingObj"], data.get("litellm_logging_obj", None))
prompt_id: Final[str | None] = data.get("prompt_id", None)
prompt_version: Final[int | None] = data.get("prompt_version", None)
## PROMPT TEMPLATE CHECK ##
@ -1760,11 +1759,13 @@ class ProxyLogging:
and prompt_id is not None
and (call_type == "completion" or call_type == "acompletion" or call_type == "aresponses")
):
from litellm.proxy.prompts.prompt_registry import parse_prompt_version
await self._process_prompt_template(
data=data,
litellm_logging_obj=litellm_logging_obj,
prompt_id=prompt_id,
prompt_version=prompt_version,
prompt_version=parse_prompt_version(data.get("prompt_version", None)),
call_type=call_type,
)

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

@ -595,16 +595,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:
@ -965,7 +969,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 {}
@ -8144,35 +8147,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

@ -2,23 +2,41 @@
Auto-Routing Strategy that works with a Semantic Router Config
"""
from collections.abc import Mapping
from typing import TYPE_CHECKING, Any, Final, Optional
from pydantic import BaseModel, ConfigDict
from litellm._logging import verbose_router_logger
from litellm.constants import DEFAULT_AUTO_ROUTER_MAX_INPUT_CHARS
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.internal_call_metadata import (
effective_turn_off_message_logging,
forwarded_internal_call_metadata,
parent_session_kwargs,
)
from litellm.types.utils import AUTOROUTER_CLASSIFIER_CALL_ORIGIN
if TYPE_CHECKING:
from semantic_router.routers import SemanticRouter
from semantic_router.routers.base import Route
from litellm.router import Router
from litellm.router_strategy.auto_router.litellm_encoder import LiteLLMRouterEncoder
from litellm.types.router import PreRoutingHookResponse
else:
Router = Any
PreRoutingHookResponse = Any
Route = Any
SemanticRouter = Any
LiteLLMRouterEncoder = Any
class _CallerMetadata(BaseModel):
model_config = ConfigDict(frozen=True, extra="ignore")
metadata: Mapping[str, object] | None = None
litellm_metadata: Mapping[str, object] | None = None
class AutoRouter(CustomLogger):
@ -50,6 +68,8 @@ class AutoRouter(CustomLogger):
"""
from semantic_router.routers import SemanticRouter
from litellm.router_strategy.auto_router.litellm_encoder import LiteLLMRouterEncoder
self.auto_router_config_path: str | None = auto_router_config_path
self.auto_router_config: str | None = auto_router_config
self.auto_sync_value = self.DEFAULT_AUTO_SYNC_VALUE
@ -59,6 +79,11 @@ class AutoRouter(CustomLogger):
self.embedding_model: str = embedding_model
self.max_input_chars: int = max_input_chars
self.litellm_router_instance: Router = litellm_router_instance
self.encoder: LiteLLMRouterEncoder = LiteLLMRouterEncoder(
litellm_router_instance=litellm_router_instance,
model_name=embedding_model,
max_input_chars=max_input_chars,
)
def _load_semantic_routing_routes(self) -> list[Route]:
from semantic_router.routers import SemanticRouter
@ -129,9 +154,6 @@ class AutoRouter(CustomLogger):
from semantic_router.routers import SemanticRouter
from litellm.litellm_core_utils.prompt_templates.factory import resolve_structured_messages
from litellm.router_strategy.auto_router.litellm_encoder import (
LiteLLMRouterEncoder,
)
from litellm.types.router import PreRoutingHookResponse
resolved_messages: Final = (
@ -149,34 +171,47 @@ class AutoRouter(CustomLogger):
#######################
routelayer = SemanticRouter(
routes=self.loaded_routes,
encoder=LiteLLMRouterEncoder(
litellm_router_instance=self.litellm_router_instance,
model_name=self.embedding_model,
max_input_chars=self.max_input_chars,
),
encoder=self.encoder,
auto_sync=self.auto_sync_value,
)
self.routelayer = routelayer
message_content: Final = self._extract_text_from_messages(resolved_messages)
route_name: Final = self._matched_route_name(routelayer, message_content)
route_name: Final = await self._matched_route_name(routelayer, message_content, request_kwargs)
return PreRoutingHookResponse(
model=route_name or self.default_model,
messages=messages,
)
def _matched_route_name(self, routelayer: "SemanticRouter", text: str) -> str | None:
async def _matched_route_name(
self, routelayer: "SemanticRouter", text: str, request_kwargs: Mapping[str, object]
) -> str | None:
"""Name of the route `text` matches, or None when nothing matched or the match failed.
The route layer embeds `text` to compare it against the routes, and that embedding call can
`text` is embedded here rather than by `routelayer(text=...)` so the caller's metadata reaches
`aembedding()` and the embedding's spend lands on the key/team that sent the request;
SemanticRouter has no way to pass kwargs through to its encoder. That embedding call can
fail (context limit, timeout, provider error). Choosing a model is a routing decision, so a
failure here falls back to the default model rather than failing the user's request.
"""
from semantic_router.schema import RouteChoice
try:
route_choice: Final = routelayer(text=text)
caller: Final = _CallerMetadata.model_validate(request_kwargs)
query_vector: Final = (
await self.encoder.aencode_queries(
[text],
metadata=forwarded_internal_call_metadata(caller.metadata, AUTOROUTER_CLASSIFIER_CALL_ORIGIN),
litellm_metadata=forwarded_internal_call_metadata(
caller.litellm_metadata, AUTOROUTER_CLASSIFIER_CALL_ORIGIN
),
proxy_server_request={"body": {"model": self.embedding_model, "input": [text]}},
turn_off_message_logging=effective_turn_off_message_logging(request_kwargs),
**parent_session_kwargs(request_kwargs),
)
)[0]
route_choice: Final = await routelayer.acall(vector=query_vector)
except Exception as e: # noqa: BLE001 -- the embedding call behind the route layer can fail many ways (context limit, timeout, provider/network error); none of them may fail the request
verbose_router_logger.warning(
"AutoRouter: semantic routing failed (%s), falling back to default model %s", e, self.default_model

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

@ -1,6 +1,8 @@
from typing import Any, Final
from collections.abc import Mapping
from typing import Any, Final, Literal
from pydantic import BaseModel, field_validator
from typing_extensions import ReadOnly, TypedDict
from litellm.proxy._types import (
LiteLLM_UserTableWithKeyCount,
@ -9,6 +11,17 @@ from litellm.proxy._types import (
)
class InsensitiveContains(TypedDict):
contains: ReadOnly[str]
mode: ReadOnly[Literal["insensitive"]]
class UserSearchWhere(TypedDict):
"""Prisma filter behind `/user/list?search=`: user_id or user_email contains the term, case-insensitive."""
OR: ReadOnly[tuple[Mapping[Literal["user_id", "user_email"], InsensitiveContains], ...]]
class UserListResponse(BaseModel):
"""
Response model for the user list endpoint

View file

@ -3630,6 +3630,7 @@ all_litellm_params = (
"litellm_system_prompt",
"provider_specific_header",
"prompt_version",
"prompt_environment",
"api_base",
"force_timeout",
"logger_fn",

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,
@ -7461,7 +7473,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

@ -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

@ -1,6 +1,13 @@
import time
from collections.abc import Iterator
from typing import Final
import httpx
from openai import OpenAI, BadRequestError, NotFoundError, APIStatusError
import pytest
from openai import APIStatusError, BadRequestError, NotFoundError, OpenAI, Stream
from openai.types.responses import ResponseStreamEvent
BACKGROUND_STREAM_ADMISSION_DEADLINE_SECONDS: Final = 90
def generate_key():
@ -153,43 +160,48 @@ def test_cancel_response():
raise e
def admitted_response_id(chunk: ResponseStreamEvent) -> str | None:
response: Final = getattr(chunk, "response", None)
return None if response is None else response.id
def events_until_admission(stream: Stream[ResponseStreamEvent], started: float) -> Iterator[ResponseStreamEvent]:
for chunk in stream:
print("stream chunk=", chunk)
yield chunk
if admitted_response_id(chunk) is not None:
return
if time.monotonic() - started > BACKGROUND_STREAM_ADMISSION_DEADLINE_SECONDS:
return
def test_cancel_streaming_response():
try:
client = get_test_client()
from litellm.types.llms.openai import ResponsesAPIResponse
client: Final = get_test_client()
started: Final = time.monotonic()
stream: Final = client.responses.create(
model="gpt-5.5",
input="count from 1 to 500, one number per line",
stream=True,
background=True,
timeout=BACKGROUND_STREAM_ADMISSION_DEADLINE_SECONDS,
)
stream = client.responses.create(
model="gpt-5.5",
input="just respond with the word 'ping'",
stream=True,
background=True,
with stream:
events: Final = tuple(events_until_admission(stream, started))
elapsed: Final = time.monotonic() - started
keepalive_events: Final = sum(1 for chunk in events if chunk.type == "keepalive")
response_id: Final = next((rid for rid in map(admitted_response_id, events) if rid is not None), None)
if response_id is None and keepalive_events:
pytest.skip(
f"OpenAI held the background stream in keepalive for {elapsed:.0f}s "
f"({keepalive_events} keepalive events) without creating the response"
)
assert response_id is not None, f"no response event within {elapsed:.0f}s of streaming a background response"
collected_chunks = []
response_id = None
for chunk in stream:
print("stream chunk=", chunk)
collected_chunks.append(chunk)
# Extract response ID from the first chunk that has it
if (
response_id is None
and hasattr(chunk, "response")
and hasattr(chunk.response, "id")
):
response_id = chunk.response.id
assert len(collected_chunks) > 0
# cancel the response if we got a response ID
if response_id:
cancel_response = client.responses.cancel(response_id)
print("CANCEL streaming response=", cancel_response)
assert hasattr(cancel_response, "id")
except Exception as e:
if "Cannot cancel a completed response" in str(e):
pass
else:
raise e
cancel_response: Final = client.responses.cancel(response_id)
print("CANCEL streaming response=", cancel_response)
assert cancel_response.status == "cancelled"
def test_cancel_invalid_response_id():

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

@ -336,3 +336,15 @@ class TestAzureResolvesTheDeclaredDefaultEffort:
drop_params=True,
)
assert ("temperature" in mapped) is temperature_survives
def test_azure_gpt_6_astra_takes_the_reasoning_series_request_shape():
params = litellm.get_optional_params(
model="gpt-6-astra",
custom_llm_provider="azure",
max_tokens=100,
reasoning_effort="max",
)
assert params["max_completion_tokens"] == 100
assert "max_tokens" not in params
assert params["reasoning_effort"] == "max"

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

@ -3270,3 +3270,90 @@ async def test_create_batch_async_validates_credentials_off_the_event_loop():
validated_on = client.post.call_args.kwargs["headers"]["x-validated-on"]
assert validated_on != str(threading.get_ident())
assert client.post.call_args.kwargs["url"] == "https://batches.example/v1/messages/batches"
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

@ -1718,6 +1718,8 @@ class TestResponsesSurfaceSharesTheEffortRule:
("gpt-5.6-sol", None, False),
("gpt-5.6-terra", "none", True),
("gpt-5.6-terra", "medium", False),
("gpt-6-astra", None, False),
("gpt-6-astra", "low", False),
],
)
def test_temperature_follows_the_resolved_effort(

View file

@ -1505,3 +1505,17 @@ class TestACatalogueOlderThanTheCodeDoesNotStripTemperature:
drop_params=True,
)
assert "temperature" not in mapped
def test_gpt_6_astra_takes_the_reasoning_series_request_shape():
params = litellm.get_optional_params(
model="gpt-6-astra",
custom_llm_provider="openai",
max_tokens=100,
reasoning_effort="max",
verbosity="low",
)
assert params["max_completion_tokens"] == 100
assert "max_tokens" not in params
assert params["reasoning_effort"] == "max"
assert params["verbosity"] == "low"

View file

@ -41,6 +41,8 @@ from litellm.llms.azure.chat.gpt_5_transformation import AzureOpenAIGPT5Config
# Models that MUST be classified as GPT-5 (routed through GPT-5 reasoning path)
GPT5_MODELS = [
"gpt-6-astra",
"openai/gpt-6-astra",
"gpt-5",
"gpt-5.1",
"gpt-5.2",
@ -120,6 +122,8 @@ class TestOpenAIGPT5ConfigIsModelGpt5Model:
# /v1/responses bridge (when reasoning_effort is set and tools are passed) on
# is_model_gpt_5_4_plus_model, so the gpt-5.6 family must land on the True side.
GPT5_4_PLUS_MODELS = [
"gpt-6-astra",
"openai/gpt-6-astra",
"gpt-5.4",
"gpt-5.5",
"gpt-5.5-pro",

View file

@ -964,6 +964,48 @@ def test_construct_target_url_with_version_prefix():
assert str(target_url) == expected_url
@pytest.mark.parametrize(
("requested_route", "expected_url"),
[
(
"/projects/test-project/locations/global/publishers/anthropic/models/claude-sonnet-4-6:streamRawPredict",
"https://aiplatform.googleapis.com/v1/projects/test-project/locations/global/publishers/anthropic/models/claude-sonnet-4-6:streamRawPredict",
),
(
"/projects/test-project/locations/global/publishers/anthropic/models/count-tokens:rawPredict",
"https://aiplatform.googleapis.com/v1/projects/test-project/locations/global/publishers/anthropic/models/count-tokens:rawPredict",
),
(
"/projects/other-project/locations/us-east5/publishers/anthropic/models/claude-sonnet-4-6:rawPredict",
"https://aiplatform.googleapis.com/v1/projects/test-project/locations/global/publishers/anthropic/models/claude-sonnet-4-6:rawPredict",
),
(
"/projects/test-project/locations/global/cachedContents",
"https://aiplatform.googleapis.com/v1beta1/projects/test-project/locations/global/cachedContents",
),
(
"/v1/projects/test-project/locations/global/publishers/anthropic/models/claude-sonnet-4-6:streamRawPredict",
"https://aiplatform.googleapis.com/v1/projects/test-project/locations/global/publishers/anthropic/models/claude-sonnet-4-6:streamRawPredict",
),
(
"/v1beta1/projects/test-project/locations/global/cachedContents",
"https://aiplatform.googleapis.com/v1beta1/projects/test-project/locations/global/cachedContents",
),
],
)
def test_construct_target_url_versionless_project_route_gets_api_version(requested_route: str, expected_url: str) -> None:
from litellm.llms.vertex_ai.common_utils import construct_target_url
target_url = construct_target_url(
base_url="https://aiplatform.googleapis.com",
requested_route=requested_route,
vertex_project="test-project",
vertex_location="global",
)
assert str(target_url) == expected_url
def test_fix_enum_types():
"""
Test _fix_enum_types function removes enum fields when type is not string.

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