Merge pull request #27559 from BerriAI/litellm_internal_staging
Some checks failed
Read Version from pyproject.toml / read-version (push) Has been cancelled
CodeQL / Analyze (actions) (push) Has been cancelled
CodeQL / Analyze (javascript-typescript) (push) Has been cancelled
CodeQL / Analyze (python) (push) Has been cancelled
CodSpeed Benchmarks / benchmarks (push) Has been cancelled
Helm unit test / unit-test (push) Has been cancelled
Scorecard supply-chain security / Scorecard analysis (push) Has been cancelled
Unit Tests: Caching (Redis) / caching-redis (push) Has been cancelled
Unit Tests: Proxy DB Operations / assert-shard-coverage (push) Has been cancelled
Unit Tests: Security / security (push) Has been cancelled
GitHub Actions Security Analysis / zizmor (push) Has been cancelled
Unit Tests: Proxy DB Operations / auth-checks (push) Has been cancelled
Unit Tests: Proxy DB Operations / budgets (push) Has been cancelled
Unit Tests: Proxy DB Operations / custom-logging (push) Has been cancelled
Unit Tests: Proxy DB Operations / db-and-spend (push) Has been cancelled
Unit Tests: Proxy DB Operations / endpoints-and-responses (push) Has been cancelled
Unit Tests: Proxy DB Operations / guardrails-hooks (push) Has been cancelled
Unit Tests: Proxy DB Operations / jwt-and-keys (push) Has been cancelled
Unit Tests: Proxy DB Operations / key-generation (push) Has been cancelled
Unit Tests: Proxy DB Operations / logging-misc (push) Has been cancelled
Unit Tests: Proxy DB Operations / proxy-runtime (push) Has been cancelled
Unit Tests: Proxy DB Operations / proxy-server-core (push) Has been cancelled
Unit Tests: Proxy DB Operations / schema-migration (push) Has been cancelled
Unit Tests: Proxy DB Operations / proxy-utils (push) Has been cancelled

[Infra] Promote Internal Staging to main
This commit is contained in:
ryan-crabbe-berri 2026-05-09 15:56:00 -07:00 • committed by GitHub
commit e182a5e0ba
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
87 changed files with 3481 additions and 3924 deletions

View file

@ -241,10 +241,27 @@ When opening issues or pull requests, follow these templates:
### Running the proxy server
Start the proxy with a config file:
Create a minimal config file and start the proxy:
```yaml
# config.yaml
model_list:
- model_name: fake-openai-endpoint
litellm_params:
model: openai/fake-model
api_key: fake-key
api_base: https://fake-api.example.com
general_settings:
master_key: sk-1234
litellm_settings:
drop_params: True
telemetry: False
```
```bash
uv run litellm --config dev_config.yaml --port 4000
uv run litellm --config config.yaml --port 4000
```
The proxy takes ~15-20 seconds to fully start (it runs Prisma migrations on boot). Wait for `/health` to return before sending requests. Without a PostgreSQL `DATABASE_URL`, the proxy connects to a default Neon dev database embedded in the `litellm-proxy-extras` package.

View file

@ -146,7 +146,7 @@ LiteLLM is a unified interface for 100+ LLM providers with two main components:
- **Bound large result sets.** Prisma materializes full results in memory. For results over ~10 MB, paginate with `take`/`skip` or `cursor`/`take`, always with an explicit `order`. Prefer cursor-based pagination (`skip` is O(n)). Don't paginate naturally small result sets.
- **Limit fetched columns on wide tables.** Use `select` to fetch only needed fields — returns a partial object, so downstream code must not access unselected fields.
- **Check index coverage.** For new or modified queries, check `schema.prisma` for a supporting index. Prefer extending an existing index (e.g. `@@index([a])` → `@@index([a, b])`) over adding a new one, unless it's a `@@unique`. Only add indexes for large/frequent queries.
- **Keep schema files in sync.** Apply schema changes to all `schema.prisma` copies (`schema.prisma`, `litellm/proxy/`, `litellm-proxy-extras/`, `litellm-js/spend-logs/` for SpendLogs) with a migration under `litellm-proxy-extras/litellm_proxy_extras/migrations/`.
- **Keep schema files in sync.** Apply schema changes to all `schema.prisma` copies (`schema.prisma`, `litellm/proxy/`, `litellm-proxy-extras/`) with a migration under `litellm-proxy-extras/litellm_proxy_extras/migrations/`.
### Setup Wizard (`litellm/setup_wizard.py`)
- The wizard is implemented as a single `SetupWizard` class with `@staticmethod` methods — keep it that way. No module-level functions except `run_setup_wizard()` (the public entrypoint) and pure helpers (color, ANSI).

View file

@ -1,18 +0,0 @@
# Use the provided base image
FROM ghcr.io/berriai/litellm:main-latest@sha256:7c311546c25e7bb6e8cafede9fcd3d0d622ac636b5c9418befaa32e85dfb0186
# Set the working directory to /app
WORKDIR /app
# Copy the configuration file into the container at /app
COPY config.yaml .
# Make sure your docker/entrypoint.sh is executable
# Convert Windows line endings to Unix
RUN sed -i 's/\r$//' docker/entrypoint.sh && chmod +x docker/entrypoint.sh
# Expose the necessary port
EXPOSE 4000/tcp
# Override the CMD instruction with your desired command and arguments
CMD ["--port", "4000", "--config", "config.yaml", "--detailed_debug", "--run_gunicorn"]

View file

@ -100,6 +100,16 @@ spec:
- name: DATABASE_URL
value: {{ .Values.db.url | quote }}
{{- end }}
{{- if and .Values.db.useExisting .Values.db.secret.readReplicaUrlKey }}
- name: DATABASE_URL_READ_REPLICA
valueFrom:
secretKeyRef:
name: {{ .Values.db.secret.name }}
key: {{ .Values.db.secret.readReplicaUrlKey }}
{{- else if .Values.db.readReplicaUrl }}
- name: DATABASE_URL_READ_REPLICA
value: {{ .Values.db.readReplicaUrl | quote }}
{{- end }}
- name: PROXY_MASTER_KEY
valueFrom:
secretKeyRef:

View file

@ -252,6 +252,26 @@ db:
passwordKey: password
# Optional: when set, DATABASE_HOST will be sourced from this secret key instead of db.endpoint
endpointKey: ""
# Optional: when set, DATABASE_URL_READ_REPLICA will be sourced from this
# secret key instead of db.readReplicaUrl. Prefer this over the plain
# value: read-replica URLs typically embed credentials, and a value
# written to db.readReplicaUrl ends up visible in the rendered pod spec
# and the Helm release secret.
readReplicaUrlKey: ""
# Optional read-replica routing. When set, the proxy sends read-only
# queries (find_*, count, group_by, query_raw/_first) to this URL while
# writes continue to go to db.url. Useful for Aurora-style clusters with
# separate reader/writer endpoints. Leave empty to keep single-DB behavior.
# When IAM_TOKEN_DB_AUTH is enabled, the reader URL is auto-refreshed
# alongside the writer (host/port/user/db are parsed from this URL once
# at startup; only the IAM token rotates).
#
# If the URL embeds credentials, prefer db.secret.readReplicaUrlKey over
# this field — the plain value is rendered into the pod spec and the
# Helm release secret. This field is intended for credential-less URLs
# only (e.g. when IAM_TOKEN_DB_AUTH supplies the token at runtime).
readReplicaUrl: ""
# Use the Stackgres Helm chart to deploy an instance of a Stackgres cluster.
# The Stackgres Operator must already be installed within the target

View file

@ -1,56 +0,0 @@
apiVersion: apps/v1
kind: Deployment
metadata:
name: litellm-deployment
spec:
replicas: 3
selector:
matchLabels:
app: litellm
template:
metadata:
labels:
app: litellm
spec:
containers:
- name: litellm-container
image: ghcr.io/berriai/litellm:main-latest
imagePullPolicy: Always
env:
- name: AZURE_API_KEY
value: "d6f****"
- name: AZURE_API_BASE
value: "https://openai"
- name: LITELLM_MASTER_KEY
value: "sk-1234"
- name: DATABASE_URL
value: "postgresql://ishaan*********"
args:
- "--config"
- "/app/proxy_config.yaml" # Update the path to mount the config file
volumeMounts: # Define volume mount for proxy_config.yaml
- name: config-volume
mountPath: /app
readOnly: true
livenessProbe:
httpGet:
path: /health/liveliness
port: 4000
initialDelaySeconds: 120
periodSeconds: 15
successThreshold: 1
failureThreshold: 3
timeoutSeconds: 10
readinessProbe:
httpGet:
path: /health/readiness
port: 4000
initialDelaySeconds: 120
periodSeconds: 15
successThreshold: 1
failureThreshold: 3
timeoutSeconds: 10
volumes: # Define volume to mount proxy_config.yaml
- name: config-volume
configMap:
name: litellm-config

View file

@ -1,12 +0,0 @@
apiVersion: v1
kind: Service
metadata:
name: litellm-service
spec:
selector:
app: litellm
ports:
- protocol: TCP
port: 4000
targetPort: 4000
type: LoadBalancer

View file

@ -1,13 +0,0 @@
model_list:
- model_name: fake-openai-endpoint
litellm_params:
model: openai/fake-model
api_key: fake-key
api_base: https://exampleopenaiendpoint-production.up.railway.app/
general_settings:
master_key: sk-1234
litellm_settings:
drop_params: True
telemetry: False

View file

@ -16,6 +16,11 @@ services:
- "4000:4000" # Map the container port to the host, change the host port if necessary
environment:
DATABASE_URL: "postgresql://llmproxy:dbpassword9090@db:5432/litellm"
# Optional: route read-only queries (find_*, count, group_by, query_raw/_first)
# to a separate reader endpoint, e.g. an Aurora reader. Leave unset for
# single-DB deployments. With IAM_TOKEN_DB_AUTH enabled, the reader URL
# is auto-refreshed alongside the writer.
# DATABASE_URL_READ_REPLICA: "postgresql://llmproxy:dbpassword9090@db-reader:5432/litellm"
STORE_MODEL_IN_DB: "True" # allows adding models to proxy via UI
env_file:
- .env # Load local .env file

View file

@ -1,68 +0,0 @@
# Base image for building
ARG LITELLM_BUILD_IMAGE=python:3.11-alpine@sha256:f07e2ace46f560f09a6eeec7b4913b80ee99546e749ef82342a419a326620856
# Runtime image
ARG LITELLM_RUNTIME_IMAGE=python:3.11-alpine@sha256:f07e2ace46f560f09a6eeec7b4913b80ee99546e749ef82342a419a326620856
ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:240fb85ab0f263ef12f492d8476aa3a2e4e1e333f7d67fbdd923d00a506a516a
FROM $UV_IMAGE AS uvbin
FROM $LITELLM_BUILD_IMAGE AS builder
WORKDIR /app
COPY --from=uvbin /uv /usr/local/bin/uv
COPY --from=uvbin /uvx /usr/local/bin/uvx
RUN apk add --no-cache gcc python3-dev musl-dev nodejs npm libsndfile
ENV PRISMA_BINARY_CACHE_DIR=/app/.cache/prisma-python/binaries \
UV_PROJECT_ENVIRONMENT=/app/.venv \
UV_LINK_MODE=copy \
XDG_CACHE_HOME=/app/.cache \
PATH="/app/.venv/bin:${PATH}"
# Copy dependency metadata first for layer caching
COPY pyproject.toml uv.lock ./
COPY enterprise/pyproject.toml enterprise/
COPY litellm-proxy-extras/pyproject.toml litellm-proxy-extras/
# Install third-party dependencies (cached unless pyproject.toml/uv.lock change)
RUN uv sync --frozen --no-install-project --no-install-workspace --no-default-groups --no-editable \
--extra proxy \
--extra proxy-runtime \
--extra extra_proxy \
--extra semantic-router \
--python python3
# Copy full source tree
COPY . .
# Install project and workspace packages (fast - deps already cached)
RUN uv sync --frozen --no-default-groups --no-editable \
--extra proxy \
--extra proxy-runtime \
--extra extra_proxy \
--extra semantic-router \
--python python3
RUN prisma generate --schema=./schema.prisma
RUN sed -i 's/\r$//' docker/entrypoint.sh && chmod +x docker/entrypoint.sh && \
sed -i 's/\r$//' docker/prod_entrypoint.sh && chmod +x docker/prod_entrypoint.sh
FROM $LITELLM_RUNTIME_IMAGE AS runtime
RUN apk upgrade --no-cache && apk add --no-cache libsndfile nodejs npm
WORKDIR /app
ENV PRISMA_BINARY_CACHE_DIR=/app/.cache/prisma-python/binaries \
XDG_CACHE_HOME=/app/.cache \
PATH="/app/.venv/bin:${PATH}"
COPY --from=builder /app /app
EXPOSE 4000/tcp
ENTRYPOINT ["docker/prod_entrypoint.sh"]
CMD ["--port", "4000"]

View file

@ -1,86 +0,0 @@
# Use the provided base image
# NOTE: This is a dev/branch-specific tag. Update digest when the base image is rebuilt.
FROM ghcr.io/berriai/litellm:litellm_fwd_server_root_path-dev
# Set the working directory to /app
WORKDIR /app
# Install Node.js and npm (adjust version as needed)
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 nodejs npm && \
npm install -g npm@11.12.1 tar@7.5.11 glob@11.1.0 @isaacs/brace-expansion@5.0.1 minimatch@10.2.4 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"; \
done && \
find "$GLOBAL/npm" -type d -name "glob" -path "*/node_modules/glob" | while read d; do \
rm -rf "$d" && cp -rL "$GLOBAL/glob" "$d"; \
done && \
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 && \
find /usr/local/lib /usr/lib -path "*/node_modules/npm/package.json" -exec \
sed -i 's/"tar": "\^7\.5\.[0-9]*"/"tar": "^7.5.10"/g; s/"minimatch": "\^10\.[0-9.]*"/"minimatch": "^10.2.4"/g' {} + 2>/dev/null && \
npm cache clean --force && \
apt-get purge -y npm
# Copy the UI source into the container
COPY ./ui/litellm-dashboard /app/ui/litellm-dashboard
# Set an environment variable for UI_BASE_PATH
# This can be overridden at build time
# set UI_BASE_PATH to "<your server root path>/ui"
ENV UI_BASE_PATH="/prod/ui"
# Build the UI with the specified UI_BASE_PATH
WORKDIR /app/ui/litellm-dashboard
RUN npm ci
RUN UI_BASE_PATH=$UI_BASE_PATH npm run build
# Create the destination directory
RUN mkdir -p /app/litellm/proxy/_experimental/out
# Move the built files to the appropriate location
# Assuming the build output is in ./out directory
RUN rm -rf /app/litellm/proxy/_experimental/out/* && \
mv ./out/* /app/litellm/proxy/_experimental/out/
# Switch back to the main app directory
WORKDIR /app
# Make sure your docker/entrypoint.sh is executable
# Convert Windows line endings to Unix for entrypoint scripts
RUN sed -i 's/\r$//' docker/entrypoint.sh && chmod +x docker/entrypoint.sh
RUN sed -i 's/\r$//' docker/prod_entrypoint.sh && chmod +x docker/prod_entrypoint.sh
# Run as non-root user
RUN groupadd --gid 1000 appuser && useradd --uid 1000 --gid 1000 --no-create-home appuser \
&& chown -R appuser:appuser /app
USER appuser
# Expose the necessary port
EXPOSE 4000/tcp
HEALTHCHECK --interval=30s --timeout=5s --start-period=10s --retries=3 \
CMD ["python", "-c", "import urllib.request; urllib.request.urlopen('http://localhost:4000/health')"]
# Override the CMD instruction with your desired command and arguments
CMD ["--port", "4000", "--config", "config.yaml", "--detailed_debug"]

View file

@ -1,121 +0,0 @@
# Base image for building
ARG LITELLM_BUILD_IMAGE=python:3.13-slim@sha256:739e7213785e88c0f702dcdc12c0973afcbd606dbf021a589cab77d6b00b579d
# Runtime image
ARG LITELLM_RUNTIME_IMAGE=python:3.13-slim@sha256:739e7213785e88c0f702dcdc12c0973afcbd606dbf021a589cab77d6b00b579d
ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:240fb85ab0f263ef12f492d8476aa3a2e4e1e333f7d67fbdd923d00a506a516a
FROM $UV_IMAGE AS uvbin
FROM $LITELLM_BUILD_IMAGE AS builder
WORKDIR /app
USER root
COPY --from=uvbin /uv /usr/local/bin/uv
COPY --from=uvbin /uvx /usr/local/bin/uvx
RUN apt-get update && apt-get install -y --no-install-recommends \
gcc \
g++ \
python3-dev \
libssl-dev \
pkg-config \
nodejs \
npm \
&& rm -rf /var/lib/apt/lists/*
ENV PRISMA_BINARY_CACHE_DIR=/app/.cache/prisma-python/binaries \
UV_PROJECT_ENVIRONMENT=/app/.venv \
UV_LINK_MODE=copy \
XDG_CACHE_HOME=/app/.cache \
PATH="/app/.venv/bin:${PATH}"
# Copy dependency metadata first for layer caching
COPY pyproject.toml uv.lock ./
COPY enterprise/pyproject.toml enterprise/
COPY litellm-proxy-extras/pyproject.toml litellm-proxy-extras/
# Install third-party dependencies (cached unless pyproject.toml/uv.lock change)
RUN uv sync --frozen --no-install-project --no-install-workspace --no-default-groups --no-editable \
--extra proxy \
--extra proxy-runtime \
--extra extra_proxy \
--extra semantic-router \
--python python
# Copy full source tree
COPY . .
# Build Admin UI before final sync
RUN sed -i 's/\r$//' docker/build_admin_ui.sh && chmod +x docker/build_admin_ui.sh && ./docker/build_admin_ui.sh
# Install project and workspace packages (fast - deps already cached)
RUN uv sync --frozen --no-default-groups --no-editable \
--extra proxy \
--extra proxy-runtime \
--extra extra_proxy \
--extra semantic-router \
--python python
RUN prisma generate --schema=./schema.prisma
RUN sed -i 's/\r$//' docker/entrypoint.sh && chmod +x docker/entrypoint.sh && \
sed -i 's/\r$//' docker/prod_entrypoint.sh && chmod +x docker/prod_entrypoint.sh
FROM $LITELLM_RUNTIME_IMAGE AS runtime
USER root
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@11.12.1 tar@7.5.11 glob@11.1.0 @isaacs/brace-expansion@5.0.1 minimatch@10.2.4 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"; \
done \
&& find "$GLOBAL/npm" -type d -name "glob" -path "*/node_modules/glob" | while read d; do \
rm -rf "$d" && cp -rL "$GLOBAL/glob" "$d"; \
done \
&& 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 \
&& find /usr/local/lib /usr/lib -path "*/node_modules/npm/package.json" -exec \
sed -i 's/"tar": "\^7\.5\.[0-9]*"/"tar": "^7.5.10"/g; s/"minimatch": "\^10\.[0-9.]*"/"minimatch": "^10.2.4"/g' {} + 2>/dev/null \
&& npm cache clean --force \
&& apt-get purge -y npm
WORKDIR /app
ENV PRISMA_BINARY_CACHE_DIR=/app/.cache/prisma-python/binaries \
XDG_CACHE_HOME=/app/.cache \
PATH="/app/.venv/bin:${PATH}"
COPY --from=builder /app /app
EXPOSE 4000/tcp
ENTRYPOINT ["docker/prod_entrypoint.sh"]
CMD ["--port", "4000"]

View file

@ -1,30 +0,0 @@
ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:240fb85ab0f263ef12f492d8476aa3a2e4e1e333f7d67fbdd923d00a506a516a
FROM $UV_IMAGE AS uvbin
FROM python:3.13-slim@sha256:739e7213785e88c0f702dcdc12c0973afcbd606dbf021a589cab77d6b00b579d
WORKDIR /app
# Copy the uv binary and the health check script.
COPY --from=uvbin /uv /usr/local/bin/uv
COPY pyproject.toml uv.lock /app/
COPY scripts/health_check/health_check_client.py /app/health_check_client.py
# Resolve and install the health-check dependencies from the project lockfile
# so the runtime image stays self-contained and reproducible.
RUN uv export --frozen --no-default-groups --only-group healthcheck --no-emit-project --no-hashes --output-file /tmp/health-check-requirements.txt \
&& uv pip install --system -r /tmp/health-check-requirements.txt \
&& rm /tmp/health-check-requirements.txt \
&& rm /app/pyproject.toml /app/uv.lock \
&& chmod +x /app/health_check_client.py
# Run as non-root user
RUN groupadd --gid 1000 appuser && useradd --uid 1000 --gid 1000 --no-create-home appuser
USER appuser
# Health check
HEALTHCHECK --interval=30s --timeout=5s --start-period=5s --retries=3 \
CMD ["python", "/app/health_check_client.py", "--help"]
# Set entrypoint
ENTRYPOINT ["python", "/app/health_check_client.py"]

View file

@ -1,108 +0,0 @@
apiVersion: v1
entries:
litellm-helm:
- apiVersion: v2
appVersion: v1.43.18
created: "2024-08-19T23:58:25.331689+08:00"
dependencies:
- condition: db.deployStandalone
name: postgresql
repository: oci://registry-1.docker.io/bitnamicharts
version: '>=13.3.0'
- condition: redis.enabled
name: redis
repository: oci://registry-1.docker.io/bitnamicharts
version: '>=18.0.0'
description: Call all LLM APIs using the OpenAI format
digest: 0411df3dc42868be8af3ad3e00cb252790e6bd7ad15f5b77f1ca5214573a8531
name: litellm-helm
type: application
urls:
- https://berriai.github.io/litellm/litellm-helm-0.2.3.tgz
version: 0.2.3
postgresql:
- annotations:
category: Database
images: |
- name: os-shell
image: docker.io/bitnami/os-shell:12-debian-12-r16
- name: postgres-exporter
image: docker.io/bitnami/postgres-exporter:0.15.0-debian-12-r14
- name: postgresql
image: docker.io/bitnami/postgresql:16.2.0-debian-12-r6
licenses: Apache-2.0
apiVersion: v2
appVersion: 16.2.0
created: "2024-08-19T23:58:25.335716+08:00"
dependencies:
- name: common
repository: oci://registry-1.docker.io/bitnamicharts
tags:
- bitnami-common
version: 2.x.x
description: PostgreSQL (Postgres) is an open source object-relational database
known for reliability and data integrity. ACID-compliant, it supports foreign
keys, joins, views, triggers and stored procedures.
digest: 3c8125526b06833df32e2f626db34aeaedb29d38f03d15349db6604027d4a167
home: https://bitnami.com
icon: https://bitnami.com/assets/stacks/postgresql/img/postgresql-stack-220x234.png
keywords:
- postgresql
- postgres
- database
- sql
- replication
- cluster
maintainers:
- name: VMware, Inc.
url: https://github.com/bitnami/charts
name: postgresql
sources:
- https://github.com/bitnami/charts/tree/main/bitnami/postgresql
urls:
- https://berriai.github.io/litellm/charts/postgresql-14.3.1.tgz
version: 14.3.1
redis:
- annotations:
category: Database
images: |
- name: kubectl
image: docker.io/bitnami/kubectl:1.29.2-debian-12-r3
- name: os-shell
image: docker.io/bitnami/os-shell:12-debian-12-r16
- name: redis
image: docker.io/bitnami/redis:7.2.4-debian-12-r9
- name: redis-exporter
image: docker.io/bitnami/redis-exporter:1.58.0-debian-12-r4
- name: redis-sentinel
image: docker.io/bitnami/redis-sentinel:7.2.4-debian-12-r7
licenses: Apache-2.0
apiVersion: v2
appVersion: 7.2.4
created: "2024-08-19T23:58:25.339392+08:00"
dependencies:
- name: common
repository: oci://registry-1.docker.io/bitnamicharts
tags:
- bitnami-common
version: 2.x.x
description: Redis(R) is an open source, advanced key-value store. It is often
referred to as a data structure server since keys can contain strings, hashes,
lists, sets and sorted sets.
digest: b2fa1835f673a18002ca864c54fadac3c33789b26f6c5e58e2851b0b14a8f984
home: https://bitnami.com
icon: https://bitnami.com/assets/stacks/redis/img/redis-stack-220x234.png
keywords:
- redis
- keyvalue
- database
maintainers:
- name: VMware, Inc.
url: https://github.com/bitnami/charts
name: redis
sources:
- https://github.com/bitnami/charts/tree/main/bitnami/redis
urls:
- https://berriai.github.io/litellm/charts/redis-18.19.1.tgz
version: 18.19.1
generated: "2024-08-19T23:58:25.322532+08:00"

View file

@ -1,5 +0,0 @@
# Supply-chain hardening
# Packages needing lifecycle scripts: npm rebuild <pkg>
ignore-scripts=true
# Protects local npm install only — npm ci (used in CI) ignores this
min-release-age=3

View file

@ -1,8 +0,0 @@
```
npm install
npm run dev
```
```
npm run deploy
```

File diff suppressed because it is too large Load diff

View file

@ -1,14 +0,0 @@
{
"scripts": {
"dev": "wrangler dev src/index.ts",
"deploy": "wrangler deploy --minify src/index.ts"
},
"dependencies": {
"hono": "4.12.16",
"openai": "4.29.2"
},
"devDependencies": {
"@cloudflare/workers-types": "4.20260501.1",
"wrangler": "4.87.0"
}
}

View file

@ -1,59 +0,0 @@
import { Hono } from 'hono'
import { Context } from 'hono';
import { bearerAuth } from 'hono/bearer-auth'
import OpenAI from "openai";
const openai = new OpenAI({
apiKey: "sk-1234",
baseURL: "https://openai-endpoint.ishaanjaffer0324.workers.dev"
});
async function call_proxy() {
const completion = await openai.chat.completions.create({
messages: [{ role: "system", content: "You are a helpful assistant." }],
model: "gpt-3.5-turbo",
});
return completion
}
const app = new Hono()
// Middleware for API Key Authentication
const apiKeyAuth = async (c: Context, next: Function) => {
const apiKey = c.req.header('Authorization');
if (!apiKey || apiKey !== 'Bearer sk-1234') {
return c.text('Unauthorized', 401);
}
await next();
};
app.use('/*', apiKeyAuth)
app.get('/', (c) => {
return c.text('Hello Hono!')
})
// Handler for chat completions
const chatCompletionHandler = async (c: Context) => {
// Assuming your logic for handling chat completion goes here
// For demonstration, just returning a simple JSON response
const response = await call_proxy()
return c.json(response);
};
// Register the above handler for different POST routes with the apiKeyAuth middleware
app.post('/v1/chat/completions', chatCompletionHandler);
app.post('/chat/completions', chatCompletionHandler);
// Example showing how you might handle dynamic segments within the URL
// Here, using ':model*' to capture the rest of the path as a parameter 'model'
app.post('/openai/deployments/:model*/chat/completions', chatCompletionHandler);
export default app

View file

@ -1,17 +0,0 @@
{
"compilerOptions": {
"target": "ESNext",
"module": "ESNext",
"moduleResolution": "Bundler",
"strict": true,
"lib": [
"ESNext"
],
"types": [
"@cloudflare/workers-types"
],
"jsx": "react-jsx",
"jsxImportSource": "hono/jsx",
"skipLibCheck": true
},
}

View file

@ -1,18 +0,0 @@
name = "my-app"
compatibility_date = "2023-12-01"
# [vars]
# MY_VAR = "my-variable"
# [[kv_namespaces]]
# binding = "MY_KV_NAMESPACE"
# id = "xxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxx"
# [[r2_buckets]]
# binding = "MY_BUCKET"
# bucket_name = "my-bucket"
# [[d1_databases]]
# binding = "DB"
# database_name = "my-database"
# database_id = ""

View file

@ -1,5 +0,0 @@
# Supply-chain hardening
# Packages needing lifecycle scripts: npm rebuild <pkg>
ignore-scripts=true
# Protects local npm install only — npm ci (used in CI) ignores this
min-release-age=3

View file

@ -1,26 +0,0 @@
# Use the specific Node.js v20.11.0 image
FROM node:20.18.1-alpine3.20
# Set the working directory inside the container
WORKDIR /app
# Copy package.json and package-lock.json to the working directory
COPY ./litellm-js/spend-logs/package*.json ./
# Install dependencies
RUN npm ci
# Install Prisma globally
RUN npm install -g prisma
# Copy the rest of the application code
COPY ./litellm-js/spend-logs .
# Generate Prisma client
RUN npx prisma generate
# Expose the port that the Node.js server will run on
EXPOSE 3000
# Command to run the Node.js app with npm run dev
CMD ["npm", "run", "dev"]

View file

@ -1,8 +0,0 @@
```
npm install
npm run dev
```
```
open http://localhost:3000
```

View file

