mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_internal_copy_38013
# Conflicts: # tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py
This commit is contained in:
commit
055eaee1f9
152 changed files with 5930 additions and 1654 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
|
||||
|
|
|
|||
|
|
@ -18,9 +18,11 @@ caller's identity metadata, minus two things that must never be forwarded as-is:
|
|||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Mapping
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
|
||||
from litellm.constants import INTERNAL_CALL_ORIGIN_METADATA_KEY, NON_INFERENCE_CALL_TYPES
|
||||
from litellm.litellm_core_utils.initialize_dynamic_callback_params import initialize_standard_callback_dynamic_params
|
||||
from litellm.types.utils import BACKGROUND_RESPONSE_COST_POLL_CALL_ORIGIN, InternalCallOrigin
|
||||
|
||||
BUDGET_RESERVATION_METADATA_KEYS: Final = frozenset({"user_api_key_budget_reservation"})
|
||||
|
|
@ -142,6 +144,19 @@ def forwarded_internal_call_metadata(
|
|||
}
|
||||
|
||||
|
||||
def parent_session_kwargs(request_kwargs: Mapping[str, object] | None) -> Mapping[str, str]:
|
||||
kwargs: Final = request_kwargs or MappingProxyType({})
|
||||
return MappingProxyType(
|
||||
{k: v for k in ("litellm_session_id", "litellm_trace_id") if isinstance(v := kwargs.get(k), str)}
|
||||
)
|
||||
|
||||
|
||||
def effective_turn_off_message_logging(request_kwargs: Mapping[str, object] | None) -> bool | None:
|
||||
return initialize_standard_callback_dynamic_params(dict(request_kwargs) if request_kwargs else None).get(
|
||||
"turn_off_message_logging"
|
||||
)
|
||||
|
||||
|
||||
def sanitized_forwardable_call_metadata(
|
||||
parent_metadata: Mapping[str, object],
|
||||
call_origin: InternalCallOrigin,
|
||||
|
|
|
|||
|
|
@ -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,6 +7,7 @@ from litellm.exceptions import UnsupportedParamsError
|
|||
from litellm.llms.openai.chat.gpt_5_transformation import (
|
||||
OpenAIGPT5Config,
|
||||
_get_effort_level,
|
||||
is_gpt_reasoning_series_name,
|
||||
)
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
|
||||
|
|
@ -35,26 +36,7 @@ class AzureOpenAIGPT5Config(AzureOpenAIConfig, OpenAIGPT5Config):
|
|||
|
||||
@classmethod
|
||||
def is_model_gpt_5_model(cls, model: str) -> bool:
|
||||
"""Check if the Azure model string refers to a gpt-5 variant.
|
||||
|
||||
Accepts both explicit gpt-5 model names and the ``gpt5_series/`` prefix
|
||||
used for manual routing.
|
||||
"""
|
||||
# The gpt-5-chat* family (gpt-5-chat, gpt-5-chat-latest, gpt-5-chat-2025-08-07,
|
||||
# …) are regular chat models: they support temperature and tool_choice but NOT
|
||||
# reasoning_effort. They must NOT be routed through the GPT-5 reasoning path.
|
||||
#
|
||||
# Versioned chat models such as gpt-5.3-chat and gpt-5.1-chat ARE reasoning
|
||||
# models and must stay on the GPT-5 path. The distinguishing feature is that
|
||||
# the gpt-5-chat family has a literal "-chat" immediately after "gpt-5"
|
||||
# (i.e. "gpt-5-chat…"), while versioned chat models interpose a minor version
|
||||
# number (i.e. "gpt-5.<digit>-chat").
|
||||
#
|
||||
# Using a startswith("gpt-5-chat") prefix check on the normalized name (rather
|
||||
# than a substring check) makes this boundary explicit and avoids any ambiguity
|
||||
# if future model names coincidentally contain "gpt-5-chat" as an interior run.
|
||||
_normalized: Final = model.split("/")[-1] # strip provider prefix, e.g. "azure/"
|
||||
return ("gpt-5" in model and not _normalized.startswith("gpt-5-chat")) or "gpt5_series" in model
|
||||
return is_gpt_reasoning_series_name(model) or "gpt5_series" in model
|
||||
|
||||
def get_supported_openai_params(self, model: str) -> list[str]:
|
||||
"""Get supported parameters for Azure OpenAI GPT-5 models.
|
||||
|
|
|
|||
|
|
@ -14,6 +14,7 @@ from litellm.litellm_core_utils.prompt_templates.factory import (
|
|||
convert_to_azure_openai_messages,
|
||||
)
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
from litellm.llms.openai.chat.gpt_5_transformation import GPT_REASONING_SERIES_MARKERS
|
||||
from litellm.types.llms.azure import (
|
||||
API_VERSION_MONTH_SUPPORTED_RESPONSE_FORMAT,
|
||||
API_VERSION_YEAR_SUPPORTED_RESPONSE_FORMAT,
|
||||
|
|
@ -139,7 +140,7 @@ class AzureOpenAIConfig(BaseConfig):
|
|||
name family needs the rename, including the ``gpt-5-chat*`` models that are excluded from
|
||||
the reasoning path by https://github.com/BerriAI/litellm/issues/13781.
|
||||
"""
|
||||
return "gpt-5" in model or "gpt5_series" in model
|
||||
return any(marker in model for marker in GPT_REASONING_SERIES_MARKERS) or "gpt5_series" in model
|
||||
|
||||
def _is_response_format_supported_model(self, model: str) -> bool:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -91,6 +91,7 @@ from litellm.llms.base_llm.vector_store_files.transformation import (
|
|||
BaseVectorStoreFilesConfig,
|
||||
)
|
||||
from litellm.llms.base_llm.videos.transformation import BaseVideoConfig
|
||||
from litellm.llms.custom_httpx.container_handler import raise_for_error_status
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
AsyncHTTPHandler,
|
||||
HTTPHandler,
|
||||
|
|
@ -8926,17 +8927,19 @@ class BaseLLMHTTPHandler:
|
|||
json=data,
|
||||
timeout=timeout,
|
||||
)
|
||||
|
||||
return container_provider_config.transform_container_create_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
raise self._handle_error(
|
||||
e=e,
|
||||
provider_config=container_provider_config,
|
||||
)
|
||||
raise_for_error_status(
|
||||
response=response,
|
||||
container_provider_config=container_provider_config,
|
||||
)
|
||||
return container_provider_config.transform_container_create_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
async def async_container_create_handler(
|
||||
self,
|
||||
|
|
@ -9002,17 +9005,19 @@ class BaseLLMHTTPHandler:
|
|||
json=data,
|
||||
timeout=timeout,
|
||||
)
|
||||
|
||||
return container_provider_config.transform_container_create_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
raise self._handle_error(
|
||||
e=e,
|
||||
provider_config=container_provider_config,
|
||||
)
|
||||
raise_for_error_status(
|
||||
response=response,
|
||||
container_provider_config=container_provider_config,
|
||||
)
|
||||
return container_provider_config.transform_container_create_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
def container_list_handler(
|
||||
self,
|
||||
|
|
@ -9092,17 +9097,19 @@ class BaseLLMHTTPHandler:
|
|||
headers=headers,
|
||||
params=params or None,
|
||||
)
|
||||
|
||||
return container_provider_config.transform_container_list_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
raise self._handle_error(
|
||||
e=e,
|
||||
provider_config=container_provider_config,
|
||||
)
|
||||
raise_for_error_status(
|
||||
response=response,
|
||||
container_provider_config=container_provider_config,
|
||||
)
|
||||
return container_provider_config.transform_container_list_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
async def async_container_list_handler(
|
||||
self,
|
||||
|
|
@ -9169,17 +9176,19 @@ class BaseLLMHTTPHandler:
|
|||
headers=headers,
|
||||
params=params or None,
|
||||
)
|
||||
|
||||
return container_provider_config.transform_container_list_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
raise self._handle_error(
|
||||
e=e,
|
||||
provider_config=container_provider_config,
|
||||
)
|
||||
raise_for_error_status(
|
||||
response=response,
|
||||
container_provider_config=container_provider_config,
|
||||
)
|
||||
return container_provider_config.transform_container_list_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
def container_retrieve_handler(
|
||||
self,
|
||||
|
|
@ -9257,17 +9266,19 @@ class BaseLLMHTTPHandler:
|
|||
headers=headers,
|
||||
params=params or None,
|
||||
)
|
||||
|
||||
return container_provider_config.transform_container_retrieve_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
raise self._handle_error(
|
||||
e=e,
|
||||
provider_config=container_provider_config,
|
||||
)
|
||||
raise_for_error_status(
|
||||
response=response,
|
||||
container_provider_config=container_provider_config,
|
||||
)
|
||||
return container_provider_config.transform_container_retrieve_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
async def async_container_retrieve_handler(
|
||||
self,
|
||||
|
|
@ -9334,17 +9345,19 @@ class BaseLLMHTTPHandler:
|
|||
headers=headers,
|
||||
params=params or None,
|
||||
)
|
||||
|
||||
return container_provider_config.transform_container_retrieve_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
raise self._handle_error(
|
||||
e=e,
|
||||
provider_config=container_provider_config,
|
||||
)
|
||||
raise_for_error_status(
|
||||
response=response,
|
||||
container_provider_config=container_provider_config,
|
||||
)
|
||||
return container_provider_config.transform_container_retrieve_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
def container_delete_handler(
|
||||
self,
|
||||
|
|
@ -9422,17 +9435,19 @@ class BaseLLMHTTPHandler:
|
|||
headers=headers,
|
||||
params=params or None,
|
||||
)
|
||||
|
||||
return container_provider_config.transform_container_delete_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
raise self._handle_error(
|
||||
e=e,
|
||||
provider_config=container_provider_config,
|
||||
)
|
||||
raise_for_error_status(
|
||||
response=response,
|
||||
container_provider_config=container_provider_config,
|
||||
)
|
||||
return container_provider_config.transform_container_delete_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
async def async_container_delete_handler(
|
||||
self,
|
||||
|
|
@ -9499,17 +9514,19 @@ class BaseLLMHTTPHandler:
|
|||
headers=headers,
|
||||
params=params or None,
|
||||
)
|
||||
|
||||
return container_provider_config.transform_container_delete_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
raise self._handle_error(
|
||||
e=e,
|
||||
provider_config=container_provider_config,
|
||||
)
|
||||
raise_for_error_status(
|
||||
response=response,
|
||||
container_provider_config=container_provider_config,
|
||||
)
|
||||
return container_provider_config.transform_container_delete_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
def container_file_list_handler(
|
||||
self,
|
||||
|
|
@ -9591,17 +9608,19 @@ class BaseLLMHTTPHandler:
|
|||
headers=headers,
|
||||
params=params or None,
|
||||
)
|
||||
|
||||
return container_provider_config.transform_container_file_list_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
raise self._handle_error(
|
||||
e=e,
|
||||
provider_config=container_provider_config,
|
||||
)
|
||||
raise_for_error_status(
|
||||
response=response,
|
||||
container_provider_config=container_provider_config,
|
||||
)
|
||||
return container_provider_config.transform_container_file_list_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
async def async_container_file_list_handler(
|
||||
self,
|
||||
|
|
@ -9670,17 +9689,19 @@ class BaseLLMHTTPHandler:
|
|||
headers=headers,
|
||||
params=params or None,
|
||||
)
|
||||
|
||||
return container_provider_config.transform_container_file_list_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
raise self._handle_error(
|
||||
e=e,
|
||||
provider_config=container_provider_config,
|
||||
)
|
||||
raise_for_error_status(
|
||||
response=response,
|
||||
container_provider_config=container_provider_config,
|
||||
)
|
||||
return container_provider_config.transform_container_file_list_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
def container_file_content_handler(
|
||||
self,
|
||||
|
|
@ -9756,17 +9777,19 @@ class BaseLLMHTTPHandler:
|
|||
headers=headers,
|
||||
params=params or None,
|
||||
)
|
||||
|
||||
return container_provider_config.transform_container_file_content_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
raise self._handle_error(
|
||||
e=e,
|
||||
provider_config=container_provider_config,
|
||||
)
|
||||
raise_for_error_status(
|
||||
response=response,
|
||||
container_provider_config=container_provider_config,
|
||||
)
|
||||
return container_provider_config.transform_container_file_content_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
async def async_container_file_content_handler(
|
||||
self,
|
||||
|
|
@ -9832,17 +9855,19 @@ class BaseLLMHTTPHandler:
|
|||
headers=headers,
|
||||
params=params or None,
|
||||
)
|
||||
|
||||
return container_provider_config.transform_container_file_content_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
raise self._handle_error(
|
||||
e=e,
|
||||
provider_config=container_provider_config,
|
||||
)
|
||||
raise_for_error_status(
|
||||
response=response,
|
||||
container_provider_config=container_provider_config,
|
||||
)
|
||||
return container_provider_config.transform_container_file_content_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
###### VECTOR STORE HANDLER ######
|
||||
@staticmethod
|
||||
|
|
|
|||
|
|
@ -61,6 +61,14 @@ def _get_effort_level(value: str | dict | None) -> str | None:
|
|||
return None
|
||||
|
||||
|
||||
GPT_REASONING_SERIES_MARKERS: Final = ("gpt-5", "gpt-6")
|
||||
|
||||
|
||||
def is_gpt_reasoning_series_name(model: str) -> bool:
|
||||
normalized: Final = model.split("/")[-1]
|
||||
return any(marker in model for marker in GPT_REASONING_SERIES_MARKERS) and not normalized.startswith("gpt-5-chat")
|
||||
|
||||
|
||||
class OpenAIGPT5Config(OpenAIGPTConfig):
|
||||
"""Configuration for gpt-5 models including GPT-5-Codex variants.
|
||||
|
||||
|
|
@ -73,21 +81,7 @@ class OpenAIGPT5Config(OpenAIGPTConfig):
|
|||
|
||||
@classmethod
|
||||
def is_model_gpt_5_model(cls, model: str) -> bool:
|
||||
# The gpt-5-chat* family (gpt-5-chat, gpt-5-chat-latest, gpt-5-chat-2025-08-07,
|
||||
# …) are regular chat models: they support temperature and tool_choice but NOT
|
||||
# reasoning_effort. They must NOT be routed through the GPT-5 reasoning path.
|
||||
#
|
||||
# Versioned chat models such as gpt-5.3-chat and gpt-5.1-chat ARE reasoning
|
||||
# models and must stay on the GPT-5 path. The distinguishing feature is that
|
||||
# the gpt-5-chat family has a literal "-chat" immediately after "gpt-5"
|
||||
# (i.e. "gpt-5-chat…"), while versioned chat models interpose a minor version
|
||||
# number (i.e. "gpt-5.<digit>-chat").
|
||||
#
|
||||
# Using a startswith("gpt-5-chat") prefix check on the normalized name (rather
|
||||
# than a substring check) makes this boundary explicit and avoids any ambiguity
|
||||
# if future model names coincidentally contain "gpt-5-chat" as an interior run.
|
||||
_normalized: Final = model.split("/")[-1] # strip provider prefix, e.g. "openai/"
|
||||
return "gpt-5" in model and not _normalized.startswith("gpt-5-chat")
|
||||
return is_gpt_reasoning_series_name(model)
|
||||
|
||||
@classmethod
|
||||
def is_model_gpt_5_search_model(cls, model: str) -> bool:
|
||||
|
|
@ -122,6 +116,8 @@ class OpenAIGPT5Config(OpenAIGPTConfig):
|
|||
def is_model_gpt_5_4_plus_model(cls, model: str) -> bool:
|
||||
"""Check if the model is gpt-5.4 or newer (5.4, 5.5, 5.6, etc., including pro)."""
|
||||
model_name: Final = model.split("/")[-1]
|
||||
if model_name.startswith("gpt-6"):
|
||||
return True
|
||||
if not model_name.startswith("gpt-5."):
|
||||
return False
|
||||
try:
|
||||
|
|
|
|||
|
|
@ -16,6 +16,7 @@ from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response impo
|
|||
)
|
||||
from litellm.litellm_core_utils.url_utils import encode_url_path_segment
|
||||
from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig
|
||||
from litellm.llms.openai.chat.gpt_5_transformation import is_gpt_reasoning_series_name
|
||||
from litellm.responses.litellm_completion_transformation.custom_tools import TOOL_CALL_ITEM_ID_PREFIX_BY_TYPE
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.llms.openai import *
|
||||
|
|
@ -89,7 +90,7 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
|
|||
parts: Final = model.split("/")
|
||||
if len(parts) > 1 and parts[0] not in ("openai",):
|
||||
return False
|
||||
return "gpt-5" in model and "gpt-5-chat" not in model
|
||||
return is_gpt_reasoning_series_name(model)
|
||||
|
||||
@staticmethod
|
||||
def _supports_reasoning_effort_none(model: str) -> bool:
|
||||
|
|
|
|||
|
|
@ -998,6 +998,16 @@ def replace_project_and_location_in_route(requested_route: str, vertex_project:
|
|||
return modified_route
|
||||
|
||||
|
||||
def _api_version_for_route(requested_route: str) -> Literal["v1", "v1beta1"]:
|
||||
return "v1beta1" if "cachedContent" in requested_route else "v1"
|
||||
|
||||
|
||||
def _with_api_version(requested_route: str) -> str:
|
||||
if not requested_route.startswith("/projects/"):
|
||||
return requested_route
|
||||
return f"/{_api_version_for_route(requested_route)}{requested_route}"
|
||||
|
||||
|
||||
def construct_target_url(
|
||||
base_url: str,
|
||||
requested_route: str,
|
||||
|
|
@ -1017,18 +1027,19 @@ def construct_target_url(
|
|||
|
||||
new_base_url: Final = httpx.URL(base_url)
|
||||
if "locations" in requested_route: # contains the target project id + location
|
||||
if vertex_project and vertex_location:
|
||||
requested_route = replace_project_and_location_in_route(requested_route, vertex_project, vertex_location)
|
||||
return new_base_url.copy_with(path=requested_route)
|
||||
targeted_route: Final = (
|
||||
replace_project_and_location_in_route(requested_route, vertex_project, vertex_location)
|
||||
if vertex_project and vertex_location
|
||||
else requested_route
|
||||
)
|
||||
return new_base_url.copy_with(path=_with_api_version(targeted_route))
|
||||
|
||||
"""
|
||||
- Add endpoint version (e.g. v1beta for cachedContent, v1 for rest)
|
||||
- Add default project id
|
||||
- Add default location
|
||||
"""
|
||||
vertex_version: Literal["v1", "v1beta1"] = "v1"
|
||||
if "cachedContent" in requested_route:
|
||||
vertex_version = "v1beta1"
|
||||
vertex_version: Literal["v1", "v1beta1"] = _api_version_for_route(requested_route)
|
||||
|
||||
# Check if the requested route starts with a version
|
||||
# e.g. /v1beta1/publishers/google/models/gemini-3-pro-preview:streamGenerateContent
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -654,23 +654,16 @@ class InMemoryGuardrailHandler:
|
|||
source: Literal["db", "config"] = "db",
|
||||
) -> None:
|
||||
"""
|
||||
Update a guardrail in memory
|
||||
|
||||
- updates the guardrail in memory
|
||||
- updates the guardrail params in litellm.callback_manager
|
||||
Update a guardrail in memory: a changed name or litellm_params rebuilds the
|
||||
live callback from the new row (fail-closed: an invalid row keeps the
|
||||
previous instance and raises), anything else only refreshes the stored row
|
||||
"""
|
||||
self.IN_MEMORY_GUARDRAILS[guardrail_id] = guardrail
|
||||
self._sources[guardrail_id] = source
|
||||
|
||||
tracked_callbacks: Final = self._tracked_callbacks(guardrail_id)
|
||||
if not tracked_callbacks:
|
||||
updated_guardrail: Final = cast(Guardrail, {**guardrail, "guardrail_id": guardrail_id})
|
||||
if self._has_guardrail_params_changed(guardrail_id, updated_guardrail):
|
||||
self.reinitialize_guardrail(guardrail=updated_guardrail, source=source)
|
||||
return
|
||||
updated_litellm_params: Final = cast(LitellmParams, guardrail.get("litellm_params", {}))
|
||||
tracked_callbacks[0].update_in_memory_litellm_params(litellm_params=updated_litellm_params)
|
||||
for sibling_callback in tracked_callbacks[1:]:
|
||||
sibling_stage = sibling_callback.event_hook
|
||||
sibling_callback.update_in_memory_litellm_params(litellm_params=updated_litellm_params)
|
||||
sibling_callback.event_hook = sibling_stage
|
||||
self.IN_MEMORY_GUARDRAILS[guardrail_id] = updated_guardrail
|
||||
self._sources[guardrail_id] = source
|
||||
|
||||
def delete_in_memory_guardrail(self, guardrail_id: str) -> None:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -17,6 +17,7 @@ import json
|
|||
import traceback
|
||||
from collections.abc import Awaitable, Mapping, Sequence
|
||||
from datetime import datetime, timezone
|
||||
from types import MappingProxyType
|
||||
from typing import Any, Final, Literal, Protocol, cast, overload
|
||||
|
||||
import fastapi
|
||||
|
|
@ -77,6 +78,7 @@ from litellm.types.proxy.management_endpoints.internal_user_endpoints import (
|
|||
BulkUpdateUserRequest,
|
||||
BulkUpdateUserResponse,
|
||||
UserListResponse,
|
||||
UserSearchWhere,
|
||||
UserUpdateResult,
|
||||
)
|
||||
from litellm.types.proxy.management_endpoints.scim_v2 import (
|
||||
|
|
@ -2080,6 +2082,22 @@ async def _authorize_user_list_request(
|
|||
return ",".join(allowed_org_ids)
|
||||
|
||||
|
||||
_NO_SEARCH_WHERE: Final[Mapping[str, object]] = MappingProxyType({})
|
||||
|
||||
|
||||
def _user_search_where(search: str | None) -> Mapping[str, object]:
|
||||
"""Prisma predicate for `/user/list?search=`: user_id or user_email contains it, case-insensitive."""
|
||||
if not search:
|
||||
return _NO_SEARCH_WHERE
|
||||
search_where: Final[UserSearchWhere] = {
|
||||
"OR": (
|
||||
{"user_id": {"contains": search, "mode": "insensitive"}},
|
||||
{"user_email": {"contains": search, "mode": "insensitive"}},
|
||||
)
|
||||
}
|
||||
return search_where
|
||||
|
||||
|
||||
@router.get(
|
||||
"/user/list",
|
||||
tags=["Internal User management"],
|
||||
|
|
@ -2091,6 +2109,10 @@ async def get_users(
|
|||
user_ids: str | None = fastapi.Query(default=None, description="Get list of users by user_ids"),
|
||||
sso_user_ids: str | None = fastapi.Query(default=None, description="Get list of users by sso_user_id"),
|
||||
user_email: str | None = fastapi.Query(default=None, description="Filter users by partial email match"),
|
||||
search: str | None = fastapi.Query(
|
||||
default=None,
|
||||
description="Combined search: matches users whose 'user_id' or 'user_email' contains the value (case-insensitive).",
|
||||
),
|
||||
team: str | None = fastapi.Query(default=None, description="Filter users by team id"),
|
||||
page: int = fastapi.Query(default=1, ge=1, description="Page number"),
|
||||
page_size: int = fastapi.Query(default=25, ge=1, le=100, description="Number of items per page"),
|
||||
|
|
@ -2121,6 +2143,8 @@ async def get_users(
|
|||
Get list of users by sso_ids. Comma separated list of sso_ids.
|
||||
user_email: Optional[str]
|
||||
Filter users by partial email match
|
||||
search: Optional[str]
|
||||
Combined search: matches users whose user_id or user_email contains the value (case-insensitive)
|
||||
team: Optional[str]
|
||||
Filter users by team id. Will match if user has this team in their teams array.
|
||||
page: int
|
||||
|
|
@ -2197,7 +2221,11 @@ async def get_users(
|
|||
where_conditions["organization_memberships"] = {"some": {"organization_id": {"in": org_id_list}}}
|
||||
|
||||
## Filter any none fastapi.Query params - e.g. where_conditions: {'user_email': {'contains': Query(None), 'mode': 'insensitive'}, 'teams': {'has': Query(None)}}
|
||||
where_conditions = {k: v for k, v in where_conditions.items() if v is not None}
|
||||
where: Final[Mapping[str, object]] = {
|
||||
key: value
|
||||
for key, value in (*where_conditions.items(), *_user_search_where(search).items())
|
||||
if value is not None
|
||||
}
|
||||
|
||||
# Build order_by conditions
|
||||
|
||||
|
|
@ -2206,14 +2234,14 @@ async def get_users(
|
|||
)
|
||||
|
||||
users: Final[Sequence[prisma_models.LiteLLM_UserTable]] = await UserRepository(prisma_client).table.find_many(
|
||||
where=where_conditions,
|
||||
where=where,
|
||||
skip=skip,
|
||||
take=page_size,
|
||||
order=(order_by if order_by else {"created_at": "desc"}), # Default to created_at desc if no sort specified
|
||||
)
|
||||
|
||||
# Get total count of user rows
|
||||
total_count: Final[int] = await UserRepository(prisma_client).table.count(where=where_conditions)
|
||||
total_count: Final[int] = await UserRepository(prisma_client).table.count(where=where)
|
||||
|
||||
# Get key count for each user
|
||||
user_key_counts: Final = await get_user_key_counts(prisma_client, [user.user_id for user in users])
|
||||
|
|
|
|||
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -29,6 +29,12 @@ from litellm.proxy._types import (
|
|||
from litellm.proxy.auth.auth_utils import is_request_body_safe
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.common_utils.path_utils import safe_filename
|
||||
from litellm.proxy.prompts.prompt_registry import (
|
||||
DEFAULT_PROMPT_ENVIRONMENT,
|
||||
get_base_prompt_id,
|
||||
get_version_number,
|
||||
prompt_environment_or_default,
|
||||
)
|
||||
from litellm.repositories.table_repositories import PromptRepository
|
||||
from litellm.types.prompts.init_prompts import (
|
||||
ListPromptsResponse,
|
||||
|
|
@ -102,165 +108,20 @@ def _prompt_table(prisma_client: "PrismaClient") -> _PromptTableActions:
|
|||
return PromptRepository(prisma_client).table
|
||||
|
||||
|
||||
def get_base_prompt_id(prompt_id: str) -> str:
|
||||
"""
|
||||
Extract the base prompt ID by stripping the version suffix if present.
|
||||
|
||||
Args:
|
||||
prompt_id: Prompt ID that may include version suffix (e.g., "jack_success.v1" or "jack_success_v1")
|
||||
|
||||
Returns:
|
||||
Base prompt ID without version suffix (e.g., "jack_success")
|
||||
|
||||
Examples:
|
||||
>>> get_base_prompt_id("jack_success.v1")
|
||||
"jack_success"
|
||||
>>> get_base_prompt_id("jack_success_v1")
|
||||
"jack_success"
|
||||
>>> get_base_prompt_id("jack_success")
|
||||
"jack_success"
|
||||
"""
|
||||
# Try dot separator first (.v)
|
||||
if ".v" in prompt_id:
|
||||
return prompt_id.split(".v")[0]
|
||||
# Try underscore separator (_v)
|
||||
if "_v" in prompt_id:
|
||||
return prompt_id.split("_v")[0]
|
||||
return prompt_id
|
||||
|
||||
|
||||
def get_version_number(prompt_id: str) -> int:
|
||||
"""
|
||||
Extract the version number from a versioned prompt ID.
|
||||
|
||||
Args:
|
||||
prompt_id: Prompt ID that may include version suffix (e.g., "jack_success.v2" or "jack_success_v2")
|
||||
|
||||
Returns:
|
||||
Version number (defaults to 1 if no version suffix or invalid format)
|
||||
|
||||
Examples:
|
||||
>>> get_version_number("jack_success.v2")
|
||||
2
|
||||
>>> get_version_number("jack_success_v2")
|
||||
2
|
||||
>>> get_version_number("jack_success")
|
||||
1
|
||||
"""
|
||||
# Try dot separator first (.v)
|
||||
if ".v" in prompt_id:
|
||||
version_str = prompt_id.split(".v")[1]
|
||||
try:
|
||||
return int(version_str)
|
||||
except ValueError:
|
||||
pass
|
||||
|
||||
# Try underscore separator (_v)
|
||||
if "_v" in prompt_id:
|
||||
version_str = prompt_id.split("_v")[1]
|
||||
try:
|
||||
return int(version_str)
|
||||
except ValueError:
|
||||
pass
|
||||
|
||||
return 1
|
||||
|
||||
|
||||
def construct_versioned_prompt_id(prompt_id: str, version: int | None = None) -> str:
|
||||
"""
|
||||
Construct a versioned prompt ID from a base prompt_id and version number.
|
||||
|
||||
Args:
|
||||
prompt_id: Base prompt ID (e.g., "jack_success")
|
||||
version: Version number (if None, returns the base prompt_id unchanged)
|
||||
|
||||
Returns:
|
||||
Versioned prompt ID (e.g., "jack_success.v4")
|
||||
|
||||
Examples:
|
||||
>>> construct_versioned_prompt_id("jack_success", 4)
|
||||
"jack_success.v4"
|
||||
>>> construct_versioned_prompt_id("jack_success", None)
|
||||
"jack_success"
|
||||
>>> construct_versioned_prompt_id("jack_success.v2", 4)
|
||||
"jack_success.v4"
|
||||
"""
|
||||
if version is None:
|
||||
return prompt_id
|
||||
|
||||
# Strip any existing version suffix first
|
||||
base_id: Final = get_base_prompt_id(prompt_id)
|
||||
return f"{base_id}.v{version}"
|
||||
|
||||
|
||||
def get_latest_version_prompt_id(prompt_id: str, all_prompt_ids: Mapping[str, object]) -> str:
|
||||
"""
|
||||
Find the latest version of a prompt from available prompt IDs.
|
||||
|
||||
Args:
|
||||
prompt_id: Base prompt ID or versioned prompt ID (e.g., "jack_success" or "jack_success.v2")
|
||||
all_prompt_ids: Dictionary of all available prompt IDs (keys are prompt IDs)
|
||||
|
||||
Returns:
|
||||
The prompt ID with the highest version number, or the original prompt_id if no versions exist
|
||||
|
||||
Examples:
|
||||
>>> all_ids = {"jack.v1": {}, "jack.v2": {}, "jack.v3": {}}
|
||||
>>> get_latest_version_prompt_id("jack", all_ids)
|
||||
"jack.v3"
|
||||
>>> get_latest_version_prompt_id("jack.v1", all_ids)
|
||||
"jack.v3"
|
||||
>>> all_ids = {"simple": {}}
|
||||
>>> get_latest_version_prompt_id("simple", all_ids)
|
||||
"simple"
|
||||
"""
|
||||
base_id: Final = get_base_prompt_id(prompt_id=prompt_id)
|
||||
|
||||
# Find all versions of this prompt
|
||||
matching_versions: Final = []
|
||||
for stored_prompt_id in all_prompt_ids:
|
||||
if get_base_prompt_id(prompt_id=stored_prompt_id) == base_id:
|
||||
version_num = get_version_number(prompt_id=stored_prompt_id)
|
||||
matching_versions.append((version_num, stored_prompt_id))
|
||||
|
||||
# Use the highest version number
|
||||
if matching_versions:
|
||||
matching_versions.sort(reverse=True)
|
||||
return matching_versions[0][1]
|
||||
else:
|
||||
# No versioned prompts found, use the base ID as-is
|
||||
return prompt_id
|
||||
|
||||
|
||||
def get_latest_prompt_versions(prompts: list[PromptSpec]) -> list[PromptSpec]:
|
||||
"""
|
||||
Filter a list of prompts to return only the latest version of each unique prompt.
|
||||
|
||||
Args:
|
||||
prompts: List of PromptSpec objects
|
||||
|
||||
Returns:
|
||||
List of PromptSpec objects with only the latest version of each prompt
|
||||
Filter prompts down to the latest version per (base prompt id, environment).
|
||||
"""
|
||||
latest_prompts: Final[dict[str, PromptSpec]] = {}
|
||||
|
||||
for prompt in prompts:
|
||||
base_id = get_base_prompt_id(prompt_id=prompt.prompt_id)
|
||||
version = get_version_number(prompt_id=prompt.prompt_id)
|
||||
|
||||
# Keep the prompt with the highest version number
|
||||
if base_id not in latest_prompts:
|
||||
latest_prompts[base_id] = prompt
|
||||
else:
|
||||
existing_version = get_version_number(prompt_id=latest_prompts[base_id].prompt_id)
|
||||
if version > existing_version:
|
||||
latest_prompts[base_id] = prompt
|
||||
|
||||
sorted_prompts: Final = sorted(prompts, key=lambda prompt: get_version_number(prompt_id=prompt.prompt_id))
|
||||
latest_prompts: Final = {
|
||||
(get_base_prompt_id(prompt_id=prompt.prompt_id), prompt_environment_or_default(prompt.environment)): prompt
|
||||
for prompt in sorted_prompts
|
||||
}
|
||||
return list(latest_prompts.values())
|
||||
|
||||
|
||||
async def get_next_version_for_prompt(
|
||||
prisma_client: "PrismaClient", prompt_id: str, environment: str = "development"
|
||||
prisma_client: "PrismaClient", prompt_id: str, environment: str = DEFAULT_PROMPT_ENVIRONMENT
|
||||
) -> int:
|
||||
"""
|
||||
Get the next version number for a prompt in a specific environment.
|
||||
|
|
@ -403,11 +264,14 @@ async def list_prompts(
|
|||
if key_metadata is not None:
|
||||
prompts: Final = cast(list[str] | None, key_metadata.get("prompts", None))
|
||||
if prompts is not None:
|
||||
all_prompts = [
|
||||
IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS[prompt_id]
|
||||
for prompt_id in prompts
|
||||
if prompt_id in IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS
|
||||
allowed_prompt_ids: Final = frozenset(prompts)
|
||||
allowed_prompts: Final = [
|
||||
spec
|
||||
for spec in IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS.values()
|
||||
if spec.prompt_id in allowed_prompt_ids
|
||||
or get_base_prompt_id(prompt_id=spec.prompt_id) in allowed_prompt_ids
|
||||
]
|
||||
all_prompts = get_latest_prompt_versions(prompts=allowed_prompts)
|
||||
if environment:
|
||||
all_prompts = [p for p in all_prompts if p.environment == environment]
|
||||
prompt_list: Final = []
|
||||
|
|
@ -576,7 +440,7 @@ def _get_prompt_template(prompt_spec: PromptSpec, base_prompt_id: str) -> Prompt
|
|||
metadata=parsed.get("metadata"),
|
||||
)
|
||||
else:
|
||||
prompt_callback: Final = IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_by_id(prompt_spec.prompt_id)
|
||||
prompt_callback: Final = IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_for_prompt(prompt=prompt_spec)
|
||||
if prompt_callback is not None:
|
||||
integration_name: Final = prompt_callback.integration_name
|
||||
if integration_name == "dotprompt":
|
||||
|
|
@ -690,15 +554,10 @@ async def get_prompt_info(
|
|||
if env_prompts:
|
||||
prompt_spec = create_versioned_prompt_spec(db_prompt=env_prompts[0])
|
||||
|
||||
# Fallback: use in-memory registry (no environment filter)
|
||||
if prompt_spec is None and environment is None:
|
||||
prompt_spec = IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id(prompt_id)
|
||||
if prompt_spec is None:
|
||||
latest_prompt_id: Final = get_latest_version_prompt_id(
|
||||
prompt_id=prompt_id,
|
||||
all_prompt_ids=IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS,
|
||||
)
|
||||
prompt_spec = IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id(latest_prompt_id)
|
||||
if prompt_spec is None:
|
||||
prompt_spec = IN_MEMORY_PROMPT_REGISTRY.resolve_prompt_spec(
|
||||
prompt_id, version=requested_version, environment=environment
|
||||
)
|
||||
|
||||
if prompt_spec is None:
|
||||
raise HTTPException(
|
||||
|
|
@ -785,7 +644,7 @@ async def create_prompt(
|
|||
environment: Final = (
|
||||
request.prompt_info.environment
|
||||
if request.prompt_info and request.prompt_info.environment
|
||||
else "development"
|
||||
else DEFAULT_PROMPT_ENVIRONMENT
|
||||
)
|
||||
|
||||
# Get next version number
|
||||
|
|
@ -885,7 +744,7 @@ async def update_prompt(
|
|||
environment: Final = (
|
||||
request.prompt_info.environment
|
||||
if request.prompt_info and request.prompt_info.environment
|
||||
else "development"
|
||||
else DEFAULT_PROMPT_ENVIRONMENT
|
||||
)
|
||||
|
||||
# Check if any version of this prompt exists (in any environment)
|
||||
|
|
@ -897,9 +756,7 @@ async def update_prompt(
|
|||
detail=f"Prompt with ID {base_prompt_id} not found",
|
||||
)
|
||||
|
||||
# Check if it's a config prompt
|
||||
existing_in_memory: Final = IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id(prompt_id)
|
||||
if existing_in_memory and existing_in_memory.prompt_info.prompt_type == "config":
|
||||
if IN_MEMORY_PROMPT_REGISTRY.has_config_prompt(base_prompt_id=base_prompt_id):
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail="Cannot update config prompts.",
|
||||
|
|
@ -988,40 +845,26 @@ async def delete_prompt(
|
|||
raise HTTPException(status_code=500, detail=CommonProxyErrors.db_not_connected_error.value)
|
||||
|
||||
try:
|
||||
# Try to get prompt directly first
|
||||
existing_prompt = IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id(prompt_id)
|
||||
|
||||
# If not found, try to find the latest version
|
||||
if existing_prompt is None:
|
||||
latest_prompt_id: Final = get_latest_version_prompt_id(
|
||||
prompt_id=prompt_id,
|
||||
all_prompt_ids=IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS,
|
||||
)
|
||||
existing_prompt = IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id(latest_prompt_id)
|
||||
# Use the resolved prompt_id for deletion
|
||||
prompt_id = latest_prompt_id
|
||||
base_prompt_id: Final = get_base_prompt_id(prompt_id=prompt_id)
|
||||
existing_prompt: Final = IN_MEMORY_PROMPT_REGISTRY.resolve_prompt_spec(prompt_id, environment=environment)
|
||||
|
||||
if existing_prompt is None:
|
||||
raise HTTPException(status_code=404, detail=f"Prompt with ID {prompt_id} not found")
|
||||
|
||||
if existing_prompt.prompt_info.prompt_type == "config":
|
||||
if IN_MEMORY_PROMPT_REGISTRY.has_config_prompt(base_prompt_id=base_prompt_id):
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail="Cannot delete config prompts.",
|
||||
)
|
||||
|
||||
# Get the base prompt ID (without version suffix) for database deletion
|
||||
base_prompt_id: Final = get_base_prompt_id(prompt_id=prompt_id)
|
||||
|
||||
# Build delete filter; scope to environment if provided
|
||||
delete_where: Final[dict[str, str]] = {"prompt_id": base_prompt_id}
|
||||
if environment:
|
||||
delete_where["environment"] = environment
|
||||
|
||||
# Delete versions from the database (scoped to environment if provided)
|
||||
delete_where: Final[dict[str, str]] = {
|
||||
"prompt_id": base_prompt_id,
|
||||
**({"environment": environment} if environment else {}),
|
||||
}
|
||||
await _prompt_table(prisma_client).delete_many(where=delete_where)
|
||||
|
||||
IN_MEMORY_PROMPT_REGISTRY.delete_prompts_by_base_id(base_prompt_id, environment=environment or None)
|
||||
IN_MEMORY_PROMPT_REGISTRY.delete_prompts_by_base_id(
|
||||
base_prompt_id=base_prompt_id, environment=environment or None
|
||||
)
|
||||
|
||||
env_msg: Final = f" from {environment}" if environment else ""
|
||||
return {"message": f"Prompt {base_prompt_id} deleted successfully{env_msg}"}
|
||||
|
|
@ -1093,7 +936,7 @@ async def patch_prompt(
|
|||
try:
|
||||
# Resolve the target row: find the latest version in the given environment
|
||||
base_prompt_id: Final = get_base_prompt_id(prompt_id=prompt_id)
|
||||
env: Final = environment or "development"
|
||||
env: Final = prompt_environment_or_default(environment)
|
||||
requested_version: Final = get_version_number(prompt_id=prompt_id) if prompt_id != base_prompt_id else None
|
||||
|
||||
# Build query to find the exact row by composite unique key
|
||||
|
|
@ -1117,11 +960,7 @@ async def patch_prompt(
|
|||
|
||||
target_row: Final = db_rows[0]
|
||||
|
||||
# Check if prompt exists in memory
|
||||
versioned_id: Final = f"{base_prompt_id}.v{target_row.version}"
|
||||
existing_prompt: Final = IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id(versioned_id)
|
||||
|
||||
if existing_prompt and existing_prompt.prompt_info.prompt_type == "config":
|
||||
if IN_MEMORY_PROMPT_REGISTRY.has_config_prompt(base_prompt_id=base_prompt_id):
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail="Cannot update config prompts.",
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
import importlib
|
||||
import os
|
||||
from collections.abc import Callable
|
||||
from collections.abc import Callable, Sequence
|
||||
from pathlib import Path
|
||||
from typing import Final
|
||||
|
||||
|
|
@ -14,6 +14,87 @@ from litellm.types.prompts.init_prompts import (
|
|||
|
||||
prompt_initializer_registry = {}
|
||||
|
||||
DEFAULT_PROMPT_ENVIRONMENT: Final = "development"
|
||||
PROMPT_ENVIRONMENT_SERVE_PRECEDENCE: Final = ("production", "staging", "development")
|
||||
|
||||
|
||||
def get_base_prompt_id(prompt_id: str) -> str:
|
||||
"""
|
||||
Extract the base prompt ID by stripping the version suffix if present.
|
||||
|
||||
Examples:
|
||||
>>> get_base_prompt_id("jack_success.v1")
|
||||
"jack_success"
|
||||
>>> get_base_prompt_id("jack_success_v1")
|
||||
"jack_success"
|
||||
>>> get_base_prompt_id("jack_success")
|
||||
"jack_success"
|
||||
"""
|
||||
if ".v" in prompt_id:
|
||||
return prompt_id.split(".v")[0]
|
||||
if "_v" in prompt_id:
|
||||
return prompt_id.split("_v")[0]
|
||||
return prompt_id
|
||||
|
||||
|
||||
def get_version_number(prompt_id: str) -> int:
|
||||
"""
|
||||
Extract the version number from a versioned prompt ID (defaults to 1).
|
||||
|
||||
Examples:
|
||||
>>> get_version_number("jack_success.v2")
|
||||
2
|
||||
>>> get_version_number("jack_success_v2")
|
||||
2
|
||||
>>> get_version_number("jack_success")
|
||||
1
|
||||
"""
|
||||
if ".v" in prompt_id:
|
||||
version_str = prompt_id.split(".v")[1]
|
||||
try:
|
||||
return int(version_str)
|
||||
except ValueError:
|
||||
pass
|
||||
|
||||
if "_v" in prompt_id:
|
||||
version_str = prompt_id.split("_v")[1]
|
||||
try:
|
||||
return int(version_str)
|
||||
except ValueError:
|
||||
pass
|
||||
|
||||
return 1
|
||||
|
||||
|
||||
def prompt_environment_or_default(environment: str | None) -> str:
|
||||
return environment or DEFAULT_PROMPT_ENVIRONMENT
|
||||
|
||||
|
||||
def registry_key_for_prompt(prompt: PromptSpec) -> str:
|
||||
return f"{prompt.prompt_id}::{prompt_environment_or_default(prompt.environment)}"
|
||||
|
||||
|
||||
def parse_prompt_version(raw_version: object) -> int | None:
|
||||
if isinstance(raw_version, bool):
|
||||
return None
|
||||
if isinstance(raw_version, int):
|
||||
return raw_version
|
||||
if isinstance(raw_version, str) and raw_version.isdigit():
|
||||
return int(raw_version)
|
||||
return None
|
||||
|
||||
|
||||
def _spec_version(prompt: PromptSpec) -> int:
|
||||
return prompt.version if prompt.version is not None else get_version_number(prompt_id=prompt.prompt_id)
|
||||
|
||||
|
||||
def _default_serve_environment(prompts: Sequence[PromptSpec]) -> str:
|
||||
present: Final = frozenset(prompt_environment_or_default(prompt.environment) for prompt in prompts)
|
||||
ladder_pick: Final = next((env for env in PROMPT_ENVIRONMENT_SERVE_PRECEDENCE if env in present), None)
|
||||
if ladder_pick is not None:
|
||||
return ladder_pick
|
||||
return min(present) if present else DEFAULT_PROMPT_ENVIRONMENT
|
||||
|
||||
|
||||
def get_prompt_initializer_from_integrations():
|
||||
"""
|
||||
|
|
@ -113,17 +194,16 @@ class InMemoryPromptRegistry:
|
|||
"""
|
||||
import litellm
|
||||
|
||||
prompt_id: Final = prompt.prompt_id
|
||||
if prompt_id in self.IN_MEMORY_PROMPTS:
|
||||
verbose_proxy_logger.debug("prompt_id already exists in IN_MEMORY_PROMPTS")
|
||||
return self.IN_MEMORY_PROMPTS[prompt_id]
|
||||
registry_key: Final = registry_key_for_prompt(prompt)
|
||||
if registry_key in self.IN_MEMORY_PROMPTS:
|
||||
verbose_proxy_logger.debug("prompt already exists in IN_MEMORY_PROMPTS")
|
||||
return self.IN_MEMORY_PROMPTS[registry_key]
|
||||
|
||||
parsed_prompt, custom_prompt_callback = self._build_prompt_callback(prompt=prompt)
|
||||
litellm.logging_callback_manager.add_litellm_callback(custom_prompt_callback)
|
||||
|
||||
# store references to the prompt in memory
|
||||
self.IN_MEMORY_PROMPTS[prompt_id] = parsed_prompt
|
||||
self.prompt_id_to_custom_prompt[prompt_id] = custom_prompt_callback
|
||||
self.IN_MEMORY_PROMPTS[registry_key] = parsed_prompt
|
||||
self.prompt_id_to_custom_prompt[registry_key] = custom_prompt_callback
|
||||
|
||||
return parsed_prompt
|
||||
|
||||
|
|
@ -166,68 +246,93 @@ class InMemoryPromptRegistry:
|
|||
import litellm
|
||||
|
||||
parsed_prompt, new_callback = self._build_prompt_callback(prompt=prompt)
|
||||
stale_callback: Final = self.prompt_id_to_custom_prompt.pop(prompt.prompt_id, None)
|
||||
self.IN_MEMORY_PROMPTS.pop(prompt.prompt_id, None)
|
||||
registry_key: Final = registry_key_for_prompt(parsed_prompt)
|
||||
stale_callback: Final = self.prompt_id_to_custom_prompt.pop(registry_key, None)
|
||||
self.IN_MEMORY_PROMPTS.pop(registry_key, None)
|
||||
if stale_callback is not None:
|
||||
litellm.logging_callback_manager.remove_callback_from_all_lists(stale_callback)
|
||||
litellm.logging_callback_manager.add_litellm_callback(new_callback)
|
||||
self.IN_MEMORY_PROMPTS[prompt.prompt_id] = parsed_prompt
|
||||
self.prompt_id_to_custom_prompt[prompt.prompt_id] = new_callback
|
||||
self.IN_MEMORY_PROMPTS[registry_key] = parsed_prompt
|
||||
self.prompt_id_to_custom_prompt[registry_key] = new_callback
|
||||
return parsed_prompt
|
||||
|
||||
def sync_prompt_from_db(self, prompt: PromptSpec) -> PromptSpec | None:
|
||||
existing: Final = self.IN_MEMORY_PROMPTS.get(prompt.prompt_id)
|
||||
existing: Final = self.IN_MEMORY_PROMPTS.get(registry_key_for_prompt(prompt))
|
||||
if existing is None:
|
||||
return self.initialize_prompt(prompt=prompt)
|
||||
if existing.litellm_params == prompt.litellm_params and existing.prompt_info == prompt.prompt_info:
|
||||
return existing
|
||||
return self.reload_prompt(prompt=prompt)
|
||||
|
||||
def get_prompt_by_id(self, prompt_id: str) -> PromptSpec | None:
|
||||
def resolve_prompt_spec(
|
||||
self,
|
||||
prompt_id: str,
|
||||
version: int | None = None,
|
||||
environment: str | None = None,
|
||||
) -> PromptSpec | None:
|
||||
"""
|
||||
Get a prompt by its ID from memory
|
||||
"""
|
||||
return self.IN_MEMORY_PROMPTS.get(prompt_id)
|
||||
Resolve a prompt spec by base prompt id, optional version, and optional environment.
|
||||
|
||||
def get_prompt_callback_by_id(self, prompt_id: str) -> CustomPromptManagement | None:
|
||||
With no environment, resolves within the default serve environment
|
||||
(production > staging > development > alphabetical first present).
|
||||
With no version, resolves to the highest version in the chosen environment.
|
||||
"""
|
||||
Get a prompt callback by its ID from memory
|
||||
"""
|
||||
return self.prompt_id_to_custom_prompt.get(prompt_id)
|
||||
base_prompt_id: Final = get_base_prompt_id(prompt_id=prompt_id)
|
||||
base_matches: Final = tuple(
|
||||
spec
|
||||
for spec in self.IN_MEMORY_PROMPTS.values()
|
||||
if get_base_prompt_id(prompt_id=spec.prompt_id) == base_prompt_id
|
||||
)
|
||||
if not base_matches:
|
||||
return None
|
||||
resolved_environment: Final = (
|
||||
environment if environment is not None else _default_serve_environment(base_matches)
|
||||
)
|
||||
env_matches: Final = tuple(
|
||||
spec for spec in base_matches if prompt_environment_or_default(spec.environment) == resolved_environment
|
||||
)
|
||||
if not env_matches:
|
||||
return None
|
||||
if version is not None:
|
||||
return next((spec for spec in env_matches if _spec_version(spec) == version), None)
|
||||
return max(env_matches, key=_spec_version)
|
||||
|
||||
def remove_prompt(self, prompt_id: str) -> None:
|
||||
def get_prompt_callback_for_prompt(self, prompt: PromptSpec) -> CustomPromptManagement | None:
|
||||
return self.prompt_id_to_custom_prompt.get(registry_key_for_prompt(prompt))
|
||||
|
||||
def has_config_prompt(self, base_prompt_id: str) -> bool:
|
||||
return any(
|
||||
spec.prompt_info.prompt_type == "config"
|
||||
for spec in self.IN_MEMORY_PROMPTS.values()
|
||||
if get_base_prompt_id(prompt_id=spec.prompt_id) == base_prompt_id
|
||||
)
|
||||
|
||||
def remove_prompt(self, registry_key: str) -> None:
|
||||
import litellm
|
||||
|
||||
self.IN_MEMORY_PROMPTS.pop(prompt_id, None)
|
||||
stale_callback: Final = self.prompt_id_to_custom_prompt.pop(prompt_id, None)
|
||||
self.IN_MEMORY_PROMPTS.pop(registry_key, None)
|
||||
stale_callback: Final = self.prompt_id_to_custom_prompt.pop(registry_key, None)
|
||||
if stale_callback is not None:
|
||||
litellm.logging_callback_manager.remove_callback_from_all_lists(stale_callback)
|
||||
|
||||
def delete_prompts_by_base_id(self, base_prompt_id: str, environment: str | None = None) -> list[str]:
|
||||
"""
|
||||
Delete all prompts matching the given base prompt ID from memory, along with their
|
||||
registered callbacks; scoped to one environment when given.
|
||||
Delete matching prompts from memory, along with their registered callbacks,
|
||||
scoped to one environment when given.
|
||||
|
||||
Args:
|
||||
base_prompt_id: The base prompt ID (without version suffix)
|
||||
environment: When set, only delete prompts deployed to this environment
|
||||
|
||||
Returns:
|
||||
List of prompt IDs that were deleted
|
||||
Returns the registry keys that were deleted.
|
||||
"""
|
||||
from litellm.proxy.prompts.prompt_endpoints import get_base_prompt_id
|
||||
|
||||
prompts_to_delete: Final = [
|
||||
pid
|
||||
for pid, prompt in self.IN_MEMORY_PROMPTS.items()
|
||||
if get_base_prompt_id(prompt_id=pid) == base_prompt_id
|
||||
and (environment is None or prompt.environment == environment)
|
||||
keys_to_delete: Final = [
|
||||
key
|
||||
for key, spec in self.IN_MEMORY_PROMPTS.items()
|
||||
if get_base_prompt_id(prompt_id=spec.prompt_id) == base_prompt_id
|
||||
and (environment is None or prompt_environment_or_default(spec.environment) == environment)
|
||||
]
|
||||
|
||||
for pid in prompts_to_delete:
|
||||
self.remove_prompt(prompt_id=pid)
|
||||
for key in keys_to_delete:
|
||||
self.remove_prompt(registry_key=key)
|
||||
|
||||
return prompts_to_delete
|
||||
return keys_to_delete
|
||||
|
||||
|
||||
IN_MEMORY_PROMPT_REGISTRY: Final = InMemoryPromptRegistry()
|
||||
|
|
|
|||
|
|
@ -7594,7 +7594,7 @@ class ProxyConfig:
|
|||
return create_versioned_prompt_spec(db_prompt=db_prompt)
|
||||
|
||||
async def _init_prompts_in_db(self, prisma_client: PrismaClient):
|
||||
from litellm.proxy.prompts.prompt_registry import IN_MEMORY_PROMPT_REGISTRY
|
||||
from litellm.proxy.prompts.prompt_registry import IN_MEMORY_PROMPT_REGISTRY, registry_key_for_prompt
|
||||
from litellm.types.prompts.init_prompts import PromptSpec
|
||||
|
||||
def parse_row(db_prompt: object) -> PromptSpec | None:
|
||||
|
|
@ -7609,21 +7609,12 @@ class ProxyConfig:
|
|||
return None
|
||||
|
||||
try:
|
||||
prompt_ids_loaded_before_db_read: Final = frozenset(IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS)
|
||||
registry_keys_loaded_before_db_read: Final = frozenset(IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS)
|
||||
prompts_in_db: Final[Sequence[object]] = await PromptRepository(prisma_client).table.find_many()
|
||||
parsed_specs: Final[tuple[PromptSpec, ...]] = tuple(
|
||||
spec for row in prompts_in_db if (spec := parse_row(row)) is not None
|
||||
)
|
||||
newest_spec_per_id: Final[Mapping[str, PromptSpec]] = MappingProxyType(
|
||||
{
|
||||
spec.prompt_id: spec
|
||||
for spec in sorted(
|
||||
parsed_specs,
|
||||
key=lambda s: s.updated_at.timestamp() if s.updated_at else float("-inf"),
|
||||
)
|
||||
}
|
||||
)
|
||||
for prompt_spec in newest_spec_per_id.values():
|
||||
for prompt_spec in parsed_specs:
|
||||
try:
|
||||
IN_MEMORY_PROMPT_REGISTRY.sync_prompt_from_db(prompt=prompt_spec)
|
||||
except Exception as prompt_sync_error: # noqa: BLE001 # one poisoned row must not block syncing the remaining prompts
|
||||
|
|
@ -7635,15 +7626,16 @@ class ProxyConfig:
|
|||
# An unparsable row still exists in the DB, so skip the sweep rather than unload its in-memory copy
|
||||
every_row_parsed: Final = len(parsed_specs) == len(prompts_in_db)
|
||||
if every_row_parsed:
|
||||
deleted_db_prompt_ids: Final = tuple(
|
||||
prompt_id
|
||||
for prompt_id in prompt_ids_loaded_before_db_read
|
||||
if (loaded_spec := IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS.get(prompt_id)) is not None
|
||||
db_registry_keys: Final = frozenset(registry_key_for_prompt(spec) for spec in parsed_specs)
|
||||
deleted_db_registry_keys: Final = tuple(
|
||||
registry_key
|
||||
for registry_key in registry_keys_loaded_before_db_read
|
||||
if (loaded_spec := IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS.get(registry_key)) is not None
|
||||
and loaded_spec.prompt_info.prompt_type == "db"
|
||||
and prompt_id not in newest_spec_per_id
|
||||
and registry_key not in db_registry_keys
|
||||
)
|
||||
for deleted_prompt_id in deleted_db_prompt_ids:
|
||||
IN_MEMORY_PROMPT_REGISTRY.remove_prompt(prompt_id=deleted_prompt_id)
|
||||
for deleted_registry_key in deleted_db_registry_keys:
|
||||
IN_MEMORY_PROMPT_REGISTRY.remove_prompt(registry_key=deleted_registry_key)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.debug("litellm.proxy.proxy_server.py::ProxyConfig:_init_prompts_in_db - %s", e)
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -1478,28 +1478,27 @@ class ProxyLogging:
|
|||
) -> None:
|
||||
"""Process prompt template if applicable."""
|
||||
|
||||
from litellm.proxy.prompts.prompt_endpoints import (
|
||||
construct_versioned_prompt_id,
|
||||
get_latest_version_prompt_id,
|
||||
)
|
||||
from litellm.proxy.prompts.prompt_registry import IN_MEMORY_PROMPT_REGISTRY
|
||||
from litellm.responses.utils import ResponsesAPIRequestUtils
|
||||
from litellm.utils import get_non_default_completion_params
|
||||
|
||||
if prompt_version is None:
|
||||
lookup_prompt_id = get_latest_version_prompt_id(
|
||||
prompt_id=prompt_id,
|
||||
all_prompt_ids=IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS,
|
||||
)
|
||||
else:
|
||||
lookup_prompt_id = construct_versioned_prompt_id(prompt_id=prompt_id, version=prompt_version)
|
||||
|
||||
custom_logger: Final = IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_by_id(lookup_prompt_id)
|
||||
prompt_spec: Final = IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id(lookup_prompt_id)
|
||||
raw_prompt_environment: Final = data.get("prompt_environment", None)
|
||||
prompt_environment: Final = raw_prompt_environment if isinstance(raw_prompt_environment, str) else None
|
||||
prompt_spec: Final = IN_MEMORY_PROMPT_REGISTRY.resolve_prompt_spec(
|
||||
prompt_id,
|
||||
version=prompt_version,
|
||||
environment=prompt_environment,
|
||||
)
|
||||
custom_logger: Final = (
|
||||
IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_for_prompt(prompt=prompt_spec)
|
||||
if prompt_spec is not None
|
||||
else None
|
||||
)
|
||||
litellm_prompt_id: str | None = None
|
||||
if prompt_spec is not None:
|
||||
litellm_prompt_id = prompt_spec.litellm_params.prompt_id
|
||||
data.pop("prompt_id", None)
|
||||
data.pop("prompt_environment", None)
|
||||
|
||||
if custom_logger and prompt_spec is not None:
|
||||
is_responses_call: Final = call_type == "aresponses"
|
||||
|
|
@ -1542,6 +1541,7 @@ class ProxyLogging:
|
|||
data.pop("prompt_variables", None)
|
||||
data.pop("prompt_label", None)
|
||||
data.pop("prompt_version", None)
|
||||
data.pop("prompt_environment", None)
|
||||
|
||||
def _process_guardrail_metadata(self, data: dict) -> None:
|
||||
"""Process guardrails from metadata and add to applied_guardrails."""
|
||||
|
|
@ -1750,7 +1750,6 @@ class ProxyLogging:
|
|||
|
||||
litellm_logging_obj: Final = cast(Optional["LiteLLMLoggingObj"], data.get("litellm_logging_obj", None))
|
||||
prompt_id: Final[str | None] = data.get("prompt_id", None)
|
||||
prompt_version: Final[int | None] = data.get("prompt_version", None)
|
||||
|
||||
## PROMPT TEMPLATE CHECK ##
|
||||
|
||||
|
|
@ -1760,11 +1759,13 @@ class ProxyLogging:
|
|||
and prompt_id is not None
|
||||
and (call_type == "completion" or call_type == "acompletion" or call_type == "aresponses")
|
||||
):
|
||||
from litellm.proxy.prompts.prompt_registry import parse_prompt_version
|
||||
|
||||
await self._process_prompt_template(
|
||||
data=data,
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
prompt_id=prompt_id,
|
||||
prompt_version=prompt_version,
|
||||
prompt_version=parse_prompt_version(data.get("prompt_version", None)),
|
||||
call_type=call_type,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
)
|
||||
|
|
|
|||
|
|
@ -595,16 +595,20 @@ set_live_deployment_replay(_replay_live_router_model_cost)
|
|||
|
||||
|
||||
# Kwargs that carry no signal about the failed attempt, so log_retry drops them from a
|
||||
# breadcrumb entirely: the request payload and the router-internal walk state. Credentials are
|
||||
# handled separately by mask_credentials_in_payload, which scrubs credential-named values from
|
||||
# whatever kwargs remain rather than trying to enumerate every credential-bearing key here.
|
||||
# breadcrumb entirely: the request payload, the proxy's snapshot of the inbound request (its body
|
||||
# aliases the live request metadata, earlier breadcrumbs included, so copying it would nest every
|
||||
# breadcrumb inside the next one), and the router-internal walk state. Credentials are handled
|
||||
# separately by mask_credentials_in_payload, which scrubs credential-named values from whatever
|
||||
# kwargs remain rather than trying to enumerate every credential-bearing key here.
|
||||
RETRY_BREADCRUMB_EXCLUDED_KWARGS: Final = frozenset(
|
||||
(
|
||||
"messages",
|
||||
"original_function",
|
||||
"attempted_targets",
|
||||
"proxy_server_request",
|
||||
)
|
||||
)
|
||||
RETRY_BREADCRUMB_LIMIT: Final = 4
|
||||
|
||||
|
||||
class Router:
|
||||
|
|
@ -965,7 +969,6 @@ class Router:
|
|||
self.total_calls: defaultdict = defaultdict(int) # dict to store total calls made to each model
|
||||
self.fail_calls: defaultdict = defaultdict(int) # dict to store fail_calls made to each model
|
||||
self.success_calls: defaultdict = defaultdict(int) # dict to store success_calls made to each model
|
||||
self.previous_models: list = [] # list to store failed calls (passed in as metadata to next call)
|
||||
|
||||
# make Router.chat.completions.create compatible for openai.chat.completions.create
|
||||
default_litellm_params = default_litellm_params or {}
|
||||
|
|
@ -8144,35 +8147,31 @@ class Router:
|
|||
"""
|
||||
When a retry or fallback happens, log the details of the just failed model call - similar to Sentry breadcrumbing
|
||||
"""
|
||||
try:
|
||||
_metadata_var: Final = "litellm_metadata" if "litellm_metadata" in kwargs else "metadata"
|
||||
# Log failed model as the previous model
|
||||
previous_model: Final = {
|
||||
_metadata_var: Final = "litellm_metadata" if "litellm_metadata" in kwargs else "metadata"
|
||||
request_metadata: Final[Mapping[str, object]] = kwargs[_metadata_var]
|
||||
attempt_kwargs: Final = MappingProxyType(
|
||||
{k: v for k, v in kwargs.items() if k != _metadata_var and k not in RETRY_BREADCRUMB_EXCLUDED_KWARGS}
|
||||
)
|
||||
attempt_metadata: Final = MappingProxyType(
|
||||
{k: v for k, v in request_metadata.items() if k != "previous_models"}
|
||||
)
|
||||
previous_model: Final = MappingProxyType(
|
||||
{
|
||||
"exception_type": type(e).__name__,
|
||||
"exception_string": str(e),
|
||||
**attempt_kwargs,
|
||||
_metadata_var: attempt_metadata,
|
||||
}
|
||||
for (
|
||||
k,
|
||||
v,
|
||||
) in kwargs.items(): # log everything in kwargs except the old previous_models value - prevent nesting
|
||||
if k != _metadata_var and k not in RETRY_BREADCRUMB_EXCLUDED_KWARGS:
|
||||
previous_model[k] = v
|
||||
elif k == _metadata_var and isinstance(v, dict):
|
||||
previous_model[_metadata_var] = {}
|
||||
for metadata_k, metadata_v in kwargs[_metadata_var].items():
|
||||
if metadata_k != "previous_models":
|
||||
previous_model[k][metadata_k] = metadata_v
|
||||
|
||||
# check current size of self.previous_models, if it's larger than 3, remove the first element
|
||||
if len(self.previous_models) > 3:
|
||||
self.previous_models.pop(0)
|
||||
|
||||
scrubbed_previous_model: Final = mask_credentials_in_payload(previous_model)
|
||||
self.previous_models.append(scrubbed_previous_model)
|
||||
kwargs[_metadata_var]["previous_models"] = self.previous_models
|
||||
return kwargs
|
||||
except Exception as e:
|
||||
raise e
|
||||
)
|
||||
earlier_breadcrumbs: Final = request_metadata.get("previous_models")
|
||||
kept_breadcrumbs: Final[tuple[object, ...]] = (
|
||||
tuple(earlier_breadcrumbs)[-(RETRY_BREADCRUMB_LIMIT - 1) :]
|
||||
if isinstance(earlier_breadcrumbs, (list, tuple))
|
||||
else ()
|
||||
)
|
||||
breadcrumbs: Final = (*kept_breadcrumbs, mask_credentials_in_payload(previous_model))
|
||||
kwargs[_metadata_var]["previous_models"] = breadcrumbs # rebind-ok: the logging object already holds this dict
|
||||
return kwargs
|
||||
|
||||
def _update_usage(self, deployment_id: str, parent_otel_span: Span | None) -> int:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -2,23 +2,41 @@
|
|||
Auto-Routing Strategy that works with a Semantic Router Config
|
||||
"""
|
||||
|
||||
from collections.abc import Mapping
|
||||
from typing import TYPE_CHECKING, Any, Final, Optional
|
||||
|
||||
from pydantic import BaseModel, ConfigDict
|
||||
|
||||
from litellm._logging import verbose_router_logger
|
||||
from litellm.constants import DEFAULT_AUTO_ROUTER_MAX_INPUT_CHARS
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.litellm_core_utils.internal_call_metadata import (
|
||||
effective_turn_off_message_logging,
|
||||
forwarded_internal_call_metadata,
|
||||
parent_session_kwargs,
|
||||
)
|
||||
from litellm.types.utils import AUTOROUTER_CLASSIFIER_CALL_ORIGIN
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from semantic_router.routers import SemanticRouter
|
||||
from semantic_router.routers.base import Route
|
||||
|
||||
from litellm.router import Router
|
||||
from litellm.router_strategy.auto_router.litellm_encoder import LiteLLMRouterEncoder
|
||||
from litellm.types.router import PreRoutingHookResponse
|
||||
else:
|
||||
Router = Any
|
||||
PreRoutingHookResponse = Any
|
||||
Route = Any
|
||||
SemanticRouter = Any
|
||||
LiteLLMRouterEncoder = Any
|
||||
|
||||
|
||||
class _CallerMetadata(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, extra="ignore")
|
||||
|
||||
metadata: Mapping[str, object] | None = None
|
||||
litellm_metadata: Mapping[str, object] | None = None
|
||||
|
||||
|
||||
class AutoRouter(CustomLogger):
|
||||
|
|
@ -50,6 +68,8 @@ class AutoRouter(CustomLogger):
|
|||
"""
|
||||
from semantic_router.routers import SemanticRouter
|
||||
|
||||
from litellm.router_strategy.auto_router.litellm_encoder import LiteLLMRouterEncoder
|
||||
|
||||
self.auto_router_config_path: str | None = auto_router_config_path
|
||||
self.auto_router_config: str | None = auto_router_config
|
||||
self.auto_sync_value = self.DEFAULT_AUTO_SYNC_VALUE
|
||||
|
|
@ -59,6 +79,11 @@ class AutoRouter(CustomLogger):
|
|||
self.embedding_model: str = embedding_model
|
||||
self.max_input_chars: int = max_input_chars
|
||||
self.litellm_router_instance: Router = litellm_router_instance
|
||||
self.encoder: LiteLLMRouterEncoder = LiteLLMRouterEncoder(
|
||||
litellm_router_instance=litellm_router_instance,
|
||||
model_name=embedding_model,
|
||||
max_input_chars=max_input_chars,
|
||||
)
|
||||
|
||||
def _load_semantic_routing_routes(self) -> list[Route]:
|
||||
from semantic_router.routers import SemanticRouter
|
||||
|
|
@ -129,9 +154,6 @@ class AutoRouter(CustomLogger):
|
|||
from semantic_router.routers import SemanticRouter
|
||||
|
||||
from litellm.litellm_core_utils.prompt_templates.factory import resolve_structured_messages
|
||||
from litellm.router_strategy.auto_router.litellm_encoder import (
|
||||
LiteLLMRouterEncoder,
|
||||
)
|
||||
from litellm.types.router import PreRoutingHookResponse
|
||||
|
||||
resolved_messages: Final = (
|
||||
|
|
@ -149,34 +171,47 @@ class AutoRouter(CustomLogger):
|
|||
#######################
|
||||
routelayer = SemanticRouter(
|
||||
routes=self.loaded_routes,
|
||||
encoder=LiteLLMRouterEncoder(
|
||||
litellm_router_instance=self.litellm_router_instance,
|
||||
model_name=self.embedding_model,
|
||||
max_input_chars=self.max_input_chars,
|
||||
),
|
||||
encoder=self.encoder,
|
||||
auto_sync=self.auto_sync_value,
|
||||
)
|
||||
self.routelayer = routelayer
|
||||
|
||||
message_content: Final = self._extract_text_from_messages(resolved_messages)
|
||||
route_name: Final = self._matched_route_name(routelayer, message_content)
|
||||
route_name: Final = await self._matched_route_name(routelayer, message_content, request_kwargs)
|
||||
|
||||
return PreRoutingHookResponse(
|
||||
model=route_name or self.default_model,
|
||||
messages=messages,
|
||||
)
|
||||
|
||||
def _matched_route_name(self, routelayer: "SemanticRouter", text: str) -> str | None:
|
||||
async def _matched_route_name(
|
||||
self, routelayer: "SemanticRouter", text: str, request_kwargs: Mapping[str, object]
|
||||
) -> str | None:
|
||||
"""Name of the route `text` matches, or None when nothing matched or the match failed.
|
||||
|
||||
The route layer embeds `text` to compare it against the routes, and that embedding call can
|
||||
`text` is embedded here rather than by `routelayer(text=...)` so the caller's metadata reaches
|
||||
`aembedding()` and the embedding's spend lands on the key/team that sent the request;
|
||||
SemanticRouter has no way to pass kwargs through to its encoder. That embedding call can
|
||||
fail (context limit, timeout, provider error). Choosing a model is a routing decision, so a
|
||||
failure here falls back to the default model rather than failing the user's request.
|
||||
"""
|
||||
from semantic_router.schema import RouteChoice
|
||||
|
||||
try:
|
||||
route_choice: Final = routelayer(text=text)
|
||||
caller: Final = _CallerMetadata.model_validate(request_kwargs)
|
||||
query_vector: Final = (
|
||||
await self.encoder.aencode_queries(
|
||||
[text],
|
||||
metadata=forwarded_internal_call_metadata(caller.metadata, AUTOROUTER_CLASSIFIER_CALL_ORIGIN),
|
||||
litellm_metadata=forwarded_internal_call_metadata(
|
||||
caller.litellm_metadata, AUTOROUTER_CLASSIFIER_CALL_ORIGIN
|
||||
),
|
||||
proxy_server_request={"body": {"model": self.embedding_model, "input": [text]}},
|
||||
turn_off_message_logging=effective_turn_off_message_logging(request_kwargs),
|
||||
**parent_session_kwargs(request_kwargs),
|
||||
)
|
||||
)[0]
|
||||
route_choice: Final = await routelayer.acall(vector=query_vector)
|
||||
except Exception as e: # noqa: BLE001 -- the embedding call behind the route layer can fail many ways (context limit, timeout, provider/network error); none of them may fail the request
|
||||
verbose_router_logger.warning(
|
||||
"AutoRouter: semantic routing failed (%s), falling back to default model %s", e, self.default_model
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -1,6 +1,8 @@
|
|||
from typing import Any, Final
|
||||
from collections.abc import Mapping
|
||||
from typing import Any, Final, Literal
|
||||
|
||||
from pydantic import BaseModel, field_validator
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
from litellm.proxy._types import (
|
||||
LiteLLM_UserTableWithKeyCount,
|
||||
|
|
@ -9,6 +11,17 @@ from litellm.proxy._types import (
|
|||
)
|
||||
|
||||
|
||||
class InsensitiveContains(TypedDict):
|
||||
contains: ReadOnly[str]
|
||||
mode: ReadOnly[Literal["insensitive"]]
|
||||
|
||||
|
||||
class UserSearchWhere(TypedDict):
|
||||
"""Prisma filter behind `/user/list?search=`: user_id or user_email contains the term, case-insensitive."""
|
||||
|
||||
OR: ReadOnly[tuple[Mapping[Literal["user_id", "user_email"], InsensitiveContains], ...]]
|
||||
|
||||
|
||||
class UserListResponse(BaseModel):
|
||||
"""
|
||||
Response model for the user list endpoint
|
||||
|
|
|
|||
|
|
@ -3630,6 +3630,7 @@ all_litellm_params = (
|
|||
"litellm_system_prompt",
|
||||
"provider_specific_header",
|
||||
"prompt_version",
|
||||
"prompt_environment",
|
||||
"api_base",
|
||||
"force_timeout",
|
||||
"logger_fn",
|
||||
|
|
|
|||
|
|
@ -83,6 +83,7 @@ from litellm.constants import (
|
|||
from litellm.litellm_core_utils.fallback_generalizations import (
|
||||
match_capability_generalizations,
|
||||
)
|
||||
from litellm.litellm_core_utils.sensitive_data_masker import redact_credentials_in_payload
|
||||
|
||||
_CachingHandlerResponse = None
|
||||
_LLMCachingHandler = None
|
||||
|
|
@ -543,6 +544,14 @@ def print_verbose(
|
|||
pass
|
||||
|
||||
|
||||
def _print_verbose_is_active() -> bool:
|
||||
"""Whether print_verbose would reach either of its two consumers, so a call site can skip
|
||||
building a payload nothing would read. _is_debugging_on() is not the same predicate: it reads
|
||||
litellm._logging.set_verbose, while print_verbose's print reads litellm.set_verbose, and
|
||||
assigning the documented litellm.set_verbose = True rebinds only the latter."""
|
||||
return litellm.set_verbose is True or verbose_logger.isEnabledFor(logging.DEBUG)
|
||||
|
||||
|
||||
####### CLIENT ###################
|
||||
# make it easy to log if completion/embedding runs succeeded or failed + see what happened | Non-Blocking
|
||||
def custom_llm_setup():
|
||||
|
|
@ -1284,16 +1293,18 @@ async def async_post_call_success_deployment_hook(
|
|||
except ValueError:
|
||||
typed_call_type = None # unknown call type
|
||||
|
||||
modified_response = response
|
||||
|
||||
CustomLogger: Final = _get_cached_custom_logger()
|
||||
for callback in litellm.callbacks:
|
||||
if isinstance(callback, CustomLogger):
|
||||
result = await callback.async_post_call_success_deployment_hook(
|
||||
request_data, cast(LLMResponseTypes, response), typed_call_type
|
||||
request_data, cast(LLMResponseTypes, modified_response), typed_call_type
|
||||
)
|
||||
if result is not None:
|
||||
return result
|
||||
modified_response = result
|
||||
|
||||
return response
|
||||
return modified_response
|
||||
|
||||
|
||||
async def async_post_call_failure_deployment_hook(
|
||||
|
|
@ -4707,7 +4718,8 @@ def get_optional_params(
|
|||
openai_params=list(DEFAULT_CHAT_COMPLETION_PARAM_VALUES.keys()),
|
||||
additional_drop_params=additional_drop_params,
|
||||
)
|
||||
print_verbose(f"Final returned optional params: {optional_params}")
|
||||
if _print_verbose_is_active():
|
||||
print_verbose(f"Final returned optional params: {redact_credentials_in_payload(optional_params)}")
|
||||
optional_params = _apply_openai_param_overrides(
|
||||
optional_params=optional_params,
|
||||
non_default_params=non_default_params,
|
||||
|
|
@ -7461,7 +7473,8 @@ def print_args_passed_to_litellm(original_function, args, kwargs):
|
|||
return
|
||||
|
||||
args_str: Final = ", ".join(map(repr, args))
|
||||
kwargs_str: Final = ", ".join(f"{key}={value!r}" for key, value in kwargs.items())
|
||||
redacted_kwargs: Final = redact_credentials_in_payload(kwargs)
|
||||
kwargs_str: Final = ", ".join(f"{key}={value!r}" for key, value in redacted_kwargs.items())
|
||||
print_verbose(
|
||||
"\n",
|
||||
) # new line before
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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]]
|
||||
|
|
|
|||
|
|
@ -1,6 +1,13 @@
|
|||
import time
|
||||
from collections.abc import Iterator
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
from openai import OpenAI, BadRequestError, NotFoundError, APIStatusError
|
||||
import pytest
|
||||
from openai import APIStatusError, BadRequestError, NotFoundError, OpenAI, Stream
|
||||
from openai.types.responses import ResponseStreamEvent
|
||||
|
||||
BACKGROUND_STREAM_ADMISSION_DEADLINE_SECONDS: Final = 90
|
||||
|
||||
|
||||
def generate_key():
|
||||
|
|
@ -153,43 +160,48 @@ def test_cancel_response():
|
|||
raise e
|
||||
|
||||
|
||||
def admitted_response_id(chunk: ResponseStreamEvent) -> str | None:
|
||||
response: Final = getattr(chunk, "response", None)
|
||||
return None if response is None else response.id
|
||||
|
||||
|
||||
def events_until_admission(stream: Stream[ResponseStreamEvent], started: float) -> Iterator[ResponseStreamEvent]:
|
||||
for chunk in stream:
|
||||
print("stream chunk=", chunk)
|
||||
yield chunk
|
||||
if admitted_response_id(chunk) is not None:
|
||||
return
|
||||
if time.monotonic() - started > BACKGROUND_STREAM_ADMISSION_DEADLINE_SECONDS:
|
||||
return
|
||||
|
||||
|
||||
def test_cancel_streaming_response():
|
||||
try:
|
||||
client = get_test_client()
|
||||
from litellm.types.llms.openai import ResponsesAPIResponse
|
||||
client: Final = get_test_client()
|
||||
started: Final = time.monotonic()
|
||||
stream: Final = client.responses.create(
|
||||
model="gpt-5.5",
|
||||
input="count from 1 to 500, one number per line",
|
||||
stream=True,
|
||||
background=True,
|
||||
timeout=BACKGROUND_STREAM_ADMISSION_DEADLINE_SECONDS,
|
||||
)
|
||||
|
||||
stream = client.responses.create(
|
||||
model="gpt-5.5",
|
||||
input="just respond with the word 'ping'",
|
||||
stream=True,
|
||||
background=True,
|
||||
with stream:
|
||||
events: Final = tuple(events_until_admission(stream, started))
|
||||
|
||||
elapsed: Final = time.monotonic() - started
|
||||
keepalive_events: Final = sum(1 for chunk in events if chunk.type == "keepalive")
|
||||
response_id: Final = next((rid for rid in map(admitted_response_id, events) if rid is not None), None)
|
||||
if response_id is None and keepalive_events:
|
||||
pytest.skip(
|
||||
f"OpenAI held the background stream in keepalive for {elapsed:.0f}s "
|
||||
f"({keepalive_events} keepalive events) without creating the response"
|
||||
)
|
||||
assert response_id is not None, f"no response event within {elapsed:.0f}s of streaming a background response"
|
||||
|
||||
collected_chunks = []
|
||||
response_id = None
|
||||
for chunk in stream:
|
||||
print("stream chunk=", chunk)
|
||||
collected_chunks.append(chunk)
|
||||
# Extract response ID from the first chunk that has it
|
||||
if (
|
||||
response_id is None
|
||||
and hasattr(chunk, "response")
|
||||
and hasattr(chunk.response, "id")
|
||||
):
|
||||
response_id = chunk.response.id
|
||||
|
||||
assert len(collected_chunks) > 0
|
||||
|
||||
# cancel the response if we got a response ID
|
||||
if response_id:
|
||||
cancel_response = client.responses.cancel(response_id)
|
||||
print("CANCEL streaming response=", cancel_response)
|
||||
assert hasattr(cancel_response, "id")
|
||||
except Exception as e:
|
||||
if "Cannot cancel a completed response" in str(e):
|
||||
pass
|
||||
else:
|
||||
raise e
|
||||
cancel_response: Final = client.responses.cancel(response_id)
|
||||
print("CANCEL streaming response=", cancel_response)
|
||||
assert cancel_response.status == "cancelled"
|
||||
|
||||
|
||||
def test_cancel_invalid_response_id():
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -336,3 +336,15 @@ class TestAzureResolvesTheDeclaredDefaultEffort:
|
|||
drop_params=True,
|
||||
)
|
||||
assert ("temperature" in mapped) is temperature_survives
|
||||
|
||||
|
||||
def test_azure_gpt_6_astra_takes_the_reasoning_series_request_shape():
|
||||
params = litellm.get_optional_params(
|
||||
model="gpt-6-astra",
|
||||
custom_llm_provider="azure",
|
||||
max_tokens=100,
|
||||
reasoning_effort="max",
|
||||
)
|
||||
assert params["max_completion_tokens"] == 100
|
||||
assert "max_tokens" not in params
|
||||
assert params["reasoning_effort"] == "max"
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -3270,3 +3270,90 @@ async def test_create_batch_async_validates_credentials_off_the_event_loop():
|
|||
validated_on = client.post.call_args.kwargs["headers"]["x-validated-on"]
|
||||
assert validated_on != str(threading.get_ident())
|
||||
assert client.post.call_args.kwargs["url"] == "https://batches.example/v1/messages/batches"
|
||||
|
||||
CONTAINER_NOT_FOUND_BODY = {
|
||||
"error": {
|
||||
"message": "Container with id 'cntr_gone' not found.",
|
||||
"type": "invalid_request_error",
|
||||
"param": None,
|
||||
"code": None,
|
||||
}
|
||||
}
|
||||
|
||||
INVALID_API_KEY_BODY = {
|
||||
"error": {
|
||||
"message": "Incorrect API key provided: sk-proj-***. You can find your API key at https://platform.openai.com/account/api-keys.",
|
||||
"type": "invalid_request_error",
|
||||
"param": None,
|
||||
"code": "invalid_api_key",
|
||||
},
|
||||
"status": 401,
|
||||
}
|
||||
|
||||
CONTAINER_LIST_BODY = {
|
||||
"object": "list",
|
||||
"data": [{"id": "cntr_a", "object": "container", "created_at": 1, "status": "running", "name": "a"}],
|
||||
"first_id": "cntr_a",
|
||||
"last_id": "cntr_a",
|
||||
"has_more": True,
|
||||
}
|
||||
|
||||
|
||||
def _container_sync_client(response: httpx.Response) -> HTTPHandler:
|
||||
client = HTTPHandler()
|
||||
client.client = httpx.Client(transport=httpx.MockTransport(lambda _request: response))
|
||||
return client
|
||||
|
||||
|
||||
def _container_async_client(response: httpx.Response) -> AsyncHTTPHandler:
|
||||
client = AsyncHTTPHandler()
|
||||
client.client = httpx.AsyncClient(transport=httpx.MockTransport(lambda _request: response))
|
||||
return client
|
||||
|
||||
|
||||
def test_container_retrieve_handler_raises_upstream_error_status_and_message():
|
||||
from litellm.llms.openai.containers.transformation import OpenAIContainerConfig
|
||||
|
||||
with pytest.raises(BaseLLMException) as exc_info:
|
||||
BaseLLMHTTPHandler().container_retrieve_handler(
|
||||
container_id="cntr_gone",
|
||||
container_provider_config=OpenAIContainerConfig(),
|
||||
litellm_params=GenericLiteLLMParams(api_key="sk-test"),
|
||||
logging_obj=Mock(),
|
||||
client=_container_sync_client(httpx.Response(404, json=CONTAINER_NOT_FOUND_BODY)),
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 404
|
||||
assert exc_info.value.message == "Container with id 'cntr_gone' not found."
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_container_list_handler_raises_upstream_error_status_and_message():
|
||||
from litellm.llms.openai.containers.transformation import OpenAIContainerConfig
|
||||
|
||||
with pytest.raises(BaseLLMException) as exc_info:
|
||||
await BaseLLMHTTPHandler().async_container_list_handler(
|
||||
container_provider_config=OpenAIContainerConfig(),
|
||||
litellm_params=GenericLiteLLMParams(api_key="sk-rejected"),
|
||||
logging_obj=Mock(),
|
||||
client=_container_async_client(httpx.Response(401, json=INVALID_API_KEY_BODY)),
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 401
|
||||
assert exc_info.value.message == INVALID_API_KEY_BODY["error"]["message"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_container_list_handler_transforms_success_response():
|
||||
from litellm.llms.openai.containers.transformation import OpenAIContainerConfig
|
||||
|
||||
response = await BaseLLMHTTPHandler().async_container_list_handler(
|
||||
container_provider_config=OpenAIContainerConfig(),
|
||||
litellm_params=GenericLiteLLMParams(api_key="sk-test"),
|
||||
logging_obj=Mock(),
|
||||
limit=1,
|
||||
client=_container_async_client(httpx.Response(200, json=CONTAINER_LIST_BODY)),
|
||||
)
|
||||
|
||||
assert [container.id for container in response.data] == ["cntr_a"]
|
||||
assert response.has_more is True
|
||||
|
|
|
|||
|
|
@ -1718,6 +1718,8 @@ class TestResponsesSurfaceSharesTheEffortRule:
|
|||
("gpt-5.6-sol", None, False),
|
||||
("gpt-5.6-terra", "none", True),
|
||||
("gpt-5.6-terra", "medium", False),
|
||||
("gpt-6-astra", None, False),
|
||||
("gpt-6-astra", "low", False),
|
||||
],
|
||||
)
|
||||
def test_temperature_follows_the_resolved_effort(
|
||||
|
|
|
|||
|
|
@ -1505,3 +1505,17 @@ class TestACatalogueOlderThanTheCodeDoesNotStripTemperature:
|
|||
drop_params=True,
|
||||
)
|
||||
assert "temperature" not in mapped
|
||||
|
||||
|
||||
def test_gpt_6_astra_takes_the_reasoning_series_request_shape():
|
||||
params = litellm.get_optional_params(
|
||||
model="gpt-6-astra",
|
||||
custom_llm_provider="openai",
|
||||
max_tokens=100,
|
||||
reasoning_effort="max",
|
||||
verbosity="low",
|
||||
)
|
||||
assert params["max_completion_tokens"] == 100
|
||||
assert "max_tokens" not in params
|
||||
assert params["reasoning_effort"] == "max"
|
||||
assert params["verbosity"] == "low"
|
||||
|
|
|
|||
|
|
@ -41,6 +41,8 @@ from litellm.llms.azure.chat.gpt_5_transformation import AzureOpenAIGPT5Config
|
|||
|
||||
# Models that MUST be classified as GPT-5 (routed through GPT-5 reasoning path)
|
||||
GPT5_MODELS = [
|
||||
"gpt-6-astra",
|
||||
"openai/gpt-6-astra",
|
||||
"gpt-5",
|
||||
"gpt-5.1",
|
||||
"gpt-5.2",
|
||||
|
|
@ -120,6 +122,8 @@ class TestOpenAIGPT5ConfigIsModelGpt5Model:
|
|||
# /v1/responses bridge (when reasoning_effort is set and tools are passed) on
|
||||
# is_model_gpt_5_4_plus_model, so the gpt-5.6 family must land on the True side.
|
||||
GPT5_4_PLUS_MODELS = [
|
||||
"gpt-6-astra",
|
||||
"openai/gpt-6-astra",
|
||||
"gpt-5.4",
|
||||
"gpt-5.5",
|
||||
"gpt-5.5-pro",
|
||||
|
|
|
|||
|
|
@ -964,6 +964,48 @@ def test_construct_target_url_with_version_prefix():
|
|||
assert str(target_url) == expected_url
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("requested_route", "expected_url"),
|
||||
[
|
||||
(
|
||||
"/projects/test-project/locations/global/publishers/anthropic/models/claude-sonnet-4-6:streamRawPredict",
|
||||
"https://aiplatform.googleapis.com/v1/projects/test-project/locations/global/publishers/anthropic/models/claude-sonnet-4-6:streamRawPredict",
|
||||
),
|
||||
(
|
||||
"/projects/test-project/locations/global/publishers/anthropic/models/count-tokens:rawPredict",
|
||||
"https://aiplatform.googleapis.com/v1/projects/test-project/locations/global/publishers/anthropic/models/count-tokens:rawPredict",
|
||||
),
|
||||
(
|
||||
"/projects/other-project/locations/us-east5/publishers/anthropic/models/claude-sonnet-4-6:rawPredict",
|
||||
"https://aiplatform.googleapis.com/v1/projects/test-project/locations/global/publishers/anthropic/models/claude-sonnet-4-6:rawPredict",
|
||||
),
|
||||
(
|
||||
"/projects/test-project/locations/global/cachedContents",
|
||||
"https://aiplatform.googleapis.com/v1beta1/projects/test-project/locations/global/cachedContents",
|
||||
),
|
||||
(
|
||||
"/v1/projects/test-project/locations/global/publishers/anthropic/models/claude-sonnet-4-6:streamRawPredict",
|
||||
"https://aiplatform.googleapis.com/v1/projects/test-project/locations/global/publishers/anthropic/models/claude-sonnet-4-6:streamRawPredict",
|
||||
),
|
||||
(
|
||||
"/v1beta1/projects/test-project/locations/global/cachedContents",
|
||||
"https://aiplatform.googleapis.com/v1beta1/projects/test-project/locations/global/cachedContents",
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_construct_target_url_versionless_project_route_gets_api_version(requested_route: str, expected_url: str) -> None:
|
||||
from litellm.llms.vertex_ai.common_utils import construct_target_url
|
||||
|
||||
target_url = construct_target_url(
|
||||
base_url="https://aiplatform.googleapis.com",
|
||||
requested_route=requested_route,
|
||||
vertex_project="test-project",
|
||||
vertex_location="global",
|
||||
)
|
||||
|
||||
assert str(target_url) == expected_url
|
||||
|
||||
|
||||
def test_fix_enum_types():
|
||||
"""
|
||||
Test _fix_enum_types function removes enum fields when type is not string.
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue