mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
Merge remote-tracking branch 'upstream/main' into litellm_fix_leakage
This commit is contained in:
commit
e150547446
135 changed files with 10488 additions and 1111 deletions
20
Dockerfile
20
Dockerfile
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 |
|
||||
|
|
|
|||
|
|
@ -6,4 +6,4 @@ metadata:
|
|||
data:
|
||||
config.yaml: |
|
||||
{{ .Values.proxy_config | toYaml | indent 6 }}
|
||||
{{- end }}
|
||||
{{- end }}
|
||||
|
|
|
|||
|
|
@ -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 }}
|
||||
|
|
|
|||
|
|
@ -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: {}
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
145
docs/my-website/blog/gpt_5_3_codex/index.md
Normal file
145
docs/my-website/blog/gpt_5_3_codex/index.md
Normal 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.
|
||||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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
|
||||
---
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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**
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
}
|
||||
}
|
||||
|
|
@ -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>
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -162,6 +162,7 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig):
|
|||
"service_tier",
|
||||
"safety_identifier",
|
||||
"prompt_cache_key",
|
||||
"prompt_cache_retention",
|
||||
"store",
|
||||
] # works across all models
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
)
|
||||
)
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
179
litellm/proxy/db/tool_registry_writer.py
Normal file
179
litellm/proxy/db/tool_registry_writer.py
Normal 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 {}
|
||||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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__}"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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=(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
149
litellm/proxy/management_endpoints/tool_management_endpoints.py
Normal file
149
litellm/proxy/management_endpoints/tool_management_endpoints.py
Normal 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))
|
||||
|
|
@ -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,
|
||||
)
|
||||
|
|
@ -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.",
|
||||
}
|
||||
)
|
||||
|
|
@ -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"},
|
||||
)
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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())
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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[
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
42
litellm/types/tool_management.py
Normal file
42
litellm/types/tool_management.py
Normal 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
|
||||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
18
package.json
18
package.json
|
|
@ -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"
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -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"
|
||||
}
|
||||
}
|
||||
|
|
@ -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"], [])
|
||||
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
@ -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}")
|
||||
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
@ -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"}
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
197
tests/test_litellm/proxy/db/test_tool_registry_writer.py
Normal file
197
tests/test_litellm/proxy/db/test_tool_registry_writer.py
Normal 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()
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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={},
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
)
|
||||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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')}"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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",
|
||||
)
|
||||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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__])
|
||||
|
|
@ -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(
|
||||
|
|
|
|||
151
ui/litellm-dashboard/package-lock.json
generated
151
ui/litellm-dashboard/package-lock.json
generated
|
|
@ -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"
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
Loading…
Add table
Reference in a new issue