@ -1,597 +0,0 @@
{
"name": "spend-logs",
"lockfileVersion": 3,
"requires": true,
"packages": {
"": {
"dependencies": {
"@hono/node-server": "1.19.13",
"hono": "4.12.16"
},
"devDependencies": {
"@types/node": "20.19.25",
"tsx": "4.20.6"
}
},
"node_modules/@esbuild/aix-ppc64": {
"version": "0.25.12",
"resolved": "https://registry.npmjs.org/@esbuild/aix-ppc64/-/aix-ppc64-0.25.12.tgz",
"integrity": "sha512-Hhmwd6CInZ3dwpuGTF8fJG6yoWmsToE+vYgD4nytZVxcu1ulHpUQRAB1UJ8+N1Am3Mz4+xOByoQoSZf4D+CpkA==",
"cpu": [
"ppc64"
],
"dev": true,
"license": "MIT",
"optional": true,
"os": [
"aix"
],
"engines": {
"node": ">=18"
}
},
"node_modules/@esbuild/android-arm": {
"version": "0.25.12",
"resolved": "https://registry.npmjs.org/@esbuild/android-arm/-/android-arm-0.25.12.tgz",
"integrity": "sha512-VJ+sKvNA/GE7Ccacc9Cha7bpS8nyzVv0jdVgwNDaR4gDMC/2TTRc33Ip8qrNYUcpkOHUT5OZ0bUcNNVZQ9RLlg==",
"cpu": [
"arm"
],
"dev": true,
"license": "MIT",
"optional": true,
"os": [
"android"
],
"engines": {
"node": ">=18"
}
},
"node_modules/@esbuild/android-arm64": {
"version": "0.25.12",
"resolved": "https://registry.npmjs.org/@esbuild/android-arm64/-/android-arm64-0.25.12.tgz",
"integrity": "sha512-6AAmLG7zwD1Z159jCKPvAxZd4y/VTO0VkprYy+3N2FtJ8+BQWFXU+OxARIwA46c5tdD9SsKGZ/1ocqBS/gAKHg==",
"cpu": [
"arm64"
],
"dev": true,
"license": "MIT",
"optional": true,
"os": [
"android"
],
"engines": {
"node": ">=18"
}
},
"node_modules/@esbuild/android-x64": {
"version": "0.25.12",
"resolved": "https://registry.npmjs.org/@esbuild/android-x64/-/android-x64-0.25.12.tgz",
"integrity": "sha512-5jbb+2hhDHx5phYR2By8GTWEzn6I9UqR11Kwf22iKbNpYrsmRB18aX/9ivc5cabcUiAT/wM+YIZ6SG9QO6a8kg==",
"cpu": [
"x64"
],
"dev": true,
"license": "MIT",
"optional": true,
"os": [
"android"
],
"engines": {
"node": ">=18"
}
},
"node_modules/@esbuild/darwin-arm64": {
"version": "0.25.12",
"resolved": "https://registry.npmjs.org/@esbuild/darwin-arm64/-/darwin-arm64-0.25.12.tgz",
"integrity": "sha512-N3zl+lxHCifgIlcMUP5016ESkeQjLj/959RxxNYIthIg+CQHInujFuXeWbWMgnTo4cp5XVHqFPmpyu9J65C1Yg==",
"cpu": [
"arm64"
],
"dev": true,
"license": "MIT",
"optional": true,
"os": [
"darwin"
],
"engines": {
"node": ">=18"
}
},
"node_modules/@esbuild/darwin-x64": {
"version": "0.25.12",
"resolved": "https://registry.npmjs.org/@esbuild/darwin-x64/-/darwin-x64-0.25.12.tgz",
"integrity": "sha512-HQ9ka4Kx21qHXwtlTUVbKJOAnmG1ipXhdWTmNXiPzPfWKpXqASVcWdnf2bnL73wgjNrFXAa3yYvBSd9pzfEIpA==",
"cpu": [
"x64"
],
"dev": true,
"license": "MIT",
"optional": true,
"os": [
"darwin"
],
"engines": {
"node": ">=18"
}
},
"node_modules/@esbuild/freebsd-arm64": {
"version": "0.25.12",
"resolved": "https://registry.npmjs.org/@esbuild/freebsd-arm64/-/freebsd-arm64-0.25.12.tgz",
"integrity": "sha512-gA0Bx759+7Jve03K1S0vkOu5Lg/85dou3EseOGUes8flVOGxbhDDh/iZaoek11Y8mtyKPGF3vP8XhnkDEAmzeg==",
"cpu": [
"arm64"
],
"dev": true,
"license": "MIT",
"optional": true,
"os": [
"freebsd"
],
"engines": {
"node": ">=18"
}
},
"node_modules/@esbuild/freebsd-x64": {
"version": "0.25.12",
"resolved": "https://registry.npmjs.org/@esbuild/freebsd-x64/-/freebsd-x64-0.25.12.tgz",
"integrity": "sha512-TGbO26Yw2xsHzxtbVFGEXBFH0FRAP7gtcPE7P5yP7wGy7cXK2oO7RyOhL5NLiqTlBh47XhmIUXuGciXEqYFfBQ==",
"cpu": [
"x64"
],
"dev": true,
"license": "MIT",
"optional": true,
"os": [
"freebsd"
],
"engines": {
"node": ">=18"
}
},
"node_modules/@esbuild/linux-arm": {
"version": "0.25.12",
"resolved": "https://registry.npmjs.org/@esbuild/linux-arm/-/linux-arm-0.25.12.tgz",
"integrity": "sha512-lPDGyC1JPDou8kGcywY0YILzWlhhnRjdof3UlcoqYmS9El818LLfJJc3PXXgZHrHCAKs/Z2SeZtDJr5MrkxtOw==",
"cpu": [
"arm"
],
"dev": true,
"license": "MIT",
"optional": true,
"os": [
"linux"
],
"engines": {
"node": ">=18"
}
},
"node_modules/@esbuild/linux-arm64": {
"version": "0.25.12",
"resolved": "https://registry.npmjs.org/@esbuild/linux-arm64/-/linux-arm64-0.25.12.tgz",
"integrity": "sha512-8bwX7a8FghIgrupcxb4aUmYDLp8pX06rGh5HqDT7bB+8Rdells6mHvrFHHW2JAOPZUbnjUpKTLg6ECyzvas2AQ==",
"cpu": [
"arm64"
],
"dev": true,
"license": "MIT",
"optional": true,
"os": [
"linux"
],
"engines": {
"node": ">=18"
}
},
"node_modules/@esbuild/linux-ia32": {
"version": "0.25.12",
"resolved": "https://registry.npmjs.org/@esbuild/linux-ia32/-/linux-ia32-0.25.12.tgz",
"integrity": "sha512-0y9KrdVnbMM2/vG8KfU0byhUN+EFCny9+8g202gYqSSVMonbsCfLjUO+rCci7pM0WBEtz+oK/PIwHkzxkyharA==",
"cpu": [
"ia32"
],
"dev": true,
"license": "MIT",
"optional": true,
"os": [
"linux"
],
"engines": {
"node": ">=18"
}
},
"node_modules/@esbuild/linux-loong64": {
"version": "0.25.12",
"resolved": "https://registry.npmjs.org/@esbuild/linux-loong64/-/linux-loong64-0.25.12.tgz",
"integrity": "sha512-h///Lr5a9rib/v1GGqXVGzjL4TMvVTv+s1DPoxQdz7l/AYv6LDSxdIwzxkrPW438oUXiDtwM10o9PmwS/6Z0Ng==",
"cpu": [
"loong64"
],
"dev": true,
"license": "MIT",
"optional": true,
"os": [
"linux"
],
"engines": {
"node": ">=18"
}
},
"node_modules/@esbuild/linux-mips64el": {
"version": "0.25.12",
"resolved": "https://registry.npmjs.org/@esbuild/linux-mips64el/-/linux-mips64el-0.25.12.tgz",
"integrity": "sha512-iyRrM1Pzy9GFMDLsXn1iHUm18nhKnNMWscjmp4+hpafcZjrr2WbT//d20xaGljXDBYHqRcl8HnxbX6uaA/eGVw==",
"cpu": [
"mips64el"
],
"dev": true,
"license": "MIT",
"optional": true,
"os": [
"linux"
],
"engines": {
"node": ">=18"
}
},
"node_modules/@esbuild/linux-ppc64": {
"version": "0.25.12",
"resolved": "https://registry.npmjs.org/@esbuild/linux-ppc64/-/linux-ppc64-0.25.12.tgz",
"integrity": "sha512-9meM/lRXxMi5PSUqEXRCtVjEZBGwB7P/D4yT8UG/mwIdze2aV4Vo6U5gD3+RsoHXKkHCfSxZKzmDssVlRj1QQA==",
"cpu": [
"ppc64"
],
"dev": true,
"license": "MIT",
"optional": true,
"os": [
"linux"
],
"engines": {
"node": ">=18"
}
},
"node_modules/@esbuild/linux-riscv64": {
"version": "0.25.12",
"resolved": "https://registry.npmjs.org/@esbuild/linux-riscv64/-/linux-riscv64-0.25.12.tgz",
"integrity": "sha512-Zr7KR4hgKUpWAwb1f3o5ygT04MzqVrGEGXGLnj15YQDJErYu/BGg+wmFlIDOdJp0PmB0lLvxFIOXZgFRrdjR0w==",
"cpu": [
"riscv64"
],
"dev": true,
"license": "MIT",
"optional": true,
"os": [
"linux"
],
"engines": {
"node": ">=18"
}
},
"node_modules/@esbuild/linux-s390x": {
"version": "0.25.12",
"resolved": "https://registry.npmjs.org/@esbuild/linux-s390x/-/linux-s390x-0.25.12.tgz",
"integrity": "sha512-MsKncOcgTNvdtiISc/jZs/Zf8d0cl/t3gYWX8J9ubBnVOwlk65UIEEvgBORTiljloIWnBzLs4qhzPkJcitIzIg==",
"cpu": [
"s390x"
],
"dev": true,
"license": "MIT",
"optional": true,
"os": [
"linux"
],
"engines": {
"node": ">=18"
}
},
"node_modules/@esbuild/linux-x64": {
"version": "0.25.12",
"resolved": "https://registry.npmjs.org/@esbuild/linux-x64/-/linux-x64-0.25.12.tgz",
"integrity": "sha512-uqZMTLr/zR/ed4jIGnwSLkaHmPjOjJvnm6TVVitAa08SLS9Z0VM8wIRx7gWbJB5/J54YuIMInDquWyYvQLZkgw==",
"cpu": [
"x64"
],
"dev": true,
"license": "MIT",
"optional": true,
"os": [
"linux"
],
"engines": {
"node": ">=18"
}
},
"node_modules/@esbuild/netbsd-arm64": {
"version": "0.25.12",
"resolved": "https://registry.npmjs.org/@esbuild/netbsd-arm64/-/netbsd-arm64-0.25.12.tgz",
"integrity": "sha512-xXwcTq4GhRM7J9A8Gv5boanHhRa/Q9KLVmcyXHCTaM4wKfIpWkdXiMog/KsnxzJ0A1+nD+zoecuzqPmCRyBGjg==",
"cpu": [
"arm64"
],
"dev": true,
"license": "MIT",
"optional": true,
"os": [
"netbsd"
],
"engines": {
"node": ">=18"
}
},
"node_modules/@esbuild/netbsd-x64": {
"version": "0.25.12",
"resolved": "https://registry.npmjs.org/@esbuild/netbsd-x64/-/netbsd-x64-0.25.12.tgz",
"integrity": "sha512-Ld5pTlzPy3YwGec4OuHh1aCVCRvOXdH8DgRjfDy/oumVovmuSzWfnSJg+VtakB9Cm0gxNO9BzWkj6mtO1FMXkQ==",
"cpu": [
"x64"
],
"dev": true,
"license": "MIT",
"optional": true,
"os": [
"netbsd"
],
"engines": {
"node": ">=18"
}
},
"node_modules/@esbuild/openbsd-arm64": {
"version": "0.25.12",
"resolved": "https://registry.npmjs.org/@esbuild/openbsd-arm64/-/openbsd-arm64-0.25.12.tgz",
"integrity": "sha512-fF96T6KsBo/pkQI950FARU9apGNTSlZGsv1jZBAlcLL1MLjLNIWPBkj5NlSz8aAzYKg+eNqknrUJ24QBybeR5A==",
"cpu": [
"arm64"
],
"dev": true,
"license": "MIT",
"optional": true,
"os": [
"openbsd"
],
"engines": {
"node": ">=18"
}
},
"node_modules/@esbuild/openbsd-x64": {
"version": "0.25.12",
"resolved": "https://registry.npmjs.org/@esbuild/openbsd-x64/-/openbsd-x64-0.25.12.tgz",
"integrity": "sha512-MZyXUkZHjQxUvzK7rN8DJ3SRmrVrke8ZyRusHlP+kuwqTcfWLyqMOE3sScPPyeIXN/mDJIfGXvcMqCgYKekoQw==",
"cpu": [
"x64"
],
"dev": true,
"license": "MIT",
"optional": true,
"os": [
"openbsd"
],
"engines": {
"node": ">=18"
}
},
"node_modules/@esbuild/openharmony-arm64": {
"version": "0.25.12",
"resolved": "https://registry.npmjs.org/@esbuild/openharmony-arm64/-/openharmony-arm64-0.25.12.tgz",
"integrity": "sha512-rm0YWsqUSRrjncSXGA7Zv78Nbnw4XL6/dzr20cyrQf7ZmRcsovpcRBdhD43Nuk3y7XIoW2OxMVvwuRvk9XdASg==",
"cpu": [
"arm64"
],
"dev": true,
"license": "MIT",
"optional": true,
"os": [
"openharmony"
],
"engines": {
"node": ">=18"
}
},
"node_modules/@esbuild/sunos-x64": {
"version": "0.25.12",
"resolved": "https://registry.npmjs.org/@esbuild/sunos-x64/-/sunos-x64-0.25.12.tgz",
"integrity": "sha512-3wGSCDyuTHQUzt0nV7bocDy72r2lI33QL3gkDNGkod22EsYl04sMf0qLb8luNKTOmgF/eDEDP5BFNwoBKH441w==",
"cpu": [
"x64"
],
"dev": true,
"license": "MIT",
"optional": true,
"os": [
"sunos"
],
"engines": {
"node": ">=18"
}
},
"node_modules/@esbuild/win32-arm64": {
"version": "0.25.12",
"resolved": "https://registry.npmjs.org/@esbuild/win32-arm64/-/win32-arm64-0.25.12.tgz",
"integrity": "sha512-rMmLrur64A7+DKlnSuwqUdRKyd3UE7oPJZmnljqEptesKM8wx9J8gx5u0+9Pq0fQQW8vqeKebwNXdfOyP+8Bsg==",
"cpu": [
"arm64"
],
"dev": true,
"license": "MIT",
"optional": true,
"os": [
"win32"
],
"engines": {
"node": ">=18"
}
},
"node_modules/@esbuild/win32-ia32": {
"version": "0.25.12",
"resolved": "https://registry.npmjs.org/@esbuild/win32-ia32/-/win32-ia32-0.25.12.tgz",
"integrity": "sha512-HkqnmmBoCbCwxUKKNPBixiWDGCpQGVsrQfJoVGYLPT41XWF8lHuE5N6WhVia2n4o5QK5M4tYr21827fNhi4byQ==",
"cpu": [
"ia32"
],
"dev": true,
"license": "MIT",
"optional": true,
"os": [
"win32"
],
"engines": {
"node": ">=18"
}
},
"node_modules/@esbuild/win32-x64": {
"version": "0.25.12",
"resolved": "https://registry.npmjs.org/@esbuild/win32-x64/-/win32-x64-0.25.12.tgz",
"integrity": "sha512-alJC0uCZpTFrSL0CCDjcgleBXPnCrEAhTBILpeAp7M/OFgoqtAetfBzX0xM00MUsVVPpVjlPuMbREqnZCXaTnA==",
"cpu": [
"x64"
],
"dev": true,
"license": "MIT",
"optional": true,
"os": [
"win32"
],
"engines": {
"node": ">=18"
}
},
"node_modules/@hono/node-server": {
"version": "1.19.13",
"resolved": "https://registry.npmjs.org/@hono/node-server/-/node-server-1.19.13.tgz",
"integrity": "sha512-TsQLe4i2gvoTtrHje625ngThGBySOgSK3Xo2XRYOdqGN1teR8+I7vchQC46uLJi8OF62YTYA3AhSpumtkhsaKQ==",
"license": "MIT",
"engines": {
"node": ">=18.14.1"
},
"peerDependencies": {
"hono": "^4"
}
},
"node_modules/@types/node": {
"version": "20.19.25",
"resolved": "https://registry.npmjs.org/@types/node/-/node-20.19.25.tgz",
"integrity": "sha512-ZsJzA5thDQMSQO788d7IocwwQbI8B5OPzmqNvpf3NY/+MHDAS759Wo0gd2WQeXYt5AAAQjzcrTVC6SKCuYgoCQ==",
"dev": true,
"license": "MIT",
"dependencies": {
"undici-types": "~6.21.0"
}
},
"node_modules/esbuild": {
"version": "0.25.12",
"resolved": "https://registry.npmjs.org/esbuild/-/esbuild-0.25.12.tgz",
"integrity": "sha512-bbPBYYrtZbkt6Os6FiTLCTFxvq4tt3JKall1vRwshA3fdVztsLAatFaZobhkBC8/BrPetoa0oksYoKXoG4ryJg==",
"dev": true,
"hasInstallScript": true,
"license": "MIT",
"bin": {
"esbuild": "bin/esbuild"
},
"engines": {
"node": ">=18"
},
"optionalDependencies": {
"@esbuild/aix-ppc64": "0.25.12",
"@esbuild/android-arm": "0.25.12",
"@esbuild/android-arm64": "0.25.12",
"@esbuild/android-x64": "0.25.12",
"@esbuild/darwin-arm64": "0.25.12",
"@esbuild/darwin-x64": "0.25.12",
"@esbuild/freebsd-arm64": "0.25.12",
"@esbuild/freebsd-x64": "0.25.12",
"@esbuild/linux-arm": "0.25.12",
"@esbuild/linux-arm64": "0.25.12",
"@esbuild/linux-ia32": "0.25.12",
"@esbuild/linux-loong64": "0.25.12",
"@esbuild/linux-mips64el": "0.25.12",
"@esbuild/linux-ppc64": "0.25.12",
"@esbuild/linux-riscv64": "0.25.12",
"@esbuild/linux-s390x": "0.25.12",
"@esbuild/linux-x64": "0.25.12",
"@esbuild/netbsd-arm64": "0.25.12",
"@esbuild/netbsd-x64": "0.25.12",
"@esbuild/openbsd-arm64": "0.25.12",
"@esbuild/openbsd-x64": "0.25.12",
"@esbuild/openharmony-arm64": "0.25.12",
"@esbuild/sunos-x64": "0.25.12",
"@esbuild/win32-arm64": "0.25.12",
"@esbuild/win32-ia32": "0.25.12",
"@esbuild/win32-x64": "0.25.12"
}
},
"node_modules/fsevents": {
"version": "2.3.3",
"resolved": "https://registry.npmjs.org/fsevents/-/fsevents-2.3.3.tgz",
"integrity": "sha512-5xoDfX+fL7faATnagmWPpbFtwh/R77WmMMqqHGS65C3vvB0YHrgF+B1YmZ3441tMj5n63k0212XNoJwzlhffQw==",
"dev": true,
"hasInstallScript": true,
"license": "MIT",
"optional": true,
"os": [
"darwin"
],
"engines": {
"node": "^8.16.0 || ^10.6.0 || >=11.0.0"
}
},
"node_modules/get-tsconfig": {
"version": "4.14.0",
"resolved": "https://registry.npmjs.org/get-tsconfig/-/get-tsconfig-4.14.0.tgz",
"integrity": "sha512-yTb+8DXzDREzgvYmh6s9vHsSVCHeC0G3PI5bEXNBHtmshPnO+S5O7qgLEOn0I5QvMy6kpZN8K1NKGyilLb93wA==",
"dev": true,
"license": "MIT",
"dependencies": {
"resolve-pkg-maps": "^1.0.0"
},
"funding": {
"url": "https://github.com/privatenumber/get-tsconfig?sponsor=1"
}
},
"node_modules/hono": {
"version": "4.12.16",
"resolved": "https://registry.npmjs.org/hono/-/hono-4.12.16.tgz",
"integrity": "sha512-jN0ZewiNAWSe5khM3EyCmBb250+b40wWbwNILNfEvq84VREWwOIkuUsFONk/3i3nqkz7Oe1PcpM2mwQEK2L9Kg==",
"license": "MIT",
"engines": {
"node": ">=16.9.0"
}
},
"node_modules/resolve-pkg-maps": {
"version": "1.0.0",
"resolved": "https://registry.npmjs.org/resolve-pkg-maps/-/resolve-pkg-maps-1.0.0.tgz",
"integrity": "sha512-seS2Tj26TBVOC2NIc2rOe2y2ZO7efxITtLZcGSOnHHNOQ7CkiUBfw0Iw2ck6xkIhPwLhKNLS8BO+hEpngQlqzw==",
"dev": true,
"license": "MIT",
"funding": {
"url": "https://github.com/privatenumber/resolve-pkg-maps?sponsor=1"
}
},
"node_modules/tsx": {
"version": "4.20.6",
"resolved": "https://registry.npmjs.org/tsx/-/tsx-4.20.6.tgz",
"integrity": "sha512-ytQKuwgmrrkDTFP4LjR0ToE2nqgy886GpvRSpU0JAnrdBYppuY5rLkRUYPU1yCryb24SsKBTL/hlDQAEFVwtZg==",
"dev": true,
"license": "MIT",
"dependencies": {
"esbuild": "~0.25.0",
"get-tsconfig": "^4.7.5"
},
"bin": {
"tsx": "dist/cli.mjs"
},
"engines": {
"node": ">=18.0.0"
},
"optionalDependencies": {
"fsevents": "~2.3.3"
}
},
"node_modules/undici-types": {
"version": "6.21.0",
"resolved": "https://registry.npmjs.org/undici-types/-/undici-types-6.21.0.tgz",
"integrity": "sha512-iwDZqg0QAGrg9Rav5H4n0M64c3mkR59cJ6wQp+7C4nI0gsmExaedaYLNO44eT4AtBBwjbTiGPMlt2Md0T9H9JQ==",
"dev": true,
"license": "MIT"
}
}
}

View file

@ -1,13 +0,0 @@
{
"scripts": {
"dev": "tsx watch src/index.ts"
},
"dependencies": {
"@hono/node-server": "1.19.13",
"hono": "4.12.16"
},
"devDependencies": {
"@types/node": "20.19.25",
"tsx": "4.20.6"
}
}

View file

@ -1,29 +0,0 @@
generator client {
provider = "prisma-client-js"
}
datasource client {
provider = "postgresql"
url = env("DATABASE_URL")
}
model LiteLLM_SpendLogs {
request_id String @id
call_type String
api_key String @default("")
spend Float @default(0.0)
total_tokens Int @default(0)
prompt_tokens Int @default(0)
completion_tokens Int @default(0)
startTime DateTime
endTime DateTime
model String @default("")
api_base String @default("")
user String @default("")
metadata Json @default("{}")
cache_hit String @default("")
cache_key String @default("")
request_tags Json @default("[]")
team_id String?
end_user String?
}

View file

@ -1,32 +0,0 @@
export type LiteLLM_IncrementSpend = {
key_transactions: Array<LiteLLM_IncrementObject>, // [{"key": spend},..]
user_transactions: Array<LiteLLM_IncrementObject>,
team_transactions: Array<LiteLLM_IncrementObject>,
spend_logs_transactions: Array<LiteLLM_SpendLogs>
}
export type LiteLLM_IncrementObject = {
key: string,
spend: number
}
export type LiteLLM_SpendLogs = {
request_id: string; // @id means it's a unique identifier
call_type: string;
api_key: string; // @default("") means it defaults to an empty string if not provided
spend: number; // Float in Prisma corresponds to number in TypeScript
total_tokens: number; // Int in Prisma corresponds to number in TypeScript
prompt_tokens: number;
completion_tokens: number;
startTime: Date; // DateTime in Prisma corresponds to Date in TypeScript
endTime: Date;
model: string; // @default("") means it defaults to an empty string if not provided
api_base: string;
user: string;
metadata: any; // Json type in Prisma is represented by any in TypeScript; could also use a more specific type if the structure of JSON is known
cache_hit: string;
cache_key: string;
request_tags: any; // Similarly, this could be an array or a more specific type depending on the expected structure
team_id?: string | null; // ? indicates it's optional and can be undefined, but could also be null if not provided
end_user?: string | null;
};

View file

@ -1,84 +0,0 @@
import { serve } from '@hono/node-server'
import { Hono } from 'hono'
import { PrismaClient } from '@prisma/client'
import {LiteLLM_SpendLogs, LiteLLM_IncrementSpend, LiteLLM_IncrementObject} from './_types'
const app = new Hono()
const prisma = new PrismaClient()
// In-memory storage for logs
let spend_logs: LiteLLM_SpendLogs[] = [];
const key_logs: LiteLLM_IncrementObject[] = [];
const user_logs: LiteLLM_IncrementObject[] = [];
const transaction_logs: LiteLLM_IncrementObject[] = [];
app.get('/', (c) => {
return c.text('Hello Hono!')
})
const MIN_LOGS = 1; // Minimum number of logs needed to initiate a flush
const FLUSH_INTERVAL = 5000; // Time in ms to wait before trying to flush again
const BATCH_SIZE = 100; // Preferred size of each batch to write to the database
const MAX_LOGS_PER_INTERVAL = 1000; // Maximum number of logs to flush in a single interval
const flushLogsToDb = async () => {
if (spend_logs.length >= MIN_LOGS) {
// Limit the logs to process in this interval to MAX_LOGS_PER_INTERVAL or less
const logsToProcess = spend_logs.slice(0, MAX_LOGS_PER_INTERVAL);
for (let i = 0; i < logsToProcess.length; i += BATCH_SIZE) {
// Create subarray for current batch, ensuring it doesn't exceed the BATCH_SIZE
const batch = logsToProcess.slice(i, i + BATCH_SIZE);
// Convert datetime strings to Date objects
const batchWithDates = batch.map(entry => ({
...entry,
startTime: new Date(entry.startTime),
endTime: new Date(entry.endTime),
// Repeat for any other DateTime fields you may have
}));
await prisma.liteLLM_SpendLogs.createMany({
data: batchWithDates,
});
console.log(`Flushed ${batch.length} logs to the DB.`);
}
// Remove the processed logs from spend_logs
spend_logs = spend_logs.slice(logsToProcess.length);
console.log(`${logsToProcess.length} logs processed. Remaining in queue: ${spend_logs.length}`);
} else {
// This will ensure it doesn't falsely claim "No logs to flush." when it's merely below the MIN_LOGS threshold.
if(spend_logs.length > 0) {
console.log(`Accumulating logs. Currently at ${spend_logs.length}, waiting for at least ${MIN_LOGS}.`);
} else {
console.log("No logs to flush.");
}
}
};
// Setup interval for attempting to flush the logs
setInterval(flushLogsToDb, FLUSH_INTERVAL);
// Route to receive log messages
app.post('/spend/update', async (c) => {
const incomingLogs = await c.req.json<LiteLLM_SpendLogs[]>();
spend_logs.push(...incomingLogs);
console.log(`Received and stored ${incomingLogs.length} logs. Total logs in memory: ${spend_logs.length}`);
return c.json({ message: `Successfully stored ${incomingLogs.length} logs` });
});
const port = 3000
console.log(`Server is running on port ${port}`)
serve({
fetch: app.fetch,
port
})

View file

@ -1,13 +0,0 @@
{
"compilerOptions": {
"target": "ESNext",
"module": "ESNext",
"moduleResolution": "Bundler",
"strict": true,
"types": [
"node"
],
"jsx": "react-jsx",
"jsxImportSource": "hono/jsx",
}
}

View file

@ -206,6 +206,7 @@ add_user_information_to_llm_headers: Optional[bool] = (
)
store_audit_logs = False # Enterprise feature, allow users to see audit logs
skip_system_message_in_guardrail: bool = False
skip_tool_message_in_guardrail: bool = False
### end of callbacks #############
email: Optional[str] = (

View file

@ -57,6 +57,17 @@ LITELLM_PROXY_REQUEST_SPAN_NAME = "Received Proxy Server Request"
RAW_REQUEST_SPAN_NAME = "raw_gen_ai_request"
LITELLM_REQUEST_SPAN_NAME = "litellm_request"
CAPTURE_MODE_NO_CONTENT = "NO_CONTENT"
CAPTURE_MODE_SPAN_ONLY = "SPAN_ONLY"
CAPTURE_MODE_EVENT_ONLY = "EVENT_ONLY"
CAPTURE_MODE_SPAN_AND_EVENT = "SPAN_AND_EVENT"
_VALID_CAPTURE_MODES = {
CAPTURE_MODE_NO_CONTENT,
CAPTURE_MODE_SPAN_ONLY,
CAPTURE_MODE_EVENT_ONLY,
CAPTURE_MODE_SPAN_AND_EVENT,
}
@dataclass
class OpenTelemetryConfig:
@ -71,6 +82,9 @@ class OpenTelemetryConfig:
ignore_context_propagation: Optional[bool] = None
# When True, create a private TracerProvider instead of reusing or setting the global one.
skip_set_global: bool = False
# Programmatic override for OTEL_INSTRUMENTATION_GENAI_CAPTURE_MESSAGE_CONTENT.
# One of NO_CONTENT, SPAN_ONLY, EVENT_ONLY, SPAN_AND_EVENT (or "true" as legacy alias).
capture_message_content: Optional[str] = None
def __post_init__(self) -> None:
# If endpoint is specified but exporter is still the default "console",
@ -182,6 +196,9 @@ class OpenTelemetry(CustomLogger):
super().__init__(**kwargs)
self._init_metrics(meter_provider)
self._init_logs(logger_provider)
# Sample env-var / config / message_logging at init so subsequent
# _capture_in_span / _capture_in_event calls are deterministic.
self._capture_mode_cached = self._compute_capture_mode_from_init_state()
self._init_otel_logger_on_litellm_proxy()
@staticmethod
@ -306,6 +323,62 @@ class OpenTelemetry(CustomLogger):
hasattr(self, "callback_name") and self.callback_name == "langfuse_otel"
)
def _compute_capture_mode_from_init_state(self) -> Optional[str]:
"""Sample explicit settings at init. Returns the resolved mode or
None if nothing explicit is set (in which case the legacy
``self.message_logging`` flag is consulted dynamically per request).
``"true"``/``"1"`` map to ``EVENT_ONLY`` per the contrib convention.
``"false"``/``"0"`` map to ``NO_CONTENT``.
Unknown values are ignored.
"""
explicit = self.config.capture_message_content or os.getenv(
"OTEL_INSTRUMENTATION_GENAI_CAPTURE_MESSAGE_CONTENT"
)
if not explicit:
return None
normalized = explicit.upper()
if normalized in ("TRUE", "1"):
return CAPTURE_MODE_EVENT_ONLY
if normalized in ("FALSE", "0"):
return CAPTURE_MODE_NO_CONTENT
if normalized in _VALID_CAPTURE_MODES:
return normalized
return None
def _resolve_capture_mode(self) -> str:
"""Return the active capture mode for this request.
Precedence:
1. ``litellm.turn_off_message_logging=True`` forces ``NO_CONTENT``
(kill-switch checked dynamically).
2. Explicit setting sampled at init from
``OpenTelemetryConfig.capture_message_content`` or
``OTEL_INSTRUMENTATION_GENAI_CAPTURE_MESSAGE_CONTENT``.
3. Legacy ``self.message_logging`` (checked dynamically).
"""
if litellm.turn_off_message_logging:
return CAPTURE_MODE_NO_CONTENT
if self._capture_mode_cached is not None:
return self._capture_mode_cached
return (
CAPTURE_MODE_SPAN_AND_EVENT
if self.message_logging
else CAPTURE_MODE_NO_CONTENT
)
def _capture_in_span(self) -> bool:
return self._resolve_capture_mode() in (
CAPTURE_MODE_SPAN_ONLY,
CAPTURE_MODE_SPAN_AND_EVENT,
)
def _capture_in_event(self) -> bool:
return self._resolve_capture_mode() in (
CAPTURE_MODE_EVENT_ONLY,
CAPTURE_MODE_SPAN_AND_EVENT,
)
def _init_tracing(self, tracer_provider):
from opentelemetry import trace
from opentelemetry.sdk.trace import TracerProvider
@ -825,8 +898,7 @@ class OpenTelemetry(CustomLogger):
from opentelemetry import trace
from opentelemetry.trace import Status, StatusCode
# only log raw LLM request/response if message_logging is on and not globally turned off
if litellm.turn_off_message_logging or not self.message_logging:
if not self._capture_in_span():
return
litellm_params = kwargs.get("litellm_params", {})
@ -1117,9 +1189,14 @@ class OpenTelemetry(CustomLogger):
}
if role == "tool" and msg.get("id"):
attrs["id"] = msg["id"]
if self.message_logging and msg.get("content"):
capture_event_content = self._capture_in_event()
if capture_event_content and msg.get("content"):
attrs["gen_ai.prompt"] = msg["content"]
body = msg.copy()
if not capture_event_content:
body.pop("content", None)
log_record = SdkLogRecord(
timestamp=self._to_ns(datetime.now()),
trace_id=parent_ctx.trace_id,
@ -1127,7 +1204,7 @@ class OpenTelemetry(CustomLogger):
trace_flags=parent_ctx.trace_flags,
severity_number=SeverityNumber.INFO,
severity_text="INFO",
body=msg.copy(),
body=body,
attributes=attrs,
)
otel_logger.emit(log_record)
@ -1141,14 +1218,15 @@ class OpenTelemetry(CustomLogger):
"finish_reason": choice.get("finish_reason"),
}
body_msg = choice.get("message", {})
if self.message_logging and body_msg.get("content"):
capture_event_content = self._capture_in_event()
if capture_event_content and body_msg.get("content"):
attrs["message.content"] = body_msg["content"]
body = {
"index": idx,
"finish_reason": choice.get("finish_reason"),
"message": {"role": body_msg.get("role", "assistant")},
}
if self.message_logging and body_msg.get("content"):
if capture_event_content and body_msg.get("content"):
body["message"]["content"] = body_msg["content"]
log_record = SdkLogRecord(
@ -1674,9 +1752,7 @@ class OpenTelemetry(CustomLogger):
########## LLM Request Medssages / tools / content Attributes ###########
#########################################################################
if litellm.turn_off_message_logging is True:
return
if self.message_logging is not True:
if not self._capture_in_span():
return
if optional_params.get("tools"):

View file

@ -23,7 +23,9 @@ from litellm.llms.anthropic.experimental_pass_through.adapters.transformation im
from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation
from litellm.llms.base_llm.guardrail_translation.utils import (
effective_skip_system_message_for_guardrail,
effective_skip_tool_message_for_guardrail,
openai_messages_without_system,
openai_messages_without_tool,
)
from litellm.proxy.pass_through_endpoints.llm_provider_handlers.anthropic_passthrough_logging_handler import (
AnthropicPassthroughLoggingHandler,
@ -108,6 +110,7 @@ class AnthropicMessagesHandler(BaseTranslation):
return data
skip_system = effective_skip_system_message_for_guardrail(guardrail_to_apply)
skip_tool = effective_skip_tool_message_for_guardrail(guardrail_to_apply)
chat_completion_compatible_request = self._translate_to_openai(data)
@ -117,6 +120,8 @@ class AnthropicMessagesHandler(BaseTranslation):
)
if skip_system:
structured_messages = openai_messages_without_system(structured_messages)
if skip_tool:
structured_messages = openai_messages_without_tool(structured_messages)
texts_to_check: List[str] = []
images_to_check: List[str] = []
@ -134,6 +139,7 @@ class AnthropicMessagesHandler(BaseTranslation):
images_to_check=images_to_check,
task_mappings=task_mappings,
skip_system_message=skip_system,
skip_tool_message=skip_tool,
)
# Step 2: Apply guardrail to all texts in batch
@ -198,13 +204,17 @@ class AnthropicMessagesHandler(BaseTranslation):
images_to_check: List[str],
task_mappings: List[Tuple[int, Optional[int]]],
skip_system_message: bool = False,
skip_tool_message: bool = False,
) -> None:
"""
Extract text content and images from a message.
Override this method to customize text/image extraction logic.
"""
if skip_system_message and str(message.get("role") or "").lower() == "system":
role = str(message.get("role") or "").lower()
if skip_system_message and role == "system":
return
if skip_tool_message and role == "tool":
return
content = message.get("content", None)

View file

@ -14,7 +14,22 @@ def effective_skip_system_message_for_guardrail(guardrail_to_apply: Any) -> bool
return bool(getattr(litellm, "skip_system_message_in_guardrail", False))
def effective_skip_tool_message_for_guardrail(guardrail_to_apply: Any) -> bool:
per = getattr(guardrail_to_apply, "skip_tool_message_in_guardrail", None)
if per is not None:
return bool(per)
import litellm
return bool(getattr(litellm, "skip_tool_message_in_guardrail", False))
def openai_messages_without_system(
messages: List[AllMessageValues],
) -> List[AllMessageValues]:
return [m for m in messages if str((m or {}).get("role") or "").lower() != "system"]
def openai_messages_without_tool(
messages: List[AllMessageValues],
) -> List[AllMessageValues]:
return [m for m in messages if str((m or {}).get("role") or "").lower() != "tool"]

View file

@ -21,7 +21,9 @@ from litellm._logging import verbose_proxy_logger
from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation
from litellm.llms.base_llm.guardrail_translation.utils import (
effective_skip_system_message_for_guardrail,
effective_skip_tool_message_for_guardrail,
openai_messages_without_system,
openai_messages_without_tool,
)
from litellm.main import stream_chunk_builder
from litellm.types.llms.openai import AllMessageValues, ChatCompletionToolParam
@ -73,6 +75,7 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
return data
skip_system = effective_skip_system_message_for_guardrail(guardrail_to_apply)
skip_tool = effective_skip_tool_message_for_guardrail(guardrail_to_apply)
texts_to_check: List[str] = []
images_to_check: List[str] = []
@ -91,6 +94,7 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
text_task_mappings=text_task_mappings,
tool_call_task_mappings=tool_call_task_mappings,
skip_system_message=skip_system,
skip_tool_message=skip_tool,
)
# Step 2: Apply guardrail to all texts and tool calls in batch
@ -102,11 +106,15 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
inputs["tool_calls"] = tool_calls_to_check # type: ignore
structured_messages = self.get_structured_messages(data)
if structured_messages:
inputs["structured_messages"] = (
openai_messages_without_system(structured_messages)
if skip_system
else structured_messages
)
if skip_system:
structured_messages = openai_messages_without_system(
structured_messages
)
if skip_tool:
structured_messages = openai_messages_without_tool(
structured_messages
)
inputs["structured_messages"] = structured_messages
# Pass tools (function definitions) to the guardrail
tools = data.get("tools")
if tools:
@ -176,13 +184,17 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
text_task_mappings: List[Tuple[int, Optional[int]]],
tool_call_task_mappings: List[Tuple[int, int]],
skip_system_message: bool = False,
skip_tool_message: bool = False,
) -> None:
"""
Extract text content, images, and tool calls from a message.
Override this method to customize text/image/tool call extraction logic.
"""
if skip_system_message and str(message.get("role") or "").lower() == "system":
role = str(message.get("role") or "").lower()
if skip_system_message and role == "system":
return
if skip_tool_message and role == "tool":
return
content = message.get("content", None)

View file

