Merge remote-tracking branch 'upstream/main' into litellm_fix_leakage

This commit is contained in:
Harshit28j 2026-02-25 18:14:55 +05:30
commit e150547446
135 changed files with 10488 additions and 1111 deletions

View file

@ -49,7 +49,7 @@ USER root
# Install runtime dependencies (libsndfile needed for audio processing on ARM64)
RUN apk add --no-cache bash openssl tzdata nodejs npm python3 py3-pip libsndfile && \
npm install -g npm@latest tar@7.5.7 glob@11.1.0 @isaacs/brace-expansion@5.0.1 && \
npm install -g npm@latest tar@7.5.8 glob@11.1.0 @isaacs/brace-expansion@5.0.1 minimatch@10.2.1 diff@8.0.3 && \
# SECURITY FIX: npm bundles tar, glob, and brace-expansion at multiple nested
# levels inside its dependency tree. `npm install -g <pkg>` only creates a
# SEPARATE global package, it does NOT replace npm's internal copies.
@ -64,6 +64,12 @@ RUN apk add --no-cache bash openssl tzdata nodejs npm python3 py3-pip libsndfile
find "$GLOBAL/npm" -type d -name "brace-expansion" -path "*/node_modules/@isaacs/brace-expansion" | while read d; do \
rm -rf "$d" && cp -rL "$GLOBAL/@isaacs/brace-expansion" "$d"; \
done && \
find "$GLOBAL/npm" -type d -name "minimatch" -path "*/node_modules/minimatch" | while read d; do \
rm -rf "$d" && cp -rL "$GLOBAL/minimatch" "$d"; \
done && \
find "$GLOBAL/npm" -type d -name "diff" -path "*/node_modules/diff" | while read d; do \
rm -rf "$d" && cp -rL "$GLOBAL/diff" "$d"; \
done && \
npm cache clean --force
WORKDIR /app
@ -90,14 +96,20 @@ RUN find /usr/lib -type f -path "*/tornado/test/*" -delete && \
# npm with old vulnerable deps at /usr/lib/python3.*/site-packages/nodejs_wheel/.
# Patch every copy of tar, glob, and brace-expansion inside that tree.
RUN GLOBAL="$(npm root -g)" && \
find /usr/lib -path "*/nodejs_wheel/*/node_modules/tar" -type d | while read d; do \
find /usr/lib -type d -name "tar" -path "*/node_modules/tar" | while read d; do \
rm -rf "$d" && cp -rL "$GLOBAL/tar" "$d"; \
done && \
find /usr/lib -path "*/nodejs_wheel/*/node_modules/glob" -type d | while read d; do \
find /usr/lib -type d -name "glob" -path "*/node_modules/glob" | while read d; do \
rm -rf "$d" && cp -rL "$GLOBAL/glob" "$d"; \
done && \
find /usr/lib -path "*/nodejs_wheel/*/node_modules/@isaacs/brace-expansion" -type d | while read d; do \
find /usr/lib -type d -name "brace-expansion" -path "*/node_modules/@isaacs/brace-expansion" | while read d; do \
rm -rf "$d" && cp -rL "$GLOBAL/@isaacs/brace-expansion" "$d"; \
done && \
find /usr/lib -type d -name "minimatch" -path "*/node_modules/minimatch" | while read d; do \
rm -rf "$d" && cp -rL "$GLOBAL/minimatch" "$d"; \
done && \
find /usr/lib -type d -name "diff" -path "*/node_modules/diff" | while read d; do \
rm -rf "$d" && cp -rL "$GLOBAL/diff" "$d"; \
done
# Install semantic_router and aurelio-sdk using script

View file

@ -203,7 +203,7 @@ curl -X POST 'http://0.0.0.0:4000/v1/chat/completions' \
{
"mcpServers": {
"LiteLLM": {
"url": "http://localhost:4000/mcp",
"url": "http://localhost:4000/mcp/",
"headers": {
"x-litellm-api-key": "Bearer sk-1234"
}

View file

@ -160,6 +160,7 @@ run_grype_scans() {
"CVE-2026-0775" # npm cli incorrect permission assignment - no fix available yet, npm is only used at build/prisma-generate time
"GHSA-3ppc-4f35-3m26" # minimatch ReDoS via repeated wildcards - from nodejs_wheel bundled npm, not used in application runtime code
"GHSA-83g3-92jg-28cx" # tar arbitrary file read/write via hardlink - from nodejs_wheel bundled npm, not used in application runtime code
"CVE-2026-25639" # axios - full fix requires 1.x major version bump; pinned to >=0.30.2 to clear other axios CVEs, upgrade to 1.x in follow-up
)
# Build JSON array of allowlisted CVE IDs for jq

View file

@ -36,6 +36,10 @@ If `db.useStackgresOperator` is used (not yet implemented):
| `serviceAccount.create` | Whether or not to create a Kubernetes Service Account for this deployment. The default is `false` because LiteLLM has no need to access the Kubernetes API. | `false` |
| `service.type` | Kubernetes Service type (e.g. `LoadBalancer`, `ClusterIP`, etc.) | `ClusterIP` |
| `service.port` | TCP port that the Kubernetes Service will listen on. Also the TCP port within the Pod that the proxy will listen on. | `4000` |
| `livenessProbe.*` | Liveness probe settings for the LiteLLM container (`path`, `periodSeconds`, `timeoutSeconds`, thresholds, and initial delay). | See `values.yaml` |
| `readinessProbe.*` | Readiness probe settings for the LiteLLM container (`path`, `periodSeconds`, `timeoutSeconds`, thresholds, and initial delay). | See `values.yaml` |
| `startupProbe.*` | Startup probe settings for the LiteLLM container (`path`, `periodSeconds`, `timeoutSeconds`, thresholds, and initial delay). | See `values.yaml` |
| `resources.*` | CPU/memory requests and limits for the LiteLLM container. | `{}` |
| `service.loadBalancerClass` | Optional LoadBalancer implementation class (only used when `service.type` is `LoadBalancer`) | `""` |
| `ingress.labels` | Additional labels for the Ingress resource | `{}` |
| `ingress.*` | See [values.yaml](./values.yaml) for example settings | N/A |

View file

@ -6,4 +6,4 @@ metadata:
data:
config.yaml: |
{{ .Values.proxy_config | toYaml | indent 6 }}
{{- end }}
{{- end }}

View file

@ -158,18 +158,31 @@ spec:
{{- end }}
livenessProbe:
httpGet:
path: /health/liveliness
path: {{ .Values.livenessProbe.path | quote }}
port: {{ if .Values.separateHealthApp }}"health"{{ else }}"http"{{ end }}
initialDelaySeconds: {{ .Values.livenessProbe.initialDelaySeconds }}
periodSeconds: {{ .Values.livenessProbe.periodSeconds }}
timeoutSeconds: {{ .Values.livenessProbe.timeoutSeconds }}
successThreshold: {{ .Values.livenessProbe.successThreshold }}
failureThreshold: {{ .Values.livenessProbe.failureThreshold }}
readinessProbe:
httpGet:
path: /health/readiness
path: {{ .Values.readinessProbe.path | quote }}
port: {{ if .Values.separateHealthApp }}"health"{{ else }}"http"{{ end }}
initialDelaySeconds: {{ .Values.readinessProbe.initialDelaySeconds }}
periodSeconds: {{ .Values.readinessProbe.periodSeconds }}
timeoutSeconds: {{ .Values.readinessProbe.timeoutSeconds }}
successThreshold: {{ .Values.readinessProbe.successThreshold }}
failureThreshold: {{ .Values.readinessProbe.failureThreshold }}
startupProbe:
httpGet:
path: /health/readiness
path: {{ .Values.startupProbe.path | quote }}
port: {{ if .Values.separateHealthApp }}"health"{{ else }}"http"{{ end }}
failureThreshold: 30
periodSeconds: 10
initialDelaySeconds: {{ .Values.startupProbe.initialDelaySeconds }}
periodSeconds: {{ .Values.startupProbe.periodSeconds }}
timeoutSeconds: {{ .Values.startupProbe.timeoutSeconds }}
successThreshold: {{ .Values.startupProbe.successThreshold }}
failureThreshold: {{ .Values.startupProbe.failureThreshold }}
resources:
{{- toYaml .Values.resources | nindent 12 }}
volumeMounts:
@ -235,4 +248,4 @@ spec:
{{- if .Values.topologySpreadConstraints }}
topologySpreadConstraints:
{{- toYaml .Values.topologySpreadConstraints | nindent 8 }}
{{- end }}
{{- end }}

View file

@ -159,4 +159,150 @@ tests:
value: -c
- equal:
path: spec.template.spec.containers[0].lifecycle.preStop.exec.command[2]
value: echo "Container stopping"
value: echo "Container stopping"
- it: should render background health check settings from proxy_config.general_settings
template: configmap-litellm.yaml
set:
proxy_config.general_settings.background_health_checks: true
proxy_config.general_settings.health_check_interval: 240
proxy_config.general_settings.health_check_concurrency: 16
proxy_config.general_settings.health_check_details: false
asserts:
- matchRegex:
path: data["config.yaml"]
pattern: '(?m)^\s*background_health_checks:\s*true$'
- matchRegex:
path: data["config.yaml"]
pattern: '(?m)^\s*health_check_interval:\s*240$'
- matchRegex:
path: data["config.yaml"]
pattern: '(?m)^\s*health_check_concurrency:\s*16$'
- matchRegex:
path: data["config.yaml"]
pattern: '(?m)^\s*health_check_details:\s*false$'
- it: should allow overriding liveness, readiness, and startup probes
template: deployment.yaml
set:
livenessProbe:
path: /custom/livez
initialDelaySeconds: 5
periodSeconds: 15
timeoutSeconds: 5
successThreshold: 1
failureThreshold: 5
readinessProbe:
path: /custom/readyz
initialDelaySeconds: 10
periodSeconds: 20
timeoutSeconds: 6
successThreshold: 1
failureThreshold: 6
startupProbe:
path: /custom/startupz
initialDelaySeconds: 15
periodSeconds: 25
timeoutSeconds: 7
successThreshold: 1
failureThreshold: 40
asserts:
- equal:
path: spec.template.spec.containers[0].livenessProbe.httpGet.path
value: /custom/livez
- equal:
path: spec.template.spec.containers[0].livenessProbe.timeoutSeconds
value: 5
- equal:
path: spec.template.spec.containers[0].readinessProbe.httpGet.path
value: /custom/readyz
- equal:
path: spec.template.spec.containers[0].readinessProbe.timeoutSeconds
value: 6
- equal:
path: spec.template.spec.containers[0].startupProbe.httpGet.path
value: /custom/startupz
- equal:
path: spec.template.spec.containers[0].startupProbe.failureThreshold
value: 40
- it: should render container resources from values
template: deployment.yaml
set:
resources:
limits:
cpu: 500m
memory: 2Gi
requests:
cpu: 250m
memory: 1Gi
asserts:
- equal:
path: spec.template.spec.containers[0].resources.limits.cpu
value: 500m
- equal:
path: spec.template.spec.containers[0].resources.limits.memory
value: 2Gi
- equal:
path: spec.template.spec.containers[0].resources.requests.cpu
value: 250m
- equal:
path: spec.template.spec.containers[0].resources.requests.memory
value: 1Gi
- it: should keep default probes and empty resources unchanged
template: deployment.yaml
asserts:
- equal:
path: spec.template.spec.containers[0].livenessProbe.httpGet.path
value: /health/liveliness
- equal:
path: spec.template.spec.containers[0].livenessProbe.initialDelaySeconds
value: 0
- equal:
path: spec.template.spec.containers[0].livenessProbe.periodSeconds
value: 10
- equal:
path: spec.template.spec.containers[0].livenessProbe.timeoutSeconds
value: 1
- equal:
path: spec.template.spec.containers[0].livenessProbe.successThreshold
value: 1
- equal:
path: spec.template.spec.containers[0].livenessProbe.failureThreshold
value: 3
- equal:
path: spec.template.spec.containers[0].readinessProbe.httpGet.path
value: /health/readiness
- equal:
path: spec.template.spec.containers[0].readinessProbe.initialDelaySeconds
value: 0
- equal:
path: spec.template.spec.containers[0].readinessProbe.periodSeconds
value: 10
- equal:
path: spec.template.spec.containers[0].readinessProbe.timeoutSeconds
value: 1
- equal:
path: spec.template.spec.containers[0].readinessProbe.successThreshold
value: 1
- equal:
path: spec.template.spec.containers[0].readinessProbe.failureThreshold
value: 3
- equal:
path: spec.template.spec.containers[0].startupProbe.httpGet.path
value: /health/readiness
- equal:
path: spec.template.spec.containers[0].startupProbe.initialDelaySeconds
value: 0
- equal:
path: spec.template.spec.containers[0].startupProbe.periodSeconds
value: 10
- equal:
path: spec.template.spec.containers[0].startupProbe.timeoutSeconds
value: 1
- equal:
path: spec.template.spec.containers[0].startupProbe.successThreshold
value: 1
- equal:
path: spec.template.spec.containers[0].startupProbe.failureThreshold
value: 30
- equal:
path: spec.template.spec.containers[0].resources
value: {}

View file

@ -84,6 +84,31 @@ service:
separateHealthApp: false
separateHealthPort: 8081
# Probe tuning for proxy container
livenessProbe:
path: /health/liveliness
initialDelaySeconds: 0
periodSeconds: 10
timeoutSeconds: 1
successThreshold: 1
failureThreshold: 3
readinessProbe:
path: /health/readiness
initialDelaySeconds: 0
periodSeconds: 10
timeoutSeconds: 1
successThreshold: 1
failureThreshold: 3
startupProbe:
path: /health/readiness
initialDelaySeconds: 0
periodSeconds: 10
timeoutSeconds: 1
successThreshold: 1
failureThreshold: 30
ingress:
enabled: false
className: "nginx"

View file

@ -5,8 +5,21 @@ FROM ghcr.io/berriai/litellm:litellm_fwd_server_root_path-dev
WORKDIR /app
# Install Node.js and npm (adjust version as needed)
RUN apt-get update && apt-get install -y nodejs npm && \
npm install -g npm@latest tar@7.5.7 glob@11.1.0 @isaacs/brace-expansion@5.0.1 && \
RUN apt-get update && apt-get upgrade -y \
libxml2 \
libexpat1 \
openssl \
libssl3 \
git \
libkrb5-3 \
libglib2.0-0 \
wget \
libaom3 \
libxslt1.1 \
libgnutls30 \
libc6 && \
apt-get install -y nodejs npm && \
npm install -g npm@latest tar@7.5.8 glob@11.1.0 @isaacs/brace-expansion@5.0.1 minimatch@10.2.1 diff@8.0.3 && \
GLOBAL="$(npm root -g)" && \
find "$GLOBAL/npm" -type d -name "tar" -path "*/node_modules/tar" | while read d; do \
rm -rf "$d" && cp -rL "$GLOBAL/tar" "$d"; \
@ -17,6 +30,12 @@ RUN apt-get update && apt-get install -y nodejs npm && \
find "$GLOBAL/npm" -type d -name "brace-expansion" -path "*/node_modules/@isaacs/brace-expansion" | while read d; do \
rm -rf "$d" && cp -rL "$GLOBAL/@isaacs/brace-expansion" "$d"; \
done && \
find "$GLOBAL/npm" -type d -name "minimatch" -path "*/node_modules/minimatch" | while read d; do \
rm -rf "$d" && cp -rL "$GLOBAL/minimatch" "$d"; \
done && \
find "$GLOBAL/npm" -type d -name "diff" -path "*/node_modules/diff" | while read d; do \
rm -rf "$d" && cp -rL "$GLOBAL/diff" "$d"; \
done && \
npm cache clean --force
# Copy the UI source into the container

View file

@ -50,7 +50,7 @@ USER root
# Install runtime dependencies
RUN apk add --no-cache bash openssl tzdata nodejs npm python3 py3-pip libsndfile && \
npm install -g npm@latest tar@7.5.7 glob@11.1.0 @isaacs/brace-expansion@5.0.1 && \
npm install -g npm@latest tar@7.5.8 glob@11.1.0 @isaacs/brace-expansion@5.0.1 minimatch@10.2.1 diff@8.0.3 && \
GLOBAL="$(npm root -g)" && \
find "$GLOBAL/npm" -type d -name "tar" -path "*/node_modules/tar" | while read d; do \
rm -rf "$d" && cp -rL "$GLOBAL/tar" "$d"; \
@ -61,6 +61,12 @@ RUN apk add --no-cache bash openssl tzdata nodejs npm python3 py3-pip libsndfile
find "$GLOBAL/npm" -type d -name "brace-expansion" -path "*/node_modules/@isaacs/brace-expansion" | while read d; do \
rm -rf "$d" && cp -rL "$GLOBAL/@isaacs/brace-expansion" "$d"; \
done && \
find "$GLOBAL/npm" -type d -name "minimatch" -path "*/node_modules/minimatch" | while read d; do \
rm -rf "$d" && cp -rL "$GLOBAL/minimatch" "$d"; \
done && \
find "$GLOBAL/npm" -type d -name "diff" -path "*/node_modules/diff" | while read d; do \
rm -rf "$d" && cp -rL "$GLOBAL/diff" "$d"; \
done && \
npm cache clean --force
WORKDIR /app
@ -79,14 +85,20 @@ RUN pip install *.whl /wheels/* --no-index --find-links=/wheels/ && rm -f *.whl
# npm with old vulnerable deps at /usr/lib/python3.*/site-packages/nodejs_wheel/.
# Patch every copy of tar, glob, and brace-expansion inside that tree.
RUN GLOBAL="$(npm root -g)" && \
find /usr/lib -path "*/nodejs_wheel/*/node_modules/tar" -type d | while read d; do \
find /usr/lib -type d -name "tar" -path "*/node_modules/tar" | while read d; do \
rm -rf "$d" && cp -rL "$GLOBAL/tar" "$d"; \
done && \
find /usr/lib -path "*/nodejs_wheel/*/node_modules/glob" -type d | while read d; do \
find /usr/lib -type d -name "glob" -path "*/node_modules/glob" | while read d; do \
rm -rf "$d" && cp -rL "$GLOBAL/glob" "$d"; \
done && \
find /usr/lib -path "*/nodejs_wheel/*/node_modules/@isaacs/brace-expansion" -type d | while read d; do \
find /usr/lib -type d -name "brace-expansion" -path "*/node_modules/@isaacs/brace-expansion" | while read d; do \
rm -rf "$d" && cp -rL "$GLOBAL/@isaacs/brace-expansion" "$d"; \
done && \
find /usr/lib -type d -name "minimatch" -path "*/node_modules/minimatch" | while read d; do \
rm -rf "$d" && cp -rL "$GLOBAL/minimatch" "$d"; \
done && \
find /usr/lib -type d -name "diff" -path "*/node_modules/diff" | while read d; do \
rm -rf "$d" && cp -rL "$GLOBAL/diff" "$d"; \
done
# Install semantic_router and aurelio-sdk using script

View file

@ -56,13 +56,26 @@ FROM $LITELLM_RUNTIME_IMAGE AS runtime
USER root
# Install only runtime dependencies
RUN apt-get update && apt-get install -y --no-install-recommends \
libssl3 \
RUN apt-get update && apt-get upgrade -y \
libxml2 \
libexpat1 \
openssl \
libssl3 \
git \
libkrb5-3 \
libglib2.0-0 \
wget \
libaom3 \
libxslt1.1 \
libgnutls30 \
libc6 \
&& apt-get install -y --no-install-recommends \
libssl3 \
libatomic1 \
nodejs \
npm \
&& rm -rf /var/lib/apt/lists/* \
&& npm install -g npm@latest tar@7.5.7 glob@11.1.0 @isaacs/brace-expansion@5.0.1 \
&& npm install -g npm@latest tar@7.5.8 glob@11.1.0 @isaacs/brace-expansion@5.0.1 minimatch@10.2.1 diff@8.0.3 \
&& GLOBAL="$(npm root -g)" \
&& find "$GLOBAL/npm" -type d -name "tar" -path "*/node_modules/tar" | while read d; do \
rm -rf "$d" && cp -rL "$GLOBAL/tar" "$d"; \
@ -73,6 +86,12 @@ RUN apt-get update && apt-get install -y --no-install-recommends \
&& find "$GLOBAL/npm" -type d -name "brace-expansion" -path "*/node_modules/@isaacs/brace-expansion" | while read d; do \
rm -rf "$d" && cp -rL "$GLOBAL/@isaacs/brace-expansion" "$d"; \
done \
&& find "$GLOBAL/npm" -type d -name "minimatch" -path "*/node_modules/minimatch" | while read d; do \
rm -rf "$d" && cp -rL "$GLOBAL/minimatch" "$d"; \
done \
&& find "$GLOBAL/npm" -type d -name "diff" -path "*/node_modules/diff" | while read d; do \
rm -rf "$d" && cp -rL "$GLOBAL/diff" "$d"; \
done \
&& npm cache clean --force
WORKDIR /app
@ -95,14 +114,20 @@ RUN pip install --no-cache-dir *.whl /wheels/* --no-index --find-links=/wheels/
# npm with old vulnerable deps at /usr/lib/python3.*/site-packages/nodejs_wheel/.
# Patch every copy of tar, glob, and brace-expansion inside that tree.
RUN GLOBAL="$(npm root -g)" && \
find /usr/lib -path "*/nodejs_wheel/*/node_modules/tar" -type d | while read d; do \
find /usr/lib -type d -name "tar" -path "*/node_modules/tar" | while read d; do \
rm -rf "$d" && cp -rL "$GLOBAL/tar" "$d"; \
done && \
find /usr/lib -path "*/nodejs_wheel/*/node_modules/glob" -type d | while read d; do \
find /usr/lib -type d -name "glob" -path "*/node_modules/glob" | while read d; do \
rm -rf "$d" && cp -rL "$GLOBAL/glob" "$d"; \
done && \
find /usr/lib -path "*/nodejs_wheel/*/node_modules/@isaacs/brace-expansion" -type d | while read d; do \
find /usr/lib -type d -name "brace-expansion" -path "*/node_modules/@isaacs/brace-expansion" | while read d; do \
rm -rf "$d" && cp -rL "$GLOBAL/@isaacs/brace-expansion" "$d"; \
done && \
find /usr/lib -type d -name "minimatch" -path "*/node_modules/minimatch" | while read d; do \
rm -rf "$d" && cp -rL "$GLOBAL/minimatch" "$d"; \
done && \
find /usr/lib -type d -name "diff" -path "*/node_modules/diff" | while read d; do \
rm -rf "$d" && cp -rL "$GLOBAL/diff" "$d"; \
done
# Generate prisma client and set permissions

View file

@ -80,7 +80,7 @@ ENV PRISMA_BINARY_CACHE_DIR=/app/.cache/prisma-python/binaries \
XDG_CACHE_HOME=/app/.cache \
PATH="/usr/lib/python3.13/site-packages/nodejs/bin:${PATH}"
RUN pip install --no-cache-dir prisma==0.11.0 nodejs-wheel-binaries==24.12.0 \
RUN pip install --no-cache-dir prisma==0.11.0 nodejs-wheel-binaries==24.13.1 \
&& mkdir -p /app/.cache/npm
RUN NPM_CONFIG_CACHE=/app/.cache/npm \
@ -105,7 +105,8 @@ RUN for i in 1 2 3; do \
&& for i in 1 2 3; do \
apk add --no-cache python3 py3-pip bash openssl tzdata nodejs npm supervisor && break || sleep 5; \
done \
&& npm install -g npm@latest tar@7.5.7 glob@11.1.0 @isaacs/brace-expansion@5.0.1 \
&& apk upgrade --no-cache nodejs \
&& npm install -g npm@latest tar@7.5.8 glob@11.1.0 @isaacs/brace-expansion@5.0.1 minimatch@10.2.1 diff@8.0.3 \
&& GLOBAL="$(npm root -g)" \
&& find "$GLOBAL/npm" -type d -name "tar" -path "*/node_modules/tar" | while read d; do \
rm -rf "$d" && cp -rL "$GLOBAL/tar" "$d"; \
@ -116,6 +117,12 @@ RUN for i in 1 2 3; do \
&& find "$GLOBAL/npm" -type d -name "brace-expansion" -path "*/node_modules/@isaacs/brace-expansion" | while read d; do \
rm -rf "$d" && cp -rL "$GLOBAL/@isaacs/brace-expansion" "$d"; \
done \
&& find "$GLOBAL/npm" -type d -name "minimatch" -path "*/node_modules/minimatch" | while read d; do \
rm -rf "$d" && cp -rL "$GLOBAL/minimatch" "$d"; \
done \
&& find "$GLOBAL/npm" -type d -name "diff" -path "*/node_modules/diff" | while read d; do \
rm -rf "$d" && cp -rL "$GLOBAL/diff" "$d"; \
done \
&& npm cache clean --force
# Copy artifacts from builder
@ -162,14 +169,20 @@ RUN pip install --no-index --find-links=/wheels/ -r requirements.txt && \
# npm with old vulnerable deps at /usr/lib/python3.*/site-packages/nodejs_wheel/.
# Patch every copy of tar, glob, and brace-expansion inside that tree.
RUN GLOBAL="$(npm root -g)" && \
find /usr/lib -path "*/nodejs_wheel/*/node_modules/tar" -type d | while read d; do \
find /usr/lib -type d -name "tar" -path "*/node_modules/tar" | while read d; do \
rm -rf "$d" && cp -rL "$GLOBAL/tar" "$d"; \
done && \
find /usr/lib -path "*/nodejs_wheel/*/node_modules/glob" -type d | while read d; do \
find /usr/lib -type d -name "glob" -path "*/node_modules/glob" | while read d; do \
rm -rf "$d" && cp -rL "$GLOBAL/glob" "$d"; \
done && \
find /usr/lib -path "*/nodejs_wheel/*/node_modules/@isaacs/brace-expansion" -type d | while read d; do \
find /usr/lib -type d -name "brace-expansion" -path "*/node_modules/@isaacs/brace-expansion" | while read d; do \
rm -rf "$d" && cp -rL "$GLOBAL/@isaacs/brace-expansion" "$d"; \
done && \
find /usr/lib -type d -name "minimatch" -path "*/node_modules/minimatch" | while read d; do \
rm -rf "$d" && cp -rL "$GLOBAL/minimatch" "$d"; \
done && \
find /usr/lib -type d -name "diff" -path "*/node_modules/diff" | while read d; do \
rm -rf "$d" && cp -rL "$GLOBAL/diff" "$d"; \
done
# Permissions, cleanup, and Prisma prep

View file

@ -0,0 +1,145 @@
---
slug: gpt_5_3_codex
title: "Day 0 Support: GPT-5.3-Codex"
date: 2026-02-24T10:00:00
authors:
- name: Sameer Kankute
title: SWE @ LiteLLM (LLM Translation)
url: https://www.linkedin.com/in/sameer-kankute/
image_url: https://pbs.twimg.com/profile_images/2001352686994907136/ONgNuSk5_400x400.jpg
- name: Krrish Dholakia
title: "CEO, LiteLLM"
url: https://www.linkedin.com/in/krish-d/
image_url: https://pbs.twimg.com/profile_images/1298587542745358340/DZv3Oj-h_400x400.jpg
- name: Ishaan Jaff
title: "CTO, LiteLLM"
url: https://www.linkedin.com/in/reffajnaahsi/
image_url: https://pbs.twimg.com/profile_images/1613813310264340481/lz54oEiB_400x400.jpg
description: "Day 0 support for GPT-5.3-Codex on LiteLLM, including phase parameter handling for Responses API."
tags: [openai, gpt-5.3-codex, codex, day 0 support]
hide_table_of_contents: false
---
import Tabs from '@theme/Tabs';
import TabItem from '@theme/TabItem';
LiteLLM now supports GPT-5.3-Codex on Day 0, including support for the new assistant `phase` metadata on Responses API output items.
## Why `phase` matters for GPT-5.3-Codex
`phase` appears on assistant output items and helps distinguish preamble/commentary turns from final closeout responses.
Reference: [Phase parameter docs](https://developers.openai.com/api/reference/overview)
Supported values:
- `null`
- `"commentary"`
- `"final_answer"`
Important:
- Persist assistant output items with `phase` exactly as returned.
- Send those assistant items back on the next turn.
- Do **not** add `phase` to user messages.
## Docker Image
```bash
docker pull ghcr.io/berriai/litellm:v1.81.12-stable.gpt-5.3
```
## Usage
<Tabs>
<TabItem value="proxy" label="LiteLLM Proxy">
**1. Setup config.yaml**
```yaml
model_list:
- model_name: gpt-5.3-codex
litellm_params:
model: openai/gpt-5.3-codex
```
**2. Start the proxy**
```bash
docker run -d \
-p 4000:4000 \
-e ANTHROPIC_API_KEY=$OPENAI_API_KEY \
-v $(pwd)/config.yaml:/app/config.yaml \
ghcr.io/berriai/litellm:v1.81.12-stable.gpt-5.3 \
--config /app/config.yaml
```
**3. Test it**
```bash
curl -X POST "http://0.0.0.0:4000/v1/responses" \
-H "Content-Type: application/json" \
-H "Authorization: Bearer $LITELLM_KEY" \
-d '{
"model": "gpt-5.3-codex",
"input": "Write a Python script that checks if a number is prime."
}'
```
</TabItem>
</Tabs>
## Python Example: Persist `phase` with OpenAI Client + LiteLLM Base URL
```python
from openai import OpenAI
client = OpenAI(
base_url="http://0.0.0.0:4000/v1", # LiteLLM Proxy
api_key="your-litellm-api-key",
)
items = [] # Persist this per conversation/thread
def _item_get(item, key, default=None):
if isinstance(item, dict):
return item.get(key, default)
return getattr(item, key, default)
def run_turn(user_text: str):
global items
# User message: no phase field
items.append(
{
"type": "message",
"role": "user",
"content": [{"type": "input_text", "text": user_text}],
}
)
resp = client.responses.create(
model="gpt-5.3-codex",
input=items,
)
# Persist assistant output items verbatim, including phase
for out_item in (resp.output or []):
items.append(out_item)
# Optional: inspect latest phase for UI/telemetry routing
latest_phase = None
for out_item in reversed(resp.output or []):
if _item_get(out_item, "type") == "output_item.done" and _item_get(out_item, "phase") is not None:
latest_phase = _item_get(out_item, "phase")
break
return resp, latest_phase
```
## Notes
- Use `/v1/responses` for GPT Codex models.
- Preserve full assistant output history for best multi-turn behavior.
- If `phase` metadata is dropped during history reconstruction, output quality can degrade on long-running tasks.

View file

@ -63,7 +63,6 @@ for _ in range(2):
}
],
},
# marked for caching with the cache_control parameter, so that this checkpoint can read from the previous cache.
{
"role": "user",
"content": [
@ -77,7 +76,6 @@ for _ in range(2):
"role": "assistant",
"content": "Certainly! the key terms and conditions are the following: the contract is 1 year long for $10/mo",
},
# The final turn is marked with cache-control, for continuing in followups.
{
"role": "user",
"content": [
@ -112,16 +110,16 @@ model_list:
api_key: os.environ/OPENAI_API_KEY
```
2. Start proxy
2. Start proxy
```bash
litellm --config /path/to/config.yaml
```
3. Test it!
3. Test it!
```python
from openai import OpenAI
from openai import OpenAI
import os
client = OpenAI(
@ -144,7 +142,6 @@ for _ in range(2):
}
],
},
# marked for caching with the cache_control parameter, so that this checkpoint can read from the previous cache.
{
"role": "user",
"content": [
@ -158,7 +155,6 @@ for _ in range(2):
"role": "assistant",
"content": "Certainly! the key terms and conditions are the following: the contract is 1 year long for $10/mo",
},
# The final turn is marked with cache-control, for continuing in followups.
{
"role": "user",
"content": [
@ -183,6 +179,78 @@ assert response.usage.prompt_tokens_details.cached_tokens > 0
</TabItem>
</Tabs>
### OpenAI `prompt_cache_key` and `prompt_cache_retention`
OpenAI prompt caching is [**automatic**](https://platform.openai.com/docs/guides/prompt-caching) — no `cache_control` message annotations are needed. Any request with 1024+ prompt tokens is eligible for caching.
OpenAI also supports two optional parameters for more control over caching behavior:
- **`prompt_cache_key`** (string) — A routing hint that improves cache hit rates for requests sharing long common prefixes. Requests with the same cache key are routed to the same backend, increasing the likelihood of a cache hit.
- **`prompt_cache_retention`** (`"in_memory"` or `"24h"`) — Controls cache TTL. Default is `"in_memory"` (5–10 min). Set to `"24h"` for extended caching that offloads KV tensors to GPU-local storage.
<Tabs>
<TabItem value="sdk" label="SDK">
```python
from litellm import completion
import os
os.environ["OPENAI_API_KEY"] = ""
response = completion(
model="gpt-4o",
messages=[
{
"role": "system",
"content": "You are an AI assistant tasked with analyzing legal documents. "
+ "Here is the full text of a complex legal agreement " * 400,
},
{
"role": "user",
"content": "What are the key terms and conditions?",
},
],
prompt_cache_key="legal-doc-analysis",
prompt_cache_retention="24h",
)
print(response.usage)
```
</TabItem>
<TabItem value="proxy" label="PROXY">
```python
from openai import OpenAI
client = OpenAI(
api_key="LITELLM_PROXY_KEY",
base_url="LITELLM_PROXY_BASE",
)
response = client.chat.completions.create(
model="gpt-4o",
messages=[
{
"role": "system",
"content": "You are an AI assistant tasked with analyzing legal documents. "
+ "Here is the full text of a complex legal agreement " * 400,
},
{
"role": "user",
"content": "What are the key terms and conditions?",
},
],
extra_body={
"prompt_cache_key": "legal-doc-analysis",
"prompt_cache_retention": "24h",
},
)
print(response.usage)
```
</TabItem>
</Tabs>
### Anthropic Example
Anthropic charges for cache writes.

View file

@ -15,6 +15,7 @@ Use LiteLLM to call Google AI's generateContent endpoints for text generation, m
| Streaming | ✅ | |
| Fallbacks | ✅ | between supported models |
| Loadbalancing | ✅ | between supported models |
| Metadata Tracking | ✅ | passes trace ID, metadata to observability callbacks (e.g. S3, Langfuse) |
## Usage
---

View file

@ -641,7 +641,7 @@ import asyncio
config = {
"mcpServers": {
"mcp_group": {
"url": "http://localhost:4000/mcp",
"url": "http://localhost:4000/mcp/",
"headers": {
"x-mcp-servers": "dev_group", # assume this gives access to github, zapier and deepwiki
"x-litellm-api-key": "Bearer sk-1234",

View file

@ -122,7 +122,7 @@ Use this to track overall LiteLLM Proxy usage.
| Metric Name | Description |
|----------------------|--------------------------------------|
| `litellm_proxy_failed_requests_metric` | Total number of failed responses from proxy - the client did not get a success response from litellm proxy. Labels: `"end_user", "hashed_api_key", "api_key_alias", "requested_model", "team", "team_alias", "user", "user_email", "exception_status", "exception_class", "route", "model_id"` |
| `litellm_proxy_total_requests_metric` | Total number of requests made to the proxy server - track number of client side requests. Labels: `"end_user", "hashed_api_key", "api_key_alias", "requested_model", "team", "team_alias", "user", "status_code", "user_email", "route", "model_id"` |
| `litellm_proxy_total_requests_metric` | Total number of requests made to the proxy server - track number of client side requests. Labels: `"end_user", "hashed_api_key", "api_key_alias", "requested_model", "team", "team_alias", "user", "status_code", "user_email", "route", "model_id"`. Optionally includes `"stream"` — see [Emit Stream Label](#emit-stream-label). |
### Callback Logging Metrics
@ -214,9 +214,31 @@ litellm_settings:
```
### Emit Stream Label
Add a `stream` label to `litellm_proxy_total_requests_metric` to split requests by streaming vs. non-streaming. Disabled by default.
```yaml title="config.yaml"
litellm_settings:
callbacks: ["prometheus"]
prometheus_emit_stream_label: true
```
When enabled, `litellm_proxy_total_requests_metric` gains a `stream` label with values `"True"`, `"False"`, or `"None"`.
```
litellm_proxy_total_requests_metric{..., stream="True"} 42
litellm_proxy_total_requests_metric{..., stream="False"} 100
```
:::note
This label is opt-in because adding a new label to an existing metric changes its cardinality and breaks existing Prometheus queries / Grafana dashboards that target this metric. Enable it only on fresh deployments or when you are ready to update your dashboards.
:::
## [BETA] Custom Metrics
Track custom metrics on prometheus on all events mentioned above.
Track custom metrics on prometheus on all events mentioned above.
### Custom Metadata Labels

View file

@ -592,6 +592,21 @@ def test_pii_masking_allows_normal_text():
## Part 7: Troubleshooting
### Issue: Guardrail failure: non-JSON response from Presidio
**Symptom:** You receive an error indicating `expected application/json Content-Type but received text/html` or similar.
**Root cause:** Your ingress controller or reverse proxy might be routing the `/analyze` or `/anonymize` POST request to a health endpoint (like `/health` or `/presidio-analyzer/health`) which returns plain text instead of JSON.
**Fix:** Ensure your `PRESIDIO_ANALYZER_API_BASE` and `PRESIDIO_ANONYMIZER_API_BASE` are correctly pointing directly to the Presidio API endpoints, or that your ingress routes the path correctly without stripping it and inadvertently forwarding to a plain-text health check endpoint.
**Verification:** You can verify your endpoints using `curl`. It should return a JSON array, not `text/html`:
```bash
curl -sv -X POST http://your-analyzer-endpoint/analyze \
-H "Content-Type: application/json" \
-d '{"text":"test","language":"en"}'
```
### Issue: Presidio Not Detecting PII
**Check 1: Language Configuration**

View file

@ -62,6 +62,8 @@
"gray-matter": "4.0.3",
"glob": ">=11.1.0",
"tar": ">=7.5.8",
"minimatch": ">=10.2.1",
"diff": ">=8.0.3",
"@isaacs/brace-expansion": ">=5.0.1",
"node-forge": ">=1.3.2",
"mdast-util-to-hast": ">=13.2.1",
@ -81,6 +83,15 @@
"url-loader": {
"ajv": "6.14.0"
},
"minimatch": "10.2.1"
"@babel/traverse": ">=7.23.2",
"ws": ">=7.5.10",
"http-proxy-middleware": ">=2.0.9",
"tar-fs": ">=2.1.4",
"webpack-dev-middleware": ">=5.3.4",
"braces": ">=3.0.3",
"axios": ">=0.30.2",
"webpack": ">=5.94.0",
"serve-static": ">=1.16.0",
"path-to-regexp": ">=0.1.12"
}
}

View file

@ -1,5 +1,5 @@
---
title: "v1.81.12-stable - Guardrail Policy Templates & Action Builder"
title: "v1.81.12-stable.1 - Guardrail Policy Templates & Action Builder"
slug: "v1-81-12"
date: 2026-02-14T00:00:00
authors:
@ -27,7 +27,7 @@ import Image from '@theme/IdealImage';
docker run \
-e STORE_MODEL_IN_DB=True \
-p 4000:4000 \
ghcr.io/berriai/litellm:main-v1.81.12-stable
ghcr.io/berriai/litellm:main-v1.81.12-stable.1
```
</TabItem>

View file

@ -12,7 +12,19 @@
},
"overrides": {
"glob": ">=11.1.0",
"tar": ">=7.5.7",
"@isaacs/brace-expansion": ">=5.0.1"
"tar": ">=7.5.8",
"minimatch": ">=10.2.1",
"diff": ">=8.0.3",
"@isaacs/brace-expansion": ">=5.0.1",
"@babel/traverse": ">=7.23.2",
"ws": ">=7.5.10",
"http-proxy-middleware": ">=2.0.9",
"tar-fs": ">=2.1.4",
"webpack-dev-middleware": ">=5.3.4",
"braces": ">=3.0.3",
"axios": ">=0.30.2",
"webpack": ">=5.94.0",
"serve-static": ">=1.16.0",
"path-to-regexp": ">=0.1.12"
}
}
}

View file

@ -374,6 +374,7 @@ enable_end_user_cost_tracking_prometheus_only: Optional[bool] = None
custom_prometheus_metadata_labels: List[str] = []
custom_prometheus_tags: List[str] = []
prometheus_metrics_config: Optional[List] = None
prometheus_emit_stream_label: bool = False
disable_add_prefix_to_prompt: bool = (
False # used by anthropic, to disable adding prefix to prompt
)

View file

@ -244,6 +244,7 @@ REDIS_DAILY_TAG_SPEND_UPDATE_BUFFER_KEY = "litellm_daily_tag_spend_update_buffer
MAX_REDIS_BUFFER_DEQUEUE_COUNT = int(os.getenv("MAX_REDIS_BUFFER_DEQUEUE_COUNT", 100))
# Bounds asyncio.Queue() instances (log queues, spend update queues, etc.) to prevent unbounded memory growth
LITELLM_ASYNCIO_QUEUE_MAXSIZE = int(os.getenv("LITELLM_ASYNCIO_QUEUE_MAXSIZE", 1000))
TOOL_POLICY_CACHE_TTL_SECONDS = int(os.getenv("TOOL_POLICY_CACHE_TTL_SECONDS", 60))
# Aggregation threshold: default to 80% of the asyncio queue maxsize so the check can always trigger.
# Must be < LITELLM_ASYNCIO_QUEUE_MAXSIZE; if set higher the aggregation logic will never fire.
MAX_SIZE_IN_MEMORY_QUEUE = int(

View file

@ -974,6 +974,9 @@ class PrometheusLogger(CustomLogger):
),
client_ip=standard_logging_payload["metadata"].get("requester_ip_address"),
user_agent=standard_logging_payload["metadata"].get("user_agent"),
stream=str(standard_logging_payload.get("stream"))
if litellm.prometheus_emit_stream_label
else None,
)
if (
@ -1624,6 +1627,9 @@ class PrometheusLogger(CustomLogger):
client_ip=_metadata.get("requester_ip_address"),
user_agent=_metadata.get("user_agent"),
model_id=model_id,
stream=str(request_data.get("stream"))
if litellm.prometheus_emit_stream_label
else None,
)
_labels = prometheus_label_factory(
supported_enum_labels=self.get_labels_for_metric(

View file

@ -203,6 +203,129 @@ class RealTimeStreaming:
return True
return False
async def _handle_provider_config_message(self, raw_response) -> None:
"""Process a backend message when a provider_config is set (transformed path)."""
returned_object = self.provider_config.transform_realtime_response( # type: ignore[union-attr]
raw_response,
self.model,
self.logging_obj,
realtime_response_transform_input={
"session_configuration_request": self.session_configuration_request,
"current_output_item_id": self.current_output_item_id,
"current_response_id": self.current_response_id,
"current_delta_chunks": self.current_delta_chunks,
"current_conversation_id": self.current_conversation_id,
"current_item_chunks": self.current_item_chunks,
"current_delta_type": self.current_delta_type,
},
)
transformed_response = returned_object["response"]
self.current_output_item_id = returned_object["current_output_item_id"]
self.current_response_id = returned_object["current_response_id"]
self.current_delta_chunks = returned_object["current_delta_chunks"]
self.current_conversation_id = returned_object["current_conversation_id"]
self.current_item_chunks = returned_object["current_item_chunks"]
self.current_delta_type = returned_object["current_delta_type"]
self.session_configuration_request = returned_object["session_configuration_request"]
events = (
transformed_response
if isinstance(transformed_response, list)
else [transformed_response]
)
for event in events:
## GUARDRAIL: inject create_response=false on session.created
if isinstance(event, dict) and event.get("type") == "session.created":
if self._has_realtime_guardrails():
await self.backend_ws.send(
json.dumps(
{
"type": "session.update",
"session": {
"turn_detection": {
"type": "server_vad",
"create_response": False,
}
},
}
)
)
for event in events:
event_str = json.dumps(event)
## GUARDRAIL: run on transcription events in provider_config path too
if (
isinstance(event, dict)
and event.get("type")
== "conversation.item.input_audio_transcription.completed"
):
transcript = event.get("transcript", "")
self.store_message(event_str)
await self.websocket.send_text(event_str)
blocked = await self.run_realtime_guardrails(
transcript, item_id=event.get("item_id")
)
if not blocked:
await self.backend_ws.send(
json.dumps({"type": "response.create"})
)
continue
## LOGGING
self.store_message(event_str)
await self.websocket.send_text(event_str)
async def _handle_raw_backend_message(self, raw_response) -> bool:
"""Process a backend message without provider_config (raw path).
Returns True if the caller should skip the default store+forward (i.e. continue the loop).
"""
try:
event_obj = json.loads(raw_response)
if event_obj.get("type") == "session.created":
# If any realtime guardrails are registered, proactively
# set create_response=false so the LLM never auto-responds
# before our guardrail has a chance to run.
if self._has_realtime_guardrails():
await self.backend_ws.send(
json.dumps(
{
"type": "session.update",
"session": {
"turn_detection": {
"type": "server_vad",
"create_response": False,
}
},
}
)
)
verbose_logger.debug(
"[realtime guardrail] injected create_response=false into session"
)
if (
event_obj.get("type")
== "conversation.item.input_audio_transcription.completed"
):
transcript = event_obj.get("transcript", "")
## LOGGING — must happen before continue below
self.store_message(raw_response)
# Forward transcript to client so user sees what they said
await self.websocket.send_text(raw_response)
blocked = await self.run_realtime_guardrails(
transcript,
item_id=event_obj.get("item_id"),
)
if not blocked:
# Clean — trigger LLM response
await self.backend_ws.send(
json.dumps({"type": "response.create"})
)
return True
except (json.JSONDecodeError, AttributeError):
pass
return False
async def backend_to_client_send_messages(self):
import websockets
@ -216,128 +339,11 @@ class RealTimeStreaming:
raw_response = await self.backend_ws.recv() # type: ignore[assignment]
if self.provider_config:
returned_object = self.provider_config.transform_realtime_response(
raw_response,
self.model,
self.logging_obj,
realtime_response_transform_input={
"session_configuration_request": self.session_configuration_request,
"current_output_item_id": self.current_output_item_id,
"current_response_id": self.current_response_id,
"current_delta_chunks": self.current_delta_chunks,
"current_conversation_id": self.current_conversation_id,
"current_item_chunks": self.current_item_chunks,
"current_delta_type": self.current_delta_type,
},
)
transformed_response = returned_object["response"]
self.current_output_item_id = returned_object[
"current_output_item_id"
]
self.current_response_id = returned_object["current_response_id"]
self.current_delta_chunks = returned_object["current_delta_chunks"]
self.current_conversation_id = returned_object[
"current_conversation_id"
]
self.current_item_chunks = returned_object["current_item_chunks"]
self.current_delta_type = returned_object["current_delta_type"]
self.session_configuration_request = returned_object[
"session_configuration_request"
]
events = (
transformed_response
if isinstance(transformed_response, list)
else [transformed_response]
)
for event in events:
## GUARDRAIL: inject create_response=false on session.created
if isinstance(event, dict) and event.get("type") == "session.created":
if self._has_realtime_guardrails():
await self.backend_ws.send(
json.dumps(
{
"type": "session.update",
"session": {
"turn_detection": {
"type": "server_vad",
"create_response": False,
}
},
}
)
)
for event in events:
event_str = json.dumps(event)
## GUARDRAIL: run on transcription events in provider_config path too
if (
isinstance(event, dict)
and event.get("type")
== "conversation.item.input_audio_transcription.completed"
):
transcript = event.get("transcript", "")
self.store_message(event_str)
await self.websocket.send_text(event_str)
blocked = await self.run_realtime_guardrails(
transcript, item_id=event.get("item_id")
)
if not blocked:
await self.backend_ws.send(
json.dumps({"type": "response.create"})
)
continue
## LOGGING
self.store_message(event_str)
await self.websocket.send_text(event_str)
await self._handle_provider_config_message(raw_response)
else:
## GUARDRAIL: intercept transcription events before triggering LLM
try:
event_obj = json.loads(raw_response)
if event_obj.get("type") == "session.created":
# If any realtime guardrails are registered, proactively
# set create_response=false so the LLM never auto-responds
# before our guardrail has a chance to run.
if self._has_realtime_guardrails():
await self.backend_ws.send(
json.dumps(
{
"type": "session.update",
"session": {
"turn_detection": {
"type": "server_vad",
"create_response": False,
}
},
}
)
)
verbose_logger.debug(
"[realtime guardrail] injected create_response=false into session"
)
if (
event_obj.get("type")
== "conversation.item.input_audio_transcription.completed"
):
transcript = event_obj.get("transcript", "")
## LOGGING — must happen before continue below
self.store_message(raw_response)
# Forward transcript to client so user sees what they said
await self.websocket.send_text(raw_response)
blocked = await self.run_realtime_guardrails(
transcript,
item_id=event_obj.get("item_id"),
)
if not blocked:
# Clean — trigger LLM response
await self.backend_ws.send(
json.dumps({"type": "response.create"})
)
continue
except (json.JSONDecodeError, AttributeError):
pass
handled = await self._handle_raw_backend_message(raw_response)
if handled:
continue
## LOGGING
self.store_message(raw_response)
await self.websocket.send_text(raw_response)

View file

@ -11,6 +11,7 @@ from typing import Any, List, Optional
import httpx
from litellm.types.utils import Usage
from litellm.llms.bedrock.chat.invoke_transformations.amazon_qwen3_transformation import (
AmazonQwen3Config,
)
@ -79,10 +80,11 @@ class AmazonQwen2Config(AmazonQwen3Config):
# Set usage information if available in response
if "usage" in response_data:
usage_data = response_data["usage"]
if hasattr(model_response, 'usage'):
model_response.usage.prompt_tokens = usage_data.get("prompt_tokens", 0)
model_response.usage.completion_tokens = usage_data.get("completion_tokens", 0)
model_response.usage.total_tokens = usage_data.get("total_tokens", 0)
model_response.usage = Usage(
prompt_tokens=usage_data.get("prompt_tokens", 0),
completion_tokens=usage_data.get("completion_tokens", 0),
total_tokens=usage_data.get("total_tokens", 0),
)
return model_response

View file

@ -10,6 +10,7 @@ from typing import Any, List, Optional
import httpx
from litellm.types.utils import Usage
from litellm.llms.base_llm.chat.transformation import BaseConfig
from litellm.llms.bedrock.chat.invoke_transformations.base_invoke_transformation import (
AmazonInvokeConfig,
@ -201,10 +202,11 @@ class AmazonQwen3Config(AmazonInvokeConfig, BaseConfig):
# Set usage information if available in response
if "usage" in response_data:
usage_data = response_data["usage"]
if hasattr(model_response, 'usage'):
model_response.usage.prompt_tokens = usage_data.get("prompt_tokens", 0)
model_response.usage.completion_tokens = usage_data.get("completion_tokens", 0)
model_response.usage.total_tokens = usage_data.get("total_tokens", 0)
model_response.usage = Usage(
prompt_tokens=usage_data.get("prompt_tokens", 0),
completion_tokens=usage_data.get("completion_tokens", 0),
total_tokens=usage_data.get("total_tokens", 0),
)
return model_response

View file

@ -29,12 +29,13 @@ class BedrockRerankHandler(BaseAWSLLM):
async def arerank(
self,
prepared_request: BedrockPreparedRequest,
timeout: Optional[Union[float, httpx.Timeout]] = None,
client: Optional[AsyncHTTPHandler] = None,
):
if client is None:
client = get_async_httpx_client(llm_provider=litellm.LlmProviders.BEDROCK)
try:
response = await client.post(url=prepared_request["endpoint_url"], headers=dict(prepared_request["prepped"].headers), data=prepared_request["body"])
response = await client.post(url=prepared_request["endpoint_url"], headers=dict(prepared_request["prepped"].headers), data=prepared_request["body"], timeout=timeout)
response.raise_for_status()
except httpx.HTTPStatusError as err:
error_code = err.response.status_code
@ -56,6 +57,7 @@ class BedrockRerankHandler(BaseAWSLLM):
return_documents: Optional[bool] = True,
max_chunks_per_doc: Optional[int] = None,
_is_async: Optional[bool] = False,
timeout: Optional[Union[float, httpx.Timeout]] = None,
api_base: Optional[str] = None,
extra_headers: Optional[dict] = None,
client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None,
@ -89,12 +91,12 @@ class BedrockRerankHandler(BaseAWSLLM):
)
if _is_async:
return self.arerank(prepared_request, client=client if client is not None and isinstance(client, AsyncHTTPHandler) else None) # type: ignore
return self.arerank(prepared_request, timeout=timeout, client=client if client is not None and isinstance(client, AsyncHTTPHandler) else None) # type: ignore
if client is None or not isinstance(client, HTTPHandler):
client = _get_httpx_client()
try:
response = client.post(url=prepared_request["endpoint_url"], headers=dict(prepared_request["prepped"].headers), data=prepared_request["body"])
response = client.post(url=prepared_request["endpoint_url"], headers=dict(prepared_request["prepped"].headers), data=prepared_request["body"], timeout=timeout)
response.raise_for_status()
except httpx.HTTPStatusError as err:
error_code = err.response.status_code

View file

@ -162,6 +162,7 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig):
"service_tier",
"safety_identifier",
"prompt_cache_key",
"prompt_cache_retention",
"store",
] # works across all models

View file

@ -131,9 +131,7 @@ class OpenAIOSeriesConfig(OpenAIGPTConfig):
def is_model_o_series_model(self, model: str) -> bool:
model = model.split("/")[-1] # could be "openai/o3" or "o3"
return model in litellm.open_ai_chat_completion_models and any(
model.startswith(pfx) for pfx in ("o1", "o3", "o4")
)
return model.startswith(("o1", "o3", "o4")) and model in litellm.open_ai_chat_completion_models
@overload
def _transform_messages(
@ -173,4 +171,4 @@ class OpenAIOSeriesConfig(OpenAIGPTConfig):
else:
return super()._transform_messages(
messages, model, is_async=cast(Literal[False], False)
)
)

View file

@ -20562,6 +20562,39 @@
"supports_tool_choice": true,
"supports_vision": true
},
"gpt-5.3-codex": {
"cache_read_input_token_cost": 1.75e-07,
"cache_read_input_token_cost_priority": 3.5e-07,
"input_cost_per_token": 1.75e-06,
"input_cost_per_token_priority": 3.5e-06,
"litellm_provider": "openai",
"max_input_tokens": 272000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "responses",
"output_cost_per_token": 1.4e-05,
"output_cost_per_token_priority": 2.8e-05,
"supported_endpoints": [
"/v1/responses"
],
"supported_modalities": [
"text",
"image"
],
"supported_output_modalities": [
"text"
],
"supports_function_calling": true,
"supports_native_streaming": true,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_system_messages": false,
"supports_tool_choice": true,
"supports_vision": true
},
"gpt-5-mini": {
"cache_read_input_token_cost": 2.5e-08,
"cache_read_input_token_cost_flex": 1.25e-08,
@ -26555,65 +26588,124 @@
"supports_function_calling": true,
"supports_tool_choice": true
},
"perplexity/preset/fast-search": {
"litellm_provider": "perplexity",
"mode": "responses",
"supports_web_search": true,
"supports_preset": true,
"supports_function_calling": true
},
"perplexity/preset/pro-search": {
"litellm_provider": "perplexity",
"mode": "responses",
"supports_web_search": true,
"supports_preset": true
"supports_preset": true,
"supports_function_calling": true
},
"perplexity/openai/gpt-4o": {
"perplexity/preset/deep-research": {
"litellm_provider": "perplexity",
"mode": "responses",
"supports_web_search": true,
"supports_reasoning": false
"supports_preset": true,
"supports_function_calling": true
},
"perplexity/openai/gpt-4o-mini": {
"perplexity/preset/advanced-deep-research": {
"litellm_provider": "perplexity",
"mode": "responses",
"supports_web_search": true,
"supports_reasoning": false
"supports_preset": true,
"supports_function_calling": true
},
"perplexity/openai/gpt-5.2": {
"litellm_provider": "perplexity",
"mode": "responses",
"supports_web_search": true,
"supports_reasoning": true
"supports_reasoning": true,
"supports_function_calling": true
},
"perplexity/anthropic/claude-3-5-sonnet-20241022": {
"perplexity/openai/gpt-5.1": {
"litellm_provider": "perplexity",
"mode": "responses",
"supports_web_search": true,
"supports_reasoning": false
"supports_reasoning": false,
"supports_function_calling": true
},
"perplexity/anthropic/claude-3-5-haiku-20241022": {
"perplexity/openai/gpt-5-mini": {
"litellm_provider": "perplexity",
"mode": "responses",
"supports_web_search": true,
"supports_reasoning": false
"supports_reasoning": false,
"supports_function_calling": true
},
"perplexity/google/gemini-2.0-flash-exp": {
"perplexity/anthropic/claude-opus-4-6": {
"litellm_provider": "perplexity",
"mode": "responses",
"supports_web_search": true,
"supports_reasoning": false
"supports_reasoning": false,
"supports_function_calling": true
},
"perplexity/google/gemini-2.0-flash-thinking-exp": {
"perplexity/anthropic/claude-opus-4-5": {
"litellm_provider": "perplexity",
"mode": "responses",
"supports_web_search": true,
"supports_reasoning": true
"supports_reasoning": false,
"supports_function_calling": true
},
"perplexity/xai/grok-2-1212": {
"perplexity/anthropic/claude-sonnet-4-5": {
"litellm_provider": "perplexity",
"mode": "responses",
"supports_web_search": true,
"supports_reasoning": false
"supports_reasoning": false,
"supports_function_calling": true
},
"perplexity/xai/grok-2-vision-1212": {
"perplexity/anthropic/claude-haiku-4-5": {
"litellm_provider": "perplexity",
"mode": "responses",
"supports_web_search": true,
"supports_reasoning": false
"supports_reasoning": false,
"supports_function_calling": true
},
"perplexity/google/gemini-3-pro-preview": {
"litellm_provider": "perplexity",
"mode": "responses",
"supports_web_search": true,
"supports_reasoning": false,
"supports_function_calling": true
},
"perplexity/google/gemini-3-flash-preview": {
"litellm_provider": "perplexity",
"mode": "responses",
"supports_web_search": true,
"supports_reasoning": false,
"supports_function_calling": true
},
"perplexity/google/gemini-2.5-pro": {
"litellm_provider": "perplexity",
"mode": "responses",
"supports_web_search": true,
"supports_reasoning": false,
"supports_function_calling": true
},
"perplexity/google/gemini-2.5-flash": {
"litellm_provider": "perplexity",
"mode": "responses",
"supports_web_search": true,
"supports_reasoning": false,
"supports_function_calling": true
},
"perplexity/xai/grok-4-1-fast-non-reasoning": {
"litellm_provider": "perplexity",
"mode": "responses",
"supports_web_search": true,
"supports_reasoning": false,
"supports_function_calling": true
},
"perplexity/perplexity/sonar": {
"litellm_provider": "perplexity",
"mode": "responses",
"supports_web_search": true,
"supports_reasoning": false,
"supports_function_calling": true
},
"publicai/aisingapore/Qwen-SEA-LION-v4-32B-IT": {
"input_cost_per_token": 0.0,

View file

@ -10,6 +10,7 @@ from litellm.proxy._experimental.mcp_server.ui_session_utils import (
)
from litellm.proxy._experimental.mcp_server.utils import merge_mcp_headers
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.common_utils.http_parsing_utils import _safe_get_request_headers
from litellm.proxy.auth.ip_address_utils import IPAddressUtils
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.types.mcp import MCPAuth
@ -685,7 +686,7 @@ if MCP_AVAILABLE:
return await _execute_with_mcp_client(
new_mcp_server_request,
_test_connection_operation,
raw_headers=dict(request.headers),
raw_headers=_safe_get_request_headers(request),
)
@router.post("/test/tools/list")
@ -744,5 +745,5 @@ if MCP_AVAILABLE:
_list_tools_operation,
mcp_auth_header=mcp_auth_header,
oauth2_headers=oauth2_headers,
raw_headers=dict(request.headers),
raw_headers=_safe_get_request_headers(request),
)

View file

@ -183,6 +183,7 @@ class LitellmTableNames(str, enum.Enum):
KEY_TABLE_NAME = "LiteLLM_VerificationToken"
PROXY_MODEL_TABLE_NAME = "LiteLLM_ProxyModelTable"
MANAGED_FILE_TABLE_NAME = "LiteLLM_ManagedFileTable"
TOOL_TABLE_NAME = "LiteLLM_ToolTable"
class Litellm_EntityType(enum.Enum):
@ -850,6 +851,7 @@ class GenerateRequestBase(LiteLLMPydanticObjectBase):
max_budget: Optional[float] = None
user_id: Optional[str] = None
team_id: Optional[str] = None
agent_id: Optional[str] = None
max_parallel_requests: Optional[int] = None
metadata: Optional[dict] = {}
tpm_limit: Optional[int] = None
@ -2079,6 +2081,13 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
health_check_interval: int = Field(
300, description="background health check interval in seconds"
)
health_check_concurrency: Optional[int] = Field(
None,
description=(
"limit concurrent health checks per cycle; when unset, "
"health checks run without a concurrency cap"
),
)
alerting: Optional[List] = Field(
None,
description="List of alerting integrations. Today, just slack - `alerting: ['slack']`",
@ -4116,6 +4125,15 @@ class SpendUpdateQueueItem(TypedDict, total=False):
response_cost: Optional[float]
class ToolDiscoveryQueueItem(TypedDict, total=False):
tool_name: str
origin: Optional[str] # MCP server name or "user_defined"
created_by: Optional[str]
key_hash: Optional[str] # hash of virtual key that triggered discovery
team_id: Optional[str] # team that triggered discovery
key_alias: Optional[str] # human-readable key alias
class LiteLLM_ManagedFileTable(LiteLLMPydanticObjectBase):
unified_file_id: str
file_object: Optional[OpenAIFileObject] = None

View file

@ -483,7 +483,7 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
parent_otel_span = (
open_telemetry_logger.create_litellm_proxy_request_started_span(
start_time=start_time,
headers=dict(request.headers),
headers=_safe_get_request_headers(request),
)
)
@ -562,7 +562,7 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
parent_otel_span=parent_otel_span,
request_headers=dict(request.headers),
request_headers=_safe_get_request_headers(request),
)
is_proxy_admin = result["is_proxy_admin"]

View file

@ -135,17 +135,29 @@ def _safe_set_request_parsed_body(
def _safe_get_request_headers(request: Optional[Request]) -> dict:
"""
[Non-Blocking] Safely get the request headers
[Non-Blocking] Safely get the request headers.
Caches the result on request.state to avoid re-creating dict(request.headers) per call.
Warning: Callers must NOT mutate the returned dict — it is shared across
all callers within the same request via the cache.
"""
if request is None:
return {}
cached = getattr(request.state, "_cached_headers", None)
if cached is not None:
return cached
try:
if request is None:
return {}
return dict(request.headers)
headers = dict(request.headers)
except Exception as e:
verbose_proxy_logger.debug(
"Unexpected error reading request headers - {}".format(e)
)
return {}
headers = {}
try:
request.state._cached_headers = headers
except Exception:
pass # request.state may not be available in all contexts
return headers
def check_file_size_under_limit(

View file

@ -3,6 +3,7 @@ from fastapi_sso.sso.base import OpenID
from litellm._logging import verbose_logger
from litellm.integrations.custom_logger import CustomLogger
from litellm.proxy.common_utils.http_parsing_utils import _safe_get_request_headers
class CustomSSOLoginHandler(CustomLogger):
@ -18,7 +19,7 @@ class CustomSSOLoginHandler(CustomLogger):
self,
request: Request,
) -> OpenID:
request_headers_dict = dict(request.headers)
request_headers_dict = _safe_get_request_headers(request)
verbose_logger.debug("inside custom ui sso sign in hook...")
return OpenID(
id=request_headers_dict.get("x-litellm-user-id") or "123",

View file

@ -13,7 +13,17 @@ import random
import time
import traceback
from datetime import datetime, timedelta, timezone
from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Union, cast, overload
from typing import (
TYPE_CHECKING,
Any,
Dict,
List,
Literal,
Optional,
Union,
cast,
overload,
)
import litellm
from litellm._logging import verbose_proxy_logger
@ -23,18 +33,19 @@ from litellm.litellm_core_utils.safe_json_loads import safe_json_loads
from litellm.proxy._types import (
DB_CONNECTION_ERROR_TYPES,
BaseDailySpendTransaction,
DailyTagSpendTransaction,
DailyOrganizationSpendTransaction,
DailyTeamSpendTransaction,
DailyEndUserSpendTransaction,
DailyUserSpendTransaction,
DailyAgentSpendTransaction,
DailyEndUserSpendTransaction,
DailyOrganizationSpendTransaction,
DailyTagSpendTransaction,
DailyTeamSpendTransaction,
DailyUserSpendTransaction,
DBSpendUpdateTransactions,
Litellm_EntityType,
LiteLLM_UserTable,
SpendLogsMetadata,
SpendLogsPayload,
SpendUpdateQueueItem,
ToolDiscoveryQueueItem,
)
from litellm.proxy.db.db_transaction_queue.daily_spend_update_queue import (
DailySpendUpdateQueue,
@ -42,6 +53,9 @@ from litellm.proxy.db.db_transaction_queue.daily_spend_update_queue import (
from litellm.proxy.db.db_transaction_queue.pod_lock_manager import PodLockManager
from litellm.proxy.db.db_transaction_queue.redis_update_buffer import RedisUpdateBuffer
from litellm.proxy.db.db_transaction_queue.spend_update_queue import SpendUpdateQueue
from litellm.proxy.db.db_transaction_queue.tool_discovery_queue import (
ToolDiscoveryQueue,
)
from litellm.proxy.route_llm_request import ROUTE_ENDPOINT_MAPPING
if TYPE_CHECKING:
@ -67,6 +81,7 @@ class DBSpendUpdateWriter:
self.redis_update_buffer = RedisUpdateBuffer(redis_cache=self.redis_cache)
self.pod_lock_manager = PodLockManager()
self.spend_update_queue = SpendUpdateQueue()
self.tool_discovery_queue = ToolDiscoveryQueue()
self.daily_spend_update_queue = DailySpendUpdateQueue()
self.daily_team_spend_update_queue = DailySpendUpdateQueue()
self.daily_end_user_spend_update_queue = DailySpendUpdateQueue()
@ -222,10 +237,124 @@ class DBSpendUpdateWriter:
)
)
self._enqueue_tool_registry_upsert(
kwargs=kwargs,
completion_response=completion_response,
hashed_token=hashed_token,
team_id=team_id,
)
verbose_proxy_logger.debug("Runs spend update on all tables")
except Exception:
verbose_proxy_logger.error(
"Spend tracking - update_database failed. Spend log insertion or daily transaction enqueue "
"may not have completed for this request. "
"response_cost=%s, token=%s, user_id=%s, team_id=%s, org_id=%s, end_user_id=%s - %s",
response_cost,
token,
user_id,
team_id,
org_id,
end_user_id,
traceback.format_exc(),
)
def _enqueue_tool_registry_upsert(
self,
kwargs: Optional[dict],
completion_response: Optional[Any],
hashed_token: Optional[str] = None,
team_id: Optional[str] = None,
) -> None:
"""
Extract tool names from the LLM request and response and enqueue them
for upsert into LiteLLM_ToolTable via ToolDiscoveryQueue.
Handles four sources:
- MCP tools: standard_logging_object.mcp_tool_call_metadata.namespaced_tool_name
- Response tool_calls (OpenAI / Anthropic pass-through converted to OpenAI format):
completion_response.choices[].message.tool_calls[].function.name
- Request tools array (OpenAI format): kwargs["tools"][].function.name
- Request tools array (Anthropic /messages format): kwargs["passthrough_logging_payload"]
["request_body"]["tools"][].name
"""
try:
if kwargs is None:
return
# Extract key_alias from kwargs metadata if available
key_alias: Optional[str] = None
_litellm_params = kwargs.get("litellm_params") or {}
_metadata = _litellm_params.get("metadata") or {}
key_alias = _metadata.get("user_api_key_alias") or None
def _enqueue(tool_name: str, origin: str = "user_defined") -> None:
self.tool_discovery_queue.add_update(
ToolDiscoveryQueueItem(
tool_name=tool_name,
origin=origin,
key_hash=hashed_token,
team_id=team_id,
key_alias=key_alias,
)
)
# --- MCP tool calls ---
sl_object = kwargs.get("standard_logging_object")
if sl_object is not None:
mcp_metadata = (
sl_object.get("metadata", {}) or {}
).get("mcp_tool_call_metadata")
if mcp_metadata and isinstance(mcp_metadata, dict):
tool_name = mcp_metadata.get("namespaced_tool_name") or mcp_metadata.get("name")
mcp_server_name = mcp_metadata.get("mcp_server_name")
if tool_name:
_enqueue(tool_name, origin=mcp_server_name or "user_defined")
# --- Tools from request body (OpenAI format: tools[].function.name) ---
request_tools = kwargs.get("tools") or []
for tool_def in request_tools:
if not isinstance(tool_def, dict):
continue
fn = tool_def.get("function") or {}
name = fn.get("name") if isinstance(fn, dict) else None
if name:
_enqueue(name)
# --- Tools from Anthropic /messages pass-through request body
# (Anthropic format: tools[].name, no "function" wrapper) ---
passthrough_payload = kwargs.get("passthrough_logging_payload") or {}
request_body = (
passthrough_payload.get("request_body")
if isinstance(passthrough_payload, dict)
else None
) or {}
for tool_def in request_body.get("tools") or []:
if not isinstance(tool_def, dict):
continue
name = tool_def.get("name")
if name:
_enqueue(name)
# --- Response tool_calls (OpenAI format; Anthropic pass-through converts tool_use here) ---
if completion_response is not None and hasattr(completion_response, "choices"):
for choice in completion_response.choices or []:
message = getattr(choice, "message", None)
if message is None:
continue
tool_calls = getattr(message, "tool_calls", None)
if not tool_calls:
continue
for tc in tool_calls:
fn = getattr(tc, "function", None)
if fn is None:
continue
tool_name = getattr(fn, "name", None)
if tool_name:
_enqueue(tool_name)
except Exception as e:
verbose_proxy_logger.debug(
f"Error updating Prisma database: {traceback.format_exc()}"
"_enqueue_tool_registry_upsert error (non-blocking): %s", e
)
async def _update_key_db(
@ -295,9 +424,14 @@ class DBSpendUpdateWriter:
)
)
except Exception as e:
verbose_proxy_logger.debug(
"\033[91m"
+ f"Update User DB call failed to execute {str(e)}\n{traceback.format_exc()}"
verbose_proxy_logger.error(
"Spend tracking - failed to enqueue user spend update. "
"user_id=%s, end_user_id=%s, response_cost=%s - %s\n%s",
user_id,
end_user_id,
response_cost,
str(e),
traceback.format_exc(),
)
async def _update_team_db(
@ -334,11 +468,24 @@ class DBSpendUpdateWriter:
response_cost=response_cost,
)
)
except Exception:
pass
except Exception as e:
verbose_proxy_logger.error(
"Spend tracking - failed to enqueue team member spend update. "
"team_id=%s, user_id=%s, response_cost=%s - %s\n%s",
team_id,
user_id,
response_cost,
str(e),
traceback.format_exc(),
)
except Exception as e:
verbose_proxy_logger.debug(
f"Update Team DB failed to execute - {str(e)}\n{traceback.format_exc()}"
verbose_proxy_logger.error(
"Spend tracking - failed to enqueue team spend update. "
"team_id=%s, response_cost=%s - %s\n%s",
team_id,
response_cost,
str(e),
traceback.format_exc(),
)
raise e
@ -363,8 +510,13 @@ class DBSpendUpdateWriter:
)
)
except Exception as e:
verbose_proxy_logger.debug(
f"Update Org DB failed to execute - {str(e)}\n{traceback.format_exc()}"
verbose_proxy_logger.error(
"Spend tracking - failed to enqueue org spend update. "
"org_id=%s, response_cost=%s - %s\n%s",
org_id,
response_cost,
str(e),
traceback.format_exc(),
)
raise e
@ -411,8 +563,13 @@ class DBSpendUpdateWriter:
)
)
except Exception as e:
verbose_proxy_logger.debug(
f"Update Tag DB failed to execute - {str(e)}\n{traceback.format_exc()}"
verbose_proxy_logger.error(
"Spend tracking - failed to enqueue tag spend update. "
"request_tags=%s, response_cost=%s - %s\n%s",
request_tags,
response_cost,
str(e),
traceback.format_exc(),
)
raise e
@ -513,6 +670,17 @@ class DBSpendUpdateWriter:
await self.redis_update_buffer.get_all_update_transactions_from_redis_buffer()
)
if db_spend_update_transactions is not None:
verbose_proxy_logger.info(
"Spend tracking - committing spend updates from Redis to DB: "
"keys=%d, users=%d, teams=%d, orgs=%d, end_users=%d, team_members=%d, tags=%d",
len(db_spend_update_transactions.get("key_list_transactions") or {}),
len(db_spend_update_transactions.get("user_list_transactions") or {}),
len(db_spend_update_transactions.get("team_list_transactions") or {}),
len(db_spend_update_transactions.get("org_list_transactions") or {}),
len(db_spend_update_transactions.get("end_user_list_transactions") or {}),
len(db_spend_update_transactions.get("team_member_list_transactions") or {}),
len(db_spend_update_transactions.get("tag_list_transactions") or {}),
)
await self._commit_spend_updates_to_db(
prisma_client=prisma_client,
n_retry_times=n_retry_times,
@ -583,7 +751,12 @@ class DBSpendUpdateWriter:
daily_spend_transactions=daily_agent_spend_update_transactions,
)
except Exception as e:
verbose_proxy_logger.error(f"Error committing spend updates: {e}")
verbose_proxy_logger.error(
"Spend tracking - failed to commit spend updates from Redis to DB. "
"Data already popped from Redis may be lost. Error: %s\n%s",
str(e),
traceback.format_exc(),
)
finally:
await self.pod_lock_manager.release_lock(
cronjob_id=DB_SPEND_UPDATE_JOB_NAME,
@ -699,6 +872,25 @@ class DBSpendUpdateWriter:
daily_spend_transactions=daily_agent_spend_update_transactions,
)
################## Tool Registry Upserts ##################
await self._flush_tool_discovery_queue(prisma_client=prisma_client)
async def _flush_tool_discovery_queue(
self,
prisma_client: PrismaClient,
) -> None:
"""Flush ToolDiscoveryQueue and batch-upsert new tools into LiteLLM_ToolTable."""
from litellm.proxy.db.tool_registry_writer import batch_upsert_tools
try:
items = self.tool_discovery_queue.flush()
if items:
await batch_upsert_tools(prisma_client=prisma_client, items=items)
except Exception as e:
verbose_proxy_logger.debug(
"_flush_tool_discovery_queue error (non-blocking): %s", e
)
async def _commit_spend_updates_to_db( # noqa: PLR0915
self,
prisma_client: PrismaClient,

View file

@ -86,6 +86,11 @@ class DailySpendUpdateQueue(BaseUpdateQueue):
) -> Dict[str, BaseDailySpendTransaction]:
"""Get all updates from the queue and return all updates aggregated by daily_transaction_key. Works for both user and team spend updates."""
updates = await self.flush_all_updates_from_in_memory_queue()
if len(updates) > 0:
verbose_proxy_logger.info(
"Spend tracking - flushed %d daily spend update items from in-memory queue",
len(updates),
)
aggregated_daily_spend_update_transactions = (
DailySpendUpdateQueue.get_aggregated_daily_spend_update_transactions(
updates

View file

@ -80,6 +80,14 @@ class PodLockManager:
)
self._emit_acquired_lock_event(cronjob_id, self.pod_id)
return True
else:
verbose_proxy_logger.info(
"Spend tracking - pod %s could not acquire lock for cronjob_id=%s, "
"held by pod %s. Spend updates in Redis will wait for the leader pod to commit.",
self.pod_id,
cronjob_id,
current_value,
)
return False
except Exception as e:
verbose_proxy_logger.error(
@ -124,10 +132,12 @@ class PodLockManager:
pod_id=self.pod_id,
)
else:
verbose_proxy_logger.debug(
"Pod %s failed to release Redis lock for cronjob_id=%s",
verbose_proxy_logger.warning(
"Spend tracking - pod %s failed to release Redis lock for cronjob_id=%s. "
"Lock will expire after TTL=%ds.",
self.pod_id,
cronjob_id,
DEFAULT_CRON_JOB_LOCK_TTL_SECONDS,
)
else:
verbose_proxy_logger.debug(

View file

@ -96,14 +96,28 @@ class RedisUpdateBuffer:
list_of_transactions = [safe_dumps(transactions)]
if self.redis_cache is None:
return
current_redis_buffer_size = await self.redis_cache.async_rpush(
key=redis_key,
values=list_of_transactions,
)
await self._emit_new_item_added_to_redis_buffer_event(
queue_size=current_redis_buffer_size,
service=service_type,
)
try:
current_redis_buffer_size = await self.redis_cache.async_rpush(
key=redis_key,
values=list_of_transactions,
)
verbose_proxy_logger.debug(
"Spend tracking - pushed spend updates to Redis buffer. "
"redis_key=%s, buffer_size=%s",
redis_key,
current_redis_buffer_size,
)
await self._emit_new_item_added_to_redis_buffer_event(
queue_size=current_redis_buffer_size,
service=service_type,
)
except Exception as e:
verbose_proxy_logger.error(
"Spend tracking - failed to push spend updates to Redis (redis_key=%s). "
"Error: %s",
redis_key,
str(e),
)
async def store_in_memory_spend_updates_in_redis(
self,
@ -305,6 +319,13 @@ class RedisUpdateBuffer:
if list_of_transactions is None:
return None
verbose_proxy_logger.info(
"Spend tracking - popped %d spend update batches from Redis buffer (key=%s). "
"These items are now removed from Redis and must be committed to DB.",
len(list_of_transactions) if isinstance(list_of_transactions, list) else 1,
REDIS_UPDATE_BUFFER_KEY,
)
# Parse the list of transactions from JSON strings
parsed_transactions = self._parse_list_of_transactions(list_of_transactions)

View file

@ -31,6 +31,11 @@ class SpendUpdateQueue(BaseUpdateQueue):
) -> DBSpendUpdateTransactions:
"""Flush all updates from the queue and return all updates aggregated by entity type."""
updates = await self.flush_all_updates_from_in_memory_queue()
if len(updates) > 0:
verbose_proxy_logger.info(
"Spend tracking - flushed %d spend update items from in-memory queue",
len(updates),
)
verbose_proxy_logger.debug("Aggregating updates by entity type: %s", updates)
return self.get_aggregated_db_spend_update_transactions(updates)

View file

@ -0,0 +1,54 @@
"""
In-memory buffer for tool registry upserts.
Unlike SpendUpdateQueue (which aggregates increments), ToolDiscoveryQueue
uses set-deduplication: each unique tool_name is only queued once per flush
cycle (~30s). The seen-set is cleared on every flush so that call_count
increments in subsequent cycles rather than stopping after the first flush.
"""
from typing import List, Set
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import ToolDiscoveryQueueItem
class ToolDiscoveryQueue:
"""
In-memory buffer for tool registry upserts.
Deduplicates by tool_name within each flush cycle: a tool is only queued
once per ~30s batch, so call_count increments once per flush cycle the
tool appears in (not once per invocation, but not once per pod lifetime
either). The seen-set is cleared on flush so subsequent batches can
re-count the same tool.
"""
def __init__(self) -> None:
self._seen_tool_names: Set[str] = set()
self._pending: List[ToolDiscoveryQueueItem] = []
def add_update(self, item: ToolDiscoveryQueueItem) -> None:
"""Enqueue a tool discovery item if tool_name has not been seen before."""
tool_name = item.get("tool_name", "")
if not tool_name:
return
if tool_name in self._seen_tool_names:
verbose_proxy_logger.debug(
"ToolDiscoveryQueue: skipping already-seen tool %s", tool_name
)
return
self._seen_tool_names.add(tool_name)
self._pending.append(item)
verbose_proxy_logger.debug(
"ToolDiscoveryQueue: queued new tool %s (origin=%s)",
tool_name,
item.get("origin"),
)
def flush(self) -> List[ToolDiscoveryQueueItem]:
"""Return and clear all pending items. Resets seen-set so the next
flush cycle can re-count the same tools."""
items, self._pending = self._pending, []
self._seen_tool_names.clear()
return items

View file

@ -0,0 +1,179 @@
"""
DB helpers for LiteLLM_ToolTable — the global tool registry.
Tools are auto-discovered from LLM responses and upserted here.
Admins use the management endpoints to read and update call_policy.
NOTE: Uses raw SQL (query_raw / execute_raw) instead of Prisma model methods
because the generated Prisma Python client may not have LiteLLM_ToolTable
when running against an older generated schema.
"""
import uuid
from datetime import datetime, timezone
from typing import TYPE_CHECKING, Dict, List, Optional
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import ToolDiscoveryQueueItem
from litellm.types.tool_management import LiteLLM_ToolTableRow, ToolCallPolicy
if TYPE_CHECKING:
from litellm.proxy.utils import PrismaClient
def _row_to_model(row: dict) -> LiteLLM_ToolTableRow:
return LiteLLM_ToolTableRow(
tool_id=row.get("tool_id", ""),
tool_name=row.get("tool_name", ""),
origin=row.get("origin"),
call_policy=row.get("call_policy", "untrusted"),
call_count=int(row.get("call_count") or 0),
assignments=row.get("assignments"),
key_hash=row.get("key_hash"),
team_id=row.get("team_id"),
key_alias=row.get("key_alias"),
created_at=row.get("created_at"),
updated_at=row.get("updated_at"),
created_by=row.get("created_by"),
updated_by=row.get("updated_by"),
)
async def batch_upsert_tools(
prisma_client: "PrismaClient",
items: List[ToolDiscoveryQueueItem],
) -> None:
"""
Batch-upsert tool registry rows via raw SQL.
On first insert: sets call_policy = "untrusted" (schema default), call_count = 1.
On conflict: increments call_count; preserves existing call_policy.
"""
if not items:
return
try:
data = [item for item in items if item.get("tool_name")]
if not data:
return
for item in data:
tool_name = item.get("tool_name", "")
origin = item.get("origin") or "user_defined"
created_by = item.get("created_by") or "system"
key_hash = item.get("key_hash")
team_id = item.get("team_id")
key_alias = item.get("key_alias")
now = datetime.now(timezone.utc).isoformat()
await prisma_client.db.execute_raw(
'INSERT INTO "LiteLLM_ToolTable" '
"(tool_id, tool_name, origin, call_policy, call_count, created_by, updated_by, key_hash, team_id, key_alias, created_at, updated_at) "
"VALUES ($7, $1, $2, 'untrusted', 1, $3, $3, $4, $5, $6, $8, $8) "
"ON CONFLICT (tool_name) DO UPDATE SET "
"call_count = \"LiteLLM_ToolTable\".call_count + 1, "
"updated_at = $8",
tool_name,
origin,
created_by,
key_hash,
team_id,
key_alias,
str(uuid.uuid4()),
now,
)
verbose_proxy_logger.debug(
"tool_registry_writer: upserted %d tool(s)", len(data)
)
except Exception as e:
verbose_proxy_logger.error("tool_registry_writer batch_upsert_tools error: %s", e)
async def list_tools(
prisma_client: "PrismaClient",
call_policy: Optional[ToolCallPolicy] = None,
) -> List[LiteLLM_ToolTableRow]:
"""Return all tools, optionally filtered by call_policy."""
try:
if call_policy is not None:
rows = await prisma_client.db.query_raw(
'SELECT tool_id, tool_name, origin, call_policy, call_count, assignments, '
'key_hash, team_id, key_alias, created_at, updated_at, created_by, updated_by '
'FROM "LiteLLM_ToolTable" WHERE call_policy = $1 ORDER BY created_at DESC',
call_policy,
)
else:
rows = await prisma_client.db.query_raw(
'SELECT tool_id, tool_name, origin, call_policy, call_count, assignments, '
'key_hash, team_id, key_alias, created_at, updated_at, created_by, updated_by '
'FROM "LiteLLM_ToolTable" ORDER BY created_at DESC',
)
return [_row_to_model(row) for row in rows]
except Exception as e:
verbose_proxy_logger.error("tool_registry_writer list_tools error: %s", e)
return []
async def get_tool(
prisma_client: "PrismaClient",
tool_name: str,
) -> Optional[LiteLLM_ToolTableRow]:
"""Return a single tool row by tool_name."""
try:
rows = await prisma_client.db.query_raw(
'SELECT tool_id, tool_name, origin, call_policy, call_count, assignments, '
'key_hash, team_id, key_alias, created_at, updated_at, created_by, updated_by '
'FROM "LiteLLM_ToolTable" WHERE tool_name = $1',
tool_name,
)
if not rows:
return None
return _row_to_model(rows[0])
except Exception as e:
verbose_proxy_logger.error("tool_registry_writer get_tool error: %s", e)
return None
async def update_tool_policy(
prisma_client: "PrismaClient",
tool_name: str,
call_policy: ToolCallPolicy,
updated_by: Optional[str],
) -> Optional[LiteLLM_ToolTableRow]:
"""Update the call_policy for a tool. Upserts the row if it does not exist yet."""
try:
_updated_by = updated_by or "system"
now = datetime.now(timezone.utc).isoformat()
await prisma_client.db.execute_raw(
'INSERT INTO "LiteLLM_ToolTable" (tool_id, tool_name, call_policy, created_by, updated_by, created_at, updated_at) '
"VALUES ($4, $1, $2, $3, $3, $5, $5) "
"ON CONFLICT (tool_name) DO UPDATE SET call_policy = $2, updated_by = $3, updated_at = $5",
tool_name,
call_policy,
_updated_by,
str(uuid.uuid4()),
now,
)
return await get_tool(prisma_client, tool_name)
except Exception as e:
verbose_proxy_logger.error("tool_registry_writer update_tool_policy error: %s", e)
return None
async def get_tools_by_names(
prisma_client: "PrismaClient",
tool_names: List[str],
) -> Dict[str, str]:
"""
Return a {tool_name: call_policy} map for the given tool names.
Used by the policy enforcement guardrail — single batch query, never N+1.
"""
if not tool_names:
return {}
try:
placeholders = ", ".join(f"${i+1}" for i in range(len(tool_names)))
rows = await prisma_client.db.query_raw(
f'SELECT tool_name, call_policy FROM "LiteLLM_ToolTable" WHERE tool_name IN ({placeholders})',
*tool_names,
)
return {row["tool_name"]: row["call_policy"] for row in rows}
except Exception as e:
verbose_proxy_logger.error("tool_registry_writer get_tools_by_names error: %s", e)
return {}

View file

@ -1,6 +1,10 @@
from datetime import datetime
from fastapi import APIRouter, Depends, HTTPException, Request, Response
from fastapi.responses import ORJSONResponse, StreamingResponse
import litellm
from litellm._uuid import uuid
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
@ -17,7 +21,8 @@ router = APIRouter(
dependencies=[Depends(user_api_key_auth)],
)
@router.post(
"/models/{model_name:path}:generateContent", dependencies=[Depends(user_api_key_auth)]
"/models/{model_name:path}:generateContent",
dependencies=[Depends(user_api_key_auth)],
)
async def google_generate_content(
request: Request,
@ -36,12 +41,12 @@ async def google_generate_content(
data = await _read_request_body(request=request)
if "model" not in data:
data["model"] = model_name
# Extract generationConfig and pass it as config parameter
generation_config = data.pop("generationConfig", None)
if generation_config:
data["config"] = generation_config
# Add user authentication metadata for cost tracking
data = await add_litellm_data_to_request(
data=data,
@ -51,7 +56,19 @@ async def google_generate_content(
general_settings=general_settings,
version=version,
)
# Create logging object with full request metadata so callbacks (e.g. S3) get user/trace_id
data["litellm_call_id"] = request.headers.get(
"x-litellm-call-id", str(uuid.uuid4())
)
logging_obj, data = litellm.utils.function_setup(
original_function="agenerate_content",
rules_obj=litellm.utils.Rules(),
start_time=datetime.now(),
**data,
)
data["litellm_logging_obj"] = logging_obj
# call router
if llm_router is None:
raise HTTPException(status_code=500, detail="Router not initialized")
@ -103,6 +120,18 @@ async def google_stream_generate_content(
version=version,
)
# Create logging object with full request metadata so streaming END callbacks (e.g. S3) get user/trace_id
data["litellm_call_id"] = request.headers.get(
"x-litellm-call-id", str(uuid.uuid4())
)
logging_obj, data = litellm.utils.function_setup(
original_function="agenerate_content_stream",
rules_obj=litellm.utils.Rules(),
start_time=datetime.now(),
**data,
)
data["litellm_logging_obj"] = logging_obj
# call router
if llm_router is None:
raise HTTPException(status_code=500, detail="Router not initialized")
@ -247,11 +276,11 @@ async def create_interaction(
)
data = await _read_request_body(request=request)
# Default to gemini provider for interactions
if "custom_llm_provider" not in data:
data["custom_llm_provider"] = "gemini"
processor = ProxyBaseLLMRequestProcessing(data=data)
try:
return await processor.base_process_llm_request(
@ -301,7 +330,7 @@ async def get_interaction(
):
"""
Get an interaction by ID.
Per OpenAPI spec: GET /{api_version}/interactions/{interaction_id}
"""
from litellm.proxy.proxy_server import (
@ -319,7 +348,7 @@ async def get_interaction(
)
data = {"interaction_id": interaction_id, "custom_llm_provider": "gemini"}
processor = ProxyBaseLLMRequestProcessing(data=data)
try:
return await processor.base_process_llm_request(
@ -369,7 +398,7 @@ async def delete_interaction(
):
"""
Delete an interaction by ID.
Per OpenAPI spec: DELETE /{api_version}/interactions/{interaction_id}
"""
from litellm.proxy.proxy_server import (
@ -387,7 +416,7 @@ async def delete_interaction(
)
data = {"interaction_id": interaction_id, "custom_llm_provider": "gemini"}
processor = ProxyBaseLLMRequestProcessing(data=data)
try:
return await processor.base_process_llm_request(
@ -437,7 +466,7 @@ async def cancel_interaction(
):
"""
Cancel an interaction by ID.
Per OpenAPI spec: POST /{api_version}/interactions/{interaction_id}:cancel
"""
from litellm.proxy.proxy_server import (
@ -455,7 +484,7 @@ async def cancel_interaction(
)
data = {"interaction_id": interaction_id, "custom_llm_provider": "gemini"}
processor = ProxyBaseLLMRequestProcessing(data=data)
try:
return await processor.base_process_llm_request(

View file

@ -322,12 +322,6 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
analyze_payload,
)
async with session.post(analyze_url, json=analyze_payload) as response:
analyze_results = await response.json()
verbose_proxy_logger.debug("analyze_results: %s", analyze_results)
# Handle error responses from Presidio (e.g., {'error': 'No text provided'})
# Presidio may return a dict instead of a list when errors occur
def _fail_on_invalid_response(
reason: str,
) -> List[PresidioAnalyzeResponseItem]:
@ -347,6 +341,36 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
)
return []
async with session.post(
analyze_url,
json=analyze_payload,
headers={"Accept": "application/json"},
) as response:
# Validate HTTP status
if response.status >= 400:
error_body = await response.text()
return _fail_on_invalid_response(
f"HTTP {response.status} from Presidio analyzer: {error_body[:200]}"
)
# Validate Content-Type is JSON
content_type = getattr(
response,
"content_type",
response.headers.get("Content-Type", ""),
)
if "application/json" not in content_type:
error_body = await response.text()
return _fail_on_invalid_response(
f"expected application/json Content-Type but received '{content_type}'; body: '{error_body[:200]}'"
)
analyze_results = await response.json()
verbose_proxy_logger.debug("analyze_results: %s", analyze_results)
# Handle error responses from Presidio (e.g., {'error': 'No text provided'})
# Presidio may return a dict instead of a list when errors occur
if isinstance(analyze_results, dict):
if "error" in analyze_results:
return _fail_on_invalid_response(
@ -423,8 +447,29 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
}
async with session.post(
anonymize_url, json=anonymize_payload
anonymize_url,
json=anonymize_payload,
headers={"Accept": "application/json"},
) as response:
# Validate HTTP status
if response.status >= 400:
error_body = await response.text()
raise Exception(
f"Presidio anonymizer returned HTTP {response.status}: {error_body[:200]}"
)
# Validate Content-Type is JSON
content_type = getattr(
response,
"content_type",
response.headers.get("Content-Type", ""),
)
if "application/json" not in content_type:
error_body = await response.text()
raise Exception(
f"Presidio anonymizer returned non-JSON Content-Type '{content_type}'; body: '{error_body[:200]}'"
)
redacted_text = await response.json()
new_text = text
@ -456,7 +501,11 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
except Exception as e:
# Sanitize exception to avoid leaking the original text (which may
# contain API keys or other secrets) in error responses.
if "Invalid anonymizer response" in str(e):
error_str = str(e)
if (
"Invalid anonymizer response" in error_str
or "Presidio anonymizer returned" in error_str
):
raise
raise Exception(
f"Presidio PII anonymization failed: {type(e).__name__}"

View file

@ -0,0 +1,16 @@
import litellm
from litellm.types.guardrails import Guardrail, LitellmParams
def initialize_guardrail(litellm_params: LitellmParams, guardrail: Guardrail):
from litellm.proxy.guardrails.guardrail_hooks.tool_policy.tool_policy_guardrail import (
ToolPolicyGuardrail,
)
_callback = ToolPolicyGuardrail(
guardrail_name=guardrail.get("guardrail_name", "tool_policy"),
event_hook=litellm_params.mode,
default_on=litellm_params.default_on,
)
litellm.logging_callback_manager.add_litellm_callback(_callback)
return _callback

View file

@ -0,0 +1,163 @@
"""
Tool Policy Guardrail
Reads call_policy from LiteLLM_ToolTable and enforces it on LLM requests/responses.
Policy values:
"trusted" - allow through (no action)
"untrusted" - allow through (no action; default for newly discovered tools)
"blocked" - raise HTTPException, preventing the tool call
"dual_llm" - (Phase 3) send to second LLM for verification; currently treated as allowed
Configuration in proxy config YAML:
guardrails:
- guardrail_name: "tool_policy"
litellm_params:
guardrail: tool_policy
mode: post_call
or both pre and post call:
- guardrail_name: "tool_policy"
litellm_params:
guardrail: tool_policy
mode: during_call # runs before LLM and on response
"""
from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional
from fastapi import HTTPException
from litellm._logging import verbose_proxy_logger
from litellm.caching.dual_cache import DualCache
from litellm.constants import TOOL_POLICY_CACHE_TTL_SECONDS
from litellm.integrations.custom_guardrail import (
CustomGuardrail,
log_guardrail_information,
)
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.utils import GenericGuardrailAPIInputs
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
GUARDRAIL_NAME = "tool_policy"
class ToolPolicyGuardrail(CustomGuardrail):
"""
Guardrail that enforces per-tool call policies stored in LiteLLM_ToolTable.
Tools with call_policy="blocked" are rejected before/after the LLM call.
Tools with call_policy="trusted" or "untrusted" pass through unchanged.
"""
def __init__(self, **kwargs: Any) -> None:
if "supported_event_hooks" not in kwargs:
kwargs["supported_event_hooks"] = [
GuardrailEventHooks.pre_call,
GuardrailEventHooks.post_call,
GuardrailEventHooks.during_call,
]
super().__init__(**kwargs)
self._policy_cache: DualCache = DualCache()
@log_guardrail_information
async def apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,
request_data: dict,
input_type: Literal["request", "response"],
logging_obj: Optional["LiteLLMLoggingObj"] = None,
) -> GenericGuardrailAPIInputs:
"""
Enforce tool policies on both request tools and response tool_calls.
- input_type="request": check inputs["tools"] (tool definitions in the LLM request)
- input_type="response": check inputs["tool_calls"] (tool_calls in the LLM response)
Raises HTTPException (400) if any tool is "blocked".
"""
if input_type == "request":
tools = inputs.get("tools") or []
tool_names = [
t["function"]["name"]
for t in tools
if isinstance(t, dict)
and isinstance(t.get("function"), dict)
and t["function"].get("name")
]
else: # response
tool_calls = inputs.get("tool_calls") or []
tool_names = []
for tc in tool_calls:
fn = None
if isinstance(tc, dict):
fn = (tc.get("function") or {}).get("name")
elif hasattr(tc, "function"):
fn = getattr(tc.function, "name", None)
if fn:
tool_names.append(fn)
if not tool_names:
return inputs
policy_map = await self._get_policies_cached(tool_names)
blocked = [name for name in tool_names if policy_map.get(name) == "blocked"]
if blocked:
verbose_proxy_logger.warning(
"ToolPolicyGuardrail: blocking tool(s) %s (policy=blocked)", blocked
)
raise HTTPException(
status_code=400,
detail={
"error": "Violated tool policy",
"blocked_tools": blocked,
"message": f"Tool(s) {blocked} are blocked by policy.",
},
)
return inputs
async def _get_policies_cached(self, tool_names: List[str]) -> Dict[str, str]:
"""
Batch-fetch call_policy for the given tool names.
Caches per individual tool name (not per combination) so that adding
a new tool to a request doesn't invalidate the cached policies for all
the other tools already in the cache.
"""
from litellm.proxy.db.tool_registry_writer import get_tools_by_names
from litellm.proxy.proxy_server import prisma_client
if not tool_names or prisma_client is None:
return {}
result: Dict[str, str] = {}
cache_misses: List[str] = []
for name in tool_names:
cached = await self._policy_cache.async_get_cache(f"tool_policy:{name}")
if cached is not None and isinstance(cached, str):
result[name] = cached
else:
cache_misses.append(name)
if cache_misses:
fetched = await get_tools_by_names(
prisma_client=prisma_client, tool_names=cache_misses
)
for name, policy in fetched.items():
result[name] = policy
await self._policy_cache.async_set_cache(
key=f"tool_policy:{name}",
value=policy,
ttl=TOOL_POLICY_CACHE_TTL_SECONDS,
)
verbose_proxy_logger.debug(
"ToolPolicyGuardrail: fetched %d policies from DB (cache hits: %d)",
len(cache_misses),
len(tool_names) - len(cache_misses),
)
return result

View file

@ -3,12 +3,15 @@
import asyncio
import logging
import random
import sys
import threading
import time
from typing import List, Optional
import litellm
logger = logging.getLogger(__name__)
from litellm.constants import HEALTH_CHECK_TIMEOUT_SECONDS, DEFAULT_HEALTH_CHECK_PROMPT
from litellm.constants import DEFAULT_HEALTH_CHECK_PROMPT, HEALTH_CHECK_TIMEOUT_SECONDS
ILLEGAL_DISPLAY_PARAMS = [
"messages",
@ -23,6 +26,29 @@ ILLEGAL_DISPLAY_PARAMS = [
MINIMAL_DISPLAY_PARAMS = ["model", "mode_error"]
def _get_process_rss_mb() -> Optional[float]:
"""
Get process RSS memory in MB.
On Linux, ru_maxrss is in KB. On macOS, ru_maxrss is in bytes.
"""
try:
import resource
ru_maxrss = resource.getrusage(resource.RUSAGE_SELF).ru_maxrss
if sys.platform == "darwin":
return float(ru_maxrss) / (1024 * 1024)
return float(ru_maxrss) / 1024
except Exception:
return None
def _rss_mb_for_log() -> str:
rss_mb = _get_process_rss_mb()
if rss_mb is None:
return "unknown"
return f"{rss_mb:.2f}"
def _get_random_llm_message():
"""
Get a random message from the LLM.
@ -67,48 +93,115 @@ async def run_with_timeout(task, timeout):
try:
return await asyncio.wait_for(task, timeout)
except asyncio.TimeoutError:
task.cancel()
# Only cancel child tasks of the current task
current_task = asyncio.current_task()
for t in asyncio.all_tasks():
if t != current_task:
t.cancel()
try:
await asyncio.wait_for(task, 0.1) # Give 100ms for cleanup
except (asyncio.TimeoutError, asyncio.CancelledError, Exception):
pass
# `asyncio.wait_for()` already cancels only the awaited task on timeout.
# Do not cancel unrelated sibling health check tasks.
return {"error": "Timeout exceeded"}
async def _perform_health_check(model_list: list, details: Optional[bool] = True):
async def _run_model_health_check(model: dict):
litellm_params = model["litellm_params"]
model_info = model.get("model_info", {})
mode = model_info.get("mode", None)
litellm_params = _update_litellm_params_for_health_check(model_info, litellm_params)
timeout = model_info.get("health_check_timeout") or HEALTH_CHECK_TIMEOUT_SECONDS
return await run_with_timeout(
litellm.ahealth_check(
litellm_params,
mode=mode,
prompt=DEFAULT_HEALTH_CHECK_PROMPT,
input=["test from litellm"],
),
timeout,
)
async def _run_health_checks_with_bounded_concurrency(
models: list, concurrency_limit: int
) -> tuple[list, int]:
"""
Run health checks with at most `concurrency_limit` active tasks.
Preserves result ordering to match `models`.
"""
results: list = [None] * len(models)
tasks_to_index: dict[asyncio.Task, int] = {}
model_iter = iter(enumerate(models))
peak_in_flight = 0
def _schedule_next() -> bool:
nonlocal peak_in_flight
try:
idx, next_model = next(model_iter)
except StopIteration:
return False
task = asyncio.create_task(_run_model_health_check(next_model))
tasks_to_index[task] = idx
peak_in_flight = max(peak_in_flight, len(tasks_to_index))
return True
for _ in range(min(concurrency_limit, len(models))):
_schedule_next()
while tasks_to_index:
done, _ = await asyncio.wait(
set(tasks_to_index.keys()),
return_when=asyncio.FIRST_COMPLETED,
)
for task in done:
idx = tasks_to_index.pop(task)
try:
results[idx] = task.result()
except Exception as e:
results[idx] = e
_schedule_next()
return results, peak_in_flight
async def _perform_health_check(
model_list: list,
details: Optional[bool] = True,
max_concurrency: Optional[int] = None,
instrumentation_context: Optional[dict] = None,
):
"""
Perform a health check for each model in the list.
max_concurrency: Optional limit on concurrent health check requests.
"""
tasks = []
for model in model_list:
litellm_params = model["litellm_params"]
model_info = model.get("model_info", {})
mode = model_info.get("mode", None)
litellm_params = _update_litellm_params_for_health_check(
model_info, litellm_params
instrumentation_context = instrumentation_context or {}
instrumentation_enabled = bool(instrumentation_context.get("enabled", False))
cycle_id = instrumentation_context.get("cycle_id", "unknown")
source = instrumentation_context.get("source", "unknown")
dispatch_mode = "unbounded"
peak_in_flight = 0
if isinstance(max_concurrency, int) and max_concurrency > 0:
dispatch_mode = "bounded"
results, peak_in_flight = await _run_health_checks_with_bounded_concurrency(
model_list, max_concurrency
)
timeout = model_info.get("health_check_timeout") or HEALTH_CHECK_TIMEOUT_SECONDS
else:
tasks = [
asyncio.create_task(_run_model_health_check(model)) for model in model_list
]
peak_in_flight = len(tasks)
results = await asyncio.gather(*tasks, return_exceptions=True)
task = run_with_timeout(
litellm.ahealth_check(
model["litellm_params"],
mode=mode,
prompt=DEFAULT_HEALTH_CHECK_PROMPT,
input=["test from litellm"],
),
timeout,
if instrumentation_enabled:
logger.debug(
"health_check_dispatch_summary source=%s cycle_id=%s mode=%s model_count=%d max_concurrency=%s peak_in_flight=%d thread_count=%d rss_mb=%s",
source,
cycle_id,
dispatch_mode,
len(model_list),
max_concurrency,
peak_in_flight,
threading.active_count(),
_rss_mb_for_log(),
)
tasks.append(task)
results = await asyncio.gather(*tasks, return_exceptions=True)
healthy_endpoints = []
unhealthy_endpoints = []
@ -190,22 +283,48 @@ async def perform_health_check(
model: Optional[str] = None,
cli_model: Optional[str] = None,
details: Optional[bool] = True,
model_id: Optional[str] = None,
max_concurrency: Optional[int] = None,
instrumentation_context: Optional[dict] = None,
):
"""
Perform a health check on the system.
When model_id is provided, only the deployment with that id is checked
(so models that share the same name but have different ids are checked separately).
When model (name) is provided, all deployments matching that name are checked.
Returns:
(bool): True if the health check passes, False otherwise.
"""
instrumentation_context = instrumentation_context or {}
instrumentation_enabled = bool(instrumentation_context.get("enabled", False))
cycle_id = instrumentation_context.get("cycle_id", "unknown")
source = instrumentation_context.get("source", "unknown")
if not model_list:
if cli_model:
model_list = [
{"model_name": cli_model, "litellm_params": {"model": cli_model}}
]
else:
if instrumentation_enabled:
logger.debug(
"health_check_cycle_skipped source=%s cycle_id=%s reason=no_models",
source,
cycle_id,
)
return [], []
if model is not None:
cycle_start_time = time.monotonic()
requested_model_count = len(model_list)
# Filter by model_id first so a single deployment is checked when id is specified
if model_id is not None:
_by_id = [x for x in model_list if (x.get("model_info") or {}).get("id") == model_id]
if _by_id:
model_list = _by_id
elif model is not None:
_new_model_list = [
x for x in model_list if x["litellm_params"]["model"] == model
]
@ -213,11 +332,56 @@ async def perform_health_check(
_new_model_list = [x for x in model_list if x["model_name"] == model]
model_list = _new_model_list
post_filter_model_count = len(model_list)
model_list = filter_deployments_by_id(
model_list=model_list
) # filter duplicate deployments (e.g. when model alias'es are used)
healthy_endpoints, unhealthy_endpoints = await _perform_health_check(
model_list, details
)
deduped_model_count = len(model_list)
if instrumentation_enabled:
logger.debug(
"health_check_cycle_start source=%s cycle_id=%s requested_model_count=%d post_model_filter_count=%d deduped_model_count=%d max_concurrency=%s thread_count=%d rss_mb=%s",
source,
cycle_id,
requested_model_count,
post_filter_model_count,
deduped_model_count,
max_concurrency,
threading.active_count(),
_rss_mb_for_log(),
)
try:
healthy_endpoints, unhealthy_endpoints = await _perform_health_check(
model_list,
details,
max_concurrency=max_concurrency,
instrumentation_context=instrumentation_context,
)
except Exception:
if instrumentation_enabled:
logger.exception(
"health_check_cycle_failed source=%s cycle_id=%s model_count=%d duration_ms=%.2f thread_count=%d rss_mb=%s",
source,
cycle_id,
deduped_model_count,
(time.monotonic() - cycle_start_time) * 1000,
threading.active_count(),
_rss_mb_for_log(),
)
raise
if instrumentation_enabled:
logger.debug(
"health_check_cycle_complete source=%s cycle_id=%s model_count=%d healthy_count=%d unhealthy_count=%d duration_ms=%.2f thread_count=%d rss_mb=%s",
source,
cycle_id,
deduped_model_count,
len(healthy_endpoints),
len(unhealthy_endpoints),
(time.monotonic() - cycle_start_time) * 1000,
threading.active_count(),
_rss_mb_for_log(),
)
return healthy_endpoints, unhealthy_endpoints

View file

@ -16,7 +16,7 @@ from litellm.proxy.health_check import perform_health_check
class SharedHealthCheckManager:
"""
Manager for coordinating health checks across multiple pods using Redis.
This class implements a shared health check state mechanism that:
- Prevents duplicate health checks across pods
- Caches health check results with configurable TTL
@ -58,7 +58,7 @@ class SharedHealthCheckManager:
async def acquire_health_check_lock(self) -> bool:
"""
Attempt to acquire the global health check lock.
Returns:
bool: True if lock was acquired, False otherwise
"""
@ -74,7 +74,7 @@ class SharedHealthCheckManager:
nx=True, # Only set if key doesn't exist
ttl=self.lock_ttl,
)
if acquired:
verbose_proxy_logger.info(
"Pod %s acquired health check lock", self.pod_id
@ -83,12 +83,10 @@ class SharedHealthCheckManager:
verbose_proxy_logger.debug(
"Pod %s failed to acquire health check lock", self.pod_id
)
return acquired
except Exception as e:
verbose_proxy_logger.error(
"Error acquiring health check lock: %s", str(e)
)
verbose_proxy_logger.error("Error acquiring health check lock: %s", str(e))
return False
async def release_health_check_lock(self) -> None:
@ -106,14 +104,12 @@ class SharedHealthCheckManager:
"Pod %s released health check lock", self.pod_id
)
except Exception as e:
verbose_proxy_logger.error(
"Error releasing health check lock: %s", str(e)
)
verbose_proxy_logger.error("Error releasing health check lock: %s", str(e))
async def get_cached_health_check_results(self) -> Optional[Dict[str, Any]]:
"""
Get cached health check results from Redis.
Returns:
Optional[Dict]: Cached health check results or None if not found/expired
"""
@ -123,7 +119,7 @@ class SharedHealthCheckManager:
try:
cache_key = self.get_health_check_cache_key()
cached_data = await self.redis_cache.async_get_cache(cache_key)
if cached_data is None:
return None
@ -136,7 +132,7 @@ class SharedHealthCheckManager:
# Check if the cache is still valid
cache_timestamp = cached_results.get("timestamp", 0)
current_time = time.time()
if current_time - cache_timestamp > self.health_check_ttl:
verbose_proxy_logger.debug("Cached health check results expired")
return None
@ -151,13 +147,13 @@ class SharedHealthCheckManager:
return None
async def cache_health_check_results(
self,
healthy_endpoints: List[Dict[str, Any]],
unhealthy_endpoints: List[Dict[str, Any]]
self,
healthy_endpoints: List[Dict[str, Any]],
unhealthy_endpoints: List[Dict[str, Any]],
) -> None:
"""
Cache health check results in Redis.
Args:
healthy_endpoints: List of healthy endpoints
unhealthy_endpoints: List of unhealthy endpoints
@ -181,7 +177,7 @@ class SharedHealthCheckManager:
safe_dumps(cache_data),
ttl=self.health_check_ttl,
)
verbose_proxy_logger.info(
"Cached health check results for %d healthy and %d unhealthy endpoints",
len(healthy_endpoints),
@ -189,29 +185,29 @@ class SharedHealthCheckManager:
)
except Exception as e:
verbose_proxy_logger.error(
"Error caching health check results: %s", str(e)
)
verbose_proxy_logger.error("Error caching health check results: %s", str(e))
async def perform_shared_health_check(
self,
model_list: List[Dict[str, Any]],
details: bool = True
self,
model_list: List[Dict[str, Any]],
details: bool = True,
max_concurrency: Optional[int] = None,
) -> Tuple[List[Dict[str, Any]], List[Dict[str, Any]]]:
"""
Perform health check with shared state coordination.
This method:
1. First checks if there are recent cached results
2. If no recent cache, tries to acquire lock to run health check
3. If lock acquired, runs health check and caches results
4. If lock not acquired, waits briefly and tries to get cached results again
5. Falls back to running health check locally if no cache available
Args:
model_list: List of models to check
details: Whether to include detailed information
max_concurrency: Optional limit on concurrent health check requests
Returns:
Tuple of (healthy_endpoints, unhealthy_endpoints)
"""
@ -225,27 +221,29 @@ class SharedHealthCheckManager:
# No recent cache, try to acquire lock
lock_acquired = await self.acquire_health_check_lock()
if lock_acquired:
try:
# We have the lock, run health check
verbose_proxy_logger.info(
"Pod %s running health check for %d models",
self.pod_id,
len(model_list)
"Pod %s running health check for %d models",
self.pod_id,
len(model_list),
)
healthy_endpoints, unhealthy_endpoints = await perform_health_check(
model_list=model_list, details=details
model_list=model_list,
details=details,
max_concurrency=max_concurrency,
)
# Cache the results
await self.cache_health_check_results(
healthy_endpoints, unhealthy_endpoints
)
return healthy_endpoints, unhealthy_endpoints
finally:
# Always release the lock
await self.release_health_check_lock()
@ -254,10 +252,10 @@ class SharedHealthCheckManager:
verbose_proxy_logger.debug(
"Pod %s waiting for other pod to complete health check", self.pod_id
)
# Wait a bit for the other pod to complete
await asyncio.sleep(2)
# Try to get cached results again
cached_results = await self.get_cached_health_check_results()
if cached_results is not None:
@ -265,19 +263,23 @@ class SharedHealthCheckManager:
cached_results.get("healthy_endpoints", []),
cached_results.get("unhealthy_endpoints", []),
)
# Still no cache, fall back to local health check
verbose_proxy_logger.warning(
"Pod %s falling back to local health check (no cache available)",
self.pod_id
"Pod %s falling back to local health check (no cache available)",
self.pod_id,
)
return await perform_health_check(
model_list=model_list,
details=details,
max_concurrency=max_concurrency,
)
return await perform_health_check(model_list=model_list, details=details)
async def is_health_check_in_progress(self) -> bool:
"""
Check if a health check is currently in progress by another pod.
Returns:
bool: True if health check is in progress, False otherwise
"""
@ -297,7 +299,7 @@ class SharedHealthCheckManager:
async def get_health_check_status(self) -> Dict[str, Any]:
"""
Get the current status of health check coordination.
Returns:
Dict containing status information
"""
@ -320,7 +322,9 @@ class SharedHealthCheckManager:
cached_results = await self.get_cached_health_check_results()
status["cache_available"] = cached_results is not None
if cached_results:
status["cache_age_seconds"] = time.time() - cached_results.get("timestamp", 0)
status["cache_age_seconds"] = time.time() - cached_results.get(
"timestamp", 0
)
status["last_checked_by"] = cached_results.get("checked_by")
except Exception as e:

View file

@ -110,26 +110,31 @@ def _resolve_os_environ_variables(params: dict) -> dict:
def get_callback_identifier(callback):
"""
Get the callback identifier string, handling both strings and objects.
This function extracts a string identifier from a callback, which can be:
- A string (returned as-is)
- An object with a callback_name attribute
- An object registered in CustomLoggerRegistry
- Falls back to callback_name() helper function
Args:
callback: The callback to identify (can be str or object)
Returns:
str: The callback identifier string
"""
if isinstance(callback, str):
return callback
if hasattr(callback, 'callback_name') and callback.callback_name:
if hasattr(callback, "callback_name") and callback.callback_name:
return callback.callback_name
if hasattr(callback, '__class__'):
callback_strs = CustomLoggerRegistry.get_all_callback_strs_from_class_type(callback.__class__)
if hasattr(callback, 'callback_name') and callback.callback_name in callback_strs:
if hasattr(callback, "__class__"):
callback_strs = CustomLoggerRegistry.get_all_callback_strs_from_class_type(
callback.__class__
)
if (
hasattr(callback, "callback_name")
and callback.callback_name in callback_strs
):
return callback.callback_name
if callback_strs:
return callback_strs[0]
@ -151,7 +156,7 @@ services = Union[
"datadog_llm_observability",
"generic_api",
"arize",
"sqs"
"sqs",
],
str,
]
@ -224,7 +229,7 @@ async def health_services_endpoint( # noqa: PLR0915
"datadog_llm_observability",
"generic_api",
"arize",
"sqs"
"sqs",
]:
raise HTTPException(
status_code=400,
@ -238,14 +243,14 @@ async def health_services_endpoint( # noqa: PLR0915
service_in_success_callbacks = True
else:
for cb in litellm.success_callback:
if hasattr(cb, 'callback_name') and cb.callback_name == service:
if getattr(cb, "callback_name", None) == service:
service_in_success_callbacks = True
break
cb_id = get_callback_identifier(cb)
if cb_id == service:
service_in_success_callbacks = True
break
if (
service == "openmeter"
or service == "braintrust"
@ -320,6 +325,7 @@ async def health_services_endpoint( # noqa: PLR0915
)
elif service == "sqs":
from litellm.integrations.sqs import SQSLogger
sqs_logger = SQSLogger()
response = await sqs_logger.async_health_check()
return {
@ -518,12 +524,12 @@ async def _save_health_check_to_db(
def _build_model_param_to_info_mapping(model_list: list) -> dict:
"""
Build a mapping from model parameter to model info (model_name, model_id).
Multiple models might share the same model parameter, so we use a list.
Args:
model_list: List of model configurations
Returns:
Dictionary mapping model parameter to list of model info dicts
"""
@ -534,14 +540,16 @@ def _build_model_param_to_info_mapping(model_list: list) -> dict:
model_id = model_info.get("id")
litellm_params = model.get("litellm_params", {})
model_param = litellm_params.get("model")
if model_param and model_name:
if model_param not in model_param_to_info:
model_param_to_info[model_param] = []
model_param_to_info[model_param].append({
"model_name": model_name,
"model_id": model_id,
})
model_param_to_info[model_param].append(
{
"model_name": model_name,
"model_id": model_id,
}
)
return model_param_to_info
@ -552,19 +560,19 @@ def _aggregate_health_check_results(
) -> dict:
"""
Aggregate health check results per unique model.
Uses (model_id, model_name) as key, or (None, model_name) if model_id is None.
Args:
model_param_to_info: Mapping from model parameter to model info
healthy_endpoints: List of healthy endpoint results
unhealthy_endpoints: List of unhealthy endpoint results
Returns:
Dictionary mapping (model_id, model_name) to aggregated health check results
"""
model_results = {}
# Process healthy endpoints
for endpoint in healthy_endpoints:
model_param = endpoint.get("model")
@ -580,7 +588,7 @@ def _aggregate_health_check_results(
"error_message": None,
}
model_results[key]["healthy_count"] += 1
# Process unhealthy endpoints
for endpoint in unhealthy_endpoints:
model_param = endpoint.get("model")
@ -600,7 +608,7 @@ def _aggregate_health_check_results(
# Use the first error message encountered
if not model_results[key]["error_message"] and error_message:
model_results[key]["error_message"] = str(error_message)[:500]
return model_results
@ -613,14 +621,14 @@ async def _save_health_check_results_if_changed(
):
"""
Save health check results to database, but only if status changed or >1 hour since last save.
OPTIMIZATION: Only saves to database if the status has changed from the last saved check.
This dramatically reduces database writes when health status remains stable.
- Stable systems: ~1 write/hour per model (instead of 12 writes/hour with 5-min intervals)
- Status changes: Immediate write (no delay)
- Result: ~92% reduction in DB writes for stable systems, while maintaining real-time updates on changes
Args:
prisma_client: Database client
model_results: Dictionary of aggregated health check results per model
@ -630,7 +638,7 @@ async def _save_health_check_results_if_changed(
"""
for result in model_results.values():
new_status = "healthy" if result["healthy_count"] > 0 else "unhealthy"
# Check if we should save this result
should_save = True
lookup_key = result["model_id"] if result["model_id"] else result["model_name"]
@ -641,6 +649,7 @@ async def _save_health_check_results_if_changed(
# Check if last check was recent (within 1 hour)
if last_check.checked_at:
from datetime import datetime, timezone
time_since_last_check = (
datetime.now(timezone.utc) - last_check.checked_at
).total_seconds()
@ -648,7 +657,7 @@ async def _save_health_check_results_if_changed(
# This ensures we still get periodic updates even if status is stable
if time_since_last_check < 3600: # 1 hour threshold
should_save = False
if should_save:
asyncio.create_task(
prisma_client.save_health_check_result(
@ -675,27 +684,27 @@ async def _save_background_health_checks_to_db(
):
"""
Save background health check results to database for each model.
Maps health check endpoints back to their original models to get model_name and model_id.
Aggregates results per unique model (by model_id if available, otherwise model_name).
OPTIMIZATION: Only saves to database if the status has changed from the last saved check.
This dramatically reduces database writes when health status remains stable.
"""
if prisma_client is None:
return
try:
# Step 1: Build mapping from model parameter to model info
model_param_to_info = _build_model_param_to_info_mapping(model_list)
# Step 2: Aggregate health check results per unique model
model_results = _aggregate_health_check_results(
model_param_to_info,
healthy_endpoints,
unhealthy_endpoints,
)
# Step 3: Get latest health checks for all models in one query to compare status
latest_checks = await prisma_client.get_all_latest_health_checks()
latest_checks_map = {}
@ -704,7 +713,7 @@ async def _save_background_health_checks_to_db(
key = check.model_id if check.model_id else check.model_name
if key not in latest_checks_map:
latest_checks_map[key] = check
# Step 4: Save aggregated results, but only if status changed
await _save_health_check_results_if_changed(
prisma_client,
@ -729,10 +738,16 @@ async def _perform_health_check_and_save(
start_time,
user_id,
model_id=None,
max_concurrency=None,
):
"""Helper function to perform health check and save results to database"""
healthy_endpoints, unhealthy_endpoints = await perform_health_check(
model_list=model_list, cli_model=cli_model, model=target_model, details=details
model_list=model_list,
cli_model=cli_model,
model=target_model,
details=details,
max_concurrency=max_concurrency,
model_id=model_id,
)
# Optionally save health check result to database (non-blocking)
@ -789,6 +804,7 @@ async def health_endpoint(
import time
from litellm.proxy.proxy_server import (
health_check_concurrency,
health_check_details,
health_check_results,
llm_model_list,
@ -841,6 +857,7 @@ async def health_endpoint(
start_time=start_time,
user_id=user_api_key_dict.user_id,
model_id=None, # CLI model doesn't have model_id
max_concurrency=health_check_concurrency,
)
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
@ -864,6 +881,7 @@ async def health_endpoint(
start_time=start_time,
user_id=user_api_key_dict.user_id,
model_id=model_id,
max_concurrency=health_check_concurrency,
)
except Exception as e:
verbose_proxy_logger.error(
@ -1420,11 +1438,11 @@ async def test_model_connection(
status_code=500,
detail={"error": CommonProxyErrors.db_not_connected_error.value},
)
# Get model name from litellm_params
request_litellm_params = litellm_params or {}
model_name = request_litellm_params.get("model")
# Look up model configuration from router if model name is provided
# This gets the litellm_params from proxy config (with resolved env vars)
config_litellm_params: dict = {}
@ -1432,34 +1450,39 @@ async def test_model_connection(
try:
# First try to find by proxy model_name (e.g., "gpt-4o")
deployments = llm_router.get_model_list(model_name=model_name)
# If not found, try to find by litellm model name (e.g., "azure/gpt-4o")
if not deployments or len(deployments) == 0:
all_deployments = llm_router.get_model_list(model_name=None)
if all_deployments:
for deployment in all_deployments:
if deployment.get("litellm_params", {}).get("model") == model_name:
if (
deployment.get("litellm_params", {}).get("model")
== model_name
):
deployments = [deployment]
break
if deployments and len(deployments) > 0:
# Use the first deployment's litellm_params as base config
# These already have resolved environment variables from proxy config
config_litellm_params = dict(deployments[0].get("litellm_params", {}))
config_litellm_params = dict(
deployments[0].get("litellm_params", {})
)
except Exception as e:
verbose_proxy_logger.debug(
f"Could not find model {model_name} in router: {e}. "
"Proceeding with request params only."
)
# Merge: config params (from proxy config) as base, request params override
# This allows users to override specific params while using config for credentials
merged_litellm_params = {**config_litellm_params, **request_litellm_params}
# Resolve os.environ/ environment variables in any remaining request params
# This handles cases where user explicitly passes os.environ/ values to override config
litellm_params = _resolve_os_environ_variables(merged_litellm_params)
## Auth check
await ModelManagementAuthChecks.can_user_make_model_call(
model_params=Deployment(

View file

@ -12,7 +12,7 @@ from litellm.litellm_core_utils.core_helpers import (
)
from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.auth_checks import log_db_metrics
from litellm.proxy.auth.auth_checks import get_key_object, get_team_object, log_db_metrics
from litellm.proxy.auth.route_checks import RouteChecks
from litellm.proxy.utils import ProxyUpdateSpend
from litellm.types.utils import (
@ -76,6 +76,10 @@ class _ProxyDBLogger(CustomLogger):
traceback_str=traceback_str,
)
_metadata = await _ProxyDBLogger._enrich_failure_metadata_with_key_info(
metadata=_metadata,
)
existing_metadata: dict = request_data.get("metadata", None) or {}
existing_metadata.update(_metadata)
@ -255,6 +259,72 @@ class _ProxyDBLogger(CustomLogger):
"Error in tracking cost callback - %s", str(e)
)
@staticmethod
async def _enrich_failure_metadata_with_key_info(metadata: dict) -> dict:
"""
Enriches failure spend log metadata by looking up the key object (and team object)
from cache/DB when key fields are missing.
This handles two scenarios:
1. Auth errors (401): UserAPIKeyAuth is created with only api_key set, all other
fields are null. We look up the full key object to fill in alias, user_id,
team_id, etc.
2. Post-auth failures (provider errors, rate limits): key fields are populated
but team_alias is missing because LiteLLM_VerificationTokenView SQL view
doesn't include it. We look up the team object to fill in team_alias.
"""
api_key_hash = metadata.get("user_api_key")
if not api_key_hash:
return metadata
from litellm.proxy.proxy_server import (
prisma_client,
proxy_logging_obj,
user_api_key_cache,
)
# Step 1: If key fields are missing, look up the full key object
if metadata.get("user_api_key_alias") is None:
try:
key_obj = await get_key_object(
hashed_token=api_key_hash,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
if metadata.get("user_api_key_alias") is None:
metadata["user_api_key_alias"] = key_obj.key_alias
if metadata.get("user_api_key_user_id") is None:
metadata["user_api_key_user_id"] = key_obj.user_id
if metadata.get("user_api_key_team_id") is None:
metadata["user_api_key_team_id"] = key_obj.team_id
if metadata.get("user_api_key_org_id") is None:
metadata["user_api_key_org_id"] = key_obj.org_id
except Exception:
verbose_proxy_logger.debug(
"Failed to enrich failure metadata with key info for api_key=%s",
api_key_hash,
)
# Step 2: If team_id is known but team_alias is missing, look up the team object
team_id = metadata.get("user_api_key_team_id")
if team_id and metadata.get("user_api_key_team_alias") is None:
try:
team_obj = await get_team_object(
team_id=team_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
if team_obj.team_alias is not None:
metadata["user_api_key_team_alias"] = team_obj.team_alias
except Exception:
verbose_proxy_logger.debug(
"Failed to enrich failure metadata with team_alias for team_id=%s",
team_id,
)
return metadata
@staticmethod
def _should_track_errors_in_db():
"""

View file

@ -14,6 +14,7 @@ from litellm.proxy._types import (AddTeamCallback, CommonProxyErrors,
LitellmDataForBackendLLMCall,
LitellmUserRoles, SpecialHeaders,
TeamCallbackMetadata, UserAPIKeyAuth)
from litellm.proxy.common_utils.http_parsing_utils import _safe_get_request_headers
# Cache special headers as a frozenset for O(1) lookup performance
_SPECIAL_HEADERS_CACHE = frozenset(
@ -824,7 +825,7 @@ async def add_litellm_data_to_request( # noqa: PLR0915
from litellm.proxy.proxy_server import llm_router, premium_user
from litellm.types.proxy.litellm_pre_call_utils import SecretFields
_raw_headers: Dict[str, str] = dict(request.headers)
_raw_headers: Dict[str, str] = _safe_get_request_headers(request)
_headers: Dict[str, str] = clean_headers(
request.headers,
litellm_key_header_name=(

View file

@ -2535,6 +2535,7 @@ async def generate_key_helper_fn( # noqa: PLR0915
user_id: Optional[str] = None,
user_alias: Optional[str] = None,
team_id: Optional[str] = None,
agent_id: Optional[str] = None,
user_email: Optional[str] = None,
user_role: Optional[str] = None,
max_parallel_requests: Optional[int] = None,
@ -2668,6 +2669,7 @@ async def generate_key_helper_fn( # noqa: PLR0915
"max_budget": key_max_budget,
"user_id": user_id,
"team_id": team_id,
"agent_id": agent_id,
"project_id": project_id,
"max_parallel_requests": max_parallel_requests,
"metadata": metadata_json,

View file

@ -21,6 +21,7 @@ from typing_extensions import TypedDict
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.proxy.common_utils.http_parsing_utils import _safe_get_request_headers
from litellm._uuid import uuid
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.proxy._types import (
@ -724,7 +725,7 @@ async def get_service_provider_config(request: Request):
"SCIM ServiceProviderConfig request: method=%s url=%s headers=%s",
request.method,
request.url,
dict(request.headers),
_safe_get_request_headers(request),
)
meta = {
"resourceType": "ServiceProviderConfig",

View file

@ -0,0 +1,149 @@
"""
TOOL POLICY MANAGEMENT
All /tool management endpoints
GET /v1/tool/list - List all discovered tools and their policies
GET /v1/tool/{tool_name} - Get a single tool's details
POST /v1/tool/policy - Update the call_policy for a tool
"""
from typing import Optional
from fastapi import APIRouter, Depends, HTTPException
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import CommonProxyErrors, UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.types.tool_management import (
LiteLLM_ToolTableRow,
ToolCallPolicy,
ToolListResponse,
ToolPolicyUpdateRequest,
ToolPolicyUpdateResponse,
)
router = APIRouter()
@router.get(
"/v1/tool/list",
tags=["tool management"],
dependencies=[Depends(user_api_key_auth)],
response_model=ToolListResponse,
)
async def list_tools(
call_policy: Optional[ToolCallPolicy] = None,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
"""
List all auto-discovered tools and their call policies.
Parameters:
- call_policy: Optional filter — one of "trusted", "untrusted", "dual_llm", "blocked"
"""
from litellm.proxy.db.tool_registry_writer import list_tools as db_list_tools
from litellm.proxy.proxy_server import prisma_client
if prisma_client is None:
raise HTTPException(
status_code=500, detail=CommonProxyErrors.db_not_connected_error.value
)
try:
tools = await db_list_tools(prisma_client=prisma_client, call_policy=call_policy)
return ToolListResponse(tools=tools, total=len(tools))
except Exception as e:
verbose_proxy_logger.exception("Error listing tools: %s", e)
raise HTTPException(status_code=500, detail=str(e))
@router.get(
"/v1/tool/{tool_name:path}",
tags=["tool management"],
dependencies=[Depends(user_api_key_auth)],
response_model=LiteLLM_ToolTableRow,
)
async def get_tool(
tool_name: str,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
"""
Get details for a single tool.
Parameters:
- tool_name: The tool name (supports namespaced names with slashes)
"""
from litellm.proxy.db.tool_registry_writer import get_tool as db_get_tool
from litellm.proxy.proxy_server import prisma_client
if prisma_client is None:
raise HTTPException(
status_code=500, detail=CommonProxyErrors.db_not_connected_error.value
)
try:
tool = await db_get_tool(prisma_client=prisma_client, tool_name=tool_name)
if tool is None:
raise HTTPException(
status_code=404, detail=f"Tool '{tool_name}' not found"
)
return tool
except HTTPException:
raise
except Exception as e:
verbose_proxy_logger.exception("Error getting tool: %s", e)
raise HTTPException(status_code=500, detail=str(e))
@router.post(
"/v1/tool/policy",
tags=["tool management"],
dependencies=[Depends(user_api_key_auth)],
response_model=ToolPolicyUpdateResponse,
)
async def update_tool_policy(
data: ToolPolicyUpdateRequest,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
"""
Set the call policy for a tool.
Parameters:
- tool_name: str - The tool to update
- call_policy: "trusted" | "untrusted" | "dual_llm" | "blocked"
Setting a tool to "blocked" will cause the ToolPolicyGuardrail to remove
that tool_call from LLM responses before returning them to the client.
"""
from litellm.proxy.db.tool_registry_writer import (
update_tool_policy as db_update_tool_policy,
)
from litellm.proxy.proxy_server import prisma_client
if prisma_client is None:
raise HTTPException(
status_code=500, detail=CommonProxyErrors.db_not_connected_error.value
)
try:
updated = await db_update_tool_policy(
prisma_client=prisma_client,
tool_name=data.tool_name,
call_policy=data.call_policy,
updated_by=user_api_key_dict.user_id,
)
if updated is None:
raise HTTPException(
status_code=500, detail=f"Failed to update policy for tool '{data.tool_name}'"
)
return ToolPolicyUpdateResponse(
tool_name=updated.tool_name,
call_policy=updated.call_policy,
updated=True,
)
except HTTPException:
raise
except Exception as e:
verbose_proxy_logger.exception("Error updating tool policy: %s", e)
raise HTTPException(status_code=500, detail=str(e))

View file

@ -0,0 +1,9 @@
"""
Usage endpoints package.
Re-exports the router from endpoints module.
"""
from litellm.proxy.management_endpoints.usage_endpoints.endpoints import ( # noqa: F401
router,
)

View file

@ -0,0 +1,578 @@
"""
AI Usage Chat - uses LLM tool calling to answer questions about
usage/spend data by querying the aggregated daily activity endpoints.
"""
import json
from datetime import date
from typing import Any, AsyncIterator, Callable, Dict, List, Literal, Optional
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.constants import DEFAULT_COMPETITOR_DISCOVERY_MODEL
from litellm.types.proxy.management_endpoints.common_daily_activity import (
SpendAnalyticsPaginatedResponse,
)
from typing_extensions import TypedDict
# ---------------------------------------------------------------------------
# Constants
# ---------------------------------------------------------------------------
USAGE_AI_TEMPERATURE = 0.2
TABLE_DAILY_USER_SPEND = "litellm_dailyuserspend"
TABLE_DAILY_TEAM_SPEND = "litellm_dailyteamspend"
TABLE_DAILY_TAG_SPEND = "litellm_dailytagspend"
ENTITY_FIELD_USER = "user_id"
ENTITY_FIELD_TEAM = "team_id"
ENTITY_FIELD_TAG = "tag"
PAGINATED_PAGE_SIZE = 200
MAX_CHAT_MESSAGES = 20
TOP_N_MODELS = 15
TOP_N_PROVIDERS = 10
TOP_N_KEYS = 10
# ---------------------------------------------------------------------------
# Types
# ---------------------------------------------------------------------------
class SSEStatusEvent(TypedDict):
type: Literal["status"]
message: str
class SSEToolCallEvent(TypedDict, total=False):
type: Literal["tool_call"]
tool_name: str
tool_label: str
arguments: Dict[str, str]
status: Literal["running", "complete", "error"]
error: str
class SSEChunkEvent(TypedDict):
type: Literal["chunk"]
content: str
class SSEDoneEvent(TypedDict):
type: Literal["done"]
class SSEErrorEvent(TypedDict):
type: Literal["error"]
message: str
SSEEvent = (
SSEStatusEvent | SSEToolCallEvent | SSEChunkEvent | SSEDoneEvent | SSEErrorEvent
)
class ToolHandler(TypedDict):
fetch: Callable[..., Any]
summarise: Callable[[Dict[str, Any]], str]
label: str
# ---------------------------------------------------------------------------
# Tool definitions (OpenAI function-calling schema)
# ---------------------------------------------------------------------------
_DATE_PARAMS = {
"start_date": {"type": "string", "description": "Start date in YYYY-MM-DD format"},
"end_date": {"type": "string", "description": "End date in YYYY-MM-DD format"},
}
_TOOL_USAGE = {
"type": "function",
"function": {
"name": "get_usage_data",
"description": (
"Fetch aggregated global usage/spend data. Returns daily spend, "
"token counts, request counts, and breakdowns by model, provider, "
"and API key. Use for overall spend, top models, top providers."
),
"parameters": {
"type": "object",
"properties": {
**_DATE_PARAMS,
"user_id": {
"type": "string",
"description": "Optional user ID filter. Omit for global view.",
},
},
"required": ["start_date", "end_date"],
},
},
}
_TOOL_TEAM = {
"type": "function",
"function": {
"name": "get_team_usage_data",
"description": (
"Fetch usage/spend data broken down by team. Use for questions "
"like 'which team spends the most' or 'show me team X usage'."
),
"parameters": {
"type": "object",
"properties": {
**_DATE_PARAMS,
"team_ids": {
"type": "string",
"description": "Optional comma-separated team IDs. Omit for all teams.",
},
},
"required": ["start_date", "end_date"],
},
},
}
_TOOL_TAG = {
"type": "function",
"function": {
"name": "get_tag_usage_data",
"description": (
"Fetch usage/spend data broken down by tag. Tags are labels "
"attached to requests (features, environments, credentials)."
),
"parameters": {
"type": "object",
"properties": {
**_DATE_PARAMS,
"tags": {
"type": "string",
"description": "Optional comma-separated tag names. Omit for all tags.",
},
},
"required": ["start_date", "end_date"],
},
},
}
TOOLS_BASE = [_TOOL_USAGE]
TOOLS_ADMIN = [_TOOL_USAGE, _TOOL_TEAM, _TOOL_TAG]
def get_tools_for_role(is_admin: bool) -> List[Dict[str, Any]]:
"""Return the tool list appropriate for the user's role."""
return TOOLS_ADMIN if is_admin else TOOLS_BASE
_SYSTEM_PROMPT_BASE = (
"You are an AI assistant embedded in the LiteLLM Usage dashboard. "
"You help users understand their LLM API spend and usage data.\n\n"
"ALWAYS call the appropriate tool(s) first to fetch data before answering. "
"You may call multiple tools if the question spans different dimensions.\n\n"
"Guidelines:\n"
"- Be concise and specific. Use exact numbers from the data.\n"
"- Format costs as dollar amounts (e.g. $12.34).\n"
"- When comparing entities, show a ranked list.\n"
"- If data is empty or no results found, say so clearly.\n"
"- Do not hallucinate data — only use what the tools return.\n"
"- Today's date will be provided below. Use it to interpret relative dates "
"like 'this week', 'this month', 'last 7 days', etc."
)
_TOOL_DESCRIPTIONS_ADMIN = (
"You have access to these tools:\n"
"- `get_usage_data`: Global/user-level usage (spend, models, providers, API keys)\n"
"- `get_team_usage_data`: Team-level usage breakdown\n"
"- `get_tag_usage_data`: Tag-level usage breakdown\n\n"
)
_TOOL_DESCRIPTIONS_BASE = (
"You have access to this tool:\n"
"- `get_usage_data`: Your usage data (spend, models, providers, API keys)\n\n"
)
def _build_system_prompt(is_admin: bool) -> str:
"""Build role-appropriate system prompt with today's date."""
tool_desc = _TOOL_DESCRIPTIONS_ADMIN if is_admin else _TOOL_DESCRIPTIONS_BASE
return (
f"{_SYSTEM_PROMPT_BASE}\n\n{tool_desc}"
f"Today's date: {date.today().isoformat()}"
)
# keep a public reference for test assertions
SYSTEM_PROMPT = _SYSTEM_PROMPT_BASE
# ---------------------------------------------------------------------------
# Data fetchers
# ---------------------------------------------------------------------------
def _parse_csv_ids(raw: Optional[str]) -> Optional[List[str]]:
if not raw:
return None
return [t.strip() for t in raw.split(",") if t.strip()]
async def _query_activity(
table_name: str,
entity_id_field: str,
entity_id: Optional[Any],
start_date: str,
end_date: str,
*,
use_aggregated: bool = False,
) -> SpendAnalyticsPaginatedResponse:
"""Shared helper that calls the daily activity query layer."""
from litellm.proxy.management_endpoints.common_daily_activity import (
get_daily_activity,
get_daily_activity_aggregated,
)
from litellm.proxy.proxy_server import prisma_client
if use_aggregated:
return await get_daily_activity_aggregated(
prisma_client=prisma_client,
table_name=table_name,
entity_id_field=entity_id_field,
entity_id=entity_id,
entity_metadata_field=None,
start_date=start_date,
end_date=end_date,
model=None,
api_key=None,
)
return await get_daily_activity(
prisma_client=prisma_client,
table_name=table_name,
entity_id_field=entity_id_field,
entity_id=entity_id,
entity_metadata_field=None,
start_date=start_date,
end_date=end_date,
model=None,
api_key=None,
page=1,
page_size=PAGINATED_PAGE_SIZE,
)
async def _fetch_usage_data(
start_date: str, end_date: str, user_id: Optional[str] = None
) -> Dict[str, Any]:
resp = await _query_activity(
TABLE_DAILY_USER_SPEND,
ENTITY_FIELD_USER,
user_id,
start_date,
end_date,
use_aggregated=True,
)
return resp.model_dump(mode="json")
async def _fetch_team_usage_data(
start_date: str, end_date: str, team_ids: Optional[str] = None
) -> Dict[str, Any]:
resp = await _query_activity(
TABLE_DAILY_TEAM_SPEND,
ENTITY_FIELD_TEAM,
_parse_csv_ids(team_ids),
start_date,
end_date,
)
return resp.model_dump(mode="json")
async def _fetch_tag_usage_data(
start_date: str, end_date: str, tags: Optional[str] = None
) -> Dict[str, Any]:
resp = await _query_activity(
TABLE_DAILY_TAG_SPEND,
ENTITY_FIELD_TAG,
_parse_csv_ids(tags),
start_date,
end_date,
)
return resp.model_dump(mode="json")
# ---------------------------------------------------------------------------
# Summarisers — convert raw JSON to concise text the LLM can reason over
# ---------------------------------------------------------------------------
def _accumulate_breakdown(
results: List[Dict[str, Any]], dimension: str, fields: List[str]
) -> Dict[str, Dict[str, float]]:
"""Aggregate a single breakdown dimension across days."""
totals: Dict[str, Dict[str, float]] = {}
for day in results:
for key, entry in day.get("breakdown", {}).get(dimension, {}).items():
if key not in totals:
totals[key] = {f: 0.0 for f in fields}
m = entry.get("metrics", {})
for f in fields:
totals[key][f] += m.get(f, 0)
return totals
def _ranked_lines(
totals: Dict[str, Dict[str, float]],
fmt: Callable[[str, Dict[str, float]], str],
limit: int,
) -> List[str]:
"""Sort by spend descending, format each entry, and truncate."""
return [
fmt(name, vals)
for name, vals in sorted(totals.items(), key=lambda x: -x[1].get("spend", 0))[
:limit
]
]
def _summarise_usage_data(data: Dict[str, Any]) -> str:
meta = data.get("metadata", {})
results = data.get("results", [])
header = (
f"Total Spend: ${meta.get('total_spend', 0):.4f}\n"
f"Total Requests: {meta.get('total_api_requests', 0)}\n"
f"Successful: {meta.get('total_successful_requests', 0)} | "
f"Failed: {meta.get('total_failed_requests', 0)}\n"
f"Total Tokens: {meta.get('total_tokens', 0)}"
)
models = _accumulate_breakdown(
results, "models", ["spend", "api_requests", "total_tokens"]
)
providers = _accumulate_breakdown(results, "providers", ["spend", "api_requests"])
model_lines = _ranked_lines(
models,
lambda n, d: f" - {n}: ${d['spend']:.4f} ({int(d['api_requests'])} reqs, {int(d['total_tokens'])} tokens)",
TOP_N_MODELS,
)
provider_lines = _ranked_lines(
providers,
lambda n, d: f" - {n}: ${d['spend']:.4f} ({int(d['api_requests'])} reqs)",
TOP_N_PROVIDERS,
)
sections = [header, ""]
sections += ["Top Models by Spend:"] + (model_lines or [" (no data)"]) + [""]
sections += ["Top Providers by Spend:"] + (provider_lines or [" (no data)"])
return "\n".join(sections)
def _summarise_entity_data(data: Dict[str, Any], entity_label: str) -> str:
"""Summarise team/tag entity usage data."""
results = data.get("results", [])
if not results:
return f"No {entity_label} usage data found for the given date range."
totals: Dict[str, Dict[str, Any]] = {}
for day in results:
for eid, entry in day.get("breakdown", {}).get("entities", {}).items():
if eid not in totals:
alias = entry.get("metadata", {}).get("alias", eid)
totals[eid] = {"alias": alias, "spend": 0.0, "requests": 0, "tokens": 0}
m = entry.get("metrics", {})
totals[eid]["spend"] += m.get("spend", 0)
totals[eid]["requests"] += m.get("api_requests", 0)
totals[eid]["tokens"] += m.get("total_tokens", 0)
lines = [f"{entity_label} Usage ({len(totals)} {entity_label.lower()}s):", ""]
for eid, d in sorted(totals.items(), key=lambda x: -x[1]["spend"]):
label = d["alias"] if d["alias"] != eid else eid
lines.append(
f"- {label} (ID: {eid}): ${d['spend']:.4f} | "
f"{int(d['requests'])} reqs | {int(d['tokens'])} tokens"
)
return "\n".join(lines)
# ---------------------------------------------------------------------------
# Tool dispatch registry
# ---------------------------------------------------------------------------
TOOL_HANDLERS: Dict[str, ToolHandler] = {
"get_usage_data": ToolHandler(
fetch=_fetch_usage_data,
summarise=_summarise_usage_data,
label="global usage data",
),
"get_team_usage_data": ToolHandler(
fetch=_fetch_team_usage_data,
summarise=lambda data: _summarise_entity_data(data, "Team"),
label="team usage data",
),
"get_tag_usage_data": ToolHandler(
fetch=_fetch_tag_usage_data,
summarise=lambda data: _summarise_entity_data(data, "Tag"),
label="tag usage data",
),
}
# ---------------------------------------------------------------------------
# SSE streaming
# ---------------------------------------------------------------------------
def _sse(event: SSEEvent) -> str:
return f"data: {json.dumps(event)}\n\n"
def _resolve_fetch_kwargs(
fn_name: str,
fn_args: Dict[str, str],
user_id: Optional[str],
is_admin: bool,
) -> Dict[str, Any]:
"""Build keyword arguments for a tool's fetch function."""
start_date = fn_args.get("start_date", "")
end_date = fn_args.get("end_date", "")
if not start_date or not end_date:
raise ValueError("Missing required start_date or end_date from tool arguments")
kwargs: Dict[str, Any] = {"start_date": start_date, "end_date": end_date}
if fn_name == "get_usage_data":
if not is_admin:
kwargs["user_id"] = user_id
elif fn_args.get("user_id"):
kwargs["user_id"] = fn_args["user_id"]
elif fn_name == "get_team_usage_data" and fn_args.get("team_ids"):
kwargs["team_ids"] = fn_args["team_ids"]
elif fn_name == "get_tag_usage_data" and fn_args.get("tags"):
kwargs["tags"] = fn_args["tags"]
return kwargs
async def _execute_tool_call(
handler: ToolHandler,
fn_name: str,
fn_args: Dict[str, str],
user_id: Optional[str],
is_admin: bool,
) -> str:
"""Run a single tool and return the summarised result text."""
kwargs = _resolve_fetch_kwargs(fn_name, fn_args, user_id, is_admin)
raw_data = await handler["fetch"](**kwargs)
return handler["summarise"](raw_data)
async def _process_tool_call(
tc: Any,
chat_messages: List[Dict[str, Any]],
user_id: Optional[str],
is_admin: bool,
) -> AsyncIterator[str]:
"""Execute a single tool call, yielding SSE events for status."""
fn_name = tc.function.name
fn_args = json.loads(tc.function.arguments)
allowed_names = {t["function"]["name"] for t in get_tools_for_role(is_admin)}
handler = TOOL_HANDLERS.get(fn_name)
if fn_name not in allowed_names or not handler:
chat_messages.append(
{
"role": "tool",
"tool_call_id": tc.id,
"content": f"Tool not available: {fn_name}",
}
)
return
tool_event_base = {
"type": "tool_call",
"tool_name": fn_name,
"tool_label": handler["label"],
"arguments": fn_args,
}
yield _sse({**tool_event_base, "status": "running"})
try:
tool_result = await _execute_tool_call(
handler, fn_name, fn_args, user_id, is_admin
)
yield _sse({**tool_event_base, "status": "complete"})
except Exception as e:
verbose_proxy_logger.error("Tool %s failed: %s", fn_name, e)
tool_result = f"Error fetching {handler['label']}. Please try again."
yield _sse({**tool_event_base, "status": "error"})
chat_messages.append(
{"role": "tool", "tool_call_id": tc.id, "content": tool_result}
)
async def _stream_final_response(
model: str, chat_messages: List[Dict[str, Any]]
) -> AsyncIterator[str]:
"""Stream the final LLM response after tool results are appended."""
yield _sse({"type": "status", "message": "Analyzing results..."})
response = await litellm.acompletion(
model=model,
messages=chat_messages,
stream=True,
temperature=USAGE_AI_TEMPERATURE,
)
async for chunk in response:
delta = chunk.choices[0].delta.content
if delta:
yield _sse({"type": "chunk", "content": delta})
async def stream_usage_ai_chat(
messages: List[Dict[str, str]],
model: Optional[str] = None,
user_id: Optional[str] = None,
is_admin: bool = False,
) -> AsyncIterator[str]:
"""Stream SSE events: status → tool_call → chunk → done."""
resolved_model = (model or "").strip() or DEFAULT_COMPETITOR_DISCOVERY_MODEL
truncated = (
messages[-MAX_CHAT_MESSAGES:] if len(messages) > MAX_CHAT_MESSAGES else messages
)
chat_messages: List[Dict[str, Any]] = [
{"role": "system", "content": _build_system_prompt(is_admin)},
*truncated,
]
try:
yield _sse({"type": "status", "message": "Thinking..."})
tools = get_tools_for_role(is_admin)
response = await litellm.acompletion(
model=resolved_model,
messages=chat_messages,
tools=tools,
temperature=USAGE_AI_TEMPERATURE,
)
choice = response.choices[0] # type: ignore
if not choice.message.tool_calls:
if choice.message.content:
yield _sse({"type": "chunk", "content": choice.message.content})
yield _sse({"type": "done"})
return
chat_messages.append(choice.message.model_dump())
for tc in choice.message.tool_calls:
async for event in _process_tool_call(tc, chat_messages, user_id, is_admin):
yield event
async for event in _stream_final_response(resolved_model, chat_messages):
yield event
yield _sse({"type": "done"})
except Exception as e:
verbose_proxy_logger.error("AI usage chat failed: %s", e)
yield _sse(
{
"type": "error",
"message": "An internal error occurred. Please try again.",
}
)

View file

@ -0,0 +1,65 @@
"""
USAGE AI CHAT ENDPOINTS
/usage/ai/chat - Stream AI chat responses about usage data
"""
from typing import List, Literal, Optional
from fastapi import APIRouter, Depends, Request
from fastapi.responses import StreamingResponse
from pydantic import BaseModel, Field
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
router = APIRouter()
class ChatMessage(BaseModel):
role: Literal["user", "assistant"]
content: str
class UsageAIChatRequest(BaseModel):
messages: List[ChatMessage] = Field(
..., description="Chat messages (user/assistant history)"
)
model: Optional[str] = Field(default=None, description="Model to use for AI chat")
@router.post(
"/usage/ai/chat",
tags=["Budget & Spend Tracking"],
dependencies=[Depends(user_api_key_auth)],
)
async def usage_ai_chat(
data: UsageAIChatRequest,
request: Request,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
"""
AI chat about usage data. Streams SSE events with the AI response.
The AI agent has access to tools that query aggregated daily activity data.
"""
from litellm.proxy.management_endpoints.common_utils import (
_user_has_admin_view,
)
from litellm.proxy.management_endpoints.usage_endpoints.ai_usage_chat import (
stream_usage_ai_chat,
)
is_admin = _user_has_admin_view(user_api_key_dict)
user_id = user_api_key_dict.user_id
messages = [{"role": m.role, "content": m.content} for m in data.messages]
return StreamingResponse(
stream_usage_ai_chat(
messages=messages,
model=data.model,
user_id=user_id,
is_admin=is_admin,
),
media_type="text/event-stream",
headers={"Cache-Control": "no-cache", "X-Accel-Buffering": "no"},
)

View file

@ -28,6 +28,7 @@ from litellm.proxy.auth.route_checks import RouteChecks
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_utils.http_parsing_utils import (
_read_request_body,
_safe_get_request_headers,
_safe_set_request_parsed_body,
get_form_data,
get_request_body,
@ -60,7 +61,7 @@ def create_request_copy(request: Request):
return {
"method": request.method,
"url": str(request.url),
"headers": dict(request.headers),
"headers": _safe_get_request_headers(request).copy(),
"cookies": request.cookies,
"query_params": dict(request.query_params),
}
@ -329,7 +330,7 @@ async def vllm_proxy_route(
method=request.method,
endpoint=endpoint,
request_query_params=request.query_params,
request_headers=dict(request.headers),
request_headers=_safe_get_request_headers(request),
stream=request_body.get("stream", False),
content=None,
data=None,
@ -1307,7 +1308,7 @@ async def azure_proxy_route(
method=request.method,
endpoint=endpoint,
request_query_params=request.query_params,
request_headers=dict(request.headers),
request_headers=_safe_get_request_headers(request),
stream=request_body.get("stream", False),
content=None,
data=None,
@ -1505,7 +1506,7 @@ def get_vertex_ai_allowed_incoming_headers(request: Request) -> dict:
Returns:
dict: Headers dictionary with only allowed headers
"""
incoming_headers = dict(request.headers) or {}
incoming_headers = _safe_get_request_headers(request)
headers = {}
for header_name in ALLOWED_VERTEX_AI_PASSTHROUGH_HEADERS:
if header_name in incoming_headers:
@ -1621,7 +1622,7 @@ async def _prepare_vertex_auth_headers(
if (
vertex_credentials is None or vertex_credentials.vertex_project is None
) and router_credentials is None:
headers = dict(request.headers) or {}
headers = _safe_get_request_headers(request).copy()
headers_passed_through = True
verbose_proxy_logger.debug(
"default_vertex_config not set, incoming request headers %s", headers

View file

@ -50,7 +50,10 @@ from litellm.proxy._types import (
)
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
from litellm.proxy.common_utils.http_parsing_utils import _read_request_body
from litellm.proxy.common_utils.http_parsing_utils import (
_read_request_body,
_safe_get_request_headers,
)
from litellm.proxy.utils import get_server_root_path
from litellm.secret_managers.main import get_secret_str
from litellm.types.llms.custom_http import httpxSpecialProvider
@ -649,7 +652,7 @@ async def pass_through_request( # noqa: PLR0915
url = httpx.URL(target)
headers = custom_headers
headers = HttpPassThroughEndpointHelpers.forward_headers_from_request(
request_headers=dict(request.headers),
request_headers=_safe_get_request_headers(request).copy(),
headers=headers,
forward_headers=forward_headers,
)

View file

@ -9,6 +9,7 @@ import secrets
import shutil
import subprocess
import sys
import threading
import time
import traceback
import warnings
@ -293,6 +294,7 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import (
from litellm.proxy.common_utils.html_forms.ui_login import build_ui_login_form
from litellm.proxy.common_utils.http_parsing_utils import (
_read_request_body,
_safe_get_request_headers,
check_file_size_under_limit,
get_form_data,
)
@ -390,6 +392,7 @@ from litellm.proxy.management_endpoints.organization_endpoints import (
router as organization_router,
)
from litellm.proxy.management_endpoints.policy_endpoints import router as policy_router
from litellm.proxy.management_endpoints.usage_endpoints import router as usage_ai_router
from litellm.proxy.management_endpoints.project_endpoints import (
router as project_router,
)
@ -408,6 +411,9 @@ from litellm.proxy.management_endpoints.team_endpoints import (
update_team,
validate_membership,
)
from litellm.proxy.management_endpoints.tool_management_endpoints import (
router as tool_management_router,
)
from litellm.proxy.management_endpoints.ui_sso import (
get_disabled_non_admin_personal_key_creation,
)
@ -658,7 +664,7 @@ _description = (
def cleanup_router_config_variables():
global master_key, user_config_file_path, otel_logging, user_custom_auth, user_custom_auth_path, user_custom_key_generate, user_custom_sso, user_custom_ui_sso_sign_in_handler, use_background_health_checks, use_shared_health_check, health_check_interval, prisma_client
global master_key, user_config_file_path, otel_logging, user_custom_auth, user_custom_auth_path, user_custom_key_generate, user_custom_sso, user_custom_ui_sso_sign_in_handler, use_background_health_checks, use_shared_health_check, health_check_interval, health_check_concurrency, prisma_client
# Set all variables to None
master_key = None
@ -672,6 +678,7 @@ def cleanup_router_config_variables():
use_background_health_checks = None
use_shared_health_check = None
health_check_interval = None
health_check_concurrency = None
prisma_client = None
@ -711,7 +718,7 @@ async def _initialize_shared_aiohttp_session():
try:
from aiohttp import ClientSession, TCPConnector
connector_kwargs = {
connector_kwargs: Dict[str, Any] = {
"keepalive_timeout": AIOHTTP_KEEPALIVE_TIMEOUT,
"ttl_dns_cache": AIOHTTP_TTL_DNS_CACHE,
}
@ -822,7 +829,9 @@ async def proxy_startup_event(app: FastAPI): # noqa: PLR0915
verbose_proxy_logger.debug("About to initialize semantic tool filter")
_config = proxy_config.get_config_state()
_litellm_settings = _config.get("litellm_settings", {})
verbose_proxy_logger.debug(f"litellm_settings keys = {list(_litellm_settings.keys())}")
verbose_proxy_logger.debug(
f"litellm_settings keys = {list(_litellm_settings.keys())}"
)
await ProxyStartupEvent._initialize_semantic_tool_filter(
llm_router=llm_router,
litellm_settings=_litellm_settings,
@ -1468,7 +1477,9 @@ redis_usage_cache: Optional[
RedisCache
] = None # redis cache used for tracking spend, tpm/rpm limits
polling_via_cache_enabled: Union[Literal["all"], List[str], bool] = False
native_background_mode: List[str] = [] # Models that should use native provider background mode instead of polling
native_background_mode: List[
str
] = [] # Models that should use native provider background mode instead of polling
polling_cache_ttl: int = 3600 # Default 1 hour TTL for polling cache
user_custom_auth = None
user_custom_key_generate = None
@ -1478,8 +1489,11 @@ use_background_health_checks = None
use_shared_health_check = None
use_queue = False
health_check_interval = None
health_check_concurrency = None
health_check_details = None
health_check_results: Dict[str, Union[int, List[Dict[str, Any]]]] = {}
background_health_check_loop_active = False
background_health_check_cycle_seq = 0
queue: List = []
litellm_proxy_budget_name = "litellm-proxy-budget"
litellm_proxy_admin_name = LITELLM_PROXY_ADMIN_NAME
@ -1759,8 +1773,14 @@ async def update_cache( # noqa: PLR0915
("{}:spend".format(litellm_proxy_admin_name), increment)
)
except Exception as e:
verbose_proxy_logger.debug(
f"An error occurred updating user cache: {str(e)}\n\n{traceback.format_exc()}"
verbose_proxy_logger.warning(
"Spend tracking - failed to update user spend in cache. "
"Budget enforcement may use stale spend values. "
"user_id=%s, response_cost=%s - %s\n%s",
user_id,
response_cost,
str(e),
traceback.format_exc(),
)
### UPDATE END-USER SPEND ###
@ -1797,8 +1817,14 @@ async def update_cache( # noqa: PLR0915
existing_spend_obj.spend = new_spend
values_to_update_in_cache.append((_id, existing_spend_obj.json()))
except Exception as e:
verbose_proxy_logger.exception(
f"An error occurred updating end user cache: {str(e)}"
verbose_proxy_logger.warning(
"Spend tracking - failed to update end user spend in cache. "
"Budget enforcement may use stale spend values. "
"end_user_id=%s, response_cost=%s - %s\n%s",
end_user_id,
response_cost,
str(e),
traceback.format_exc(),
)
### UPDATE TEAM SPEND ###
@ -1839,8 +1865,14 @@ async def update_cache( # noqa: PLR0915
existing_spend_obj.spend = new_spend
values_to_update_in_cache.append((_id, existing_spend_obj))
except Exception as e:
verbose_proxy_logger.exception(
f"An error occurred updating end user cache: {str(e)}"
verbose_proxy_logger.warning(
"Spend tracking - failed to update team spend in cache. "
"Budget enforcement may use stale spend values. "
"team_id=%s, response_cost=%s - %s\n%s",
team_id,
response_cost,
str(e),
traceback.format_exc(),
)
### UPDATE TAG SPEND ###
@ -1885,8 +1917,14 @@ async def update_cache( # noqa: PLR0915
existing_tag_obj.spend = new_spend
values_to_update_in_cache.append((cache_key, existing_tag_obj))
except Exception as e:
verbose_proxy_logger.exception(
f"An error occurred updating tag cache: {str(e)}"
verbose_proxy_logger.warning(
"Spend tracking - failed to update tag spend in cache. "
"Budget enforcement may use stale spend values. "
"tags=%s, response_cost=%s - %s\n%s",
tags,
response_cost,
str(e),
traceback.format_exc(),
)
if token is not None and response_cost is not None:
@ -1927,6 +1965,88 @@ def run_ollama_serve():
)
def _get_process_rss_mb() -> Optional[float]:
"""
Get process RSS memory in MB.
On Linux, ru_maxrss is in KB. On macOS, ru_maxrss is in bytes.
"""
try:
import resource
ru_maxrss = resource.getrusage(resource.RUSAGE_SELF).ru_maxrss
if sys.platform == "darwin":
return float(ru_maxrss) / (1024 * 1024)
return float(ru_maxrss) / 1024
except Exception:
return None
def _rss_mb_for_log() -> str:
rss_mb = _get_process_rss_mb()
if rss_mb is None:
return "unknown"
return f"{rss_mb:.2f}"
async def _run_direct_health_check_with_instrumentation(
model_list: list,
details: Optional[bool],
max_concurrency: Optional[int],
instrumentation_context: dict,
):
try:
return await perform_health_check(
model_list=model_list,
details=details,
max_concurrency=max_concurrency,
instrumentation_context=instrumentation_context,
)
except TypeError as e:
if "instrumentation_context" not in str(e):
raise
# Backward compatibility for monkeypatched or wrapped callables
# that do not accept instrumentation_context.
return await perform_health_check(
model_list=model_list,
details=details,
max_concurrency=max_concurrency,
)
def _schedule_background_health_check_db_save(
prisma_client,
shared_health_manager,
model_list: list,
healthy_endpoints: list,
unhealthy_endpoints: list,
):
"""Fire-and-forget: persist health check results to DB if prisma is available."""
if prisma_client is None:
return
import time as time_module
from litellm.proxy.health_endpoints._health_endpoints import (
_save_background_health_checks_to_db,
)
checked_by = (
shared_health_manager.pod_id
if shared_health_manager is not None
else "background_health_check"
)
start_time = time_module.time()
asyncio.create_task(
_save_background_health_checks_to_db(
prisma_client,
model_list,
healthy_endpoints,
unhealthy_endpoints,
start_time,
checked_by=checked_by,
)
)
async def _run_background_health_check():
"""
Periodically run health checks in the background on the endpoints.
@ -1934,7 +2054,10 @@ async def _run_background_health_check():
Update health_check_results, based on this.
Uses shared health check state when Redis is available to coordinate across pods.
"""
global health_check_results, llm_model_list, health_check_interval, health_check_details, use_shared_health_check, redis_usage_cache, prisma_client
global health_check_results, llm_model_list, health_check_interval
global health_check_concurrency, health_check_details, use_shared_health_check
global redis_usage_cache, prisma_client
global background_health_check_loop_active, background_health_check_cycle_seq
if (
health_check_interval is None
@ -1943,6 +2066,24 @@ async def _run_background_health_check():
):
return
if background_health_check_loop_active:
verbose_proxy_logger.warning(
"background_health_check_loop_overlap_detected existing_loop_active=true interval_seconds=%s max_concurrency=%s shared=%s",
health_check_interval,
health_check_concurrency,
use_shared_health_check,
)
background_health_check_loop_active = True
verbose_proxy_logger.info(
"background_health_check_loop_started interval_seconds=%s max_concurrency=%s shared=%s details=%s thread_count=%d rss_mb=%s",
health_check_interval,
health_check_concurrency,
use_shared_health_check,
health_check_details,
threading.active_count(),
_rss_mb_for_log(),
)
# Initialize shared health check manager if Redis is available and feature is enabled
shared_health_manager = None
if use_shared_health_check and redis_usage_cache is not None:
@ -1958,8 +2099,13 @@ async def _run_background_health_check():
verbose_proxy_logger.info("Initialized shared health check manager")
while True:
background_health_check_cycle_seq += 1
cycle_id = f"bg-{background_health_check_cycle_seq}"
cycle_start_time = time.monotonic()
# make 1 deep copy of llm_model_list on every health check iteration
_llm_model_list = copy.deepcopy(llm_model_list) or []
model_count_total = len(_llm_model_list)
# filter out models that have disabled background health checks
_llm_model_list = [
@ -1967,6 +2113,33 @@ async def _run_background_health_check():
for m in _llm_model_list
if not m.get("model_info", {}).get("disable_background_health_check", False)
]
model_count_enabled = len(_llm_model_list)
expected_peak_in_flight = model_count_enabled
if (
isinstance(health_check_concurrency, int)
and health_check_concurrency > 0
and model_count_enabled > 0
):
expected_peak_in_flight = min(model_count_enabled, health_check_concurrency)
verbose_proxy_logger.debug(
"background_health_check_cycle_start cycle_id=%s model_count_total=%d model_count_enabled=%d interval_seconds=%s max_concurrency=%s expected_peak_in_flight=%d shared=%s thread_count=%d rss_mb=%s",
cycle_id,
model_count_total,
model_count_enabled,
health_check_interval,
health_check_concurrency,
expected_peak_in_flight,
shared_health_manager is not None,
threading.active_count(),
_rss_mb_for_log(),
)
instrumentation_context = {
"enabled": True,
"source": "proxy_background_loop",
"cycle_id": cycle_id,
}
# Use shared health check if available, otherwise fall back to direct health check
# Convert health_check_details to bool for perform_shared_health_check (defaults to True if None)
@ -1980,19 +2153,31 @@ async def _run_background_health_check():
healthy_endpoints,
unhealthy_endpoints,
) = await shared_health_manager.perform_shared_health_check(
model_list=_llm_model_list, details=details_bool
model_list=_llm_model_list,
details=details_bool,
max_concurrency=health_check_concurrency,
)
except Exception as e:
verbose_proxy_logger.error(
"Error in shared health check, falling back to direct health check: %s",
str(e),
)
healthy_endpoints, unhealthy_endpoints = await perform_health_check(
model_list=_llm_model_list, details=health_check_details
healthy_endpoints, unhealthy_endpoints = (
await _run_direct_health_check_with_instrumentation(
_llm_model_list,
health_check_details,
health_check_concurrency,
instrumentation_context,
)
)
else:
healthy_endpoints, unhealthy_endpoints = await perform_health_check(
model_list=_llm_model_list, details=health_check_details
healthy_endpoints, unhealthy_endpoints = (
await _run_direct_health_check_with_instrumentation(
_llm_model_list,
health_check_details,
health_check_concurrency,
instrumentation_context,
)
)
# Update the global variable with the health check results
@ -2000,34 +2185,34 @@ async def _run_background_health_check():
health_check_results["unhealthy_endpoints"] = unhealthy_endpoints
health_check_results["healthy_count"] = len(healthy_endpoints)
health_check_results["unhealthy_count"] = len(unhealthy_endpoints)
cycle_duration_ms = (time.monotonic() - cycle_start_time) * 1000
verbose_proxy_logger.debug(
"background_health_check_cycle_complete cycle_id=%s model_count_enabled=%d healthy_count=%d unhealthy_count=%d duration_ms=%.2f interval_seconds=%s thread_count=%d rss_mb=%s",
cycle_id,
model_count_enabled,
len(healthy_endpoints),
len(unhealthy_endpoints),
cycle_duration_ms,
health_check_interval,
threading.active_count(),
_rss_mb_for_log(),
)
if cycle_duration_ms > (health_check_interval * 1000):
verbose_proxy_logger.warning(
"background_health_check_cycle_duration_exceeded_interval cycle_id=%s duration_ms=%.2f interval_seconds=%s",
cycle_id,
cycle_duration_ms,
health_check_interval,
)
# Save background health checks to database (non-blocking)
if prisma_client is not None:
import time as time_module
from litellm.proxy.health_endpoints._health_endpoints import (
_save_background_health_checks_to_db,
)
# Use pod_id or a system identifier for checked_by if shared health check is enabled
checked_by = None
if shared_health_manager is not None:
checked_by = shared_health_manager.pod_id
else:
# Use a system identifier for background health checks
checked_by = "background_health_check"
start_time = time_module.time()
asyncio.create_task(
_save_background_health_checks_to_db(
prisma_client,
_llm_model_list,
healthy_endpoints,
unhealthy_endpoints,
start_time,
checked_by=checked_by,
)
)
_schedule_background_health_check_db_save(
prisma_client,
shared_health_manager,
_llm_model_list,
healthy_endpoints,
unhealthy_endpoints,
)
await asyncio.sleep(health_check_interval)
@ -2480,7 +2665,7 @@ class ProxyConfig:
"""
Load config values into proxy global state
"""
global master_key, user_config_file_path, otel_logging, user_custom_auth, user_custom_auth_path, user_custom_key_generate, user_custom_sso, user_custom_ui_sso_sign_in_handler, use_background_health_checks, use_shared_health_check, health_check_interval, use_queue, proxy_budget_rescheduler_max_time, proxy_budget_rescheduler_min_time, ui_access_mode, litellm_master_key_hash, proxy_batch_write_at, disable_spend_logs, prompt_injection_detection_obj, redis_usage_cache, store_model_in_db, premium_user, open_telemetry_logger, health_check_details, proxy_batch_polling_interval, config_passthrough_endpoints
global master_key, user_config_file_path, otel_logging, user_custom_auth, user_custom_auth_path, user_custom_key_generate, user_custom_sso, user_custom_ui_sso_sign_in_handler, use_background_health_checks, use_shared_health_check, health_check_interval, health_check_concurrency, use_queue, proxy_budget_rescheduler_max_time, proxy_budget_rescheduler_min_time, ui_access_mode, litellm_master_key_hash, proxy_batch_write_at, disable_spend_logs, prompt_injection_detection_obj, redis_usage_cache, store_model_in_db, premium_user, open_telemetry_logger, health_check_details, proxy_batch_polling_interval, config_passthrough_endpoints
config: dict = await self.get_config(config_file_path=config_file_path)
@ -2905,7 +3090,18 @@ class ProxyConfig:
health_check_interval = general_settings.get(
"health_check_interval", DEFAULT_HEALTH_CHECK_INTERVAL
)
health_check_concurrency = general_settings.get(
"health_check_concurrency", None
)
health_check_details = general_settings.get("health_check_details", True)
verbose_proxy_logger.info(
"background_health_check_config enabled=%s shared=%s interval_seconds=%s max_concurrency=%s details=%s",
use_background_health_checks,
use_shared_health_check,
health_check_interval,
health_check_concurrency,
health_check_details,
)
### RBAC ###
rbac_role_permissions = general_settings.get("role_permissions", None)
@ -2999,7 +3195,7 @@ class ProxyConfig:
for k, v in router_settings.items():
if k in available_args:
router_params[k] = v
elif k == "health_check_interval":
elif k in {"health_check_interval", "health_check_concurrency"}:
raise ValueError(
f"'{k}' is NOT a valid router_settings parameter. Please move it to 'general_settings'."
)
@ -4201,9 +4397,7 @@ class ProxyConfig:
)
if self._should_load_db_object(object_type="semantic_filter_settings"):
await self._init_semantic_filter_settings_in_db(
prisma_client=prisma_client
)
await self._init_semantic_filter_settings_in_db(prisma_client=prisma_client)
async def _init_semantic_filter_settings_in_db(self, prisma_client: PrismaClient):
"""
@ -5259,30 +5453,38 @@ class ProxyStartupEvent:
):
"""Initialize MCP semantic tool filter if configured"""
from litellm.proxy.hooks.mcp_semantic_filter import SemanticToolFilterHook
mcp_semantic_filter_config = litellm_settings.get("mcp_semantic_tool_filter", None)
mcp_semantic_filter_config = litellm_settings.get(
"mcp_semantic_tool_filter", None
)
# Only proceed if the feature is configured and enabled
if not mcp_semantic_filter_config or not mcp_semantic_filter_config.get("enabled", False):
verbose_proxy_logger.debug("Semantic tool filter not configured or not enabled, skipping initialization")
if not mcp_semantic_filter_config or not mcp_semantic_filter_config.get(
"enabled", False
):
verbose_proxy_logger.debug(
"Semantic tool filter not configured or not enabled, "
"skipping initialization"
)
return
verbose_proxy_logger.debug(
f"Initializing semantic tool filter: llm_router={llm_router is not None}, "
f"config={mcp_semantic_filter_config}"
)
hook = await SemanticToolFilterHook.initialize_from_config(
config=mcp_semantic_filter_config,
llm_router=llm_router,
)
if hook:
verbose_proxy_logger.debug("Semantic tool filter hook registered")
litellm.logging_callback_manager.add_litellm_callback(hook)
else:
# Only warn if the feature was configured but failed to initialize
verbose_proxy_logger.warning("Semantic tool filter hook was configured but failed to initialize")
verbose_proxy_logger.warning(
"Semantic tool filter hook was configured but failed to initialize"
)
@classmethod
def _initialize_jwt_auth(
@ -8706,7 +8908,8 @@ async def _apply_search_filter_to_models(
# Fetch database models if we need more for the current page
if router_models_count < models_needed_for_page:
models_to_fetch = min(
models_needed_for_page - router_models_count, db_models_total_count
models_needed_for_page - router_models_count,
db_models_total_count,
)
if models_to_fetch > 0:
@ -8742,21 +8945,21 @@ async def _apply_search_filter_to_models(
def _normalize_datetime_for_sorting(dt: Any) -> Optional[datetime]:
"""
Normalize a datetime value to a timezone-aware UTC datetime for sorting.
This function handles:
- None values: returns None
- String values: parses ISO format strings and converts to UTC-aware datetime
- Datetime objects: converts naive datetimes to UTC-aware, and aware datetimes to UTC
Args:
dt: Datetime value (None, str, or datetime object)
Returns:
UTC-aware datetime object, or None if input is None or cannot be parsed
"""
if dt is None:
return None
if isinstance(dt, str):
try:
# Handle ISO format strings, including 'Z' suffix
@ -8770,14 +8973,14 @@ def _normalize_datetime_for_sorting(dt: Any) -> Optional[datetime]:
return parsed_dt
except (ValueError, AttributeError):
return None
if isinstance(dt, datetime):
# If naive, assume UTC and make it aware
if dt.tzinfo is None:
return dt.replace(tzinfo=timezone.utc)
# If aware, convert to UTC
return dt.astimezone(timezone.utc)
return None
@ -8797,46 +9000,60 @@ def _sort_models(
Returns:
Sorted list of models
"""
if not sort_by or sort_by not in ["model_name", "created_at", "updated_at", "costs", "status"]:
if not sort_by or sort_by not in [
"model_name",
"created_at",
"updated_at",
"costs",
"status",
]:
return all_models
reverse = sort_order.lower() == "desc"
def get_sort_key(model: Dict[str, Any]) -> Any:
model_info = model.get("model_info", {})
if sort_by == "model_name":
return model.get("model_name", "").lower()
elif sort_by == "created_at":
created_at = model_info.get("created_at")
normalized_dt = _normalize_datetime_for_sorting(created_at)
if normalized_dt is None:
# Put None values at the end for asc, at the start for desc
return (datetime.max.replace(tzinfo=timezone.utc) if not reverse else datetime.min.replace(tzinfo=timezone.utc))
return (
datetime.max.replace(tzinfo=timezone.utc)
if not reverse
else datetime.min.replace(tzinfo=timezone.utc)
)
return normalized_dt
elif sort_by == "updated_at":
updated_at = model_info.get("updated_at")
normalized_dt = _normalize_datetime_for_sorting(updated_at)
if normalized_dt is None:
return (datetime.max.replace(tzinfo=timezone.utc) if not reverse else datetime.min.replace(tzinfo=timezone.utc))
return (
datetime.max.replace(tzinfo=timezone.utc)
if not reverse
else datetime.min.replace(tzinfo=timezone.utc)
)
return normalized_dt
elif sort_by == "costs":
input_cost = model_info.get("input_cost_per_token", 0) or 0
output_cost = model_info.get("output_cost_per_token", 0) or 0
total_cost = input_cost + output_cost
# Put 0 or None costs at the end for asc, at the start for desc
if total_cost == 0:
return (float("inf") if not reverse else float("-inf"))
return float("inf") if not reverse else float("-inf")
return total_cost
elif sort_by == "status":
# False (config) comes before True (db) for asc
db_model = model_info.get("db_model", False)
return db_model
return None
try:
@ -9032,9 +9249,7 @@ async def _find_model_by_id(
)
if db_model:
# Convert database model to router format
decrypted_models = proxy_config.decrypt_model_list_from_db(
[db_model]
)
decrypted_models = proxy_config.decrypt_model_list_from_db([db_model])
if decrypted_models:
found_model = decrypted_models[0]
except Exception as e:
@ -9208,13 +9423,13 @@ async def model_info_v2(
)
verbose_proxy_logger.debug("all_models: %s", all_models)
# Append A2A agents to models list
all_models = await append_agents_to_model_info(
models=all_models,
user_api_key_dict=user_api_key_dict,
)
# Update total count to include agents
search_total_count = len(all_models)
@ -10057,7 +10272,7 @@ async def model_group_info(
model_groups: List[ModelGroupInfoProxy] = _get_model_group_info(
llm_router=llm_router, all_models_str=all_models_str, model_group=model_group
)
# Append A2A agents to model groups
model_groups = await append_agents_to_model_group(
model_groups=model_groups,
@ -10238,7 +10453,7 @@ async def async_queue_request(
data["proxy_server_request"] = {
"url": str(request.url),
"method": request.method,
"headers": dict(request.headers),
"headers": _safe_get_request_headers(request).copy(),
"body": copy.copy(data), # use copy instead of deepcopy
}
@ -10259,7 +10474,7 @@ async def async_queue_request(
data["metadata"] = {}
data["metadata"]["user_api_key"] = user_api_key_dict.api_key
data["metadata"]["user_api_key_metadata"] = user_api_key_dict.metadata
_headers = dict(request.headers)
_headers = _safe_get_request_headers(request).copy()
_headers.pop(
"authorization", None
) # do not store the original `sk-..` api key in the db
@ -12658,6 +12873,7 @@ app.include_router(caching_router)
app.include_router(analytics_router)
app.include_router(guardrails_router)
app.include_router(policy_router)
app.include_router(usage_ai_router)
app.include_router(policy_crud_router)
app.include_router(policy_resolve_router)
app.include_router(search_tool_management_router)
@ -12671,6 +12887,7 @@ app.include_router(budget_management_router)
app.include_router(model_management_router)
app.include_router(model_access_group_management_router)
app.include_router(tag_management_router)
app.include_router(tool_management_router)
app.include_router(cost_tracking_settings_router)
app.include_router(router_settings_router)
app.include_router(fallback_management_router)

View file

@ -64,6 +64,8 @@ model LiteLLM_AgentsTable {
litellm_params Json?
agent_card_params Json
agent_access_groups String[] @default([])
object_permission_id String?
object_permission LiteLLM_ObjectPermissionTable? @relation(fields: [object_permission_id], references: [object_permission_id])
created_at DateTime @default(now()) @map("created_at")
created_by String
updated_at DateTime @default(now()) @updatedAt @map("updated_at")
@ -264,6 +266,7 @@ model LiteLLM_ObjectPermissionTable {
organizations LiteLLM_OrganizationTable[]
users LiteLLM_UserTable[]
end_users LiteLLM_EndUserTable[]
agents_table LiteLLM_AgentsTable[]
}
// Holds the MCP server configuration
@ -314,6 +317,7 @@ model LiteLLM_VerificationToken {
router_settings Json? @default("{}")
user_id String?
team_id String?
agent_id String?
project_id String?
permissions Json @default("{}")
max_parallel_requests Int?
@ -1051,6 +1055,26 @@ model LiteLLM_PolicyAttachmentTable {
updated_by String?
}
// Global tool registry - auto-discovered from LLM responses; admins set call_policy here
model LiteLLM_ToolTable {
tool_id String @id @default(uuid())
tool_name String @unique // e.g. "huggingface_remote-mcp__dynamic_space"
origin String? // MCP server name or "user_defined"
call_policy String @default("untrusted") // "trusted" | "untrusted" | "dual_llm" | "blocked"
call_count Int @default(0) // cumulative number of times this tool was seen
assignments Json? @default("{}")
key_hash String? // hash of the virtual key that first called this tool
team_id String? // team that first called this tool
key_alias String? // human-readable alias of the virtual key
created_at DateTime @default(now())
created_by String?
updated_at DateTime @default(now()) @updatedAt
updated_by String?
@@index([call_policy])
@@index([team_id])
}
//Unified Access Groups table for storing unified access groups
model LiteLLM_AccessGroupTable {
access_group_id String @id @default(uuid())

View file

@ -4457,6 +4457,11 @@ class ProxyUpdateSpend:
len(logs_to_process) :
]
popped_batch = True
if len(logs_to_process) > 0:
verbose_proxy_logger.info(
"Spend tracking - processing %d spend logs for DB write",
len(logs_to_process),
)
start_time = time.time()
try:
for i in range(n_retry_times + 1):
@ -4503,9 +4508,17 @@ class ProxyUpdateSpend:
f"{len(logs_to_process)} logs processed. Remaining in queue: {remaining_count}"
)
break
except DB_CONNECTION_ERROR_TYPES:
except DB_CONNECTION_ERROR_TYPES as e:
if i is None:
i = 0
verbose_proxy_logger.warning(
"Spend tracking - DB connection error writing spend logs, "
"retry %d/%d. logs_count=%d, error=%s",
i + 1,
n_retry_times,
len(logs_to_process),
str(e),
)
if i >= n_retry_times:
raise
await asyncio.sleep(2**i)
@ -4620,8 +4633,8 @@ async def update_spend_logs_job(
logs_to_process=logs_to_process,
)
except Exception as guardrail_tracking_err:
verbose_proxy_logger.debug(
"Guardrail usage tracking failed (non-fatal): %s",
verbose_proxy_logger.warning(
"Spend tracking - guardrail usage tracking failed (non-fatal): %s",
guardrail_tracking_err,
)

View file

@ -19,6 +19,7 @@ from fastapi import APIRouter, Request, Response
import litellm
from litellm.proxy._types import *
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_utils.http_parsing_utils import _safe_get_request_headers
from litellm.proxy.litellm_pre_call_utils import _get_dynamic_logging_metadata
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
create_pass_through_route,
@ -32,7 +33,7 @@ def create_request_copy(request: Request):
return {
"method": request.method,
"url": str(request.url),
"headers": dict(request.headers),
"headers": _safe_get_request_headers(request).copy(),
"cookies": request.cookies,
"query_params": dict(request.query_params),
}

View file

@ -365,6 +365,7 @@ def rerank( # noqa: PLR0915
max_chunks_per_doc=max_chunks_per_doc,
_is_async=_is_async,
optional_params=optional_params.model_dump(exclude_unset=True),
timeout=optional_params.timeout,
api_base=api_base,
extra_headers=merged_headers,
logging_obj=litellm_logging_obj,

View file

@ -163,7 +163,11 @@ from litellm.types.utils import (
)
from litellm.types.utils import ModelInfo
from litellm.types.utils import ModelInfo as ModelMapInfo
from litellm.types.utils import ModelResponseStream, StandardLoggingPayload, Usage
from litellm.types.utils import (
ModelResponseStream,
StandardLoggingPayload,
Usage,
)
from litellm.utils import (
CustomStreamWrapper,
EmbeddingResponse,
@ -1555,6 +1559,9 @@ class Router:
logging_obj=model_response.logging_obj,
)
self._async_generator = async_generator
# Preserve hidden params (including litellm_overhead_time_ms) from original response
if hasattr(model_response, "_hidden_params"):
self._hidden_params = model_response._hidden_params.copy()
def __aiter__(self):
return self
@ -6416,6 +6423,7 @@ class Router:
self.model_list = []
self.model_id_to_deployment_index_map = {} # Reset the index
self.model_name_to_deployment_indices = {} # Reset the model_name index
self._invalidate_model_group_info_cache()
self._invalidate_access_groups_cache()
# we add api_base/api_key each model so load balancing between azure/gpt on api_base1 and api_base2 works
@ -6726,6 +6734,7 @@ class Router:
"""
idx = len(self.model_list)
self.model_list.append(model)
self._invalidate_model_group_info_cache()
self._invalidate_access_groups_cache()
# Update model_id index for O(1) lookup
@ -6774,6 +6783,7 @@ class Router:
if removal_idx is not None:
self.model_list.pop(removal_idx)
self._invalidate_model_group_info_cache()
self._invalidate_access_groups_cache()
self._update_deployment_indices_after_removal(
model_id=deployment_id, removal_idx=removal_idx
@ -6808,6 +6818,7 @@ class Router:
if deployment_idx is not None:
# Pop the item from the list first
item = self.model_list.pop(deployment_idx)
self._invalidate_model_group_info_cache()
self._invalidate_access_groups_cache()
self._update_deployment_indices_after_removal(
model_id=id, removal_idx=deployment_idx
@ -6978,9 +6989,9 @@ class Router:
raise ValueError("Deployment not found")
## GET BASE MODEL
base_model = deployment.get("model_info", {}).get("base_model", None)
base_model = (deployment.get("model_info") or {}).get("base_model", None)
if base_model is None:
base_model = deployment.get("litellm_params", {}).get("base_model", None)
base_model = (deployment.get("litellm_params") or {}).get("base_model", None)
model = base_model
@ -6995,7 +7006,7 @@ class Router:
raise ValueError(
f"Deployment missing valid litellm_params. "
f"Got: {type(litellm_params_data).__name__}, "
f"deployment_id: {deployment.get('model_info', {}).get('id', 'unknown')}"
f"deployment_id: {(deployment.get('model_info') or {}).get('id', 'unknown')}"
)
_model, custom_llm_provider, _, _ = litellm.get_llm_provider(
model=litellm_params.model,
@ -7015,10 +7026,10 @@ class Router:
if potential_models is not None:
for potential_model in potential_models:
try:
if potential_model.get("model_info", {}).get(
if (potential_model.get("model_info") or {}).get(
"id"
) == deployment.get("model_info", {}).get("id"):
model = potential_model.get("litellm_params", {}).get(
) == (deployment.get("model_info") or {}).get("id"):
model = (potential_model.get("litellm_params") or {}).get(
"model"
)
break
@ -7039,9 +7050,10 @@ class Router:
model_info = litellm.get_model_info(model=model_info_name)
## CHECK USER SET MODEL INFO
user_model_info = deployment.get("model_info", {})
user_model_info = deployment.get("model_info") or {}
model_info.update(user_model_info)
if model_info is not None:
model_info.update(user_model_info)
return model_info
@ -7568,6 +7580,7 @@ class Router:
"""
# First populate the model_list
self.model_list = []
self._invalidate_model_group_info_cache()
self._invalidate_access_groups_cache()
for _, model in enumerate(model_list):
# Extract model_info from the model dict
@ -7912,6 +7925,13 @@ class Router:
return returned_models
def _invalidate_model_group_info_cache(self) -> None:
"""Invalidate the cached model group info.
Call this whenever self.model_list is modified to ensure the cache is rebuilt.
"""
self._cached_get_model_group_info.cache_clear()
def _invalidate_access_groups_cache(self) -> None:
"""Invalidate the cached access groups.

View file

@ -55,22 +55,21 @@ def _sanitize_prometheus_label_value(value: Optional[Any]) -> Optional[str]:
return None
# Coerce non-string values (int, bool, etc.) to str before sanitizing
if not isinstance(value, str):
value = str(value)
str_value: str = value if isinstance(value, str) else str(value)
# Remove Unicode line/paragraph separators that break text format
value = value.replace("\u2028", "").replace("\u2029", "")
str_value = str_value.replace("\u2028", "").replace("\u2029", "")
# Remove carriage returns
value = value.replace("\r", "")
str_value = str_value.replace("\r", "")
# Replace newlines with spaces
value = value.replace("\n", " ")
str_value = str_value.replace("\n", " ")
# Escape backslashes and double quotes per Prometheus exposition format
value = value.replace("\\", "\\\\").replace('"', '\\"')
str_value = str_value.replace("\\", "\\\\").replace('"', '\\"')
return value
return str_value
@dataclass
@ -185,6 +184,7 @@ class UserAPIKeyLabelNames(Enum):
CLIENT_IP = "client_ip"
USER_AGENT = "user_agent"
CALLBACK_NAME = "callback_name"
STREAM = "stream"
DEFINED_PROMETHEUS_METRICS = Literal[
@ -638,6 +638,14 @@ class PrometheusMetricLabels:
]
)
# Conditionally add stream label to litellm_proxy_total_requests_metric
if (
label_name == "litellm_proxy_total_requests_metric"
and litellm.prometheus_emit_stream_label is True
and UserAPIKeyLabelNames.STREAM.value not in default_labels
):
custom_labels.append(UserAPIKeyLabelNames.STREAM.value)
return default_labels + custom_labels
@ -709,6 +717,9 @@ class UserAPIKeyLabelValues(BaseModel):
user_agent: Annotated[
Optional[str], Field(..., alias=UserAPIKeyLabelNames.USER_AGENT.value)
] = None
stream: Annotated[
Optional[str], Field(..., alias=UserAPIKeyLabelNames.STREAM.value)
] = None
class PrometheusMetricsConfig(BaseModel):

View file

@ -1125,6 +1125,7 @@ class ResponsesAPIOptionalRequestParams(TypedDict, total=False):
prompt: Optional[PromptObject]
max_tool_calls: Optional[int]
prompt_cache_key: Optional[str]
prompt_cache_retention: Optional[str]
stream_options: Optional[dict]
top_logprobs: Optional[int]
partial_images: Optional[

View file

@ -6,6 +6,7 @@ from typing_extensions import Any, List, Optional, TypedDict
from litellm.types.llms.base import BaseLiteLLMOpenAIResponseObject
Phase = Optional[Literal["commentary", "final_answer"]] # TODO: Once openai sdk has updated, we can remove this and use the openai sdk type
class GenericResponseOutputItemContentAnnotation(BaseLiteLLMOpenAIResponseObject):
"""Annotation for content in a message"""
@ -35,6 +36,7 @@ class OutputFunctionToolCall(BaseLiteLLMOpenAIResponseObject):
type: Optional[str] # "function_call"
id: Optional[str]
status: Literal["in_progress", "completed", "incomplete"]
phase: Phase = None
class OutputImageGenerationCall(BaseLiteLLMOpenAIResponseObject):
@ -57,6 +59,7 @@ class GenericResponseOutputItem(BaseLiteLLMOpenAIResponseObject):
status: str # "completed", "in_progress", etc.
role: str # "assistant", "user", etc.
content: List[OutputText]
phase: Phase = None
class DeleteResponseResult(BaseLiteLLMOpenAIResponseObject):

View file

@ -0,0 +1,42 @@
"""
Pydantic models for Tool Policy management endpoints.
"""
from datetime import datetime
from typing import Dict, List, Literal, Optional
from pydantic import BaseModel
ToolCallPolicy = Literal["trusted", "untrusted", "dual_llm", "blocked"]
class LiteLLM_ToolTableRow(BaseModel):
tool_id: str
tool_name: str
origin: Optional[str] = None
call_policy: ToolCallPolicy = "untrusted"
call_count: int = 0
assignments: Optional[Dict] = None
key_hash: Optional[str] = None
team_id: Optional[str] = None
key_alias: Optional[str] = None
created_at: Optional[datetime] = None
updated_at: Optional[datetime] = None
created_by: Optional[str] = None
updated_by: Optional[str] = None
class ToolListResponse(BaseModel):
tools: List[LiteLLM_ToolTableRow]
total: int
class ToolPolicyUpdateRequest(BaseModel):
tool_name: str
call_policy: ToolCallPolicy
class ToolPolicyUpdateResponse(BaseModel):
tool_name: str
call_policy: ToolCallPolicy
updated: bool

View file

@ -1826,7 +1826,7 @@ class ModelResponse(ModelResponseBase):
else:
usage = usage
elif stream is None or stream is False:
usage = Usage()
usage = None # avoid constructing throwaway Usage; set by convert_to_model_response_object
if hidden_params:
self._hidden_params = hidden_params

View file

@ -20562,6 +20562,39 @@
"supports_tool_choice": true,
"supports_vision": true
},
"gpt-5.3-codex": {
"cache_read_input_token_cost": 1.75e-07,
"cache_read_input_token_cost_priority": 3.5e-07,
"input_cost_per_token": 1.75e-06,
"input_cost_per_token_priority": 3.5e-06,
"litellm_provider": "openai",
"max_input_tokens": 272000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "responses",
"output_cost_per_token": 1.4e-05,
"output_cost_per_token_priority": 2.8e-05,
"supported_endpoints": [
"/v1/responses"
],
"supported_modalities": [
"text",
"image"
],
"supported_output_modalities": [
"text"
],
"supports_function_calling": true,
"supports_native_streaming": true,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_system_messages": false,
"supports_tool_choice": true,
"supports_vision": true
},
"gpt-5-mini": {
"cache_read_input_token_cost": 2.5e-08,
"cache_read_input_token_cost_flex": 1.25e-08,

View file

@ -12,7 +12,19 @@
},
"overrides": {
"glob": ">=11.1.0",
"tar": ">=7.5.7",
"@isaacs/brace-expansion": ">=5.0.1"
"tar": ">=7.5.8",
"minimatch": ">=10.2.1",
"diff": ">=8.0.3",
"@isaacs/brace-expansion": ">=5.0.1",
"@babel/traverse": ">=7.23.2",
"ws": ">=7.5.10",
"http-proxy-middleware": ">=2.0.9",
"tar-fs": ">=2.1.4",
"webpack-dev-middleware": ">=5.3.4",
"braces": ">=3.0.3",
"axios": ">=0.30.2",
"webpack": ">=5.94.0",
"serve-static": ">=1.16.0",
"path-to-regexp": ">=0.1.12"
}
}
}

View file

@ -3,6 +3,8 @@
urllib3>=2.6.0 # CVE-2025-66471, CVE-2025-66418, CVE-2026-21441
tornado>=6.5.3 # CVE-2025-67725, CVE-2025-67726, CVE-2025-67724
filelock>=3.20.1 # CVE-2025-68146
h11>=0.16.0 # CVE-2025-43859, GHSA-vqfr-h8mv-ghfj — HTTP request smuggling
wheel>=0.46.2 # CVE-2026-24049 — path traversal
Pillow==12.1.1 #GHSA-cfh3-3jmp-rvhc
cryptography==46.0.5 #GHSA-r6ph-v2qm-q3c2
@ -21,7 +23,7 @@ boto3==1.40.53 # aws bedrock/sagemaker calls (has bedrock-agentcore-control, com
redis==5.2.1 # redis caching
redisvl==0.4.1 ## redis semantic caching
prisma==0.11.0 # for db
nodejs-wheel-binaries==24.12.0 ## required by prisma for migrations, prevents runtime download (updated from nodejs-bin for security fixes)
nodejs-wheel-binaries==24.13.1 ## required by prisma for migrations, prevents runtime download (updated from nodejs-bin for security fixes)
mangum==0.17.0 # for aws lambda functions
pynacl==1.6.2 # for encrypting keys
google-cloud-aiplatform==1.133.0 # for vertex ai calls

View file

@ -92,7 +92,7 @@ async def test_azure_img_gen_health_check():
litellm._turn_on_debug()
max_retries = 3
retry_delay = 1 # Start with 1 second delay
for attempt in range(max_retries):
response = await litellm.ahealth_check(
model_params={
@ -103,11 +103,11 @@ async def test_azure_img_gen_health_check():
mode="image_generation",
prompt="cute baby sea otter",
)
# Check if response is successful (no error)
if isinstance(response, dict) and "error" not in response:
return response
# Check if error is a transient Azure internal server error
error_str = str(response.get("error", "")).lower()
is_transient_error = (
@ -116,16 +116,18 @@ async def test_azure_img_gen_health_check():
or "internalfailure" in error_str
or "internal failure" in error_str
)
# If it's the last attempt or not a transient error, fail the test
if attempt == max_retries - 1 or not is_transient_error:
assert isinstance(response, dict) and "error" not in response, f"Health check failed: {response.get('error', 'Unknown error')}"
assert (
isinstance(response, dict) and "error" not in response
), f"Health check failed: {response.get('error', 'Unknown error')}"
return response
# Wait before retrying with exponential backoff
await asyncio.sleep(retry_delay)
retry_delay *= 2 # Exponential backoff
# Should not reach here, but just in case
assert False, "Health check failed after all retries"
@ -471,6 +473,50 @@ def test_update_litellm_params_for_health_check():
)
@pytest.mark.asyncio
async def test_perform_health_check_filters_by_model_id():
"""
When model_id is passed, only that deployment is checked (not all deployments
that share the same model name).
"""
from litellm.proxy.health_check import perform_health_check
# Two deployments with same model_name but different ids
model_list = [
{
"model_name": "gpt-4",
"model_info": {"id": "deployment-id-1"},
"litellm_params": {"model": "gpt-4", "api_key": "fake-key-1"},
},
{
"model_name": "gpt-4",
"model_info": {"id": "deployment-id-2"},
"litellm_params": {"model": "gpt-4", "api_key": "fake-key-2"},
},
]
captured_list = []
async def mock_perform_health_check(m_list, details=True):
captured_list.append(m_list)
return [{"model": "gpt-4", "api_key": m_list[0]["litellm_params"]["api_key"]}], []
with patch(
"litellm.proxy.health_check._perform_health_check",
side_effect=mock_perform_health_check,
):
healthy_endpoints, unhealthy_endpoints = await perform_health_check(
model_list=model_list, model_id="deployment-id-2", details=True
)
# Only one deployment (deployment-id-2) should have been passed to _perform_health_check
assert len(captured_list) == 1
assert len(captured_list[0]) == 1
assert (captured_list[0][0].get("model_info") or {}).get("id") == "deployment-id-2"
assert len(healthy_endpoints) == 1
assert healthy_endpoints[0]["api_key"] == "fake-key-2"
@pytest.mark.asyncio
async def test_perform_health_check_with_health_check_model():
"""
@ -562,6 +608,99 @@ async def test_health_check_bad_model():
), "Health check took longer than health_check_timeout"
@pytest.mark.asyncio
async def test_health_check_respects_concurrency_limit():
from litellm.proxy.health_check import _perform_health_check
model_list = [
{"litellm_params": {"model": f"openai/gpt-4o-mini-{i}", "api_key": "fake-key"}}
for i in range(6)
]
active = 0
max_active = 0
async def mock_health_check(litellm_params, **kwargs):
nonlocal active, max_active
active += 1
max_active = max(max_active, active)
await asyncio.sleep(0.05)
active -= 1
return {"status": "healthy"}
with patch("litellm.ahealth_check", side_effect=mock_health_check):
await _perform_health_check(model_list, max_concurrency=2)
assert max_active <= 2
@pytest.mark.asyncio
async def test_health_check_creates_only_bounded_initial_tasks():
from litellm.proxy.health_check import _perform_health_check
model_list = [
{"litellm_params": {"model": f"openai/gpt-4o-mini-{i}", "api_key": "fake-key"}}
for i in range(10)
]
release_event = asyncio.Event()
create_task_call_count = 0
real_create_task = asyncio.create_task
async def mock_health_check(litellm_params, **kwargs):
await release_event.wait()
return {"status": "healthy"}
def tracked_create_task(coro):
nonlocal create_task_call_count
create_task_call_count += 1
return real_create_task(coro)
with patch("litellm.ahealth_check", side_effect=mock_health_check), patch(
"litellm.proxy.health_check.asyncio.create_task", side_effect=tracked_create_task
):
perform_task = real_create_task(
_perform_health_check(model_list, max_concurrency=2)
)
await asyncio.sleep(0.05)
assert create_task_call_count == 2
release_event.set()
await perform_task
@pytest.mark.asyncio
async def test_timeout_does_not_cancel_other_health_checks():
from litellm.proxy.health_check import _perform_health_check
model_list = [
{
"litellm_params": {"model": "openai/slow-model", "api_key": "fake-key"},
"model_info": {"health_check_timeout": 0.05},
},
{
"litellm_params": {"model": "openai/fast-model", "api_key": "fake-key"},
"model_info": {"health_check_timeout": 1},
},
]
async def mock_health_check(litellm_params, **kwargs):
if litellm_params["model"] == "openai/slow-model":
await asyncio.sleep(0.2)
return {"status": "healthy"}
await asyncio.sleep(0.01)
return {"status": "healthy"}
with patch("litellm.ahealth_check", side_effect=mock_health_check):
healthy_endpoints, unhealthy_endpoints = await _perform_health_check(
model_list, max_concurrency=1
)
healthy_models = {endpoint["model"] for endpoint in healthy_endpoints}
unhealthy_models = {endpoint["model"] for endpoint in unhealthy_endpoints}
assert "openai/fast-model" in healthy_models
assert "openai/slow-model" in unhealthy_models
@pytest.mark.asyncio
async def test_ahealth_check_ocr():
litellm._turn_on_debug()
@ -643,20 +782,20 @@ async def test_image_generation_health_check_prompt(monkeypatch):
async def test_health_check_with_custom_llm_provider():
"""
Test that ahealth_check correctly uses custom_llm_provider from model_params.
This test verifies the fix for the issue where the UI's "Test connect" button
failed with "LLM Provider NOT provided" error for OpenAI-compatible self-hosted
providers, even when a provider was selected in the dropdown.
The fix ensures that when custom_llm_provider is passed in model_params,
it's properly forwarded to get_llm_provider() to identify the correct provider.
"""
from unittest.mock import MagicMock
# Mock the completion call to avoid making real API calls
mock_response = MagicMock()
mock_response._hidden_params = {"headers": {"x-ratelimit-remaining-tokens": "1000"}}
with patch("litellm.acompletion", return_value=mock_response):
# Test with a custom model name that wouldn't be recognized without custom_llm_provider
response = await litellm.ahealth_check(
@ -668,7 +807,7 @@ async def test_health_check_with_custom_llm_provider():
},
mode="chat",
)
# Should succeed without "LLM Provider NOT provided" error
assert "error" not in response
assert isinstance(response, dict)

View file

@ -1453,3 +1453,47 @@ def test_convert_to_model_response_object_falsy_id_preserves_auto_generated(fals
)
assert result.id == original_id
assert result.id.startswith("chatcmpl-")
def test_convert_to_model_response_object_default_usage_overwritten():
"""
Regression test: convert_to_model_response_object must properly set Usage
on a ModelResponse that only has the default Usage from ModelResponse.__init__()
(i.e. no extra litellm.Usage() set via setattr beforehand).
This validates the optimization of removing the redundant
`setattr(model_response, "usage", litellm.Usage())` in completion().
"""
mr = ModelResponse()
# usage is not set by default (optimization: avoid constructing throwaway Usage)
assert not hasattr(mr, "usage")
response_object = {
"id": "chatcmpl-usage-test",
"choices": [
{
"index": 0,
"message": {"role": "assistant", "content": "Hello"},
"finish_reason": "stop",
}
],
"usage": {
"prompt_tokens": 15,
"completion_tokens": 7,
"total_tokens": 22,
},
"model": "gpt-4o",
}
result = convert_to_model_response_object(
model_response_object=mr,
response_object=response_object,
stream=False,
start_time=datetime.now(),
end_time=datetime.now(),
)
assert isinstance(result, ModelResponse)
assert result.usage.prompt_tokens == 15
assert result.usage.completion_tokens == 7
assert result.usage.total_tokens == 22

View file

@ -13,7 +13,19 @@
},
"overrides": {
"glob": ">=11.1.0",
"tar": ">=7.5.7",
"@isaacs/brace-expansion": ">=5.0.1"
"tar": ">=7.5.8",
"minimatch": ">=10.2.1",
"diff": ">=8.0.3",
"@isaacs/brace-expansion": ">=5.0.1",
"@babel/traverse": ">=7.23.2",
"ws": ">=7.5.10",
"http-proxy-middleware": ">=2.0.9",
"tar-fs": ">=2.1.4",
"webpack-dev-middleware": ">=5.3.4",
"braces": ">=3.0.3",
"axios": ">=0.30.2",
"webpack": ">=5.94.0",
"serve-static": ">=1.16.0",
"path-to-regexp": ">=0.1.12"
}
}
}

View file

@ -25,7 +25,19 @@
},
"overrides": {
"glob": ">=11.1.0",
"tar": ">=7.5.7",
"@isaacs/brace-expansion": ">=5.0.1"
"tar": ">=7.5.8",
"minimatch": ">=10.2.1",
"diff": ">=8.0.3",
"@isaacs/brace-expansion": ">=5.0.1",
"@babel/traverse": ">=7.23.2",
"ws": ">=7.5.10",
"http-proxy-middleware": ">=2.0.9",
"tar-fs": ">=2.1.4",
"webpack-dev-middleware": ">=5.3.4",
"braces": ">=3.0.3",
"axios": ">=0.30.2",
"webpack": ">=5.94.0",
"serve-static": ">=1.16.0",
"path-to-regexp": ">=0.1.12"
}
}

View file

@ -2330,7 +2330,9 @@ async def test_run_background_health_check_reflects_llm_model_list(monkeypatch):
test_model_list_2 = [{"model_name": "model-b"}]
called_model_lists = []
async def fake_perform_health_check(model_list, details):
async def fake_perform_health_check(
model_list, details, max_concurrency=None
):
called_model_lists.append(copy.deepcopy(model_list))
return (["healthy"], ["unhealthy"])
@ -2378,7 +2380,9 @@ async def test_background_health_check_skip_disabled_models(monkeypatch):
]
called_model_lists = []
async def fake_perform_health_check(model_list, details):
async def fake_perform_health_check(
model_list, details, max_concurrency=None
):
called_model_lists.append(copy.deepcopy(model_list))
return (["healthy"], [])

View file

@ -0,0 +1,81 @@
"""
Unit tests for prometheus_emit_stream_label opt-in setting.
Tests that:
- stream label is NOT added to litellm_proxy_total_requests_metric by default
- stream label IS added when litellm.prometheus_emit_stream_label = True
- stream value is populated correctly from standard_logging_payload
"""
import pytest
import litellm
from litellm.types.integrations.prometheus import (
PrometheusMetricLabels,
UserAPIKeyLabelNames,
)
def test_stream_label_not_present_by_default():
"""stream label should NOT appear in litellm_proxy_total_requests_metric unless opted in"""
litellm.prometheus_emit_stream_label = False
labels = PrometheusMetricLabels.get_labels("litellm_proxy_total_requests_metric")
assert UserAPIKeyLabelNames.STREAM.value not in labels
def test_stream_label_present_when_opted_in():
"""stream label SHOULD appear in litellm_proxy_total_requests_metric when opted in"""
litellm.prometheus_emit_stream_label = True
try:
labels = PrometheusMetricLabels.get_labels("litellm_proxy_total_requests_metric")
assert UserAPIKeyLabelNames.STREAM.value in labels
finally:
litellm.prometheus_emit_stream_label = False
def test_stream_label_not_in_other_metrics_when_opted_in():
"""stream label should NOT be added to other metrics even when opted in"""
litellm.prometheus_emit_stream_label = True
try:
other_metrics = [
"litellm_proxy_failed_requests_metric",
"litellm_spend_metric",
"litellm_input_tokens_metric",
"litellm_output_tokens_metric",
"litellm_llm_api_latency_metric",
]
for metric in other_metrics:
labels = PrometheusMetricLabels.get_labels(metric)
assert UserAPIKeyLabelNames.STREAM.value not in labels, (
f"stream label should not be in {metric}"
)
finally:
litellm.prometheus_emit_stream_label = False
def test_stream_label_name():
"""STREAM label name should be 'stream'"""
assert UserAPIKeyLabelNames.STREAM.value == "stream"
def test_user_api_key_label_values_has_stream_field():
"""UserAPIKeyLabelValues should accept stream field"""
from litellm.types.integrations.prometheus import UserAPIKeyLabelValues
values = UserAPIKeyLabelValues(stream="True")
assert values.stream == "True"
values_false = UserAPIKeyLabelValues(stream="False")
assert values_false.stream == "False"
values_none = UserAPIKeyLabelValues()
assert values_none.stream is None
def test_stream_label_in_model_dump():
"""stream field appears in model_dump() output for use in prometheus_label_factory"""
from litellm.types.integrations.prometheus import UserAPIKeyLabelValues
values = UserAPIKeyLabelValues(stream="True")
dumped = values.model_dump()
assert "stream" in dumped
assert dumped["stream"] == "True"

View file

@ -147,7 +147,8 @@ class TestResponseCompliance:
"""Verify status enum values match spec."""
schema = spec_dict["components"]["schemas"]["CreateModelInteractionParams"]
status_prop = schema["properties"]["status"]
expected_statuses = ["UNSPECIFIED", "IN_PROGRESS", "REQUIRES_ACTION", "COMPLETED", "FAILED", "CANCELLED", "INCOMPLETE"]
# Google Interactions API uses lowercase status values (updated Feb 2026)
expected_statuses = ["in_progress", "requires_action", "completed", "failed", "cancelled", "incomplete"]
assert status_prop["enum"] == expected_statuses
print(f"✓ Status enum values: {expected_statuses}")

View file

@ -237,6 +237,83 @@ async def test_bedrock_rerank_header_forwarding_async(model):
pytest.fail(f"Failed to forward headers to {model}: {str(e)}")
def test_bedrock_rerank_timeout_sync():
"""
Test that the timeout parameter is passed through to the HTTP client for Bedrock rerank (sync).
"""
client = HTTPHandler()
model = "bedrock/arn:aws:bedrock:us-east-1::foundation-model/cohere.rerank-v3-5:0"
mock_credentials_info = create_mock_credentials()
with patch.object(client, "post") as mock_post, \
patch("litellm.llms.bedrock.rerank.handler.BedrockRerankHandler._get_boto_credentials_from_optional_params", return_value=mock_credentials_info), \
patch("botocore.auth.SigV4Auth") as mock_sigv4:
mock_sigv4.return_value = MagicMock()
mock_response = Mock()
mock_response.status_code = 200
mock_response.text = json.dumps(bedrock_rerank_response)
mock_response.json = lambda: json.loads(mock_response.text)
mock_response.raise_for_status = lambda: None
mock_post.return_value = mock_response
litellm.rerank(
model=model,
query=test_query,
documents=test_documents,
top_n=3,
client=client,
timeout=0.001,
aws_region_name="us-east-1",
aws_bedrock_runtime_endpoint="https://bedrock-runtime.us-east-1.amazonaws.com",
)
assert mock_post.called
call_kwargs = mock_post.call_args.kwargs
assert call_kwargs.get("timeout") == 0.001, (
f"Expected timeout=0.001, got timeout={call_kwargs.get('timeout')}"
)
@pytest.mark.asyncio
async def test_bedrock_rerank_timeout_async():
"""
Test that the timeout parameter is passed through to the HTTP client for Bedrock rerank (async).
"""
client = AsyncHTTPHandler()
model = "bedrock/arn:aws:bedrock:us-east-1::foundation-model/cohere.rerank-v3-5:0"
mock_credentials_info = create_mock_credentials()
with patch.object(client, "post", new_callable=AsyncMock) as mock_post, \
patch("litellm.llms.bedrock.rerank.handler.BedrockRerankHandler._get_boto_credentials_from_optional_params", return_value=mock_credentials_info), \
patch("botocore.auth.SigV4Auth") as mock_sigv4:
mock_sigv4.return_value = MagicMock()
mock_response = AsyncMock()
mock_response.status_code = 200
mock_response.text = json.dumps(bedrock_rerank_response)
mock_response.json = lambda: json.loads(mock_response.text)
mock_response.raise_for_status = lambda: None
mock_post.return_value = mock_response
await litellm.arerank(
model=model,
query=test_query,
documents=test_documents,
top_n=3,
client=client,
timeout=0.001,
aws_region_name="us-east-1",
aws_bedrock_runtime_endpoint="https://bedrock-runtime.us-east-1.amazonaws.com",
)
assert mock_post.called
call_kwargs = mock_post.call_args.kwargs
assert call_kwargs.get("timeout") == 0.001, (
f"Expected timeout=0.001, got timeout={call_kwargs.get('timeout')}"
)
def test_bedrock_rerank_extra_headers_and_headers_merge():
"""
Test that both extra_headers and headers parameters are correctly merged for Bedrock rerank.

View file

@ -207,10 +207,10 @@ class TestOpenAIChatCompletionStreamingHandler:
def test_chunk_parser_maps_reasoning_to_reasoning_content(self):
"""
Test that chunk_parser maps 'reasoning' field to 'reasoning_content'.
Some OpenAI-compatible providers (e.g., GLM-5, hosted_vllm) return
delta.reasoning, but LiteLLM expects delta.reasoning_content.
Regression test for: Streaming responses with delta.reasoning field
coming back empty when using openai/ or hosted_vllm/ providers.
"""
@ -293,3 +293,34 @@ class TestPromptCacheKeyIntegration:
prompt_cache_key="test-cache-key-123",
)
assert optional_params.get("prompt_cache_key") == "test-cache-key-123"
class TestPromptCacheParams:
"""Tests for prompt_cache_key and prompt_cache_retention support."""
def setup_method(self):
self.config = OpenAIGPTConfig()
def test_prompt_cache_key_in_supported_params(self):
"""Test that prompt_cache_key is in supported params for OpenAI models."""
supported_params = self.config.get_supported_openai_params("gpt-4o")
assert "prompt_cache_key" in supported_params
def test_prompt_cache_retention_in_supported_params(self):
"""Test that prompt_cache_retention is in supported params for OpenAI models."""
supported_params = self.config.get_supported_openai_params("gpt-4o")
assert "prompt_cache_retention" in supported_params
def test_prompt_cache_params_passed_through(self):
"""Test that prompt_cache_key and prompt_cache_retention are passed through by map_openai_params."""
optional_params = self.config.map_openai_params(
non_default_params={
"prompt_cache_key": "my-cache-key",
"prompt_cache_retention": "24h",
},
optional_params={},
model="gpt-4o",
drop_params=False,
)
assert optional_params.get("prompt_cache_key") == "my-cache-key"
assert optional_params.get("prompt_cache_retention") == "24h"

View file

@ -925,4 +925,318 @@ def test_get_supported_openai_params():
assert "temperature" in params
assert "stream" in params
assert "background" in params
assert "stream" in params
assert "stream" in params
class TestPhaseParameter:
"""Tests for the `phase` parameter on assistant output items (gpt-5.3-codex)."""
def setup_method(self):
self.config = OpenAIResponsesAPIConfig()
self.model = "gpt-5.3-codex"
self.logging_obj = MagicMock()
@staticmethod
def _make_output_text(text: str):
from litellm.types.responses.main import OutputText
return OutputText(type="output_text", text=text, annotations=[])
def test_generic_response_output_item_accepts_phase_commentary(self):
from litellm.types.responses.main import GenericResponseOutputItem
item = GenericResponseOutputItem(
type="message",
id="msg_001",
status="completed",
role="assistant",
content=[self._make_output_text("Thinking...")],
phase="commentary",
)
assert item.phase == "commentary"
def test_generic_response_output_item_accepts_phase_final_answer(self):
from litellm.types.responses.main import GenericResponseOutputItem
item = GenericResponseOutputItem(
type="message",
id="msg_002",
status="completed",
role="assistant",
content=[self._make_output_text("The answer is 42.")],
phase="final_answer",
)
assert item.phase == "final_answer"
def test_generic_response_output_item_phase_defaults_to_none(self):
from litellm.types.responses.main import GenericResponseOutputItem
item = GenericResponseOutputItem(
type="message",
id="msg_003",
status="completed",
role="assistant",
content=[self._make_output_text("Hello")],
)
assert item.phase is None
def test_output_function_tool_call_accepts_phase(self):
from litellm.types.responses.main import OutputFunctionToolCall
item = OutputFunctionToolCall(
type="function_call",
id="fc_001",
arguments='{"query": "test"}',
call_id="call_001",
name="search",
status="completed",
phase="commentary",
)
assert item.phase == "commentary"
def test_input_passthrough_dict_preserves_phase(self):
"""Dict input items (the normal HTTP flow) must preserve phase verbatim."""
input_items = [
{
"type": "message",
"role": "user",
"content": [{"type": "input_text", "text": "Hi"}],
},
{
"type": "message",
"role": "assistant",
"content": [{"type": "output_text", "text": "Preamble..."}],
"phase": "commentary",
},
{
"type": "message",
"role": "assistant",
"content": [{"type": "output_text", "text": "Done."}],
"phase": "final_answer",
},
{
"type": "message",
"role": "assistant",
"content": [{"type": "output_text", "text": "Neutral."}],
"phase": None,
},
]
result = self.config._validate_input_param(input_items)
assert isinstance(result, list)
assert "phase" not in result[0]
assert result[1]["phase"] == "commentary"
assert result[2]["phase"] == "final_answer"
assert result[3]["phase"] is None
def test_input_passthrough_pydantic_preserves_non_null_phase(self):
"""Pydantic input items must preserve non-null phase values."""
from litellm.types.responses.main import GenericResponseOutputItem
item = GenericResponseOutputItem(
type="message",
id="msg_010",
status="completed",
role="assistant",
content=[self._make_output_text("commentary")],
phase="commentary",
)
result = self.config._validate_input_param([item])
assert isinstance(result, list)
assert result[0]["phase"] == "commentary"
def test_response_parsing_preserves_phase_on_output(self):
"""Non-streaming response must preserve phase on output items."""
raw_json = {
"id": "resp_001",
"created_at": 1700000000,
"model": "gpt-5.3-codex",
"object": "response",
"status": "completed",
"output": [
{
"type": "message",
"id": "msg_001",
"status": "completed",
"role": "assistant",
"content": [{"type": "output_text", "text": "preamble"}],
"phase": "commentary",
},
{
"type": "message",
"id": "msg_002",
"status": "completed",
"role": "assistant",
"content": [{"type": "output_text", "text": "answer"}],
"phase": "final_answer",
},
],
"usage": {"input_tokens": 10, "output_tokens": 20, "total_tokens": 30},
}
response = ResponsesAPIResponse(**raw_json)
assert len(response.output) == 2
for idx, output_item in enumerate(response.output):
if isinstance(output_item, dict):
phase = output_item.get("phase")
else:
phase = getattr(output_item, "phase", None)
expected = "commentary" if idx == 0 else "final_answer"
assert phase == expected, (
f"output[{idx}] phase={phase!r}, expected {expected!r}"
)
def test_streaming_output_item_done_preserves_phase(self):
"""OutputItemDoneEvent must preserve phase on its item."""
from litellm.types.llms.openai import (
OutputItemDoneEvent,
ResponsesAPIStreamEvents,
)
chunk = {
"type": "response.output_item.done",
"output_index": 0,
"sequence_number": 3,
"item": {
"type": "message",
"id": "msg_100",
"status": "completed",
"role": "assistant",
"content": [{"type": "output_text", "text": "done"}],
"phase": "final_answer",
},
}
result = self.config.transform_streaming_response(
model=self.model, parsed_chunk=chunk, logging_obj=self.logging_obj
)
assert isinstance(result, OutputItemDoneEvent)
assert result.type == ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE
assert getattr(result.item, "phase", None) == "final_answer"
def test_streaming_output_item_added_preserves_phase(self):
"""OutputItemAddedEvent must preserve phase on its item."""
from litellm.types.llms.openai import (
OutputItemAddedEvent,
ResponsesAPIStreamEvents,
)
chunk = {
"type": "response.output_item.added",
"output_index": 0,
"item": {
"type": "message",
"id": "msg_200",
"role": "assistant",
"phase": "commentary",
},
}
result = self.config.transform_streaming_response(
model=self.model, parsed_chunk=chunk, logging_obj=self.logging_obj
)
assert isinstance(result, OutputItemAddedEvent)
assert result.type == ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED
assert getattr(result.item, "phase", None) == "commentary"
def test_streaming_response_completed_preserves_phase(self):
"""ResponseCompletedEvent must preserve phase on output items inside the response."""
completed_chunk = {
"type": "response.completed",
"response": {
"id": "resp_300",
"created_at": 1700000000,
"model": "gpt-5.3-codex",
"object": "response",
"status": "completed",
"output": [
{
"type": "message",
"id": "msg_300",
"status": "completed",
"role": "assistant",
"content": [{"type": "output_text", "text": "final"}],
"phase": "final_answer",
}
],
"usage": {
"input_tokens": 5,
"output_tokens": 10,
"total_tokens": 15,
},
},
}
result = self.config.transform_streaming_response(
model=self.model,
parsed_chunk=completed_chunk,
logging_obj=self.logging_obj,
)
assert result.type == ResponsesAPIStreamEvents.RESPONSE_COMPLETED
output_item = result.response.output[0]
if isinstance(output_item, dict):
assert output_item["phase"] == "final_answer"
else:
assert getattr(output_item, "phase", None) == "final_answer"
def test_phase_roundtrip_output_to_input(self):
"""Simulate full round-trip: parse response output, then send items back as input."""
raw_json = {
"id": "resp_rt",
"created_at": 1700000000,
"model": "gpt-5.3-codex",
"object": "response",
"status": "completed",
"output": [
{
"type": "message",
"id": "msg_rt1",
"status": "completed",
"role": "assistant",
"content": [{"type": "output_text", "text": "preamble"}],
"phase": "commentary",
},
{
"type": "message",
"id": "msg_rt2",
"status": "completed",
"role": "assistant",
"content": [{"type": "output_text", "text": "answer"}],
"phase": "final_answer",
},
],
"usage": {"input_tokens": 10, "output_tokens": 20, "total_tokens": 30},
}
response = ResponsesAPIResponse(**raw_json)
input_items = []
for item in response.output:
if isinstance(item, dict):
input_items.append(item)
else:
input_items.append(
item.model_dump() if hasattr(item, "model_dump") else dict(item)
)
input_items.append(
{
"type": "message",
"role": "user",
"content": [{"type": "input_text", "text": "next question"}],
}
)
validated = self.config._validate_input_param(input_items)
assert isinstance(validated, list)
assert validated[0]["phase"] == "commentary"
assert validated[1]["phase"] == "final_answer"
assert "phase" not in validated[2]

View file

@ -761,3 +761,70 @@ async def test_request_body_with_html_script_tags():
f"Message content with HTML was modified during parsing: "
f"expected={msg['content']!r}, got={result['messages'][2]['content']!r}"
)
def test_safe_get_request_headers_caches_on_request_state():
"""
Test that _safe_get_request_headers caches the result on request.state
and returns the same object on subsequent calls.
"""
mock_request = MagicMock()
mock_request.headers = {"content-type": "application/json", "authorization": "Bearer sk-123"}
mock_request.state = MagicMock(spec=[]) # empty spec so getattr returns default
# First call — should create and cache
result1 = _safe_get_request_headers(mock_request)
assert result1 == {"content-type": "application/json", "authorization": "Bearer sk-123"}
assert mock_request.state._cached_headers is result1
# Second call — should return the cached object (same identity)
result2 = _safe_get_request_headers(mock_request)
assert result2 is result1
def test_safe_get_request_headers_none_request():
"""
Test that _safe_get_request_headers returns empty dict for None request.
"""
result = _safe_get_request_headers(None)
assert result == {}
def test_safe_get_request_headers_copy_protects_cache():
"""
Test that callers using .copy() before mutation do not corrupt the cache.
"""
mock_request = MagicMock()
mock_request.headers = {"authorization": "Bearer sk-123", "host": "localhost"}
mock_request.state = MagicMock(spec=[])
original = _safe_get_request_headers(mock_request)
# Simulate what mutation call sites do: copy then pop
mutable = _safe_get_request_headers(mock_request).copy()
mutable.pop("authorization", None)
# Cache must be unaffected
assert "authorization" in _safe_get_request_headers(mock_request)
assert _safe_get_request_headers(mock_request) is original
def test_safe_get_request_headers_state_unavailable():
"""
Test that _safe_get_request_headers still returns headers when
request.state rejects attribute writes (the except path on the cache-write).
"""
class ReadOnlyState:
"""State object that allows reads but raises on writes."""
def __setattr__(self, name, value):
raise AttributeError("read-only state")
def __getattr__(self, name):
return None # _cached_headers not found → triggers fresh read
mock_request = MagicMock()
mock_request.headers = {"content-type": "application/json"}
mock_request.state = ReadOnlyState()
result = _safe_get_request_headers(mock_request)
assert result == {"content-type": "application/json"}

View file

@ -0,0 +1,75 @@
"""
Unit tests for ToolDiscoveryQueue.
"""
import os
import sys
import pytest
sys.path.insert(0, os.path.abspath("../../.."))
from litellm.proxy.db.db_transaction_queue.tool_discovery_queue import (
ToolDiscoveryQueue,
)
@pytest.fixture
def queue():
return ToolDiscoveryQueue()
def test_add_single_tool(queue):
queue.add_update({"tool_name": "my_tool", "origin": "user_defined"})
items = queue.flush()
assert len(items) == 1
assert items[0]["tool_name"] == "my_tool"
assert items[0]["origin"] == "user_defined"
def test_deduplication_same_name(queue):
"""Adding the same tool_name twice should only keep the first."""
queue.add_update({"tool_name": "tool_a", "origin": "mcp_server"})
queue.add_update({"tool_name": "tool_a", "origin": "user_defined"})
items = queue.flush()
assert len(items) == 1
assert items[0]["origin"] == "mcp_server" # first wins
def test_deduplication_different_names(queue):
queue.add_update({"tool_name": "tool_a"})
queue.add_update({"tool_name": "tool_b"})
items = queue.flush()
assert len(items) == 2
names = {i["tool_name"] for i in items}
assert names == {"tool_a", "tool_b"}
def test_flush_clears_pending(queue):
queue.add_update({"tool_name": "tool_x"})
items1 = queue.flush()
assert len(items1) == 1
items2 = queue.flush()
assert len(items2) == 0
def test_seen_names_reset_after_flush(queue):
"""Seen-set is cleared on flush so the same tool can re-enter the next cycle."""
queue.add_update({"tool_name": "tool_a"})
queue.flush()
queue.add_update({"tool_name": "tool_a"}) # same tool, new cycle
items = queue.flush()
assert len(items) == 1
assert items[0]["tool_name"] == "tool_a"
def test_empty_tool_name_ignored(queue):
queue.add_update({"tool_name": ""})
queue.add_update({"tool_name": None}) # type: ignore[arg-type]
items = queue.flush()
assert len(items) == 0
def test_flush_returns_list(queue):
result = queue.flush()
assert isinstance(result, list)

View file

@ -0,0 +1,197 @@
"""
Unit tests for tool_registry_writer.py — uses a mock prisma client
that exposes execute_raw / query_raw (matching the actual raw-SQL implementation).
"""
import os
import sys
from datetime import datetime, timezone
from unittest.mock import AsyncMock, MagicMock
import pytest
sys.path.insert(0, os.path.abspath("../../.."))
from litellm.proxy.db.tool_registry_writer import (
batch_upsert_tools,
get_tool,
get_tools_by_names,
list_tools,
update_tool_policy,
)
def _make_prisma(query_rows=None):
"""Return a minimal mock prisma_client with execute_raw / query_raw."""
default_row = {
"tool_id": "uuid-1",
"tool_name": "my_tool",
"origin": "user_defined",
"call_policy": "untrusted",
"call_count": 1,
"assignments": {},
"key_hash": None,
"team_id": None,
"key_alias": None,
"created_at": datetime.now(timezone.utc),
"updated_at": datetime.now(timezone.utc),
"created_by": None,
"updated_by": None,
}
rows = query_rows if query_rows is not None else [default_row]
prisma = MagicMock()
prisma.db.execute_raw = AsyncMock(return_value=None)
prisma.db.query_raw = AsyncMock(return_value=rows)
return prisma
@pytest.mark.asyncio
async def test_batch_upsert_tools_calls_execute_raw():
prisma = _make_prisma()
items = [{"tool_name": "tool_a", "origin": "mcp_server", "created_by": None}]
await batch_upsert_tools(prisma, items)
prisma.db.execute_raw.assert_awaited_once()
call_args = prisma.db.execute_raw.call_args
sql = call_args.args[0]
assert "LiteLLM_ToolTable" in sql
assert "ON CONFLICT" in sql
@pytest.mark.asyncio
async def test_batch_upsert_tools_empty_list():
prisma = _make_prisma()
await batch_upsert_tools(prisma, [])
prisma.db.execute_raw.assert_not_awaited()
@pytest.mark.asyncio
async def test_batch_upsert_tools_skips_empty_names():
prisma = _make_prisma()
items = [{"tool_name": "", "origin": None}, {"tool_name": None}] # type: ignore[list-item]
await batch_upsert_tools(prisma, items)
prisma.db.execute_raw.assert_not_awaited()
@pytest.mark.asyncio
async def test_batch_upsert_multiple_tools_calls_execute_raw_per_tool():
prisma = _make_prisma()
items = [
{"tool_name": "tool_a", "origin": "mcp_server", "created_by": None},
{"tool_name": "tool_b", "origin": "user_defined", "created_by": "alice"},
]
await batch_upsert_tools(prisma, items)
assert prisma.db.execute_raw.await_count == 2
@pytest.mark.asyncio
async def test_list_tools_no_filter():
row = {
"tool_id": "id1",
"tool_name": "tool_a",
"origin": "mcp",
"call_policy": "untrusted",
"call_count": 5,
"assignments": {},
"key_hash": None,
"team_id": None,
"key_alias": None,
"created_at": datetime.now(timezone.utc),
"updated_at": datetime.now(timezone.utc),
"created_by": None,
"updated_by": None,
}
prisma = _make_prisma(query_rows=[row])
result = await list_tools(prisma)
assert len(result) == 1
assert result[0].tool_name == "tool_a"
assert result[0].call_count == 5
prisma.db.query_raw.assert_awaited_once()
@pytest.mark.asyncio
async def test_list_tools_with_policy_filter():
row = {
"tool_id": "id1",
"tool_name": "blocked_tool",
"origin": None,
"call_policy": "blocked",
"call_count": 2,
"assignments": None,
"key_hash": None,
"team_id": None,
"key_alias": None,
"created_at": datetime.now(timezone.utc),
"updated_at": datetime.now(timezone.utc),
"created_by": None,
"updated_by": None,
}
prisma = _make_prisma(query_rows=[row])
result = await list_tools(prisma, call_policy="blocked")
assert result[0].call_policy == "blocked"
call_args = prisma.db.query_raw.call_args
sql = call_args.args[0]
assert "WHERE call_policy" in sql
@pytest.mark.asyncio
async def test_get_tool_found():
prisma = _make_prisma()
result = await get_tool(prisma, "my_tool")
assert result is not None
assert result.tool_name == "my_tool"
prisma.db.query_raw.assert_awaited_once()
@pytest.mark.asyncio
async def test_get_tool_not_found():
prisma = _make_prisma(query_rows=[])
result = await get_tool(prisma, "nonexistent")
assert result is None
@pytest.mark.asyncio
async def test_update_tool_policy_calls_execute_raw():
row = {
"tool_id": "uuid-1",
"tool_name": "my_tool",
"origin": "user_defined",
"call_policy": "blocked",
"call_count": 1,
"assignments": {},
"key_hash": None,
"team_id": None,
"key_alias": None,
"created_at": datetime.now(timezone.utc),
"updated_at": datetime.now(timezone.utc),
"created_by": None,
"updated_by": "admin",
}
prisma = _make_prisma(query_rows=[row])
result = await update_tool_policy(prisma, "my_tool", "blocked", "admin")
assert result is not None
assert result.call_policy == "blocked"
prisma.db.execute_raw.assert_awaited_once()
call_args = prisma.db.execute_raw.call_args
sql = call_args.args[0]
assert "ON CONFLICT" in sql
assert "call_policy" in sql
@pytest.mark.asyncio
async def test_get_tools_by_names_returns_policy_map():
rows = [
{"tool_name": "tool_a", "call_policy": "trusted"},
{"tool_name": "tool_b", "call_policy": "blocked"},
]
prisma = _make_prisma(query_rows=rows)
result = await get_tools_by_names(prisma, ["tool_a", "tool_b"])
assert result == {"tool_a": "trusted", "tool_b": "blocked"}
@pytest.mark.asyncio
async def test_get_tools_by_names_empty_list():
prisma = _make_prisma()
result = await get_tools_by_names(prisma, [])
assert result == {}
prisma.db.query_raw.assert_not_awaited()

View file

@ -2,10 +2,9 @@
"""
Test to verify the Google GenAI proxy API endpoints
"""
import asyncio
import os
import sys
from unittest.mock import AsyncMock, MagicMock, patch
from unittest.mock import AsyncMock, patch
import pytest
@ -13,7 +12,6 @@ sys.path.insert(
0, os.path.abspath("../../..")
) # Adds the parent directory to the system path
import litellm
def test_google_generate_content_endpoint():
@ -401,3 +399,123 @@ def test_google_generate_content_with_image_config():
assert "contents" in called_data
assert len(called_data["contents"]) == 1
assert called_data["contents"][0]["role"] == "user"
def test_google_generate_content_metadata_and_trace_id_callbacks():
"""Test that google_generate_content sets litellm_call_id and logging_obj for callbacks (e.g. S3, Langfuse)"""
try:
from fastapi import FastAPI
from fastapi.testclient import TestClient
from litellm.proxy.google_endpoints.endpoints import router as google_router
except ImportError as e:
pytest.skip(f"Skipping test due to missing dependency: {e}")
# Create a FastAPI app and include the router
app = FastAPI()
app.include_router(google_router)
# Create a test client
client = TestClient(app)
# Mock all required proxy server dependencies
with patch("litellm.proxy.proxy_server.llm_router") as mock_router, patch(
"litellm.proxy.proxy_server.general_settings", {}
), patch("litellm.proxy.proxy_server.proxy_config") as mock_proxy_config, patch(
"litellm.proxy.proxy_server.version", "1.0.0"
), patch(
"litellm.proxy.litellm_pre_call_utils.add_litellm_data_to_request"
) as mock_add_data:
mock_router.agenerate_content = AsyncMock(return_value={"test": "response"})
# Mock add_litellm_data_to_request to return data with metadata
async def mock_add_litellm_data(
data, request, user_api_key_dict, proxy_config, general_settings, version
):
# Simulate adding user metadata
data["litellm_metadata"] = {
"user_api_key_user_id": "test-user-id",
}
return data
mock_add_data.side_effect = mock_add_litellm_data
# Send a request to the endpoint with x-litellm-call-id header
test_call_id = "test-custom-call-id"
response = client.post(
"/v1beta/models/test-model:generateContent",
json={"contents": [{"role": "user", "parts": [{"text": "Hello"}]}]},
headers={
"Authorization": "Bearer sk-test-key",
"x-litellm-call-id": test_call_id,
},
)
assert response.status_code == 200
mock_router.agenerate_content.assert_called_once()
call_args = mock_router.agenerate_content.call_args
called_data = call_args[1]
# Verify that the litellm_logging_obj got assigned in the final called_data to router
assert "litellm_logging_obj" in called_data
assert "litellm_call_id" in called_data
assert called_data["litellm_call_id"] == test_call_id
def test_google_stream_generate_content_metadata_and_trace_id_callbacks():
"""Test that google_stream_generate_content sets litellm_call_id and logging_obj for callbacks"""
try:
from fastapi import FastAPI
from fastapi.testclient import TestClient
from litellm.proxy.google_endpoints.endpoints import router as google_router
except ImportError as e:
pytest.skip(f"Skipping test due to missing dependency: {e}")
app = FastAPI()
app.include_router(google_router)
client = TestClient(app)
mock_stream = AsyncMock()
mock_stream.__aiter__ = lambda self: mock_stream
mock_stream.__anext__.side_effect = StopAsyncIteration
with patch("litellm.proxy.proxy_server.llm_router") as mock_router, patch(
"litellm.proxy.proxy_server.general_settings", {}
), patch("litellm.proxy.proxy_server.proxy_config") as mock_proxy_config, patch(
"litellm.proxy.proxy_server.version", "1.0.0"
), patch(
"litellm.proxy.litellm_pre_call_utils.add_litellm_data_to_request"
) as mock_add_data:
mock_router.agenerate_content_stream = AsyncMock(return_value=mock_stream)
async def mock_add_litellm_data(
data, request, user_api_key_dict, proxy_config, general_settings, version
):
data["litellm_metadata"] = {
"user_api_key_user_id": "test-user-id",
}
return data
mock_add_data.side_effect = mock_add_litellm_data
test_call_id = "test-custom-stream-call-id"
response = client.post(
"/v1beta/models/test-model:streamGenerateContent",
json={"contents": [{"role": "user", "parts": [{"text": "Hello stream"}]}]},
headers={
"Authorization": "Bearer sk-test-key",
"x-litellm-call-id": test_call_id,
},
)
assert response.status_code == 200
mock_router.agenerate_content_stream.assert_called_once()
call_args = mock_router.agenerate_content_stream.call_args
called_data = call_args[1]
assert "litellm_logging_obj" in called_data
assert "litellm_call_id" in called_data
assert called_data["litellm_call_id"] == test_call_id

View file

@ -24,12 +24,29 @@ from litellm.types.guardrails import LitellmParams, PiiAction, PiiEntityType
from litellm.types.utils import Choices, Message, ModelResponse
def _make_mock_session_iterator(json_response):
def _make_mock_session_iterator(
json_response, status=200, content_type="application/json", text_response=""
):
"""Create a mock _get_session_iterator that yields a session returning json_response."""
@asynccontextmanager
async def mock_iterator():
class MockResponse:
def __init__(self):
self.status = status
self.content_type = content_type
self.headers = {"Content-Type": content_type}
async def text(self):
if text_response:
return text_response
import json
try:
return json.dumps(json_response)
except Exception:
return str(json_response)
async def json(self):
return json_response
@ -41,6 +58,7 @@ def _make_mock_session_iterator(json_response):
class MockSession:
def post(self, *args, **kwargs):
self.last_kwargs = kwargs
return MockResponse()
async def __aenter__(self):
@ -1444,3 +1462,149 @@ def test_deny_list_and_score_threshold_combined():
# EMAIL_ADDRESS passes both filters
assert len(filtered) == 1
assert filtered[0]["entity_type"] == "EMAIL_ADDRESS"
@pytest.mark.asyncio
async def test_analyze_text_non_json_content_type_fail_closed():
"""
Test that analyze_text raises GuardrailRaisedException when Presidio health
endpoint returns text/html and fail-closed is enabled.
"""
guardrail = _OPTIONAL_PresidioPIIMasking(
presidio_analyzer_api_base="http://test-analyzer/",
presidio_anonymizer_api_base="http://test-anonymizer/",
pii_entities_config={"PERSON": PiiAction.BLOCK},
mock_testing=False,
)
mock_iterator = _make_mock_session_iterator(
json_response=None,
status=200,
content_type="text/html; charset=utf-8",
text_response="Presidio Analyzer service is up.",
)
with patch.object(guardrail, "_get_session_iterator", mock_iterator):
with pytest.raises(GuardrailRaisedException) as exc_info:
await guardrail.analyze_text(
text="Hello world",
presidio_config=None,
request_data={},
)
assert "expected application/json Content-Type" in str(exc_info.value)
assert "text/html" in str(exc_info.value)
@pytest.mark.asyncio
async def test_analyze_text_non_json_content_type_fail_open():
"""
Test that analyze_text returns empty list when Presidio returns text/html
and fail-closed is NOT enabled.
"""
guardrail = _OPTIONAL_PresidioPIIMasking(
presidio_analyzer_api_base="http://test-analyzer/",
presidio_anonymizer_api_base="http://test-anonymizer/",
mock_testing=False,
)
mock_iterator = _make_mock_session_iterator(
json_response=None,
status=200,
content_type="text/html; charset=utf-8",
text_response="Presidio Analyzer service is up.",
)
with patch.object(guardrail, "_get_session_iterator", mock_iterator):
results = await guardrail.analyze_text(
text="Hello world",
presidio_config=None,
request_data={},
)
assert results == []
@pytest.mark.asyncio
async def test_analyze_text_http_error_status():
"""
Test that analyze_text handles 5xx HTTP errors properly.
"""
guardrail = _OPTIONAL_PresidioPIIMasking(
presidio_analyzer_api_base="http://test-analyzer/",
presidio_anonymizer_api_base="http://test-anonymizer/",
pii_entities_config={"PERSON": PiiAction.BLOCK},
mock_testing=False,
)
mock_iterator = _make_mock_session_iterator(
json_response=None,
status=500,
content_type="text/plain",
text_response="Internal Server Error",
)
with patch.object(guardrail, "_get_session_iterator", mock_iterator):
with pytest.raises(GuardrailRaisedException) as exc_info:
await guardrail.analyze_text(
text="Hello world",
presidio_config=None,
request_data={},
)
assert "HTTP 500" in str(exc_info.value)
@pytest.mark.asyncio
async def test_anonymize_text_non_json_content_type():
"""
Test that anonymize_text raises Exception for non-JSON responses.
"""
guardrail = _OPTIONAL_PresidioPIIMasking(
presidio_analyzer_api_base="http://test-analyzer/",
presidio_anonymizer_api_base="http://test-anonymizer/",
mock_testing=False,
)
mock_iterator = _make_mock_session_iterator(
json_response=None,
status=200,
content_type="text/html",
text_response="Presidio Anonymizer service is up.",
)
with patch.object(guardrail, "_get_session_iterator", mock_iterator):
with pytest.raises(
Exception, match="Presidio anonymizer returned non-JSON Content-Type"
):
await guardrail.anonymize_text(
text="Hello world",
analyze_results=[{"start": 0, "end": 5, "entity_type": "PERSON"}],
output_parse_pii=False,
masked_entity_count={},
)
@pytest.mark.asyncio
async def test_anonymize_text_http_error_status():
"""
Test that anonymize_text raises Exception on HTTP error.
"""
guardrail = _OPTIONAL_PresidioPIIMasking(
presidio_analyzer_api_base="http://test-analyzer/",
presidio_anonymizer_api_base="http://test-anonymizer/",
mock_testing=False,
)
mock_iterator = _make_mock_session_iterator(
json_response=None,
status=502,
content_type="text/plain",
text_response="Bad Gateway",
)
with patch.object(guardrail, "_get_session_iterator", mock_iterator):
with pytest.raises(Exception, match="Presidio anonymizer returned HTTP 502"):
await guardrail.anonymize_text(
text="Hello world",
analyze_results=[{"start": 0, "end": 5, "entity_type": "PERSON"}],
output_parse_pii=False,
masked_entity_count={},
)

View file

@ -0,0 +1,181 @@
"""
Unit tests for ToolPolicyGuardrail.
"""
import os
import sys
from typing import Any
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from fastapi import HTTPException
sys.path.insert(0, os.path.abspath("../../../../../.."))
from litellm.proxy.guardrails.guardrail_hooks.tool_policy.tool_policy_guardrail import (
ToolPolicyGuardrail,
)
from litellm.types.guardrails import GuardrailEventHooks
@pytest.fixture
def guardrail():
return ToolPolicyGuardrail()
# --- helpers ---
def _tool_request_inputs(tool_names: list) -> dict:
return {
"tools": [
{"type": "function", "function": {"name": name, "description": ""}}
for name in tool_names
]
}
def _tool_response_inputs(tool_names: list) -> dict:
return {
"tool_calls": [
{"type": "function", "function": {"name": name}}
for name in tool_names
]
}
# --- tests ---
def test_guardrail_supports_pre_and_post_call(guardrail):
hooks = guardrail.supported_event_hooks
assert GuardrailEventHooks.pre_call in hooks
assert GuardrailEventHooks.post_call in hooks
@pytest.mark.asyncio
async def test_no_tools_in_request_passes_through(guardrail):
inputs: Any = {"tools": []}
result = await guardrail.apply_guardrail(
inputs=inputs, request_data={}, input_type="request"
)
assert result is inputs
@pytest.mark.asyncio
async def test_no_tool_calls_in_response_passes_through(guardrail):
inputs: Any = {"tool_calls": []}
result = await guardrail.apply_guardrail(
inputs=inputs, request_data={}, input_type="response"
)
assert result is inputs
@pytest.mark.asyncio
async def test_untrusted_tools_pass_through(guardrail):
policy_map = {"search": "untrusted", "read_file": "trusted"}
with patch.object(guardrail, "_get_policies_cached", new=AsyncMock(return_value=policy_map)):
inputs: Any = _tool_request_inputs(["search", "read_file"])
result = await guardrail.apply_guardrail(
inputs=inputs, request_data={}, input_type="request"
)
assert result is inputs
@pytest.mark.asyncio
async def test_blocked_tool_in_request_raises_http_exception(guardrail):
policy_map = {"dangerous_tool": "blocked"}
with patch.object(guardrail, "_get_policies_cached", new=AsyncMock(return_value=policy_map)):
inputs: Any = _tool_request_inputs(["dangerous_tool"])
with pytest.raises(HTTPException) as exc_info:
await guardrail.apply_guardrail(
inputs=inputs, request_data={}, input_type="request"
)
assert exc_info.value.status_code == 400
assert "dangerous_tool" in exc_info.value.detail["blocked_tools"]
@pytest.mark.asyncio
async def test_blocked_tool_in_response_raises_http_exception(guardrail):
policy_map = {"exfil_tool": "blocked"}
with patch.object(guardrail, "_get_policies_cached", new=AsyncMock(return_value=policy_map)):
inputs: Any = _tool_response_inputs(["exfil_tool"])
with pytest.raises(HTTPException) as exc_info:
await guardrail.apply_guardrail(
inputs=inputs, request_data={}, input_type="response"
)
assert exc_info.value.status_code == 400
assert "exfil_tool" in exc_info.value.detail["blocked_tools"]
@pytest.mark.asyncio
async def test_mixed_blocked_and_allowed_raises_for_blocked(guardrail):
policy_map = {"safe_tool": "trusted", "bad_tool": "blocked"}
with patch.object(guardrail, "_get_policies_cached", new=AsyncMock(return_value=policy_map)):
inputs: Any = _tool_request_inputs(["safe_tool", "bad_tool"])
with pytest.raises(HTTPException) as exc_info:
await guardrail.apply_guardrail(
inputs=inputs, request_data={}, input_type="request"
)
blocked = exc_info.value.detail["blocked_tools"]
assert "bad_tool" in blocked
assert "safe_tool" not in blocked
@pytest.mark.asyncio
async def test_tool_not_in_db_passes_through(guardrail):
"""Tools not found in the DB (no entry) should not be blocked."""
with patch.object(guardrail, "_get_policies_cached", new=AsyncMock(return_value={})):
inputs: Any = _tool_request_inputs(["unknown_tool"])
result = await guardrail.apply_guardrail(
inputs=inputs, request_data={}, input_type="request"
)
assert result is inputs
@pytest.mark.asyncio
async def test_get_policies_cached_uses_cache(guardrail):
"""Second call with same tool names should return the cached result."""
policy_map = {"tool_a": "trusted"}
with patch(
"litellm.proxy.db.tool_registry_writer.get_tools_by_names",
new=AsyncMock(return_value=policy_map),
) as mock_db, patch(
"litellm.proxy.proxy_server.prisma_client",
new=MagicMock(),
):
# first call — should hit DB
result1 = await guardrail._get_policies_cached(["tool_a"])
assert result1 == policy_map
# second call — should hit cache, not DB again
result2 = await guardrail._get_policies_cached(["tool_a"])
assert result2 == policy_map
assert mock_db.call_count == 1
@pytest.mark.asyncio
async def test_get_policies_cached_no_prisma(guardrail):
"""Without a prisma client, returns empty dict."""
with patch(
"litellm.proxy.proxy_server.prisma_client",
None,
):
result = await guardrail._get_policies_cached(["tool_a"])
assert result == {}
@pytest.mark.asyncio
async def test_response_tool_calls_as_objects(guardrail):
"""tool_calls that are objects (not dicts) with .function.name should work."""
policy_map = {"obj_tool": "blocked"}
with patch.object(guardrail, "_get_policies_cached", new=AsyncMock(return_value=policy_map)):
fn = MagicMock()
fn.name = "obj_tool"
tc = MagicMock()
tc.function = fn
inputs: Any = {"tool_calls": [tc]}
with pytest.raises(HTTPException):
await guardrail.apply_guardrail(
inputs=inputs, request_data={}, input_type="response"
)

View file

@ -169,6 +169,223 @@ async def test_track_cost_callback_skips_when_no_standard_logging_object():
mock_proxy_logging.failed_tracking_alert.assert_not_called()
@pytest.mark.asyncio
async def test_enrich_failure_metadata_with_team_alias():
"""
When team_id is set but team_alias is missing (and key_alias is present),
_enrich_failure_metadata_with_key_info should look up the team from cache
and populate user_api_key_team_alias.
"""
mock_team_obj = MagicMock()
mock_team_obj.team_alias = "my-team-alias"
with patch(
"litellm.proxy.hooks.proxy_track_cost_callback.get_team_object",
new_callable=AsyncMock,
return_value=mock_team_obj,
):
metadata = {
"user_api_key": "hashed_key",
"user_api_key_alias": "my-key-alias", # already set
"user_api_key_team_id": "test_team_id",
"user_api_key_team_alias": None,
}
result = await _ProxyDBLogger._enrich_failure_metadata_with_key_info(metadata)
assert result["user_api_key_team_alias"] == "my-team-alias"
@pytest.mark.asyncio
async def test_enrich_failure_metadata_with_full_key_lookup():
"""
When all key fields are null (auth error 401 scenario), _enrich_failure_metadata_with_key_info
should look up the key object from cache/DB and populate alias, user_id, team_id,
then look up the team to get team_alias.
"""
mock_key_obj = MagicMock()
mock_key_obj.key_alias = "fetched-key-alias"
mock_key_obj.user_id = "fetched-user-id"
mock_key_obj.team_id = "fetched-team-id"
mock_key_obj.org_id = "fetched-org-id"
mock_team_obj = MagicMock()
mock_team_obj.team_alias = "fetched-team-alias"
with patch(
"litellm.proxy.hooks.proxy_track_cost_callback.get_key_object",
new_callable=AsyncMock,
return_value=mock_key_obj,
), patch(
"litellm.proxy.hooks.proxy_track_cost_callback.get_team_object",
new_callable=AsyncMock,
return_value=mock_team_obj,
):
metadata = {
"user_api_key": "hashed_key",
"user_api_key_alias": None, # all null - simulates auth error path
"user_api_key_user_id": None,
"user_api_key_team_id": None,
"user_api_key_team_alias": None,
"user_api_key_org_id": None,
}
result = await _ProxyDBLogger._enrich_failure_metadata_with_key_info(metadata)
assert result["user_api_key_alias"] == "fetched-key-alias"
assert result["user_api_key_user_id"] == "fetched-user-id"
assert result["user_api_key_team_id"] == "fetched-team-id"
assert result["user_api_key_org_id"] == "fetched-org-id"
assert result["user_api_key_team_alias"] == "fetched-team-alias"
@pytest.mark.asyncio
async def test_enrich_failure_metadata_skips_when_team_alias_present():
"""
When team_alias is already populated, _enrich_failure_metadata_with_key_info
should not perform a team cache lookup.
"""
with patch(
"litellm.proxy.hooks.proxy_track_cost_callback.get_key_object",
new_callable=AsyncMock,
) as mock_get_key, patch(
"litellm.proxy.hooks.proxy_track_cost_callback.get_team_object",
new_callable=AsyncMock,
) as mock_get_team:
metadata = {
"user_api_key": "hashed_key",
"user_api_key_alias": "existing-alias",
"user_api_key_team_id": "test_team_id",
"user_api_key_team_alias": "already-set",
}
result = await _ProxyDBLogger._enrich_failure_metadata_with_key_info(metadata)
assert result["user_api_key_team_alias"] == "already-set"
mock_get_key.assert_not_called()
mock_get_team.assert_not_called()
@pytest.mark.asyncio
async def test_enrich_failure_metadata_skips_when_no_api_key():
"""
When api_key hash is absent, _enrich_failure_metadata_with_key_info should
not perform any lookups.
"""
with patch(
"litellm.proxy.hooks.proxy_track_cost_callback.get_key_object",
new_callable=AsyncMock,
) as mock_get_key:
metadata = {
"user_api_key": None,
"user_api_key_alias": None,
"user_api_key_team_id": None,
"user_api_key_team_alias": None,
}
result = await _ProxyDBLogger._enrich_failure_metadata_with_key_info(metadata)
mock_get_key.assert_not_called()
@pytest.mark.asyncio
async def test_async_post_call_failure_hook_enriches_auth_error_metadata():
"""
Simulates a 401 ProxyException (e.g. can_key_call_model). In this case
UserAPIKeyAuth is created with only api_key set. The failure hook should
look up the key and team from cache/DB to populate all missing fields.
"""
logger = _ProxyDBLogger()
# This is what auth_exception_handler creates for 401 errors
user_api_key_dict = UserAPIKeyAuth(
api_key="hashed_key",
# key_alias, user_id, team_id, team_alias are all None
)
request_data = {
"model": "claude-haiku-4-5",
"messages": [{"role": "user", "content": "Hello"}],
"metadata": {},
"litellm_params": {},
}
mock_key_obj = MagicMock()
mock_key_obj.key_alias = "my-key-alias"
mock_key_obj.user_id = "my-user-id"
mock_key_obj.team_id = "my-team-id"
mock_key_obj.org_id = None
mock_team_obj = MagicMock()
mock_team_obj.team_alias = "my-team-alias"
with patch(
"litellm.proxy.db.db_spend_update_writer.DBSpendUpdateWriter.update_database",
new_callable=AsyncMock,
) as mock_update_database, patch(
"litellm.proxy.hooks.proxy_track_cost_callback.get_key_object",
new_callable=AsyncMock,
return_value=mock_key_obj,
), patch(
"litellm.proxy.hooks.proxy_track_cost_callback.get_team_object",
new_callable=AsyncMock,
return_value=mock_team_obj,
):
await logger.async_post_call_failure_hook(
request_data=request_data,
original_exception=Exception("401 - model not allowed"),
user_api_key_dict=user_api_key_dict,
)
mock_update_database.assert_called_once()
call_args = mock_update_database.call_args[1]
metadata = call_args["kwargs"]["litellm_params"]["metadata"]
assert metadata["user_api_key_alias"] == "my-key-alias"
assert metadata["user_api_key_user_id"] == "my-user-id"
assert metadata["user_api_key_team_id"] == "my-team-id"
assert metadata["user_api_key_team_alias"] == "my-team-alias"
@pytest.mark.asyncio
async def test_async_post_call_failure_hook_enriches_missing_team_alias():
"""
When user_api_key_dict has a team_id but no team_alias, async_post_call_failure_hook
should look up the team from cache and populate user_api_key_team_alias in the
spend log metadata written to the DB.
"""
logger = _ProxyDBLogger()
user_api_key_dict = UserAPIKeyAuth(
api_key="test_api_key",
key_alias="test_alias",
user_id="test_user_id",
team_id="test_team_id",
team_alias=None, # Missing - simulates regular key auth where SQL view omits team_alias
)
request_data = {
"model": "gpt-4",
"messages": [{"role": "user", "content": "Hello"}],
"metadata": {},
"litellm_params": {},
}
mock_team_obj = MagicMock()
mock_team_obj.team_alias = "enriched-team-alias"
with patch(
"litellm.proxy.db.db_spend_update_writer.DBSpendUpdateWriter.update_database",
new_callable=AsyncMock,
) as mock_update_database, patch(
"litellm.proxy.hooks.proxy_track_cost_callback.get_team_object",
new_callable=AsyncMock,
return_value=mock_team_obj,
):
await logger.async_post_call_failure_hook(
request_data=request_data,
original_exception=Exception("Provider rate limit"),
user_api_key_dict=user_api_key_dict,
)
mock_update_database.assert_called_once()
call_args = mock_update_database.call_args[1]
metadata = call_args["kwargs"]["litellm_params"]["metadata"]
assert metadata["user_api_key_team_alias"] == "enriched-team-alias"
assert metadata["user_api_key_team_id"] == "test_team_id"
@pytest.mark.asyncio
@pytest.mark.parametrize("model_value", [None, ""])
async def test_track_cost_callback_skips_for_falsy_model_and_no_slo(model_value):

View file

@ -5623,6 +5623,7 @@ async def test_rotate_master_key_model_data_valid_for_prisma(
litellm_params/model_info are JSON strings (create_many expects dicts).
"""
from unittest.mock import AsyncMock, MagicMock
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy.management_endpoints.key_management_endpoints import (
_rotate_master_key,
@ -6157,3 +6158,56 @@ async def test_get_member_team_ids():
# Should return team-A and team-B (user is a member of both)
# Should NOT return team-C (user is not in members list)
assert sorted(result) == ["team-A", "team-B"]
@pytest.mark.asyncio
async def test_generate_key_with_agent_id():
"""Test that agent_id is accepted in GenerateKeyRequest and passed to generate_key_helper_fn."""
from litellm.proxy._types import GenerateKeyRequest
# Verify GenerateKeyRequest accepts agent_id
request = GenerateKeyRequest(
key_alias="agent-test-key",
agent_id="test-agent-123",
models=[],
)
assert request.agent_id == "test-agent-123"
data_json = request.model_dump(exclude_unset=True, exclude_none=True)
assert data_json["agent_id"] == "test-agent-123"
@pytest.mark.asyncio
async def test_generate_key_helper_fn_agent_id():
"""Test that generate_key_helper_fn passes agent_id into the insert_data call."""
from unittest.mock import AsyncMock, MagicMock, call, patch
import litellm.proxy.management_endpoints.key_management_endpoints as km
mock_prisma_client = AsyncMock()
mock_insert = AsyncMock(
return_value=MagicMock(
token="sk-test",
created_at=None,
updated_at=None,
litellm_budget_table=None,
)
)
mock_prisma_client.insert_data = mock_insert
with patch.object(km, "prisma_client", mock_prisma_client):
with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client):
await generate_key_helper_fn(
request_type="key",
agent_id="test-agent-456",
key_alias="test-agent-key",
models=[],
table_name="key",
)
assert mock_insert.called, "insert_data was never called"
# insert_data is called as insert_data(data=key_data, ...)
call_kwargs = mock_insert.call_args.kwargs
key_data = call_kwargs.get("data", {})
assert key_data.get("agent_id") == "test-agent-456", (
f"Expected agent_id='test-agent-456' in key_data, got: {key_data.get('agent_id')}"
)

View file

@ -0,0 +1,149 @@
"""
Unit tests for tool management endpoints (/v1/tool/*).
Uses FastAPI TestClient with mocked DB functions.
Patches target the source modules (litellm.proxy.db.tool_registry_writer.*
and litellm.proxy.proxy_server.prisma_client) because the endpoint code
imports these inside function bodies to avoid circular imports.
"""
import os
import sys
from datetime import datetime, timezone
from typing import Optional
from unittest.mock import AsyncMock, MagicMock, patch
from fastapi import FastAPI
from fastapi.testclient import TestClient
sys.path.insert(0, os.path.abspath("../../.."))
from litellm.proxy.management_endpoints.tool_management_endpoints import router
from litellm.types.tool_management import LiteLLM_ToolTableRow
# --- helpers ---
def _make_tool_row(
tool_name: str = "my_tool",
call_policy: str = "untrusted",
origin: Optional[str] = None,
) -> LiteLLM_ToolTableRow:
now = datetime.now(timezone.utc)
return LiteLLM_ToolTableRow(
tool_id="uuid-1",
tool_name=tool_name,
origin=origin,
call_policy=call_policy, # type: ignore[arg-type]
assignments={},
created_at=now,
updated_at=now,
)
def _make_app() -> FastAPI:
"""Build a minimal FastAPI app with the tool management router."""
app = FastAPI()
app.include_router(router)
return app
# Stub the auth dependency so we don't need a real proxy running.
def _override_auth():
from litellm.proxy._types import UserAPIKeyAuth
return UserAPIKeyAuth(api_key="sk-test", user_id="admin")
# A real (non-None) prisma stub for truthiness checks.
_MOCK_PRISMA = MagicMock()
# --- test class ---
class TestToolManagementEndpoints:
def setup_method(self):
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
app = _make_app()
app.dependency_overrides[user_api_key_auth] = _override_auth
self.client = TestClient(app, raise_server_exceptions=True)
@patch(
"litellm.proxy.db.tool_registry_writer.list_tools",
new_callable=AsyncMock,
)
@patch("litellm.proxy.proxy_server.prisma_client", _MOCK_PRISMA)
def test_list_tools_returns_200(self, mock_db_list):
mock_db_list.return_value = [_make_tool_row()]
resp = self.client.get("/v1/tool/list")
assert resp.status_code == 200
body = resp.json()
assert body["total"] == 1
assert body["tools"][0]["tool_name"] == "my_tool"
@patch(
"litellm.proxy.db.tool_registry_writer.list_tools",
new_callable=AsyncMock,
)
@patch("litellm.proxy.proxy_server.prisma_client", _MOCK_PRISMA)
def test_list_tools_with_policy_filter(self, mock_db_list):
mock_db_list.return_value = [_make_tool_row(call_policy="blocked")]
resp = self.client.get("/v1/tool/list?call_policy=blocked")
assert resp.status_code == 200
assert resp.json()["tools"][0]["call_policy"] == "blocked"
@patch(
"litellm.proxy.db.tool_registry_writer.get_tool",
new_callable=AsyncMock,
)
@patch("litellm.proxy.proxy_server.prisma_client", _MOCK_PRISMA)
def test_get_tool_found(self, mock_db_get):
mock_db_get.return_value = _make_tool_row(tool_name="tool_a")
resp = self.client.get("/v1/tool/tool_a")
assert resp.status_code == 200
assert resp.json()["tool_name"] == "tool_a"
@patch(
"litellm.proxy.db.tool_registry_writer.get_tool",
new_callable=AsyncMock,
)
@patch("litellm.proxy.proxy_server.prisma_client", _MOCK_PRISMA)
def test_get_tool_not_found_returns_404(self, mock_db_get):
mock_db_get.return_value = None
resp = self.client.get("/v1/tool/nonexistent", follow_redirects=True)
assert resp.status_code == 404
@patch(
"litellm.proxy.db.tool_registry_writer.update_tool_policy",
new_callable=AsyncMock,
)
@patch("litellm.proxy.proxy_server.prisma_client", _MOCK_PRISMA)
def test_update_tool_policy_blocked(self, mock_db_update):
mock_db_update.return_value = _make_tool_row(call_policy="blocked")
resp = self.client.post(
"/v1/tool/policy",
json={"tool_name": "my_tool", "call_policy": "blocked"},
)
assert resp.status_code == 200
body = resp.json()
assert body["call_policy"] == "blocked"
assert body["updated"] is True
@patch("litellm.proxy.proxy_server.prisma_client", None)
def test_list_tools_no_db_returns_500(self):
resp = self.client.get("/v1/tool/list")
assert resp.status_code == 500
def test_update_tool_policy_invalid_policy_returns_422(self):
resp = self.client.post(
"/v1/tool/policy",
json={"tool_name": "my_tool", "call_policy": "invalid_value"},
)
assert resp.status_code == 422

View file

@ -0,0 +1,402 @@
"""
Tests for AI Usage Chat module.
"""
import json
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from litellm.proxy.management_endpoints.usage_endpoints.ai_usage_chat import (
TOOL_HANDLERS,
TOOLS_ADMIN,
TOOLS_BASE,
_build_system_prompt,
_summarise_entity_data,
_summarise_usage_data,
stream_usage_ai_chat,
)
SAMPLE_AGGREGATED_RESPONSE = {
"results": [
{
"date": "2025-01-15",
"metrics": {
"spend": 50.25,
"prompt_tokens": 20000,
"completion_tokens": 10000,
"total_tokens": 30000,
"api_requests": 500,
"successful_requests": 480,
"failed_requests": 20,
"cache_read_input_tokens": 0,
"cache_creation_input_tokens": 0,
},
"breakdown": {
"models": {
"gpt-4": {
"metrics": {
"spend": 40.0,
"api_requests": 300,
"total_tokens": 25000,
},
"metadata": {},
"api_key_breakdown": {},
},
},
"providers": {
"openai": {
"metrics": {"spend": 50.25, "api_requests": 500},
"metadata": {},
"api_key_breakdown": {},
},
},
"api_keys": {
"sk-test123": {
"metrics": {"spend": 50.25},
"metadata": {"key_alias": "Production Key"},
},
},
"model_groups": {},
"mcp_servers": {},
"entities": {},
},
},
],
"metadata": {
"total_spend": 50.25,
"total_api_requests": 500,
"total_successful_requests": 480,
"total_failed_requests": 20,
"total_tokens": 30000,
},
}
SAMPLE_TEAM_RESPONSE = {
"results": [
{
"date": "2025-01-15",
"metrics": {"spend": 100.0, "api_requests": 1000, "total_tokens": 50000},
"breakdown": {
"entities": {
"team-1": {
"metrics": {
"spend": 60.0,
"api_requests": 600,
"total_tokens": 30000,
},
"metadata": {"alias": "Engineering"},
"api_key_breakdown": {},
},
"team-2": {
"metrics": {
"spend": 40.0,
"api_requests": 400,
"total_tokens": 20000,
},
"metadata": {"alias": "Marketing"},
"api_key_breakdown": {},
},
},
"models": {},
"providers": {},
"api_keys": {},
"model_groups": {},
"mcp_servers": {},
},
},
],
"metadata": {"total_spend": 100.0, "total_api_requests": 1000},
}
class TestToolSchemas:
def test_admin_tools_include_all(self):
assert len(TOOLS_ADMIN) == 3
names = {t["function"]["name"] for t in TOOLS_ADMIN}
assert "get_usage_data" in names
assert "get_team_usage_data" in names
assert "get_tag_usage_data" in names
def test_base_tools_restricted_to_usage_only(self):
assert len(TOOLS_BASE) == 1
assert TOOLS_BASE[0]["function"]["name"] == "get_usage_data"
def test_admin_prompt_mentions_all_tools(self):
prompt = _build_system_prompt(is_admin=True)
assert "get_usage_data" in prompt
assert "get_team_usage_data" in prompt
assert "get_tag_usage_data" in prompt
def test_non_admin_prompt_only_mentions_usage_tool(self):
prompt = _build_system_prompt(is_admin=False)
assert "get_usage_data" in prompt
assert "get_team_usage_data" not in prompt
assert "get_tag_usage_data" not in prompt
def test_system_prompt_includes_todays_date(self):
from datetime import date
prompt = _build_system_prompt(is_admin=True)
assert date.today().isoformat() in prompt
class TestSummariseUsageData:
def test_summarise_includes_totals(self):
summary = _summarise_usage_data(SAMPLE_AGGREGATED_RESPONSE)
assert "$50.25" in summary
assert "500" in summary
def test_summarise_includes_models(self):
summary = _summarise_usage_data(SAMPLE_AGGREGATED_RESPONSE)
assert "gpt-4" in summary
def test_summarise_includes_providers(self):
summary = _summarise_usage_data(SAMPLE_AGGREGATED_RESPONSE)
assert "openai" in summary
def test_summarise_handles_empty_data(self):
empty = {"results": [], "metadata": {}}
summary = _summarise_usage_data(empty)
assert "no data" in summary.lower()
class TestSummariseEntityData:
def test_team_summary_includes_teams(self):
summary = _summarise_entity_data(SAMPLE_TEAM_RESPONSE, "Team")
assert "Engineering" in summary
assert "Marketing" in summary
assert "$60.0" in summary
assert "$40.0" in summary
def test_team_summary_empty(self):
empty = {"results": [], "metadata": {}}
summary = _summarise_entity_data(empty, "Team")
assert "No Team usage data" in summary
class TestStreamUsageAiChat:
@pytest.mark.asyncio
async def test_stream_emits_status_events(self):
mock_tool_call = MagicMock()
mock_tool_call.id = "call_123"
mock_tool_call.function.name = "get_usage_data"
mock_tool_call.function.arguments = json.dumps(
{
"start_date": "2025-01-01",
"end_date": "2025-01-31",
}
)
mock_first_response = MagicMock()
mock_first_response.choices = [MagicMock()]
mock_first_response.choices[0].message.tool_calls = [mock_tool_call]
mock_first_response.choices[0].message.model_dump.return_value = {
"role": "assistant",
"content": None,
"tool_calls": [
{
"id": "call_123",
"type": "function",
"function": {
"name": "get_usage_data",
"arguments": '{"start_date":"2025-01-01","end_date":"2025-01-31"}',
},
}
],
}
async def mock_stream():
chunk = MagicMock()
chunk.choices = [MagicMock()]
chunk.choices[0].delta.content = "Total spend is $50.25"
yield chunk
with patch(
"litellm.proxy.management_endpoints.usage_endpoints.ai_usage_chat.litellm"
) as mock_litellm, patch(
"litellm.proxy.management_endpoints.usage_endpoints.ai_usage_chat._fetch_usage_data",
new_callable=AsyncMock,
) as mock_fetch:
mock_litellm.acompletion = AsyncMock(
side_effect=[
mock_first_response,
mock_stream(),
]
)
mock_fetch.return_value = SAMPLE_AGGREGATED_RESPONSE
events = []
async for event in stream_usage_ai_chat(
messages=[{"role": "user", "content": "What is my total spend?"}],
model="gpt-4o-mini",
user_id="user-123",
is_admin=True,
):
events.append(json.loads(event.replace("data: ", "").strip()))
status_events = [e for e in events if e["type"] == "status"]
tool_call_events = [e for e in events if e["type"] == "tool_call"]
chunk_events = [e for e in events if e["type"] == "chunk"]
done_events = [e for e in events if e["type"] == "done"]
assert len(status_events) >= 1
assert "Thinking" in status_events[0]["message"]
assert len(tool_call_events) >= 1
assert tool_call_events[0]["tool_name"] == "get_usage_data"
assert tool_call_events[0]["status"] in ("running", "complete")
assert len(chunk_events) >= 1
assert len(done_events) == 1
@pytest.mark.asyncio
async def test_stream_handles_team_tool(self):
mock_tool_call = MagicMock()
mock_tool_call.id = "call_team"
mock_tool_call.function.name = "get_team_usage_data"
mock_tool_call.function.arguments = json.dumps(
{
"start_date": "2025-01-01",
"end_date": "2025-01-31",
}
)
mock_first_response = MagicMock()
mock_first_response.choices = [MagicMock()]
mock_first_response.choices[0].message.tool_calls = [mock_tool_call]
mock_first_response.choices[0].message.model_dump.return_value = {
"role": "assistant",
"content": None,
"tool_calls": [
{
"id": "call_team",
"type": "function",
"function": {
"name": "get_team_usage_data",
"arguments": '{"start_date":"2025-01-01","end_date":"2025-01-31"}',
},
}
],
}
async def mock_stream():
chunk = MagicMock()
chunk.choices = [MagicMock()]
chunk.choices[0].delta.content = "Engineering is the top team."
yield chunk
with patch(
"litellm.proxy.management_endpoints.usage_endpoints.ai_usage_chat.litellm"
) as mock_litellm, patch(
"litellm.proxy.management_endpoints.usage_endpoints.ai_usage_chat._fetch_team_usage_data",
new_callable=AsyncMock,
) as mock_fetch:
mock_litellm.acompletion = AsyncMock(
side_effect=[
mock_first_response,
mock_stream(),
]
)
mock_fetch.return_value = SAMPLE_TEAM_RESPONSE
events = []
async for event in stream_usage_ai_chat(
messages=[{"role": "user", "content": "Which team spends the most?"}],
model="gpt-4o-mini",
is_admin=True,
):
events.append(json.loads(event.replace("data: ", "").strip()))
chunk_events = [e for e in events if e["type"] == "chunk"]
assert len(chunk_events) >= 1
assert "Engineering" in chunk_events[0]["content"]
@pytest.mark.asyncio
async def test_stream_handles_error(self):
with patch(
"litellm.proxy.management_endpoints.usage_endpoints.ai_usage_chat.litellm"
) as mock_litellm:
mock_litellm.acompletion = AsyncMock(side_effect=Exception("LLM error"))
events = []
async for event in stream_usage_ai_chat(
messages=[{"role": "user", "content": "test"}],
):
events.append(json.loads(event.replace("data: ", "").strip()))
error_events = [e for e in events if e["type"] == "error"]
assert len(error_events) == 1
assert "internal error" in error_events[0]["message"].lower()
@pytest.mark.asyncio
async def test_non_admin_enforces_user_id(self):
mock_tool_call = MagicMock()
mock_tool_call.id = "call_456"
mock_tool_call.function.name = "get_usage_data"
mock_tool_call.function.arguments = json.dumps(
{
"start_date": "2025-01-01",
"end_date": "2025-01-31",
"user_id": "other-user",
}
)
mock_first_response = MagicMock()
mock_first_response.choices = [MagicMock()]
mock_first_response.choices[0].message.tool_calls = [mock_tool_call]
mock_first_response.choices[0].message.model_dump.return_value = {
"role": "assistant",
"content": None,
"tool_calls": [
{
"id": "call_456",
"type": "function",
"function": {
"name": "get_usage_data",
"arguments": '{"start_date":"2025-01-01","end_date":"2025-01-31","user_id":"other-user"}',
},
}
],
}
async def mock_stream():
chunk = MagicMock()
chunk.choices = [MagicMock()]
chunk.choices[0].delta.content = "Data."
yield chunk
mock_fetch = AsyncMock(return_value=SAMPLE_AGGREGATED_RESPONSE)
with patch(
"litellm.proxy.management_endpoints.usage_endpoints.ai_usage_chat.litellm"
) as mock_litellm, patch.dict(
"litellm.proxy.management_endpoints.usage_endpoints.ai_usage_chat.TOOL_HANDLERS",
{
"get_usage_data": {
"fetch": mock_fetch,
"summarise": _summarise_usage_data,
"label": "global usage data",
}
},
):
mock_litellm.acompletion = AsyncMock(
side_effect=[
mock_first_response,
mock_stream(),
]
)
events = []
async for event in stream_usage_ai_chat(
messages=[{"role": "user", "content": "Show data"}],
model="gpt-4o-mini",
user_id="my-user-id",
is_admin=False,
):
events.append(event)
mock_fetch.assert_called_once_with(
start_date="2025-01-01",
end_date="2025-01-31",
user_id="my-user-id",
)

View file

@ -1,4 +1,6 @@
import copy
import datetime
from typing import AsyncGenerator
from unittest.mock import AsyncMock, MagicMock
import pytest
@ -1348,3 +1350,202 @@ class TestOverrideOpenAIResponseModel:
# Verify the model was not changed
assert response_obj.model == fallback_model
class TestStreamingOverheadHeader:
"""
Tests that x-litellm-overhead-duration-ms is emitted in streaming responses.
Regression tests for: streaming requests not including overhead header.
"""
def test_get_custom_headers_includes_overhead_when_set(self):
"""
get_custom_headers() returns x-litellm-overhead-duration-ms
when litellm_overhead_time_ms is in hidden_params.
"""
mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth)
mock_user_api_key_dict.tpm_limit = None
mock_user_api_key_dict.rpm_limit = None
mock_user_api_key_dict.max_budget = None
mock_user_api_key_dict.spend = 0.0
mock_user_api_key_dict.allowed_model_region = None
hidden_params = {
"litellm_overhead_time_ms": 42.5,
"_response_ms": 500.0,
"model_id": "test-model-id",
"api_base": "https://api.openai.com",
}
headers = ProxyBaseLLMRequestProcessing.get_custom_headers(
user_api_key_dict=mock_user_api_key_dict,
call_id="test-call-id",
model_id="test-model-id",
cache_key="",
api_base="https://api.openai.com",
version="1.0.0",
response_cost=0.001,
model_region="",
hidden_params=hidden_params,
)
assert "x-litellm-overhead-duration-ms" in headers
assert headers["x-litellm-overhead-duration-ms"] == "42.5"
def test_get_custom_headers_omits_overhead_when_none(self):
"""
get_custom_headers() omits x-litellm-overhead-duration-ms
when litellm_overhead_time_ms is not in hidden_params.
"""
mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth)
mock_user_api_key_dict.tpm_limit = None
mock_user_api_key_dict.rpm_limit = None
mock_user_api_key_dict.max_budget = None
mock_user_api_key_dict.spend = 0.0
mock_user_api_key_dict.allowed_model_region = None
hidden_params = {
"_response_ms": 500.0,
"model_id": "test-model-id",
}
headers = ProxyBaseLLMRequestProcessing.get_custom_headers(
user_api_key_dict=mock_user_api_key_dict,
call_id="test-call-id",
model_id="test-model-id",
cache_key="",
api_base="https://api.openai.com",
version="1.0.0",
response_cost=0.001,
model_region="",
hidden_params=hidden_params,
)
# Should be absent (None gets filtered by exclude_values)
assert "x-litellm-overhead-duration-ms" not in headers
def test_update_response_metadata_sets_overhead_on_stream_wrapper(self):
"""
update_response_metadata() sets litellm_overhead_time_ms on
a streaming response's _hidden_params when llm_api_duration_ms is available.
"""
from litellm.litellm_core_utils.llm_response_utils.response_metadata import (
update_response_metadata,
)
# Mock the logging object with llm_api_duration_ms set
mock_logging_obj = MagicMock()
mock_logging_obj.model_call_details = {
"llm_api_duration_ms": 200.0,
"litellm_params": {},
}
mock_logging_obj.caching_details = None
mock_logging_obj.callback_duration_ms = None
mock_logging_obj.litellm_call_id = "test-call-id"
mock_logging_obj._response_cost_calculator = MagicMock(return_value=0.001)
# Simulate a streaming result object with _hidden_params (like CustomStreamWrapper)
stream_result = MagicMock()
stream_result._hidden_params = {
"model_id": "test-model-id",
"api_base": "https://api.openai.com",
"additional_headers": {},
}
start_time = datetime.datetime.now() - datetime.timedelta(milliseconds=300)
end_time = datetime.datetime.now()
update_response_metadata(
result=stream_result,
logging_obj=mock_logging_obj,
model="gpt-4o",
kwargs={},
start_time=start_time,
end_time=end_time,
)
assert "litellm_overhead_time_ms" in stream_result._hidden_params
overhead = stream_result._hidden_params["litellm_overhead_time_ms"]
assert overhead is not None
assert isinstance(overhead, float)
# overhead = total_response_ms (~300ms) - llm_api_duration_ms (200ms) = ~100ms
assert overhead > 0
@pytest.mark.asyncio
async def test_streaming_response_includes_overhead_header(self):
"""
StreamingResponse returned by create_response() includes
x-litellm-overhead-duration-ms in its headers.
"""
async def mock_generator() -> AsyncGenerator[str, None]:
yield 'data: {"id":"chatcmpl-test","choices":[{"delta":{"content":"hi"}}]}\n\n'
yield "data: [DONE]\n\n"
headers = {
"x-litellm-overhead-duration-ms": "42.5",
"x-litellm-call-id": "test-call-id",
"x-litellm-model-id": "test-model-id",
}
response = await create_response(
generator=mock_generator(),
media_type="text/event-stream",
headers=headers,
)
assert isinstance(response, StreamingResponse)
assert response.headers.get("x-litellm-overhead-duration-ms") == "42.5"
def test_streaming_overhead_header_in_custom_headers_from_stream_hidden_params(
self,
):
"""
Verifies that when get_custom_headers() is called with a streaming
response's hidden_params (containing litellm_overhead_time_ms),
the x-litellm-overhead-duration-ms header is correctly populated.
This tests the critical path: update_response_metadata sets the value
→ get_custom_headers reads it → StreamingResponse header is set.
"""
mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth)
mock_user_api_key_dict.tpm_limit = None
mock_user_api_key_dict.rpm_limit = None
mock_user_api_key_dict.max_budget = None
mock_user_api_key_dict.spend = 0.0
mock_user_api_key_dict.allowed_model_region = None
# This is what CustomStreamWrapper._hidden_params looks like after
# update_response_metadata() has been called on it
hidden_params = {
"model_id": "openai-gpt4o-deployment",
"api_base": "https://api.openai.com",
"additional_headers": {},
"litellm_overhead_time_ms": 55.3, # set by update_response_metadata
"_response_ms": 280.0,
"litellm_call_id": "test-call-id",
"response_cost": 0.002,
"cache_key": None,
"fastest_response_batch_completion": None,
"callback_duration_ms": None,
}
custom_headers = ProxyBaseLLMRequestProcessing.get_custom_headers(
user_api_key_dict=mock_user_api_key_dict,
call_id="test-call-id",
model_id=hidden_params.get("model_id"),
cache_key=hidden_params.get("cache_key") or "",
api_base=hidden_params.get("api_base") or "",
version="1.0.0",
response_cost=hidden_params.get("response_cost"),
model_region="",
hidden_params=hidden_params,
)
# The overhead header must be present and correct
assert "x-litellm-overhead-duration-ms" in custom_headers, (
"x-litellm-overhead-duration-ms header must be emitted during streaming. "
"It was missing — this is the streaming overhead header regression."
)
assert custom_headers["x-litellm-overhead-duration-ms"] == "55.3"

View file

@ -12,6 +12,7 @@ sys.path.insert(0, os.path.abspath("../../.."))
from litellm.proxy.health_endpoints._health_endpoints import (
_aggregate_health_check_results,
_build_model_param_to_info_mapping,
_perform_health_check_and_save,
_save_background_health_checks_to_db,
_save_health_check_results_if_changed,
_save_health_check_to_db,
@ -466,5 +467,43 @@ async def test_get_all_latest_health_checks_without_model_id(mock_prisma):
assert result[0].checked_at == mock_check2.checked_at # Latest
@pytest.mark.asyncio
async def test_perform_health_check_and_save_passes_model_id_to_perform_health_check():
"""Test that _perform_health_check_and_save passes model_id to perform_health_check so health checks run by model id."""
model_list = [
{
"model_name": "gpt-4",
"model_info": {"id": "deployment-abc"},
"litellm_params": {"model": "gpt-4"},
},
]
healthy = [{"model": "gpt-4"}]
unhealthy = []
async def mock_perform_health_check(model_list, model=None, cli_model=None, details=True, model_id=None):
return healthy, unhealthy
with patch(
"litellm.proxy.health_endpoints._health_endpoints.perform_health_check",
side_effect=mock_perform_health_check,
) as mock_perform:
result = await _perform_health_check_and_save(
model_list=model_list,
target_model=None,
cli_model=None,
details=True,
prisma_client=None,
start_time=0.0,
user_id="user-1",
model_id="deployment-abc",
)
mock_perform.assert_called_once()
call_kwargs = mock_perform.call_args[1]
assert call_kwargs["model_id"] == "deployment-abc"
assert result["healthy_count"] == 1
assert result["unhealthy_count"] == 0
if __name__ == "__main__":
pytest.main([__file__])

View file

@ -925,6 +925,73 @@ def test_router_get_model_access_groups_team_only_models():
assert list(access_groups.keys()) == ["default-models"]
def test_cached_get_model_group_info():
"""
Test that _cached_get_model_group_info caches results and
invalidates on deployment changes.
"""
from litellm.types.router import Deployment, LiteLLM_Params
router = litellm.Router(
model_list=[
{
"model_name": "gpt-4",
"litellm_params": {"model": "gpt-4", "api_key": "fake"},
"model_info": {"tpm": 1000, "rpm": 100},
},
]
)
# First call should compute and cache
result1 = router._cached_get_model_group_info("gpt-4")
assert result1 is not None
assert result1.tpm == 1000
# Second call should hit cache (same object)
result2 = router._cached_get_model_group_info("gpt-4")
assert result1 is result2
# Add a deployment — cache should be invalidated
router.add_deployment(
Deployment(
model_name="gpt-4",
litellm_params=LiteLLM_Params(model="gpt-4", api_key="fake2"),
model_info={"tpm": 2000, "rpm": 200},
)
)
result3 = router._cached_get_model_group_info("gpt-4")
assert result3 is not result2
assert result3 is not None
assert result3.tpm == 3000 # 1000 + 2000
# Delete a deployment — cache should be invalidated
deployment_id = router.model_list[-1]["model_info"]["id"]
router.delete_deployment(id=deployment_id)
result4 = router._cached_get_model_group_info("gpt-4")
assert result4 is not result3
assert result4 is not None
assert result4.tpm == 1000
# set_model_list — cache should be invalidated
router.set_model_list(
[
{
"model_name": "gpt-4",
"litellm_params": {"model": "gpt-4", "api_key": "fake"},
"model_info": {"tpm": 5000},
},
]
)
result5 = router._cached_get_model_group_info("gpt-4")
assert result5 is not result4
assert result5 is not None
assert result5.tpm == 5000
# Verify cache still works after invalidation
result6 = router._cached_get_model_group_info("gpt-4")
assert result5 is result6
def test_get_model_access_groups_caching():
"""
Test that get_model_access_groups caches the no-args result
@ -1297,6 +1364,61 @@ async def test_acompletion_streaming_iterator_edge_cases():
print("✓ Edge case tests passed!")
@pytest.mark.asyncio
async def test_acompletion_streaming_iterator_preserves_hidden_params():
"""
Regression test: FallbackStreamWrapper must copy _hidden_params from the
original CustomStreamWrapper so that x-litellm-overhead-duration-ms (and
other hidden params) are present in the proxy response headers for streaming.
"""
from unittest.mock import MagicMock
router = litellm.Router(
model_list=[
{
"model_name": "gpt-4",
"litellm_params": {"model": "gpt-4", "api_key": "fake-key"},
}
],
)
# Simulate a CustomStreamWrapper that already has timing metadata set by
# update_response_metadata (litellm_overhead_time_ms, _response_ms, etc.)
mock_response = MagicMock()
mock_response.model = "gpt-4"
mock_response.custom_llm_provider = "openai"
mock_response.logging_obj = MagicMock()
mock_response._hidden_params = {
"litellm_overhead_time_ms": 12.34,
"_response_ms": 500.0,
"litellm_call_id": "test-call-id",
"api_base": "https://api.openai.com",
"additional_headers": {},
}
# Make the mock iterable (yields nothing — we only care about hidden_params)
async def _empty():
return
yield # make it an async generator
mock_response.__aiter__ = lambda self: _empty().__aiter__()
result = await router._acompletion_streaming_iterator(
model_response=mock_response,
messages=[{"role": "user", "content": "hi"}],
initial_kwargs={"model": "gpt-4", "stream": True},
)
# The returned FallbackStreamWrapper must carry the original _hidden_params
assert hasattr(result, "_hidden_params"), "result must have _hidden_params"
assert result._hidden_params.get("litellm_overhead_time_ms") == 12.34, (
"litellm_overhead_time_ms must be preserved — "
"this is what drives x-litellm-overhead-duration-ms in streaming responses"
)
assert result._hidden_params.get("litellm_call_id") == "test-call-id"
assert result._hidden_params.get("_response_ms") == 500.0
@pytest.mark.asyncio
async def test_async_function_with_fallbacks_common_utils():
"""Test the async_function_with_fallbacks_common_utils method"""
@ -1858,7 +1980,7 @@ def test_get_deployment_credentials_with_provider_resolves_credential_name():
litellm_credential_name to actual credential values (for UI-created models).
"""
from litellm.types.utils import CredentialItem
# Setup credential list with a test credential
litellm.credential_list = [
CredentialItem(

View file

@ -1757,29 +1757,6 @@
"url": "https://opencollective.com/libvips"
}
},
"node_modules/@isaacs/balanced-match": {
"version": "4.0.1",
"resolved": "https://registry.npmjs.org/@isaacs/balanced-match/-/balanced-match-4.0.1.tgz",
"integrity": "sha512-yzMTt9lEb8Gv7zRioUilSglI0c0smZ9k5D65677DLWLtWJaXIS3CqcGyUFByYKlnUj6TkjLVs54fBl6+TiGQDQ==",
"dev": true,
"license": "MIT",
"engines": {
"node": "20 || >=22"
}
},
"node_modules/@isaacs/brace-expansion": {
"version": "5.0.1",
"resolved": "https://registry.npmjs.org/@isaacs/brace-expansion/-/brace-expansion-5.0.1.tgz",
"integrity": "sha512-WMz71T1JS624nWj2n2fnYAuPovhv7EUhk69R6i9dsVyzxt5eM3bjwvgk9L+APE1TRscGysAVMANkB0jh0LQZrQ==",
"dev": true,
"license": "MIT",
"dependencies": {
"@isaacs/balanced-match": "^4.0.1"
},
"engines": {
"node": "20 || >=22"
}
},
"node_modules/@istanbuljs/schema": {
"version": "0.1.3",
"resolved": "https://registry.npmjs.org/@istanbuljs/schema/-/schema-0.1.3.tgz",
@ -3696,32 +3673,6 @@
"typescript": ">=4.8.4 <6.0.0"
}
},
"node_modules/@typescript-eslint/typescript-estree/node_modules/brace-expansion": {
"version": "2.0.2",
"resolved": "https://registry.npmjs.org/brace-expansion/-/brace-expansion-2.0.2.tgz",
"integrity": "sha512-Jt0vHyM+jmUBqojB7E1NIYadt0vI0Qxjxd2TErW94wDz+E2LAm5vKMXXwg6ZZBTHPuUlDgQHKXvjGBdfcF1ZDQ==",
"dev": true,
"license": "MIT",
"dependencies": {
"balanced-match": "^1.0.0"
}
},
"node_modules/@typescript-eslint/typescript-estree/node_modules/minimatch": {
"version": "9.0.5",
"resolved": "https://registry.npmjs.org/minimatch/-/minimatch-9.0.5.tgz",
"integrity": "sha512-G6T0ZX48xgozx7587koeX9Ys2NYy6Gmv//P89sEte9V9whIapMNF4idKxnW2QtCcLiTWlb/wfCabAtAFWhhBow==",
"dev": true,
"license": "ISC",
"dependencies": {
"brace-expansion": "^2.0.1"
},
"engines": {
"node": ">=16 || 14 >=14.17"
},
"funding": {
"url": "https://github.com/sponsors/isaacs"
}
},
"node_modules/@typescript-eslint/utils": {
"version": "8.54.0",
"resolved": "https://registry.npmjs.org/@typescript-eslint/utils/-/utils-8.54.0.tgz",
@ -4749,11 +4700,14 @@
}
},
"node_modules/balanced-match": {
"version": "1.0.2",
"resolved": "https://registry.npmjs.org/balanced-match/-/balanced-match-1.0.2.tgz",
"integrity": "sha512-3oSeUO0TMV67hN1AmbXsK4yaqU7tjiHlbxRDZOpH0KW9+CeX4bRAaX0Anxt0tx2MrpRpWwQaPwIlISEJhYU5Pw==",
"version": "4.0.4",
"resolved": "https://registry.npmjs.org/balanced-match/-/balanced-match-4.0.4.tgz",
"integrity": "sha512-BLrgEcRTwX2o6gGxGOCNyMvGSp35YofuYzw9h1IMTRmKqttAZZVU67bdb9Pr2vUHA8+j3i2tJfjO6C6+4myGTA==",
"dev": true,
"license": "MIT"
"license": "MIT",
"engines": {
"node": "18 || 20 || >=22"
}
},
"node_modules/baseline-browser-mapping": {
"version": "2.9.19",
@ -4787,14 +4741,16 @@
}
},
"node_modules/brace-expansion": {
"version": "1.1.12",
"resolved": "https://registry.npmjs.org/brace-expansion/-/brace-expansion-1.1.12.tgz",
"integrity": "sha512-9T9UjW3r0UW5c1Q7GTwllptXwhvYmEzFhzMfZ9H7FQWt+uZePjZPjBP/W1ZEyZ1twGWom5/56TF4lPcqjnDHcg==",
"version": "5.0.3",
"resolved": "https://registry.npmjs.org/brace-expansion/-/brace-expansion-5.0.3.tgz",
"integrity": "sha512-fy6KJm2RawA5RcHkLa1z/ScpBeA762UF9KmZQxwIbDtRJrgLzM10depAiEQ+CXYcoiqW1/m96OAAoke2nE9EeA==",
"dev": true,
"license": "MIT",
"dependencies": {
"balanced-match": "^1.0.0",
"concat-map": "0.0.1"
"balanced-match": "^4.0.2"
},
"engines": {
"node": "18 || 20 || >=22"
}
},
"node_modules/braces": {
@ -5149,13 +5105,6 @@
"integrity": "sha512-VRhuHOLoKYOy4UbilLbUzbYg93XLjv2PncJC50EuTWPA3gaja1UjBsUP/D/9/juV3vQFr6XBEzn9KCAHdUvOHw==",
"license": "MIT"
},
"node_modules/concat-map": {
"version": "0.0.1",
"resolved": "https://registry.npmjs.org/concat-map/-/concat-map-0.0.1.tgz",
"integrity": "sha512-/Srv4dswyQNBfohGpz9o6Yb3Gz3SrUDqBH5rTuhGR7ahtlbYKnVxw2bCFMRljaA7EXHaXZ8wsHdodFvbkhKmqg==",
"dev": true,
"license": "MIT"
},
"node_modules/copy-to-clipboard": {
"version": "3.3.3",
"resolved": "https://registry.npmjs.org/copy-to-clipboard/-/copy-to-clipboard-3.3.3.tgz",
@ -6924,22 +6873,6 @@
"node": ">=10.13.0"
}
},
"node_modules/glob/node_modules/minimatch": {
"version": "10.1.1",
"resolved": "https://registry.npmjs.org/minimatch/-/minimatch-10.1.1.tgz",
"integrity": "sha512-enIvLvRAFZYXJzkCYG5RKmPfrFArdLv+R+lbQ53BmIMLIry74bjKzX6iHAm8WYamJkhSSEabrWN5D97XnKObjQ==",
"dev": true,
"license": "BlueOak-1.0.0",
"dependencies": {
"@isaacs/brace-expansion": "^5.0.0"
},
"engines": {
"node": "20 || >=22"
},
"funding": {
"url": "https://github.com/sponsors/isaacs"
}
},
"node_modules/globals": {
"version": "14.0.0",
"resolved": "https://registry.npmjs.org/globals/-/globals-14.0.0.tgz",
@ -9035,16 +8968,19 @@
}
},
"node_modules/minimatch": {
"version": "3.1.2",
"resolved": "https://registry.npmjs.org/minimatch/-/minimatch-3.1.2.tgz",
"integrity": "sha512-J7p63hRiAjw1NDEww1W7i37+ByIrOWO5XQQAzZ3VOcL0PNybwpfmV/N05zFAzwQ9USyEcX6t3UO+K5aqBQOIHw==",
"version": "10.2.2",
"resolved": "https://registry.npmjs.org/minimatch/-/minimatch-10.2.2.tgz",
"integrity": "sha512-+G4CpNBxa5MprY+04MbgOw1v7So6n5JY166pFi9KfYwT78fxScCeSNQSNzp6dpPSW2rONOps6Ocam1wFhCgoVw==",
"dev": true,
"license": "ISC",
"license": "BlueOak-1.0.0",
"dependencies": {
"brace-expansion": "^1.1.7"
"brace-expansion": "^5.0.2"
},
"engines": {
"node": "*"
"node": "18 || 20 || >=22"
},
"funding": {
"url": "https://github.com/sponsors/isaacs"
}
},
"node_modules/minimist": {
@ -12004,32 +11940,6 @@
"node": ">=18"
}
},
"node_modules/test-exclude/node_modules/brace-expansion": {
"version": "2.0.2",
"resolved": "https://registry.npmjs.org/brace-expansion/-/brace-expansion-2.0.2.tgz",
"integrity": "sha512-Jt0vHyM+jmUBqojB7E1NIYadt0vI0Qxjxd2TErW94wDz+E2LAm5vKMXXwg6ZZBTHPuUlDgQHKXvjGBdfcF1ZDQ==",
"dev": true,
"license": "MIT",
"dependencies": {
"balanced-match": "^1.0.0"
}
},
"node_modules/test-exclude/node_modules/minimatch": {
"version": "9.0.5",
"resolved": "https://registry.npmjs.org/minimatch/-/minimatch-9.0.5.tgz",
"integrity": "sha512-G6T0ZX48xgozx7587koeX9Ys2NYy6Gmv//P89sEte9V9whIapMNF4idKxnW2QtCcLiTWlb/wfCabAtAFWhhBow==",
"dev": true,
"license": "ISC",
"dependencies": {
"brace-expansion": "^2.0.1"
},
"engines": {
"node": ">=16 || 14 >=14.17"
},
"funding": {
"url": "https://github.com/sponsors/isaacs"
}
},
"node_modules/thenify": {
"version": "3.3.1",
"resolved": "https://registry.npmjs.org/thenify/-/thenify-3.3.1.tgz",
@ -13085,21 +12995,6 @@
"type": "github",
"url": "https://github.com/sponsors/wooorm"
}
},
"node_modules/@next/swc-win32-ia32-msvc": {
"version": "14.2.33",
"resolved": "https://registry.npmjs.org/@next/swc-win32-ia32-msvc/-/swc-win32-ia32-msvc-14.2.33.tgz",
"integrity": "sha512-pc9LpGNKhJ0dXQhZ5QMmYxtARwwmWLpeocFmVG5Z0DzWq5Uf0izcI8tLc+qOpqxO1PWqZ5A7J1blrUIKrIFc7Q==",
"cpu": [
"ia32"
],
"optional": true,
"os": [
"win32"
],
"engines": {
"node": ">= 10"
}
}
}
}

View file

@ -86,14 +86,25 @@
"mermaid": ">=11.10.0",
"js-yaml": ">=4.1.1",
"glob": ">=11.1.0",
"tar": ">=7.5.7",
"tar": ">=7.5.8",
"minimatch": ">=10.2.1",
"@isaacs/brace-expansion": ">=5.0.1",
"node-forge": ">=1.3.2",
"lodash-es": ">=4.17.23",
"lodash": ">=4.17.23"
"lodash": ">=4.17.23",
"@babel/traverse": ">=7.23.2",
"ws": ">=7.5.10",
"http-proxy-middleware": ">=2.0.9",
"tar-fs": ">=2.1.4",
"webpack-dev-middleware": ">=5.3.4",
"braces": ">=3.0.3",
"axios": ">=0.30.2",
"webpack": ">=5.94.0",
"serve-static": ">=1.16.0",
"path-to-regexp": ">=0.1.12"
},
"engines": {
"node": ">=18.17.0",
"npm": ">=8.3.0"
}
}
}

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