mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-01 02:02:20 +00:00
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:
commit
2392f81b93
74 changed files with 5932 additions and 1586 deletions
20
Dockerfile
20
Dockerfile
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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 }}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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 \
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
6
litellm-js/spend-logs/package-lock.json
generated
6
litellm-js/spend-logs/package-lock.json
generated
|
|
@ -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": {
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 [
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
196
litellm/proxy/_experimental/mcp_server/auth/token_exchange.py
Normal file
196
litellm/proxy/_experimental/mcp_server/auth/token_exchange.py
Normal 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()
|
||||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
),
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
213
litellm/proxy/db/routing_prisma_wrapper.py
Normal file
213
litellm/proxy/db/routing_prisma_wrapper.py
Normal 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
|
||||
|
|
@ -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
|
||||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
{
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
],
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
220
tests/llm_translation/test_vcr_filters.py
Normal file
220
tests/llm_translation/test_vcr_filters.py
Normal 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
|
||||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
]
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"""
|
||||
|
||||
|
|
|
|||
887
tests/test_litellm/proxy/db/test_routing_prisma_wrapper.py
Normal file
887
tests/test_litellm/proxy/db/test_routing_prisma_wrapper.py
Normal 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
|
||||
)
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
):
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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."
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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="")
|
||||
|
|
|
|||
62
tests/test_litellm/test_xai_grok_4_3_model_metadata.py
Normal file
62
tests/test_litellm/test_xai_grok_4_3_model_metadata.py
Normal 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"
|
||||
1630
ui/litellm-dashboard/package-lock.json
generated
1630
ui/litellm-dashboard/package-lock.json
generated
File diff suppressed because it is too large
Load diff
|
|
@ -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);
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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, {
|
||||
|
|
|
|||
|
|
@ -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
28
uv.lock
generated
|
|
@ -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" },
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue