mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_fix_responses_queued_id_encryption
# Conflicts: # type-discipline-budget.json
This commit is contained in:
commit
2df5f4a7c8
100 changed files with 4428 additions and 906 deletions
24
.github/workflows/test-litellm-ui-build.yml
vendored
24
.github/workflows/test-litellm-ui-build.yml
vendored
|
|
@ -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 .
|
||||
|
|
|
|||
7
.github/workflows/test-rust.yml
vendored
7
.github/workflows/test-rust.yml
vendored
|
|
@ -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
|
||||
|
|
|
|||
1
.github/workflows/test-unit.yml
vendored
1
.github/workflows/test-unit.yml
vendored
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 -}}
|
||||
|
|
|
|||
|
|
@ -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 }}
|
||||
|
|
|
|||
205
helm/litellm/tests/ingress_controller_tests.yaml
Normal file
205
helm/litellm/tests/ingress_controller_tests.yaml
Normal 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
|
||||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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 "
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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).
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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/`.
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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).
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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": [],
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
||||
|
|
|
|||
|
|
@ -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] = {}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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}"
|
||||
|
|
|
|||
|
|
@ -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 ###
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -77,6 +77,7 @@ from litellm.llms.base_llm.vector_store_files.transformation import (
|
|||
BaseVectorStoreFilesConfig,
|
||||
)
|
||||
from litellm.llms.base_llm.videos.transformation import BaseVideoConfig
|
||||
from litellm.llms.custom_httpx.container_handler import raise_for_error_status
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
AsyncHTTPHandler,
|
||||
HTTPHandler,
|
||||
|
|
@ -8763,17 +8764,19 @@ class BaseLLMHTTPHandler:
|
|||
json=data,
|
||||
timeout=timeout,
|
||||
)
|
||||
|
||||
return container_provider_config.transform_container_create_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
raise self._handle_error(
|
||||
e=e,
|
||||
provider_config=container_provider_config,
|
||||
)
|
||||
raise_for_error_status(
|
||||
response=response,
|
||||
container_provider_config=container_provider_config,
|
||||
)
|
||||
return container_provider_config.transform_container_create_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
async def async_container_create_handler(
|
||||
self,
|
||||
|
|
@ -8839,17 +8842,19 @@ class BaseLLMHTTPHandler:
|
|||
json=data,
|
||||
timeout=timeout,
|
||||
)
|
||||
|
||||
return container_provider_config.transform_container_create_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
raise self._handle_error(
|
||||
e=e,
|
||||
provider_config=container_provider_config,
|
||||
)
|
||||
raise_for_error_status(
|
||||
response=response,
|
||||
container_provider_config=container_provider_config,
|
||||
)
|
||||
return container_provider_config.transform_container_create_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
def container_list_handler(
|
||||
self,
|
||||
|
|
@ -8929,17 +8934,19 @@ class BaseLLMHTTPHandler:
|
|||
headers=headers,
|
||||
params=params or None,
|
||||
)
|
||||
|
||||
return container_provider_config.transform_container_list_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
raise self._handle_error(
|
||||
e=e,
|
||||
provider_config=container_provider_config,
|
||||
)
|
||||
raise_for_error_status(
|
||||
response=response,
|
||||
container_provider_config=container_provider_config,
|
||||
)
|
||||
return container_provider_config.transform_container_list_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
async def async_container_list_handler(
|
||||
self,
|
||||
|
|
@ -9006,17 +9013,19 @@ class BaseLLMHTTPHandler:
|
|||
headers=headers,
|
||||
params=params or None,
|
||||
)
|
||||
|
||||
return container_provider_config.transform_container_list_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
raise self._handle_error(
|
||||
e=e,
|
||||
provider_config=container_provider_config,
|
||||
)
|
||||
raise_for_error_status(
|
||||
response=response,
|
||||
container_provider_config=container_provider_config,
|
||||
)
|
||||
return container_provider_config.transform_container_list_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
def container_retrieve_handler(
|
||||
self,
|
||||
|
|
@ -9094,17 +9103,19 @@ class BaseLLMHTTPHandler:
|
|||
headers=headers,
|
||||
params=params or None,
|
||||
)
|
||||
|
||||
return container_provider_config.transform_container_retrieve_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
raise self._handle_error(
|
||||
e=e,
|
||||
provider_config=container_provider_config,
|
||||
)
|
||||
raise_for_error_status(
|
||||
response=response,
|
||||
container_provider_config=container_provider_config,
|
||||
)
|
||||
return container_provider_config.transform_container_retrieve_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
async def async_container_retrieve_handler(
|
||||
self,
|
||||
|
|
@ -9171,17 +9182,19 @@ class BaseLLMHTTPHandler:
|
|||
headers=headers,
|
||||
params=params or None,
|
||||
)
|
||||
|
||||
return container_provider_config.transform_container_retrieve_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
raise self._handle_error(
|
||||
e=e,
|
||||
provider_config=container_provider_config,
|
||||
)
|
||||
raise_for_error_status(
|
||||
response=response,
|
||||
container_provider_config=container_provider_config,
|
||||
)
|
||||
return container_provider_config.transform_container_retrieve_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
def container_delete_handler(
|
||||
self,
|
||||
|
|
@ -9259,17 +9272,19 @@ class BaseLLMHTTPHandler:
|
|||
headers=headers,
|
||||
params=params or None,
|
||||
)
|
||||
|
||||
return container_provider_config.transform_container_delete_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
raise self._handle_error(
|
||||
e=e,
|
||||
provider_config=container_provider_config,
|
||||
)
|
||||
raise_for_error_status(
|
||||
response=response,
|
||||
container_provider_config=container_provider_config,
|
||||
)
|
||||
return container_provider_config.transform_container_delete_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
async def async_container_delete_handler(
|
||||
self,
|
||||
|
|
@ -9336,17 +9351,19 @@ class BaseLLMHTTPHandler:
|
|||
headers=headers,
|
||||
params=params or None,
|
||||
)
|
||||
|
||||
return container_provider_config.transform_container_delete_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
raise self._handle_error(
|
||||
e=e,
|
||||
provider_config=container_provider_config,
|
||||
)
|
||||
raise_for_error_status(
|
||||
response=response,
|
||||
container_provider_config=container_provider_config,
|
||||
)
|
||||
return container_provider_config.transform_container_delete_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
def container_file_list_handler(
|
||||
self,
|
||||
|
|
@ -9428,17 +9445,19 @@ class BaseLLMHTTPHandler:
|
|||
headers=headers,
|
||||
params=params or None,
|
||||
)
|
||||
|
||||
return container_provider_config.transform_container_file_list_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
raise self._handle_error(
|
||||
e=e,
|
||||
provider_config=container_provider_config,
|
||||
)
|
||||
raise_for_error_status(
|
||||
response=response,
|
||||
container_provider_config=container_provider_config,
|
||||
)
|
||||
return container_provider_config.transform_container_file_list_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
async def async_container_file_list_handler(
|
||||
self,
|
||||
|
|
@ -9507,17 +9526,19 @@ class BaseLLMHTTPHandler:
|
|||
headers=headers,
|
||||
params=params or None,
|
||||
)
|
||||
|
||||
return container_provider_config.transform_container_file_list_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
raise self._handle_error(
|
||||
e=e,
|
||||
provider_config=container_provider_config,
|
||||
)
|
||||
raise_for_error_status(
|
||||
response=response,
|
||||
container_provider_config=container_provider_config,
|
||||
)
|
||||
return container_provider_config.transform_container_file_list_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
def container_file_content_handler(
|
||||
self,
|
||||
|
|
@ -9593,17 +9614,19 @@ class BaseLLMHTTPHandler:
|
|||
headers=headers,
|
||||
params=params or None,
|
||||
)
|
||||
|
||||
return container_provider_config.transform_container_file_content_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
raise self._handle_error(
|
||||
e=e,
|
||||
provider_config=container_provider_config,
|
||||
)
|
||||
raise_for_error_status(
|
||||
response=response,
|
||||
container_provider_config=container_provider_config,
|
||||
)
|
||||
return container_provider_config.transform_container_file_content_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
async def async_container_file_content_handler(
|
||||
self,
|
||||
|
|
@ -9669,17 +9692,19 @@ class BaseLLMHTTPHandler:
|
|||
headers=headers,
|
||||
params=params or None,
|
||||
)
|
||||
|
||||
return container_provider_config.transform_container_file_content_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
raise self._handle_error(
|
||||
e=e,
|
||||
provider_config=container_provider_config,
|
||||
)
|
||||
raise_for_error_status(
|
||||
response=response,
|
||||
container_provider_config=container_provider_config,
|
||||
)
|
||||
return container_provider_config.transform_container_file_content_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
###### VECTOR STORE HANDLER ######
|
||||
@staticmethod
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
)
|
||||
|
|
|
|||
|
|
@ -594,16 +594,20 @@ set_live_deployment_replay(_replay_live_router_model_cost)
|
|||
|
||||
|
||||
# Kwargs that carry no signal about the failed attempt, so log_retry drops them from a
|
||||
# breadcrumb entirely: the request payload and the router-internal walk state. Credentials are
|
||||
# handled separately by mask_credentials_in_payload, which scrubs credential-named values from
|
||||
# whatever kwargs remain rather than trying to enumerate every credential-bearing key here.
|
||||
# breadcrumb entirely: the request payload, the proxy's snapshot of the inbound request (its body
|
||||
# aliases the live request metadata, earlier breadcrumbs included, so copying it would nest every
|
||||
# breadcrumb inside the next one), and the router-internal walk state. Credentials are handled
|
||||
# separately by mask_credentials_in_payload, which scrubs credential-named values from whatever
|
||||
# kwargs remain rather than trying to enumerate every credential-bearing key here.
|
||||
RETRY_BREADCRUMB_EXCLUDED_KWARGS: Final = frozenset(
|
||||
(
|
||||
"messages",
|
||||
"original_function",
|
||||
"attempted_targets",
|
||||
"proxy_server_request",
|
||||
)
|
||||
)
|
||||
RETRY_BREADCRUMB_LIMIT: Final = 4
|
||||
|
||||
|
||||
class Router:
|
||||
|
|
@ -964,7 +968,6 @@ class Router:
|
|||
self.total_calls: defaultdict = defaultdict(int) # dict to store total calls made to each model
|
||||
self.fail_calls: defaultdict = defaultdict(int) # dict to store fail_calls made to each model
|
||||
self.success_calls: defaultdict = defaultdict(int) # dict to store success_calls made to each model
|
||||
self.previous_models: list = [] # list to store failed calls (passed in as metadata to next call)
|
||||
|
||||
# make Router.chat.completions.create compatible for openai.chat.completions.create
|
||||
default_litellm_params = default_litellm_params or {}
|
||||
|
|
@ -8143,35 +8146,31 @@ class Router:
|
|||
"""
|
||||
When a retry or fallback happens, log the details of the just failed model call - similar to Sentry breadcrumbing
|
||||
"""
|
||||
try:
|
||||
_metadata_var: Final = "litellm_metadata" if "litellm_metadata" in kwargs else "metadata"
|
||||
# Log failed model as the previous model
|
||||
previous_model: Final = {
|
||||
_metadata_var: Final = "litellm_metadata" if "litellm_metadata" in kwargs else "metadata"
|
||||
request_metadata: Final[Mapping[str, object]] = kwargs[_metadata_var]
|
||||
attempt_kwargs: Final = MappingProxyType(
|
||||
{k: v for k, v in kwargs.items() if k != _metadata_var and k not in RETRY_BREADCRUMB_EXCLUDED_KWARGS}
|
||||
)
|
||||
attempt_metadata: Final = MappingProxyType(
|
||||
{k: v for k, v in request_metadata.items() if k != "previous_models"}
|
||||
)
|
||||
previous_model: Final = MappingProxyType(
|
||||
{
|
||||
"exception_type": type(e).__name__,
|
||||
"exception_string": str(e),
|
||||
**attempt_kwargs,
|
||||
_metadata_var: attempt_metadata,
|
||||
}
|
||||
for (
|
||||
k,
|
||||
v,
|
||||
) in kwargs.items(): # log everything in kwargs except the old previous_models value - prevent nesting
|
||||
if k != _metadata_var and k not in RETRY_BREADCRUMB_EXCLUDED_KWARGS:
|
||||
previous_model[k] = v
|
||||
elif k == _metadata_var and isinstance(v, dict):
|
||||
previous_model[_metadata_var] = {}
|
||||
for metadata_k, metadata_v in kwargs[_metadata_var].items():
|
||||
if metadata_k != "previous_models":
|
||||
previous_model[k][metadata_k] = metadata_v
|
||||
|
||||
# check current size of self.previous_models, if it's larger than 3, remove the first element
|
||||
if len(self.previous_models) > 3:
|
||||
self.previous_models.pop(0)
|
||||
|
||||
scrubbed_previous_model: Final = mask_credentials_in_payload(previous_model)
|
||||
self.previous_models.append(scrubbed_previous_model)
|
||||
kwargs[_metadata_var]["previous_models"] = self.previous_models
|
||||
return kwargs
|
||||
except Exception as e:
|
||||
raise e
|
||||
)
|
||||
earlier_breadcrumbs: Final = request_metadata.get("previous_models")
|
||||
kept_breadcrumbs: Final[tuple[object, ...]] = (
|
||||
tuple(earlier_breadcrumbs)[-(RETRY_BREADCRUMB_LIMIT - 1) :]
|
||||
if isinstance(earlier_breadcrumbs, (list, tuple))
|
||||
else ()
|
||||
)
|
||||
breadcrumbs: Final = (*kept_breadcrumbs, mask_credentials_in_payload(previous_model))
|
||||
kwargs[_metadata_var]["previous_models"] = breadcrumbs # rebind-ok: the logging object already holds this dict
|
||||
return kwargs
|
||||
|
||||
def _update_usage(self, deployment_id: str, parent_otel_span: Span | None) -> int:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -83,6 +83,7 @@ from litellm.constants import (
|
|||
from litellm.litellm_core_utils.fallback_generalizations import (
|
||||
match_capability_generalizations,
|
||||
)
|
||||
from litellm.litellm_core_utils.sensitive_data_masker import redact_credentials_in_payload
|
||||
|
||||
_CachingHandlerResponse = None
|
||||
_LLMCachingHandler = None
|
||||
|
|
@ -543,6 +544,14 @@ def print_verbose(
|
|||
pass
|
||||
|
||||
|
||||
def _print_verbose_is_active() -> bool:
|
||||
"""Whether print_verbose would reach either of its two consumers, so a call site can skip
|
||||
building a payload nothing would read. _is_debugging_on() is not the same predicate: it reads
|
||||
litellm._logging.set_verbose, while print_verbose's print reads litellm.set_verbose, and
|
||||
assigning the documented litellm.set_verbose = True rebinds only the latter."""
|
||||
return litellm.set_verbose is True or verbose_logger.isEnabledFor(logging.DEBUG)
|
||||
|
||||
|
||||
####### CLIENT ###################
|
||||
# make it easy to log if completion/embedding runs succeeded or failed + see what happened | Non-Blocking
|
||||
def custom_llm_setup():
|
||||
|
|
@ -1284,16 +1293,18 @@ async def async_post_call_success_deployment_hook(
|
|||
except ValueError:
|
||||
typed_call_type = None # unknown call type
|
||||
|
||||
modified_response = response
|
||||
|
||||
CustomLogger: Final = _get_cached_custom_logger()
|
||||
for callback in litellm.callbacks:
|
||||
if isinstance(callback, CustomLogger):
|
||||
result = await callback.async_post_call_success_deployment_hook(
|
||||
request_data, cast(LLMResponseTypes, response), typed_call_type
|
||||
request_data, cast(LLMResponseTypes, modified_response), typed_call_type
|
||||
)
|
||||
if result is not None:
|
||||
return result
|
||||
modified_response = result
|
||||
|
||||
return response
|
||||
return modified_response
|
||||
|
||||
|
||||
async def async_post_call_failure_deployment_hook(
|
||||
|
|
@ -4707,7 +4718,8 @@ def get_optional_params(
|
|||
openai_params=list(DEFAULT_CHAT_COMPLETION_PARAM_VALUES.keys()),
|
||||
additional_drop_params=additional_drop_params,
|
||||
)
|
||||
print_verbose(f"Final returned optional params: {optional_params}")
|
||||
if _print_verbose_is_active():
|
||||
print_verbose(f"Final returned optional params: {redact_credentials_in_payload(optional_params)}")
|
||||
optional_params = _apply_openai_param_overrides(
|
||||
optional_params=optional_params,
|
||||
non_default_params=non_default_params,
|
||||
|
|
@ -7462,7 +7474,8 @@ def print_args_passed_to_litellm(original_function, args, kwargs):
|
|||
return
|
||||
|
||||
args_str: Final = ", ".join(map(repr, args))
|
||||
kwargs_str: Final = ", ".join(f"{key}={value!r}" for key, value in kwargs.items())
|
||||
redacted_kwargs: Final = redact_credentials_in_payload(kwargs)
|
||||
kwargs_str: Final = ", ".join(f"{key}={value!r}" for key, value in redacted_kwargs.items())
|
||||
print_verbose(
|
||||
"\n",
|
||||
) # new line before
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -42,7 +42,7 @@
|
|||
"limit": 52
|
||||
},
|
||||
"B010": {
|
||||
"limit": 188
|
||||
"limit": 187
|
||||
},
|
||||
"B018": {
|
||||
"limit": 2
|
||||
|
|
@ -144,7 +144,7 @@
|
|||
"limit": 1
|
||||
},
|
||||
"PLR1704": {
|
||||
"limit": 3
|
||||
"limit": 1
|
||||
},
|
||||
"PLR1714": {
|
||||
"limit": 253
|
||||
|
|
@ -240,13 +240,13 @@
|
|||
"limit": 96
|
||||
},
|
||||
"TRY201": {
|
||||
"limit": 403
|
||||
"limit": 401
|
||||
},
|
||||
"TRY203": {
|
||||
"limit": 111
|
||||
"limit": 109
|
||||
},
|
||||
"TRY300": {
|
||||
"limit": 854
|
||||
"limit": 852
|
||||
},
|
||||
"UP028": {
|
||||
"limit": 2
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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]]
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -3098,3 +3098,91 @@ async def test_a_provider_that_keeps_rejecting_is_not_retried_forever_on_the_asy
|
|||
)
|
||||
|
||||
assert len(recorder.bodies) == 2
|
||||
|
||||
|
||||
CONTAINER_NOT_FOUND_BODY = {
|
||||
"error": {
|
||||
"message": "Container with id 'cntr_gone' not found.",
|
||||
"type": "invalid_request_error",
|
||||
"param": None,
|
||||
"code": None,
|
||||
}
|
||||
}
|
||||
|
||||
INVALID_API_KEY_BODY = {
|
||||
"error": {
|
||||
"message": "Incorrect API key provided: sk-proj-***. You can find your API key at https://platform.openai.com/account/api-keys.",
|
||||
"type": "invalid_request_error",
|
||||
"param": None,
|
||||
"code": "invalid_api_key",
|
||||
},
|
||||
"status": 401,
|
||||
}
|
||||
|
||||
CONTAINER_LIST_BODY = {
|
||||
"object": "list",
|
||||
"data": [{"id": "cntr_a", "object": "container", "created_at": 1, "status": "running", "name": "a"}],
|
||||
"first_id": "cntr_a",
|
||||
"last_id": "cntr_a",
|
||||
"has_more": True,
|
||||
}
|
||||
|
||||
|
||||
def _container_sync_client(response: httpx.Response) -> HTTPHandler:
|
||||
client = HTTPHandler()
|
||||
client.client = httpx.Client(transport=httpx.MockTransport(lambda _request: response))
|
||||
return client
|
||||
|
||||
|
||||
def _container_async_client(response: httpx.Response) -> AsyncHTTPHandler:
|
||||
client = AsyncHTTPHandler()
|
||||
client.client = httpx.AsyncClient(transport=httpx.MockTransport(lambda _request: response))
|
||||
return client
|
||||
|
||||
|
||||
def test_container_retrieve_handler_raises_upstream_error_status_and_message():
|
||||
from litellm.llms.openai.containers.transformation import OpenAIContainerConfig
|
||||
|
||||
with pytest.raises(BaseLLMException) as exc_info:
|
||||
BaseLLMHTTPHandler().container_retrieve_handler(
|
||||
container_id="cntr_gone",
|
||||
container_provider_config=OpenAIContainerConfig(),
|
||||
litellm_params=GenericLiteLLMParams(api_key="sk-test"),
|
||||
logging_obj=Mock(),
|
||||
client=_container_sync_client(httpx.Response(404, json=CONTAINER_NOT_FOUND_BODY)),
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 404
|
||||
assert exc_info.value.message == "Container with id 'cntr_gone' not found."
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_container_list_handler_raises_upstream_error_status_and_message():
|
||||
from litellm.llms.openai.containers.transformation import OpenAIContainerConfig
|
||||
|
||||
with pytest.raises(BaseLLMException) as exc_info:
|
||||
await BaseLLMHTTPHandler().async_container_list_handler(
|
||||
container_provider_config=OpenAIContainerConfig(),
|
||||
litellm_params=GenericLiteLLMParams(api_key="sk-rejected"),
|
||||
logging_obj=Mock(),
|
||||
client=_container_async_client(httpx.Response(401, json=INVALID_API_KEY_BODY)),
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 401
|
||||
assert exc_info.value.message == INVALID_API_KEY_BODY["error"]["message"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_container_list_handler_transforms_success_response():
|
||||
from litellm.llms.openai.containers.transformation import OpenAIContainerConfig
|
||||
|
||||
response = await BaseLLMHTTPHandler().async_container_list_handler(
|
||||
container_provider_config=OpenAIContainerConfig(),
|
||||
litellm_params=GenericLiteLLMParams(api_key="sk-test"),
|
||||
logging_obj=Mock(),
|
||||
limit=1,
|
||||
client=_container_async_client(httpx.Response(200, json=CONTAINER_LIST_BODY)),
|
||||
)
|
||||
|
||||
assert [container.id for container in response.data] == ["cntr_a"]
|
||||
assert response.has_more is True
|
||||
|
|
|
|||
|
|
@ -1,4 +1,6 @@
|
|||
import io
|
||||
import json
|
||||
from typing import get_type_hints
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import orjson
|
||||
|
|
@ -18,9 +20,11 @@ from litellm.proxy.common_utils.http_parsing_utils import (
|
|||
_safe_get_request_parsed_body,
|
||||
_safe_get_request_query_params,
|
||||
_safe_set_request_parsed_body,
|
||||
coerce_numeric_form_fields,
|
||||
get_form_data,
|
||||
get_request_body,
|
||||
get_tags_from_request_body,
|
||||
numeric_form_fields,
|
||||
populate_request_with_path_params,
|
||||
)
|
||||
|
||||
|
|
@ -1029,3 +1033,79 @@ class TestGetRequestBody:
|
|||
mock_request = MagicMock()
|
||||
mock_request.method = "GET"
|
||||
assert await get_request_body(mock_request) == {}
|
||||
|
||||
|
||||
class TestNumericFormFields:
|
||||
def test_image_edit_schema_yields_only_n(self):
|
||||
from litellm.types.images.main import ImageEditRequestParams
|
||||
|
||||
assert dict(numeric_form_fields(get_type_hints(ImageEditRequestParams))) == {"n": int}
|
||||
|
||||
def test_qualifiers_and_optionality_are_unwrapped(self):
|
||||
from typing import Optional
|
||||
|
||||
from typing_extensions import Annotated, NotRequired, ReadOnly, Required, TypedDict
|
||||
|
||||
class Schema(TypedDict, total=False):
|
||||
plain: int
|
||||
optional: Optional[int]
|
||||
piped: int | None
|
||||
read_only: ReadOnly[int | None]
|
||||
not_required: NotRequired[ReadOnly[int]]
|
||||
required: Required[ReadOnly[Annotated[float, "meta"]]]
|
||||
|
||||
assert dict(numeric_form_fields(get_type_hints(Schema))) == {
|
||||
"plain": int,
|
||||
"optional": int,
|
||||
"piped": int,
|
||||
"read_only": int,
|
||||
"not_required": int,
|
||||
"required": float,
|
||||
}
|
||||
|
||||
def test_non_scalar_and_bool_fields_are_skipped(self):
|
||||
from typing import Any, Literal, Optional, Union
|
||||
|
||||
from typing_extensions import TypedDict
|
||||
|
||||
class Schema(TypedDict, total=False):
|
||||
flag: bool
|
||||
optional_flag: Optional[bool]
|
||||
text: str
|
||||
choice: Optional[Literal["high", "low"]]
|
||||
numbers: list[int]
|
||||
mapping: Optional[dict[str, Any]]
|
||||
ambiguous: Union[int, str]
|
||||
|
||||
assert dict(numeric_form_fields(get_type_hints(Schema))) == {}
|
||||
|
||||
|
||||
class TestCoerceNumericFormFields:
|
||||
numeric_fields = {"n": int, "temperature": float}
|
||||
|
||||
def test_numeric_strings_are_parsed(self):
|
||||
assert coerce_numeric_form_fields(
|
||||
parsed_body={"n": "2", "temperature": "0.5"},
|
||||
numeric_fields=self.numeric_fields,
|
||||
) == {"n": 2, "temperature": 0.5}
|
||||
|
||||
def test_other_fields_keep_their_string_values(self):
|
||||
result = coerce_numeric_form_fields(
|
||||
parsed_body={"size": "1024x1024", "prompt": "2", "quality": "high"},
|
||||
numeric_fields=self.numeric_fields,
|
||||
)
|
||||
assert result == {"size": "1024x1024", "prompt": "2", "quality": "high"}
|
||||
|
||||
def test_unparseable_value_is_left_for_the_provider_to_reject(self):
|
||||
assert coerce_numeric_form_fields(
|
||||
parsed_body={"n": "two", "temperature": ""},
|
||||
numeric_fields=self.numeric_fields,
|
||||
) == {"n": "two", "temperature": ""}
|
||||
|
||||
def test_already_typed_and_non_string_values_pass_through(self):
|
||||
buffer = io.BytesIO(b"png")
|
||||
result = coerce_numeric_form_fields(
|
||||
parsed_body={"n": 3, "temperature": None, "image": buffer},
|
||||
numeric_fields=self.numeric_fields,
|
||||
)
|
||||
assert result == {"n": 3, "temperature": None, "image": buffer}
|
||||
|
|
|
|||
0
tests/test_litellm/proxy/container_endpoints/__init__.py
Normal file
0
tests/test_litellm/proxy/container_endpoints/__init__.py
Normal file
134
tests/test_litellm/proxy/container_endpoints/test_endpoints.py
Normal file
134
tests/test_litellm/proxy/container_endpoints/test_endpoints.py
Normal file
|
|
@ -0,0 +1,134 @@
|
|||
import sys
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
from fastapi import FastAPI
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.container_endpoints import endpoints, ownership
|
||||
from litellm.types.containers.main import ContainerListResponse, ContainerObject
|
||||
|
||||
PROXY_SERVER_STUB = SimpleNamespace(
|
||||
general_settings={},
|
||||
prisma_client=None,
|
||||
llm_router=None,
|
||||
proxy_config=None,
|
||||
proxy_logging_obj=None,
|
||||
select_data_generator=None,
|
||||
user_api_base=None,
|
||||
user_max_tokens=None,
|
||||
user_model=None,
|
||||
user_request_timeout=None,
|
||||
user_temperature=None,
|
||||
version="test",
|
||||
)
|
||||
ADMIN = UserAPIKeyAuth(user_id="admin-1", user_role=LitellmUserRoles.PROXY_ADMIN)
|
||||
NON_ADMIN = UserAPIKeyAuth(user_id="user-1")
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def clear_allowed_container_ids_cache():
|
||||
ownership._ALLOWED_CONTAINER_IDS_CACHE.cache_dict.clear()
|
||||
ownership._ALLOWED_CONTAINER_IDS_CACHE.ttl_dict.clear()
|
||||
yield
|
||||
ownership._ALLOWED_CONTAINER_IDS_CACHE.cache_dict.clear()
|
||||
ownership._ALLOWED_CONTAINER_IDS_CACHE.ttl_dict.clear()
|
||||
|
||||
|
||||
def _client(auth: UserAPIKeyAuth) -> TestClient:
|
||||
app = FastAPI()
|
||||
app.include_router(endpoints.router)
|
||||
app.dependency_overrides[user_api_key_auth] = lambda: auth
|
||||
return TestClient(app)
|
||||
|
||||
|
||||
def _container(container_id: str) -> ContainerObject:
|
||||
return ContainerObject(id=container_id, object="container", created_at=1, status="active")
|
||||
|
||||
|
||||
def _page(*container_ids: str, has_more: bool) -> ContainerListResponse:
|
||||
return ContainerListResponse(
|
||||
object="list",
|
||||
data=[_container(container_id) for container_id in container_ids],
|
||||
has_more=has_more,
|
||||
)
|
||||
|
||||
|
||||
def _upstream_pages(monkeypatch, pages_by_after) -> MagicMock:
|
||||
processor_cls = MagicMock(
|
||||
side_effect=lambda data: SimpleNamespace(
|
||||
base_process_llm_request=AsyncMock(return_value=pages_by_after[data["after"]])
|
||||
)
|
||||
)
|
||||
monkeypatch.setattr(endpoints, "ProxyBaseLLMRequestProcessing", processor_cls)
|
||||
return processor_cls
|
||||
|
||||
|
||||
def _forwarded_pages(processor_cls: MagicMock):
|
||||
return [(call.kwargs["data"]["after"], call.kwargs["data"]["limit"]) for call in processor_cls.call_args_list]
|
||||
|
||||
|
||||
def test_list_containers_forwards_typed_pagination_params_for_admins(monkeypatch):
|
||||
monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", PROXY_SERVER_STUB)
|
||||
processor_cls = _upstream_pages(monkeypatch, {"cntr_prev": _page("cntr_next", has_more=True)})
|
||||
|
||||
response = _client(ADMIN).get(
|
||||
"/v1/containers",
|
||||
params={"limit": "1", "order": "desc", "after": "cntr_prev"},
|
||||
headers={"Authorization": "Bearer sk-test"},
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
assert [item["id"] for item in response.json()["data"]] == ["cntr_next"]
|
||||
assert response.json()["has_more"] is True
|
||||
assert _forwarded_pages(processor_cls) == [("cntr_prev", 1)]
|
||||
assert processor_cls.call_args.kwargs["data"]["order"] == "desc"
|
||||
|
||||
|
||||
def test_list_containers_rejects_a_non_integer_limit(monkeypatch):
|
||||
monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", PROXY_SERVER_STUB)
|
||||
processor_cls = _upstream_pages(monkeypatch, {})
|
||||
|
||||
response = _client(ADMIN).get(
|
||||
"/v1/containers",
|
||||
params={"limit": "abc"},
|
||||
headers={"Authorization": "Bearer sk-test"},
|
||||
)
|
||||
|
||||
assert response.status_code == 422
|
||||
processor_cls.assert_not_called()
|
||||
|
||||
|
||||
def test_list_containers_pages_upstream_until_non_admin_keys_see_their_containers(monkeypatch):
|
||||
monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", PROXY_SERVER_STUB)
|
||||
table = AsyncMock()
|
||||
table.find_many.return_value = [SimpleNamespace(model_object_id="container:openai:cntr_owned")]
|
||||
monkeypatch.setattr(
|
||||
ownership,
|
||||
"_get_prisma_client",
|
||||
AsyncMock(return_value=SimpleNamespace(db=SimpleNamespace(litellm_managedobjecttable=table))),
|
||||
)
|
||||
processor_cls = _upstream_pages(
|
||||
monkeypatch,
|
||||
{
|
||||
None: _page("cntr_other", has_more=True),
|
||||
"cntr_other": _page("cntr_owned", has_more=False),
|
||||
},
|
||||
)
|
||||
|
||||
response = _client(NON_ADMIN).get(
|
||||
"/v1/containers",
|
||||
params={"limit": "1"},
|
||||
headers={"Authorization": "Bearer sk-test"},
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
body = response.json()
|
||||
assert [item["id"] for item in body["data"]] == ["cntr_owned"]
|
||||
assert body["first_id"] == "cntr_owned"
|
||||
assert body["last_id"] == "cntr_owned"
|
||||
assert body["has_more"] is False
|
||||
assert _forwarded_pages(processor_cls) == [(None, 100), ("cntr_other", 100)]
|
||||
|
|
@ -0,0 +1,62 @@
|
|||
import sys
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
from fastapi import FastAPI
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.container_endpoints import endpoints, handler_factory
|
||||
|
||||
PROXY_SERVER_STUB = SimpleNamespace(
|
||||
general_settings={},
|
||||
prisma_client=None,
|
||||
llm_router=None,
|
||||
proxy_config=None,
|
||||
proxy_logging_obj=None,
|
||||
select_data_generator=None,
|
||||
user_api_base=None,
|
||||
user_max_tokens=None,
|
||||
user_model=None,
|
||||
user_request_timeout=None,
|
||||
user_temperature=None,
|
||||
version="test",
|
||||
)
|
||||
|
||||
|
||||
def _client() -> TestClient:
|
||||
app = FastAPI()
|
||||
app.include_router(endpoints.router)
|
||||
app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(user_id="user-1")
|
||||
return TestClient(app)
|
||||
|
||||
|
||||
def test_list_container_files_forwards_declared_query_params(monkeypatch):
|
||||
monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", PROXY_SERVER_STUB)
|
||||
monkeypatch.setattr(
|
||||
handler_factory,
|
||||
"assert_user_can_access_container",
|
||||
AsyncMock(return_value=("cntr_123", "openai")),
|
||||
)
|
||||
processor_cls = MagicMock()
|
||||
processor_cls.return_value.base_process_llm_request = AsyncMock(
|
||||
return_value={"object": "list", "data": [], "has_more": True}
|
||||
)
|
||||
monkeypatch.setattr(handler_factory, "ProxyBaseLLMRequestProcessing", processor_cls)
|
||||
|
||||
response = _client().get(
|
||||
"/v1/containers/cntr_123/files",
|
||||
params={"limit": "1", "order": "desc", "after": "cfile_prev", "unknown": "x"},
|
||||
headers={"Authorization": "Bearer sk-test"},
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
assert response.json()["has_more"] is True
|
||||
assert processor_cls.return_value.base_process_llm_request.await_args.kwargs["route_type"] == "alist_container_files"
|
||||
forwarded = processor_cls.call_args.kwargs["data"]
|
||||
assert forwarded["container_id"] == "cntr_123"
|
||||
assert forwarded["limit"] == "1"
|
||||
assert forwarded["order"] == "desc"
|
||||
assert forwarded["after"] == "cfile_prev"
|
||||
assert "unknown" not in forwarded
|
||||
|
|
@ -29,7 +29,7 @@ def test_hands_the_alter_statement_to_the_prisma_cli():
|
|||
return subprocess.CompletedProcess(cmd, 0)
|
||||
|
||||
with patch(
|
||||
"litellm_proxy_extras.replica_identity.subprocess.run", side_effect=capture
|
||||
"litellm_proxy_extras.replica_identity.run_prisma", side_effect=capture
|
||||
):
|
||||
applied = apply_replica_identity_full(
|
||||
schema_path="/somewhere/schema.prisma",
|
||||
|
|
@ -60,7 +60,7 @@ def test_hands_the_alter_statement_to_the_prisma_cli():
|
|||
)
|
||||
def test_every_failure_is_reported_instead_of_raised(failure):
|
||||
with patch(
|
||||
"litellm_proxy_extras.replica_identity.subprocess.run", side_effect=failure
|
||||
"litellm_proxy_extras.replica_identity.run_prisma", side_effect=failure
|
||||
):
|
||||
assert (
|
||||
apply_replica_identity_full(
|
||||
|
|
|
|||
|
|
@ -5527,10 +5527,6 @@ async def test_streaming_end_of_stream_block_emits_error_frame_instead_of_trunca
|
|||
error frame instead. The finish chunk is withheld while the end-of-stream
|
||||
scan runs, so on a block it is dropped rather than relayed before the
|
||||
frame."""
|
||||
from litellm.llms import load_guardrail_translation_mappings
|
||||
from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail import (
|
||||
unified_guardrail as unified_module,
|
||||
)
|
||||
from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import (
|
||||
UnifiedLLMGuardrails,
|
||||
)
|
||||
|
|
@ -5569,20 +5565,16 @@ async def test_streaming_end_of_stream_block_emits_error_frame_instead_of_trunca
|
|||
yield _chunk("the forbidden ")
|
||||
yield _chunk("topic answer", finish_reason="stop")
|
||||
|
||||
unified_module.endpoint_guardrail_translation_mappings = load_guardrail_translation_mappings()
|
||||
try:
|
||||
with patch.object(guardrail, "make_bedrock_api_request", new_callable=AsyncMock) as mock_api:
|
||||
mock_api.side_effect = guardrail._get_http_exception_for_blocked_guardrail(blocked_response)
|
||||
with patch.object(guardrail, "make_bedrock_api_request", new_callable=AsyncMock) as mock_api:
|
||||
mock_api.side_effect = guardrail._get_http_exception_for_blocked_guardrail(blocked_response)
|
||||
|
||||
out = []
|
||||
async for item in UnifiedLLMGuardrails().async_post_call_streaming_iterator_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="test", request_route="/v1/chat/completions"),
|
||||
response=_mock_stream(),
|
||||
request_data={"guardrail_to_apply": guardrail, "model": "gpt-4"},
|
||||
):
|
||||
out.append(item)
|
||||
finally:
|
||||
unified_module.endpoint_guardrail_translation_mappings = None
|
||||
out = []
|
||||
async for item in UnifiedLLMGuardrails().async_post_call_streaming_iterator_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="test", request_route="/v1/chat/completions"),
|
||||
response=_mock_stream(),
|
||||
request_data={"guardrail_to_apply": guardrail, "model": "gpt-4"},
|
||||
):
|
||||
out.append(item)
|
||||
|
||||
assert len(out) == 2
|
||||
assert isinstance(out[0], ModelResponseStream)
|
||||
|
|
@ -5792,3 +5784,25 @@ async def test_apply_guardrail_debug_log_masks_signed_request_headers():
|
|||
assert header_lines, "expected the signed-request debug line to be logged"
|
||||
assert any("X-Amz-Security-Token" in message for message in header_lines)
|
||||
assert all(session_token not in message for message in rendered_messages)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bearer_token_never_runs_the_sigv4_credential_chain(monkeypatch):
|
||||
"""The guardrail's AWS profile does not exist, so resolving SigV4 credentials
|
||||
raises; with a bearer token configured the guardrail must still run, since
|
||||
the bearer token alone signs the request."""
|
||||
monkeypatch.setenv("AWS_BEARER_TOKEN_BEDROCK", "env-bearer-token-12345")
|
||||
guardrail = BedrockGuardrail(
|
||||
guardrailIdentifier="test-guardrail",
|
||||
guardrailVersion="DRAFT",
|
||||
aws_profile_name="litellm-no-such-aws-profile",
|
||||
)
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.json.return_value = {"action": "NONE", "assessments": []}
|
||||
|
||||
with patch.object(guardrail.async_handler, "post", new_callable=AsyncMock, return_value=mock_response) as mock_post:
|
||||
response = await guardrail.make_bedrock_api_request(source="INPUT", messages=[{"role": "user", "content": "hello"}])
|
||||
|
||||
assert response["action"] == "NONE"
|
||||
assert mock_post.call_args.kwargs["headers"]["Authorization"] == "Bearer env-bearer-token-12345"
|
||||
|
|
|
|||
|
|
@ -833,3 +833,31 @@ async def test_many_blocks_scanned_at_request_level_and_can_block():
|
|||
sent_texts = [c["text"] for m in body_messages for c in m["content"]]
|
||||
assert sent_texts == [f"b{i}" for i in range(25)]
|
||||
assert all(len(m["content"]) <= 10 for m in body_messages)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_checks_bearer_token_never_runs_the_sigv4_credential_chain(monkeypatch):
|
||||
"""Same bearer-token rule as ApplyGuardrail: the guardrail's AWS profile does
|
||||
not exist, yet the InvokeGuardrailChecks call still goes out on the bearer
|
||||
token and its verdict is enforced."""
|
||||
monkeypatch.setenv("AWS_BEARER_TOKEN_BEDROCK", "env-bearer-token-12345")
|
||||
g = BedrockGuardrail(
|
||||
checks=CONTENT_FILTER_CHECKS,
|
||||
content_filter_threshold=0.5,
|
||||
aws_profile_name="litellm-no-such-aws-profile",
|
||||
)
|
||||
payload = {"results": {"contentFilter": {"results": [{"category": "VIOLENCE", "severityScore": 0.8}]}}}
|
||||
post = AsyncMock(return_value=_mock_http_response(200, payload))
|
||||
|
||||
with patch.object(g.async_handler, "post", new=post):
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await g.make_bedrock_api_request(
|
||||
source="INPUT",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
request_data={"messages": []},
|
||||
)
|
||||
|
||||
assert exc.value.detail["bedrock_guardrail_checks"] == [
|
||||
{"check": "contentFilter", "category": "VIOLENCE", "severityScore": 0.8}
|
||||
]
|
||||
assert post.call_args.kwargs["headers"]["Authorization"] == "Bearer env-bearer-token-12345"
|
||||
|
|
|
|||
|
|
@ -75,19 +75,29 @@ class _NoopTranslation(BaseTranslation):
|
|||
return response
|
||||
|
||||
|
||||
def _patch_translation_mappings(monkeypatch, mappings):
|
||||
"""Point the unified guardrail at ``mappings`` for one test, restored by pytest.
|
||||
|
||||
Every override goes through this one seam: competing writers to the same state
|
||||
are what leaked a stale handler map into unrelated test files (LIT-6834).
|
||||
"""
|
||||
monkeypatch.setattr(unified_module, "load_guardrail_translation_mappings", lambda: mappings)
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _inject_mcp_handler_mapping():
|
||||
def _inject_mcp_handler_mapping(monkeypatch):
|
||||
"""Inject MCP handler mapping so the unified guardrail can run inside tests."""
|
||||
unified_module.endpoint_guardrail_translation_mappings = {
|
||||
CallTypes.call_mcp_tool: MCPGuardrailTranslationHandler,
|
||||
CallTypes.anthropic_messages: _NoopTranslation,
|
||||
CallTypes.ocr: OCRHandler,
|
||||
CallTypes.aocr: OCRHandler,
|
||||
CallTypes.responses: OpenAIResponsesHandler,
|
||||
CallTypes.aresponses: OpenAIResponsesHandler,
|
||||
}
|
||||
yield
|
||||
unified_module.endpoint_guardrail_translation_mappings = None
|
||||
_patch_translation_mappings(
|
||||
monkeypatch,
|
||||
{
|
||||
CallTypes.call_mcp_tool: MCPGuardrailTranslationHandler,
|
||||
CallTypes.anthropic_messages: _NoopTranslation,
|
||||
CallTypes.ocr: OCRHandler,
|
||||
CallTypes.aocr: OCRHandler,
|
||||
CallTypes.responses: OpenAIResponsesHandler,
|
||||
CallTypes.aresponses: OpenAIResponsesHandler,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
class TestUnifiedLLMGuardrails:
|
||||
|
|
@ -396,7 +406,7 @@ class TestUnifiedLLMGuardrails:
|
|||
|
||||
class TestAsyncPostCallStreamingIteratorHook:
|
||||
@pytest.mark.asyncio
|
||||
async def test_streaming_content_not_lost_on_sampled_chunks(self):
|
||||
async def test_streaming_content_not_lost_on_sampled_chunks(self, monkeypatch):
|
||||
"""
|
||||
Verify that every chunk's content is preserved in the output stream.
|
||||
|
||||
|
|
@ -442,10 +452,7 @@ class TestUnifiedLLMGuardrails:
|
|||
|
||||
return responses_so_far
|
||||
|
||||
# Override the mapping to use our content-clearing translation
|
||||
unified_module.endpoint_guardrail_translation_mappings = {
|
||||
CallTypes.acompletion: _ContentClearingTranslation,
|
||||
}
|
||||
_patch_translation_mappings(monkeypatch, {CallTypes.acompletion: _ContentClearingTranslation})
|
||||
|
||||
handler = UnifiedLLMGuardrails()
|
||||
guardrail = RecordingGuardrail()
|
||||
|
|
@ -885,12 +892,8 @@ class TestStreamingTransform:
|
|||
completions streaming surface."""
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _use_openai_handler_mapping(self):
|
||||
unified_module.endpoint_guardrail_translation_mappings = {
|
||||
CallTypes.acompletion: OpenAIChatCompletionsHandler,
|
||||
}
|
||||
yield
|
||||
unified_module.endpoint_guardrail_translation_mappings = None
|
||||
def _use_openai_handler_mapping(self, monkeypatch):
|
||||
_patch_translation_mappings(monkeypatch, {CallTypes.acompletion: OpenAIChatCompletionsHandler})
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_block_only_drops_text_rewrites(self):
|
||||
|
|
@ -1719,6 +1722,10 @@ class TestAppliedGuardrailsReflectsExecution:
|
|||
decision and marks itself only when it actually ran (LIT-4650). Ordinary
|
||||
guardrails are still auto-marked by the hook after dispatch."""
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _use_texts_only_mapping(self, monkeypatch):
|
||||
_patch_translation_mappings(monkeypatch, {CallTypes.pass_through: _TextsOnlyTranslation})
|
||||
|
||||
@staticmethod
|
||||
def _data(guardrail):
|
||||
return {
|
||||
|
|
@ -1728,7 +1735,6 @@ class TestAppliedGuardrailsReflectsExecution:
|
|||
}
|
||||
|
||||
async def _run(self, guardrail):
|
||||
unified_module.endpoint_guardrail_translation_mappings = {CallTypes.pass_through: _TextsOnlyTranslation}
|
||||
data = self._data(guardrail)
|
||||
await UnifiedLLMGuardrails().async_pre_call_hook(
|
||||
user_api_key_dict=None,
|
||||
|
|
@ -1830,10 +1836,8 @@ class TestStreamingHttpErrorFrames:
|
|||
silently truncates the SSE stream (PR #38722 defect 1)."""
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _use_real_mappings(self):
|
||||
unified_module.endpoint_guardrail_translation_mappings = load_guardrail_translation_mappings()
|
||||
yield
|
||||
unified_module.endpoint_guardrail_translation_mappings = None
|
||||
def _use_real_mappings(self, monkeypatch):
|
||||
_patch_translation_mappings(monkeypatch, load_guardrail_translation_mappings())
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_chat_eos_block_emits_data_error_frame(self):
|
||||
|
|
@ -1938,10 +1942,8 @@ class TestStreamingGuardrailInformationBucket:
|
|||
guardrail_information write was diverted and /spend/logs showed null."""
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _use_real_mappings(self):
|
||||
unified_module.endpoint_guardrail_translation_mappings = load_guardrail_translation_mappings()
|
||||
yield
|
||||
unified_module.endpoint_guardrail_translation_mappings = None
|
||||
def _use_real_mappings(self, monkeypatch):
|
||||
_patch_translation_mappings(monkeypatch, load_guardrail_translation_mappings())
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_chat_eos_scan_writes_guardrail_information_to_metadata(self):
|
||||
|
|
@ -2038,11 +2040,7 @@ class TestStreamingScanDedup:
|
|||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _use_real_mappings(self, monkeypatch):
|
||||
monkeypatch.setattr(
|
||||
unified_module,
|
||||
"endpoint_guardrail_translation_mappings",
|
||||
load_guardrail_translation_mappings(),
|
||||
)
|
||||
_patch_translation_mappings(monkeypatch, load_guardrail_translation_mappings())
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_chat_terminal_chunk_on_sampled_index_is_scanned_once(self):
|
||||
|
|
@ -2239,3 +2237,55 @@ class TestStreamingScanDedup:
|
|||
|
||||
assert out == chunks
|
||||
assert [scan["texts"] for scan in guardrail.scans] == [["abc"]]
|
||||
|
||||
|
||||
class TestTranslationMappingsAreReadLive:
|
||||
"""The hooks must read the handler map on every call, never memoize it on the module.
|
||||
|
||||
A second module-level cache is what let one test's handler map outlive its own
|
||||
teardown and decide how unrelated files translated their streams (LIT-6834).
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def _ocr_request(guardrail):
|
||||
return {
|
||||
"guardrail_to_apply": guardrail,
|
||||
"model": "mistral/mistral-ocr-latest",
|
||||
"document": {
|
||||
"type": "document_url",
|
||||
"document_url": "https://arxiv.org/pdf/2201.04234",
|
||||
},
|
||||
}
|
||||
|
||||
async def _run_pre_call(self, guardrail):
|
||||
await UnifiedLLMGuardrails().async_pre_call_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="test-key"),
|
||||
cache=DualCache(),
|
||||
data=self._ocr_request(guardrail),
|
||||
call_type=CallTypes.aocr.value,
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_remapping_between_calls_changes_which_handler_runs(self, monkeypatch):
|
||||
_patch_translation_mappings(monkeypatch, {CallTypes.completion: _NoopTranslation})
|
||||
unmapped = RecordingGuardrail()
|
||||
await self._run_pre_call(unmapped)
|
||||
assert unmapped.apply_calls == []
|
||||
|
||||
_patch_translation_mappings(monkeypatch, {CallTypes.aocr: OCRHandler})
|
||||
mapped = RecordingGuardrail()
|
||||
await self._run_pre_call(mapped)
|
||||
assert [call["input_type"] for call in mapped.apply_calls] == ["request"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_module_exposes_no_second_assignable_handler_map(self, monkeypatch):
|
||||
_patch_translation_mappings(monkeypatch, {CallTypes.aocr: OCRHandler})
|
||||
guardrail = RecordingGuardrail()
|
||||
await self._run_pre_call(guardrail)
|
||||
|
||||
assert len(guardrail.apply_calls) == 1
|
||||
assert not [
|
||||
name
|
||||
for name, value in vars(unified_module).items()
|
||||
if isinstance(value, dict) and CallTypes.aocr in value
|
||||
]
|
||||
|
|
|
|||
|
|
@ -802,7 +802,7 @@ async def test_bedrock_guardrail_prepare_request_with_api_key():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bedrock_guardrail_prepare_request_without_api_key():
|
||||
async def test_bedrock_guardrail_prepare_request_without_api_key(monkeypatch):
|
||||
"""Test _prepare_request method falls back to SigV4 when no api_key is provided"""
|
||||
from unittest.mock import Mock, patch
|
||||
|
||||
|
|
@ -820,18 +820,13 @@ async def test_bedrock_guardrail_prepare_request_without_api_key():
|
|||
|
||||
# Test data without api_key
|
||||
test_data = {"source": "INPUT", "content": [{"text": {"text": "test content"}}]}
|
||||
monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False)
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.guardrails.guardrail_hooks.bedrock_guardrails.get_secret_str"
|
||||
) as mock_get_secret,
|
||||
patch("botocore.auth.SigV4Auth") as mock_sigv4_auth,
|
||||
patch("botocore.awsrequest.AWSRequest") as mock_aws_request,
|
||||
):
|
||||
|
||||
# Mock no AWS_BEARER_TOKEN_BEDROCK
|
||||
mock_get_secret.return_value = None
|
||||
|
||||
# Mock SigV4Auth
|
||||
mock_sigv4_instance = Mock()
|
||||
mock_sigv4_auth.return_value = mock_sigv4_instance
|
||||
|
|
@ -857,7 +852,7 @@ async def test_bedrock_guardrail_prepare_request_without_api_key():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bedrock_guardrail_prepare_request_with_bearer_token_env():
|
||||
async def test_bedrock_guardrail_prepare_request_with_bearer_token_env(monkeypatch):
|
||||
"""Test _prepare_request method uses Bearer token from environment when available"""
|
||||
from unittest.mock import Mock, patch
|
||||
|
||||
|
|
@ -875,15 +870,9 @@ async def test_bedrock_guardrail_prepare_request_with_bearer_token_env():
|
|||
|
||||
# Test data without api_key
|
||||
test_data = {"source": "INPUT", "content": [{"text": {"text": "test content"}}]}
|
||||
monkeypatch.setenv("AWS_BEARER_TOKEN_BEDROCK", "env-bearer-token-456")
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.guardrails.guardrail_hooks.bedrock_guardrails.get_secret_str"
|
||||
) as mock_get_secret,
|
||||
patch("botocore.awsrequest.AWSRequest") as mock_aws_request,
|
||||
):
|
||||
|
||||
mock_get_secret.return_value = "env-bearer-token-456"
|
||||
with patch("botocore.awsrequest.AWSRequest") as mock_aws_request:
|
||||
mock_request_instance = Mock()
|
||||
mock_request_instance.prepare.return_value = Mock()
|
||||
mock_aws_request.return_value = mock_request_instance
|
||||
|
|
|
|||
|
|
@ -5,10 +5,13 @@ from typing import Any, Dict
|
|||
|
||||
import orjson
|
||||
import pytest
|
||||
from fastapi import FastAPI
|
||||
from fastapi.testclient import TestClient
|
||||
from starlette.requests import Request
|
||||
from starlette.responses import Response
|
||||
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.image_endpoints import endpoints
|
||||
|
||||
|
||||
|
|
@ -115,3 +118,52 @@ async def test_image_generation_prompt_rerouting(monkeypatch):
|
|||
assert captured_route_request_data["prompt"] == "sanitized prompt"
|
||||
assert "messages" not in captured_route_request_data
|
||||
assert response.headers.get("x-callback-test") == "value"
|
||||
|
||||
|
||||
def _image_edit_client(monkeypatch, captured: Dict[str, Any]) -> TestClient:
|
||||
class CaptureProcessing:
|
||||
def __init__(self, data: Dict[str, Any]) -> None:
|
||||
captured.update(data)
|
||||
|
||||
async def base_process_llm_request(self, **_: Any) -> Dict[str, Any]:
|
||||
return {"data": [{"b64_json": "aGk="}]}
|
||||
|
||||
monkeypatch.setattr(endpoints, "ProxyBaseLLMRequestProcessing", CaptureProcessing)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {})
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.user_model", None)
|
||||
|
||||
app = FastAPI()
|
||||
app.include_router(endpoints.router)
|
||||
app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth()
|
||||
return TestClient(app)
|
||||
|
||||
|
||||
def test_image_edit_multipart_n_reaches_the_provider_as_an_int(monkeypatch):
|
||||
"""A multipart `n` must not arrive as the string Starlette parsed it into."""
|
||||
captured: Dict[str, Any] = {}
|
||||
|
||||
response = _image_edit_client(monkeypatch, captured).post(
|
||||
"/v1/images/edits",
|
||||
files={"image": ("tree.png", b"\x89PNG\r\n\x1a\n", "image/png")},
|
||||
data={"model": "nova-canvas", "prompt": "add a hat", "n": "2", "size": "1024x1024"},
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
assert captured["n"] == 2
|
||||
assert isinstance(captured["n"], int)
|
||||
assert captured["size"] == "1024x1024"
|
||||
assert captured["prompt"] == "add a hat"
|
||||
|
||||
|
||||
def test_image_edit_multipart_n_that_is_not_a_number_is_left_alone(monkeypatch):
|
||||
"""An unparseable `n` still reaches the provider, which rejects it as before."""
|
||||
captured: Dict[str, Any] = {}
|
||||
|
||||
response = _image_edit_client(monkeypatch, captured).post(
|
||||
"/v1/images/edits",
|
||||
files={"image": ("tree.png", b"\x89PNG\r\n\x1a\n", "image/png")},
|
||||
data={"model": "nova-canvas", "prompt": "add a hat", "n": "two"},
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
assert captured["n"] == "two"
|
||||
|
|
|
|||
|
|
@ -292,8 +292,8 @@ class TestUnifiedGuardrailCallTypeResolution:
|
|||
|
||||
with patch.object(
|
||||
unified_guardrail_module,
|
||||
"endpoint_guardrail_translation_mappings",
|
||||
{CallTypes.pass_through: mock_handler_class},
|
||||
"load_guardrail_translation_mappings",
|
||||
lambda: {CallTypes.pass_through: mock_handler_class},
|
||||
):
|
||||
result = await unified.async_post_call_success_hook(
|
||||
data=data,
|
||||
|
|
|
|||
|
|
@ -421,6 +421,76 @@ def test_rag_query_store_params_win_over_user_retrieval_config(client_internal_u
|
|||
assert forwarded_config["aws_region_name"] == "eu-west-1"
|
||||
|
||||
|
||||
def test_rag_query_forwards_managed_store_credentials_to_search(client_internal_user):
|
||||
"""
|
||||
Regression for LIT-6773: the registry store's api_key / api_base and its
|
||||
provider extras (Milvus outputFields, milvus_text_field) must reach the
|
||||
vector store search the way the direct /v1/vector_stores/{id}/search
|
||||
endpoint forwards them. Pre-fix the RAG path allowlisted them away and a
|
||||
managed Milvus store 500'd with "MILVUS_API_KEY is not set".
|
||||
"""
|
||||
import litellm
|
||||
from litellm import Router
|
||||
from litellm.types.vector_stores import VectorStoreSearchResponse
|
||||
|
||||
mock_vector_store = {
|
||||
"vector_store_id": "customer_kb",
|
||||
"custom_llm_provider": "milvus",
|
||||
"litellm_params": {
|
||||
"vector_store_id": "customer_kb",
|
||||
"custom_llm_provider": "milvus",
|
||||
"api_base": "http://127.0.0.1:19530",
|
||||
"api_key": "root:Milvus",
|
||||
"litellm_embedding_model": "multilingual-e5-large",
|
||||
"milvus_text_field": "book_intro_text",
|
||||
"outputFields": ["book_intro_text"],
|
||||
},
|
||||
}
|
||||
mock_registry = MagicMock()
|
||||
mock_registry.get_litellm_managed_vector_store_from_registry.return_value = mock_vector_store
|
||||
fake_search = AsyncMock(
|
||||
return_value=VectorStoreSearchResponse(object="vector_store.search_results.page", search_query="q", data=[])
|
||||
)
|
||||
router = Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "gpt-4o-mini",
|
||||
"litellm_params": {"model": "openai/gpt-4o-mini", "api_key": "sk-test", "mock_response": "hi"},
|
||||
}
|
||||
]
|
||||
)
|
||||
|
||||
with (
|
||||
patch("litellm.vector_stores.asearch", new=fake_search), # test-quality-ok: the search boundary under test
|
||||
patch.object(litellm, "vector_store_registry", mock_registry), # test-quality-ok: seeds the store under test
|
||||
patch("litellm.proxy.proxy_server.llm_router", router), # test-quality-ok: mock-response router for completion
|
||||
patch( # test-quality-ok: store access is not under test, so the request reaches the search boundary
|
||||
"litellm.proxy.vector_store_endpoints.utils.can_user_access_vector_store",
|
||||
new=AsyncMock(return_value=True),
|
||||
),
|
||||
):
|
||||
response = client_internal_user.post(
|
||||
"/v1/rag/query",
|
||||
json={
|
||||
"model": "gpt-4o-mini",
|
||||
"messages": [{"role": "user", "content": "which database is built for similarity search?"}],
|
||||
"retrieval_config": {"vector_store_id": "customer_kb", "custom_llm_provider": "milvus", "top_k": 2},
|
||||
},
|
||||
)
|
||||
|
||||
assert response.status_code == 200, response.json()
|
||||
fake_search.assert_awaited_once()
|
||||
search_kwargs = fake_search.await_args.kwargs
|
||||
assert search_kwargs["vector_store_id"] == "customer_kb"
|
||||
assert search_kwargs["custom_llm_provider"] == "milvus"
|
||||
assert search_kwargs["max_num_results"] == 2
|
||||
assert search_kwargs["api_base"] == "http://127.0.0.1:19530"
|
||||
assert search_kwargs["api_key"] == "root:Milvus"
|
||||
assert search_kwargs["litellm_embedding_model"] == "multilingual-e5-large"
|
||||
assert search_kwargs["milvus_text_field"] == "book_intro_text"
|
||||
assert search_kwargs["outputFields"] == ["book_intro_text"]
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"blocked_key",
|
||||
["embedding_model", "litellm_embedding_model", "litellm_embedding_config", "litellm_credential_name"],
|
||||
|
|
|
|||
|
|
@ -3399,6 +3399,138 @@ def test_get_spend_logs_id_prefers_the_response_id_over_the_standard_logging_id(
|
|||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_spend_log_request_id_is_the_message_id_a_bridged_streaming_caller_was_streamed():
|
||||
"""A streaming /v1/messages call against a non-Anthropic model is served a msg_ id the
|
||||
adapter mints itself, and it is the only request id that call ever shows the caller, so
|
||||
GET /spend/logs?request_id=msg_... has to land on the row."""
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
from litellm.llms.anthropic.experimental_pass_through.responses_adapters.streaming_iterator import (
|
||||
AnthropicResponsesStreamWrapper,
|
||||
)
|
||||
from litellm.types.llms.openai import (
|
||||
ResponseAPIUsage,
|
||||
ResponseCompletedEvent,
|
||||
ResponsesAPIResponse,
|
||||
)
|
||||
|
||||
logging_obj = Logging(
|
||||
model="gpt-5.6",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
stream=True,
|
||||
call_type="anthropic_messages",
|
||||
start_time=datetime.datetime.now(timezone.utc),
|
||||
litellm_call_id="6825cafe-0000-4000-8000-000000000001",
|
||||
function_id="1234",
|
||||
)
|
||||
logging_obj.optional_params = {}
|
||||
|
||||
completed_response = ResponsesAPIResponse(
|
||||
id="resp_01Lit6825Bridged",
|
||||
object="response",
|
||||
created_at=1767225600,
|
||||
model="gpt-5.6",
|
||||
status="completed",
|
||||
output=[
|
||||
{
|
||||
"id": "msg_bridged_output",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"status": "completed",
|
||||
"content": [{"type": "output_text", "text": "epsilon", "annotations": []}],
|
||||
}
|
||||
],
|
||||
usage=ResponseAPIUsage(input_tokens=12, output_tokens=5, total_tokens=17),
|
||||
)
|
||||
|
||||
async def _responses_stream():
|
||||
yield {"type": "response.created"}
|
||||
yield {"type": "response.output_text.delta", "item_id": "msg_bridged_output", "delta": "epsilon"}
|
||||
yield ResponseCompletedEvent(type="response.completed", response=completed_response)
|
||||
|
||||
wrapper = AnthropicResponsesStreamWrapper(
|
||||
responses_stream=_responses_stream(),
|
||||
model="gpt-5.6",
|
||||
litellm_logging_obj=logging_obj,
|
||||
)
|
||||
sse_frames = [frame.decode() async for frame in wrapper.async_anthropic_sse_wrapper()]
|
||||
|
||||
message_start_frames = [f for f in sse_frames if f.startswith("event: message_start\n")]
|
||||
assert len(message_start_frames) == 1
|
||||
streamed_message_id = json.loads(message_start_frames[0].split("data: ", 1)[1])["message"]["id"]
|
||||
assert streamed_message_id.startswith("msg_")
|
||||
|
||||
_, _, logged_response = logging_obj._success_handler_helper_fn(
|
||||
result=ResponseCompletedEvent(type="response.completed", response=completed_response),
|
||||
start_time=datetime.datetime.now(timezone.utc),
|
||||
end_time=datetime.datetime.now(timezone.utc),
|
||||
)
|
||||
|
||||
assert logged_response.id == streamed_message_id
|
||||
payload = get_logging_payload(
|
||||
kwargs={
|
||||
"call_type": "anthropic_messages",
|
||||
"model": "gpt-5.6",
|
||||
"litellm_call_id": "6825cafe-0000-4000-8000-000000000001",
|
||||
"litellm_params": {"metadata": {"user_api_key": "test-key"}},
|
||||
},
|
||||
response_obj=logged_response,
|
||||
start_time=datetime.datetime.now(timezone.utc),
|
||||
end_time=datetime.datetime.now(timezone.utc),
|
||||
)
|
||||
assert payload["request_id"] == streamed_message_id
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_spend_log_request_id_is_untouched_when_no_message_id_was_streamed():
|
||||
"""Only the bridged streaming adapter mints a msg_ id of its own, so every other
|
||||
/v1/messages call must keep the id its own response carried."""
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
from litellm.types.llms.openai import (
|
||||
ResponseAPIUsage,
|
||||
ResponseCompletedEvent,
|
||||
ResponsesAPIResponse,
|
||||
)
|
||||
|
||||
logging_obj = Logging(
|
||||
model="gpt-5.6",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
stream=False,
|
||||
call_type="anthropic_messages",
|
||||
start_time=datetime.datetime.now(timezone.utc),
|
||||
litellm_call_id="6825cafe-0000-4000-8000-000000000002",
|
||||
function_id="1234",
|
||||
)
|
||||
logging_obj.optional_params = {}
|
||||
|
||||
completed_response = ResponsesAPIResponse(
|
||||
id="resp_01Lit6825Unbridged",
|
||||
object="response",
|
||||
created_at=1767225600,
|
||||
model="gpt-5.6",
|
||||
status="completed",
|
||||
output=[
|
||||
{
|
||||
"id": "msg_unbridged_output",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"status": "completed",
|
||||
"content": [{"type": "output_text", "text": "epsilon", "annotations": []}],
|
||||
}
|
||||
],
|
||||
usage=ResponseAPIUsage(input_tokens=12, output_tokens=5, total_tokens=17),
|
||||
)
|
||||
|
||||
_, _, logged_response = logging_obj._success_handler_helper_fn(
|
||||
result=ResponseCompletedEvent(type="response.completed", response=completed_response),
|
||||
start_time=datetime.datetime.now(timezone.utc),
|
||||
end_time=datetime.datetime.now(timezone.utc),
|
||||
)
|
||||
|
||||
assert logged_response.id
|
||||
assert not logged_response.id.startswith("msg_")
|
||||
|
||||
|
||||
def test_batch_cost_row_does_not_collide_with_the_batch_creation_row():
|
||||
"""Creating a batch writes a row keyed by the batch's own id, so keying the cost row
|
||||
the same way makes the insert a duplicate of it. request_id is the primary key and the
|
||||
|
|
@ -3956,3 +4088,207 @@ def test_caller_forged_router_metadata_is_discarded(bucket):
|
|||
)
|
||||
metadata = json.loads(payload["metadata"])
|
||||
assert metadata["router_metadata"] is None
|
||||
|
||||
|
||||
ANTHROPIC_MESSAGES_RESPONSE: Final = {
|
||||
"id": "msg_01Lit6806NonStreaming",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"model": "claude-haiku-4-5",
|
||||
"content": [{"type": "text", "text": "epsilon"}],
|
||||
"stop_reason": "end_turn",
|
||||
"stop_sequence": None,
|
||||
"usage": {"input_tokens": 14, "output_tokens": 4},
|
||||
}
|
||||
|
||||
ANTHROPIC_MESSAGES_SSE_CHUNKS: Final = (
|
||||
'event: message_start\ndata: {"type":"message_start","message":{"id":"msg_01Lit6806Streaming",'
|
||||
'"type":"message","role":"assistant","model":"claude-haiku-4-5","content":[],'
|
||||
'"usage":{"input_tokens":14,"output_tokens":1}}}\n\n',
|
||||
'event: content_block_start\ndata: {"type":"content_block_start","index":0,'
|
||||
'"content_block":{"type":"text","text":""}}\n\n',
|
||||
'event: content_block_delta\ndata: {"type":"content_block_delta","index":0,'
|
||||
'"delta":{"type":"text_delta","text":"epsilon"}}\n\n',
|
||||
'event: content_block_stop\ndata: {"type":"content_block_stop","index":0}\n\n',
|
||||
'event: message_delta\ndata: {"type":"message_delta","delta":{"stop_reason":"end_turn"},'
|
||||
'"usage":{"output_tokens":4}}\n\n',
|
||||
"event: message_stop\ndata: {\"type\":\"message_stop\"}\n\n",
|
||||
)
|
||||
|
||||
|
||||
def _anthropic_messages_logging_obj(*, stream: bool) -> Any:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
|
||||
logging_obj = Logging(
|
||||
model="claude-haiku-4-5",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
stream=stream,
|
||||
call_type="anthropic_messages",
|
||||
start_time=datetime.datetime.now(timezone.utc),
|
||||
litellm_call_id="6806cafe-0000-4000-8000-000000000001",
|
||||
function_id="1234",
|
||||
)
|
||||
logging_obj.optional_params = {}
|
||||
logging_obj.model_call_details["custom_llm_provider"] = "anthropic"
|
||||
return logging_obj
|
||||
|
||||
|
||||
def _spend_log_request_id(response_obj: Any, kwargs: dict) -> str:
|
||||
payload = get_logging_payload(
|
||||
kwargs=kwargs,
|
||||
response_obj=response_obj,
|
||||
start_time=datetime.datetime.now(timezone.utc),
|
||||
end_time=datetime.datetime.now(timezone.utc),
|
||||
)
|
||||
return payload["request_id"]
|
||||
|
||||
|
||||
def test_spend_log_request_id_is_the_message_id_a_non_streaming_messages_caller_received():
|
||||
"""
|
||||
POST /v1/messages hands the caller `id: msg_...`, the only request id they ever see, so
|
||||
GET /spend/logs?request_id=msg_... has to find the row.
|
||||
"""
|
||||
logging_obj = _anthropic_messages_logging_obj(stream=False)
|
||||
|
||||
logged_response = logging_obj._handle_anthropic_messages_response_logging(
|
||||
result=ANTHROPIC_MESSAGES_RESPONSE
|
||||
)
|
||||
|
||||
assert logged_response.id == "msg_01Lit6806NonStreaming"
|
||||
assert (
|
||||
_spend_log_request_id(
|
||||
response_obj=logged_response,
|
||||
kwargs={
|
||||
"call_type": "anthropic_messages",
|
||||
"model": "claude-haiku-4-5",
|
||||
"litellm_call_id": "6806cafe-0000-4000-8000-000000000001",
|
||||
"litellm_params": {"metadata": {"user_api_key": "test-key"}},
|
||||
},
|
||||
)
|
||||
== "msg_01Lit6806NonStreaming"
|
||||
)
|
||||
|
||||
|
||||
def test_spend_log_request_id_is_the_message_id_a_streaming_messages_caller_received():
|
||||
"""
|
||||
The streaming leg of /v1/messages logs through the Anthropic passthrough handler, which used
|
||||
to stamp litellm_call_id over the msg_ id carried by the message_start event.
|
||||
"""
|
||||
from litellm.proxy.pass_through_endpoints.llm_provider_handlers.anthropic_passthrough_logging_handler import (
|
||||
AnthropicPassthroughLoggingHandler,
|
||||
)
|
||||
from litellm.types.passthrough_endpoints.pass_through_endpoints import EndpointType
|
||||
|
||||
logging_obj = _anthropic_messages_logging_obj(stream=True)
|
||||
logging_obj.model_call_details["stream"] = True
|
||||
|
||||
logged = AnthropicPassthroughLoggingHandler._handle_logging_anthropic_collected_chunks(
|
||||
litellm_logging_obj=logging_obj,
|
||||
passthrough_success_handler_obj=MagicMock(),
|
||||
url_route="/v1/messages",
|
||||
request_body={"model": "claude-haiku-4-5"},
|
||||
endpoint_type=EndpointType.ANTHROPIC,
|
||||
start_time=datetime.datetime.now(timezone.utc),
|
||||
all_chunks=list(ANTHROPIC_MESSAGES_SSE_CHUNKS),
|
||||
end_time=datetime.datetime.now(timezone.utc),
|
||||
)
|
||||
|
||||
assert logged["result"].id == "msg_01Lit6806Streaming"
|
||||
assert (
|
||||
_spend_log_request_id(
|
||||
response_obj=logged["result"],
|
||||
kwargs={
|
||||
**logged["kwargs"],
|
||||
"call_type": "anthropic_messages",
|
||||
"litellm_call_id": "6806cafe-0000-4000-8000-000000000001",
|
||||
"litellm_params": {"metadata": {"user_api_key": "test-key"}},
|
||||
},
|
||||
)
|
||||
== "msg_01Lit6806Streaming"
|
||||
)
|
||||
|
||||
|
||||
def test_spend_log_request_id_still_falls_back_to_litellm_call_id_without_a_provider_id():
|
||||
"""
|
||||
Anthropic-compatible upstreams that omit `id` must keep landing on litellm_call_id rather
|
||||
than on a fresh chatcmpl- uuid nobody can look up.
|
||||
"""
|
||||
logging_obj = _anthropic_messages_logging_obj(stream=True)
|
||||
logging_obj.model_call_details["stream"] = True
|
||||
|
||||
from litellm.proxy.pass_through_endpoints.llm_provider_handlers.anthropic_passthrough_logging_handler import (
|
||||
AnthropicPassthroughLoggingHandler,
|
||||
)
|
||||
|
||||
AnthropicPassthroughLoggingHandler._create_anthropic_response_logging_payload(
|
||||
litellm_model_response=litellm.ModelResponse(id="chatcmpl-generated"),
|
||||
model="claude-haiku-4-5",
|
||||
kwargs={},
|
||||
start_time=datetime.datetime.now(timezone.utc),
|
||||
end_time=datetime.datetime.now(timezone.utc),
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
assert logging_obj.model_call_details["complete_streaming_response"].id == (
|
||||
"6806cafe-0000-4000-8000-000000000001"
|
||||
)
|
||||
|
||||
|
||||
def test_spend_log_request_id_for_chat_completions_is_untouched():
|
||||
"""
|
||||
/v1/chat/completions callers look their rows up by the chatcmpl- id in the response body.
|
||||
"""
|
||||
assert (
|
||||
_spend_log_request_id(
|
||||
response_obj=litellm.ModelResponse(id="chatcmpl-EJvWIw3DAhuKYuwp3jJI4Pnhp2vjv", choices=[]),
|
||||
kwargs={
|
||||
"call_type": "acompletion",
|
||||
"model": "gpt-5.6",
|
||||
"litellm_call_id": "6806cafe-0000-4000-8000-000000000002",
|
||||
"litellm_params": {"metadata": {"user_api_key": "test-key"}},
|
||||
},
|
||||
)
|
||||
== "chatcmpl-EJvWIw3DAhuKYuwp3jJI4Pnhp2vjv"
|
||||
)
|
||||
|
||||
|
||||
def test_spend_log_request_id_is_the_response_id_a_bridged_messages_caller_received():
|
||||
"""
|
||||
/v1/messages against a non-Anthropic model answers with the Responses id the caller then
|
||||
looks their row up by, so the row must not fall back to a fresh chatcmpl- uuid.
|
||||
"""
|
||||
from litellm.types.llms.openai import ResponseAPIUsage, ResponsesAPIResponse
|
||||
|
||||
logging_obj = _anthropic_messages_logging_obj(stream=False)
|
||||
bridged_response = ResponsesAPIResponse(
|
||||
id="resp_01Lit6806Bridged",
|
||||
object="response",
|
||||
created_at=1767225600,
|
||||
model="gpt-5.6",
|
||||
status="completed",
|
||||
output=[
|
||||
{
|
||||
"id": "msg_bridged_output",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"status": "completed",
|
||||
"content": [{"type": "output_text", "text": "delta", "annotations": []}],
|
||||
}
|
||||
],
|
||||
usage=ResponseAPIUsage(input_tokens=13, output_tokens=5, total_tokens=18),
|
||||
)
|
||||
|
||||
logged_response = logging_obj._handle_anthropic_messages_response_logging(result=bridged_response)
|
||||
|
||||
assert logged_response.id == "resp_01Lit6806Bridged"
|
||||
assert (
|
||||
_spend_log_request_id(
|
||||
response_obj=logged_response,
|
||||
kwargs={
|
||||
"call_type": "anthropic_messages",
|
||||
"model": "gpt-5.6",
|
||||
"litellm_call_id": "6806cafe-0000-4000-8000-000000000003",
|
||||
"litellm_params": {"metadata": {"user_api_key": "test-key"}},
|
||||
},
|
||||
)
|
||||
== "resp_01Lit6806Bridged"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -68,12 +68,10 @@ async def test_success_hook_attaches_original_response_on_block():
|
|||
user_api_key_dict = UserAPIKeyAuth(api_key="test", request_route="/chat/completions")
|
||||
data = {"guardrail_to_apply": guardrail, "model": "gpt-4o"}
|
||||
|
||||
# Inject our translation for the inferred call type (the module global is
|
||||
# cached across tests, so patch it directly rather than the loader).
|
||||
with patch.object(
|
||||
ug,
|
||||
"endpoint_guardrail_translation_mappings",
|
||||
{
|
||||
"load_guardrail_translation_mappings",
|
||||
lambda: {
|
||||
CallTypes.acompletion: lambda: translation,
|
||||
CallTypes.completion: lambda: translation,
|
||||
},
|
||||
|
|
|
|||
|
|
@ -1540,8 +1540,8 @@ class TestCommonRequestProcessingHelpers:
|
|||
expected_error_data = {
|
||||
"error": {
|
||||
"message": "Error processing stream start",
|
||||
"type": "None",
|
||||
"param": "None",
|
||||
"type": "internal_server_error",
|
||||
"param": None,
|
||||
"code": str(status.HTTP_500_INTERNAL_SERVER_ERROR),
|
||||
}
|
||||
}
|
||||
|
|
@ -1569,8 +1569,8 @@ class TestCommonRequestProcessingHelpers:
|
|||
expected_error_data = {
|
||||
"error": {
|
||||
"message": "Content blocked by guardrail",
|
||||
"type": "None",
|
||||
"param": "None",
|
||||
"type": "invalid_request_error",
|
||||
"param": None,
|
||||
"code": "400",
|
||||
}
|
||||
}
|
||||
|
|
@ -1934,6 +1934,104 @@ class TestCommonRequestProcessingHelpers:
|
|||
assert mock_tracer.trace.call_count == 0
|
||||
|
||||
|
||||
def _stringified_none_paths(node: object, path: str = "error") -> tuple[str, ...]:
|
||||
if isinstance(node, dict):
|
||||
return tuple(
|
||||
found
|
||||
for key, value in node.items()
|
||||
for found in _stringified_none_paths(value, f"{path}.{key}")
|
||||
)
|
||||
if isinstance(node, (list, tuple)):
|
||||
return tuple(
|
||||
found
|
||||
for index, value in enumerate(node)
|
||||
for found in _stringified_none_paths(value, f"{path}[{index}]")
|
||||
)
|
||||
return (path,) if node == "None" else ()
|
||||
|
||||
|
||||
def _blocked_guardrail_exception() -> HTTPException:
|
||||
return HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": "Violated guardrail policy",
|
||||
"bedrock_guardrail_response": {"action": "GUARDRAIL_INTERVENED"},
|
||||
"guardrailIdentifier": "gf3sc1mzinjw",
|
||||
"guardrailVersion": "DRAFT",
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
class TestGuardrailBlockErrorPayloadNeverStringifiesNone:
|
||||
"""Regression for LIT-6808: a blocked-guardrail error body carried the literal string
|
||||
"None" for type and param instead of a real error type and JSON null."""
|
||||
|
||||
def test_non_streaming_block_payload_carries_a_real_type_and_null_param(self):
|
||||
from litellm.proxy.common_request_processing import (
|
||||
proxy_exception_from_http_exception,
|
||||
)
|
||||
|
||||
payload = json.loads(
|
||||
json.dumps(proxy_exception_from_http_exception(_blocked_guardrail_exception(), {}).to_dict())
|
||||
)
|
||||
|
||||
assert _stringified_none_paths(payload) == ()
|
||||
assert payload["type"] == "invalid_request_error"
|
||||
assert payload["param"] is None
|
||||
assert payload["code"] == "400"
|
||||
assert payload["message"] == "Violated guardrail policy"
|
||||
|
||||
def test_streaming_block_frame_carries_a_real_type_and_null_param(self):
|
||||
from litellm.proxy.common_request_processing import sse_error_payload
|
||||
|
||||
error_status, error_obj = sse_error_payload(_blocked_guardrail_exception())
|
||||
frame = json.loads(json.dumps({"error": dict(error_obj)}))
|
||||
|
||||
assert error_status == 400
|
||||
assert _stringified_none_paths(frame["error"]) == ()
|
||||
assert frame["error"]["type"] == "invalid_request_error"
|
||||
assert frame["error"]["param"] is None
|
||||
assert frame["error"]["code"] == "400"
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"status_code, expected_type",
|
||||
[
|
||||
(400, "invalid_request_error"),
|
||||
(401, "authentication_error"),
|
||||
(403, "permission_error"),
|
||||
(404, "invalid_request_error"),
|
||||
(429, "rate_limit_error"),
|
||||
(500, "internal_server_error"),
|
||||
(503, "internal_server_error"),
|
||||
],
|
||||
)
|
||||
def test_status_code_decides_the_type_when_the_exception_carries_none(self, status_code, expected_type):
|
||||
from litellm.proxy.common_request_processing import (
|
||||
proxy_exception_from_http_exception,
|
||||
)
|
||||
|
||||
payload = proxy_exception_from_http_exception(
|
||||
HTTPException(status_code=status_code, detail="blocked"), {}
|
||||
).to_dict()
|
||||
|
||||
assert payload["type"] == expected_type
|
||||
assert payload["param"] is None
|
||||
|
||||
def test_a_type_and_param_the_exception_carries_win_over_the_fallback(self):
|
||||
from litellm.proxy.common_request_processing import (
|
||||
proxy_exception_from_http_exception,
|
||||
)
|
||||
|
||||
exc = HTTPException(status_code=400, detail="unknown model")
|
||||
exc.type = "authentication_error"
|
||||
exc.param = "model"
|
||||
|
||||
payload = proxy_exception_from_http_exception(exc, {}).to_dict()
|
||||
|
||||
assert payload["type"] == "authentication_error"
|
||||
assert payload["param"] == "model"
|
||||
|
||||
|
||||
class TestExtractErrorFromSSEChunk:
|
||||
"""Tests for _extract_error_from_sse_chunk function"""
|
||||
|
||||
|
|
@ -2999,6 +3097,25 @@ class TestHandleLLMApiExceptionDictDetail:
|
|||
assert proxy_exc.message == "Content blocked by guardrail"
|
||||
assert proxy_exc.provider_specific_fields is None
|
||||
|
||||
async def test_blocked_guardrail_error_body_never_carries_the_string_none(self):
|
||||
"""Regression for LIT-6808: the error body a blocked request returns must carry a real
|
||||
error type and JSON null rather than the literal string "None"."""
|
||||
proxy_exc = await self._invoke(_blocked_guardrail_exception())
|
||||
payload = json.loads(json.dumps(proxy_exc.to_dict()))
|
||||
|
||||
assert _stringified_none_paths(payload) == ()
|
||||
assert payload["type"] == "invalid_request_error"
|
||||
assert payload["param"] is None
|
||||
|
||||
async def test_unclassified_exception_error_body_never_carries_the_string_none(self):
|
||||
"""The same holds on the generic fallback, where nothing carries a type at all."""
|
||||
proxy_exc = await self._invoke(ValueError("Something broke"))
|
||||
payload = json.loads(json.dumps(proxy_exc.to_dict()))
|
||||
|
||||
assert _stringified_none_paths(payload) == ()
|
||||
assert payload["type"] == "internal_server_error"
|
||||
assert payload["param"] is None
|
||||
|
||||
async def test_not_found_error_preserves_404(self):
|
||||
"""NotFoundError with status_code=404 should map to ProxyException code=404."""
|
||||
from litellm.exceptions import NotFoundError
|
||||
|
|
|
|||
|
|
@ -388,6 +388,72 @@ async def test_aquery_does_not_forward_connection_override_keys_to_search():
|
|||
assert not (blocked & set(search_kwargs.keys()))
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aquery_forwards_vector_store_params_to_search_but_not_completion():
|
||||
"""
|
||||
Regression for LIT-6773: the server-trusted vector_store_params (a managed
|
||||
store's litellm_params) must reach the search call wholesale, including the
|
||||
connection keys the caller allowlist blocks, while the caller's own
|
||||
retrieval_config overrides stay blocked, the caller's top-level api_key and
|
||||
api_base stay on the completion only, and the completion never inherits the
|
||||
store's connection params.
|
||||
"""
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
from litellm.types.vector_stores import VectorStoreSearchResponse
|
||||
|
||||
fake_search = AsyncMock(
|
||||
return_value=VectorStoreSearchResponse(
|
||||
object="vector_store.search_results.page", search_query="q", data=[]
|
||||
)
|
||||
)
|
||||
fake_completion = AsyncMock(
|
||||
return_value=ModelResponse(
|
||||
id="chatcmpl-test",
|
||||
choices=[{"index": 0, "message": {"role": "assistant", "content": "hi"}, "finish_reason": "stop"}],
|
||||
model="gpt-4o-mini",
|
||||
)
|
||||
)
|
||||
with (
|
||||
patch("litellm.vector_stores.asearch", new=fake_search), # test-quality-ok: the search boundary under test
|
||||
patch("litellm.acompletion", new=fake_completion), # test-quality-ok: the completion boundary under test
|
||||
):
|
||||
await litellm.aquery(
|
||||
model="gpt-4o-mini",
|
||||
messages=[{"role": "user", "content": "hello"}],
|
||||
api_key="sk-llm-key",
|
||||
api_base="https://llm.example.com",
|
||||
retrieval_config={
|
||||
"vector_store_id": "customer_kb",
|
||||
"custom_llm_provider": "milvus",
|
||||
"api_base": "https://attacker.example.com",
|
||||
"api_key": "attacker-key",
|
||||
},
|
||||
vector_store_params={
|
||||
"vector_store_id": "customer_kb",
|
||||
"custom_llm_provider": "milvus",
|
||||
"api_base": "http://127.0.0.1:19530",
|
||||
"api_key": "root:Milvus",
|
||||
"milvus_text_field": "book_intro_text",
|
||||
"outputFields": ["book_intro_text"],
|
||||
},
|
||||
)
|
||||
|
||||
fake_search.assert_awaited_once()
|
||||
search_kwargs = fake_search.await_args.kwargs
|
||||
assert search_kwargs["vector_store_id"] == "customer_kb"
|
||||
assert search_kwargs["custom_llm_provider"] == "milvus"
|
||||
assert search_kwargs["api_base"] == "http://127.0.0.1:19530"
|
||||
assert search_kwargs["api_key"] == "root:Milvus"
|
||||
assert search_kwargs["milvus_text_field"] == "book_intro_text"
|
||||
assert search_kwargs["outputFields"] == ["book_intro_text"]
|
||||
fake_completion.assert_awaited_once()
|
||||
completion_kwargs = fake_completion.await_args.kwargs
|
||||
assert completion_kwargs["api_key"] == "sk-llm-key"
|
||||
assert completion_kwargs["api_base"] == "https://llm.example.com"
|
||||
assert not ({"milvus_text_field", "outputFields"} & set(completion_kwargs))
|
||||
|
||||
|
||||
def test_rag_call_types_are_registered():
|
||||
"""
|
||||
query/aquery/ingest/aingest are @client-decorated entry points, so their
|
||||
|
|
|
|||
55
tests/test_litellm/test_azure_ai_grok_4_6_model_metadata.py
Normal file
55
tests/test_litellm/test_azure_ai_grok_4_6_model_metadata.py
Normal file
|
|
@ -0,0 +1,55 @@
|
|||
from pathlib import Path
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
from litellm import cost_per_token, get_model_info
|
||||
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
|
||||
|
||||
REPO_ROOT: Final = Path(__file__).parents[2]
|
||||
MODEL: Final = "azure_ai/grok-4.6"
|
||||
SOURCE: Final = (
|
||||
"https://techcommunity.microsoft.com/blog/azure-ai-foundry-blog/"
|
||||
"grok-4-6-comes-to-microsoft-foundry-models-built-for-long-horizon-reasoning-and-/4547578"
|
||||
)
|
||||
COST_MAP_ADAPTER: Final = TypeAdapter(dict[str, dict[str, object]])
|
||||
|
||||
|
||||
def _cost_map_entry(path: Path) -> dict[str, object]:
|
||||
return COST_MAP_ADAPTER.validate_json(path.read_bytes())[MODEL]
|
||||
|
||||
|
||||
@pytest.mark.usefixtures("local_model_cost_map")
|
||||
def test_azure_ai_grok_4_6_is_priced_and_routed() -> None:
|
||||
routed_model, provider, _, _ = get_llm_provider(model=MODEL)
|
||||
assert (routed_model, provider) == ("grok-4.6", "azure_ai")
|
||||
|
||||
info = get_model_info(model=routed_model, custom_llm_provider=provider)
|
||||
assert info["litellm_provider"] == "azure_ai"
|
||||
assert info["mode"] == "chat"
|
||||
assert info["input_cost_per_token"] == 2e-06
|
||||
assert info["output_cost_per_token"] == 6e-06
|
||||
assert info["cache_read_input_token_cost"] == 5e-07
|
||||
assert info["max_input_tokens"] == 200000
|
||||
assert info["max_output_tokens"] == 128000
|
||||
assert info["max_tokens"] == 128000
|
||||
assert info["supports_function_calling"] is True
|
||||
assert info["supports_prompt_caching"] is True
|
||||
assert info["supports_reasoning"] is True
|
||||
assert info["supports_response_schema"] is True
|
||||
assert info["supports_tool_choice"] is True
|
||||
assert info["supports_vision"] is True
|
||||
assert info["supports_web_search"] is True
|
||||
|
||||
prompt_cost, completion_cost = cost_per_token(model=MODEL, prompt_tokens=1_000_000, completion_tokens=1_000_000)
|
||||
assert prompt_cost == pytest.approx(2.0)
|
||||
assert completion_cost == pytest.approx(6.0)
|
||||
|
||||
|
||||
def test_azure_ai_grok_4_6_entry_source_and_backup_match() -> None:
|
||||
main_entry = _cost_map_entry(REPO_ROOT / "model_prices_and_context_window.json")
|
||||
backup_entry = _cost_map_entry(REPO_ROOT / "litellm" / "model_prices_and_context_window_backup.json")
|
||||
|
||||
assert main_entry["source"] == SOURCE
|
||||
assert backup_entry == main_entry
|
||||
|
|
@ -9275,9 +9275,11 @@ class _FallbackAttemptRecorder(CustomLogger):
|
|||
def __init__(self):
|
||||
super().__init__()
|
||||
self.failed_targets = []
|
||||
self.breadcrumbs_per_target = []
|
||||
|
||||
async def log_failure_fallback_event(self, original_model_group, kwargs, original_exception):
|
||||
self.failed_targets.append(kwargs.get("model"))
|
||||
self.breadcrumbs_per_target.append(kwargs.get("metadata", {}).get("previous_models", ()))
|
||||
|
||||
|
||||
def _cyclic_fallback_router(num_retries=0):
|
||||
|
|
@ -9348,14 +9350,16 @@ async def test_retry_breadcrumbs_do_not_carry_the_walk_state():
|
|||
A retry has to be configured for the walk state to reach log_retry at all."""
|
||||
router = _cyclic_fallback_router(num_retries=1)
|
||||
capture = _LogCapture(logging.ERROR)
|
||||
recorder = _FallbackAttemptRecorder()
|
||||
|
||||
await _drive_cyclic_fallback(router, capture)
|
||||
await _drive_cyclic_fallback(router, capture, recorder)
|
||||
|
||||
assert router.previous_models, "no retry breadcrumbs were recorded"
|
||||
breadcrumbs = [breadcrumb for hop in recorder.breadcrumbs_per_target for breadcrumb in hop]
|
||||
assert breadcrumbs, "no retry breadcrumbs were recorded"
|
||||
assert any(
|
||||
"fallback_depth" in breadcrumb for breadcrumb in router.previous_models
|
||||
"fallback_depth" in breadcrumb for breadcrumb in breadcrumbs
|
||||
), "no breadcrumb carried router walk state, so this test cannot see the leak"
|
||||
for breadcrumb in router.previous_models:
|
||||
for breadcrumb in breadcrumbs:
|
||||
assert "attempted_targets" not in breadcrumb
|
||||
|
||||
|
||||
|
|
@ -9393,15 +9397,94 @@ async def test_retry_breadcrumbs_never_carry_a_forwarded_credential(container_ke
|
|||
container still reaches the breadcrumb, but the raw secret never does, whatever key holds it."""
|
||||
router = _cyclic_fallback_router(num_retries=1)
|
||||
capture = _LogCapture(logging.ERROR)
|
||||
metadata = {}
|
||||
|
||||
await _drive_cyclic_fallback(router, capture, **request_kwargs)
|
||||
await _drive_cyclic_fallback(router, capture, metadata=metadata, **request_kwargs)
|
||||
|
||||
assert router.previous_models, "no retry breadcrumbs were recorded"
|
||||
dumped = json.dumps(router.previous_models, default=str)
|
||||
breadcrumbs = metadata["previous_models"]
|
||||
assert breadcrumbs, "no retry breadcrumbs were recorded"
|
||||
dumped = json.dumps(breadcrumbs, default=str)
|
||||
assert container_key in dumped, "the credential-bearing kwarg never reached the breadcrumb, so this test cannot see the leak"
|
||||
assert _BREADCRUMB_CREDENTIAL_CANARY not in dumped
|
||||
|
||||
|
||||
def _always_failing_router(num_retries):
|
||||
return litellm.Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "broken-group",
|
||||
"litellm_params": {
|
||||
"model": "openai/gpt-4o-mini",
|
||||
"api_key": "sk-fake",
|
||||
"mock_response": "litellm.InternalServerError",
|
||||
},
|
||||
}
|
||||
],
|
||||
num_retries=num_retries,
|
||||
)
|
||||
|
||||
|
||||
async def _fail_one_proxy_shaped_request(router, request_marker):
|
||||
"""The proxy hands the router a metadata dict and a proxy_server_request whose body is a
|
||||
shallow copy of the request, so body["metadata"] is the very same dict the router later
|
||||
stamps previous_models onto."""
|
||||
metadata = {"request_marker": request_marker}
|
||||
with pytest.raises(litellm.InternalServerError):
|
||||
await router.acompletion(
|
||||
model="broken-group",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
metadata=metadata,
|
||||
proxy_server_request={
|
||||
"url": "http://localhost:4000/v1/chat/completions",
|
||||
"method": "POST",
|
||||
"headers": {},
|
||||
"body": {"model": "broken-group", "metadata": metadata},
|
||||
},
|
||||
)
|
||||
return metadata["previous_models"]
|
||||
|
||||
|
||||
def _nested_breadcrumb_lists(node):
|
||||
if isinstance(node, dict):
|
||||
return [v for k, v in node.items() if k == "previous_models"] + [
|
||||
found for v in node.values() for found in _nested_breadcrumb_lists(v)
|
||||
]
|
||||
if isinstance(node, (list, tuple)):
|
||||
return [found for item in node for found in _nested_breadcrumb_lists(item)]
|
||||
return []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_retry_breadcrumbs_stay_per_request_and_flat_across_failing_requests():
|
||||
"""Every failed attempt appends a breadcrumb to metadata["previous_models"], and the proxy's
|
||||
request snapshot aliases that same metadata dict. Kept on the Router and copied wholesale,
|
||||
each breadcrumb embedded every earlier one from every earlier request, so the breadcrumb
|
||||
tree, and with it the debug repr of the kwargs, roughly doubled on each failed attempt until
|
||||
a single-worker proxy spent minutes in the redaction regex and stopped answering."""
|
||||
router = _always_failing_router(num_retries=2)
|
||||
|
||||
breadcrumbs_per_request = [
|
||||
await _fail_one_proxy_shaped_request(router, f"request-{request_number}") for request_number in range(1, 7)
|
||||
]
|
||||
|
||||
for request_number, breadcrumbs in enumerate(breadcrumbs_per_request, start=1):
|
||||
assert len(breadcrumbs) == 3, "one initial attempt plus two retries failed, each leaving one breadcrumb"
|
||||
assert {breadcrumb["metadata"]["request_marker"] for breadcrumb in breadcrumbs} == {f"request-{request_number}"}
|
||||
for breadcrumb in breadcrumbs:
|
||||
assert _nested_breadcrumb_lists(breadcrumb) == []
|
||||
assert len({len(repr(breadcrumbs)) for breadcrumbs in breadcrumbs_per_request}) == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_retry_breadcrumbs_keep_only_the_last_four_attempts():
|
||||
router = _always_failing_router(num_retries=6)
|
||||
|
||||
breadcrumbs = await _fail_one_proxy_shaped_request(router, "request-1")
|
||||
|
||||
assert len(breadcrumbs) == 4
|
||||
assert [breadcrumb["metadata"]["attempted_retries"] for breadcrumb in breadcrumbs] == [3, 4, 5, 6]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fallback_traceback_stays_available_at_debug_level():
|
||||
"""Dropping the stack from the ERROR line is only safe because the fallback path still
|
||||
|
|
|
|||
|
|
@ -40,6 +40,7 @@ from litellm.utils import (
|
|||
_is_streaming_request,
|
||||
_snapshot_exception_for_hook,
|
||||
async_post_call_failure_deployment_hook,
|
||||
async_post_call_success_deployment_hook,
|
||||
client,
|
||||
get_llm_provider,
|
||||
get_non_default_completion_params,
|
||||
|
|
@ -5808,6 +5809,90 @@ class TestHuggingFaceConfigFetch:
|
|||
assert request_timeout["read"] == HF_CONFIG_FETCH_TIMEOUT_SECONDS
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_success_deployment_hook_chains_past_callback_returning_response(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
"""Regression (LIT-5863): the dispatcher must run every callback, chaining each non-None
|
||||
result into the next call, instead of returning at the first callback answering non-None.
|
||||
A guardrail answering with the unmodified response used to starve every callback after it."""
|
||||
from litellm.types.utils import ModelResponse
|
||||
|
||||
original = ModelResponse()
|
||||
replacement = ModelResponse()
|
||||
|
||||
class PassthroughLogger(CustomLogger):
|
||||
async def async_post_call_success_deployment_hook(self, request_data, response, call_type):
|
||||
return response
|
||||
|
||||
class ReplacingLogger(CustomLogger):
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
self.seen: list = []
|
||||
|
||||
async def async_post_call_success_deployment_hook(self, request_data, response, call_type):
|
||||
self.seen.append(response)
|
||||
return replacement
|
||||
|
||||
class ObservingLogger(CustomLogger):
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
self.seen: list = []
|
||||
|
||||
async def async_post_call_success_deployment_hook(self, request_data, response, call_type):
|
||||
self.seen.append(response)
|
||||
return None
|
||||
|
||||
replacer = ReplacingLogger()
|
||||
observer = ObservingLogger()
|
||||
monkeypatch.setattr(litellm, "callbacks", [PassthroughLogger(), replacer, observer])
|
||||
|
||||
result = await async_post_call_success_deployment_hook(
|
||||
request_data={}, response=original, call_type=CallTypes.acompletion
|
||||
)
|
||||
|
||||
assert replacer.seen == [original]
|
||||
assert observer.seen == [replacement]
|
||||
assert result is replacement
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_registered_guardrail_does_not_starve_vector_store_search_results(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
"""Regression (LIT-5863): with any guardrail registered ahead of the lazily-appended
|
||||
VectorStorePreCallHook, /v1/chat/completions responses lost
|
||||
provider_specific_fields["search_results"] because the guardrail answered the unmodified
|
||||
response and the dispatcher stopped there."""
|
||||
from types import SimpleNamespace
|
||||
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.integrations.vector_store_integrations.vector_store_pre_call_hook import (
|
||||
VectorStorePreCallHook,
|
||||
)
|
||||
from litellm.types.utils import ModelResponse
|
||||
|
||||
search_results: Final = [{"search_query": "coolant", "data": [{"content": [{"text": "Cryoline-9", "type": "text"}]}]}]
|
||||
logging_obj = SimpleNamespace(model_call_details={"search_results": search_results})
|
||||
response = ModelResponse(choices=[{"message": {"role": "assistant", "content": "Cryoline-9"}}])
|
||||
|
||||
monkeypatch.setattr(
|
||||
litellm,
|
||||
"callbacks",
|
||||
[CustomGuardrail(guardrail_name="dummy-guardrail"), VectorStorePreCallHook()],
|
||||
)
|
||||
|
||||
result = await async_post_call_success_deployment_hook(
|
||||
request_data={"litellm_logging_obj": logging_obj},
|
||||
response=response,
|
||||
call_type=CallTypes.acompletion,
|
||||
)
|
||||
|
||||
provider_fields = result.choices[0].message.provider_specific_fields
|
||||
assert provider_fields is not None
|
||||
assert provider_fields["search_results"] == search_results
|
||||
|
||||
|
||||
class TestIsVisionExplicitlyDisabled:
|
||||
"""github_copilot and chatgpt run an OAuth device flow inside get_llm_provider; the
|
||||
explicit-disable lookup must adopt the declared prefix instead of resolving it, exactly
|
||||
|
|
@ -5836,3 +5921,135 @@ class TestIsVisionExplicitlyDisabled:
|
|||
is_vision_explicitly_disabled("fireworks_ai/accounts/fireworks/models/deepseek-v4-flash-0731") is True
|
||||
)
|
||||
assert is_vision_explicitly_disabled("anthropic/claude-sonnet-4-5") is False
|
||||
|
||||
|
||||
class TestVerboseRequestLineRedaction:
|
||||
"""`litellm.set_verbose = True` echoes the caller's kwargs back as a `litellm.completion(...)`
|
||||
line on stdout, so a credential kwarg lands in whatever collects stdout: a terminal, a
|
||||
container log drain, a CI job log. Credential-named kwargs must not survive that echo,
|
||||
at any nesting depth, while ordinary params still must, or the line stops telling the
|
||||
developer what they called."""
|
||||
|
||||
FAKE_API_KEY: Final = "sk-fake-lit6823-0000000000000000"
|
||||
|
||||
def _verbose_request_line(self, capsys, monkeypatch, **kwargs) -> str:
|
||||
monkeypatch.setattr(litellm, "set_verbose", True)
|
||||
monkeypatch.setattr("litellm._logging.set_verbose", True)
|
||||
capsys.readouterr()
|
||||
litellm.completion(
|
||||
model="gpt-3.5-turbo",
|
||||
messages=[{"role": "user", "content": "hello"}],
|
||||
mock_response="hi",
|
||||
**kwargs,
|
||||
)
|
||||
captured: Final = capsys.readouterr()
|
||||
return "\n".join(line for line in (captured.out + captured.err).splitlines() if "litellm.completion(" in line)
|
||||
|
||||
def test_api_key_never_reaches_the_request_line(self, capsys, monkeypatch):
|
||||
printed: Final = self._verbose_request_line(capsys, monkeypatch, api_key=self.FAKE_API_KEY)
|
||||
|
||||
assert "litellm.completion(" in printed
|
||||
assert self.FAKE_API_KEY not in printed
|
||||
assert "api_key='REDACTED'" in printed
|
||||
|
||||
def test_credential_headers_never_reach_the_request_line(self, capsys, monkeypatch):
|
||||
printed: Final = self._verbose_request_line(
|
||||
capsys,
|
||||
monkeypatch,
|
||||
api_key=self.FAKE_API_KEY,
|
||||
extra_headers={"Authorization": "Bearer fake-lit6823-header", "x-request-id": "abc123"},
|
||||
)
|
||||
|
||||
assert "fake-lit6823-header" not in printed
|
||||
assert "'Authorization': 'REDACTED'" in printed
|
||||
assert "'x-request-id': 'abc123'" in printed
|
||||
|
||||
def test_credentials_nested_in_a_list_never_reach_the_request_line(self, capsys, monkeypatch):
|
||||
printed: Final = self._verbose_request_line(
|
||||
capsys,
|
||||
monkeypatch,
|
||||
api_key=self.FAKE_API_KEY,
|
||||
extra_body={"providers": [{"name": "openai", "api_key": "sk-fake-lit6823-nested"}]},
|
||||
)
|
||||
|
||||
assert "sk-fake-lit6823-nested" not in printed
|
||||
assert "'name': 'openai'" in printed
|
||||
|
||||
def test_ordinary_params_still_printed(self, capsys, monkeypatch):
|
||||
printed: Final = self._verbose_request_line(
|
||||
capsys, monkeypatch, api_key=self.FAKE_API_KEY, max_tokens=17, temperature=0.25
|
||||
)
|
||||
|
||||
assert "model='gpt-3.5-turbo'" in printed
|
||||
assert "max_tokens=17" in printed
|
||||
assert "temperature=0.25" in printed
|
||||
|
||||
|
||||
class TestFinalOptionalParamsLineRedaction:
|
||||
"""A verbose run echoes the fully built optional params too, and `extra_body` carries whatever the
|
||||
caller nested inside it straight onto that line, so a credential tucked in there lands in a terminal
|
||||
or a log drain in plaintext. It has to be redacted on both surfaces `print_verbose` writes to, and the
|
||||
line has to keep printing on both, because `litellm.set_verbose` and the DEBUG logger are independent
|
||||
switches and neither implies the other."""
|
||||
|
||||
FAKE_NESTED_KEY: Final = "sk-fake-lit6835-nested-0000000000"
|
||||
|
||||
def _complete(self, **kwargs) -> None:
|
||||
litellm.completion(
|
||||
model="gpt-3.5-turbo",
|
||||
messages=[{"role": "user", "content": "hello"}],
|
||||
mock_response="hi",
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
def _printed_line(self, capsys) -> str:
|
||||
captured: Final = capsys.readouterr()
|
||||
return "\n".join(
|
||||
line for line in (captured.out + captured.err).splitlines() if "Final returned optional params" in line
|
||||
)
|
||||
|
||||
def test_nested_credential_is_redacted_when_only_set_verbose_is_on(self, capsys, caplog, monkeypatch):
|
||||
monkeypatch.setattr(litellm, "set_verbose", True)
|
||||
with caplog.at_level(logging.WARNING, logger=verbose_logger.name):
|
||||
capsys.readouterr()
|
||||
self._complete(extra_body={"providers": [{"name": "openai", "api_key": self.FAKE_NESTED_KEY}]})
|
||||
printed: Final = self._printed_line(capsys)
|
||||
|
||||
assert printed
|
||||
assert self.FAKE_NESTED_KEY not in printed
|
||||
assert "'api_key': 'REDACTED'" in printed
|
||||
assert "'name': 'openai'" in printed
|
||||
|
||||
def test_line_still_reaches_the_logger_when_only_the_debug_logger_is_on(self, capsys, caplog, monkeypatch):
|
||||
monkeypatch.setattr(litellm, "set_verbose", False)
|
||||
with caplog.at_level(logging.DEBUG, logger=verbose_logger.name):
|
||||
self._complete(extra_body={"providers": [{"name": "openai", "api_key": self.FAKE_NESTED_KEY}]})
|
||||
logged: Final = "\n".join(
|
||||
record.getMessage()
|
||||
for record in caplog.records
|
||||
if "Final returned optional params" in record.getMessage()
|
||||
)
|
||||
|
||||
assert logged
|
||||
assert self.FAKE_NESTED_KEY not in logged
|
||||
assert "'name': 'openai'" in logged
|
||||
|
||||
def test_nothing_is_emitted_when_neither_verbose_switch_is_on(self, capsys, caplog, monkeypatch):
|
||||
monkeypatch.setattr(litellm, "set_verbose", False)
|
||||
with caplog.at_level(logging.WARNING, logger=verbose_logger.name):
|
||||
capsys.readouterr()
|
||||
self._complete(extra_body={"providers": [{"name": "openai", "api_key": self.FAKE_NESTED_KEY}]})
|
||||
captured: Final = capsys.readouterr()
|
||||
|
||||
assert "Final returned optional params" not in captured.out + captured.err
|
||||
assert self.FAKE_NESTED_KEY not in captured.out + captured.err
|
||||
|
||||
def test_ordinary_optional_params_still_reach_the_line(self, capsys, caplog, monkeypatch):
|
||||
monkeypatch.setattr(litellm, "set_verbose", True)
|
||||
with caplog.at_level(logging.WARNING, logger=verbose_logger.name):
|
||||
capsys.readouterr()
|
||||
self._complete(max_tokens=17, temperature=0.25)
|
||||
printed: Final = self._printed_line(capsys)
|
||||
|
||||
assert "'max_tokens': 17" in printed
|
||||
assert "'temperature': 0.25" in printed
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@
|
|||
"limit": 22328
|
||||
},
|
||||
"LIT002": {
|
||||
"limit": 26763
|
||||
"limit": 26760
|
||||
},
|
||||
"LIT003": {
|
||||
"limit": 261
|
||||
|
|
@ -30,7 +30,7 @@
|
|||
"limit": 16470
|
||||
},
|
||||
"LIT011": {
|
||||
"limit": 5517
|
||||
"limit": 5516
|
||||
},
|
||||
"LIT012": {
|
||||
"limit": 4489
|
||||
|
|
|
|||
30
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
30
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -44070,7 +44070,11 @@ export interface operations {
|
|||
};
|
||||
list_containers_containers_get: {
|
||||
parameters: {
|
||||
query?: never;
|
||||
query?: {
|
||||
after?: string | null;
|
||||
limit?: number | null;
|
||||
order?: string | null;
|
||||
};
|
||||
header?: never;
|
||||
path?: never;
|
||||
cookie?: never;
|
||||
|
|
@ -44086,6 +44090,15 @@ export interface operations {
|
|||
"application/json": unknown;
|
||||
};
|
||||
};
|
||||
/** @description Validation Error */
|
||||
422: {
|
||||
headers: {
|
||||
[name: string]: unknown;
|
||||
};
|
||||
content: {
|
||||
"application/json": components["schemas"]["HTTPValidationError"];
|
||||
};
|
||||
};
|
||||
};
|
||||
};
|
||||
create_container_containers_post: {
|
||||
|
|
@ -61114,7 +61127,11 @@ export interface operations {
|
|||
};
|
||||
list_containers_v1_containers_get: {
|
||||
parameters: {
|
||||
query?: never;
|
||||
query?: {
|
||||
after?: string | null;
|
||||
limit?: number | null;
|
||||
order?: string | null;
|
||||
};
|
||||
header?: never;
|
||||
path?: never;
|
||||
cookie?: never;
|
||||
|
|
@ -61130,6 +61147,15 @@ export interface operations {
|
|||
"application/json": unknown;
|
||||
};
|
||||
};
|
||||
/** @description Validation Error */
|
||||
422: {
|
||||
headers: {
|
||||
[name: string]: unknown;
|
||||
};
|
||||
content: {
|
||||
"application/json": components["schemas"]["HTTPValidationError"];
|
||||
};
|
||||
};
|
||||
};
|
||||
};
|
||||
create_container_v1_containers_post: {
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue