Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_openapi_mcp_extra_headers_forward

Co-authored-by: Mateo Wang <mateo-berri@users.noreply.github.com>
This commit is contained in:
Cursor Agent 2026-05-09 18:00:46 +00:00
commit 2392f81b93
No known key found for this signature in database
74 changed files with 5932 additions and 1586 deletions

View file

@ -1,8 +1,8 @@
# Base image for building
ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:3258be472764337fd13095bcbb3182da170243b5819fd67ad4c0754590588b31
ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:31da6565f35af6401031c1d7aa91dc84ac76c5c48edd17fb90f0ed9e3173c7a9
# Runtime image
ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:3258be472764337fd13095bcbb3182da170243b5819fd67ad4c0754590588b31
ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:31da6565f35af6401031c1d7aa91dc84ac76c5c48edd17fb90f0ed9e3173c7a9
ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:240fb85ab0f263ef12f492d8476aa3a2e4e1e333f7d67fbdd923d00a506a516a
FROM $UV_IMAGE AS uvbin
@ -68,8 +68,8 @@ FROM $LITELLM_RUNTIME_IMAGE AS runtime
USER root
RUN apk add --no-cache bash openssl tzdata nodejs npm python3 libsndfile supervisor && \
npm install -g npm@11.12.1 tar@7.5.11 glob@13.0.6 @isaacs/brace-expansion@5.0.1 brace-expansion@5.0.5 minimatch@10.2.4 diff@8.0.3 picomatch@4.0.4 && \
RUN apk add --no-cache bash openssl tzdata nodejs npm python3 libsndfile && \
npm install -g npm@11.14.0 tar@7.5.11 glob@13.0.6 @isaacs/brace-expansion@5.0.1 brace-expansion@5.0.5 minimatch@10.2.4 diff@8.0.3 picomatch@4.0.4 && \
GLOBAL="$(npm root -g)" && \
for pkg in tar glob @isaacs/brace-expansion brace-expansion minimatch diff picomatch; do \
name="${pkg##*/}"; \
@ -85,17 +85,17 @@ ENV PATH="/app/.venv/bin:${PATH}"
COPY --from=builder /app /app
# Prisma binaries live in $HOME/.cache (default prisma-python location),
# which is /root/.cache here. Copy them from the builder so they survive
# deployments that volume-mount /app/.cache (e.g. readOnlyRootFilesystem
# + emptyDir) — otherwise the mount would shadow the baked-in query engine.
COPY --from=builder /root/.cache /root/.cache
# which is /root/.cache here. Copy only the Prisma subdirs — copying the
# whole /root/.cache drags in the uv build cache (~660 MB, includes a
# setuptools wheel that surfaces as a CVE finding even though it's not
# on the runtime sys.path).
COPY --from=builder /root/.cache/prisma /root/.cache/prisma
COPY --from=builder /root/.cache/prisma-python /root/.cache/prisma-python
RUN find /app/.venv -type f -path "*/tornado/test/*" -delete && \
find /app/.venv -type d -path "*/tornado/test" -delete
EXPOSE 4000/tcp
COPY docker/supervisord.conf /etc/supervisord.conf
ENTRYPOINT ["docker/prod_entrypoint.sh"]
CMD ["--port", "4000"]

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:
@ -129,12 +139,6 @@ spec:
value: {{ $val | quote }}
{{- end }}
{{- end }}
{{- if .Values.separateHealthApp }}
- name: SEPARATE_HEALTH_APP
value: "1"
- name: SEPARATE_HEALTH_PORT
value: {{ .Values.separateHealthPort | default "8081" | quote }}
{{- end }}
{{- with .Values.extraEnvVars }}
{{- toYaml . | nindent 12 }}
{{- end }}
@ -175,15 +179,10 @@ spec:
- name: http
containerPort: {{ .Values.service.port }}
protocol: TCP
{{- if .Values.separateHealthApp }}
- name: health
containerPort: {{ .Values.separateHealthPort | default 8081 }}
protocol: TCP
{{- end }}
livenessProbe:
httpGet:
path: {{ .Values.livenessProbe.path | quote }}
port: {{ if .Values.separateHealthApp }}"health"{{ else }}"http"{{ end }}
port: "http"
initialDelaySeconds: {{ .Values.livenessProbe.initialDelaySeconds }}
periodSeconds: {{ .Values.livenessProbe.periodSeconds }}
timeoutSeconds: {{ .Values.livenessProbe.timeoutSeconds }}
@ -192,7 +191,7 @@ spec:
readinessProbe:
httpGet:
path: {{ .Values.readinessProbe.path | quote }}
port: {{ if .Values.separateHealthApp }}"health"{{ else }}"http"{{ end }}
port: "http"
initialDelaySeconds: {{ .Values.readinessProbe.initialDelaySeconds }}
periodSeconds: {{ .Values.readinessProbe.periodSeconds }}
timeoutSeconds: {{ .Values.readinessProbe.timeoutSeconds }}
@ -201,7 +200,7 @@ spec:
startupProbe:
httpGet:
path: {{ .Values.startupProbe.path | quote }}
port: {{ if .Values.separateHealthApp }}"health"{{ else }}"http"{{ end }}
port: "http"
initialDelaySeconds: {{ .Values.startupProbe.initialDelaySeconds }}
periodSeconds: {{ .Values.startupProbe.periodSeconds }}
timeoutSeconds: {{ .Values.startupProbe.timeoutSeconds }}

View file

@ -88,12 +88,6 @@ service:
# optionally specify loadBalancerClass
# loadBalancerClass: tailscale
# Separate health app configuration
# When enabled, health checks will use a separate port and the application
# will receive SEPARATE_HEALTH_APP=1 and SEPARATE_HEALTH_PORT from environment variables
separateHealthApp: false
separateHealthPort: 8081
# Probes for LiteLLM gateway container
livenessProbe:
path: /health/liveliness
@ -258,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

@ -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

@ -66,7 +66,7 @@ FROM $LITELLM_RUNTIME_IMAGE AS runtime
USER root
RUN apk add --no-cache bash openssl tzdata nodejs npm python3 libsndfile supervisor && \
RUN apk add --no-cache bash openssl tzdata nodejs npm python3 libsndfile && \
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 \
@ -102,7 +102,5 @@ RUN find /app/.venv -type f -path "*/tornado/test/*" -delete && \
EXPOSE 4000/tcp
COPY docker/supervisord.conf /etc/supervisord.conf
ENTRYPOINT ["docker/prod_entrypoint.sh"]
CMD ["--port", "4000"]

View file

@ -103,13 +103,12 @@ RUN for i in 1 2 3; do \
apk upgrade --no-cache && break || sleep 5; \
done && \
for i in 1 2 3; do \
apk add --no-cache python3 bash openssl tzdata supervisor libsndfile nodejs && break || sleep 5; \
apk add --no-cache python3 bash openssl tzdata libsndfile nodejs && break || sleep 5; \
done
COPY --from=builder /app /app
COPY --from=builder /var/lib/litellm/ui /var/lib/litellm/ui
COPY --from=builder /var/lib/litellm/assets /var/lib/litellm/assets
COPY --from=builder /app/docker/supervisord.conf /etc/supervisord.conf
ENV PATH="/app/.venv/bin:${PATH}" \
PRISMA_BINARY_CACHE_DIR=/app/.cache/prisma-python/binaries \

View file

@ -1,14 +1,8 @@
#!/bin/sh
if [ "$SEPARATE_HEALTH_APP" = "1" ]; then
export LITELLM_ARGS="$@"
export SUPERVISORD_STOPWAITSECS="${SUPERVISORD_STOPWAITSECS:-3600}"
exec supervisord -c /etc/supervisord.conf
fi
if [ "$USE_DDTRACE" = "true" ]; then
export DD_TRACE_OPENAI_ENABLED="False"
exec ddtrace-run litellm "$@"
else
exec litellm "$@"
fi
fi

View file

@ -1,46 +0,0 @@
[supervisord]
nodaemon=true
loglevel=info
logfile=/tmp/supervisord.log
pidfile=/tmp/supervisord.pid
[group:litellm]
programs=main,health
[program:main]
command=sh -c 'if [ "$USE_DDTRACE" = "true" ]; then export DD_TRACE_OPENAI_ENABLED="False"; exec ddtrace-run python -m litellm.proxy.proxy_cli --host 0.0.0.0 --port=4000 $LITELLM_ARGS; else exec python -m litellm.proxy.proxy_cli --host 0.0.0.0 --port=4000 $LITELLM_ARGS; fi'
autostart=true
autorestart=true
startretries=3
priority=1
exitcodes=0
stopasgroup=true
killasgroup=true
stopwaitsecs=%(ENV_SUPERVISORD_STOPWAITSECS)s
stdout_logfile=/dev/stdout
stderr_logfile=/dev/stderr
stdout_logfile_maxbytes = 0
stderr_logfile_maxbytes = 0
environment=PYTHONUNBUFFERED=true
[program:health]
command=sh -c '[ "$SEPARATE_HEALTH_APP" = "1" ] && exec uvicorn litellm.proxy.health_endpoints.health_app_factory:build_health_app --factory --host 0.0.0.0 --port=${SEPARATE_HEALTH_PORT:-4001} || exit 0'
autostart=true
autorestart=true
startretries=3
priority=2
exitcodes=0
stopasgroup=true
killasgroup=true
stopwaitsecs=%(ENV_SUPERVISORD_STOPWAITSECS)s
stdout_logfile=/dev/stdout
stderr_logfile=/dev/stderr
stdout_logfile_maxbytes = 0
stderr_logfile_maxbytes = 0
environment=PYTHONUNBUFFERED=true
[eventlistener:process_monitor]
command=python -c "from supervisor import childutils; import os, signal; [os.kill(os.getppid(), signal.SIGTERM) for h,p in iter(lambda: childutils.listener.wait(), None) if h['eventname'] in ['PROCESS_STATE_FATAL', 'PROCESS_STATE_EXITED'] and dict([x.split(':') for x in p.split(' ')])['processname'] in ['main', 'health'] or childutils.listener.ok()]"
events=PROCESS_STATE_EXITED,PROCESS_STATE_FATAL
autostart=true
autorestart=true

View file

@ -535,9 +535,9 @@
}
},
"node_modules/get-tsconfig": {
"version": "4.13.0",
"resolved": "https://registry.npmjs.org/get-tsconfig/-/get-tsconfig-4.13.0.tgz",
"integrity": "sha512-1VKTZJCwBrvbd+Wn3AOgQP/2Av+TfTCOlE4AcRJE72W1ksZXbAx8PPBR9RzgTeSPzlPMHrbANMH3LbltH73wxQ==",
"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": {

View file

@ -388,6 +388,7 @@ anthropic_beta_headers_url: str = os.getenv(
suppress_debug_info = False
dynamodb_table_name: Optional[str] = None
s3_callback_params: Optional[Dict] = None
s3_audit_callback_params: Optional[Dict] = None
datadog_llm_observability_params: Optional[Union[DatadogLLMObsInitParams, Dict]] = None
datadog_params: Optional[Union[DatadogInitParams, Dict]] = None
aws_sqs_callback_params: Optional[Dict] = None

View file

@ -161,6 +161,11 @@ MCP_STDIO_ALLOWED_COMMANDS: frozenset = frozenset(
| (set(_MCP_STDIO_EXTRA_COMMANDS.split(",")) - {""})
)
# MCP OAuth2 Token Exchange (OBO) Defaults
MCP_TOKEN_EXCHANGE_CACHE_MAX_SIZE = int(
os.getenv("MCP_TOKEN_EXCHANGE_CACHE_MAX_SIZE", "500")
)
LITELLM_UI_ALLOW_HEADERS = [
"x-litellm-semantic-filter",
"x-litellm-semantic-filter-tools",

View file

@ -366,6 +366,8 @@ class MCPClient:
headers["Authorization"] = f"Bearer {self._mcp_auth_value}"
elif self.auth_type == MCPAuth.token:
headers["Authorization"] = f"token {self._mcp_auth_value}"
elif self.auth_type == MCPAuth.oauth2_token_exchange:
headers["Authorization"] = f"Bearer {self._mcp_auth_value}"
elif isinstance(self._mcp_auth_value, dict):
headers.update(self._mcp_auth_value)
# Note: aws_sigv4 auth is not handled here — SigV4 requires per-request

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

@ -16,6 +16,7 @@ from litellm._logging import print_verbose, verbose_logger
from litellm.constants import DEFAULT_S3_BATCH_SIZE, DEFAULT_S3_FLUSH_INTERVAL_SECONDS
from litellm.integrations.s3 import get_s3_object_key
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
from litellm.llms.custom_httpx.http_handler import (
_get_httpx_client,
@ -53,15 +54,25 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
s3_strip_base64_files: bool = False,
s3_use_key_prefix: bool = False,
s3_use_virtual_hosted_style: bool = False,
s3_callback_params_override: Optional[dict] = None,
**kwargs,
):
try:
verbose_logger.debug(
f"in init s3 logger - s3_callback_params {litellm.s3_callback_params}"
)
_masker = SensitiveDataMasker()
if s3_callback_params_override is not None:
verbose_logger.debug(
f"in init s3 logger (audit override) - "
f"{_masker.mask_dict(dict(s3_callback_params_override))}"
)
else:
verbose_logger.debug(
f"in init s3 logger - s3_callback_params "
f"{_masker.mask_dict(dict(litellm.s3_callback_params or {}))}"
)
# Initialize S3 params first to get the correct s3_verify value
self._init_s3_params(
params_source=s3_callback_params_override,
s3_bucket_name=s3_bucket_name,
s3_region_name=s3_region_name,
s3_api_version=s3_api_version,
@ -139,94 +150,85 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
s3_strip_base64_files: bool = False,
s3_use_key_prefix: bool = False,
s3_use_virtual_hosted_style: bool = False,
params_source: Optional[dict] = None,
):
"""
Initialize the s3 params for this logging callback
Initialize the s3 params for this logging callback. Reads from
`params_source` if given (e.g. `s3_audit_callback_params` for the
audit-log instance), otherwise falls back to `litellm.s3_callback_params`.
Resolves `os.environ/X` markers into a local dict; never mutates the source.
"""
litellm.s3_callback_params = litellm.s3_callback_params or {}
# read in .env variables - example os.environ/AWS_BUCKET_NAME
for key, value in litellm.s3_callback_params.items():
if isinstance(value, str) and value.startswith("os.environ/"):
litellm.s3_callback_params[key] = litellm.get_secret(value)
if params_source is None:
params_source = litellm.s3_callback_params or {}
params: dict = {
key: (
litellm.get_secret(value)
if isinstance(value, str) and value.startswith("os.environ/")
else value
)
for key, value in params_source.items()
}
self.s3_bucket_name = (
litellm.s3_callback_params.get("s3_bucket_name") or s3_bucket_name
)
self.s3_region_name = (
litellm.s3_callback_params.get("s3_region_name") or s3_region_name
)
self.s3_api_version = (
litellm.s3_callback_params.get("s3_api_version") or s3_api_version
)
self.s3_bucket_name = params.get("s3_bucket_name") or s3_bucket_name
self.s3_region_name = params.get("s3_region_name") or s3_region_name
self.s3_api_version = params.get("s3_api_version") or s3_api_version
self.s3_use_ssl = (
litellm.s3_callback_params.get("s3_use_ssl", True)
if litellm.s3_callback_params.get("s3_use_ssl") is not None
params.get("s3_use_ssl", True)
if params.get("s3_use_ssl") is not None
else s3_use_ssl
)
self.s3_verify = (
litellm.s3_callback_params.get("s3_verify")
if litellm.s3_callback_params.get("s3_verify") is not None
params.get("s3_verify")
if params.get("s3_verify") is not None
else s3_verify
)
self.s3_endpoint_url = (
litellm.s3_callback_params.get("s3_endpoint_url") or s3_endpoint_url
)
self.s3_endpoint_url = params.get("s3_endpoint_url") or s3_endpoint_url
self.s3_aws_access_key_id = (
litellm.s3_callback_params.get("s3_aws_access_key_id")
or s3_aws_access_key_id
params.get("s3_aws_access_key_id") or s3_aws_access_key_id
)
self.s3_aws_secret_access_key = (
litellm.s3_callback_params.get("s3_aws_secret_access_key")
or s3_aws_secret_access_key
params.get("s3_aws_secret_access_key") or s3_aws_secret_access_key
)
self.s3_aws_session_token = (
litellm.s3_callback_params.get("s3_aws_session_token")
or s3_aws_session_token
params.get("s3_aws_session_token") or s3_aws_session_token
)
self.s3_aws_session_name = (
litellm.s3_callback_params.get("s3_aws_session_name") or s3_aws_session_name
params.get("s3_aws_session_name") or s3_aws_session_name
)
self.s3_aws_profile_name = (
litellm.s3_callback_params.get("s3_aws_profile_name") or s3_aws_profile_name
params.get("s3_aws_profile_name") or s3_aws_profile_name
)
self.s3_aws_role_name = (
litellm.s3_callback_params.get("s3_aws_role_name") or s3_aws_role_name
)
self.s3_aws_role_name = params.get("s3_aws_role_name") or s3_aws_role_name
self.s3_aws_web_identity_token = (
litellm.s3_callback_params.get("s3_aws_web_identity_token")
or s3_aws_web_identity_token
params.get("s3_aws_web_identity_token") or s3_aws_web_identity_token
)
self.s3_aws_sts_endpoint = (
litellm.s3_callback_params.get("s3_aws_sts_endpoint") or s3_aws_sts_endpoint
params.get("s3_aws_sts_endpoint") or s3_aws_sts_endpoint
)
self.s3_config = litellm.s3_callback_params.get("s3_config") or s3_config
self.s3_path = litellm.s3_callback_params.get("s3_path") or s3_path
# done reading litellm.s3_callback_params
self.s3_config = params.get("s3_config") or s3_config
self.s3_path = params.get("s3_path") or s3_path
self.s3_use_team_prefix = (
bool(litellm.s3_callback_params.get("s3_use_team_prefix", False))
or s3_use_team_prefix
bool(params.get("s3_use_team_prefix", False)) or s3_use_team_prefix
)
self.s3_use_key_prefix = (
bool(litellm.s3_callback_params.get("s3_use_key_prefix", False))
or s3_use_key_prefix
bool(params.get("s3_use_key_prefix", False)) or s3_use_key_prefix
)
self.s3_strip_base64_files = (
bool(litellm.s3_callback_params.get("s3_strip_base64_files", False))
or s3_strip_base64_files
bool(params.get("s3_strip_base64_files", False)) or s3_strip_base64_files
)
self.s3_use_virtual_hosted_style = (
bool(litellm.s3_callback_params.get("s3_use_virtual_hosted_style", False))
bool(params.get("s3_use_virtual_hosted_style", False))
or s3_use_virtual_hosted_style
)

View file

@ -436,12 +436,21 @@ def update_messages_with_model_file_ids(
"""
Updates messages with model file ids.
For managed files (unified file IDs), uses model_file_id_mapping if it
resolves the id, otherwise decodes the base64-encoded unified file ID
and extracts the llm_output_file_id directly. Mirrors the Responses-API
sibling `update_responses_input_with_model_file_ids`.
model_file_id_mapping: Dict[str, Dict[str, str]] = {
"litellm_proxy/file_id": {
"model_id": "provider_file_id"
}
}
"""
from litellm.proxy.openai_files_endpoints.common_utils import (
_is_base64_encoded_unified_file_id,
convert_b64_uid_to_unified_uid,
)
for message in messages:
if message.get("role") == "user":
@ -450,7 +459,13 @@ def update_messages_with_model_file_ids(
if isinstance(content, str):
continue
for c in content:
if c["type"] == "file":
if not isinstance(c, dict):
# Content list items aren't always dicts. e.g.
# text_completion forwards a token-ids list/list-of-
# lists through this path. Skip non-dict items
# instead of indexing into them.
continue
if c.get("type") == "file":
file_object = cast(ChatCompletionFileObject, c)
file_object_file_field = file_object.get("file")
if not isinstance(file_object_file_field, dict):
@ -468,9 +483,23 @@ def update_messages_with_model_file_ids(
if file_id:
provider_file_id = (
model_file_id_mapping.get(file_id, {}).get(model_id)
or file_id
if model_file_id_mapping
else None
)
if (
not provider_file_id
and _is_base64_encoded_unified_file_id(file_id)
):
unified_file_id = convert_b64_uid_to_unified_uid(
file_id
)
if "llm_output_file_id," in unified_file_id:
provider_file_id = unified_file_id.split(
"llm_output_file_id,"
)[1].split(";")[0]
file_object_file_field["file_id"] = (
provider_file_id or file_id
)
file_object_file_field["file_id"] = provider_file_id
if format:
file_object_file_field["format"] = format
return messages

View file

@ -1459,14 +1459,14 @@ def completion( # type: ignore # noqa: PLR0915
if eos_token:
custom_prompt_dict[model]["eos_token"] = eos_token
if kwargs.get("model_file_id_mapping"):
messages = update_messages_with_model_file_ids(
messages=messages,
model_id=kwargs.get("model_info", {}).get("id", None),
model_file_id_mapping=cast(
Dict[str, Dict[str, str]], kwargs.get("model_file_id_mapping")
),
)
messages = update_messages_with_model_file_ids(
messages=messages,
model_id=kwargs.get("model_info", {}).get("id", None),
model_file_id_mapping=cast(
Dict[str, Dict[str, str]],
kwargs.get("model_file_id_mapping") or {},
),
)
provider_config: Optional[BaseConfig] = None
if custom_llm_provider is not None and custom_llm_provider in [

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",
@ -34940,6 +34954,48 @@
"supports_vision": true,
"supports_web_search": true
},
"xai/grok-4.3": {
"cache_read_input_token_cost": 2e-07,
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
"input_cost_per_token": 1.25e-06,
"input_cost_per_token_above_200k_tokens": 2.5e-06,
"litellm_provider": "xai",
"max_input_tokens": 1000000,
"max_output_tokens": 1000000,
"max_tokens": 1000000,
"mode": "chat",
"output_cost_per_token": 2.5e-06,
"output_cost_per_token_above_200k_tokens": 5e-06,
"source": "https://docs.x.ai/docs/models",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
"supports_web_search": true
},
"xai/grok-4.3-latest": {
"cache_read_input_token_cost": 2e-07,
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
"input_cost_per_token": 1.25e-06,
"input_cost_per_token_above_200k_tokens": 2.5e-06,
"litellm_provider": "xai",
"max_input_tokens": 1000000,
"max_output_tokens": 1000000,
"max_tokens": 1000000,
"mode": "chat",
"output_cost_per_token": 2.5e-06,
"output_cost_per_token_above_200k_tokens": 5e-06,
"source": "https://docs.x.ai/docs/models",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
"supports_web_search": true
},
"xai/grok-beta": {
"input_cost_per_token": 5e-06,
"litellm_provider": "xai",

View file

@ -0,0 +1,196 @@
"""
OAuth 2.0 Token Exchange (RFC 8693) handler for MCP servers.
Exchanges a user's incoming JWT (subject_token) for a scoped access token
at an IDP's token exchange endpoint. The exchanged token is then used to
authenticate requests to the upstream MCP server.
See: https://datatracker.ietf.org/doc/html/rfc8693
"""
import asyncio
import hashlib
import weakref
from typing import TYPE_CHECKING, Dict, Tuple
import httpx
from litellm._logging import verbose_logger
from litellm.caching.in_memory_cache import InMemoryCache
from litellm.constants import (
MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL,
MCP_OAUTH2_TOKEN_CACHE_MIN_TTL,
MCP_OAUTH2_TOKEN_EXPIRY_BUFFER_SECONDS,
MCP_TOKEN_EXCHANGE_CACHE_MAX_SIZE,
)
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.types.llms.custom_http import httpxSpecialProvider
if TYPE_CHECKING:
from litellm.types.mcp_server.mcp_server_manager import MCPServer
# RFC 8693 grant type constant
TOKEN_EXCHANGE_GRANT_TYPE = "urn:ietf:params:oauth:grant-type:token-exchange"
DEFAULT_SUBJECT_TOKEN_TYPE = "urn:ietf:params:oauth:token-type:access_token"
class TokenExchangeHandler:
"""Handles OAuth 2.0 Token Exchange (RFC 8693) for MCP servers.
Caches exchanged tokens keyed by ``hash(subject_token + server_id)`` so
repeated calls with the same user token skip the IDP round-trip.
"""
def __init__(self) -> None:
self._cache = InMemoryCache(
max_size_in_memory=MCP_TOKEN_EXCHANGE_CACHE_MAX_SIZE,
default_ttl=MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL,
)
# WeakValueDictionary so locks are GC'd once no coroutine holds a reference,
# preventing unbounded growth with many rotating user tokens.
self._locks: weakref.WeakValueDictionary[str, asyncio.Lock] = (
weakref.WeakValueDictionary()
)
def _get_lock(self, cache_key: str) -> asyncio.Lock:
lock = self._locks.get(cache_key)
if lock is None:
lock = asyncio.Lock()
self._locks[cache_key] = lock
return lock
@staticmethod
def _cache_key(subject_token: str, server_id: str) -> str:
raw = f"{subject_token}:{server_id}"
return hashlib.sha256(raw.encode()).hexdigest()
async def exchange_token(
self,
subject_token: str,
server: "MCPServer",
) -> str:
"""Exchange *subject_token* for a scoped access token.
Returns the exchanged ``access_token`` string (suitable for a
``Bearer`` header).
Raises ``ValueError`` on configuration or IDP errors.
"""
cache_key = self._cache_key(subject_token, server.server_id)
# Fast path
cached = self._cache.get_cache(cache_key)
if cached is not None:
return cached
# Slow path — one exchange at a time per (user, server) pair
async with self._get_lock(cache_key):
cached = self._cache.get_cache(cache_key)
if cached is not None:
return cached
token, ttl = await self._do_exchange(subject_token, server)
self._cache.set_cache(cache_key, token, ttl=ttl)
return token
async def _do_exchange(
self,
subject_token: str,
server: "MCPServer",
) -> Tuple[str, int]:
"""POST to the token exchange endpoint with RFC 8693 parameters.
Returns ``(access_token, ttl_seconds)``.
"""
endpoint = server.token_exchange_endpoint or server.token_url
if not endpoint:
raise ValueError(
f"MCP server '{server.server_id}' has auth_type=oauth2_token_exchange "
f"but no token_exchange_endpoint or token_url configured"
)
if not server.client_id or not server.client_secret:
raise ValueError(
f"MCP server '{server.server_id}' has auth_type=oauth2_token_exchange "
f"but missing client_id or client_secret"
)
data: Dict[str, str] = {
"grant_type": TOKEN_EXCHANGE_GRANT_TYPE,
"subject_token": subject_token,
"subject_token_type": server.subject_token_type
or DEFAULT_SUBJECT_TOKEN_TYPE,
"client_id": server.client_id,
"client_secret": server.client_secret,
}
if server.audience:
data["audience"] = server.audience
if server.scopes:
data["scope"] = " ".join(server.scopes)
verbose_logger.debug(
"Exchanging token for MCP server %s at %s (audience=%s)",
server.server_id,
endpoint,
server.audience,
)
client = get_async_httpx_client(llm_provider=httpxSpecialProvider.MCP)
try:
response = await client.post(endpoint, data=data)
response.raise_for_status()
except httpx.HTTPStatusError as exc:
verbose_logger.debug(
"Token exchange IDP error for MCP server %s (status %d)",
server.server_id,
exc.response.status_code,
)
raise ValueError(
f"Token exchange for MCP server '{server.server_id}' "
f"failed with status {exc.response.status_code}"
) from exc
body = response.json()
if not isinstance(body, dict):
raise ValueError(
f"Token exchange response for MCP server '{server.server_id}' "
f"returned non-object JSON (got {type(body).__name__})"
)
access_token = body.get("access_token")
if not access_token:
raise ValueError(
f"Token exchange response for MCP server '{server.server_id}' "
f"missing 'access_token'"
)
raw_expires_in = body.get("expires_in")
try:
expires_in = (
int(raw_expires_in)
if raw_expires_in is not None
else MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL
)
except (TypeError, ValueError):
expires_in = MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL
ttl = max(
expires_in - MCP_OAUTH2_TOKEN_EXPIRY_BUFFER_SECONDS,
MCP_OAUTH2_TOKEN_CACHE_MIN_TTL,
)
verbose_logger.info(
"Token exchange succeeded for MCP server %s (expires in %ds)",
server.server_id,
expires_in,
)
return access_token, ttl
def invalidate(self, subject_token: str, server_id: str) -> None:
"""Remove a cached exchanged token (e.g. after a 401)."""
cache_key = self._cache_key(subject_token, server_id)
self._cache.delete_cache(cache_key)
# Module-level singleton
mcp_token_exchange_handler = TokenExchangeHandler()

View file

@ -411,6 +411,15 @@ class MCPServerManager:
aws_role_name=server_config.get("aws_role_name", None),
aws_session_name=server_config.get("aws_session_name", None),
instructions=server_config.get("instructions", None),
# Token Exchange (OBO) fields
token_exchange_endpoint=server_config.get(
"token_exchange_endpoint", None
),
audience=server_config.get("audience", None),
subject_token_type=server_config.get(
"subject_token_type",
"urn:ietf:params:oauth:token-type:access_token",
),
)
self._assign_unique_short_prefix(new_server)
self.config_mcp_servers[server_id] = new_server
@ -766,6 +775,17 @@ class MCPServerManager:
aws_role_name=aws_creds.get("aws_role_name"),
aws_session_name=aws_creds.get("aws_session_name"),
instructions=mcp_server.instructions,
# Token Exchange (OBO) fields — read from credentials JSON blob
token_exchange_endpoint=(
credentials_dict.get("token_exchange_endpoint")
if credentials_dict
else None
),
audience=(credentials_dict.get("audience") if credentials_dict else None),
subject_token_type=(
credentials_dict.get("subject_token_type") if credentials_dict else None
)
or "urn:ietf:params:oauth:token-type:access_token",
)
return new_server
@ -1140,6 +1160,29 @@ class MCPServerManager:
#########################################################
# Methods that call the upstream MCP servers
#########################################################
@staticmethod
def _extract_bearer_token(
oauth2_headers: Optional[Dict[str, str]],
raw_headers: Optional[Dict[str, str]],
) -> Optional[str]:
"""Extract the bare Bearer token from oauth2_headers or raw_headers.
Returns the token string without the ``Bearer `` prefix, or ``None``
if no Authorization header is found.
"""
auth_value: Optional[str] = None
if oauth2_headers and "Authorization" in oauth2_headers:
auth_value = oauth2_headers["Authorization"]
elif raw_headers:
# raw_headers may have lowercase keys depending on the ASGI server
normalized = {k.lower(): v for k, v in raw_headers.items()}
auth_value = normalized.get("authorization")
if auth_value:
if auth_value.startswith("Bearer "):
return auth_value[len("Bearer ") :]
return auth_value
return None
def _build_stdio_env(
self,
server: MCPServer,
@ -1173,25 +1216,30 @@ class MCPServerManager:
mcp_auth_header: Optional[Union[str, Dict[str, str]]] = None,
extra_headers: Optional[Dict[str, str]] = None,
stdio_env: Optional[Dict[str, str]] = None,
subject_token: Optional[str] = None,
) -> MCPClient:
"""
Create an MCPClient instance for the given server.
Auth resolution (single place for all auth logic):
1. ``mcp_auth_header`` — per-request/per-user override
2. OAuth2 client_credentials token — auto-fetched and cached
3. ``server.authentication_token`` — static token from config/DB
2. OAuth2 Token Exchange (OBO) — exchange user token for scoped token
3. OAuth2 client_credentials token — auto-fetched and cached
4. ``server.authentication_token`` — static token from config/DB
Args:
server: The server configuration.
mcp_auth_header: Optional per-request auth override.
extra_headers: Additional headers to forward.
stdio_env: Environment variables for stdio transport.
subject_token: Optional user JWT for token exchange (OBO) flow.
Returns:
Configured MCP client instance.
"""
auth_value = await resolve_mcp_auth(server, mcp_auth_header)
auth_value = await resolve_mcp_auth(
server, mcp_auth_header, subject_token=subject_token
)
transport = server.transport or MCPTransport.sse
@ -2543,9 +2591,12 @@ class MCPServerManager:
if server_auth_header is None:
server_auth_header = mcp_auth_header
# oauth2 headers
# Extract subject token for OAuth2 Token Exchange (OBO) flow
subject_token: Optional[str] = None
extra_headers: Optional[Dict[str, str]] = None
if mcp_server.auth_type == MCPAuth.oauth2:
if mcp_server.auth_type == MCPAuth.oauth2_token_exchange:
subject_token = self._extract_bearer_token(oauth2_headers, raw_headers)
elif mcp_server.auth_type == MCPAuth.oauth2:
if mcp_server.has_client_credentials:
# For M2M OAuth servers, Authorization must come from token fetch.
extra_headers = None
@ -2613,6 +2664,7 @@ class MCPServerManager:
mcp_auth_header=server_auth_header,
extra_headers=extra_headers,
stdio_env=stdio_env,
subject_token=subject_token,
)
call_tool_params = MCPCallToolRequestParams(

View file

@ -26,6 +26,7 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import (
decrypt_value_helper,
encrypt_value_helper,
)
from litellm.proxy._experimental.mcp_server.auth import token_exchange
from litellm.types.llms.custom_http import httpxSpecialProvider
if TYPE_CHECKING:
@ -50,12 +51,23 @@ class MCPOAuth2TokenCache(InMemoryCache):
def _get_lock(self, server_id: str) -> asyncio.Lock:
return self._locks.setdefault(server_id, asyncio.Lock())
async def async_get_token(self, server: "MCPServer") -> Optional[str]:
@staticmethod
def _has_client_credentials_config(server: "MCPServer") -> bool:
return bool(server.client_id and server.client_secret and server.token_url)
async def async_get_token(
self,
server: "MCPServer",
*,
require_client_credentials_flow: bool = True,
) -> Optional[str]:
"""Return a valid access token, fetching or refreshing as needed.
Returns ``None`` when the server lacks client credentials config.
"""
if not server.has_client_credentials:
if require_client_credentials_flow and not server.has_client_credentials:
return None
if not self._has_client_credentials_config(server):
return None
server_id = server.server_id
@ -263,16 +275,38 @@ mcp_per_user_token_cache = MCPPerUserTokenCache()
async def resolve_mcp_auth(
server: "MCPServer",
mcp_auth_header: Optional[Union[str, Dict[str, str]]] = None,
subject_token: Optional[str] = None,
) -> Optional[Union[str, Dict[str, str]]]:
"""Resolve the auth value for an MCP server.
Priority:
1. ``mcp_auth_header`` — per-request/per-user override
2. OAuth2 client_credentials token — auto-fetched and cached
3. ``server.authentication_token`` — static token from config/DB
2. OAuth2 Token Exchange (OBO / RFC 8693) — exchange user token for scoped token
3. OAuth2 client_credentials token — auto-fetched and cached
4. ``server.authentication_token`` — static token from config/DB
"""
if mcp_auth_header:
return mcp_auth_header
if server.has_token_exchange_config:
if subject_token:
return await token_exchange.mcp_token_exchange_handler.exchange_token(
subject_token, server
)
# No subject_token — fall back to client_credentials using the same client
# credentials and token_url so M2M scenarios still work.
if server.client_id and server.client_secret and server.token_url:
return await mcp_oauth2_token_cache.async_get_token(
server,
require_client_credentials_flow=False,
)
# OBO configured but no subject_token and missing client credentials — warn
# rather than silently proceeding unauthenticated.
verbose_logger.warning(
"MCP server '%s' is configured for token exchange (OBO) but no subject_token "
"was provided and client credentials (client_id/client_secret/token_url) are "
"incomplete. The request will proceed without authentication.",
server.server_id,
)
if server.has_client_credentials:
return await mcp_oauth2_token_cache.async_get_token(server)
return server.authentication_token

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",
@ -656,6 +658,13 @@ class LiteLLMRoutes(enum.Enum):
"/health/services",
] + info_routes
# Stateless validators on caller-supplied log data; source logs are
# already accessible via spend_tracking_routes, so no scope expansion.
compliance_check_routes = [
"/compliance/eu-ai-act",
"/compliance/gdpr",
]
# Routes in `global_spend_tracking_routes` return proxy-wide spend across
# every team, customer, and api_key. They are intentionally NOT included
# here — non-admin roles must not see other tenants' spend. Admin roles go
@ -675,6 +684,7 @@ class LiteLLMRoutes(enum.Enum):
]
+ spend_tracking_routes
+ key_management_routes
+ compliance_check_routes
)
internal_user_view_only_routes = spend_tracking_routes
@ -699,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

@ -1512,7 +1512,7 @@ class ProxyBaseLLMRequestProcessing:
status_code=result.status_code,
headers=HttpPassThroughEndpointHelpers.get_response_headers(
headers=result.headers,
custom_headers=None,
custom_headers=dict(fastapi_response.headers),
),
)

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

@ -1,8 +0,0 @@
from fastapi import FastAPI
from litellm.proxy.health_endpoints._health_endpoints import router as health_router
def build_health_app():
health_app = FastAPI(title="LiteLLM Health Endpoints")
health_app.include_router(health_router)
return health_app

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

@ -46,25 +46,46 @@ def get_audit_log_changed_by(
def _resolve_audit_log_callback(name: str) -> Optional[CustomLogger]:
"""Resolve a string callback name to a CustomLogger instance, with caching."""
"""Resolve a string callback name to a CustomLogger instance, with caching.
For "s3_v2" with `litellm.s3_audit_callback_params` set, constructs a
dedicated `S3Logger` so audit logs can target a different bucket than the
normal-log singleton served by `_init_custom_logger_compatible_class`.
"""
if name in _audit_log_callback_cache:
return _audit_log_callback_cache[name]
from litellm.litellm_core_utils.litellm_logging import (
_init_custom_logger_compatible_class,
)
instance: Optional[CustomLogger]
if (
name == "s3_v2"
and getattr(litellm, "s3_audit_callback_params", None) is not None
):
from litellm.integrations.s3_v2 import S3Logger as S3V2Logger
instance = _init_custom_logger_compatible_class(
logging_integration=name, # type: ignore
internal_usage_cache=None,
llm_router=None,
)
instance = S3V2Logger(
s3_callback_params_override=litellm.s3_audit_callback_params
)
else:
from litellm.litellm_core_utils.litellm_logging import (
_init_custom_logger_compatible_class,
)
instance = _init_custom_logger_compatible_class(
logging_integration=name, # type: ignore
internal_usage_cache=None,
llm_router=None,
)
if instance is not None:
_audit_log_callback_cache[name] = instance
return instance
def reset_audit_log_callback_cache() -> None:
"""Clear cached audit-log callback instances. Call on config reload."""
_audit_log_callback_cache.clear()
def _build_audit_log_payload(
request_data: LiteLLM_AuditLogs,
) -> StandardAuditLogPayload:

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")
@ -1050,11 +1055,6 @@ def run_server( # noqa: PLR0915
litellm_settings=litellm_settings if config else None, # type: ignore[possibly-unbound]
)
# --- SEPARATE HEALTH APP LOGIC ---
# To run the health app separately, use:
# uvicorn litellm.proxy.health_app_factory:build_health_app --factory --host 0.0.0.0 --port=4001
# This is compatible with the SEPARATE_HEALTH_APP Docker/supervisord pattern.
# --- END SEPARATE HEALTH APP LOGIC ---
# Skip server startup if requested (after all setup is done)
if skip_server_startup:
print( # noqa

View file

@ -3823,6 +3823,11 @@ class ProxyConfig:
f"{blue_color_code} Initialized Failure Callbacks - {litellm.failure_callback} {reset_color_code}"
) # noqa
elif key == "audit_log_callbacks":
from litellm.proxy.management_helpers.audit_logs import (
reset_audit_log_callback_cache,
)
reset_audit_log_callback_cache()
litellm.audit_log_callbacks = []
for callback in value:
@ -3904,6 +3909,21 @@ class ProxyConfig:
f"{blue_color_code} setting litellm.{key}={value}{reset_color_code}"
)
setattr(litellm, key, value)
if key in {"s3_audit_callback_params", "s3_callback_params"}:
from litellm.proxy.management_helpers.audit_logs import (
reset_audit_log_callback_cache,
)
from litellm.litellm_core_utils.litellm_logging import (
_in_memory_loggers,
)
from litellm.integrations.s3_v2 import S3Logger as S3V2Logger
reset_audit_log_callback_cache()
_in_memory_loggers[:] = [
cb
for cb in _in_memory_loggers
if not isinstance(cb, S3V2Logger)
]
## GENERAL SERVER SETTINGS (e.g. master key,..) # do this after initializing litellm, to ensure sentry logging works for proxylogging
general_settings = config.get("general_settings", {})
@ -8808,6 +8828,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,
)
@ -2274,6 +2278,13 @@ class ProxyLogging:
Covers:
1. /chat/completions
"""
from litellm.proxy.proxy_server import llm_router
# Merge model-level guardrails before checking which guardrails to run
request_data = _check_and_merge_model_level_guardrails(
data=request_data, llm_router=llm_router
)
current_response = response
for callback in litellm.callbacks:
@ -2562,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
@ -2617,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"]:
@ -4265,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

@ -37,6 +37,7 @@ class MCPAuth(str, enum.Enum):
oauth2 = "oauth2"
aws_sigv4 = "aws_sigv4"
token = "token"
oauth2_token_exchange = "oauth2_token_exchange"
# MCP Literals
@ -54,6 +55,7 @@ MCPAuthType = Optional[
MCPAuth.oauth2,
MCPAuth.aws_sigv4,
MCPAuth.token,
MCPAuth.oauth2_token_exchange,
]
]
@ -117,6 +119,22 @@ class MCPCredentials(TypedDict, total=False):
aws_session_name: Optional[str]
"""Session name for STS AssumeRole (used in CloudTrail). Not a secret — stored unencrypted."""
audience: Optional[str]
"""
Target audience for OAuth 2.0 Token Exchange (RFC 8693)
"""
token_exchange_endpoint: Optional[str]
"""
IDP token endpoint for OAuth 2.0 Token Exchange (RFC 8693)
"""
subject_token_type: Optional[str]
"""
Subject token type for OAuth 2.0 Token Exchange (RFC 8693).
Default: urn:ietf:params:oauth:token-type:access_token
"""
class MCPServerCostInfo(TypedDict, total=False):
default_cost_per_query: Optional[float]

View file

@ -57,6 +57,10 @@ class MCPServer(BaseModel):
aws_service_name: Optional[str] = None # defaults to "bedrock-agentcore"
aws_role_name: Optional[str] = None # IAM role ARN for STS AssumeRole
aws_session_name: Optional[str] = None # session name for CloudTrail auditing
# Token Exchange (OBO) fields — RFC 8693
token_exchange_endpoint: Optional[str] = None
audience: Optional[str] = None
subject_token_type: str = "urn:ietf:params:oauth:token-type:access_token"
# Stdio-specific fields
command: Optional[str] = None
args: Optional[List[str]] = None
@ -127,3 +131,12 @@ class MCPServer(BaseModel):
return any(h.lower() in auth_header_names for h in self.extra_headers)
return False
@property
def has_token_exchange_config(self) -> bool:
"""True if this server is configured for OAuth2 token exchange (OBO / RFC 8693)."""
return (
self.auth_type == MCPAuth.oauth2_token_exchange
and bool(self.client_id and self.client_secret)
and bool(self.token_exchange_endpoint or self.token_url)
)

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",
@ -34945,6 +34959,48 @@
"supports_vision": true,
"supports_web_search": true
},
"xai/grok-4.3": {
"cache_read_input_token_cost": 2e-07,
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
"input_cost_per_token": 1.25e-06,
"input_cost_per_token_above_200k_tokens": 2.5e-06,
"litellm_provider": "xai",
"max_input_tokens": 1000000,
"max_output_tokens": 1000000,
"max_tokens": 1000000,
"mode": "chat",
"output_cost_per_token": 2.5e-06,
"output_cost_per_token_above_200k_tokens": 5e-06,
"source": "https://docs.x.ai/docs/models",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
"supports_web_search": true
},
"xai/grok-4.3-latest": {
"cache_read_input_token_cost": 2e-07,
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
"input_cost_per_token": 1.25e-06,
"input_cost_per_token_above_200k_tokens": 2.5e-06,
"litellm_provider": "xai",
"max_input_tokens": 1000000,
"max_output_tokens": 1000000,
"max_tokens": 1000000,
"mode": "chat",
"output_cost_per_token": 2.5e-06,
"output_cost_per_token_above_200k_tokens": 5e-06,
"source": "https://docs.x.ai/docs/models",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
"supports_web_search": true
},
"xai/grok-beta": {
"input_cost_per_token": 5e-06,
"litellm_provider": "xai",

View file

@ -1,6 +1,6 @@
[project]
name = "litellm"
version = "1.84.0"
version = "1.85.0"
description = "Library to easily interface with LLM API providers"
readme = "README.md"
requires-python = ">=3.10, <3.14"
@ -10,18 +10,22 @@ authors = [
{ name = "BerriAI" },
]
dependencies = [
"fastuuid==0.14.0",
"httpx==0.28.1",
"openai==2.33.0",
"python-dotenv==1.2.2",
"tiktoken==0.12.0",
"importlib-metadata==8.5.0",
"tokenizers==0.23.1",
"click==8.1.8",
"jinja2==3.1.6",
"aiohttp==3.13.4",
"pydantic==2.12.5",
"jsonschema==4.23.0",
# Ranges (not exact pins) so SDK consumers can coexist with their other
# deps. Reproducibility for our Docker/CI comes from `uv.lock`.
# When changing a floor, verify it installs + imports on every supported
# Python with: `uv pip install --resolution=lowest-direct .`
"fastuuid>=0.14.0,<1.0",
"httpx>=0.28.0,<1.0",
"openai>=2.20.0,<3.0.0",
"python-dotenv>=1.0.0,<2.0",
"tiktoken>=0.8.0,<1.0",
"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",
"aiohttp>=3.10,<4.0",
"pydantic>=2.10.0,<3.0.0",
"jsonschema>=4.0.0,<5.0",
]
[project.urls]
@ -246,7 +250,7 @@ source-exclude = [
profile = "black"
[tool.commitizen]
version = "1.84.0"
version = "1.85.0"
version_files = [
"pyproject.toml:^version",
]

View file

@ -7,7 +7,9 @@ from __future__ import annotations
import atexit
import hashlib
import json
import os
import re
import sys
from typing import Iterable
@ -74,6 +76,17 @@ FILTERED_RESPONSE_HEADERS = (
"date",
)
# Tiny placeholder used to replace base64 image payloads in cassettes.
# Decodes to b"test" — short, valid base64 so test code that decodes
# the field still succeeds.
VCR_IMAGE_B64_PLACEHOLDER = "dGVzdA=="
# Fixed boundary substituted into multipart request bodies so the
# ``safe_body`` matcher sees the same bytes across record and replay.
# httpx generates a fresh random boundary per request via os.urandom,
# which otherwise turns every multipart cassette into a permanent miss.
VCR_FIXED_MULTIPART_BOUNDARY = "vcr-static-boundary"
def _scrub_response(response):
if not isinstance(response, dict):
@ -86,8 +99,88 @@ def _scrub_response(response):
return response
def _replace_b64_json_in_place(obj) -> bool:
"""Recursively replace ``b64_json`` string values in a JSON tree.
Returns ``True`` if any value was rewritten. The check on the
existing value's length keeps the function idempotent — once a
value has been swapped to the placeholder, subsequent invocations
are no-ops.
"""
changed = False
if isinstance(obj, dict):
for key, value in obj.items():
if (
key == "b64_json"
and isinstance(value, str)
and len(value) > len(VCR_IMAGE_B64_PLACEHOLDER)
):
obj[key] = VCR_IMAGE_B64_PLACEHOLDER
changed = True
elif _replace_b64_json_in_place(value):
changed = True
elif isinstance(obj, list):
for item in obj:
if _replace_b64_json_in_place(item):
changed = True
return changed
def _strip_image_b64_payloads(response):
"""Replace ``b64_json`` payloads in image-gen responses before save.
Image-edit and image-generation responses carry the full base64
PNG/JPEG (1-10+ MB) in ``data[*].b64_json``. The image_gen tests
only assert response shape — the field decodes, schema validates —
they never inspect pixel content. Swapping to a 4-byte placeholder
preserves all those checks while shrinking cassettes by ~99%.
"""
if not isinstance(response, dict):
return response
body = response.get("body")
if not isinstance(body, dict):
return response
raw = body.get("string")
if raw is None:
return response
if isinstance(raw, (bytes, bytearray)):
try:
text = bytes(raw).decode("utf-8")
except UnicodeDecodeError:
return response
was_bytes = True
elif isinstance(raw, str):
text = raw
was_bytes = False
else:
return response
try:
payload = json.loads(text)
except (ValueError, TypeError):
return response
if not _replace_b64_json_in_place(payload):
return response
new_text = json.dumps(payload, separators=(",", ":"))
body["string"] = new_text.encode("utf-8") if was_bytes else new_text
headers = response.get("headers")
if isinstance(headers, dict):
new_len_value = str(len(new_text.encode("utf-8")))
for key in list(headers):
if str(key).lower() == "content-length":
value = headers[key]
headers[key] = (
[new_len_value] if isinstance(value, list) else new_len_value
)
return response
def _before_record_response(response):
return filter_non_2xx_response(_scrub_response(response))
return filter_non_2xx_response(_scrub_response(_strip_image_b64_payloads(response)))
def _safe_body_matcher(r1, r2) -> None:
@ -172,8 +265,84 @@ def _strip_headers(headers, names: Iterable[str]) -> None:
pass
def _normalize_multipart_boundary(request) -> None:
"""Rewrite random multipart boundaries to a fixed string in-place.
httpx generates a fresh ``boundary=<random hex>`` for every
multipart request via ``os.urandom``. Without normalization, the
request body bytes differ across runs even when everything else is
identical, the ``safe_body`` matcher misses, and the persister
keeps appending new episodes until ``MAX_EPISODES_PER_CASSETTE``
refuses the save — leaving audio-transcription tests effectively
unmocked. Replacing the boundary in both the Content-Type header
and the body bytes makes the request deterministic.
Idempotent — vcrpy invokes this hook multiple times per request,
so the second invocation sees ``boundary=vcr-static-boundary``
already and short-circuits.
"""
headers = getattr(request, "headers", None)
if headers is None:
return
content_type_key = None
content_type_value = None
try:
for key in list(headers.keys()):
if str(key).lower() == "content-type":
content_type_key = key
value = headers[key]
content_type_value = value if isinstance(value, str) else str(value)
break
except AttributeError:
return
if not content_type_value or "multipart/" not in content_type_value.lower():
return
fixed_param = f"boundary={VCR_FIXED_MULTIPART_BOUNDARY}"
if fixed_param in content_type_value:
return
match = re.search(r"boundary=([^\s;]+)", content_type_value)
if not match:
return
current_boundary = match.group(1).strip('"')
if current_boundary == VCR_FIXED_MULTIPART_BOUNDARY:
return
try:
headers[content_type_key] = content_type_value.replace(
match.group(0), fixed_param
)
except (TypeError, AttributeError):
return
body = getattr(request, "body", None)
if body is None:
return
if isinstance(body, (bytes, bytearray)):
try:
new_body = bytes(body).replace(
current_boundary.encode("utf-8"),
VCR_FIXED_MULTIPART_BOUNDARY.encode("utf-8"),
)
except (TypeError, ValueError):
return
elif isinstance(body, str):
new_body = body.replace(current_boundary, VCR_FIXED_MULTIPART_BOUNDARY)
else:
return
try:
request.body = new_body
except (AttributeError, TypeError):
pass
def _before_record_request(request):
"""Fingerprint API keys, then scrub them.
"""Fingerprint API keys, scrub them, and normalize multipart boundaries.
Order matters in two ways:
@ -187,7 +356,8 @@ def _before_record_request(request):
auth headers we already stripped, so re-hashing would yield
``"no-key"`` and the stored vs. incoming fingerprints would
diverge. Skip the recompute when the header is already set so
this hook is idempotent.
this hook is idempotent. The boundary normalizer is also
idempotent for the same reason.
"""
headers = getattr(request, "headers", None)
if headers is None:
@ -199,6 +369,7 @@ def _before_record_request(request):
except (TypeError, AttributeError):
pass
_strip_headers(headers, FILTERED_REQUEST_HEADERS)
_normalize_multipart_boundary(request)
return request

View file

@ -1,7 +1,6 @@
# conftest.py
import asyncio
import importlib
import os
import sys
@ -12,16 +11,6 @@ sys.path.insert(
) # Adds the parent directory to the system path
import litellm # noqa: E402,F401
from tests._vcr_conftest_common import ( # noqa: E402
VerboseReporterState,
apply_vcr_auto_marker_to_items,
record_vcr_outcome,
register_persister_if_enabled,
vcr_config_dict,
)
_verbose_state = VerboseReporterState()
@pytest.fixture(scope="session")
def event_loop():
@ -31,37 +20,3 @@ def event_loop():
loop = asyncio.new_event_loop()
yield loop
loop.close()
@pytest.fixture(scope="module")
def vcr_config():
return vcr_config_dict()
def pytest_recording_configure(config, vcr):
register_persister_if_enabled(vcr)
@pytest.hookimpl(hookwrapper=True)
def pytest_runtest_makereport(item, call):
outcome = yield
rep = outcome.get_result()
setattr(item, f"rep_{rep.when}", rep)
@pytest.fixture(autouse=True)
def _vcr_outcome_gate(request, vcr):
yield
record_vcr_outcome(request, vcr)
def pytest_configure(config):
_verbose_state.remember_pluginmanager(config)
def pytest_runtest_logreport(report):
_verbose_state.maybe_emit_verdict(report)
def pytest_collection_modifyitems(config, items):
apply_vcr_auto_marker_to_items(items)

View file

@ -47,6 +47,37 @@ class TestCustomLogger(CustomLogger):
self.standard_logging_object = kwargs["standard_logging_object"]
async def _acreate_fine_tuning_job_with_propagation_retry(
*, max_attempts: int = 12, initial_delay: float = 1.0, **kwargs
):
"""
Wrap litellm.acreate_fine_tuning_job and retry on the eventual-consistency
400 OpenAI returns when a freshly-uploaded training file isn't yet visible
to the fine-tuning endpoint (`'file-... does not exist'`).
Polling the files-retrieve endpoint or `FileObject.status` doesn't help —
OpenAI's `status` field is deprecated, and the retrieve and fine-tuning
endpoints don't share a consistency model. Retrying the operation itself
is the only reliable signal that propagation has finished.
Total budget with defaults: ~70s across 12 attempts (exp backoff capped at
8s).
"""
delay = initial_delay
last_error: Optional[openai.BadRequestError] = None
for _ in range(max_attempts):
try:
return await litellm.acreate_fine_tuning_job(**kwargs)
except openai.BadRequestError as e:
if "does not exist" not in str(e):
raise
last_error = e
await asyncio.sleep(delay)
delay = min(delay * 1.5, 8.0)
assert last_error is not None
raise last_error
@pytest.mark.asyncio
async def test_create_fine_tune_jobs_async():
try:
@ -64,9 +95,11 @@ async def test_create_fine_tune_jobs_async():
)
print("Response from creating file=", file_obj)
create_fine_tuning_response = await litellm.acreate_fine_tuning_job(
model="gpt-3.5-turbo-0125",
training_file=file_obj.id,
create_fine_tuning_response = (
await _acreate_fine_tuning_job_with_propagation_retry(
model="gpt-4o-mini-2024-07-18",
training_file=file_obj.id,
)
)
print(
@ -74,7 +107,7 @@ async def test_create_fine_tune_jobs_async():
)
assert create_fine_tuning_response.id is not None
assert create_fine_tuning_response.model == "gpt-3.5-turbo-0125"
assert create_fine_tuning_response.model == "gpt-4o-mini-2024-07-18"
await asyncio.sleep(2)
_logged_standard_logging_object = custom_logger.standard_logging_object
@ -83,7 +116,7 @@ async def test_create_fine_tune_jobs_async():
"custom_logger.standard_logging_object=",
json.dumps(_logged_standard_logging_object, indent=4),
)
assert _logged_standard_logging_object["model"] == "gpt-3.5-turbo-0125"
assert _logged_standard_logging_object["model"] == "gpt-4o-mini-2024-07-18"
assert _logged_standard_logging_object["id"] == create_fine_tuning_response.id
# list fine tuning jobs
@ -427,10 +460,10 @@ async def test_mock_openai_create_fine_tune_job():
with patch.object(client.fine_tuning.jobs, "create") as mock_create:
mock_create.return_value = FineTuningJob(
id="ft-123",
model="gpt-3.5-turbo-0125",
model="gpt-4o-mini-2024-07-18",
created_at=1677610602,
status="validating_files",
fine_tuned_model="ft:gpt-3.5-turbo-0125:org:custom_suffix:id",
fine_tuned_model="ft:gpt-4o-mini-2024-07-18:org:custom_suffix:id",
object="fine_tuning.job",
hyperparameters=Hyperparameters(
n_epochs=3,
@ -442,7 +475,7 @@ async def test_mock_openai_create_fine_tune_job():
)
response = await litellm.acreate_fine_tuning_job(
model="gpt-3.5-turbo-0125",
model="gpt-4o-mini-2024-07-18",
training_file="file-123",
hyperparameters={"n_epochs": 3},
suffix="custom_suffix",
@ -453,16 +486,19 @@ async def test_mock_openai_create_fine_tune_job():
mock_create.assert_called_once()
request_params = mock_create.call_args.kwargs
assert request_params["model"] == "gpt-3.5-turbo-0125"
assert request_params["model"] == "gpt-4o-mini-2024-07-18"
assert request_params["training_file"] == "file-123"
assert request_params["hyperparameters"] == {"n_epochs": 3}
assert request_params["suffix"] == "custom_suffix"
# Verify the response
assert response.id == "ft-123"
assert response.model == "gpt-3.5-turbo-0125"
assert response.model == "gpt-4o-mini-2024-07-18"
assert response.status == "validating_files"
assert response.fine_tuned_model == "ft:gpt-3.5-turbo-0125:org:custom_suffix:id"
assert (
response.fine_tuned_model
== "ft:gpt-4o-mini-2024-07-18:org:custom_suffix:id"
)
@pytest.mark.asyncio

View file

@ -853,7 +853,11 @@ class BaseLLMChatTest(ABC):
@pytest.mark.parametrize(
"image_url",
[
"http://img1.etsystatic.com/260/0/7813604/il_fullxfull.4226713999_q86e.jpg",
# In-repo logo served via jsdelivr (sha-pinned, immutable).
# Bedrock fetches the URL and base64-embeds it in the
# Converse request body; using a multi-MB hosted product
# photo here previously bloated cassettes to ~60 MB each.
"https://cdn.jsdelivr.net/gh/BerriAI/litellm@d769e81c90d453240c61fc572cdb27fae06a89d0/ui/litellm-dashboard/public/assets/logos/litellm_logo.jpg",
"https://awsmp-logos.s3.amazonaws.com/seller-xw5kijmvmzasy/c233c9ade2ccb5491072ae232c814942.png",
],
)

View file

@ -101,7 +101,7 @@ async def test_openai_realtime_direct_call_no_intent():
try:
await litellm._arealtime(
model="openai/gpt-4o-realtime-preview-2024-10-01",
model="openai/gpt-4o-realtime-preview",
websocket=websocket_client,
api_key=os.environ.get("OPENAI_API_KEY"),
timeout=60,
@ -250,13 +250,13 @@ async def test_openai_realtime_direct_call_with_intent():
caught_exception = None
query_params: RealtimeQueryParams = {
"model": "openai/gpt-4o-realtime-preview-2024-10-01",
"model": "openai/gpt-4o-realtime-preview",
"intent": "chat",
}
try:
await litellm._arealtime(
model="openai/gpt-4o-realtime-preview-2024-10-01",
model="openai/gpt-4o-realtime-preview",
websocket=websocket_client,
api_key=os.environ.get("OPENAI_API_KEY"),
query_params=query_params,
@ -331,7 +331,7 @@ def test_realtime_query_params_construction():
from litellm.types.realtime import RealtimeQueryParams
# Test case 1: intent is None (should not be included)
model = "gpt-4o-realtime-preview-2024-10-01"
model = "gpt-4o-realtime-preview"
intent = None
query_params: RealtimeQueryParams = {"model": model}
@ -369,17 +369,17 @@ async def test_realtime_query_params_use_normalized_model_name(monkeypatch):
)
def fake_get_llm_provider(model, api_base=None, api_key=None):
return ("gpt-4o-realtime-preview-2024-10-01", "openai", None, None)
return ("gpt-4o-realtime-preview", "openai", None, None)
monkeypatch.setattr(realtime_main, "get_llm_provider", fake_get_llm_provider)
query_params: RealtimeQueryParams = {
"model": "openai/gpt-4o-realtime-preview-2024-10-01",
"model": "openai/gpt-4o-realtime-preview",
"intent": "chat",
}
await realtime_main._arealtime(
model="openai/gpt-4o-realtime-preview-2024-10-01",
model="openai/gpt-4o-realtime-preview",
websocket=MagicMock(),
api_key="sk-test",
query_params=query_params,
@ -387,7 +387,5 @@ async def test_realtime_query_params_use_normalized_model_name(monkeypatch):
)
called_kwargs = mock_async_realtime.call_args.kwargs
assert (
called_kwargs["query_params"]["model"] == "gpt-4o-realtime-preview-2024-10-01"
)
assert called_kwargs["query_params"]["model"] == "gpt-4o-realtime-preview"
assert called_kwargs["query_params"]["intent"] == "chat"

View file

@ -2,6 +2,7 @@
Tests for Evals API operations across providers
"""
import hashlib
import os
import sys
from abc import ABC, abstractmethod
@ -20,6 +21,46 @@ from litellm.types.llms.openai_evals import (
)
def _stable_eval_name(test_node_name: str, suffix: str = "") -> str:
"""Deterministic eval name keyed off the test's node name.
The previous ``f"Test Eval {int(time.time())}"`` pattern embedded a
fresh value into the request body every run, defeating VCR's
``safe_body`` matcher and forcing a real OpenAI ``create`` call on
every CI run. With a stable per-test name the cassette matches on
replay, and provider-side resources stay bounded because each test
deletes the eval it owns on teardown.
"""
nonce = hashlib.sha1(test_node_name.encode()).hexdigest()[:12]
return f"vcr-managed-{nonce}{suffix}"
_TESTING_CRITERIA = [
{
"type": "label_model",
"model": "gpt-4o",
"input": [
{
"role": "developer",
"content": "Classify the sentiment as 'positive' or 'negative'",
},
{"role": "user", "content": "Statement: {{item.input}}"},
],
"passing_labels": ["positive"],
"labels": ["positive", "negative"],
"name": "Sentiment grader",
}
]
_PROVIDER_FLAKINESS = (
litellm.InternalServerError,
litellm.APIConnectionError,
litellm.Timeout,
litellm.ServiceUnavailableError,
)
class BaseEvalsAPITest(ABC):
"""
Base test class for Evals API operations.
@ -41,13 +82,64 @@ class BaseEvalsAPITest(ABC):
"""Return the API base URL for the provider"""
pass
@pytest.fixture
def managed_eval(self, request):
"""Create a stable-named eval for this test; delete on teardown.
Function-scoped so each cassette captures the full
create→test→delete cycle. A class-scoped fixture would push
the create into whichever test ran first and the delete into
whichever ran last, which is fragile under reordering.
Replaces the prior ``list_evals().data[0].id`` pattern, which
made the URL of ``get_eval`` / ``update_eval`` vary across
runs (the "first" eval depends on what other runs left
behind).
"""
custom_llm_provider = self.get_custom_llm_provider()
api_key = self.get_api_key()
api_base = self.get_api_base()
if not api_key:
pytest.skip(f"No API key provided for {custom_llm_provider}")
try:
created = litellm.create_eval(
name=_stable_eval_name(request.node.name),
data_source_config={
"type": "stored_completions",
"metadata": {"usecase": "chatbot", "vcr": "managed"},
},
testing_criteria=_TESTING_CRITERIA,
custom_llm_provider=custom_llm_provider,
api_key=api_key,
api_base=api_base,
)
except _PROVIDER_FLAKINESS:
pytest.skip("Provider service unavailable")
except litellm.RateLimitError:
pytest.skip("Rate limit exceeded")
yield created
# Best-effort cleanup. OpenAI eval names are not unique-keyed
# (only IDs are), so a failed delete doesn't block the next
# run's create.
try:
litellm.delete_eval(
eval_id=created.id,
custom_llm_provider=custom_llm_provider,
api_key=api_key,
api_base=api_base,
)
except Exception:
pass
@pytest.mark.flaky(retries=3, delay=2)
def test_create_eval(self):
def test_create_eval(self, request):
"""
Test creating an evaluation.
"""
import time
custom_llm_provider = self.get_custom_llm_provider()
api_key = self.get_api_key()
api_base = self.get_api_base()
@ -56,53 +148,45 @@ class BaseEvalsAPITest(ABC):
pytest.skip(f"No API key provided for {custom_llm_provider}")
litellm.set_verbose = True
unique_name = _stable_eval_name(request.node.name)
# Create eval with stored_completions data source
unique_name = f"Test Eval {int(time.time())}"
created_id = None
try:
response = litellm.create_eval(
name=unique_name,
data_source_config={
"type": "stored_completions",
"metadata": {"usecase": "chatbot"},
},
testing_criteria=[
{
"type": "label_model",
"model": "gpt-4o",
"input": [
{
"role": "developer",
"content": "Classify the sentiment as 'positive' or 'negative'",
},
{"role": "user", "content": "Statement: {{item.input}}"},
],
"passing_labels": ["positive"],
"labels": ["positive", "negative"],
"name": "Sentiment grader",
}
],
custom_llm_provider=custom_llm_provider,
api_key=api_key,
api_base=api_base,
)
except (
litellm.InternalServerError,
litellm.APIConnectionError,
litellm.Timeout,
litellm.ServiceUnavailableError,
):
pytest.skip("Provider service unavailable")
except litellm.RateLimitError:
pytest.skip("Rate limit exceeded")
try:
response = litellm.create_eval(
name=unique_name,
data_source_config={
"type": "stored_completions",
"metadata": {"usecase": "chatbot"},
},
testing_criteria=_TESTING_CRITERIA,
custom_llm_provider=custom_llm_provider,
api_key=api_key,
api_base=api_base,
)
except _PROVIDER_FLAKINESS:
pytest.skip("Provider service unavailable")
except litellm.RateLimitError:
pytest.skip("Rate limit exceeded")
assert response is not None
assert isinstance(response, Eval)
assert response.id is not None
assert response.name == unique_name
print(f"Created eval: {response}")
print(f"Eval ID: {response.id}")
assert response is not None
assert isinstance(response, Eval)
assert response.id is not None
assert response.name == unique_name
created_id = response.id
print(f"Created eval: {response}")
print(f"Eval ID: {response.id}")
finally:
if created_id is not None:
try:
litellm.delete_eval(
eval_id=created_id,
custom_llm_provider=custom_llm_provider,
api_key=api_key,
api_base=api_base,
)
except Exception:
pass
def test_list_evals(self):
"""
@ -130,7 +214,7 @@ class BaseEvalsAPITest(ABC):
assert hasattr(response, "has_more")
print(f"Listed evals: {len(response.data)} evaluations")
def test_get_eval(self):
def test_get_eval(self, managed_eval):
"""
Test getting a specific evaluation by ID.
"""
@ -138,89 +222,54 @@ class BaseEvalsAPITest(ABC):
api_key = self.get_api_key()
api_base = self.get_api_base()
if not api_key:
pytest.skip(f"No API key provided for {custom_llm_provider}")
litellm.set_verbose = True
# First list existing evals to get an ID
list_response = litellm.list_evals(
limit=1,
response = litellm.get_eval(
eval_id=managed_eval.id,
custom_llm_provider=custom_llm_provider,
api_key=api_key,
api_base=api_base,
)
assert isinstance(list_response, ListEvalsResponse)
assert response is not None
assert isinstance(response, Eval)
assert response.id == managed_eval.id
print(f"Retrieved eval: {response}")
if list_response.data and len(list_response.data) > 0:
eval_id = list_response.data[0].id
print(f"Testing with eval ID: {eval_id}")
# Get the eval
response = litellm.get_eval(
eval_id=eval_id,
custom_llm_provider=custom_llm_provider,
api_key=api_key,
api_base=api_base,
)
assert response is not None
assert isinstance(response, Eval)
assert response.id == eval_id
print(f"Retrieved eval: {response}")
else:
pytest.skip("No existing evals to test with")
def test_update_eval(self):
@pytest.mark.flaky(retries=3, delay=2)
def test_update_eval(self, request, managed_eval):
"""
Test updating an evaluation.
"""
import time
custom_llm_provider = self.get_custom_llm_provider()
api_key = self.get_api_key()
api_base = self.get_api_base()
if not api_key:
pytest.skip(f"No API key provided for {custom_llm_provider}")
litellm.set_verbose = True
updated_name = _stable_eval_name(request.node.name, suffix="-updated")
# First list existing evals
list_response = litellm.list_evals(
limit=1,
response = litellm.update_eval(
eval_id=managed_eval.id,
name=updated_name,
custom_llm_provider=custom_llm_provider,
api_key=api_key,
api_base=api_base,
)
assert isinstance(list_response, ListEvalsResponse)
if list_response.data and len(list_response.data) > 0:
eval_id = list_response.data[0].id
updated_name = f"Updated Eval {int(time.time())}"
# Update the eval
response = litellm.update_eval(
eval_id=eval_id,
name=updated_name,
custom_llm_provider=custom_llm_provider,
api_key=api_key,
api_base=api_base,
)
assert response is not None
assert isinstance(response, Eval)
assert response.id == eval_id
assert response.name == updated_name
print(f"Updated eval: {response}")
else:
pytest.skip("No existing evals to test with")
assert response is not None
assert isinstance(response, Eval)
assert response.id == managed_eval.id
assert response.name == updated_name
print(f"Updated eval: {response}")
def test_delete_eval(self):
"""
Test deleting an evaluation.
Real delete coverage now lives in the ``managed_eval`` fixture
teardown and in ``test_create_eval``'s ``finally`` block, so
this stays a no-op skip rather than creating a fresh resource
just to delete it.
"""
custom_llm_provider = self.get_custom_llm_provider()
api_key = self.get_api_key()
@ -229,8 +278,7 @@ class BaseEvalsAPITest(ABC):
if not api_key:
pytest.skip(f"No API key provided for {custom_llm_provider}")
# Skip this test to avoid deleting production evals
pytest.skip("Skipping delete test to preserve existing evals")
pytest.skip("Delete is exercised via managed_eval fixture teardown.")
class TestOpenAIEvalsAPI(BaseEvalsAPITest):

View file

@ -0,0 +1,220 @@
"""Unit tests for the VCR record-time filters that keep cassettes small.
Covers:
- ``_strip_image_b64_payloads`` — replaces base64 image bodies in
image-gen responses so cassettes don't carry MB-class PNG payloads.
- ``_normalize_multipart_boundary`` — rewrites random multipart
boundaries to a fixed string so audio-transcription request bodies
match across record and replay.
"""
from __future__ import annotations
import json
import os
import sys
from vcr.request import Request
sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "..", "..")))
from tests._vcr_conftest_common import ( # noqa: E402
VCR_FIXED_MULTIPART_BOUNDARY,
VCR_IMAGE_B64_PLACEHOLDER,
_normalize_multipart_boundary,
_strip_image_b64_payloads,
)
# ---------------------------------------------------------------------------
# Image b64 stripper
# ---------------------------------------------------------------------------
def _image_response(b64_payload: str, body_type: str = "bytes") -> dict:
body_text = json.dumps({"data": [{"b64_json": b64_payload}]})
body_string = body_text.encode("utf-8") if body_type == "bytes" else body_text
return {
"status": {"code": 200, "message": "OK"},
"headers": {
"content-type": ["application/json"],
"content-length": [str(len(body_text.encode("utf-8")))],
},
"body": {"string": body_string},
}
def test_strip_image_b64_replaces_payload_when_body_is_bytes():
response = _image_response("A" * 5000, body_type="bytes")
out = _strip_image_b64_payloads(response)
payload = json.loads(out["body"]["string"].decode("utf-8"))
assert payload["data"][0]["b64_json"] == VCR_IMAGE_B64_PLACEHOLDER
def test_strip_image_b64_replaces_payload_when_body_is_str():
response = _image_response("A" * 5000, body_type="str")
out = _strip_image_b64_payloads(response)
payload = json.loads(out["body"]["string"])
assert payload["data"][0]["b64_json"] == VCR_IMAGE_B64_PLACEHOLDER
def test_strip_image_b64_updates_content_length():
response = _image_response("A" * 5000)
out = _strip_image_b64_payloads(response)
expected_len = len(out["body"]["string"])
assert out["headers"]["content-length"] == [str(expected_len)]
def test_strip_image_b64_is_idempotent():
response = _image_response("A" * 5000)
once = _strip_image_b64_payloads(response)
twice = _strip_image_b64_payloads(once)
assert once["body"]["string"] == twice["body"]["string"]
def test_strip_image_b64_handles_nested_data():
body_text = json.dumps(
{
"outer": {
"data": [
{"b64_json": "X" * 4000, "label": "first"},
{"b64_json": "Y" * 4000, "label": "second"},
]
}
}
)
response = {
"status": {"code": 200, "message": "OK"},
"headers": {"content-type": ["application/json"]},
"body": {"string": body_text.encode("utf-8")},
}
out = _strip_image_b64_payloads(response)
payload = json.loads(out["body"]["string"].decode("utf-8"))
assert payload["outer"]["data"][0]["b64_json"] == VCR_IMAGE_B64_PLACEHOLDER
assert payload["outer"]["data"][1]["b64_json"] == VCR_IMAGE_B64_PLACEHOLDER
assert payload["outer"]["data"][0]["label"] == "first"
def test_strip_image_b64_leaves_non_image_response_unchanged():
body_text = json.dumps({"choices": [{"message": {"content": "hello"}}]})
response = {
"status": {"code": 200, "message": "OK"},
"headers": {"content-type": ["application/json"]},
"body": {"string": body_text.encode("utf-8")},
}
out = _strip_image_b64_payloads(response)
assert json.loads(out["body"]["string"].decode("utf-8")) == json.loads(body_text)
def test_strip_image_b64_leaves_invalid_json_unchanged():
response = {
"status": {"code": 200, "message": "OK"},
"headers": {"content-type": ["application/octet-stream"]},
"body": {"string": b"\x89PNG\r\n\x1a\n binary stuff not json"},
}
out = _strip_image_b64_payloads(response)
assert out["body"]["string"] == b"\x89PNG\r\n\x1a\n binary stuff not json"
def test_strip_image_b64_skips_short_values():
"""Already-placeholder values aren't re-replaced (idempotency guard)."""
body_text = json.dumps({"data": [{"b64_json": VCR_IMAGE_B64_PLACEHOLDER}]})
response = {
"status": {"code": 200, "message": "OK"},
"headers": {"content-type": ["application/json"]},
"body": {"string": body_text.encode("utf-8")},
}
out = _strip_image_b64_payloads(response)
payload = json.loads(out["body"]["string"].decode("utf-8"))
assert payload["data"][0]["b64_json"] == VCR_IMAGE_B64_PLACEHOLDER
# ---------------------------------------------------------------------------
# Multipart boundary normalizer
# ---------------------------------------------------------------------------
def _multipart_request(boundary: str):
body_text = (
f"--{boundary}\r\n"
'Content-Disposition: form-data; name="file"; filename="audio.wav"\r\n'
"Content-Type: audio/wav\r\n"
"\r\n"
"fake-audio-bytes\r\n"
f"--{boundary}--\r\n"
)
return Request(
method="POST",
uri="https://api.openai.com/v1/audio/transcriptions",
body=body_text.encode("utf-8"),
headers={
"content-type": f"multipart/form-data; boundary={boundary}",
},
)
def test_normalize_multipart_rewrites_header_and_body():
req = _multipart_request("abc123random")
_normalize_multipart_boundary(req)
assert (
req.headers["content-type"]
== f"multipart/form-data; boundary={VCR_FIXED_MULTIPART_BOUNDARY}"
)
assert b"abc123random" not in req.body
assert VCR_FIXED_MULTIPART_BOUNDARY.encode("utf-8") in req.body
def test_normalize_multipart_is_idempotent():
req = _multipart_request("abc123random")
_normalize_multipart_boundary(req)
body_first = req.body
header_first = req.headers["content-type"]
_normalize_multipart_boundary(req)
assert req.body == body_first
assert req.headers["content-type"] == header_first
def test_normalize_multipart_two_distinct_boundaries_match_after_normalize():
"""Whisper-style: two requests with different random boundaries should
end up with byte-identical bodies after normalization."""
req1 = _multipart_request("boundaryAAA")
req2 = _multipart_request("boundaryBBB")
_normalize_multipart_boundary(req1)
_normalize_multipart_boundary(req2)
assert req1.body == req2.body
assert req1.headers["content-type"] == req2.headers["content-type"]
def test_normalize_multipart_skips_non_multipart_requests():
req = Request(
method="POST",
uri="https://api.openai.com/v1/chat/completions",
body=b'{"model":"gpt-4o"}',
headers={"content-type": "application/json"},
)
_normalize_multipart_boundary(req)
assert req.headers["content-type"] == "application/json"
assert req.body == b'{"model":"gpt-4o"}'
def test_normalize_multipart_skips_request_without_content_type():
req = Request(
method="POST",
uri="https://api.openai.com/v1/chat/completions",
body=b"unknown body",
headers={},
)
_normalize_multipart_boundary(req)
assert req.body == b"unknown body"
def test_normalize_multipart_handles_quoted_boundary():
req = Request(
method="POST",
uri="https://api.openai.com/v1/audio/transcriptions",
body=b"--quoted-boundary--body content--quoted-boundary--",
headers={"content-type": 'multipart/form-data; boundary="quoted-boundary"'},
)
_normalize_multipart_boundary(req)
assert b"quoted-boundary" not in req.body
assert VCR_FIXED_MULTIPART_BOUNDARY.encode("utf-8") in req.body

View file

@ -101,6 +101,7 @@ _SCALAR_ATTRS = (
"redact_messages_in_exceptions",
"redact_user_api_key_info",
"s3_callback_params",
"s3_audit_callback_params",
"datadog_params",
"vector_store_registry",
)
@ -128,6 +129,7 @@ def isolate_litellm_state():
leaking across tests within the same xdist worker.
"""
from litellm.litellm_core_utils import litellm_logging as ll_logging
from litellm.proxy.management_helpers import audit_logs as ll_audit_logs
# Flush cache and clear internal logger instances before test
if hasattr(litellm, "in_memory_llm_clients_cache"):
@ -135,6 +137,7 @@ def isolate_litellm_state():
# Clear cached logger instances (LangsmithLogger, SlackAlerting, etc.)
ll_logging._in_memory_loggers.clear()
ll_audit_logs._audit_log_callback_cache.clear()
# Reset ALL attrs to their true defaults before the test runs.
# This undoes any module-level mutations from test file imports.
@ -156,6 +159,7 @@ def isolate_litellm_state():
litellm.in_memory_llm_clients_cache.flush_cache()
ll_logging._in_memory_loggers.clear()
ll_audit_logs._audit_log_callback_cache.clear()
for attr in _LIST_ATTRS:
if attr in _DEFAULTS:

View file

@ -12,7 +12,15 @@ from abc import ABC, abstractmethod
# Test resources
TEST_IMAGE_PATH = "test_image_edit.png"
TEST_PDF_URL = "https://arxiv.org/pdf/2201.04234"
# Tiny in-repo PDF served via jsdelivr (sha-pinned, immutable). The arxiv
# PDF previously used here was several MB — once base64-encoded into the
# Vertex OCR request it ballooned cassettes past 100 MB per test. Keep
# the URL stable across runs so cassettes don't churn.
TEST_PDF_URL = (
"https://cdn.jsdelivr.net/gh/BerriAI/litellm"
"@d769e81c90d453240c61fc572cdb27fae06a89d0"
"/tests/llm_translation/fixtures/dummy.pdf"
)
class BaseOCRTest(ABC):

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

@ -1123,3 +1123,74 @@ async def test_combined_prefix_reflects_in_s3_object_key():
result = logger.create_s3_batch_logging_element(datetime.utcnow(), payload)
key = result.s3_object_key
assert "myteam/apikey/" in key, f"Expected both prefixes in key: {key}"
# --------------------------------------------------------------
# params_source / s3_callback_params_override (audit-log decoupling)
# --------------------------------------------------------------
def test_s3_callback_params_override_uses_alternate_dict():
"""`s3_callback_params_override` makes the logger read its config from
the override dict instead of `litellm.s3_callback_params`."""
import litellm
original = litellm.s3_callback_params
litellm.s3_callback_params = {"s3_bucket_name": "normal-bucket"}
try:
logger = S3Logger(
s3_callback_params_override={
"s3_bucket_name": "audit-bucket",
"s3_path": "audit-prefix",
"s3_region_name": "us-west-2",
}
)
assert logger.s3_bucket_name == "audit-bucket"
assert logger.s3_path == "audit-prefix"
assert logger.s3_region_name == "us-west-2"
finally:
litellm.s3_callback_params = original
def test_s3_callback_params_override_does_not_mutate_inputs(monkeypatch):
"""Resolving `os.environ/X` markers must not mutate the override dict
or `litellm.s3_callback_params`."""
import litellm
monkeypatch.setenv("MY_AUDIT_BUCKET", "resolved-bucket")
override = {"s3_bucket_name": "os.environ/MY_AUDIT_BUCKET"}
original_global = litellm.s3_callback_params
litellm.s3_callback_params = {"s3_bucket_name": "os.environ/MY_AUDIT_BUCKET"}
try:
logger = S3Logger(s3_callback_params_override=override)
assert logger.s3_bucket_name == "resolved-bucket"
assert override["s3_bucket_name"] == "os.environ/MY_AUDIT_BUCKET"
assert (
litellm.s3_callback_params["s3_bucket_name"] == "os.environ/MY_AUDIT_BUCKET"
)
finally:
litellm.s3_callback_params = original_global
def test_s3_callback_params_override_none_falls_back_to_global():
"""No override → behaves exactly as today (reads `litellm.s3_callback_params`)."""
import litellm
original = litellm.s3_callback_params
litellm.s3_callback_params = {"s3_bucket_name": "from-global"}
try:
logger = S3Logger()
assert logger.s3_bucket_name == "from-global"
finally:
litellm.s3_callback_params = original
def test_s3_callback_params_override_empty_dict_is_opt_in():
"""An empty override dict skips the global entirely (env/IAM-only config)."""
import litellm
original = litellm.s3_callback_params
litellm.s3_callback_params = {"s3_bucket_name": "from-global"}
try:
logger = S3Logger(s3_callback_params_override={})
assert logger.s3_bucket_name is None
finally:
litellm.s3_callback_params = original

View file

@ -157,14 +157,15 @@ class TestResponseCompliance:
# Check CreateModelInteractionParams which includes output fields
schema = spec_dict["components"]["schemas"]["CreateModelInteractionParams"]
# Output fields (readOnly)
# Output fields (readOnly). Google renamed `outputs` → `steps` in the
# upstream spec; keep this list aligned with the live schema.
output_fields = [
"id",
"status",
"created",
"updated",
"role",
"outputs",
"steps",
"usage",
]

View file

@ -367,3 +367,128 @@ def test_update_messages_with_model_file_ids_skips_non_openai_file_blocks():
# Messages pass through unchanged when there is no `file` sub-dict to remap.
assert updated == messages
# Reusable fixture (decodes to: litellm_proxy:application/pdf;unified_id,...;
# target_model_names,gpt-4o;llm_output_file_id,file-ECBPW7ML9g7XHdwGgUPZaM;
# llm_output_file_model_id,...)
UNIFIED_FILE_ID_B64 = (
"bGl0ZWxsbV9wcm94eTphcHBsaWNhdGlvbi9wZGY7dW5pZmllZF9pZCw2YzBiNTg5MC04OTE0"
"LTQ4ZTAtYjhmNC0wYWU1ZWQzYzE0YTU7dGFyZ2V0X21vZGVsX25hbWVzLGdwdC00bztsbG1f"
"b3V0cHV0X2ZpbGVfaWQsZmlsZS1FQ0JQVzdNTDlnN1hIZHdHZ1VQWmFNO2xsbV9vdXRwdXRf"
"ZmlsZV9tb2RlbF9pZCxlMjY0NTNmOWU3NmU3OTkzNjgwZDAwNjhkOThjMWY0Y2MyMDViYmFk"
"MDk2N2EzM2M2NjQ4OTM1NjhjYTc0M2My"
)
def test_update_messages_with_model_file_ids_decodes_unified_id_when_mapping_empty():
"""When the mapping is empty (e.g. multi-replica cache miss), the function
must decode the base64-encoded unified file id and substitute the embedded
llm_output_file_id — mirroring the Responses-API sibling."""
messages = [
{
"role": "user",
"content": [
{"type": "text", "text": "What is in this recording?"},
{
"type": "file",
"file": {
"file_id": UNIFIED_FILE_ID_B64,
"format": "audio/wav",
},
},
],
}
]
updated = update_messages_with_model_file_ids(messages, "any-model-id", {})
assert updated[0]["content"][1]["file"]["file_id"] == "file-ECBPW7ML9g7XHdwGgUPZaM"
# Customer-supplied format is preserved (this is the field whose absence
# the misleading error message used to complain about).
assert updated[0]["content"][1]["file"]["format"] == "audio/wav"
def test_update_messages_with_model_file_ids_mapping_takes_precedence_over_decode():
"""When both mapping and decode would resolve, the mapping must win
(preserves per-deployment routing precision)."""
mapping = {UNIFIED_FILE_ID_B64: {"model-A": "mapped-provider-file-id"}}
messages = [
{
"role": "user",
"content": [
{
"type": "file",
"file": {
"file_id": UNIFIED_FILE_ID_B64,
"format": "application/pdf",
},
},
],
}
]
updated = update_messages_with_model_file_ids(messages, "model-A", mapping)
assert updated[0]["content"][0]["file"]["file_id"] == "mapped-provider-file-id"
def test_update_messages_with_model_file_ids_non_unified_passes_through():
"""A raw provider id (e.g. gs:// URI or a random string) must be left
untouched when the mapping doesn't resolve it. The decode fallback must
not corrupt non-unified ids."""
raw_id = "gs://my-bucket/uploads/abc-123.wav"
messages = [
{
"role": "user",
"content": [
{"type": "file", "file": {"file_id": raw_id, "format": "audio/wav"}},
],
}
]
updated = update_messages_with_model_file_ids(messages, "model-A", {})
assert updated[0]["content"][0]["file"]["file_id"] == raw_id
def test_update_messages_with_model_file_ids_mapping_miss_falls_back_to_decode():
"""A mapping that exists but doesn't contain this file_id should still
trigger the decode fallback — covers the case where the hook resolved
*some* ids but not this one."""
other_id = "some-other-file-id"
mapping = {other_id: {"model-A": "other-provider-id"}}
messages = [
{
"role": "user",
"content": [
{
"type": "file",
"file": {"file_id": UNIFIED_FILE_ID_B64, "format": "audio/wav"},
},
],
}
]
updated = update_messages_with_model_file_ids(messages, "model-A", mapping)
assert updated[0]["content"][0]["file"]["file_id"] == "file-ECBPW7ML9g7XHdwGgUPZaM"
def test_update_messages_with_model_file_ids_tolerates_non_dict_content_items():
"""Content list items aren't always dicts. text_completion forwards
token-ids (list of ints, or list of list of ints for batch) through
this path. The function must skip non-dict items instead of indexing
into them."""
messages_token_ids = [{"role": "user", "content": [15496, 995]}]
messages_token_ids_batch = [{"role": "user", "content": [[15496, 995], [9906, 0]]}]
# Both should pass through unchanged without raising.
assert (
update_messages_with_model_file_ids(messages_token_ids, "model-A", {})
== messages_token_ids
)
assert (
update_messages_with_model_file_ids(messages_token_ids_batch, "model-A", {})
== messages_token_ids_batch
)

View file

@ -0,0 +1,511 @@
"""
Tests for OAuth 2.0 Token Exchange (RFC 8693) handler for MCP servers.
Covers: exchange flow, caching, error handling, resolve_mcp_auth integration,
bearer token extraction, and config loading.
"""
from unittest.mock import AsyncMock, MagicMock, patch
import httpx
import pytest
from litellm.proxy._experimental.mcp_server.auth.token_exchange import (
TOKEN_EXCHANGE_GRANT_TYPE,
TokenExchangeHandler,
)
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
MCPServerManager,
)
from litellm.proxy._experimental.mcp_server.oauth2_token_cache import (
resolve_mcp_auth,
)
from litellm.proxy._types import LiteLLM_MCPServerTable, MCPTransport
from litellm.types.mcp import MCPAuth
from litellm.types.mcp_server.mcp_server_manager import MCPServer
def _obo_server(**overrides) -> MCPServer:
defaults = dict(
server_id="srv-obo-1",
name="test-obo",
url="https://mcp.example.com/mcp",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2_token_exchange,
client_id="litellm-client-id",
client_secret="litellm-client-secret",
token_exchange_endpoint="https://idp.example.com/oauth2/token",
audience="api://mcp-server",
scopes=["mcp.tools.read", "mcp.tools.execute"],
)
defaults.update(overrides)
return MCPServer(**defaults)
def _exchange_response(token="exchanged-tok-abc", expires_in=3600):
resp = MagicMock()
resp.json.return_value = {
"access_token": token,
"token_type": "Bearer",
"expires_in": expires_in,
}
resp.raise_for_status = MagicMock()
resp.text = ""
return resp
# ── Exchange Flow ──
@pytest.mark.asyncio
async def test_exchange_token_success():
"""Token exchange sends correct RFC 8693 parameters and returns access_token."""
handler = TokenExchangeHandler()
server = _obo_server()
mock_client = AsyncMock()
mock_client.post.return_value = _exchange_response("scoped-token-1")
with patch(
"litellm.proxy._experimental.mcp_server.auth.token_exchange.get_async_httpx_client",
return_value=mock_client,
):
result = await handler.exchange_token("user-jwt-xyz", server)
assert result == "scoped-token-1"
mock_client.post.assert_called_once()
_, kwargs = mock_client.post.call_args
data = kwargs["data"]
assert data["grant_type"] == TOKEN_EXCHANGE_GRANT_TYPE
assert data["subject_token"] == "user-jwt-xyz"
assert data["subject_token_type"] == "urn:ietf:params:oauth:token-type:access_token"
assert data["audience"] == "api://mcp-server"
assert data["scope"] == "mcp.tools.read mcp.tools.execute"
assert data["client_id"] == "litellm-client-id"
assert data["client_secret"] == "litellm-client-secret"
@pytest.mark.asyncio
async def test_exchange_token_no_audience():
"""When audience is None, it is omitted from the request."""
handler = TokenExchangeHandler()
server = _obo_server(audience=None)
mock_client = AsyncMock()
mock_client.post.return_value = _exchange_response()
with patch(
"litellm.proxy._experimental.mcp_server.auth.token_exchange.get_async_httpx_client",
return_value=mock_client,
):
await handler.exchange_token("user-jwt", server)
_, kwargs = mock_client.post.call_args
assert "audience" not in kwargs["data"]
@pytest.mark.asyncio
async def test_exchange_token_no_scopes():
"""When scopes is None, scope param is omitted from the request."""
handler = TokenExchangeHandler()
server = _obo_server(scopes=None)
mock_client = AsyncMock()
mock_client.post.return_value = _exchange_response()
with patch(
"litellm.proxy._experimental.mcp_server.auth.token_exchange.get_async_httpx_client",
return_value=mock_client,
):
await handler.exchange_token("user-jwt", server)
_, kwargs = mock_client.post.call_args
assert "scope" not in kwargs["data"]
# ── Caching ──
@pytest.mark.asyncio
async def test_exchange_token_cached():
"""Second call with same user token uses cache — only 1 HTTP POST."""
handler = TokenExchangeHandler()
server = _obo_server()
mock_client = AsyncMock()
mock_client.post.return_value = _exchange_response("cached-exchange-tok")
with patch(
"litellm.proxy._experimental.mcp_server.auth.token_exchange.get_async_httpx_client",
return_value=mock_client,
):
t1 = await handler.exchange_token("same-jwt", server)
t2 = await handler.exchange_token("same-jwt", server)
assert t1 == t2 == "cached-exchange-tok"
assert mock_client.post.call_count == 1
@pytest.mark.asyncio
async def test_different_user_tokens_not_shared():
"""Different user JWTs get different exchanged tokens."""
handler = TokenExchangeHandler()
server = _obo_server()
call_count = 0
async def mock_post(url, data=None):
nonlocal call_count
call_count += 1
resp = MagicMock()
resp.json.return_value = {
"access_token": f"exchanged-{call_count}",
"expires_in": 3600,
}
resp.raise_for_status = MagicMock()
return resp
mock_client = AsyncMock()
mock_client.post = mock_post
with patch(
"litellm.proxy._experimental.mcp_server.auth.token_exchange.get_async_httpx_client",
return_value=mock_client,
):
t1 = await handler.exchange_token("user-a-jwt", server)
t2 = await handler.exchange_token("user-b-jwt", server)
assert t1 == "exchanged-1"
assert t2 == "exchanged-2"
assert call_count == 2
# ── Error Handling ──
@pytest.mark.asyncio
async def test_exchange_token_http_error():
"""HTTP errors from the IDP are wrapped in a ValueError."""
handler = TokenExchangeHandler()
server = _obo_server()
mock_response = MagicMock()
mock_response.status_code = 400
mock_response.text = "invalid_grant"
mock_response.raise_for_status.side_effect = httpx.HTTPStatusError(
"Bad Request",
request=MagicMock(),
response=mock_response,
)
mock_client = AsyncMock()
mock_client.post.return_value = mock_response
with (
patch(
"litellm.proxy._experimental.mcp_server.auth.token_exchange.get_async_httpx_client",
return_value=mock_client,
),
pytest.raises(ValueError, match="failed with status 400"),
):
await handler.exchange_token("bad-jwt", server)
@pytest.mark.asyncio
async def test_exchange_token_http_error_does_not_log_response_body():
"""Raw IDP error bodies are not logged because they can contain credentials."""
handler = TokenExchangeHandler()
server = _obo_server()
raw_response_body = "client_secret=do-not-log"
mock_response = MagicMock()
mock_response.status_code = 401
mock_response.text = raw_response_body
mock_response.raise_for_status.side_effect = httpx.HTTPStatusError(
"Unauthorized",
request=MagicMock(),
response=mock_response,
)
mock_client = AsyncMock()
mock_client.post.return_value = mock_response
with (
patch(
"litellm.proxy._experimental.mcp_server.auth.token_exchange.get_async_httpx_client",
return_value=mock_client,
),
patch(
"litellm.proxy._experimental.mcp_server.auth.token_exchange.verbose_logger.debug"
) as mock_debug,
pytest.raises(ValueError, match="failed with status 401"),
):
await handler.exchange_token("bad-jwt", server)
logged_values = " ".join(
str(value)
for call in mock_debug.call_args_list
for value in [*call.args, *call.kwargs.values()]
)
assert raw_response_body not in logged_values
@pytest.mark.asyncio
async def test_exchange_token_missing_access_token():
"""Response without access_token raises ValueError."""
handler = TokenExchangeHandler()
server = _obo_server()
resp = MagicMock()
resp.json.return_value = {"token_type": "Bearer"}
resp.raise_for_status = MagicMock()
mock_client = AsyncMock()
mock_client.post.return_value = resp
with (
patch(
"litellm.proxy._experimental.mcp_server.auth.token_exchange.get_async_httpx_client",
return_value=mock_client,
),
pytest.raises(ValueError, match="missing 'access_token'"),
):
await handler.exchange_token("jwt", server)
@pytest.mark.asyncio
async def test_exchange_token_missing_endpoint():
"""Missing token_exchange_endpoint and token_url raises ValueError."""
handler = TokenExchangeHandler()
server = _obo_server(token_exchange_endpoint=None, token_url=None)
with pytest.raises(ValueError, match="no token_exchange_endpoint or token_url"):
await handler.exchange_token("jwt", server)
@pytest.mark.asyncio
async def test_exchange_token_missing_credentials():
"""Missing client_id or client_secret raises ValueError."""
handler = TokenExchangeHandler()
server = _obo_server(client_id=None, client_secret=None)
# has_token_exchange_config will be False, so we call _do_exchange directly
with pytest.raises(ValueError, match="missing client_id or client_secret"):
await handler._do_exchange("jwt", server)
# ── resolve_mcp_auth Integration ──
@pytest.mark.asyncio
async def test_resolve_mcp_auth_with_token_exchange():
"""resolve_mcp_auth delegates to token exchange when server has OBO config and subject_token provided."""
server = _obo_server()
mock_handler = AsyncMock()
mock_handler.exchange_token.return_value = "obo-scoped-token"
with patch(
"litellm.proxy._experimental.mcp_server.auth.token_exchange.mcp_token_exchange_handler",
mock_handler,
):
result = await resolve_mcp_auth(server, subject_token="user-jwt")
assert result == "obo-scoped-token"
mock_handler.exchange_token.assert_called_once_with("user-jwt", server)
@pytest.mark.asyncio
async def test_resolve_mcp_auth_obo_without_subject_token_falls_through():
"""Without a subject_token, resolve_mcp_auth falls through to client_credentials."""
server = _obo_server(
token_url="https://auth.example.com/token",
)
mock_client = AsyncMock()
mock_client.post.return_value = _exchange_response("cc-token")
with patch(
"litellm.proxy._experimental.mcp_server.oauth2_token_cache.get_async_httpx_client",
return_value=mock_client,
):
result = await resolve_mcp_auth(server, subject_token=None)
# Falls through to client_credentials since subject_token is None
# The server has client_id/client_secret/token_url so has_client_credentials is True
assert result == "cc-token"
@pytest.mark.asyncio
async def test_resolve_mcp_auth_obo_without_subject_token_uses_cached_client_credentials():
"""The M2M fallback for OBO servers reuses the client_credentials cache."""
server = _obo_server(
server_id="srv-obo-m2m-cache",
token_url="https://auth.example.com/token",
)
mock_client = AsyncMock()
mock_client.post.return_value = _exchange_response("cached-cc-token")
with patch(
"litellm.proxy._experimental.mcp_server.oauth2_token_cache.get_async_httpx_client",
return_value=mock_client,
):
first = await resolve_mcp_auth(server, subject_token=None)
second = await resolve_mcp_auth(server, subject_token=None)
assert first == second == "cached-cc-token"
mock_client.post.assert_called_once()
@pytest.mark.asyncio
async def test_resolve_mcp_auth_header_beats_obo():
"""An explicit mcp_auth_header takes priority over OBO token exchange."""
server = _obo_server()
result = await resolve_mcp_auth(
server, mcp_auth_header="Bearer override", subject_token="user-jwt"
)
assert result == "Bearer override"
# ── Bearer Token Extraction ──
def test_extract_bearer_token_from_oauth2_headers():
"""Extracts token from oauth2_headers Authorization header."""
result = MCPServerManager._extract_bearer_token(
oauth2_headers={"Authorization": "Bearer my-jwt-token"},
raw_headers=None,
)
assert result == "my-jwt-token"
def test_extract_bearer_token_from_raw_headers():
"""Falls back to raw_headers when oauth2_headers missing."""
result = MCPServerManager._extract_bearer_token(
oauth2_headers=None,
raw_headers={"authorization": "Bearer raw-jwt"},
)
assert result == "raw-jwt"
def test_extract_bearer_token_no_bearer_prefix():
"""Returns token as-is when no Bearer prefix."""
result = MCPServerManager._extract_bearer_token(
oauth2_headers={"Authorization": "some-opaque-token"},
raw_headers=None,
)
assert result == "some-opaque-token"
def test_extract_bearer_token_none():
"""Returns None when no auth headers present."""
result = MCPServerManager._extract_bearer_token(
oauth2_headers=None,
raw_headers=None,
)
assert result is None
# ── MCPServer Properties ──
def test_has_token_exchange_config_true():
"""has_token_exchange_config is True for a fully configured OBO server."""
server = _obo_server()
assert server.has_token_exchange_config is True
def test_has_token_exchange_config_false_wrong_auth_type():
"""has_token_exchange_config is False when auth_type is not oauth2_token_exchange."""
server = _obo_server(auth_type=MCPAuth.oauth2)
assert server.has_token_exchange_config is False
def test_has_token_exchange_config_false_missing_creds():
"""has_token_exchange_config is False when client_id/client_secret missing."""
server = _obo_server(client_id=None)
assert server.has_token_exchange_config is False
def test_has_token_exchange_config_uses_token_url_fallback():
"""has_token_exchange_config is True when token_url is set instead of token_exchange_endpoint."""
server = _obo_server(
token_exchange_endpoint=None,
token_url="https://idp.example.com/token",
)
assert server.has_token_exchange_config is True
# ── Config Loading ──
@pytest.mark.asyncio
async def test_config_loading_token_exchange_fields():
"""load_servers_from_config correctly maps OBO config fields to MCPServer."""
manager = MCPServerManager()
config = {
"my_obo_server": {
"url": "https://mcp.example.com/mcp",
"transport": "http",
"auth_type": "oauth2_token_exchange",
"client_id": "my-client",
"client_secret": "my-secret",
"token_exchange_endpoint": "https://idp.example.com/oauth2/token",
"audience": "api://my-mcp",
"scopes": ["read", "write"],
"subject_token_type": "urn:ietf:params:oauth:token-type:jwt",
}
}
await manager.load_servers_from_config(config)
servers = list(manager.config_mcp_servers.values())
assert len(servers) == 1
server = servers[0]
assert server.auth_type == MCPAuth.oauth2_token_exchange
assert server.token_exchange_endpoint == "https://idp.example.com/oauth2/token"
assert server.audience == "api://my-mcp"
assert server.subject_token_type == "urn:ietf:params:oauth:token-type:jwt"
assert server.client_id == "my-client"
assert server.client_secret == "my-secret"
assert server.scopes == ["read", "write"]
assert server.has_token_exchange_config is True
@pytest.mark.asyncio
async def test_config_loading_default_subject_token_type():
"""subject_token_type defaults to access_token when not specified in config."""
manager = MCPServerManager()
config = {
"obo_defaults": {
"url": "https://mcp.example.com/mcp",
"transport": "http",
"auth_type": "oauth2_token_exchange",
"client_id": "cid",
"client_secret": "csec",
"token_exchange_endpoint": "https://idp.example.com/token",
}
}
await manager.load_servers_from_config(config)
server = list(manager.config_mcp_servers.values())[0]
assert server.subject_token_type == "urn:ietf:params:oauth:token-type:access_token"
@pytest.mark.asyncio
async def test_database_loading_token_exchange_scopes_from_credentials():
"""DB-loaded OBO server credentials retain configured scopes."""
manager = MCPServerManager()
db_server = LiteLLM_MCPServerTable(
server_id="srv-obo-db",
server_name="obo_db_server",
url="https://mcp.example.com/mcp",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2_token_exchange,
credentials={
"client_id": "db-client",
"client_secret": "db-secret",
"token_exchange_endpoint": "https://idp.example.com/oauth2/token",
"audience": "api://db-mcp",
"scopes": ["db.read", "db.write"],
},
)
server = await manager.build_mcp_server_from_table(
db_server,
credentials_are_encrypted=False,
)
assert server.auth_type == MCPAuth.oauth2_token_exchange
assert server.client_id == "db-client"
assert server.client_secret == "db-secret"
assert server.token_exchange_endpoint == "https://idp.example.com/oauth2/token"
assert server.audience == "api://db-mcp"
assert server.scopes == ["db.read", "db.write"]

View file

@ -549,7 +549,11 @@ class TestHookHeaderMergePriority:
captured_extra_headers: Dict[str, Any] = {}
async def fake_create_mcp_client(
server, mcp_auth_header=None, extra_headers=None, stdio_env=None
server,
mcp_auth_header=None,
extra_headers=None,
stdio_env=None,
subject_token=None,
):
captured_extra_headers["value"] = extra_headers
mock_client = MagicMock()
@ -589,7 +593,11 @@ class TestHookHeaderMergePriority:
captured_extra_headers: Dict[str, Any] = {}
async def fake_create_mcp_client(
server, mcp_auth_header=None, extra_headers=None, stdio_env=None
server,
mcp_auth_header=None,
extra_headers=None,
stdio_env=None,
subject_token=None,
):
captured_extra_headers["value"] = extra_headers
mock_client = MagicMock()
@ -635,7 +643,11 @@ class TestHookHeaderMergePriority:
captured_extra_headers: Dict[str, Any] = {}
async def fake_create_mcp_client(
server, mcp_auth_header=None, extra_headers=None, stdio_env=None
server,
mcp_auth_header=None,
extra_headers=None,
stdio_env=None,
subject_token=None,
):
captured_extra_headers["value"] = extra_headers
mock_client = MagicMock()
@ -691,7 +703,11 @@ class TestHookHeaderMergePriority:
captured_extra_headers: Dict[str, Any] = {}
async def fake_create_mcp_client(
server, mcp_auth_header=None, extra_headers=None, stdio_env=None
server,
mcp_auth_header=None,
extra_headers=None,
stdio_env=None,
subject_token=None,
):
captured_extra_headers["value"] = extra_headers
mock_client = MagicMock()
@ -739,7 +755,11 @@ class TestHookHeaderMergePriority:
captured_extra_headers: Dict[str, Any] = {}
async def fake_create_mcp_client(
server, mcp_auth_header=None, extra_headers=None, stdio_env=None
server,
mcp_auth_header=None,
extra_headers=None,
stdio_env=None,
subject_token=None,
):
captured_extra_headers["value"] = extra_headers
mock_client = MagicMock()

View file

@ -1228,6 +1228,7 @@ async def test_oauth2_headers_passed_to_mcp_client():
mcp_auth_header=None,
extra_headers=None,
stdio_env=None,
subject_token=None,
):
# Capture the arguments for verification
captured_client_args.update(
@ -1236,6 +1237,7 @@ async def test_oauth2_headers_passed_to_mcp_client():
"mcp_auth_header": mcp_auth_header,
"extra_headers": extra_headers,
"stdio_env": stdio_env,
"subject_token": subject_token,
}
)
# Return a mock client that doesn't actually connect

View file

@ -450,7 +450,7 @@ class TestMCPServerManager:
captured_extra_headers = None
async def capture_create_mcp_client(
server, mcp_auth_header, extra_headers, stdio_env
server, mcp_auth_header, extra_headers, stdio_env, subject_token=None
): # pragma: no cover - helper
nonlocal captured_extra_headers
captured_extra_headers = extra_headers

View file

@ -53,6 +53,83 @@ def test_non_admin_config_update_route_rejected():
assert "Your role=internal_user" in str(exc_info.value)
@pytest.mark.parametrize(
"route",
["/compliance/eu-ai-act", "/compliance/gdpr"],
)
def test_compliance_routes_open_to_internal_user(route):
"""Compliance routes are stateless validators on caller-supplied log data
- non-admin internal_user roles can call them."""
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=route,
request=request,
valid_token=valid_token,
request_data={},
)
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"],
)
def test_compliance_routes_blocked_for_internal_user_view_only(route):
"""Deprecated internal_user_viewer role must not gain compliance route access."""
role = LitellmUserRoles.INTERNAL_USER_VIEW_ONLY.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 = {}
with pytest.raises(Exception) as exc_info:
RouteChecks.non_proxy_admin_allowed_routes_check(
user_obj=user_obj,
_user_role=role,
route=route,
request=request,
valid_token=valid_token,
request_data={},
)
assert "Only proxy admin can be used" in str(exc_info.value)
def test_proxy_admin_viewer_config_update_route_rejected():
"""Test that proxy admin viewer users are rejected when trying to call /config/update"""

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

@ -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

@ -346,3 +346,126 @@ class TestS3LoggerAuditLogEvent:
element = logger.log_queue[0]
assert element.s3_object_key.startswith("audit_logs/")
assert "audit-456" in element.s3_object_key
class TestS3AuditCallbackParamsDecoupling:
"""`s3_audit_callback_params` should give the audit-log path its own
S3Logger instance, distinct from the singleton serving normal logs."""
@pytest.fixture(autouse=True)
def _isolate_caches_and_globals(self):
from litellm.litellm_core_utils import litellm_logging as ll_logging
from litellm.proxy.management_helpers import audit_logs as ll_audit_logs
original_s3 = litellm.s3_callback_params
original_audit = getattr(litellm, "s3_audit_callback_params", None)
ll_audit_logs._audit_log_callback_cache.clear()
ll_logging._in_memory_loggers.clear()
yield
litellm.s3_callback_params = original_s3
litellm.s3_audit_callback_params = original_audit
ll_audit_logs._audit_log_callback_cache.clear()
ll_logging._in_memory_loggers.clear()
def test_opt_in_constructs_separate_instance_with_audit_config(self):
"""Audit config set → audit resolver returns a fresh S3Logger pointing
at the audit bucket, distinct from the normal-log singleton."""
from litellm.integrations.s3_v2 import S3Logger
from litellm.litellm_core_utils.litellm_logging import (
_init_custom_logger_compatible_class,
)
from litellm.proxy.management_helpers.audit_logs import (
_resolve_audit_log_callback,
)
litellm.s3_callback_params = {"s3_bucket_name": "normal-bucket"}
litellm.s3_audit_callback_params = {"s3_bucket_name": "audit-bucket"}
with patch("asyncio.create_task"):
audit_instance = _resolve_audit_log_callback("s3_v2")
normal_instance = _init_custom_logger_compatible_class(
logging_integration="s3_v2",
internal_usage_cache=None,
llm_router=None,
)
assert isinstance(audit_instance, S3Logger)
assert isinstance(normal_instance, S3Logger)
assert id(audit_instance) != id(normal_instance)
assert audit_instance.s3_bucket_name == "audit-bucket"
assert normal_instance.s3_bucket_name == "normal-bucket"
def test_opt_out_preserves_singleton_behavior(self):
"""No `s3_audit_callback_params` → audit and normal share the singleton
(existing behavior, regression guard)."""
from litellm.integrations.s3_v2 import S3Logger
from litellm.litellm_core_utils.litellm_logging import (
_init_custom_logger_compatible_class,
)
from litellm.proxy.management_helpers.audit_logs import (
_resolve_audit_log_callback,
)
litellm.s3_callback_params = {"s3_bucket_name": "shared-bucket"}
litellm.s3_audit_callback_params = None
with patch("asyncio.create_task"):
normal_instance = _init_custom_logger_compatible_class(
logging_integration="s3_v2",
internal_usage_cache=None,
llm_router=None,
)
audit_instance = _resolve_audit_log_callback("s3_v2")
assert isinstance(audit_instance, S3Logger)
assert id(audit_instance) == id(normal_instance)
assert audit_instance.s3_bucket_name == "shared-bucket"
def test_empty_dict_opts_in(self):
"""`s3_audit_callback_params = {}` is opt-in (truthy-by-presence) and
produces a separate instance with no bucket configured (env/IAM-only)."""
from litellm.integrations.s3_v2 import S3Logger
from litellm.litellm_core_utils.litellm_logging import (
_init_custom_logger_compatible_class,
)
from litellm.proxy.management_helpers.audit_logs import (
_resolve_audit_log_callback,
)
litellm.s3_callback_params = {"s3_bucket_name": "normal-bucket"}
litellm.s3_audit_callback_params = {}
with patch("asyncio.create_task"):
audit_instance = _resolve_audit_log_callback("s3_v2")
normal_instance = _init_custom_logger_compatible_class(
logging_integration="s3_v2",
internal_usage_cache=None,
llm_router=None,
)
assert id(audit_instance) != id(normal_instance)
assert audit_instance.s3_bucket_name is None
assert normal_instance.s3_bucket_name == "normal-bucket"
def test_reset_audit_log_callback_cache_clears_audit_instance(self):
"""`reset_audit_log_callback_cache()` must drop the cached audit
instance so a config reload picks up the new params."""
from litellm.proxy.management_helpers.audit_logs import (
_audit_log_callback_cache,
_resolve_audit_log_callback,
reset_audit_log_callback_cache,
)
litellm.s3_audit_callback_params = {"s3_bucket_name": "first"}
with patch("asyncio.create_task"):
first = _resolve_audit_log_callback("s3_v2")
assert first is not None and "s3_v2" in _audit_log_callback_cache
reset_audit_log_callback_cache()
assert "s3_v2" not in _audit_log_callback_cache
litellm.s3_audit_callback_params = {"s3_bucket_name": "second"}
second = _resolve_audit_log_callback("s3_v2")
assert second is not None
assert id(second) != id(first)
assert second.s3_bucket_name == "second"

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

@ -3,8 +3,9 @@ import datetime
from typing import AsyncGenerator
from unittest.mock import AsyncMock, MagicMock, patch
import httpx
import pytest
from fastapi import HTTPException, Request, status
from fastapi import HTTPException, Request, Response, status
from fastapi.responses import JSONResponse, StreamingResponse
import litellm
@ -26,6 +27,48 @@ from litellm.proxy.utils import ProxyLogging
class TestProxyBaseLLMRequestProcessing:
@pytest.mark.asyncio
async def test_base_passthrough_process_llm_request_preserves_litellm_headers_for_non_streaming_response(
self, monkeypatch
):
processing_obj = ProxyBaseLLMRequestProcessing(data={})
async def fake_base_process_llm_request(**kwargs):
passthrough_response = kwargs["fastapi_response"]
passthrough_response.headers["x-litellm-call-id"] = "test-call-id"
passthrough_response.headers["x-litellm-version"] = "test-version"
return httpx.Response(
status_code=200,
content=b'{"ok":true}',
headers={
"content-type": "application/json",
"x-amzn-requestid": "bedrock-request-id",
},
)
monkeypatch.setattr(
processing_obj,
"base_process_llm_request",
fake_base_process_llm_request,
)
result = await processing_obj.base_passthrough_process_llm_request(
request=MagicMock(spec=Request),
fastapi_response=Response(),
user_api_key_dict=MagicMock(spec=UserAPIKeyAuth),
proxy_logging_obj=MagicMock(spec=ProxyLogging),
general_settings={},
proxy_config=MagicMock(spec=ProxyConfig),
select_data_generator=MagicMock(),
model="bedrock-test-model",
)
assert result.status_code == 200
assert result.body == b'{"ok":true}'
assert result.headers["x-amzn-requestid"] == "bedrock-request-id"
assert result.headers["x-litellm-call-id"] == "test-call-id"
assert result.headers["x-litellm-version"] == "test-version"
@pytest.mark.asyncio
async def test_common_processing_pre_call_logic_pre_call_hook_receives_litellm_call_id(
self, monkeypatch

View file

@ -294,3 +294,134 @@ async def test_post_call_success_hook_skips_guardrail_not_on_model():
)
assert guardrail.was_called is False
# ---------------------------------------------------------------------------
# Integration test: async_post_call_streaming_iterator_hook with model-level guardrails
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_streaming_iterator_hook_runs_model_level_guardrail():
"""
Model-level guardrails configured on a deployment should execute in
async_post_call_streaming_iterator_hook (streaming path) — even when
`default_on: false` and the guardrail is not in the request body.
"""
from litellm.caching.caching import DualCache
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.utils import ProxyLogging
from litellm.types.guardrails import GuardrailEventHooks
class TestStreamingGuardrail(CustomGuardrail):
def __init__(self):
super().__init__(
guardrail_name="test-model-guardrail",
event_hook=GuardrailEventHooks.post_call,
)
self.was_called = False
async def async_post_call_streaming_iterator_hook(
self, user_api_key_dict, response, request_data
):
self.was_called = True
async for chunk in response:
yield chunk
guardrail = TestStreamingGuardrail()
mock_router = MagicMock()
mock_deployment = MagicMock()
mock_deployment.litellm_params.get.return_value = ["test-model-guardrail"]
mock_router.get_deployment.return_value = mock_deployment
async def fake_response():
yield "chunk-1"
yield "chunk-2"
with (
patch("litellm.callbacks", [guardrail]),
patch("litellm.proxy.proxy_server.llm_router", mock_router),
):
proxy_logging = ProxyLogging(user_api_key_cache=DualCache())
request_data = {
"model": "gpt-4",
"metadata": {"model_info": {"id": "model-uuid-123"}},
}
user_api_key_dict = UserAPIKeyAuth(api_key="test-key")
chunks = []
async for chunk in proxy_logging.async_post_call_streaming_iterator_hook(
response=fake_response(),
user_api_key_dict=user_api_key_dict,
request_data=request_data,
):
chunks.append(chunk)
assert guardrail.was_called is True
assert chunks == ["chunk-1", "chunk-2"]
@pytest.mark.asyncio
async def test_streaming_iterator_hook_skips_guardrail_not_on_model():
"""
Streaming guardrails NOT configured on the model (and not in the request
body / key / team) should not execute, even after the dispatcher merge
runs. Confirms the gate stays closed for unrelated guardrails.
"""
from litellm.caching.caching import DualCache
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.utils import ProxyLogging
from litellm.types.guardrails import GuardrailEventHooks
class TestStreamingGuardrail(CustomGuardrail):
def __init__(self):
super().__init__(
guardrail_name="unrelated-guardrail",
event_hook=GuardrailEventHooks.post_call,
)
self.was_called = False
async def async_post_call_streaming_iterator_hook(
self, user_api_key_dict, response, request_data
):
self.was_called = True
async for chunk in response:
yield chunk
guardrail = TestStreamingGuardrail()
# Deployment has a DIFFERENT guardrail configured
mock_router = MagicMock()
mock_deployment = MagicMock()
mock_deployment.litellm_params.get.return_value = ["some-other-guardrail"]
mock_router.get_deployment.return_value = mock_deployment
async def fake_response():
yield "chunk-1"
with (
patch("litellm.callbacks", [guardrail]),
patch("litellm.proxy.proxy_server.llm_router", mock_router),
):
proxy_logging = ProxyLogging(user_api_key_cache=DualCache())
request_data = {
"model": "gpt-4",
"metadata": {"model_info": {"id": "model-uuid-123"}},
}
user_api_key_dict = UserAPIKeyAuth(api_key="test-key")
chunks = []
async for chunk in proxy_logging.async_post_call_streaming_iterator_hook(
response=fake_response(),
user_api_key_dict=user_api_key_dict,
request_data=request_data,
):
chunks.append(chunk)
assert guardrail.was_called is False
assert chunks == ["chunk-1"]

View file

@ -2,7 +2,6 @@ import os
import sys
from unittest.mock import MagicMock, patch
import fastapi
import pytest
sys.path.insert(
@ -12,7 +11,6 @@ sys.path.insert(
import builtins
import types
from litellm.proxy.health_endpoints.health_app_factory import build_health_app
from litellm.proxy.proxy_cli import ProxyInitializationHelpers
@ -771,62 +769,8 @@ class TestProxyInitializationHelpers:
mock_uvicorn_run.assert_called_once()
class TestHealthAppFactory:
"""Test cases for the health app factory module"""
def test_build_health_app(self):
"""Test that build_health_app creates a FastAPI app with the correct title and includes the health router"""
# Execute
health_app = build_health_app()
# Assert
assert health_app.title == "LiteLLM Health Endpoints"
assert isinstance(health_app, fastapi.FastAPI)
# Verify that the app has the expected health endpoints by checking route paths
# When a router is included, its routes are flattened into the main app's routes
route_paths = []
for route in health_app.routes:
if hasattr(route, "path"):
route_paths.append(route.path)
# Check for some expected health endpoints
expected_paths = [
"/test",
"/health/services",
"/health",
"/health/history",
"/health/latest",
"/settings",
"/active/callbacks",
"/health/readiness",
"/health/liveliness",
"/health/liveness",
"/health/test_connection",
]
# At least some of the expected health endpoints should be present
found_paths = [path for path in expected_paths if path in route_paths]
assert (
len(found_paths) > 0
), f"Expected to find health endpoints, but found: {route_paths}"
# Verify that the app has routes (indicating the router was included)
assert (
len(health_app.routes) > 0
), "Health app should have routes from the included router"
def test_build_health_app_returns_different_instances(self):
"""Test that build_health_app returns different FastAPI instances on each call"""
# Execute
health_app_1 = build_health_app()
health_app_2 = build_health_app()
# Assert
assert health_app_1 is not health_app_2
assert health_app_1.title == health_app_2.title
assert isinstance(health_app_1, fastapi.FastAPI)
assert isinstance(health_app_2, fastapi.FastAPI)
class TestRunServerDbSetup:
"""Tests for run_server's prisma setup_database behavior."""
@patch("subprocess.run")
@patch("atexit.register")

View file

@ -6513,3 +6513,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

@ -0,0 +1,62 @@
import json
from pathlib import Path
import pytest
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
@pytest.mark.parametrize("model", ["xai/grok-4.3", "xai/grok-4.3-latest"])
def test_xai_grok_4_3_model_info(model):
json_path = Path(__file__).parents[2] / "model_prices_and_context_window.json"
with open(json_path) as f:
model_cost = json.load(f)
info = model_cost.get(model)
assert (
info is not None
), f"{model} not found in model_prices_and_context_window.json"
assert info["litellm_provider"] == "xai"
assert info["mode"] == "chat"
assert info["input_cost_per_token"] == 1.25e-06
assert info["output_cost_per_token"] == 2.5e-06
assert info["cache_read_input_token_cost"] == 2e-07
assert info["input_cost_per_token_above_200k_tokens"] == 2.5e-06
assert info["output_cost_per_token_above_200k_tokens"] == 5e-06
assert info["cache_read_input_token_cost_above_200k_tokens"] == 4e-07
assert info["max_input_tokens"] == 1000000
assert info["max_output_tokens"] == 1000000
assert info["max_tokens"] == 1000000
assert info["supports_function_calling"] is True
assert info["supports_prompt_caching"] is True
assert info["supports_reasoning"] is True
assert info["supports_response_schema"] is True
assert info["supports_tool_choice"] is True
assert info["supports_vision"] is True
assert info["supports_web_search"] is True
routed_model, provider, _, _ = get_llm_provider(model=model)
assert routed_model == model.split("/", 1)[1]
assert provider == "xai"
def test_xai_grok_4_3_backup_matches_main():
"""Ensure the bundled model cost map stays in sync with the canonical file."""
repo_root = Path(__file__).parents[2]
main_path = repo_root / "model_prices_and_context_window.json"
backup_path = repo_root / "litellm" / "model_prices_and_context_window_backup.json"
with open(main_path) as f:
main_cost = json.load(f)
with open(backup_path) as f:
backup_cost = json.load(f)
for model in ("xai/grok-4.3", "xai/grok-4.3-latest"):
assert backup_cost.get(model) == main_cost.get(
model
), f"{model} differs between main and backup model cost maps"

File diff suppressed because it is too large Load diff

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

@ -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>

28
uv.lock generated
View file

@ -9,7 +9,7 @@ resolution-markers = [
]
[options]
exclude-newer = "2026-05-02T11:18:44.200141Z"
exclude-newer = "0001-01-01T00:00:00Z" # This has no effect and is included for backwards compatibility when using relative exclude-newer values.
exclude-newer-span = "P3D"
[manifest]
@ -3189,7 +3189,7 @@ wheels = [
[[package]]
name = "litellm"
version = "1.84.0"
version = "1.85.0"
source = { editable = "." }
dependencies = [
{ name = "aiohttp" },
@ -3374,7 +3374,7 @@ proxy-dev = [
[package.metadata]
requires-dist = [
{ name = "a2a-sdk", marker = "extra == 'extra-proxy'", specifier = "==0.3.24" },
{ name = "aiohttp", specifier = "==3.13.4" },
{ name = "aiohttp", specifier = ">=3.10,<4.0" },
{ name = "anthropic", extras = ["vertex"], marker = "extra == 'proxy-runtime'", specifier = "==0.84.0" },
{ name = "apscheduler", marker = "extra == 'proxy'", specifier = "==3.11.2" },
{ name = "audioread", marker = "extra == 'stt-nvidia-riva'", specifier = ">=3.0.1" },
@ -3387,14 +3387,14 @@ requires-dist = [
{ name = "azure-storage-file-datalake", marker = "extra == 'proxy-runtime'", specifier = "==12.20.0" },
{ name = "backoff", marker = "extra == 'proxy'", specifier = "==2.2.1" },
{ name = "boto3", marker = "extra == 'proxy'", specifier = "==1.43.1" },
{ name = "click", specifier = "==8.1.8" },
{ name = "click", specifier = ">=8.0.0,<9.0" },
{ name = "cryptography", marker = "extra == 'proxy'", specifier = "==46.0.7" },
{ name = "ddtrace", marker = "extra == 'proxy-runtime'", specifier = "==2.19.0" },
{ name = "detect-secrets", marker = "extra == 'proxy-runtime'", specifier = "==1.5.0" },
{ name = "diskcache", marker = "extra == 'caching'", specifier = "==5.6.3" },
{ name = "fastapi", marker = "extra == 'proxy'", specifier = "==0.124.4" },
{ name = "fastapi-sso", marker = "extra == 'proxy'", specifier = "==0.19.0" },
{ name = "fastuuid", specifier = "==0.14.0" },
{ name = "fastuuid", specifier = ">=0.14.0,<1.0" },
{ name = "google-cloud-aiplatform", marker = "extra == 'google'", specifier = "==1.133.0" },
{ name = "google-cloud-aiplatform", marker = "extra == 'proxy-runtime'", specifier = "==1.133.0" },
{ name = "google-cloud-iam", marker = "extra == 'extra-proxy'", specifier = "==2.19.1" },
@ -3403,10 +3403,10 @@ requires-dist = [
{ name = "grpcio", marker = "extra == 'grpc'", specifier = "==1.78.0" },
{ name = "grpcio", marker = "extra == 'proxy-runtime'", specifier = "==1.78.0" },
{ name = "gunicorn", marker = "extra == 'proxy'", specifier = "==23.0.0" },
{ name = "httpx", specifier = "==0.28.1" },
{ name = "importlib-metadata", specifier = "==8.5.0" },
{ name = "jinja2", specifier = "==3.1.6" },
{ name = "jsonschema", specifier = "==4.23.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 = "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" },
{ name = "litellm-proxy-extras", marker = "extra == 'proxy'", editable = "litellm-proxy-extras" },
@ -3417,7 +3417,7 @@ requires-dist = [
{ name = "numpy", marker = "extra == 'stt-nvidia-riva'", specifier = ">=1.26.0" },
{ name = "numpydoc", marker = "extra == 'utils'", specifier = "==1.8.0" },
{ name = "nvidia-riva-client", marker = "extra == 'stt-nvidia-riva'", specifier = ">=2.15.0" },
{ name = "openai", specifier = "==2.33.0" },
{ name = "openai", specifier = ">=2.20.0,<3.0.0" },
{ name = "opentelemetry-api", marker = "extra == 'proxy-runtime'", specifier = "==1.28.0" },
{ name = "opentelemetry-exporter-otlp", marker = "extra == 'proxy-runtime'", specifier = "==1.28.0" },
{ name = "opentelemetry-sdk", marker = "extra == 'proxy-runtime'", specifier = "==1.28.0" },
@ -3425,12 +3425,12 @@ requires-dist = [
{ name = "polars", marker = "extra == 'proxy'", specifier = "==1.38.1" },
{ name = "prisma", marker = "extra == 'extra-proxy'", specifier = "==0.11.0" },
{ name = "prometheus-client", marker = "extra == 'proxy-runtime'", specifier = "==0.20.0" },
{ name = "pydantic", specifier = "==2.12.5" },
{ name = "pydantic", specifier = ">=2.10.0,<3.0.0" },
{ name = "pyjwt", marker = "extra == 'proxy'", specifier = "==2.12.0" },
{ name = "pynacl", marker = "extra == 'proxy'", specifier = "==1.6.2" },
{ name = "pypdf", marker = "python_full_version < '3.14' and extra == 'proxy-runtime'", specifier = "==6.10.2" },
{ name = "pyroscope-io", marker = "sys_platform != 'win32' and extra == 'proxy'", specifier = "==0.8.16" },
{ name = "python-dotenv", specifier = "==1.2.2" },
{ name = "python-dotenv", specifier = ">=1.0.0,<2.0" },
{ name = "python-multipart", marker = "extra == 'proxy'", specifier = "==0.0.27" },
{ name = "pyyaml", marker = "extra == 'proxy'", specifier = "==6.0.3" },
{ name = "redisvl", marker = "python_full_version < '3.14' and extra == 'extra-proxy'", specifier = "==0.4.1" },
@ -3442,8 +3442,8 @@ requires-dist = [
{ name = "sentry-sdk", marker = "extra == 'proxy-runtime'", specifier = "==2.21.0" },
{ name = "soundfile", marker = "extra == 'proxy'", specifier = "==0.12.1" },
{ name = "soundfile", marker = "extra == 'stt-nvidia-riva'", specifier = ">=0.12.1" },
{ name = "tiktoken", specifier = "==0.12.0" },
{ name = "tokenizers", specifier = "==0.23.1" },
{ name = "tiktoken", specifier = ">=0.8.0,<1.0" },
{ name = "tokenizers", specifier = ">=0.21.0,<1.0" },
{ name = "uvicorn", marker = "extra == 'proxy'", specifier = "==0.33.0" },
{ name = "uvloop", marker = "sys_platform != 'win32' and extra == 'proxy'", specifier = "==0.21.0" },
{ name = "websockets", marker = "extra == 'proxy'", specifier = "==15.0.1" },