@ -27187,6 +27187,20 @@
"supports_reasoning": true,
"supports_tool_choice": true
},
"openrouter/qwen/qwen3.6-plus": {
"input_cost_per_token": 3.25e-07,
"litellm_provider": "openrouter",
"max_input_tokens": 1000000,
"max_output_tokens": 65536,
"max_tokens": 65536,
"mode": "chat",
"output_cost_per_token": 1.95e-06,
"source": "https://openrouter.ai/qwen/qwen3.6-plus",
"supports_function_calling": true,
"supports_reasoning": true,
"supports_tool_choice": true,
"supports_vision": true
},
"openrouter/qwen/qwen3.5-35b-a3b": {
"input_cost_per_token": 2.5e-07,
"litellm_provider": "openrouter",

View file

@ -506,7 +506,8 @@ class MCPServerManager:
# Add any static headers from server config.
#
# Note: `extra_headers` on MCPServer is a List[str] of header names to forward
# from the client request (not available in this OpenAPI tool generation step).
# from each client MCP request; values are applied at call time via
# `_request_extra_headers` in server.py (not baked in here).
# `static_headers` is a dict of concrete headers to always send.
headers = (
merge_mcp_headers(

View file

@ -55,6 +55,13 @@ _request_auth_header: contextvars.ContextVar[Optional[str]] = contextvars.Contex
"_request_auth_header", default=None
)
# Per-request extra headers forwarded from the client request.
# Populated from MCPServer.extra_headers names matched against raw request
# headers in server.py before dispatching to a local/OpenAPI tool handler.
_request_extra_headers: contextvars.ContextVar[Optional[Dict[str, str]]] = (
contextvars.ContextVar("_request_extra_headers", default=None)
)
def _sanitize_path_parameter_value(param_value: Any, param_name: str) -> str:
"""Ensure path params cannot introduce directory traversal."""
@ -297,6 +304,46 @@ def build_input_schema(operation: Dict[str, Any]) -> Dict[str, Any]:
}
def _merge_openapi_tool_request_headers(
static_headers: Dict[str, str]
) -> Dict[str, str]:
"""Merge static closure headers with per-request ContextVar overrides.
Precedence (highest to lowest):
1. ``_request_auth_header`` — BYOK override of ``Authorization``
2. ``static_headers`` — operator-configured headers baked into the
tool closure at registration time
3. ``_request_extra_headers`` — per-request headers forwarded from
the MCP caller (allowlisted by ``MCPServer.extra_headers``)
This matches the existing MCP invariant in
:func:`litellm.proxy._experimental.mcp_server.utils.merge_mcp_headers`
and the managed MCP path, where ``static_headers`` always wins over
caller-forwarded headers. Keeping the same precedence here prevents an
authenticated caller from overriding an operator-configured value
(e.g. a tenant id or upstream API key) by sending the same header name.
Header names are compared case-insensitively so different casing cannot
bypass the precedence rules.
"""
request_extra = _request_extra_headers.get() or {}
static = static_headers or {}
static_lower_names = {k.lower() for k in static}
effective_headers: Dict[str, str] = {
k: v for k, v in request_extra.items() if k.lower() not in static_lower_names
}
effective_headers.update(static)
override_auth = _request_auth_header.get()
if override_auth:
for existing in [k for k in effective_headers if k.lower() == "authorization"]:
del effective_headers[existing]
effective_headers["Authorization"] = override_auth
return effective_headers
def create_tool_function(
path: str,
method: str,
@ -334,14 +381,7 @@ def create_tool_function(
The function safely handles parameter names that aren't valid Python identifiers
by using **kwargs instead of named parameters.
"""
# Allow per-request auth override (e.g. BYOK credential set via ContextVar).
# The ContextVar holds the full Authorization header value, including the
# correct prefix (Bearer / ApiKey / Basic) formatted by the caller in
# server.py based on the server's configured auth_type.
effective_headers = dict(headers)
override_auth = _request_auth_header.get()
if override_auth:
effective_headers["Authorization"] = override_auth
effective_headers = _merge_openapi_tool_request_headers(headers)
# Build URL from base_url and path
url = base_url + path

View file

@ -158,6 +158,7 @@ if MCP_AVAILABLE:
)
from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import (
_request_auth_header,
_request_extra_headers,
)
from litellm.proxy._experimental.mcp_server.sse_transport import SseServerTransport
from litellm.proxy._experimental.mcp_server.tool_registry import (
@ -2195,11 +2196,40 @@ if MCP_AVAILABLE:
auth_header_value = f"Basic {mcp_auth_header}"
else:
auth_header_value = f"Bearer {mcp_auth_header}"
# Forward named client headers to OpenAPI tool upstream requests.
# MCPServer.extra_headers lists header names to copy from raw_headers.
# OAuth2 M2M: never take Authorization from the caller (matches
# _prepare_mcp_server_headers for managed MCP).
forwarded_headers: Optional[Dict[str, str]] = None
if mcp_server and mcp_server.extra_headers and raw_headers:
normalized_raw = {
str(k).lower(): v
for k, v in raw_headers.items()
if isinstance(k, str)
}
skip_caller_authorization = bool(mcp_server.has_client_credentials)
for header_name in mcp_server.extra_headers:
if not isinstance(header_name, str):
continue
if (
skip_caller_authorization
and header_name.lower() == "authorization"
):
continue
value = normalized_raw.get(header_name.lower())
if value is not None:
if forwarded_headers is None:
forwarded_headers = {}
forwarded_headers[header_name] = value
_auth_token = _request_auth_header.set(auth_header_value)
_extra_token = _request_extra_headers.set(forwarded_headers)
try:
local_content = await _handle_local_mcp_tool(name, arguments)
finally:
_request_auth_header.reset(_auth_token)
_request_extra_headers.reset(_extra_token)
response = CallToolResult(content=cast(Any, local_content), isError=False)
# Try managed MCP server tool (pass the full prefixed name)

View file

@ -353,8 +353,10 @@ class LiteLLMRoutes(enum.Enum):
# realtime
"/realtime",
"/v1/realtime",
"/openai/v1/realtime",
"/realtime?{model}",
"/v1/realtime?{model}",
"/openai/v1/realtime?{model}",
# responses API
"/responses",
"/v1/responses",
@ -707,6 +709,8 @@ class LiteLLMRoutes(enum.Enum):
# Project read routes - endpoint scopes results to caller's teams (non-admin)
"/project/list",
"/project/info",
# Endpoint enforces proxy-admin vs team-admin model access itself.
"/health/test_connection",
# Invitation routes - org/team admins checked in endpoint via _user_has_admin_privileges
"/invitation/new",
"/invitation/delete",

View file

@ -2849,7 +2849,7 @@ def _can_object_call_model(
object_type=object_type
),
param="model",
code=status.HTTP_401_UNAUTHORIZED,
code=status.HTTP_403_FORBIDDEN,
)
@ -3082,7 +3082,7 @@ async def can_user_call_model(
message=f"User not allowed to access model. No default model access, only team models allowed. Tried to access {model}",
type=ProxyErrorTypes.key_model_access_denied,
param="model",
code=status.HTTP_401_UNAUTHORIZED,
code=status.HTTP_403_FORBIDDEN,
)
return _can_object_call_model(
@ -3625,7 +3625,7 @@ async def _check_team_member_model_access(
message=f"Team member not allowed to access model. User={valid_token.user_id}, Team={team_object.team_id}, Model={model}. Allowed member models = {member_allowed_models}",
type=ProxyErrorTypes.team_model_access_denied,
param="model",
code=status.HTTP_401_UNAUTHORIZED,
code=status.HTTP_403_FORBIDDEN,
)

View file

@ -123,7 +123,7 @@ class UserAPIKeyAuthExceptionHandler:
message=e.message,
type=ProxyErrorTypes.budget_exceeded,
param=None,
code=400,
code=getattr(e, "status_code", status.HTTP_429_TOO_MANY_REQUESTS),
)
if isinstance(e, HTTPException):
raise ProxyException(

View file

@ -1107,7 +1107,7 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
raise ProxyException(
message=f"Authentication Error - Expired Key. Key Expiry time {expiry_time} and current time {current_time}",
type=ProxyErrorTypes.expired_key,
code=400,
code=status.HTTP_401_UNAUTHORIZED,
param=abbreviate_api_key(api_key=api_key),
)
valid_token = update_valid_token_with_end_user_params(
@ -1432,7 +1432,7 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
raise ProxyException(
message=f"Authentication Error - Expired Key. Key Expiry time {expiry_time} and current time {current_time}",
type=ProxyErrorTypes.expired_key,
code=400,
code=status.HTTP_401_UNAUTHORIZED,
param=abbreviate_api_key(api_key=api_key),
)
@ -2417,7 +2417,7 @@ async def _run_post_custom_auth_checks(
raise ProxyException(
message=f"Authentication Error - Expired Key. Key Expiry time {expiry_time} and current time {current_time}",
type=ProxyErrorTypes.expired_key,
code=400,
code=status.HTTP_401_UNAUTHORIZED,
param=(
abbreviate_api_key(api_key=valid_token.token)
if valid_token.token

View file

@ -14,7 +14,7 @@ jwt_display_template = """
padding: 20px;
display: flex;
justify-content: center;
align-items: center;
align-items: flex-start;
min-height: 100vh;
color: #333;
}
@ -27,18 +27,18 @@ jwt_display_template = """
width: 800px;
max-width: 100%;
}
.logo-container {
text-align: center;
margin-bottom: 30px;
}
.logo {
font-size: 24px;
font-weight: 600;
color: #1e293b;
}
h2 {
margin: 0 0 10px;
color: #1e293b;
@ -46,7 +46,14 @@ jwt_display_template = """
font-weight: 600;
text-align: center;
}
h3 {
margin: 0 0 12px;
color: #1e293b;
font-size: 18px;
font-weight: 600;
}
.subtitle {
color: #64748b;
margin: 0 0 20px;
@ -58,15 +65,15 @@ jwt_display_template = """
background-color: #f1f5f9;
border-radius: 6px;
padding: 20px;
margin-bottom: 30px;
margin-bottom: 20px;
border-left: 4px solid #2563eb;
}
.success-box {
background-color: #f0fdf4;
border-radius: 6px;
padding: 20px;
margin-bottom: 30px;
margin-bottom: 20px;
border-left: 4px solid #16a34a;
}
@ -78,7 +85,7 @@ jwt_display_template = """
font-weight: 600;
font-size: 16px;
}
.success-header {
display: flex;
align-items: center;
@ -87,46 +94,53 @@ jwt_display_template = """
font-weight: 600;
font-size: 16px;
}
.info-header svg, .success-header svg {
margin-right: 8px;
}
.data-container {
margin-top: 20px;
}
.data-row {
display: flex;
border-bottom: 1px solid #e2e8f0;
padding: 12px 0;
}
.data-row:last-child {
border-bottom: none;
}
.data-label {
font-weight: 500;
color: #334155;
width: 180px;
width: 220px;
flex-shrink: 0;
}
.data-value {
color: #475569;
word-break: break-all;
}
.empty-note {
color: #64748b;
font-style: italic;
margin: 0;
font-size: 14px;
}
.jwt-container {
background-color: #f8fafc;
border-radius: 6px;
padding: 15px;
margin-top: 20px;
margin-top: 12px;
overflow-x: auto;
border: 1px solid #e2e8f0;
}
.jwt-text {
font-family: monospace;
white-space: pre-wrap;
@ -134,7 +148,7 @@ jwt_display_template = """
margin: 0;
color: #334155;
}
.back-button {
display: inline-block;
background-color: #6466E9;
@ -146,18 +160,18 @@ jwt_display_template = """
margin-top: 20px;
text-align: center;
}
.back-button:hover {
background-color: #4138C2;
text-decoration: none;
}
.buttons {
display: flex;
gap: 10px;
margin-top: 20px;
margin-top: 12px;
}
.copy-button {
background-color: #e2e8f0;
color: #334155;
@ -169,11 +183,11 @@ jwt_display_template = """
display: flex;
align-items: center;
}
.copy-button:hover {
background-color: #cbd5e1;
}
.copy-button svg {
margin-right: 6px;
}
@ -188,7 +202,7 @@ jwt_display_template = """
</div>
<h2>SSO Debug Information</h2>
<p class="subtitle">Results from the SSO authentication process.</p>
<div class="success-box">
<div class="success-header">
<svg xmlns="http://www.w3.org/2000/svg" width="16" height="16" viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2" stroke-linecap="round" stroke-linejoin="round">
@ -199,11 +213,7 @@ jwt_display_template = """
</div>
<p>The SSO authentication completed successfully. Below is the information returned by the provider.</p>
</div>
<div class="data-container" id="userData">
<!-- Data will be inserted here by JavaScript -->
</div>
<div class="info-box">
<div class="info-header">
<svg xmlns="http://www.w3.org/2000/svg" width="16" height="16" viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2" stroke-linecap="round" stroke-linejoin="round">
@ -211,22 +221,62 @@ jwt_display_template = """
<line x1="12" y1="16" x2="12" y2="12"></line>
<line x1="12" y1="8" x2="12.01" y2="8"></line>
</svg>
JSON Representation
Parsed by Proxy
</div>
<p class="empty-note">Fields the proxy extracted into its internal user model.</p>
<div class="data-container" id="parsedByProxy">
<!-- Populated by JavaScript -->
</div>
</div>
<div class="info-box">
<div class="info-header">
<svg xmlns="http://www.w3.org/2000/svg" width="16" height="16" viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2" stroke-linecap="round" stroke-linejoin="round">
<circle cx="12" cy="12" r="10"></circle>
<line x1="12" y1="16" x2="12" y2="12"></line>
<line x1="12" y1="8" x2="12.01" y2="8"></line>
</svg>
Raw Claims (userinfo)
</div>
<p class="empty-note">Complete set of claims returned by the IdP's userinfo endpoint.</p>
<div class="jwt-container">
<pre class="jwt-text" id="jsonData">Loading...</pre>
<pre class="jwt-text" id="rawClaims">Loading...</pre>
</div>
<div class="buttons">
<button class="copy-button" onclick="copyToClipboard('jsonData')">
<button class="copy-button" onclick="copyToClipboard('rawClaims')">
<svg xmlns="http://www.w3.org/2000/svg" width="14" height="14" viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2" stroke-linecap="round" stroke-linejoin="round">
<rect x="9" y="9" width="13" height="13" rx="2" ry="2"></rect>
<path d="M5 15H4a2 2 0 0 1-2-2V4a2 2 0 0 1 2-2h9a2 2 0 0 1 2 2v1"></path>
</svg>
Copy to Clipboard
Copy
</button>
</div>
</div>
<div class="info-box">
<div class="info-header">
<svg xmlns="http://www.w3.org/2000/svg" width="16" height="16" viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2" stroke-linecap="round" stroke-linejoin="round">
<circle cx="12" cy="12" r="10"></circle>
<line x1="12" y1="16" x2="12" y2="12"></line>
<line x1="12" y1="8" x2="12.01" y2="8"></line>
</svg>
Access Token Claims
</div>
<p class="empty-note">Decoded payload of the access token JWT (when the IdP issues one).</p>
<div class="jwt-container">
<pre class="jwt-text" id="accessTokenClaims">Loading...</pre>
</div>
<div class="buttons">
<button class="copy-button" onclick="copyToClipboard('accessTokenClaims')">
<svg xmlns="http://www.w3.org/2000/svg" width="14" height="14" viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2" stroke-linecap="round" stroke-linejoin="round">
<rect x="9" y="9" width="13" height="13" rx="2" ry="2"></rect>
<path d="M5 15H4a2 2 0 0 1-2-2V4a2 2 0 0 1 2-2h9a2 2 0 0 1 2 2v1"></path>
</svg>
Copy
</button>
</div>
</div>
<a href="/sso/debug/login" class="back-button">
Try Another SSO Login
</a>
@ -234,39 +284,58 @@ jwt_display_template = """
<script>
// This will be populated with the actual data from the server
const userData = SSO_DATA;
function renderUserData() {
const container = document.getElementById('userData');
const jsonDisplay = document.getElementById('jsonData');
// Format JSON with indentation for display
jsonDisplay.textContent = JSON.stringify(userData, null, 2);
// Clear container
const ssoData = SSO_DATA;
function renderParsed(container, parsed) {
container.innerHTML = '';
// Add each key-value pair to the UI
for (const [key, value] of Object.entries(userData)) {
if (typeof value !== 'object' || value === null) {
const row = document.createElement('div');
row.className = 'data-row';
const label = document.createElement('div');
label.className = 'data-label';
label.textContent = key;
const dataValue = document.createElement('div');
dataValue.className = 'data-value';
dataValue.textContent = value !== null ? value : 'null';
row.appendChild(label);
row.appendChild(dataValue);
container.appendChild(row);
const entries = Object.entries(parsed || {});
if (entries.length === 0) {
const note = document.createElement('p');
note.className = 'empty-note';
note.textContent = 'No fields available.';
container.appendChild(note);
return;
}
for (const [key, value] of entries) {
const row = document.createElement('div');
row.className = 'data-row';
const label = document.createElement('div');
label.className = 'data-label';
label.textContent = key;
const dataValue = document.createElement('div');
dataValue.className = 'data-value';
if (value === null || value === undefined) {
dataValue.textContent = 'null';
} else if (typeof value === 'object') {
dataValue.textContent = JSON.stringify(value);
} else {
dataValue.textContent = String(value);
}
row.appendChild(label);
row.appendChild(dataValue);
container.appendChild(row);
}
}
function renderJson(elementId, value) {
const el = document.getElementById(elementId);
const obj = value || {};
if (Object.keys(obj).length === 0) {
el.textContent = '(empty — provider returned no claims for this section)';
} else {
el.textContent = JSON.stringify(obj, null, 2);
}
}
function renderUserData() {
renderParsed(document.getElementById('parsedByProxy'), ssoData.parsed_by_proxy);
renderJson('rawClaims', ssoData.raw_claims);
renderJson('accessTokenClaims', ssoData.access_token_claims);
}
function copyToClipboard(elementId) {
const text = document.getElementById(elementId).textContent;
navigator.clipboard.writeText(text).then(() => {
@ -275,7 +344,7 @@ jwt_display_template = """
console.error('Could not copy text: ', err);
});
}
// Render the data when the page loads
document.addEventListener('DOMContentLoaded', renderUserData);
</script>

View file

@ -10,13 +10,64 @@ import subprocess
import time
import urllib
import urllib.parse
from dataclasses import dataclass
from datetime import datetime, timedelta
from typing import Any, Optional, Union
from typing import Any, Dict, Optional, Union
from litellm._logging import verbose_proxy_logger
from litellm.secret_managers.main import str_to_bool
@dataclass(frozen=True)
class IAMEndpoint:
"""Static parts of an RDS IAM-authenticated Postgres connection.
The IAM token rotates every ~15 minutes; everything else (host, port, user,
database name, schema) stays fixed. We capture the static fields once so
refresh just regenerates the token and reassembles the URL.
"""
host: str
port: str
user: str
name: str
schema: Optional[str] = None
def build_url(self, token: str) -> str:
url = f"postgresql://{self.user}:{token}@{self.host}:{self.port}/{self.name}"
if self.schema:
url += f"?schema={self.schema}"
return url
def parse_iam_endpoint_from_url(url: str) -> IAMEndpoint:
"""Parse an IAMEndpoint from a Postgres URL.
Used so a reader URL can drive its own IAM refresh without requiring
callers to set parallel DATABASE_HOST_READ_REPLICA / etc. env vars.
"""
parsed = urllib.parse.urlparse(url)
if not parsed.hostname or not parsed.username:
raise ValueError("Cannot parse IAM endpoint from URL: missing host or username")
name = (parsed.path or "/").lstrip("/")
if not name:
raise ValueError("Cannot parse IAM endpoint from URL: missing database name")
port = str(parsed.port) if parsed.port else "5432"
schema: Optional[str] = None
if parsed.query:
qs = urllib.parse.parse_qs(parsed.query)
schema_vals = qs.get("schema")
if schema_vals:
schema = schema_vals[0]
return IAMEndpoint(
host=parsed.hostname,
port=port,
user=parsed.username,
name=name,
schema=schema,
)
class PrismaWrapper:
"""
Wrapper around Prisma client that handles RDS IAM token authentication.
@ -37,10 +88,33 @@ class PrismaWrapper:
# Fallback refresh interval if token parsing fails (10 minutes)
FALLBACK_REFRESH_INTERVAL_SECONDS = 600
def __init__(self, original_prisma: Any, iam_token_db_auth: bool):
def __init__(
self,
original_prisma: Any,
iam_token_db_auth: bool,
*,
db_url_env_var: str = "DATABASE_URL",
iam_endpoint: Optional[IAMEndpoint] = None,
recreate_uses_datasource: bool = False,
log_prefix: str = "",
):
self._original_prisma = original_prisma
self.iam_token_db_auth = iam_token_db_auth
# Per-connection knobs so the same wrapper can be used for the writer
# (defaults: DATABASE_URL env, IAM endpoint from DATABASE_HOST/etc.,
# recreate via env reload) or for a reader (DATABASE_URL_READ_REPLICA
# env, IAM endpoint parsed from that URL, recreate via datasource
# override since Prisma only auto-reads DATABASE_URL).
self._db_url_env_var = db_url_env_var
self._iam_endpoint = iam_endpoint
self._recreate_uses_datasource = recreate_uses_datasource
# Tag every log line emitted by this wrapper instance so writer and
# reader can be told apart in interleaved output (e.g. "[writer] RDS
# IAM token refresh scheduled in 720 seconds"). Empty string (default)
# keeps backward-compatible logs for the single-DB case.
self._log_prefix = f"{log_prefix} " if log_prefix else ""
# Background token refresh task management
self._token_refresh_task: Optional[asyncio.Task] = None
self._reconnection_lock = asyncio.Lock()
@ -157,7 +231,7 @@ class PrismaWrapper:
Returns 0 if token should be refreshed immediately.
Returns FALLBACK_REFRESH_INTERVAL_SECONDS if parsing fails.
"""
db_url = os.getenv("DATABASE_URL")
db_url = os.getenv(self._db_url_env_var)
token = self._extract_token_from_db_url(db_url)
expiration_time = self._parse_token_expiration(token)
@ -199,12 +273,30 @@ class PrismaWrapper:
return datetime.utcnow() > expiration_time
def get_rds_iam_token(self) -> Optional[str]:
"""Generate a new RDS IAM token and update DATABASE_URL."""
if self.iam_token_db_auth:
from litellm.proxy.auth.rds_iam_token import generate_iam_auth_token
"""Generate a new RDS IAM token and update the configured DB URL env var.
When the wrapper was constructed with an explicit `iam_endpoint`
(typical for a reader wrapper whose host/port/user came from a parsed
URL), use that. Otherwise fall back to the legacy DATABASE_HOST/PORT/
USER/NAME/SCHEMA env vars (writer behavior).
"""
if not self.iam_token_db_auth:
return None
from litellm.proxy.auth.rds_iam_token import generate_iam_auth_token
if self._iam_endpoint is not None:
endpoint = self._iam_endpoint
token = generate_iam_auth_token(
db_host=endpoint.host, db_port=endpoint.port, db_user=endpoint.user
)
_db_url = endpoint.build_url(token)
else:
db_host = os.getenv("DATABASE_HOST")
db_port = os.getenv("DATABASE_PORT")
# Default to the Postgres standard port; passing None to
# `generate_iam_auth_token` makes botocore embed the literal
# string "None" in the presigned URL, which then fails to parse.
db_port = os.getenv("DATABASE_PORT", "5432")
db_user = os.getenv("DATABASE_USER")
db_name = os.getenv("DATABASE_NAME")
db_schema = os.getenv("DATABASE_SCHEMA")
@ -217,9 +309,8 @@ class PrismaWrapper:
if db_schema:
_db_url += f"?schema={db_schema}"
os.environ["DATABASE_URL"] = _db_url
return _db_url
return None
os.environ[self._db_url_env_var] = _db_url
return _db_url
async def recreate_prisma_client(
self, new_db_url: str, http_client: Optional[Any] = None
@ -231,6 +322,11 @@ class PrismaWrapper:
synchronous `subprocess.Popen.wait()` that can freeze the asyncio event
loop for 30-120+ seconds when the engine is stuck on TCP close,
breaking `/health/liveliness` and causing Kubernetes pod restarts.
The writer wrapper relies on Prisma re-reading `DATABASE_URL` from env;
the reader wrapper opts into `recreate_uses_datasource=True` so the
new URL is passed explicitly via `datasource={"url": ...}` (Prisma
does not auto-read alternate env vars like DATABASE_URL_READ_REPLICA).
"""
from prisma import Prisma # type: ignore
@ -238,10 +334,12 @@ class PrismaWrapper:
if old_engine_pid > 0:
await self._kill_engine_process(old_engine_pid)
kwargs: Dict[str, Any] = {}
if http_client is not None:
self._original_prisma = Prisma(http=http_client)
else:
self._original_prisma = Prisma()
kwargs["http"] = http_client
if self._recreate_uses_datasource:
kwargs["datasource"] = {"url": new_db_url}
self._original_prisma = Prisma(**kwargs)
await self._original_prisma.connect()
@ -265,7 +363,8 @@ class PrismaWrapper:
self._token_refresh_task = asyncio.create_task(self._token_refresh_loop())
verbose_proxy_logger.info(
"Started RDS IAM token proactive refresh background task"
"%sStarted RDS IAM token proactive refresh background task",
self._log_prefix,
)
async def stop_token_refresh_task(self) -> None:
@ -283,7 +382,9 @@ class PrismaWrapper:
except asyncio.CancelledError:
pass
self._token_refresh_task = None
verbose_proxy_logger.info("Stopped RDS IAM token refresh background task")
verbose_proxy_logger.info(
"%sStopped RDS IAM token refresh background task", self._log_prefix
)
async def _token_refresh_loop(self) -> None:
"""
@ -294,7 +395,7 @@ class PrismaWrapper:
This is more efficient than polling, requiring only 1 wake-up per token cycle.
"""
verbose_proxy_logger.info(
f"RDS IAM token refresh loop started. "
f"{self._log_prefix}RDS IAM token refresh loop started. "
f"Tokens will be refreshed {self.TOKEN_REFRESH_BUFFER_SECONDS}s before expiration."
)
@ -305,21 +406,25 @@ class PrismaWrapper:
if sleep_seconds > 0:
verbose_proxy_logger.info(
f"RDS IAM token refresh scheduled in {sleep_seconds:.0f} seconds "
f"({sleep_seconds / 60:.1f} minutes)"
f"{self._log_prefix}RDS IAM token refresh scheduled in "
f"{sleep_seconds:.0f} seconds ({sleep_seconds / 60:.1f} minutes)"
)
await asyncio.sleep(sleep_seconds)
# Refresh the token
verbose_proxy_logger.info("Proactively refreshing RDS IAM token...")
verbose_proxy_logger.info(
"%sProactively refreshing RDS IAM token...", self._log_prefix
)
await self._safe_refresh_token()
except asyncio.CancelledError:
verbose_proxy_logger.info("RDS IAM token refresh loop cancelled")
verbose_proxy_logger.info(
"%sRDS IAM token refresh loop cancelled", self._log_prefix
)
break
except Exception as e:
verbose_proxy_logger.error(
f"Error in RDS IAM token refresh loop: {e}. "
f"{self._log_prefix}Error in RDS IAM token refresh loop: {e}. "
f"Retrying in {self.FALLBACK_REFRESH_INTERVAL_SECONDS}s..."
)
# On error, wait before retrying to avoid tight error loops
@ -341,65 +446,75 @@ class PrismaWrapper:
await self.recreate_prisma_client(new_db_url)
self._last_refresh_time = datetime.utcnow()
verbose_proxy_logger.info(
"RDS IAM token refreshed successfully. New token valid for ~15 minutes."
"%sRDS IAM token refreshed successfully. New token valid for ~15 minutes.",
self._log_prefix,
)
else:
verbose_proxy_logger.error(
"Failed to generate new RDS IAM token during proactive refresh"
"%sFailed to generate new RDS IAM token during proactive refresh",
self._log_prefix,
)
def __getattr__(self, name: str):
"""
Proxy attribute access to the underlying Prisma client.
If IAM token auth is enabled and the token is expired, this method
provides a synchronous fallback to refresh the token. However, this
should rarely be needed since the background task proactively refreshes
tokens before they expire.
If IAM token auth is enabled and the token is found expired here, the
proactive refresh task has missed its window. Behavior depends on
whether we're called from inside a running event loop:
FIXED: Now properly waits for reconnection to complete before returning,
instead of the previous fire-and-forget pattern that caused the bug.
- Inside the loop (typical: from a coroutine): schedule a refresh as a
background task and return the (stale) attribute. The caller's await
will likely fail with a connection error and be retried by upper
layers (`call_with_db_reconnect_retry`); by that time the refresh
has either completed or escalated to the proactive loop's error
path. We CANNOT block here — `run_coroutine_threadsafe(...)` +
`future.result()` from inside the same loop deadlocks the loop
(loop thread is blocked, scheduled coroutine never runs, 30s timeout).
- No running loop (sync caller, mostly tests): run the refresh in a
fresh loop and re-fetch the attribute.
"""
original_attr = getattr(self._original_prisma, name)
if self.iam_token_db_auth:
db_url = os.getenv("DATABASE_URL")
db_url = os.getenv(self._db_url_env_var)
# Check if token is expired (should be rare if background task is running)
if self.is_token_expired(db_url):
verbose_proxy_logger.warning(
"RDS IAM token expired in __getattr__ - proactive refresh may have failed. "
"Triggering synchronous fallback refresh..."
)
try:
running_loop = asyncio.get_running_loop()
except RuntimeError:
running_loop = None
new_db_url = self.get_rds_iam_token()
if new_db_url:
loop = asyncio.get_event_loop()
if loop.is_running():
# FIXED: Actually wait for the reconnection to complete!
# The previous code used fire-and-forget which caused the bug.
future = asyncio.run_coroutine_threadsafe(
self.recreate_prisma_client(new_db_url), loop
)
try:
# Wait up to 30 seconds for reconnection
future.result(timeout=30)
verbose_proxy_logger.info(
"Synchronous token refresh completed successfully"
)
except Exception as e:
verbose_proxy_logger.error(
f"Failed to refresh token synchronously: {e}"
)
raise
else:
asyncio.run(self.recreate_prisma_client(new_db_url))
# Get the NEW attribute after reconnection
original_attr = getattr(self._original_prisma, name)
if running_loop is not None:
verbose_proxy_logger.warning(
"%sRDS IAM token expired in __getattr__ — proactive refresh "
"may have failed. Scheduling async refresh; the current "
"request may fail and be retried with the fresh token.",
self._log_prefix,
)
# Non-blocking: schedule the locked refresh on the
# running loop. The reconnection lock inside
# `_safe_refresh_token` coalesces concurrent triggers.
running_loop.create_task(self._safe_refresh_token())
else:
raise ValueError("Failed to get RDS IAM token")
verbose_proxy_logger.warning(
"%sRDS IAM token expired in __getattr__ — proactive refresh "
"may have failed. Triggering synchronous fallback refresh...",
self._log_prefix,
)
new_db_url = self.get_rds_iam_token()
if new_db_url:
asyncio.run(self.recreate_prisma_client(new_db_url))
# Re-fetch attribute against the recreated Prisma instance.
original_attr = getattr(self._original_prisma, name)
verbose_proxy_logger.info(
"%sSynchronous token refresh completed successfully",
self._log_prefix,
)
else:
raise ValueError("Failed to get RDS IAM token")
return original_attr

View file

@ -0,0 +1,213 @@
"""
RoutingPrismaWrapper: routes Prisma reads to a read-replica client and writes
to a writer client. Used when DATABASE_URL_READ_REPLICA is configured;
otherwise PrismaClient uses the writer-only PrismaWrapper directly.
"""
import os
from typing import Any, Callable, Optional
from litellm._logging import verbose_proxy_logger
from litellm.proxy.db.prisma_client import PrismaWrapper
# Per-model action methods that read from the database. These are routed to
# the read replica when one is configured.
_MODEL_READ_METHODS = frozenset(
{
"find_first",
"find_first_or_raise",
"find_many",
"find_unique",
"find_unique_or_raise",
"count",
"group_by",
"query_first",
"query_raw",
}
)
# Top-level Prisma client methods that read from the database.
_TOP_LEVEL_READ_METHODS = frozenset({"query_first", "query_raw"})
class _RoutedActions:
"""Per-model accessor that sends reads to the reader and writes to the writer.
`should_use_reader` is consulted on every read dispatch so a mid-call flip
of the routing wrapper's reader-availability flag (e.g. after the reader
fails a recreate) is observed without re-fetching the actions accessor.
"""
__slots__ = ("_writer_actions", "_reader_actions", "_should_use_reader")
def __init__(
self,
writer_actions: Any,
reader_actions: Any,
should_use_reader: Callable[[], bool],
):
self._writer_actions = writer_actions
self._reader_actions = reader_actions
self._should_use_reader = should_use_reader
def __getattr__(self, name: str) -> Any:
if name in _MODEL_READ_METHODS and self._should_use_reader():
return getattr(self._reader_actions, name)
return getattr(self._writer_actions, name)
class RoutingPrismaWrapper:
"""
Routes Prisma operations between a writer and a reader Prisma client.
Reads (find_*, count, group_by, query_raw, query_first) go to the reader;
everything else (writes, transactions, raw execute) goes to the writer.
Lifecycle methods (connect, disconnect, IAM token refresh) act on both
clients so callers do not need to know about the split. When
IAM_TOKEN_DB_AUTH is enabled, both writer and reader refresh their tokens
independently on their own ~12-minute cadence.
Reader degradation: a reader-side failure (failed connect, failed
recreate) is non-fatal — the wrapper sets `_reader_unavailable=True`, logs
a warning, and routes subsequent reads to the writer. The next successful
`connect()` or `recreate_prisma_client()` clears the flag. This keeps the
proxy serving traffic during transient reader outages instead of failing
startup or returning errors for read-heavy endpoints.
"""
def __init__(self, writer: PrismaWrapper, reader: PrismaWrapper):
self._writer = writer
self._reader = reader
# When True, reads fall back to the writer. Flipped on by reader
# connect/recreate failures and flipped off on the next reader recovery.
self._reader_unavailable: bool = False
@property
def writer(self) -> PrismaWrapper:
return self._writer
@property
def reader(self) -> PrismaWrapper:
return self._reader
@property
def reader_unavailable(self) -> bool:
return self._reader_unavailable
def _should_use_reader(self) -> bool:
return not self._reader_unavailable
async def connect(self, *args: Any, **kwargs: Any) -> None:
await self._writer.connect(*args, **kwargs)
verbose_proxy_logger.info("[writer] DB connected")
try:
await self._reader.connect(*args, **kwargs)
self._reader_unavailable = False
verbose_proxy_logger.info("[reader] DB connected")
except Exception as e:
# Degrade gracefully: the proxy keeps serving traffic with reads
# routed to the writer until the reader endpoint is reachable.
# Aborting startup here would tie proxy availability to an
# opt-in, best-effort reader endpoint.
self._reader_unavailable = True
verbose_proxy_logger.warning(
"Failed to connect to read replica DB: %s. "
"Falling back to the writer for reads until the reader is reachable.",
e,
)
async def disconnect(self, *args: Any, **kwargs: Any) -> None:
first_error: Optional[BaseException] = None
for client in (self._writer, self._reader):
try:
await client.disconnect(*args, **kwargs)
except Exception as e:
if first_error is None:
first_error = e
verbose_proxy_logger.warning("Error disconnecting Prisma client: %s", e)
if first_error is not None:
raise first_error
def is_connected(self) -> bool:
# Reflects writer health only. The reader is best-effort; its
# availability is tracked via `_reader_unavailable` and a degraded
# reader must NOT cause a writer reconnect (would loop indefinitely
# since recreate_prisma_client only fixes writer-side problems).
return bool(self._writer.is_connected())
async def start_token_refresh_task(self) -> None:
await self._writer.start_token_refresh_task()
await self._reader.start_token_refresh_task()
async def stop_token_refresh_task(self) -> None:
await self._writer.stop_token_refresh_task()
await self._reader.stop_token_refresh_task()
async def recreate_prisma_client(
self, new_db_url: str, http_client: Optional[Any] = None
) -> None:
"""Recreate both writer and reader Prisma clients.
The writer reconnect path in PrismaClient calls
`self.db.recreate_prisma_client(...)`. Without this method, a DB-wide
connectivity event would only re-create the writer; the reader engine
would stay broken and every routed read would fail. We always recreate
the writer first (its URL is the one passed in), then best-effort
recreate the reader. A reader failure flips `_reader_unavailable=True`
so reads transparently fall through to the writer.
"""
await self._writer.recreate_prisma_client(new_db_url, http_client=http_client)
try:
await self._recreate_reader(http_client=http_client)
self._reader_unavailable = False
except Exception as e:
self._reader_unavailable = True
verbose_proxy_logger.warning(
"Failed to recreate reader Prisma client: %s. "
"Reads will fall back to the writer until the reader recovers.",
e,
)
async def _recreate_reader(self, http_client: Optional[Any] = None) -> None:
"""Resolve the reader URL and recreate its Prisma client.
IAM-enabled readers regenerate their token (host/port/user came from
the parsed reader URL at construction time). Non-IAM readers reuse
the URL stored in `DATABASE_URL_READ_REPLICA`.
"""
if self._reader.iam_token_db_auth:
new_reader_url = self._reader.get_rds_iam_token()
if not new_reader_url:
raise RuntimeError(
"Failed to generate fresh IAM token for read replica"
)
await self._reader.recreate_prisma_client(
new_reader_url, http_client=http_client
)
return
reader_url = os.getenv("DATABASE_URL_READ_REPLICA", "")
if not reader_url:
raise RuntimeError(
"DATABASE_URL_READ_REPLICA not set; cannot recreate read replica client"
)
await self._reader.recreate_prisma_client(reader_url, http_client=http_client)
def __getattr__(self, name: str) -> Any:
if name in _TOP_LEVEL_READ_METHODS:
target = self._writer if self._reader_unavailable else self._reader
return getattr(target, name)
writer_attr = getattr(self._writer, name)
# Per-model action accessors are non-callable instances that expose
# both `find_many` and `create`. Methods like execute_raw / batch_ /
# tx are callables and stay on the writer untouched.
if (
not callable(writer_attr)
and hasattr(writer_attr, "find_many")
and hasattr(writer_attr, "create")
):
try:
reader_attr = getattr(self._reader, name)
except AttributeError:
return writer_attr
return _RoutedActions(writer_attr, reader_attr, self._should_use_reader)
return writer_attr

View file

@ -482,6 +482,11 @@ class InMemoryGuardrailHandler:
"skip_system_message_in_guardrail",
getattr(litellm_params, "skip_system_message_in_guardrail", None),
)
setattr(
custom_guardrail_callback,
"skip_tool_message_in_guardrail",
getattr(litellm_params, "skip_tool_message_in_guardrail", None),
)
parsed_guardrail = Guardrail(
guardrail_id=guardrail.get("guardrail_id"),

View file

@ -319,6 +319,9 @@ class _PROXY_VirtualKeyModelMaxBudgetLimiter(RouterBudgetLimiting):
response_cost=response_cost,
)
if self.dual_cache.redis_cache is not None:
await self._push_in_memory_increments_to_redis()
verbose_proxy_logger.debug(
"current state of in memory cache %s",
json.dumps(

View file

@ -4078,6 +4078,8 @@ async def debug_sso_callback(request: Request):
redirect_url += "/sso/debug/callback"
result = None
received_response: Optional[dict] = None
access_token_payload: Optional[dict] = None
if google_client_id is not None:
result = await GoogleSSOHandler.get_google_callback_response(
request=request,
@ -4094,12 +4096,14 @@ async def debug_sso_callback(request: Request):
)
elif generic_client_id is not None:
result, _, _ = await get_generic_sso_response(
request=request,
jwt_handler=jwt_handler,
generic_client_id=generic_client_id,
redirect_url=redirect_url,
sso_jwt_handler=sso_jwt_handler,
result, received_response, access_token_payload = (
await get_generic_sso_response(
request=request,
jwt_handler=jwt_handler,
generic_client_id=generic_client_id,
redirect_url=redirect_url,
sso_jwt_handler=sso_jwt_handler,
)
)
# If result is None, return a basic error message
@ -4128,10 +4132,32 @@ async def debug_sso_callback(request: Request):
except Exception as e:
filtered_result[key] = f"Complex value (not displayable): {str(e)}"
# Defense-in-depth: ensure no bearer tokens leak into the rendered HTML even if
# a non-conforming IdP places them in its userinfo response.
safe_raw_claims = {
k: v
for k, v in (received_response or {}).items()
if k not in _OAUTH_TOKEN_FIELDS
}
safe_access_token_claims = {
k: v
for k, v in (access_token_payload or {}).items()
if k not in _OAUTH_TOKEN_FIELDS
}
sso_payload = {
"parsed_by_proxy": filtered_result,
"raw_claims": safe_raw_claims,
"access_token_claims": safe_access_token_claims,
}
# Replace the placeholder in the template with the actual data
sso_payload_json = json.dumps(sso_payload, indent=2, default=str).replace(
"</", "<\\/"
)
html_content = jwt_display_template.replace(
"const userData = SSO_DATA;",
f"const userData = {json.dumps(filtered_result, indent=2)};",
"const ssoData = SSO_DATA;",
f"const ssoData = {sso_payload_json};",
)
return HTMLResponse(content=html_content)

View file

@ -79,7 +79,10 @@ class PrometheusAuthMiddleware:
# Send 401 response directly via ASGI protocol
error_message = getattr(e, "message", str(e))
body = json.dumps(
f"Unauthorized access to metrics endpoint: {error_message}"
f"Unauthorized access to metrics endpoint: {error_message} "
f"To allow unauthenticated access, set "
f"`litellm_settings.require_auth_for_metrics_endpoint: false` "
f"in your proxy_config.yaml."
).encode("utf-8")
await send(
{

View file

@ -813,7 +813,12 @@ def run_server( # noqa: PLR0915
from litellm.proxy.auth.rds_iam_token import generate_iam_auth_token
db_host = os.getenv("DATABASE_HOST")
db_port = os.getenv("DATABASE_PORT")
# Default to the Postgres standard port. Without a default,
# `db_port=None` flows into `boto.generate_db_auth_token(Port=None)`
# and botocore stringifies it to `"None"` while building the
# presigned URL, which then blows up with `ValueError: Port could
# not be cast to integer value as 'None'` during signing.
db_port = os.getenv("DATABASE_PORT", "5432")
db_user = os.getenv("DATABASE_USER")
db_name = os.getenv("DATABASE_NAME")
db_schema = os.getenv("DATABASE_SCHEMA")

View file

@ -211,6 +211,7 @@ from litellm import Router
from litellm._logging import verbose_proxy_logger, verbose_router_logger
from litellm.caching.caching import DualCache, RedisCache
from litellm.caching.redis_cluster_cache import RedisClusterCache
from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
from litellm.constants import (
_REALTIME_BODY_CACHE_SIZE,
@ -6750,27 +6751,64 @@ class ProxyStartupEvent:
"budget_duration not set on Proxy. budget_duration is required to use max_budget."
)
# add proxy budget to db in the user table
asyncio.create_task(
generate_key_helper_fn( # type: ignore
request_type="user",
table_name="user",
user_id=litellm_proxy_budget_name,
duration=None,
models=[],
aliases={},
config={},
spend=0,
max_budget=litellm.max_budget,
budget_duration=litellm.budget_duration,
query_type="update_data",
update_key_values={
"max_budget": litellm.max_budget,
"budget_duration": litellm.budget_duration,
},
)
cls._upsert_proxy_budget_with_reset_at_backfill(litellm_proxy_budget_name)
)
@classmethod
async def _upsert_proxy_budget_with_reset_at_backfill(
cls, litellm_proxy_budget_name: str
) -> None:
"""
Upsert the proxy admin user row with the configured max_budget /
budget_duration, then backfill budget_reset_at if currently NULL.
The backfill uses `WHERE budget_reset_at IS NULL` so it only fires
when the row pre-existed without a reset schedule (e.g. row created
via a different path before the proxy budget was configured). On
subsequent restarts it no-ops, so an active reset window is never
slid forward.
"""
await generate_key_helper_fn( # type: ignore
request_type="user",
table_name="user",
user_id=litellm_proxy_budget_name,
duration=None,
models=[],
aliases={},
config={},
spend=0,
max_budget=litellm.max_budget,
budget_duration=litellm.budget_duration,
query_type="update_data",
update_key_values={
"max_budget": litellm.max_budget,
"budget_duration": litellm.budget_duration,
},
)
# Without this, the upsert leaves budget_reset_at=NULL on rows that
# took the UPDATE path, and reset_budget_for_litellm_users never
# matches them (NULL < now() is unknown in SQL) — so the proxy-wide
# spend cap blocks forever once it's hit.
if prisma_client is not None and litellm.budget_duration is not None:
try:
await prisma_client.db.litellm_usertable.update_many(
where={
"user_id": litellm_proxy_budget_name,
"budget_reset_at": None,
},
data={
"budget_reset_at": get_budget_reset_time(
budget_duration=litellm.budget_duration
)
},
)
except Exception as e:
verbose_proxy_logger.warning(
"Failed to backfill budget_reset_at on proxy admin row: %s", e
)
@classmethod
async def _warm_global_spend_cache(
cls,
@ -8828,6 +8866,7 @@ def _realtime_query_params_template(
return tuple(params)
@app.websocket("/openai/v1/realtime")
@app.websocket("/v1/realtime")
@app.websocket("/realtime")
async def realtime_websocket_endpoint(

View file

@ -95,11 +95,9 @@ async def reserve_budget_for_request(
route=route,
llm_router=llm_router,
)
if reservation_cost is None:
reservation_cost = await _get_smallest_remaining_budget(
counters=counters,
current_spend_by_counter_key=current_spend_by_counter_key,
)
# estimate_request_max_cost still returns None when the model is unknown
# to the cost map (no token-priced cost fields, e.g. image/audio routes).
# In that case we fall back to read-time enforcement only.
if reservation_cost is None or reservation_cost <= 0:
return None
@ -553,32 +551,6 @@ def _coerce_window(window: Any) -> dict:
return {}
async def _get_smallest_remaining_budget(
counters: List[_BudgetCounter],
current_spend_by_counter_key: Dict[str, float],
) -> Optional[float]:
remaining_budget: Optional[float] = None
for counter in counters:
current_spend = await _get_current_counter_value(counter=counter)
current_spend_by_counter_key[counter.counter_key] = current_spend
remaining = counter.max_budget - current_spend
if remaining <= 0:
raise litellm.BudgetExceededError(
current_cost=current_spend,
max_budget=counter.max_budget,
message=(
"Budget has been exceeded! "
f"{counter.entity_type}={counter.entity_id} "
f"Current cost: {current_spend}, "
f"Max budget: {counter.max_budget}"
),
)
remaining_budget = (
remaining if remaining_budget is None else min(remaining_budget, remaining)
)
return remaining_budget
async def _reserve_counter(
counter: _BudgetCounter,
reservation_cost: float,
@ -855,6 +827,13 @@ def _estimate_request_max_cost_for_model(
if model_info is None:
return None
image_cost = _estimate_image_generation_cost(
request_body=request_body,
model_info=model_info,
)
if image_cost is not None:
return image_cost
input_cost_per_token = _to_float(model_info.get("input_cost_per_token"))
output_cost_per_token = _to_float(model_info.get("output_cost_per_token"))
input_tokens = _estimate_input_tokens(
@ -886,6 +865,44 @@ def _estimate_request_max_cost_for_model(
return cost
def _estimate_image_generation_cost(
request_body: dict,
model_info: Dict[str, Any],
) -> Optional[float]:
"""
Reserve `n × per-image cost` for image-generation requests so concurrent
requests against a depleted budget cannot all slip past the admission gate
onto the provider. Token-based pricing (e.g. gpt-image-1) is handled by
the chat-route token path; per-pixel and size/quality-tiered pricing
(DALL-E 2 size variants, premium tiers) are not handled here and fall
through to read-time enforcement.
The "output" vs "input" cost-per-image naming is inconsistent across
providers — OpenAI's dall-e-3 entry uses ``input_cost_per_image`` while
aiml/dall-e-3 uses ``output_cost_per_image`` — so both are summed.
"""
# Gate strictly on `mode`. Several chat and embedding models carry
# ``input_cost_per_image`` / ``output_cost_per_image`` to price multimodal
# *vision input* (e.g. ``gemini-3.1-pro-preview``, ``azure/gpt-realtime-*``,
# ``amazon.titan-embed-image-v1``). Falling back to "treat as image-gen if
# an image cost field is present" would short-circuit the token-priced
# path for those models and reserve a fraction of a cent instead of the
# true per-token cost. All real image-generation entries in
# ``model_prices_and_context_window.json`` carry ``mode: image_generation``
# or ``mode: image_edit``, so the field-presence fallback is unnecessary.
if model_info.get("mode") not in ("image_generation", "image_edit"):
return None
output_cost_per_image = _to_float(model_info.get("output_cost_per_image"))
input_cost_per_image = _to_float(model_info.get("input_cost_per_image"))
cost_per_image = (output_cost_per_image or 0.0) + (input_cost_per_image or 0.0)
if cost_per_image <= 0:
return None
n = _to_int(request_body.get("n")) or 1
return cost_per_image * max(n, 1)
def _get_model_cost_info(
model: str,
llm_router: Optional[Router],
@ -946,6 +963,9 @@ def _estimate_input_tokens(
return None
DEFAULT_MAX_OUTPUT_TOKENS_FALLBACK = 16384
def _estimate_output_tokens(
request_body: dict,
route: str,
@ -954,15 +974,27 @@ def _estimate_output_tokens(
if _is_input_only_route(route=route):
return 0
requested: Optional[int] = None
for key in ("max_completion_tokens", "max_tokens", "max_output_tokens"):
max_tokens = _to_int(request_body.get(key))
if max_tokens is not None:
return max_tokens
requested = _to_int(request_body.get(key))
if requested is not None:
break
# If the caller did not cap output tokens, avoid reserving a model's
# theoretical maximum context. The caller can still admit one request by
# reserving the smallest remaining budget in reserve_budget_for_request().
return None
# Clamp at min(requested-or-default, model_max-or-default). Two purposes:
# (1) Without an explicit cap we still need a finite reservation so the
# atomic admission counter actually bounds concurrent in-flight cost
# (mirrors parallel_request_limiter_v3's DEFAULT_MAX_TOKENS_ESTIMATE).
# (2) An adversarial caller cannot send max_tokens=999999999 to inflate
# the reservation up to remaining team headroom and pin the counter
# at the cap — the model can only physically emit max_output_tokens
# anyway, so reserving more is both wasteful and a DoS surface.
model_ceiling = (
_to_int(model_info.get("max_output_tokens"))
or DEFAULT_MAX_OUTPUT_TOKENS_FALLBACK
)
if requested is None:
requested = DEFAULT_MAX_OUTPUT_TOKENS_FALLBACK
return min(requested, model_ceiling)
def _count_text_tokens(model: str, text: Any) -> int:

View file

@ -113,7 +113,11 @@ from litellm.proxy.db.exception_handler import (
call_with_db_reconnect_retry,
)
from litellm.proxy.db.log_db_metrics import log_db_metrics
from litellm.proxy.db.prisma_client import PrismaWrapper
from litellm.proxy.db.prisma_client import (
PrismaWrapper,
parse_iam_endpoint_from_url,
)
from litellm.proxy.db.routing_prisma_wrapper import RoutingPrismaWrapper
from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import (
UnifiedLLMGuardrails,
)
@ -2569,24 +2573,101 @@ class PrismaClient:
raise Exception(
"Unable to find Prisma binaries. Please run 'prisma generate' first."
)
iam_flag = (
self.iam_token_db_auth if self.iam_token_db_auth is not None else False
)
# When read-replica routing is on, tag log lines with [writer]/[reader]
# so the two wrappers' interleaved IAM refresh logs can be told apart.
# Single-DB deployments get an empty prefix (logs unchanged).
read_replica_url = os.getenv("DATABASE_URL_READ_REPLICA")
writer_log_prefix = "[writer]" if read_replica_url else ""
if http_client is not None:
self.db = PrismaWrapper(
writer_wrapper = PrismaWrapper(
original_prisma=Prisma(http=http_client),
iam_token_db_auth=(
self.iam_token_db_auth
if self.iam_token_db_auth is not None
else False
),
iam_token_db_auth=iam_flag,
log_prefix=writer_log_prefix,
)
else:
self.db = PrismaWrapper(
writer_wrapper = PrismaWrapper(
original_prisma=Prisma(),
iam_token_db_auth=(
self.iam_token_db_auth
if self.iam_token_db_auth is not None
else False
),
) # Client to connect to Prisma db
iam_token_db_auth=iam_flag,
log_prefix=writer_log_prefix,
)
# Optional read-replica routing. When DATABASE_URL_READ_REPLICA is set,
# reads (find_*, count, group_by, query_raw/_first) are routed to the
# reader endpoint and writes stay on the writer. Falls back to the
# writer-only wrapper when the env var is unset, preserving existing
# single-DB deployments.
self.db: Union[PrismaWrapper, RoutingPrismaWrapper]
if read_replica_url:
try:
# If IAM auth is enabled, the reader refreshes its own token on
# the same cadence as the writer. We parse the static endpoint
# pieces (host/port/user/db) once from the reader URL — only
# the IAM token rotates after that.
reader_iam_endpoint = (
parse_iam_endpoint_from_url(read_replica_url) if iam_flag else None
)
# Mint a fresh IAM token for the reader BEFORE constructing the
# Prisma client. Mirrors what `proxy_cli.py` already does for
# the writer (proxy_cli.py:812-832) — without this, the reader
# Prisma is built with whatever placeholder URL the user
# supplied (no real token), and the first query falls through
# to the synchronous fallback path in
# `PrismaWrapper.__getattr__`, which deadlocks the event loop
# and times out after 30s.
if iam_flag and reader_iam_endpoint is not None:
from litellm.proxy.auth.rds_iam_token import (
generate_iam_auth_token,
)
reader_token = generate_iam_auth_token(
db_host=reader_iam_endpoint.host,
db_port=reader_iam_endpoint.port,
db_user=reader_iam_endpoint.user,
)
read_replica_url = reader_iam_endpoint.build_url(reader_token)
os.environ["DATABASE_URL_READ_REPLICA"] = read_replica_url
reader_kwargs: Dict[str, Any] = {
"datasource": {"url": read_replica_url}
}
if http_client is not None:
reader_prisma = Prisma(http=http_client, **reader_kwargs)
else:
reader_prisma = Prisma(**reader_kwargs)
reader_wrapper = PrismaWrapper(
original_prisma=reader_prisma,
iam_token_db_auth=iam_flag,
db_url_env_var="DATABASE_URL_READ_REPLICA",
iam_endpoint=reader_iam_endpoint,
recreate_uses_datasource=True,
log_prefix="[reader]",
)
self.db = RoutingPrismaWrapper(
writer=writer_wrapper, reader=reader_wrapper
)
verbose_proxy_logger.info(
"PrismaClient: read-replica routing enabled via DATABASE_URL_READ_REPLICA"
+ (" (with IAM token auto-refresh)" if iam_flag else "")
)
except Exception as e:
# Reader is opt-in; never let its construction fail proxy
# startup. Mirrors the runtime contract from
# `RoutingPrismaWrapper.connect`: reader-side failures are
# logged and we keep serving traffic via the writer alone.
# This recovers from transient AWS STS hiccups during the
# reader IAM token mint, malformed DATABASE_URL_READ_REPLICA,
# and Prisma construction errors. Operator restart is required
# to retry read-routing once the underlying issue is resolved.
verbose_proxy_logger.warning(
"Failed to initialize read replica Prisma client: %s. "
"Falling back to writer-only mode (no read routing) until proxy restart.",
e,
)
self.db = writer_wrapper
else:
self.db = writer_wrapper # Client to connect to Prisma db
self._db_reconnect_lock = asyncio.Lock()
self._db_health_watchdog_task: Optional[asyncio.Task] = None
self._db_last_reconnect_attempt_ts: float = 0.0
@ -2624,6 +2705,13 @@ class PrismaClient:
self._engine_wait_thread: Optional[threading.Thread] = None
verbose_proxy_logger.debug("Success - Created Prisma Client")
@property
def writer_db(self) -> PrismaWrapper:
"""Underlying writer Prisma wrapper, regardless of read-replica routing."""
if isinstance(self.db, RoutingPrismaWrapper):
return self.db.writer
return self.db
def get_request_status(
self, payload: Union[dict, SpendLogsPayload]
) -> Literal["success", "failure"]:
@ -4272,7 +4360,10 @@ class PrismaClient:
self._cleanup_engine_watcher()
await self.db.recreate_prisma_client(db_url)
await self._start_engine_watcher()
await self.db.query_raw("SELECT 1")
# Smoke-test the writer specifically; query_raw on the routing
# wrapper sends to the reader, which would not validate the
# newly-recreated writer engine.
await self.writer_db.query_raw("SELECT 1")
await asyncio.wait_for(_do_direct_reconnect(), timeout=effective_timeout)

View file

@ -633,6 +633,16 @@ class BaseLitellmParams(
),
)
skip_tool_message_in_guardrail: Optional[bool] = Field(
default=None,
description=(
"When True, unified guardrails skip tool-role messages when building "
"evaluation inputs (texts and structured_messages). When False, tool "
"messages are included even if litellm_settings sets a global skip. When "
"None, use the global litellm.skip_tool_message_in_guardrail setting."
),
)
# Lakera specific params
category_thresholds: Optional[LakeraCategoryThresholds] = Field(
default=None,

View file

@ -825,6 +825,7 @@ API_ROUTE_TO_CALL_TYPES = {
# Realtime API
"/realtime": [CallTypes.arealtime],
"/v1/realtime": [CallTypes.arealtime],
"/openai/v1/realtime": [CallTypes.arealtime],
# Provider-specific routes
"/anthropic/v1/messages": [CallTypes.anthropic_messages],
# Google GenAI routes

View file

@ -27192,6 +27192,20 @@
"supports_reasoning": true,
"supports_tool_choice": true
},
"openrouter/qwen/qwen3.6-plus": {
"input_cost_per_token": 3.25e-07,
"litellm_provider": "openrouter",
"max_input_tokens": 1000000,
"max_output_tokens": 65536,
"max_tokens": 65536,
"mode": "chat",
"output_cost_per_token": 1.95e-06,
"source": "https://openrouter.ai/qwen/qwen3.6-plus",
"supports_function_calling": true,
"supports_reasoning": true,
"supports_tool_choice": true,
"supports_vision": true
},
"openrouter/qwen/qwen3.5-35b-a3b": {
"input_cost_per_token": 2.5e-07,
"litellm_provider": "openrouter",

View file

@ -22,7 +22,7 @@ dependencies = [
"importlib-metadata>=8.0.0,<9.0",
"tokenizers>=0.21.0,<1.0",
"click>=8.0.0,<9.0",
"jinja2>=3.1.0,<4.0",
"jinja2>=3.1.6,<4.0",
"aiohttp>=3.10,<4.0",
"pydantic>=2.10.0,<3.0.0",
"jsonschema>=4.0.0,<5.0",

View file

@ -25,8 +25,8 @@ async def make_calls_until_budget_exceeded(session, key: str, call_function, **k
# Check error structure and values that should be consistent
assert (
error_dict["code"] == "400"
), f"Expected error code 400, got: {error_dict['code']}"
error_dict["code"] == "429"
), f"Expected error code 429, got: {error_dict['code']}"
assert (
error_dict["type"] == "budget_exceeded"
), f"Expected error type budget_exceeded, got: {error_dict['type']}"

View file

@ -99,7 +99,7 @@ async def test_model_access_patterns(key_models, test_model, expect_success):
# Assert error structure and values
assert _error_body["type"] == "key_model_access_denied"
assert _error_body["param"] == "model"
assert _error_body["code"] == "401"
assert _error_body["code"] == "403"
assert "key not allowed to access model" in _error_body["message"]
@ -297,7 +297,7 @@ def _validate_model_access_exception(
# Assert error structure and values
assert _error_body["type"] == expected_type
assert _error_body["param"] == "model"
assert _error_body["code"] == "401"
assert _error_body["code"] == "403"
if expected_type == "key_model_access_denied":
assert "key not allowed to access model" in _error_body["message"]
elif expected_type == "team_model_access_denied":

View file

@ -413,3 +413,71 @@ async def test_async_log_success_event_uses_end_user_model_budget_duration(
f"{END_USER_SPEND_CACHE_KEY_PREFIX}:{end_user_id}:{model}:{budget_duration}"
)
assert call_kwargs["response_cost"] == 0.05
@pytest.mark.asyncio
async def test_async_log_success_event_pushes_redis_increments_when_redis_configured():
"""
Virtual-key model max budget limiter does not run RouterBudgetLimiting.__init__,
so the periodic Redis flush task never starts. After logging spend we must call
_push_in_memory_increments_to_redis when Redis is wired so other workers see spend.
"""
dual_cache = DualCache()
dual_cache.redis_cache = object() # truthy placeholder; push only checks is not None
limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache)
model = "gpt-4"
kwargs = {
"standard_logging_object": {
"response_cost": 0.01,
"model": model,
"metadata": {"user_api_key_hash": "vk-hash"},
},
"litellm_params": {
"metadata": {
"user_api_key_model_max_budget": {
model: {"budget_limit": 10.0, "time_period": "1d"},
},
},
},
}
with patch.object(limiter, "_increment_spend_for_key", new_callable=AsyncMock):
with patch.object(
limiter,
"_push_in_memory_increments_to_redis",
new_callable=AsyncMock,
) as mock_push:
await limiter.async_log_success_event(
kwargs, response_obj=None, start_time=None, end_time=None
)
mock_push.assert_awaited_once()
@pytest.mark.asyncio
async def test_async_log_success_event_skips_redis_push_without_redis(budget_limiter):
"""When dual_cache has no Redis backend, do not await _push_in_memory_increments_to_redis."""
assert budget_limiter.dual_cache.redis_cache is None
model = "gpt-4"
kwargs = {
"standard_logging_object": {
"response_cost": 0.01,
"model": model,
"metadata": {"user_api_key_hash": "vk-hash"},
},
"litellm_params": {
"metadata": {
"user_api_key_model_max_budget": {
model: {"budget_limit": 10.0, "time_period": "1d"},
},
},
},
}
with patch.object(budget_limiter, "_increment_spend_for_key", new_callable=AsyncMock):
with patch.object(
budget_limiter,
"_push_in_memory_increments_to_redis",
new_callable=AsyncMock,
) as mock_push:
await budget_limiter.async_log_success_event(
kwargs, response_obj=None, start_time=None, end_time=None
)
mock_push.assert_not_awaited()

View file

@ -442,6 +442,109 @@ class TestOpenTelemetryDualHandlerIsolation(unittest.TestCase):
)
class TestOpenTelemetryCaptureMessageContent(unittest.TestCase):
"""OTEL_INSTRUMENTATION_GENAI_CAPTURE_MESSAGE_CONTENT and the
OpenTelemetryConfig.capture_message_content programmatic override
drive what the handler captures in spans vs events."""
@staticmethod
def _make(env=None, config_value=None, message_logging=True):
env_dict = (
{"OTEL_INSTRUMENTATION_GENAI_CAPTURE_MESSAGE_CONTENT": env}
if env is not None
else {"OTEL_INSTRUMENTATION_GENAI_CAPTURE_MESSAGE_CONTENT": ""}
)
with patch.dict(os.environ, env_dict):
handler = OpenTelemetry(
config=OpenTelemetryConfig(
exporter="console", capture_message_content=config_value
)
)
handler.message_logging = message_logging
return handler, handler._resolve_capture_mode()
def test_no_explicit_setting_falls_back_to_message_logging_true(self):
_, mode = self._make()
self.assertEqual(mode, "SPAN_AND_EVENT")
def test_no_explicit_setting_falls_back_to_message_logging_false(self):
_, mode = self._make(message_logging=False)
self.assertEqual(mode, "NO_CONTENT")
def test_env_var_no_content(self):
_, mode = self._make(env="NO_CONTENT")
self.assertEqual(mode, "NO_CONTENT")
def test_env_var_span_only(self):
_, mode = self._make(env="SPAN_ONLY")
self.assertEqual(mode, "SPAN_ONLY")
def test_env_var_event_only(self):
_, mode = self._make(env="EVENT_ONLY")
self.assertEqual(mode, "EVENT_ONLY")
def test_env_var_span_and_event(self):
_, mode = self._make(env="SPAN_AND_EVENT")
self.assertEqual(mode, "SPAN_AND_EVENT")
def test_env_var_legacy_true_maps_to_event_only(self):
_, mode = self._make(env="true")
self.assertEqual(mode, "EVENT_ONLY")
def test_env_var_legacy_false_maps_to_no_content(self):
for env in ("false", "0"):
with self.subTest(env=env):
_, mode = self._make(env=env)
self.assertEqual(mode, "NO_CONTENT")
def test_env_var_unknown_value_falls_through_to_legacy(self):
_, mode = self._make(env="garbage", message_logging=True)
self.assertEqual(mode, "SPAN_AND_EVENT")
def test_config_field_overrides_env(self):
_, mode = self._make(env="EVENT_ONLY", config_value="SPAN_ONLY")
self.assertEqual(mode, "SPAN_ONLY")
def test_turn_off_message_logging_forces_no_content(self):
with patch("litellm.turn_off_message_logging", True):
_, mode = self._make(env="SPAN_AND_EVENT", message_logging=True)
self.assertEqual(mode, "NO_CONTENT")
def test_capture_in_span_and_event_predicates(self):
cases = {
"NO_CONTENT": (False, False),
"SPAN_ONLY": (True, False),
"EVENT_ONLY": (False, True),
"SPAN_AND_EVENT": (True, True),
}
for mode, (in_span, in_event) in cases.items():
handler, _ = self._make(env=mode)
self.assertEqual(handler._capture_in_span(), in_span, msg=mode)
self.assertEqual(handler._capture_in_event(), in_event, msg=mode)
def test_two_handlers_can_have_different_modes(self):
# FIL's stated requirement: one handler strips content, the other keeps it.
with patch.dict(
os.environ, {"OTEL_INSTRUMENTATION_GENAI_CAPTURE_MESSAGE_CONTENT": ""}
):
stripped = OpenTelemetry(
config=OpenTelemetryConfig(
exporter="console", capture_message_content="NO_CONTENT"
)
)
kept = OpenTelemetry(
config=OpenTelemetryConfig(
exporter="console", capture_message_content="SPAN_AND_EVENT"
)
)
self.assertEqual(stripped._resolve_capture_mode(), "NO_CONTENT")
self.assertEqual(kept._resolve_capture_mode(), "SPAN_AND_EVENT")
self.assertFalse(stripped._capture_in_span())
self.assertFalse(stripped._capture_in_event())
self.assertTrue(kept._capture_in_span())
self.assertTrue(kept._capture_in_event())
class TestOpenTelemetry(unittest.TestCase):
POLL_INTERVAL = 0.05
POLL_TIMEOUT = 2.0
@ -1067,6 +1170,7 @@ class TestOpenTelemetry(unittest.TestCase):
result = otel._get_span_name(kwargs)
self.assertEqual(result, LITELLM_REQUEST_SPAN_NAME)
@patch.dict(os.environ, {"OTEL_INSTRUMENTATION_GENAI_CAPTURE_MESSAGE_CONTENT": ""})
@patch("litellm.turn_off_message_logging", False)
def test_maybe_log_raw_request_creates_span(self):
"""Test _maybe_log_raw_request creates span when logging enabled"""
@ -2194,6 +2298,19 @@ class TestOpenTelemetrySemanticConventions138(unittest.TestCase):
See: https://github.com/BerriAI/litellm/issues/17794
"""
def setUp(self):
# Insulate from a shell-set OTEL_INSTRUMENTATION_GENAI_CAPTURE_MESSAGE_CONTENT
# so these tests exercise the legacy default path (message_logging=True).
self._prev = os.environ.pop(
"OTEL_INSTRUMENTATION_GENAI_CAPTURE_MESSAGE_CONTENT", None
)
def tearDown(self):
if self._prev is not None:
os.environ["OTEL_INSTRUMENTATION_GENAI_CAPTURE_MESSAGE_CONTENT"] = (
self._prev
)
def test_input_messages_uses_parts_structure(self):
"""
Test that gen_ai.input.messages uses the OTEL 1.38 parts array structure.

View file

@ -15,6 +15,8 @@ from unittest.mock import AsyncMock, patch
import pytest
from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import (
_request_auth_header,
_request_extra_headers,
_resolve_param_list,
_resolve_ref,
build_input_schema,
@ -1011,3 +1013,197 @@ class TestRegisterToolsFromOpenAPI:
assert re.match(
r"^[a-zA-Z0-9_-]+$", name
), f"fallback tool name {name!r} not sanitized"
class TestRequestExtraHeaders:
"""Tests for _request_extra_headers ContextVar forwarding in tool_function."""
@pytest.mark.asyncio
async def test_extra_headers_forwarded_to_upstream(self):
"""Extra headers set via ContextVar are included in the upstream request."""
operation = {}
func = create_tool_function(
path="/data",
method="get",
operation=operation,
base_url="https://api.example.com",
)
with patch(GET_ASYNC_CLIENT_TARGET) as mock_client:
async_client = _create_mock_client("get", "ok")
mock_client.return_value = async_client
token = _request_extra_headers.set({"X-TOKEN": "secret-value"})
try:
result = await func()
finally:
_request_extra_headers.reset(token)
assert result == "ok"
call_args = async_client.get.call_args
headers_sent = call_args[1]["headers"]
assert headers_sent.get("X-TOKEN") == "secret-value"
@pytest.mark.asyncio
async def test_no_extra_headers_by_default(self):
"""Without setting _request_extra_headers, no extra headers are injected."""
operation = {}
func = create_tool_function(
path="/data",
method="get",
operation=operation,
base_url="https://api.example.com",
headers={"X-Static": "static-value"},
)
with patch(GET_ASYNC_CLIENT_TARGET) as mock_client:
async_client = _create_mock_client("get", "ok")
mock_client.return_value = async_client
result = await func()
assert result == "ok"
call_args = async_client.get.call_args
headers_sent = call_args[1]["headers"]
assert headers_sent == {"X-Static": "static-value"}
assert "X-TOKEN" not in headers_sent
@pytest.mark.asyncio
async def test_extra_headers_merged_with_static_headers(self):
"""Forwarded headers are passed through alongside non-conflicting static headers."""
operation = {}
func = create_tool_function(
path="/data",
method="post",
operation=operation,
base_url="https://api.example.com",
headers={"X-Static": "static-value"},
)
with patch(GET_ASYNC_CLIENT_TARGET) as mock_client:
async_client = _create_mock_client("post", "created")
mock_client.return_value = async_client
token = _request_extra_headers.set({"X-TOKEN": "dynamic-value"})
try:
result = await func()
finally:
_request_extra_headers.reset(token)
assert result == "created"
call_args = async_client.post.call_args
headers_sent = call_args[1]["headers"]
assert headers_sent.get("X-Static") == "static-value"
assert headers_sent.get("X-TOKEN") == "dynamic-value"
@pytest.mark.asyncio
async def test_static_headers_win_over_forwarded_on_conflict(self):
"""Static (operator) headers must override forwarded (caller) headers on name conflict."""
operation = {}
func = create_tool_function(
path="/data",
method="get",
operation=operation,
base_url="https://api.example.com",
headers={"X-Tenant": "operator-tenant"},
)
with patch(GET_ASYNC_CLIENT_TARGET) as mock_client:
async_client = _create_mock_client("get", "ok")
mock_client.return_value = async_client
token = _request_extra_headers.set({"X-Tenant": "caller-spoofed"})
try:
result = await func()
finally:
_request_extra_headers.reset(token)
assert result == "ok"
call_args = async_client.get.call_args
headers_sent = call_args[1]["headers"]
assert headers_sent.get("X-Tenant") == "operator-tenant"
assert "caller-spoofed" not in headers_sent.values()
@pytest.mark.asyncio
async def test_static_headers_win_case_insensitively(self):
"""Forwarded header with different casing must not bypass the static-wins rule."""
operation = {}
func = create_tool_function(
path="/data",
method="get",
operation=operation,
base_url="https://api.example.com",
headers={"X-Tenant": "operator-tenant"},
)
with patch(GET_ASYNC_CLIENT_TARGET) as mock_client:
async_client = _create_mock_client("get", "ok")
mock_client.return_value = async_client
token = _request_extra_headers.set({"x-tenant": "caller-spoofed"})
try:
result = await func()
finally:
_request_extra_headers.reset(token)
assert result == "ok"
call_args = async_client.get.call_args
headers_sent = call_args[1]["headers"]
assert headers_sent.get("X-Tenant") == "operator-tenant"
assert "x-tenant" not in headers_sent
assert "caller-spoofed" not in headers_sent.values()
@pytest.mark.asyncio
async def test_auth_header_still_overrides_extra_headers(self):
"""_request_auth_header takes precedence for Authorization over extra headers."""
operation = {}
func = create_tool_function(
path="/secure",
method="get",
operation=operation,
base_url="https://api.example.com",
)
with patch(GET_ASYNC_CLIENT_TARGET) as mock_client:
async_client = _create_mock_client("get", "secure-data")
mock_client.return_value = async_client
extra_token = _request_extra_headers.set(
{"Authorization": "Bearer extra", "X-TOKEN": "token-value"}
)
auth_token = _request_auth_header.set("Bearer byok-credential")
try:
result = await func()
finally:
_request_auth_header.reset(auth_token)
_request_extra_headers.reset(extra_token)
assert result == "secure-data"
call_args = async_client.get.call_args
headers_sent = call_args[1]["headers"]
assert headers_sent.get("Authorization") == "Bearer byok-credential"
assert headers_sent.get("X-TOKEN") == "token-value"
@pytest.mark.asyncio
async def test_extra_headers_not_leaked_between_calls(self):
"""After resetting the ContextVar, subsequent calls do not see the headers."""
operation = {}
func = create_tool_function(
path="/data",
method="get",
operation=operation,
base_url="https://api.example.com",
)
with patch(GET_ASYNC_CLIENT_TARGET) as mock_client:
async_client = _create_mock_client("get", "ok")
mock_client.return_value = async_client
token = _request_extra_headers.set({"X-TOKEN": "first-call"})
_request_extra_headers.reset(token)
await func()
call_args = async_client.get.call_args
headers_sent = call_args[1]["headers"]
assert "X-TOKEN" not in headers_sent

View file

@ -12,6 +12,7 @@ from datetime import datetime, timedelta
import httpx
import pytest
from fastapi import status
import litellm
from litellm.proxy._types import (
@ -31,6 +32,7 @@ from litellm.proxy._types import (
)
from litellm.proxy.auth.auth_checks import (
ExperimentalUIJWTToken,
_can_object_call_model,
_can_object_call_vector_stores,
_check_end_user_budget,
_check_team_member_budget,
@ -206,6 +208,52 @@ def test_get_key_object_from_ui_hash_key_invalid():
assert key_object is None
@pytest.mark.parametrize(
"object_type,expected_error_type",
[
("key", ProxyErrorTypes.key_model_access_denied),
("team", ProxyErrorTypes.team_model_access_denied),
("user", ProxyErrorTypes.user_model_access_denied),
("org", ProxyErrorTypes.org_model_access_denied),
("project", ProxyErrorTypes.project_model_access_denied),
],
)
def test_can_object_call_model_denials_return_forbidden(
object_type, expected_error_type
):
with pytest.raises(ProxyException) as exc_info:
_can_object_call_model(
model="restricted-model",
llm_router=None,
models=["allowed-model"],
object_type=object_type,
)
assert exc_info.value.type == expected_error_type
assert int(exc_info.value.code) == status.HTTP_403_FORBIDDEN
@pytest.mark.asyncio
async def test_can_user_call_model_no_default_models_returns_forbidden():
from litellm.proxy._types import SpecialModelNames
from litellm.proxy.auth.auth_checks import can_user_call_model
user_object = LiteLLM_UserTable(
user_id="test-user",
models=[SpecialModelNames.no_default_models.value],
)
with pytest.raises(ProxyException) as exc_info:
await can_user_call_model(
model="restricted-model",
llm_router=None,
user_object=user_object,
)
assert exc_info.value.type == ProxyErrorTypes.key_model_access_denied
assert int(exc_info.value.code) == status.HTTP_403_FORBIDDEN
@pytest.mark.asyncio
async def test_get_key_object_should_reconnect_once_on_db_connection_error():
mock_prisma_client = MagicMock()
@ -1144,6 +1192,7 @@ async def test_check_team_member_model_access_denied_model():
proxy_logging_obj=MagicMock(),
)
assert exc_info.value.type == ProxyErrorTypes.team_model_access_denied
assert int(exc_info.value.code) == status.HTTP_403_FORBIDDEN
@pytest.mark.asyncio

View file

@ -140,6 +140,7 @@ async def test_handle_authentication_error_budget_exceeded():
)
assert exc_info.value.type == ProxyErrorTypes.budget_exceeded
assert int(exc_info.value.code) == status.HTTP_429_TOO_MANY_REQUESTS
@pytest.mark.asyncio

View file

@ -80,6 +80,28 @@ def test_compliance_routes_open_to_internal_user(route):
)
def test_health_test_connection_route_delegates_internal_user_auth_to_endpoint():
"""Team model test-connection requests are authorized by the endpoint."""
role = LitellmUserRoles.INTERNAL_USER.value
user_obj = LiteLLM_UserTable(
user_id="test_user",
user_email="test@example.com",
user_role=role,
)
valid_token = UserAPIKeyAuth(user_id="test_user", user_role=role)
request = MagicMock(spec=Request)
request.query_params = {}
RouteChecks.non_proxy_admin_allowed_routes_check(
user_obj=user_obj,
_user_role=role,
route="/health/test_connection",
request=request,
valid_token=valid_token,
request_data={},
)
@pytest.mark.parametrize(
"route",
["/compliance/eu-ai-act", "/compliance/gdpr"],

View file

@ -1,6 +1,7 @@
import json
import os
import sys
from datetime import datetime, timedelta
from types import SimpleNamespace
from unittest.mock import ANY, AsyncMock, MagicMock, patch
@ -9,6 +10,7 @@ sys.path.insert(
) # Adds the parent directory to the system path
import pytest
from fastapi import status
import litellm
import litellm.proxy.proxy_server
@ -178,6 +180,26 @@ async def test_custom_auth_does_not_enforce_key_model_access_by_default():
mock_can_key.assert_not_awaited()
@pytest.mark.asyncio
async def test_post_custom_auth_expired_key_returns_unauthorized():
expired_token = UserAPIKeyAuth(
token="test_token",
expires=datetime.now() - timedelta(minutes=1),
)
with pytest.raises(ProxyException) as exc_info:
await _run_post_custom_auth_checks(
valid_token=expired_token,
request=MagicMock(),
request_data={},
route="/v1/chat/completions",
parent_otel_span=None,
)
assert exc_info.value.type == ProxyErrorTypes.expired_key
assert int(exc_info.value.code) == status.HTTP_401_UNAUTHORIZED
@pytest.mark.asyncio
async def test_custom_auth_honors_key_level_model_access_restriction_allowed_with_opt_in():
valid_token = UserAPIKeyAuth(token="test_token", models=["gpt-4o-mini"])
@ -934,6 +956,7 @@ async def test_proxy_admin_expired_key_from_cache():
assert (
exc_info.value.type == ProxyErrorTypes.expired_key
), f"Expected expired_key error type, got {exc_info.value.type}"
assert int(exc_info.value.code) == status.HTTP_401_UNAUTHORIZED
assert "Expired Key" in str(
exc_info.value.message
), f"Exception message should mention 'Expired Key', got: {exc_info.value.message}"

View file

@ -0,0 +1,887 @@
import asyncio
import logging
import os
import sys
from typing import Any, Dict
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
sys.path.insert(0, os.path.abspath("../../../.."))
# NOTE: do NOT patch sys.modules["prisma"] file-wide via an autouse fixture.
# Doing so leaks across pytest-xdist test scheduling: when a worker runs a
# routing test, then later runs test_exception_handler.py, the cached MagicMock
# attribute references break `isinstance(e, prisma.errors.X)` in
# `is_database_transport_error`. The two tests below that actually need to
# stub the prisma SDK do so per-test via monkeypatch, which is properly scoped.
def _make_wrappers():
from litellm.proxy.db.prisma_client import PrismaWrapper
writer_inner = MagicMock(name="writer_prisma")
reader_inner = MagicMock(name="reader_prisma")
writer = PrismaWrapper(original_prisma=writer_inner, iam_token_db_auth=False)
reader = PrismaWrapper(original_prisma=reader_inner, iam_token_db_auth=False)
return writer, writer_inner, reader, reader_inner
class _FakeActions:
"""Stand-in for a Prisma per-model Actions instance (non-callable, has find_many/create)."""
def __init__(self, name: str):
self._name = name
for method in (
"find_many",
"find_unique",
"find_first",
"count",
"group_by",
"create",
"update",
"upsert",
"delete",
"delete_many",
"update_many",
):
setattr(self, method, MagicMock(name=f"{name}.{method}"))
def _model_actions_mock(name: str) -> _FakeActions:
return _FakeActions(name)
def test_top_level_query_raw_routes_to_reader():
from litellm.proxy.db.routing_prisma_wrapper import RoutingPrismaWrapper
writer, writer_inner, reader, reader_inner = _make_wrappers()
routing = RoutingPrismaWrapper(writer=writer, reader=reader)
# query_raw should resolve to the reader's underlying client.
assert routing.query_raw is reader_inner.query_raw
assert routing.query_first is reader_inner.query_first
def test_top_level_execute_raw_routes_to_writer():
from litellm.proxy.db.routing_prisma_wrapper import RoutingPrismaWrapper
writer, writer_inner, reader, reader_inner = _make_wrappers()
routing = RoutingPrismaWrapper(writer=writer, reader=reader)
# execute_raw, batch_, tx are write-side and must hit the writer.
assert routing.execute_raw is writer_inner.execute_raw
assert routing.batch_ is writer_inner.batch_
assert routing.tx is writer_inner.tx
def test_per_model_reads_route_to_reader_writes_to_writer():
from litellm.proxy.db.routing_prisma_wrapper import RoutingPrismaWrapper
writer, writer_inner, reader, reader_inner = _make_wrappers()
writer_inner.litellm_usertable = _model_actions_mock("writer_users")
reader_inner.litellm_usertable = _model_actions_mock("reader_users")
routing = RoutingPrismaWrapper(writer=writer, reader=reader)
actions = routing.litellm_usertable
# Reads → reader actions.
assert actions.find_many is reader_inner.litellm_usertable.find_many
assert actions.find_unique is reader_inner.litellm_usertable.find_unique
assert actions.find_first is reader_inner.litellm_usertable.find_first
assert actions.count is reader_inner.litellm_usertable.count
assert actions.group_by is reader_inner.litellm_usertable.group_by
# Writes → writer actions.
assert actions.create is writer_inner.litellm_usertable.create
assert actions.update is writer_inner.litellm_usertable.update
assert actions.upsert is writer_inner.litellm_usertable.upsert
assert actions.delete is writer_inner.litellm_usertable.delete
assert actions.update_many is writer_inner.litellm_usertable.update_many
assert actions.delete_many is writer_inner.litellm_usertable.delete_many
@pytest.mark.asyncio
async def test_connect_invokes_both_clients():
from litellm.proxy.db.routing_prisma_wrapper import RoutingPrismaWrapper
writer, writer_inner, reader, reader_inner = _make_wrappers()
writer_inner.connect = AsyncMock()
reader_inner.connect = AsyncMock()
routing = RoutingPrismaWrapper(writer=writer, reader=reader)
await routing.connect()
writer_inner.connect.assert_awaited_once()
reader_inner.connect.assert_awaited_once()
@pytest.mark.asyncio
async def test_connect_logs_writer_and_reader_success(caplog):
"""Successful startup emits a positive INFO confirmation for both writer
and reader so operators can verify connectivity without inspecting the URL
in logs."""
from litellm.proxy.db.routing_prisma_wrapper import RoutingPrismaWrapper
writer, writer_inner, reader, reader_inner = _make_wrappers()
writer_inner.connect = AsyncMock()
reader_inner.connect = AsyncMock()
routing = RoutingPrismaWrapper(writer=writer, reader=reader)
with caplog.at_level(logging.INFO, logger="LiteLLM Proxy"):
await routing.connect()
messages = [r.getMessage() for r in caplog.records]
assert "[writer] DB connected" in messages
assert "[reader] DB connected" in messages
@pytest.mark.asyncio
async def test_disconnect_continues_when_one_side_fails():
from litellm.proxy.db.routing_prisma_wrapper import RoutingPrismaWrapper
writer, writer_inner, reader, reader_inner = _make_wrappers()
writer_inner.disconnect = AsyncMock(side_effect=RuntimeError("writer down"))
reader_inner.disconnect = AsyncMock()
routing = RoutingPrismaWrapper(writer=writer, reader=reader)
with pytest.raises(RuntimeError, match="writer down"):
await routing.disconnect()
# Reader still attempted even though writer raised.
reader_inner.disconnect.assert_awaited_once()
def test_is_connected_reflects_writer_only():
"""is_connected() must NOT depend on reader health — a healthy writer with
a degraded reader should report True so that PrismaClient.connect()'s
health check does not re-trigger a writer reconnect (which only fixes
writer-side problems and would loop indefinitely)."""
from litellm.proxy.db.routing_prisma_wrapper import RoutingPrismaWrapper
writer, writer_inner, reader, reader_inner = _make_wrappers()
routing = RoutingPrismaWrapper(writer=writer, reader=reader)
writer_inner.is_connected = MagicMock(return_value=True)
reader_inner.is_connected = MagicMock(return_value=True)
assert routing.is_connected() is True
# Reader down → still True (reader degradation is tracked separately).
reader_inner.is_connected = MagicMock(return_value=False)
assert routing.is_connected() is True
# Writer down → False.
writer_inner.is_connected = MagicMock(return_value=False)
assert routing.is_connected() is False
def test_token_refresh_delegates_to_both_writer_and_reader():
from litellm.proxy.db.routing_prisma_wrapper import RoutingPrismaWrapper
writer = MagicMock()
writer.start_token_refresh_task = AsyncMock()
writer.stop_token_refresh_task = AsyncMock()
reader = MagicMock()
reader.start_token_refresh_task = AsyncMock()
reader.stop_token_refresh_task = AsyncMock()
routing = RoutingPrismaWrapper(writer=writer, reader=reader)
asyncio.run(routing.start_token_refresh_task())
asyncio.run(routing.stop_token_refresh_task())
# Both wrappers get start/stop — each manages its own IAM token. When
# IAM is disabled on a wrapper its task body is a no-op.
writer.start_token_refresh_task.assert_awaited_once()
writer.stop_token_refresh_task.assert_awaited_once()
reader.start_token_refresh_task.assert_awaited_once()
reader.stop_token_refresh_task.assert_awaited_once()
def test_routed_actions_falls_back_to_writer_for_unknown_methods():
from litellm.proxy.db.routing_prisma_wrapper import _RoutedActions
writer_actions = _model_actions_mock("writer")
writer_actions.some_custom_method = "writer-custom"
reader_actions = _model_actions_mock("reader")
reader_actions.some_custom_method = "reader-custom"
routed = _RoutedActions(writer_actions, reader_actions, lambda: True)
# Unknown method → defaults to writer (safe fallback for write-like ops).
assert routed.some_custom_method == "writer-custom"
def test_routed_actions_respects_should_use_reader_flag():
"""When the routing wrapper marks the reader unavailable, _RoutedActions
must redirect reads to the writer instead — without needing to re-fetch
the actions accessor."""
from litellm.proxy.db.routing_prisma_wrapper import _RoutedActions
writer_actions = _model_actions_mock("writer")
reader_actions = _model_actions_mock("reader")
use_reader = {"value": True}
routed = _RoutedActions(writer_actions, reader_actions, lambda: use_reader["value"])
# Reader healthy → reads to reader.
assert routed.find_many is reader_actions.find_many
# Reader degrades mid-flight → next read goes to writer.
use_reader["value"] = False
assert routed.find_many is writer_actions.find_many
# ---------------------------------------------------------------------------
# Reader graceful degradation
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_connect_swallows_reader_failure_and_falls_back_to_writer():
"""A reader connect failure must NOT abort proxy startup. The wrapper
flips into degraded mode so subsequent reads route to the writer."""
from litellm.proxy.db.routing_prisma_wrapper import RoutingPrismaWrapper
writer, writer_inner, reader, reader_inner = _make_wrappers()
writer_inner.connect = AsyncMock()
reader_inner.connect = AsyncMock(side_effect=RuntimeError("reader unreachable"))
routing = RoutingPrismaWrapper(writer=writer, reader=reader)
# Must not raise — reader failure is non-fatal.
await routing.connect()
assert routing.reader_unavailable is True
writer_inner.connect.assert_awaited_once()
reader_inner.connect.assert_awaited_once()
@pytest.mark.asyncio
async def test_reads_route_to_writer_when_reader_unavailable():
"""Top-level read methods and per-model reads must fall through to the
writer while the reader is degraded."""
from litellm.proxy.db.routing_prisma_wrapper import RoutingPrismaWrapper
writer, writer_inner, reader, reader_inner = _make_wrappers()
writer_inner.litellm_usertable = _model_actions_mock("writer_users")
reader_inner.litellm_usertable = _model_actions_mock("reader_users")
routing = RoutingPrismaWrapper(writer=writer, reader=reader)
routing._reader_unavailable = True
# Top-level reads → writer.
assert routing.query_raw is writer_inner.query_raw
assert routing.query_first is writer_inner.query_first
# Per-model reads → writer actions.
actions = routing.litellm_usertable
assert actions.find_many is writer_inner.litellm_usertable.find_many
assert actions.find_unique is writer_inner.litellm_usertable.find_unique
@pytest.mark.asyncio
async def test_recreate_prisma_client_recreates_both_writer_and_reader():
"""Writer reconnect path calls recreate_prisma_client. The routing wrapper
must recreate BOTH clients so a DB-wide event doesn't leave a stale reader."""
from litellm.proxy.db.routing_prisma_wrapper import RoutingPrismaWrapper
writer = MagicMock()
writer.recreate_prisma_client = AsyncMock()
reader = MagicMock()
reader.iam_token_db_auth = False
reader.recreate_prisma_client = AsyncMock()
routing = RoutingPrismaWrapper(writer=writer, reader=reader)
with patch.dict(os.environ, {"DATABASE_URL_READ_REPLICA": "reader-url"}):
await routing.recreate_prisma_client("writer-url", http_client=None)
writer.recreate_prisma_client.assert_awaited_once_with(
"writer-url", http_client=None
)
reader.recreate_prisma_client.assert_awaited_once_with(
"reader-url", http_client=None
)
assert routing.reader_unavailable is False
@pytest.mark.asyncio
async def test_recreate_recovers_reader_after_prior_degradation():
"""If a previous connect/recreate degraded the reader, a successful
recreate must clear the flag so reads start hitting the reader again."""
from litellm.proxy.db.routing_prisma_wrapper import RoutingPrismaWrapper
writer = MagicMock()
writer.recreate_prisma_client = AsyncMock()
reader = MagicMock()
reader.iam_token_db_auth = False
reader.recreate_prisma_client = AsyncMock()
routing = RoutingPrismaWrapper(writer=writer, reader=reader)
routing._reader_unavailable = True
with patch.dict(os.environ, {"DATABASE_URL_READ_REPLICA": "reader-url"}):
await routing.recreate_prisma_client("writer-url")
assert routing.reader_unavailable is False
@pytest.mark.asyncio
async def test_recreate_degrades_reader_if_reader_recreate_fails():
"""If the reader recreate fails, writer recreate still succeeds and the
routing wrapper degrades (does not raise)."""
from litellm.proxy.db.routing_prisma_wrapper import RoutingPrismaWrapper
writer = MagicMock()
writer.recreate_prisma_client = AsyncMock()
reader = MagicMock()
reader.iam_token_db_auth = False
reader.recreate_prisma_client = AsyncMock(
side_effect=RuntimeError("reader still down")
)
routing = RoutingPrismaWrapper(writer=writer, reader=reader)
with patch.dict(os.environ, {"DATABASE_URL_READ_REPLICA": "reader-url"}):
# Must not raise — writer was recreated, reader is best-effort.
await routing.recreate_prisma_client("writer-url")
writer.recreate_prisma_client.assert_awaited_once()
assert routing.reader_unavailable is True
@pytest.mark.asyncio
async def test_recreate_degrades_reader_when_replica_url_missing():
"""Non-IAM reader needs DATABASE_URL_READ_REPLICA. If it's missing
(configuration drift), the wrapper degrades instead of raising."""
from litellm.proxy.db.routing_prisma_wrapper import RoutingPrismaWrapper
writer = MagicMock()
writer.recreate_prisma_client = AsyncMock()
reader = MagicMock()
reader.iam_token_db_auth = False
reader.recreate_prisma_client = AsyncMock()
routing = RoutingPrismaWrapper(writer=writer, reader=reader)
# Ensure env var is absent.
with patch.dict(os.environ, {}, clear=False):
os.environ.pop("DATABASE_URL_READ_REPLICA", None)
await routing.recreate_prisma_client("writer-url")
writer.recreate_prisma_client.assert_awaited_once()
reader.recreate_prisma_client.assert_not_awaited()
assert routing.reader_unavailable is True
@pytest.mark.asyncio
async def test_recreate_iam_reader_refreshes_token():
"""IAM-enabled readers must refresh their token (reader has its own parsed
endpoint) and pass the fresh URL to recreate_prisma_client."""
from litellm.proxy.db.routing_prisma_wrapper import RoutingPrismaWrapper
writer = MagicMock()
writer.recreate_prisma_client = AsyncMock()
reader = MagicMock()
reader.iam_token_db_auth = True
reader.get_rds_iam_token = MagicMock(return_value="postgresql://u:fresh@h:5432/db")
reader.recreate_prisma_client = AsyncMock()
routing = RoutingPrismaWrapper(writer=writer, reader=reader)
await routing.recreate_prisma_client("writer-url")
reader.get_rds_iam_token.assert_called_once()
reader.recreate_prisma_client.assert_awaited_once_with(
"postgresql://u:fresh@h:5432/db", http_client=None
)
assert routing.reader_unavailable is False
@pytest.mark.asyncio
async def test_recreate_degrades_when_iam_token_generation_returns_none():
"""If `get_rds_iam_token` returns None (e.g. AWS-side failure), the wrapper
must degrade rather than crash — this exercises the explicit `raise
RuntimeError` inside `_recreate_reader`'s IAM branch."""
from litellm.proxy.db.routing_prisma_wrapper import RoutingPrismaWrapper
writer = MagicMock()
writer.recreate_prisma_client = AsyncMock()
reader = MagicMock()
reader.iam_token_db_auth = True
reader.get_rds_iam_token = MagicMock(return_value=None)
reader.recreate_prisma_client = AsyncMock()
routing = RoutingPrismaWrapper(writer=writer, reader=reader)
await routing.recreate_prisma_client("writer-url")
writer.recreate_prisma_client.assert_awaited_once()
reader.recreate_prisma_client.assert_not_awaited()
assert routing.reader_unavailable is True
def test_writer_and_reader_properties_expose_underlying_wrappers():
"""The `writer` and `reader` properties are used by PrismaClient.writer_db
to smoke-test the writer specifically during reconnect — they must return
the exact wrappers passed in."""
from litellm.proxy.db.routing_prisma_wrapper import RoutingPrismaWrapper
writer, _, reader, _ = _make_wrappers()
routing = RoutingPrismaWrapper(writer=writer, reader=reader)
assert routing.writer is writer
assert routing.reader is reader
def test_per_model_accessor_falls_back_when_reader_lacks_attr():
"""If the reader Prisma client somehow lacks a model accessor that the
writer has (older client / partial mock), the wrapper must fall back to
the writer accessor instead of raising AttributeError to the caller."""
from litellm.proxy.db.prisma_client import PrismaWrapper
from litellm.proxy.db.routing_prisma_wrapper import RoutingPrismaWrapper
# Plain class with only the accessor set on the writer side. Using a real
# class instead of MagicMock so attribute access raises AttributeError
# naturally instead of auto-creating mock attributes.
class _PartialPrisma:
pass
writer_inner = _PartialPrisma()
writer_inner.litellm_usertable = _model_actions_mock("writer_users")
reader_inner = _PartialPrisma() # deliberately missing litellm_usertable
writer = PrismaWrapper(original_prisma=writer_inner, iam_token_db_auth=False)
reader = PrismaWrapper(original_prisma=reader_inner, iam_token_db_auth=False)
routing = RoutingPrismaWrapper(writer=writer, reader=reader)
actions = routing.litellm_usertable
# Falls back to the writer's accessor verbatim — not a _RoutedActions wrapper.
assert actions is writer_inner.litellm_usertable
@pytest.mark.asyncio
async def test_writer_recreate_passes_http_client_through(monkeypatch):
"""When PrismaClient is constructed with an http_client, recreate must
forward it to the new Prisma() so connection settings persist across
reconnects."""
from litellm.proxy.db.prisma_client import PrismaWrapper
captured_kwargs: Dict[str, Any] = {}
class FakePrisma:
def __init__(self, **kwargs):
captured_kwargs.update(kwargs)
async def connect(self):
return None
fake_module = MagicMock()
fake_module.Prisma = FakePrisma
monkeypatch.setitem(sys.modules, "prisma", fake_module)
writer = PrismaWrapper(original_prisma=MagicMock(), iam_token_db_auth=False)
sentinel_http = object()
await writer.recreate_prisma_client(
"postgresql://u:p@h:5432/db", http_client=sentinel_http
)
assert captured_kwargs == {"http": sentinel_http}
# ---------------------------------------------------------------------------
# IAM endpoint parsing + reader IAM refresh
# ---------------------------------------------------------------------------
def test_parse_iam_endpoint_from_url_extracts_all_fields():
from litellm.proxy.db.prisma_client import parse_iam_endpoint_from_url
ep = parse_iam_endpoint_from_url(
"postgresql://litellm_user:initial-token@aurora-reader.example.com:6543/litellm?schema=public"
)
assert ep.host == "aurora-reader.example.com"
assert ep.port == "6543"
assert ep.user == "litellm_user"
assert ep.name == "litellm"
assert ep.schema == "public"
def test_parse_iam_endpoint_defaults_port_to_5432_and_skips_schema():
from litellm.proxy.db.prisma_client import parse_iam_endpoint_from_url
ep = parse_iam_endpoint_from_url("postgresql://u@host/dbname")
assert ep.host == "host"
assert ep.port == "5432"
assert ep.user == "u"
assert ep.name == "dbname"
assert ep.schema is None
def test_parse_iam_endpoint_rejects_url_without_user_or_dbname():
from litellm.proxy.db.prisma_client import parse_iam_endpoint_from_url
with pytest.raises(ValueError, match="missing host or username"):
parse_iam_endpoint_from_url("postgresql://host:5432/db")
with pytest.raises(ValueError, match="missing database name"):
parse_iam_endpoint_from_url("postgresql://u@host:5432/")
def test_iam_endpoint_build_url_inserts_token_verbatim():
from litellm.proxy.db.prisma_client import IAMEndpoint
# `generate_iam_auth_token` already URL-encodes the presigned token, so
# `build_url` must NOT encode again — double-encoding turned `%3D` into
# `%253D` and broke RDS auth on the reader path.
ep = IAMEndpoint(host="h", port="5432", user="u", name="db", schema="public")
pre_encoded_token = "token%2Fwith%3Fweird%26chars%3Dyes"
url = ep.build_url(pre_encoded_token)
assert url == f"postgresql://u:{pre_encoded_token}@h:5432/db?schema=public"
# Sanity check: no `%25` (the encoding of `%`), confirming we didn't re-encode.
assert "%25" not in url
@pytest.mark.asyncio
async def test_iam_refresh_logs_carry_log_prefix(caplog):
"""When `log_prefix` is set on a PrismaWrapper, every IAM-related log
line emitted by that wrapper must start with the prefix so writer and
reader can be told apart in interleaved output."""
from litellm.proxy.db.prisma_client import PrismaWrapper
wrapper = PrismaWrapper(
original_prisma=MagicMock(),
iam_token_db_auth=True,
log_prefix="[reader]",
)
with caplog.at_level(logging.INFO, logger="LiteLLM Proxy"):
await wrapper.start_token_refresh_task()
# Loop emits "RDS IAM token refresh loop started..." on first tick.
# Cancel immediately so the loop body runs once and we can assert.
await wrapper.stop_token_refresh_task()
messages = [r.getMessage() for r in caplog.records]
# Both start and stop notifications carry the prefix.
assert any(
m.startswith("[reader] Started RDS IAM token proactive refresh")
for m in messages
)
assert any(
m.startswith("[reader] Stopped RDS IAM token refresh background task")
for m in messages
)
def test_get_rds_iam_token_returns_none_when_iam_disabled():
"""`get_rds_iam_token` short-circuits to None when iam_token_db_auth is
False — covers the early-return guard at the top of the method."""
from litellm.proxy.db.prisma_client import PrismaWrapper
wrapper = PrismaWrapper(original_prisma=MagicMock(), iam_token_db_auth=False)
assert wrapper.get_rds_iam_token() is None
@pytest.mark.asyncio
async def test_getattr_does_not_block_inside_running_loop_on_expired_token(monkeypatch):
"""When `__getattr__` runs inside a running event loop and the IAM token
is expired, it MUST schedule the refresh as a background task and return
immediately. The previous `run_coroutine_threadsafe` + `future.result()`
pattern deadlocks the loop (loop thread blocks waiting for a coroutine
that needs the loop to run) and times out at 30s — exactly what was
breaking the reader on first query."""
from litellm.proxy.db.prisma_client import PrismaWrapper
# Stale URL — `is_token_expired` returns True because the password isn't
# a parseable IAM token, so we exercise the expired branch.
monkeypatch.setenv(
"DATABASE_URL_READ_REPLICA",
"postgresql://reader:placeholder@reader.aurora.local:5432/litellm",
)
inner = MagicMock()
inner.query_raw = MagicMock(name="query_raw_attr")
wrapper = PrismaWrapper(
original_prisma=inner,
iam_token_db_auth=True,
db_url_env_var="DATABASE_URL_READ_REPLICA",
)
# Replace the heavy refresh coroutine with a no-op AsyncMock so we can
# observe whether it was scheduled without actually doing the recreate.
refresh_calls = {"count": 0}
async def fake_refresh():
refresh_calls["count"] += 1
monkeypatch.setattr(wrapper, "_safe_refresh_token", fake_refresh)
# Direct attribute access from inside this async test runs __getattr__
# on the loop thread, exercising the in-loop branch. If the previous
# `run_coroutine_threadsafe` + `future.result()` pattern were back, this
# line would deadlock the loop and the test would hang (and pytest's
# per-test timeout would catch it).
attr = wrapper.query_raw
# Yield once so the scheduled refresh task gets a chance to run.
await asyncio.sleep(0)
assert attr is inner.query_raw
assert refresh_calls["count"] == 1
def test_writer_get_rds_iam_token_defaults_port_when_unset(monkeypatch):
"""When DATABASE_PORT is unset, the writer must default to the Postgres
standard port instead of passing `None` through. Passing None to
`generate_iam_auth_token` makes botocore embed the literal string
\"None\" in the presigned URL during signing and crashes with
`ValueError: Port could not be cast to integer value as 'None'`."""
from litellm.proxy.db.prisma_client import PrismaWrapper
monkeypatch.setenv("DATABASE_HOST", "writer.aurora.local")
monkeypatch.delenv("DATABASE_PORT", raising=False)
monkeypatch.setenv("DATABASE_USER", "litellm")
monkeypatch.setenv("DATABASE_NAME", "litellm")
monkeypatch.delenv("DATABASE_SCHEMA", raising=False)
monkeypatch.delenv("DATABASE_URL", raising=False)
captured: Dict[str, Any] = {}
def fake_generate(db_host=None, db_port=None, db_user=None):
captured["port"] = db_port
return "TOKEN"
fake_module = MagicMock()
fake_module.generate_iam_auth_token = fake_generate
monkeypatch.setitem(sys.modules, "litellm.proxy.auth.rds_iam_token", fake_module)
writer = PrismaWrapper(
original_prisma=MagicMock(),
iam_token_db_auth=True,
)
new_url = writer.get_rds_iam_token()
assert captured["port"] == "5432" # default applied, NOT None
assert ":5432/litellm" in (new_url or "")
def test_writer_get_rds_iam_token_uses_database_host_env_vars(monkeypatch):
"""Writer's IAM path (no iam_endpoint configured) reads host/port/user/db
from the legacy DATABASE_HOST/PORT/USER/NAME env vars and writes the URL
back to DATABASE_URL — this is the pre-read-replica behavior the patch
must preserve."""
from litellm.proxy.db.prisma_client import PrismaWrapper
monkeypatch.setenv("DATABASE_HOST", "writer.aurora.local")
monkeypatch.setenv("DATABASE_PORT", "5432")
monkeypatch.setenv("DATABASE_USER", "litellm")
monkeypatch.setenv("DATABASE_NAME", "litellm")
monkeypatch.setenv("DATABASE_SCHEMA", "public")
monkeypatch.delenv("DATABASE_URL", raising=False)
captured: Dict[str, Any] = {}
def fake_generate(db_host=None, db_port=None, db_user=None):
captured["host"] = db_host
captured["port"] = db_port
captured["user"] = db_user
return "WRITER-TOKEN"
fake_module = MagicMock()
fake_module.generate_iam_auth_token = fake_generate
monkeypatch.setitem(sys.modules, "litellm.proxy.auth.rds_iam_token", fake_module)
writer = PrismaWrapper(
original_prisma=MagicMock(),
iam_token_db_auth=True,
# No iam_endpoint → legacy DATABASE_HOST/etc. path.
)
new_url = writer.get_rds_iam_token()
assert captured == {
"host": "writer.aurora.local",
"port": "5432",
"user": "litellm",
}
assert new_url == (
"postgresql://litellm:WRITER-TOKEN@writer.aurora.local:5432/litellm?schema=public"
)
# Writer updates its own env var (DATABASE_URL by default), not the reader's.
assert os.environ["DATABASE_URL"] == new_url
def test_reader_iam_refresh_uses_parsed_endpoint(monkeypatch):
"""The reader generates fresh tokens against its parsed endpoint and
writes the new URL to DATABASE_URL_READ_REPLICA — not DATABASE_URL."""
from litellm.proxy.db.prisma_client import IAMEndpoint, PrismaWrapper
# Pre-seed env vars so we can prove the reader does NOT touch DATABASE_URL.
monkeypatch.setenv("DATABASE_URL", "writer-url-untouched")
monkeypatch.setenv("DATABASE_URL_READ_REPLICA", "stale-reader-url")
captured: Dict[str, Any] = {}
def fake_generate(db_host=None, db_port=None, db_user=None):
captured["host"] = db_host
captured["port"] = db_port
captured["user"] = db_user
return "FRESH-TOKEN"
fake_module = MagicMock()
fake_module.generate_iam_auth_token = fake_generate
monkeypatch.setitem(sys.modules, "litellm.proxy.auth.rds_iam_token", fake_module)
endpoint = IAMEndpoint(
host="reader.aurora.local",
port="5432",
user="lit",
name="litellm",
schema=None,
)
reader = PrismaWrapper(
original_prisma=MagicMock(),
iam_token_db_auth=True,
db_url_env_var="DATABASE_URL_READ_REPLICA",
iam_endpoint=endpoint,
recreate_uses_datasource=True,
)
new_url = reader.get_rds_iam_token()
# IAM token generator was called with the reader's parsed endpoint, not
# the writer's DATABASE_HOST/PORT/USER env vars.
assert captured == {
"host": "reader.aurora.local",
"port": "5432",
"user": "lit",
}
assert new_url is not None
assert new_url.startswith(
"postgresql://lit:FRESH-TOKEN@reader.aurora.local:5432/litellm"
)
# The reader updates its OWN env var; writer's DATABASE_URL is left alone.
assert os.environ["DATABASE_URL_READ_REPLICA"] == new_url
assert os.environ["DATABASE_URL"] == "writer-url-untouched"
@pytest.mark.asyncio
async def test_reader_recreate_uses_datasource_override(monkeypatch):
"""Reader recreate must pass `datasource={"url": ...}` to Prisma() — Prisma
only auto-reads DATABASE_URL, so without the override the new reader URL
would be silently ignored."""
from litellm.proxy.db.prisma_client import IAMEndpoint, PrismaWrapper
captured_kwargs: Dict[str, Any] = {}
class FakePrisma:
def __init__(self, **kwargs):
captured_kwargs.update(kwargs)
async def connect(self):
return None
fake_module = MagicMock()
fake_module.Prisma = FakePrisma
monkeypatch.setitem(sys.modules, "prisma", fake_module)
reader = PrismaWrapper(
original_prisma=MagicMock(),
iam_token_db_auth=True,
db_url_env_var="DATABASE_URL_READ_REPLICA",
iam_endpoint=IAMEndpoint(host="h", port="5432", user="u", name="db"),
recreate_uses_datasource=True,
)
await reader.recreate_prisma_client(
"postgresql://u:newtoken@h:5432/db", http_client=None
)
assert captured_kwargs == {
"datasource": {"url": "postgresql://u:newtoken@h:5432/db"}
}
@pytest.mark.asyncio
async def test_writer_recreate_does_not_use_datasource(monkeypatch):
"""Writer keeps relying on Prisma reading DATABASE_URL from env — datasource
override must NOT leak into the writer path (would override the freshly
rotated env var)."""
from litellm.proxy.db.prisma_client import PrismaWrapper
captured_kwargs: Dict[str, Any] = {}
class FakePrisma:
def __init__(self, **kwargs):
captured_kwargs.update(kwargs)
async def connect(self):
return None
fake_module = MagicMock()
fake_module.Prisma = FakePrisma
monkeypatch.setitem(sys.modules, "prisma", fake_module)
writer = PrismaWrapper(
original_prisma=MagicMock(),
iam_token_db_auth=True,
)
await writer.recreate_prisma_client(
"postgresql://u:newtoken@h:5432/db", http_client=None
)
assert "datasource" not in captured_kwargs
def test_prisma_client_init_falls_back_to_writer_when_reader_iam_token_fails(
monkeypatch, caplog
):
"""A transient AWS STS error (or any other failure) during the reader
IAM token mint must NOT abort proxy startup. The reader is opt-in, so
`PrismaClient.__init__` should log a warning and fall back to the
writer-only `PrismaWrapper`. The runtime contract in
`RoutingPrismaWrapper.connect` already says reader-side failures are
non-fatal — but that code never runs if construction throws first."""
from litellm.proxy.db.prisma_client import PrismaWrapper
from litellm.proxy.db.routing_prisma_wrapper import RoutingPrismaWrapper
monkeypatch.setenv("IAM_TOKEN_DB_AUTH", "true")
monkeypatch.setenv(
"DATABASE_URL_READ_REPLICA",
"postgresql://reader_user@reader.aurora.local:5432/litellm",
)
class FakePrisma:
def __init__(self, **kwargs):
self.kwargs = kwargs
async def connect(self):
return None
fake_prisma_module = MagicMock()
fake_prisma_module.Prisma = FakePrisma
monkeypatch.setitem(sys.modules, "prisma", fake_prisma_module)
fake_iam_module = MagicMock()
def boom(**_kwargs):
raise RuntimeError("simulated AWS STS hiccup")
fake_iam_module.generate_iam_auth_token = boom
monkeypatch.setitem(
sys.modules, "litellm.proxy.auth.rds_iam_token", fake_iam_module
)
from litellm.proxy.utils import PrismaClient
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
client = PrismaClient(
database_url="postgresql://writer@writer.aurora.local:5432/litellm",
proxy_logging_obj=MagicMock(),
)
# Construction did not raise, and the proxy is in writer-only mode —
# NOT a RoutingPrismaWrapper, so reads will go to the writer.
assert isinstance(client.db, PrismaWrapper)
assert not isinstance(client.db, RoutingPrismaWrapper)
# And the operator gets a clear warning.
assert any(
"Failed to initialize read replica Prisma client" in r.getMessage()
for r in caplog.records
)

View file

@ -8,7 +8,9 @@ from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation
from litellm.llms.base_llm.guardrail_translation.utils import (
effective_skip_system_message_for_guardrail,
effective_skip_tool_message_for_guardrail,
openai_messages_without_system,
openai_messages_without_tool,
)
from litellm.llms.openai.chat.guardrail_translation.handler import (
OpenAIChatCompletionsHandler,
@ -180,6 +182,136 @@ class TestUnifiedLLMGuardrails:
}
assert "system" in roles
class TestSkipToolMessageForChatCompletions:
def test_openai_messages_without_tool(self):
msgs = [
{"role": "user", "content": "hi"},
{
"role": "assistant",
"content": None,
"tool_calls": [
{
"id": "call_1",
"type": "function",
"function": {"name": "f", "arguments": "{}"},
}
],
},
{"role": "tool", "content": "tool result", "tool_call_id": "call_1"},
]
out = openai_messages_without_tool(msgs)
assert len(out) == 2
assert all(m["role"] != "tool" for m in out)
assert msgs[2]["content"] == "tool result"
def test_effective_skip_tool_respects_per_guardrail_over_global(
self, monkeypatch
):
monkeypatch.setattr(
litellm, "skip_tool_message_in_guardrail", True, raising=False
)
class G:
skip_tool_message_in_guardrail = False
assert effective_skip_tool_message_for_guardrail(G()) is False
class G2:
skip_tool_message_in_guardrail = None
assert effective_skip_tool_message_for_guardrail(G2()) is True
@pytest.mark.asyncio
async def test_openai_handler_skips_tool_in_guardrail_inputs(self, monkeypatch):
monkeypatch.setattr(
litellm, "skip_tool_message_in_guardrail", True, raising=False
)
captured = {}
class MockGuardrail:
skip_tool_message_in_guardrail = None
async def apply_guardrail(
self, inputs, request_data, input_type, logging_obj=None
):
captured["inputs"] = inputs
return inputs
data = {
"messages": [
{"role": "user", "content": "hello"},
{
"role": "assistant",
"content": None,
"tool_calls": [
{
"id": "call_1",
"type": "function",
"function": {"name": "f", "arguments": "{}"},
}
],
},
{
"role": "tool",
"content": "secret tool result",
"tool_call_id": "call_1",
},
],
"model": "gpt-4o",
}
handler = OpenAIChatCompletionsHandler()
await handler.process_input_messages(
data=data,
guardrail_to_apply=MockGuardrail(),
litellm_logging_obj=None,
)
assert "secret tool result" not in captured["inputs"]["texts"]
sm = captured["inputs"].get("structured_messages") or []
assert all(m.get("role") != "tool" for m in sm)
assert data["messages"][2]["content"] == "secret tool result"
@pytest.mark.asyncio
async def test_openai_handler_per_guardrail_skip_tool_false_overrides_global(
self, monkeypatch
):
monkeypatch.setattr(
litellm, "skip_tool_message_in_guardrail", True, raising=False
)
captured = {}
class MockGuardrail:
skip_tool_message_in_guardrail = False
async def apply_guardrail(
self, inputs, request_data, input_type, logging_obj=None
):
captured["inputs"] = inputs
return inputs
data = {
"messages": [
{"role": "user", "content": "u"},
{"role": "tool", "content": "tr", "tool_call_id": "call_1"},
],
}
await OpenAIChatCompletionsHandler().process_input_messages(
data=data,
guardrail_to_apply=MockGuardrail(),
litellm_logging_obj=None,
)
assert "tr" in captured["inputs"]["texts"]
roles = {
m.get("role")
for m in (captured["inputs"].get("structured_messages") or [])
}
assert "tool" in roles
class TestAsyncPreCallHook:
@pytest.mark.asyncio
async def test_uses_mcp_event_type(self):

View file

@ -6129,3 +6129,163 @@ class TestPKCEStateCookieBinding:
# State-cookie check passed, so the function got past the early
# ProxyException raise and produced an SSO result object.
assert result is not None
@pytest.mark.asyncio
async def test_debug_sso_callback_renders_full_jwt_claims():
"""
/sso/debug/callback should render the complete set of claims returned by the
IdP — both the raw userinfo response and the decoded access-token JWT — in
addition to the proxy-parsed OpenID fields. Bearer tokens must be stripped
even if a non-conforming IdP places them in its userinfo response.
"""
from litellm.proxy.management_endpoints.ui_sso import debug_sso_callback
mock_request = MagicMock(spec=Request)
mock_request.base_url = "http://proxy.example.com/"
mock_request.cookies = {}
mock_request.query_params = {}
parsed_openid = CustomOpenID(
id="user_123",
email="philip@example.com",
first_name="Philip",
last_name="Schwartz",
display_name="Philip Schwartz",
provider="generic",
team_ids=["ord-engineering-high"],
user_role=None,
)
raw_userinfo_with_leaked_token = {
"sub": "user_123",
"email": "philip@example.com",
"team_id": "ord-engineering-high",
"team_alias": "ord-engineering-high",
"teams": ["ord-engineering-high"],
"roles": ["litellm.api.user"],
# Defense-in-depth: a non-conforming IdP could shove a bearer token
# into userinfo. The debug endpoint must strip it before rendering.
"access_token": "should-not-render",
"id_token": "should-not-render-either",
}
access_token_payload = {
"sub": "user_123",
"scope": "openid profile email",
"groups": ["litellm-users"],
}
async def fake_get_generic_sso_response(**kwargs):
return parsed_openid, raw_userinfo_with_leaked_token, access_token_payload
with (
patch.dict(
os.environ,
{"GENERIC_CLIENT_ID": "test_client_id"},
clear=False,
),
patch(
"litellm.proxy.management_endpoints.ui_sso.get_generic_sso_response",
side_effect=fake_get_generic_sso_response,
),
patch("litellm.proxy.proxy_server.general_settings", {}),
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()),
patch("litellm.proxy.proxy_server.jwt_handler", MagicMock(spec=JWTHandler)),
):
# Microsoft / Google envs may leak in from other tests — ensure only
# the generic path runs.
for var in ("MICROSOFT_CLIENT_ID", "GOOGLE_CLIENT_ID"):
os.environ.pop(var, None)
response = await debug_sso_callback(mock_request)
body = response.body.decode()
# The embedded JSON payload drives the rendered page. Extract and parse it
# so we can assert on shape, not on cosmetic HTML details.
marker = "const ssoData = "
start = body.index(marker) + len(marker)
end = body.index(";", start)
while body[end - 1] not in "}]": # handle ';' inside string values
end = body.index(";", end + 1)
payload = json.loads(body[start:end])
assert set(payload.keys()) == {
"parsed_by_proxy",
"raw_claims",
"access_token_claims",
}
# Parsed OpenID fields are shown
assert payload["parsed_by_proxy"]["email"] == "philip@example.com"
assert payload["parsed_by_proxy"]["id"] == "user_123"
# Raw IdP claims surface fields the OpenID model drops (the original LIT-2838 ask)
assert payload["raw_claims"]["team_id"] == "ord-engineering-high"
assert payload["raw_claims"]["team_alias"] == "ord-engineering-high"
assert payload["raw_claims"]["teams"] == ["ord-engineering-high"]
assert payload["raw_claims"]["roles"] == ["litellm.api.user"]
# Defense-in-depth: bearer tokens must never appear in the rendered HTML
assert "access_token" not in payload["raw_claims"]
assert "id_token" not in payload["raw_claims"]
assert "should-not-render" not in body
# Decoded access-token JWT claims are surfaced
assert payload["access_token_claims"]["groups"] == ["litellm-users"]
@pytest.mark.asyncio
async def test_debug_sso_callback_handles_missing_raw_response():
"""
Microsoft and Google paths don't return a raw response or access-token
payload. The debug endpoint must still render successfully with empty
sections instead of crashing.
"""
from litellm.proxy.management_endpoints.ui_sso import debug_sso_callback
mock_request = MagicMock(spec=Request)
mock_request.base_url = "http://proxy.example.com/"
mock_request.cookies = {}
mock_request.query_params = {}
parsed_openid = CustomOpenID(
id="user_456",
email="user@example.com",
first_name="Some",
last_name="User",
display_name="Some User",
provider="microsoft",
team_ids=[],
user_role=None,
)
async def fake_microsoft_callback(**kwargs):
return parsed_openid
with (
patch.dict(
os.environ,
{"MICROSOFT_CLIENT_ID": "test_microsoft_id"},
clear=False,
),
patch.object(
MicrosoftSSOHandler,
"get_microsoft_callback_response",
side_effect=fake_microsoft_callback,
),
patch("litellm.proxy.proxy_server.general_settings", {}),
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()),
patch("litellm.proxy.proxy_server.jwt_handler", MagicMock(spec=JWTHandler)),
):
for var in ("GENERIC_CLIENT_ID", "GOOGLE_CLIENT_ID"):
os.environ.pop(var, None)
response = await debug_sso_callback(mock_request)
assert response.status_code == 200
body = response.body.decode()
assert '"raw_claims": {}' in body
assert '"access_token_claims": {}' in body
assert "user@example.com" in body

View file

@ -121,6 +121,26 @@ def test_invalid_auth_metrics(app_with_middleware, monkeypatch):
assert "Unauthorized access to metrics endpoint" in response.text
def test_invalid_auth_metrics_includes_optout_hint(app_with_middleware, monkeypatch):
"""
The 401 body must tell operators how to restore the previous unauthenticated
behavior, otherwise a Prometheus scraper that worked pre-upgrade just sees
"Malformed API Key" with no actionable migration path.
"""
monkeypatch.setattr(litellm, "require_auth_for_metrics_endpoint", True)
monkeypatch.setattr(
"litellm.proxy.middleware.prometheus_auth_middleware.user_api_key_auth",
fake_invalid_auth,
)
client = TestClient(app_with_middleware)
response = client.get("/metrics")
assert response.status_code == 401, response.text
assert "require_auth_for_metrics_endpoint" in response.text
assert "false" in response.text
def test_metrics_auth_uses_real_auth_when_route_is_public(
app_with_middleware, monkeypatch
):

View file

@ -585,15 +585,24 @@ async def test_should_cap_known_estimate_to_remaining_budget(
@pytest.mark.asyncio
async def test_should_reserve_remaining_budget_when_output_cap_missing(
async def test_should_clamp_reservation_to_default_when_output_cap_missing(
spend_counter_state,
):
"""When max_tokens is not specified, _estimate_output_tokens falls back to
DEFAULT_MAX_OUTPUT_TOKENS_FALLBACK (16K), clamped by the model's
max_output_tokens. Reservation must be a bounded per-request amount
(mirroring parallel_request_limiter_v3's DEFAULT_MAX_TOKENS_ESTIMATE),
not the entire remaining headroom."""
from litellm.proxy.spend_tracking.budget_reservation import (
DEFAULT_MAX_OUTPUT_TOKENS_FALLBACK,
)
counter_cache, key_cache = spend_counter_state
proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache)
valid_token = UserAPIKeyAuth(
token="key-budget-uncapped",
spend=0.2,
max_budget=1.0,
max_budget=10000.0,
)
await key_cache.async_set_cache(
key="key-budget-uncapped",
@ -602,22 +611,24 @@ async def test_should_reserve_remaining_budget_when_output_cap_missing(
request_body = _request_body()
request_body.pop("max_tokens")
output_cost_per_token = 1e-5 # roughly Opus 4.5/4.7 output rate
expected_cost = DEFAULT_MAX_OUTPUT_TOKENS_FALLBACK * output_cost_per_token
with patch(
"litellm.proxy.spend_tracking.budget_reservation._get_model_cost_info",
return_value={
"input_cost_per_token": 0.0,
"output_cost_per_token": 100.0,
"max_output_tokens": 200000,
"output_cost_per_token": output_cost_per_token,
"max_output_tokens": 200000, # well above the 16K fallback
},
):
assert (
estimate_request_max_cost(
request_body=request_body,
route="/chat/completions",
llm_router=None,
)
is None
estimated = estimate_request_max_cost(
request_body=request_body,
route="/chat/completions",
llm_router=None,
)
assert estimated == pytest.approx(expected_cost)
reservation = await reserve_budget_for_request(
request_body=request_body,
route="/chat/completions",
@ -631,47 +642,45 @@ async def test_should_reserve_remaining_budget_when_output_cap_missing(
)
assert reservation is not None
assert reservation["reserved_cost"] == pytest.approx(0.8)
assert counter_cache.in_memory_cache.get_cache(
key="spend:key:key-budget-uncapped"
) == pytest.approx(1.0)
assert reservation["reserved_cost"] == pytest.approx(expected_cost)
await release_budget_reservation(reservation)
@pytest.mark.asyncio
async def test_should_shrink_uncapped_reservation_when_counter_advances(
async def test_should_clamp_reservation_to_model_ceiling_when_caller_overrequests(
spend_counter_state,
monkeypatch,
):
"""An adversarial caller sending max_tokens=999_999_999 must not be able
to inflate the per-request reservation up to the entire remaining team
headroom. _estimate_output_tokens clamps the explicit value at the
model's max_output_tokens — the model can only physically emit that
many tokens anyway, so anything more is both wasteful and a DoS surface."""
counter_cache, key_cache = spend_counter_state
proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache)
valid_token = UserAPIKeyAuth(
token="key-budget-uncapped-race",
spend=0.2,
max_budget=1.0,
token="key-budget-overrequest",
spend=0.0,
max_budget=10000.0,
)
await key_cache.async_set_cache(
key="key-budget-overrequest",
value=valid_token,
)
request_body = _request_body()
request_body.pop("max_tokens")
request_body["max_tokens"] = 999_999_999
from litellm.proxy.spend_tracking import budget_reservation
async def stale_counter_read(counter):
await counter_cache.async_increment_cache(
key=counter.counter_key,
value=0.3,
)
return 0.2
monkeypatch.setattr(
budget_reservation,
"_get_current_counter_value",
stale_counter_read,
)
output_cost_per_token = 1e-5
model_ceiling = 128_000
expected_cost = model_ceiling * output_cost_per_token
with patch(
"litellm.proxy.spend_tracking.budget_reservation.estimate_request_max_cost",
return_value=None,
"litellm.proxy.spend_tracking.budget_reservation._get_model_cost_info",
return_value={
"input_cost_per_token": 0.0,
"output_cost_per_token": output_cost_per_token,
"max_output_tokens": model_ceiling,
},
):
reservation = await reserve_budget_for_request(
request_body=request_body,
@ -686,66 +695,91 @@ async def test_should_shrink_uncapped_reservation_when_counter_advances(
)
assert reservation is not None
assert reservation["reserved_cost"] == pytest.approx(0.7)
assert counter_cache.in_memory_cache.get_cache(
key="spend:key:key-budget-uncapped-race"
) == pytest.approx(1.0)
assert reservation["reserved_cost"] == pytest.approx(expected_cost)
await release_budget_reservation(reservation)
assert counter_cache.in_memory_cache.get_cache(
key="spend:key:key-budget-uncapped-race"
) == pytest.approx(0.3)
@pytest.mark.asyncio
async def test_should_shrink_uncapped_reservation_multiple_times(
async def test_should_reserve_image_generation_cost_per_image(
spend_counter_state,
monkeypatch,
):
"""Image-generation requests reserve `n × per-image cost` so concurrent
requests against a depleted budget cannot all bypass the admission gate.
The OpenAI ``dall-e-3`` entry exposes the per-image price as
``input_cost_per_image`` (a naming quirk), while other providers use
``output_cost_per_image`` — both must be honored."""
counter_cache, key_cache = spend_counter_state
proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache)
valid_token = UserAPIKeyAuth(
token="key-budget-double-resize",
spend=0.2,
max_budget=1.0,
team_id="team-budget-double-resize",
token="key-image-gen",
spend=0.0,
max_budget=10.0,
)
await key_cache.async_set_cache(key="key-image-gen", value=valid_token)
request_body = {"model": "dall-e-3", "prompt": "a cat", "n": 3}
with patch(
"litellm.proxy.spend_tracking.budget_reservation._get_model_cost_info",
return_value={
"mode": "image_generation",
"input_cost_per_image": 0.04,
},
):
reservation = await reserve_budget_for_request(
request_body=request_body,
route="/v1/images/generations",
llm_router=None,
valid_token=valid_token,
team_object=None,
user_object=None,
prisma_client=None,
user_api_key_cache=key_cache,
proxy_logging_obj=proxy_logging_obj,
)
assert reservation is not None
assert reservation["reserved_cost"] == pytest.approx(0.12) # 3 × $0.04
await release_budget_reservation(reservation)
@pytest.mark.asyncio
async def test_should_reject_concurrent_image_request_against_depleted_budget(
spend_counter_state,
):
"""Greptile P1 regression: with image-gen reservation in place, a second
concurrent image request against a budget already pinned at the cap by
the first reservation must raise BudgetExceededError instead of
silently reaching the provider."""
counter_cache, key_cache = spend_counter_state
proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache)
valid_token = UserAPIKeyAuth(
token="key-image-deplete",
spend=0.0,
team_id="team-image-deplete",
)
team_object = LiteLLM_TeamTable(
team_id="team-budget-double-resize",
spend=0.2,
max_budget=1.0,
team_id="team-image-deplete",
max_budget=0.04,
spend=0.0,
)
request_body = _request_body()
request_body.pop("max_tokens")
from litellm.proxy.spend_tracking import budget_reservation
stale_spend_by_counter_key = {
"spend:key:key-budget-double-resize": 0.3,
"spend:team:team-budget-double-resize": 0.4,
}
async def stale_counter_read(counter):
await counter_cache.async_increment_cache(
key=counter.counter_key,
value=stale_spend_by_counter_key[counter.counter_key],
)
return 0.2
monkeypatch.setattr(
budget_reservation,
"_get_current_counter_value",
stale_counter_read,
await key_cache.async_set_cache(
key=f"team_id:{team_object.team_id}",
value=team_object,
)
request_body = {"model": "dall-e-3", "prompt": "a cat"}
with patch(
"litellm.proxy.spend_tracking.budget_reservation.estimate_request_max_cost",
return_value=None,
"litellm.proxy.spend_tracking.budget_reservation._get_model_cost_info",
return_value={
"mode": "image_generation",
"input_cost_per_image": 0.04,
},
):
reservation = await reserve_budget_for_request(
first = await reserve_budget_for_request(
request_body=request_body,
route="/chat/completions",
route="/v1/images/generations",
llm_router=None,
valid_token=valid_token,
team_object=team_object,
@ -754,32 +788,167 @@ async def test_should_shrink_uncapped_reservation_multiple_times(
user_api_key_cache=key_cache,
proxy_logging_obj=proxy_logging_obj,
)
assert first is not None
with pytest.raises(litellm.BudgetExceededError):
await reserve_budget_for_request(
request_body=request_body,
route="/v1/images/generations",
llm_router=None,
valid_token=valid_token,
team_object=team_object,
user_object=None,
prisma_client=None,
user_api_key_cache=key_cache,
proxy_logging_obj=proxy_logging_obj,
)
await release_budget_reservation(first)
@pytest.mark.asyncio
async def test_should_skip_reservation_for_per_pixel_image_model(
spend_counter_state,
):
"""DALL-E 2-style per-pixel pricing depends on the requested ``size``,
which we don't decode here. Fall through to read-time enforcement
rather than guess."""
counter_cache, key_cache = spend_counter_state
proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache)
valid_token = UserAPIKeyAuth(
token="key-image-per-pixel",
spend=0.0,
max_budget=1.0,
)
await key_cache.async_set_cache(key="key-image-per-pixel", value=valid_token)
request_body = {"model": "dall-e-2", "prompt": "a cat", "size": "256x256"}
with patch(
"litellm.proxy.spend_tracking.budget_reservation._get_model_cost_info",
return_value={
"mode": "image_generation",
"input_cost_per_pixel": 2.4414e-07,
"output_cost_per_pixel": 0.0,
},
):
reservation = await reserve_budget_for_request(
request_body=request_body,
route="/v1/images/generations",
llm_router=None,
valid_token=valid_token,
team_object=None,
user_object=None,
prisma_client=None,
user_api_key_cache=key_cache,
proxy_logging_obj=proxy_logging_obj,
)
assert reservation is None
@pytest.mark.asyncio
async def test_should_use_token_pricing_for_chat_model_with_image_cost_field(
spend_counter_state,
):
"""Several chat and embedding models carry ``input_cost_per_image`` /
``output_cost_per_image`` to price multimodal vision *input*, not image
generation (e.g. gemini-3.1-pro-preview, azure/gpt-realtime-*,
amazon.titan-embed-image-v1). _estimate_image_generation_cost must gate
on ``mode`` so these models still go through the token-priced path —
otherwise a long chat reserves a fraction of a cent instead of the true
token cost."""
counter_cache, key_cache = spend_counter_state
proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache)
valid_token = UserAPIKeyAuth(
token="key-multimodal-chat",
spend=0.0,
max_budget=10.0,
)
await key_cache.async_set_cache(key="key-multimodal-chat", value=valid_token)
# Roughly the gemini-3.1-pro-preview shape: chat-mode model that
# carries an output_cost_per_image alongside token pricing.
output_cost_per_token = 1.2e-5
request_body = {
"model": "gemini-3.1-pro-preview",
"messages": [{"role": "user", "content": "hello"}],
"max_tokens": 1000,
}
expected_cost = 1000 * output_cost_per_token # token-priced path, not 1 × $0.00012
with patch(
"litellm.proxy.spend_tracking.budget_reservation._get_model_cost_info",
return_value={
"mode": "chat",
"input_cost_per_token": 2e-6,
"output_cost_per_token": output_cost_per_token,
"output_cost_per_image": 0.00012,
"max_output_tokens": 64000,
},
):
reservation = await reserve_budget_for_request(
request_body=request_body,
route="/chat/completions",
llm_router=None,
valid_token=valid_token,
team_object=None,
user_object=None,
prisma_client=None,
user_api_key_cache=key_cache,
proxy_logging_obj=proxy_logging_obj,
)
assert reservation is not None
assert reservation["reserved_cost"] == pytest.approx(0.6)
assert [entry["reserved_cost"] for entry in reservation["entries"]] == [
pytest.approx(0.6),
pytest.approx(0.6),
]
assert [entry["applied_adjustment"] for entry in reservation["entries"]] == [
pytest.approx(0.0),
pytest.approx(0.0),
]
assert counter_cache.in_memory_cache.get_cache(
key="spend:key:key-budget-double-resize"
) == pytest.approx(0.9)
assert counter_cache.in_memory_cache.get_cache(
key="spend:team:team-budget-double-resize"
) == pytest.approx(1.0)
# Token-priced path: reservation ≈ output_tokens × output_cost_per_token,
# plus a small input-token contribution. Must NOT collapse to the
# per-image price ($0.00012) which would indicate the image-gen branch
# incorrectly fired for this chat model.
assert reservation["reserved_cost"] == pytest.approx(expected_cost, rel=0.05)
assert reservation["reserved_cost"] > 0.001 # well above per-image price
await release_budget_reservation(reservation)
assert counter_cache.in_memory_cache.get_cache(
key="spend:key:key-budget-double-resize"
) == pytest.approx(0.3)
assert counter_cache.in_memory_cache.get_cache(
key="spend:team:team-budget-double-resize"
) == pytest.approx(0.4)
@pytest.mark.asyncio
async def test_should_reserve_image_edit_cost_per_image(
spend_counter_state,
):
"""``image_edit`` models (Flux Kontext, Stability inpaint/outpaint, etc.)
bill per generated image just like ``image_generation`` and must get
the same atomic per-image reservation."""
counter_cache, key_cache = spend_counter_state
proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache)
valid_token = UserAPIKeyAuth(
token="key-image-edit",
spend=0.0,
max_budget=10.0,
)
await key_cache.async_set_cache(key="key-image-edit", value=valid_token)
request_body = {"model": "stability/inpaint", "prompt": "a cat", "n": 2}
with patch(
"litellm.proxy.spend_tracking.budget_reservation._get_model_cost_info",
return_value={
"mode": "image_edit",
"output_cost_per_image": 0.05,
},
):
reservation = await reserve_budget_for_request(
request_body=request_body,
route="/v1/images/edits",
llm_router=None,
valid_token=valid_token,
team_object=None,
user_object=None,
prisma_client=None,
user_api_key_cache=key_cache,
proxy_logging_obj=proxy_logging_obj,
)
assert reservation is not None
assert reservation["reserved_cost"] == pytest.approx(0.10) # 2 × $0.05
await release_budget_reservation(reservation)
def test_should_start_window_without_reset_at_at_duration_boundary():
@ -1047,62 +1216,6 @@ async def test_should_release_tracked_entry_when_reservation_fails_after_increme
) == pytest.approx(0.0)
@pytest.mark.asyncio
async def test_should_not_re_read_uncapped_budget_after_reservation_fallback(
spend_counter_state,
monkeypatch,
):
_, key_cache = spend_counter_state
proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache)
valid_token = UserAPIKeyAuth(
token="key-budget-uncapped-read-once",
spend=0.2,
max_budget=1.0,
)
from litellm.proxy.spend_tracking import budget_reservation
current_counter_reads = []
async def mock_get_current_counter_value(counter):
current_counter_reads.append(counter.counter_key)
return counter.fallback_spend
async def mock_reserve_counter(counter, reservation_cost):
return None
monkeypatch.setattr(
budget_reservation,
"_get_current_counter_value",
mock_get_current_counter_value,
)
monkeypatch.setattr(
budget_reservation,
"_reserve_counter",
mock_reserve_counter,
)
with patch(
"litellm.proxy.spend_tracking.budget_reservation.estimate_request_max_cost",
return_value=None,
):
reservation = await reserve_budget_for_request(
request_body=_request_body(),
route="/chat/completions",
llm_router=None,
valid_token=valid_token,
team_object=None,
user_object=None,
prisma_client=None,
user_api_key_cache=key_cache,
proxy_logging_obj=proxy_logging_obj,
)
assert reservation is not None
assert reservation["reserved_cost"] == pytest.approx(0.8)
assert current_counter_reads == ["spend:key:key-budget-uncapped-read-once"]
@pytest.mark.asyncio
async def test_should_reconcile_reserved_counter_to_actual_spend(
spend_counter_state,
@ -1492,4 +1605,94 @@ async def test_should_reserve_all_budgeted_counters(spend_counter_state):
counter_cache.in_memory_cache.get_cache(key="spend:team:team-budget-all") == 0.3
)
await release_budget_reservation(reservation)
@pytest.mark.asyncio
async def test_should_not_block_concurrent_team_request_when_first_request_lacks_max_tokens(
spend_counter_state,
):
"""
Regression test: a team-bound request with no max_tokens must not pin the
team's spend counter at max_budget for the duration of the request.
Repro of the integration-test team being falsely budget-blocked at the
$2000 cap while DB spend is $0.144: the first request without max_tokens
used to reserve the entire remaining headroom, leaving any subsequent
request stuck behind a counter sitting at the cap until the success
callback finished reconciling.
"""
counter_cache, key_cache = spend_counter_state
proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache)
valid_token = UserAPIKeyAuth(
token="key-team-integration-tests",
spend=0.0,
team_id="team-integration-tests",
)
team_object = LiteLLM_TeamTable(
team_id="team-integration-tests",
max_budget=2000.0,
spend=0.144,
)
await key_cache.async_set_cache(
key=f"team_id:{team_object.team_id}",
value=team_object,
)
request_body = _request_body()
request_body.pop("max_tokens")
# Realistic Opus 4.7 output pricing — the 16K fallback × $25/M ≈ $0.40
# reservation per request, leaving ~5000 admittable concurrent requests
# against a $2000 team budget.
with patch(
"litellm.proxy.spend_tracking.budget_reservation._get_model_cost_info",
return_value={
"input_cost_per_token": 5e-6,
"output_cost_per_token": 2.5e-5,
"max_output_tokens": 128000,
},
):
first_reservation = await reserve_budget_for_request(
request_body=request_body,
route="/chat/completions",
llm_router=None,
valid_token=valid_token,
team_object=team_object,
user_object=None,
prisma_client=None,
user_api_key_cache=key_cache,
proxy_logging_obj=proxy_logging_obj,
)
# The team counter must not be pinned at max_budget while the first
# request is in flight, otherwise concurrent requests false-positive.
team_counter_after_first = (
counter_cache.in_memory_cache.get_cache(
key=f"spend:team:{team_object.team_id}"
)
or 0.0
)
assert team_counter_after_first < team_object.max_budget, (
f"Team counter sat at {team_counter_after_first} after one uncapped "
f"reservation against a {team_object.max_budget} budget — concurrent "
"requests will be falsely blocked."
)
# Second request — same shape — must succeed without raising.
second_reservation = await reserve_budget_for_request(
request_body=request_body,
route="/chat/completions",
llm_router=None,
valid_token=valid_token,
team_object=team_object,
user_object=None,
prisma_client=None,
user_api_key_cache=key_cache,
proxy_logging_obj=proxy_logging_obj,
)
assert second_reservation is not None
if first_reservation is not None:
await release_budget_reservation(first_reservation)
if second_reservation is not None:
await release_budget_reservation(second_reservation)

View file

@ -1728,6 +1728,67 @@ async def test_add_proxy_budget_to_db_only_creates_user_no_keys():
assert call_args.kwargs["query_type"] == "update_data"
@pytest.mark.asyncio
async def test_add_proxy_budget_to_db_backfills_budget_reset_at():
"""
Test that _upsert_proxy_budget_with_reset_at_backfill issues a conditional
update_many with `WHERE budget_reset_at IS NULL` to backfill the column on
rows that pre-existed without a reset schedule. Without this, the proxy
admin row stays at NULL and reset_budget_for_litellm_users never matches
it (NULL < now() is unknown in SQL), so the global proxy budget never
resets.
"""
from unittest.mock import AsyncMock, MagicMock, patch
import litellm
from litellm.proxy.proxy_server import ProxyStartupEvent
litellm.budget_duration = "30d"
litellm.max_budget = 100.0
litellm_proxy_budget_name = "litellm-proxy-budget"
mock_prisma = MagicMock()
mock_prisma.db.litellm_usertable.update_many = AsyncMock(return_value={"count": 1})
mock_generate_key_helper = AsyncMock(
return_value={
"user_id": litellm_proxy_budget_name,
"max_budget": 100.0,
"budget_duration": "30d",
"spend": 0,
"models": [],
}
)
with (
patch(
"litellm.proxy.proxy_server.generate_key_helper_fn",
mock_generate_key_helper,
),
patch("litellm.proxy.proxy_server.prisma_client", mock_prisma),
):
await ProxyStartupEvent._upsert_proxy_budget_with_reset_at_backfill(
litellm_proxy_budget_name
)
# Upsert ran with the configured budget
mock_generate_key_helper.assert_called_once()
# Backfill update_many ran with the conditional WHERE
mock_prisma.db.litellm_usertable.update_many.assert_called_once()
backfill_call = mock_prisma.db.litellm_usertable.update_many.call_args
assert backfill_call.kwargs["where"]["user_id"] == litellm_proxy_budget_name
assert backfill_call.kwargs["where"]["budget_reset_at"] is None
# The backfilled value must be a real future datetime — anything else and
# reset_budget_for_litellm_users would still skip the row.
from datetime import datetime, timezone
backfilled_reset_at = backfill_call.kwargs["data"]["budget_reset_at"]
assert isinstance(backfilled_reset_at, datetime)
assert backfilled_reset_at > datetime.now(timezone.utc)
@pytest.mark.asyncio
async def test_custom_ui_sso_sign_in_handler_config_loading():
"""
@ -6513,3 +6574,37 @@ async def test_get_current_spend_redis_error_falls_back_to_in_memory():
finally:
ps.spend_counter_cache = orig_counter
ps.prisma_client = orig_prisma
def test_realtime_websocket_route_aliases_registered():
"""Realtime sessions reach the proxy via three path aliases stacked on
`realtime_websocket_endpoint`. Dropping any of them silently 405s
WebSocket upgrades because the catch-all `/openai/{endpoint:path}`
HTTP passthrough only declares HTTP methods. The aliases must also be
in `LiteLLMRoutes.openai_routes` (so non-admin / team / key-scoped
auth allows them) and in `API_ROUTE_TO_CALL_TYPES` (so call-type-aware
logic such as guardrails can resolve the realtime call type)."""
from starlette.routing import WebSocketRoute
from litellm.proxy._types import LiteLLMRoutes
from litellm.proxy.proxy_server import app
from litellm.types.utils import API_ROUTE_TO_CALL_TYPES, CallTypes
websocket_paths = {
route.path for route in app.routes if isinstance(route, WebSocketRoute)
}
openai_routes = LiteLLMRoutes.openai_routes.value
for expected in ("/openai/v1/realtime", "/v1/realtime", "/realtime"):
assert expected in websocket_paths, (
f"{expected!r} missing from registered WebSocket routes; the "
f"realtime endpoint will 405 for clients hitting this path."
)
assert expected in openai_routes, (
f"{expected!r} missing from LiteLLMRoutes.openai_routes; "
f"non-admin / team / key-scoped users will get 403 on this path."
)
assert API_ROUTE_TO_CALL_TYPES.get(expected) == [CallTypes.arealtime], (
f"{expected!r} missing from API_ROUTE_TO_CALL_TYPES; call-type "
f"resolution will return None and break call-type-aware features."
)

View file

@ -110,6 +110,25 @@ def test_wandb_model_api_pricing_entries():
assert model_info["output_cost_per_token"] == output_cost
def test_openrouter_qwen36_plus_model_info():
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
litellm.model_cost = litellm.get_model_cost_map(url="")
model_info = litellm.model_cost.get("openrouter/qwen/qwen3.6-plus")
assert model_info is not None
assert model_info["litellm_provider"] == "openrouter"
assert model_info["mode"] == "chat"
assert model_info["max_input_tokens"] == 1000000
assert model_info["max_output_tokens"] == 65536
assert model_info["input_cost_per_token"] == 3.25e-07
assert model_info["output_cost_per_token"] == 1.95e-06
assert model_info["supports_function_calling"] is True
assert model_info["supports_tool_choice"] is True
assert model_info["supports_reasoning"] is True
assert model_info["supports_vision"] is True
def test_cost_calculator_with_usage(monkeypatch):
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
litellm.model_cost = litellm.get_model_cost_map(url="")

View file

@ -303,7 +303,7 @@ async def test_chat_completion():
api_key=key_gen["key"],
api_version="2024-02-15-preview",
)
with pytest.raises(openai.AuthenticationError) as e:
with pytest.raises(openai.PermissionDeniedError) as e:
response = await azure_client.chat.completions.create(
model="gpt-4",
messages=[{"role": "user", "content": "Hello!"}],

View file

@ -302,14 +302,14 @@ async def test_user_model_access():
model="good-model",
)
with pytest.raises(openai.AuthenticationError):
with pytest.raises(openai.PermissionDeniedError):
await chat_completion(
session=session,
key=key,
model="bedrock/anthropic.claude-3-sonnet-20240229-v1:0",
)
with pytest.raises(openai.AuthenticationError):
with pytest.raises(openai.PermissionDeniedError):
await chat_completion(
session=session,
key=key,

View file

@ -5,6 +5,7 @@ import { createGuardrailCall, getGuardrailProviderSpecificParams, getGuardrailUI
import ContentFilterConfiguration from "./content_filter/ContentFilterConfiguration";
import {
choiceToSkipSystemForCreate,
choiceToSkipToolForCreate,
getGuardrailProviders,
guardrail_provider_map,
guardrailLogoMap,
@ -188,6 +189,7 @@ const AddGuardrailForm: React.FC<AddGuardrailFormProps> = ({ visible, onClose, a
mode: preset.mode,
default_on: preset.defaultOn,
skip_system_message_choice: "inherit",
skip_tool_message_choice: "inherit",
};
if (preset.provider === "BlockCodeExecution") {
baseValues.confidence_threshold = 0.5;
@ -433,6 +435,11 @@ const AddGuardrailForm: React.FC<AddGuardrailFormProps> = ({ visible, onClose, a
guardrailData.litellm_params.skip_system_message_in_guardrail = skipForCreate;
}
const skipToolForCreate = choiceToSkipToolForCreate(values.skip_tool_message_choice);
if (skipToolForCreate !== undefined) {
guardrailData.litellm_params.skip_tool_message_in_guardrail = skipToolForCreate;
}
// For Presidio PII, add the entity and action configurations
if (values.provider === "PresidioPII" && selectedEntities.length > 0) {
const piiEntitiesConfig: { [key: string]: string } = {};
@ -804,6 +811,18 @@ const AddGuardrailForm: React.FC<AddGuardrailFormProps> = ({ visible, onClose, a
</Select>
</Form.Item>
<Form.Item
name="skip_tool_message_choice"
label="Skip tool messages in guardrail"
tooltip="Unified guardrails only: omit role: tool from guardrail evaluation input (OpenAI chat + Anthropic messages). The model still receives full messages. Use global default follows litellm_settings.skip_tool_message_in_guardrail."
>
<Select>
<Select.Option value="inherit">Use global default</Select.Option>
<Select.Option value="yes">Yes — exclude from guardrail scan</Select.Option>
<Select.Option value="no">No — always include in scan</Select.Option>
</Select>
</Form.Item>
{/* Use the GuardrailProviderFields component to render provider-specific fields */}
{!isToolPermissionProvider && !shouldRenderContentFilterConfigSettings(selectedProvider) && !shouldRenderLLMJudgeFields(selectedProvider) && (
<GuardrailProviderFields
@ -1155,6 +1174,7 @@ const AddGuardrailForm: React.FC<AddGuardrailFormProps> = ({ visible, onClose, a
mode: "pre_call",
default_on: false,
skip_system_message_choice: "inherit",
skip_tool_message_choice: "inherit",
}}
>
{stepConfigs.map((step, index) => {

View file

@ -6,6 +6,7 @@ import {
guardrailLogoMap,
getGuardrailProviders,
type SkipSystemMessageChoice,
type SkipToolMessageChoice,
} from "./guardrail_info_helpers";
import { getGuardrailUISettings, getGlobalLitellmHeaderName } from "../networking";
import PiiConfiguration from "./pii_configuration";
@ -29,6 +30,7 @@ interface EditGuardrailFormProps {
default_on: boolean;
pii_entities_config?: { [key: string]: string };
skip_system_message_choice?: SkipSystemMessageChoice;
skip_tool_message_choice?: SkipToolMessageChoice;
[key: string]: any;
};
}
@ -138,6 +140,15 @@ const EditGuardrailForm: React.FC<EditGuardrailFormProps> = ({
delete litellm_params.skip_system_message_in_guardrail;
}
const skipToolChoice = values.skip_tool_message_choice as SkipToolMessageChoice | undefined;
if (skipToolChoice === "yes") {
litellm_params.skip_tool_message_in_guardrail = true;
} else if (skipToolChoice === "no") {
litellm_params.skip_tool_message_in_guardrail = false;
} else {
delete litellm_params.skip_tool_message_in_guardrail;
}
let guardrail_info: any = {};
// For Presidio PII, add the entity and action configurations
@ -432,6 +443,18 @@ const EditGuardrailForm: React.FC<EditGuardrailFormProps> = ({
</Select>
</Form.Item>
<Form.Item
name="skip_tool_message_choice"
label="Skip tool messages in guardrail"
tooltip="Unified guardrails only: whether role: tool content is omitted from guardrail input (LLM still receives full messages). Use global default follows litellm_settings.skip_tool_message_in_guardrail."
>
<Select>
<Option value="inherit">Use global default</Option>
<Option value="yes">Yes — exclude from guardrail scan</Option>
<Option value="no">No — always include in scan</Option>
</Select>
</Form.Item>
{renderProviderSpecificFields()}
<div className="flex justify-end space-x-2 mt-4">

View file

@ -29,7 +29,9 @@ import {
getGuardrailLogoAndName,
guardrail_provider_map,
skipSystemMessageToChoice,
skipToolMessageToChoice,
type SkipSystemMessageChoice,
type SkipToolMessageChoice,
} from "./guardrail_info_helpers";
import GuardrailOptionalParams from "./guardrail_optional_params";
import GuardrailProviderFields from "./guardrail_provider_fields";
@ -214,12 +216,16 @@ const GuardrailInfoView: React.FC<GuardrailInfoProps> = ({ guardrailId, onClose,
if (guardrailData && form) {
const lp = { ...(guardrailData.litellm_params || {}) };
delete lp.skip_system_message_in_guardrail;
delete lp.skip_tool_message_in_guardrail;
form.setFieldsValue({
guardrail_name: guardrailData.guardrail_name,
...lp,
skip_system_message_choice: skipSystemMessageToChoice(
guardrailData.litellm_params?.skip_system_message_in_guardrail,
),
skip_tool_message_choice: skipToolMessageToChoice(
guardrailData.litellm_params?.skip_tool_message_in_guardrail,
),
guardrail_info: guardrailData.guardrail_info ? JSON.stringify(guardrailData.guardrail_info, null, 2) : "",
// Include any optional_params if they exist
...(guardrailData.litellm_params?.optional_params && {
@ -302,6 +308,20 @@ const GuardrailInfoView: React.FC<GuardrailInfoProps> = ({ guardrailId, onClose,
}
}
const prevSkipToolChoice = skipToolMessageToChoice(
guardrailData.litellm_params?.skip_tool_message_in_guardrail,
);
const nextSkipToolChoice = values.skip_tool_message_choice as SkipToolMessageChoice | undefined;
if (nextSkipToolChoice !== undefined && nextSkipToolChoice !== prevSkipToolChoice) {
if (nextSkipToolChoice === "inherit") {
updateData.litellm_params.skip_tool_message_in_guardrail = null;
} else if (nextSkipToolChoice === "yes") {
updateData.litellm_params.skip_tool_message_in_guardrail = true;
} else {
updateData.litellm_params.skip_tool_message_in_guardrail = false;
}
}
// Only include guardrail_info if it has changed
const originalGuardrailInfo = guardrailData.guardrail_info;
const newGuardrailInfo = values.guardrail_info ? JSON.parse(values.guardrail_info) : undefined;
@ -674,11 +694,15 @@ const GuardrailInfoView: React.FC<GuardrailInfoProps> = ({ guardrailId, onClose,
...(() => {
const lp = { ...(guardrailData.litellm_params || {}) };
delete lp.skip_system_message_in_guardrail;
delete lp.skip_tool_message_in_guardrail;
return lp;
})(),
skip_system_message_choice: skipSystemMessageToChoice(
guardrailData.litellm_params?.skip_system_message_in_guardrail,
),
skip_tool_message_choice: skipToolMessageToChoice(
guardrailData.litellm_params?.skip_tool_message_in_guardrail,
),
guardrail_info: guardrailData.guardrail_info
? JSON.stringify(guardrailData.guardrail_info, null, 2)
: "",
@ -716,6 +740,18 @@ const GuardrailInfoView: React.FC<GuardrailInfoProps> = ({ guardrailId, onClose,
</Select>
</Form.Item>
<Form.Item
label="Skip tool messages in guardrail"
name="skip_tool_message_choice"
tooltip="Unified guardrails: omit role: tool from guardrail input (LLM still gets full messages). Use global default follows litellm_settings.skip_tool_message_in_guardrail."
>
<Select>
<Select.Option value="inherit">Use global default</Select.Option>
<Select.Option value="yes">Yes — exclude from guardrail scan</Select.Option>
<Select.Option value="no">No — always include in scan</Select.Option>
</Select>
</Form.Item>
{guardrailData.litellm_params?.guardrail === "presidio" && (
<>
<Divider orientation="left">PII Protection</Divider>

View file

@ -12,6 +12,8 @@ import {
GuardrailProviders,
skipSystemMessageToChoice,
choiceToSkipSystemForCreate,
skipToolMessageToChoice,
choiceToSkipToolForCreate,
} from "./guardrail_info_helpers";
describe("guardrail_info_helpers", () => {
@ -215,4 +217,18 @@ describe("guardrail_info_helpers", () => {
expect(choiceToSkipSystemForCreate("no")).toBe(false);
});
});
describe("skipToolMessageToChoice / choiceToSkipToolForCreate", () => {
it("maps API values to form choices and back for create", () => {
expect(skipToolMessageToChoice(undefined)).toBe("inherit");
expect(skipToolMessageToChoice(null)).toBe("inherit");
expect(skipToolMessageToChoice(true)).toBe("yes");
expect(skipToolMessageToChoice(false)).toBe("no");
expect(choiceToSkipToolForCreate("inherit")).toBeUndefined();
expect(choiceToSkipToolForCreate(undefined)).toBeUndefined();
expect(choiceToSkipToolForCreate("yes")).toBe(true);
expect(choiceToSkipToolForCreate("no")).toBe(false);
});
});
});

View file

@ -179,3 +179,19 @@ export function choiceToSkipSystemForCreate(choice: SkipSystemMessageChoice | un
if (choice === "no") return false;
return undefined;
}
/** Tri-state UI value for `litellm_params.skip_tool_message_in_guardrail` (inherit = use global). */
export type SkipToolMessageChoice = "inherit" | "yes" | "no";
export function skipToolMessageToChoice(v: boolean | null | undefined): SkipToolMessageChoice {
if (v === true) return "yes";
if (v === false) return "no";
return "inherit";
}
/** Create flow: omit key when inheriting global default. */
export function choiceToSkipToolForCreate(choice: SkipToolMessageChoice | undefined): boolean | undefined {
if (choice === "yes") return true;
if (choice === "no") return false;
return undefined;
}

View file

@ -11,7 +11,12 @@ import {
SortingState,
useReactTable,
} from "@tanstack/react-table";
import { getGuardrailLogoAndName, guardrail_provider_map, skipSystemMessageToChoice } from "./guardrail_info_helpers";
import {
getGuardrailLogoAndName,
guardrail_provider_map,
skipSystemMessageToChoice,
skipToolMessageToChoice,
} from "./guardrail_info_helpers";
import EditGuardrailForm from "./edit_guardrail_form";
import { Guardrail, GuardrailDefinitionLocation } from "./types";
@ -304,6 +309,9 @@ const GuardrailTable: React.FC<GuardrailTableProps> = ({
skip_system_message_choice: skipSystemMessageToChoice(
selectedGuardrail.litellm_params?.skip_system_message_in_guardrail,
),
skip_tool_message_choice: skipToolMessageToChoice(
selectedGuardrail.litellm_params?.skip_tool_message_in_guardrail,
),
...selectedGuardrail.guardrail_info,
}}
/>

View file

@ -404,3 +404,52 @@ describe("individualModelHealthCheckCall", () => {
expect(parsed.searchParams.get("model_id")).toBe("id/with/slashes");
});
});
describe("teamInfoCall", () => {
const originalFetch = global.fetch;
beforeEach(() => {
vi.clearAllMocks();
});
afterEach(() => {
global.fetch = originalFetch;
});
it("should URL-encode team_id query param to handle special characters safely", async () => {
const mockFetch = vi.fn().mockResolvedValue({
ok: true,
json: vi.fn().mockResolvedValue({ team_id: "team with spaces & special?chars" }),
} as any);
global.fetch = mockFetch as any;
const teamID = "team with spaces & special?chars";
await Networking.teamInfoCall("token", teamID);
expect(mockFetch).toHaveBeenCalledOnce();
const [url] = mockFetch.mock.calls[0];
const urlStr = typeof url === "string" ? url : (url as Request).url;
const parsed = typeof url === "string" ? new URL(url, "http://example.com") : new URL((url as Request).url);
expect(urlStr).toContain("/team/info");
// Encoded value is present in the raw URL string (verifies encodeURIComponent was used)
expect(urlStr).toContain(`team_id=${encodeURIComponent(teamID)}`);
// Round-trip parse returns the original team_id
expect(parsed.searchParams.get("team_id")).toBe(teamID);
});
it("should not append team_id when teamID is null", async () => {
const mockFetch = vi.fn().mockResolvedValue({
ok: true,
json: vi.fn().mockResolvedValue({}),
} as any);
global.fetch = mockFetch as any;
await Networking.teamInfoCall("token", null);
expect(mockFetch).toHaveBeenCalledOnce();
const [url] = mockFetch.mock.calls[0];
const parsed = typeof url === "string" ? new URL(url, "http://example.com") : new URL((url as Request).url);
expect(parsed.searchParams.has("team_id")).toBe(false);
});
});

View file

@ -1386,7 +1386,7 @@ export const teamInfoCall = async (accessToken: string, teamID: string | null) =
try {
let url = proxyBaseUrl ? `${proxyBaseUrl}/team/info` : `/team/info`;
if (teamID) {
url = `${url}?team_id=${teamID}`;
url = `${url}?team_id=${encodeURIComponent(teamID)}`;
}
console.log("in teamInfoCall");
const response = await fetch(url, {

View file

@ -155,6 +155,15 @@ vi.mock("antd", () => {
const Button = ({ children, htmlType, ...props }: { children?: any; htmlType?: string }) =>
React.createElement("button", { ...props, type: htmlType ?? props.type }, children);
const Typography = ({ children, ...props }: { children?: any }) =>
React.createElement("div", props, children);
Typography.Text = ({ children, ...props }: { children?: any }) =>
React.createElement("span", props, children);
Typography.Paragraph = ({ children, ...props }: { children?: any }) =>
React.createElement("p", props, children);
Typography.Title = ({ children, ...props }: { children?: any }) =>
React.createElement("h1", props, children);
return {
Button,
Form,
@ -171,6 +180,7 @@ vi.mock("antd", () => {
Switch,
Tag,
Tooltip,
Typography,
};
});

View file

@ -9,7 +9,7 @@ import { formatNumberWithCommas } from "@/utils/dataUtils";
import { InfoCircleOutlined } from "@ant-design/icons";
import { useQueryClient } from "@tanstack/react-query";
import { Accordion, AccordionBody, AccordionHeader, Button, Col, Grid, Text, TextInput, Title } from "@tremor/react";
import { Button as Button2, Form, Input, Modal, Radio, Select, Switch, Tag, Tooltip } from "antd";
import { Button as Button2, Form, Input, Modal, Radio, Select, Switch, Tag, Tooltip, Typography } from "antd";
import debounce from "lodash/debounce";
import React, { useCallback, useEffect, useState } from "react";
import { rolesWithWriteAccess } from "../../utils/roles";
@ -979,28 +979,28 @@ const CreateKey: React.FC<CreateKeyProps> = ({ team, teams, data, addKey, autoOp
}
}}
>
<Option value="default" label="Default">
<div style={{ padding: "4px 0" }}>
<div style={{ fontWeight: 500 }}>Default</div>
<div style={{ fontSize: "11px", color: "#6b7280", marginTop: "2px" }}>
Can call AI APIs + Management routes
</div>
</div>
</Option>
<Option value="llm_api" label="AI APIs">
<div style={{ padding: "4px 0" }}>
<div style={{ fontWeight: 500 }}>AI APIs</div>
<div style={{ fontSize: "11px", color: "#6b7280", marginTop: "2px" }}>
<Typography.Text strong>AI APIs</Typography.Text>
<Typography.Paragraph type="secondary" style={{ fontSize: 11, margin: "2px 0 0" }}>
Can call only AI API routes (chat/completions, embeddings, etc.)
</div>
</Typography.Paragraph>
</div>
</Option>
<Option value="management" label="Management">
<div style={{ padding: "4px 0" }}>
<div style={{ fontWeight: 500 }}>Management</div>
<div style={{ fontSize: "11px", color: "#6b7280", marginTop: "2px" }}>
<Typography.Text strong>Management</Typography.Text>
<Typography.Paragraph type="secondary" style={{ fontSize: 11, margin: "2px 0 0" }}>
Can call only management routes (user/team/key management)
</div>
</Typography.Paragraph>
</div>
</Option>
<Option value="default" label="Full Access">
<div style={{ padding: "4px 0" }}>
<Typography.Text strong>Full Access</Typography.Text>
<Typography.Paragraph type="secondary" style={{ fontSize: 11, margin: "2px 0 0" }}>
Can call all routes (AI APIs, Management, and read-only)
</Typography.Paragraph>
</div>
</Option>
</Select>

2
uv.lock generated
View file

@ -3405,7 +3405,7 @@ requires-dist = [
{ name = "gunicorn", marker = "extra == 'proxy'", specifier = "==23.0.0" },
{ name = "httpx", specifier = ">=0.28.0,<1.0" },
{ name = "importlib-metadata", specifier = ">=8.0.0,<9.0" },
{ name = "jinja2", specifier = ">=3.1.0,<4.0" },
{ name = "jinja2", specifier = ">=3.1.6,<4.0" },
{ name = "jsonschema", specifier = ">=4.0.0,<5.0" },
{ name = "langfuse", marker = "extra == 'proxy-runtime'", specifier = "==2.59.7" },
{ name = "litellm-enterprise", marker = "extra == 'proxy'", editable = "enterprise" },