merge(team): resolve conflict with litellm_internal_staging

Keep both Sequence (bulk budget snapshots) and Mapping (team general_settings)
imports in team_endpoints.py
This commit is contained in:
mubashir1osmani 2026-07-16 16:47:21 -07:00
commit a3768b489f
356 changed files with 12582 additions and 10825 deletions

View file

@ -5,6 +5,8 @@ on:
branches:
- main
- litellm_internal_staging
- litellm_oss_staging
- "litellm_**"
permissions:
contents: read

View file

@ -4,7 +4,11 @@ on:
push:
branches: [main, litellm_internal_staging]
pull_request:
branches: [main, litellm_internal_staging]
branches:
- main
- litellm_internal_staging
- litellm_oss_staging
- "litellm_**"
concurrency:
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }}

View file

@ -114,6 +114,7 @@ COPY --from=builder /app/litellm/proxy/prisma_migration.py /app/litellm/proxy/pr
# working directory on sys.path; litellm/proxy/hooks resolves
# enterprise.enterprise_hooks from it)
COPY --from=builder /app/enterprise /app/enterprise
COPY --from=builder /app/litellm-proxy-extras /app/litellm-proxy-extras
# Prisma binaries live in $HOME/.cache (default prisma-python location),
# 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

View file

@ -111,6 +111,7 @@ COPY --from=builder /app/litellm/proxy/prisma_migration.py /app/litellm/proxy/pr
# working directory on sys.path; litellm/proxy/hooks resolves
# enterprise.enterprise_hooks from it)
COPY --from=builder /app/enterprise /app/enterprise
COPY --from=builder /app/litellm-proxy-extras /app/litellm-proxy-extras
# 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

View file

@ -137,6 +137,7 @@ COPY --from=builder /app/litellm/proxy/prisma_migration.py /app/litellm/proxy/pr
# working directory on sys.path; litellm/proxy/hooks resolves
# enterprise.enterprise_hooks from it)
COPY --from=builder /app/enterprise /app/enterprise
COPY --from=builder /app/litellm-proxy-extras /app/litellm-proxy-extras
COPY --from=builder /app/.cache /app/.cache
COPY --from=builder /var/lib/litellm/ui /var/lib/litellm/ui
COPY --from=builder /var/lib/litellm/assets /var/lib/litellm/assets

View file

@ -113,6 +113,10 @@ class PagerDutyAlerting(SlackAlerting):
user_api_key_spend=_meta.get("user_api_key_spend"),
user_api_key_max_budget=_meta.get("user_api_key_max_budget"),
user_api_key_budget_reset_at=_meta.get("user_api_key_budget_reset_at"),
user_api_key_user_spend=_meta.get("user_api_key_user_spend"),
user_api_key_user_max_budget=_meta.get("user_api_key_user_max_budget"),
user_api_key_team_spend=_meta.get("user_api_key_team_spend"),
user_api_key_team_max_budget=_meta.get("user_api_key_team_max_budget"),
user_api_key_org_id=_meta.get("user_api_key_org_id"),
user_api_key_org_alias=_meta.get("user_api_key_org_alias"),
user_api_key_team_id=_meta.get("user_api_key_team_id"),
@ -196,6 +200,10 @@ class PagerDutyAlerting(SlackAlerting):
if user_api_key_dict.budget_reset_at
else None
),
user_api_key_user_spend=user_api_key_dict.user_spend,
user_api_key_user_max_budget=user_api_key_dict.user_max_budget,
user_api_key_team_spend=user_api_key_dict.team_spend,
user_api_key_team_max_budget=user_api_key_dict.team_max_budget,
user_api_key_org_id=user_api_key_dict.org_id,
user_api_key_org_alias=user_api_key_dict.organization_alias,
user_api_key_team_id=user_api_key_dict.team_id,

View file

@ -1,6 +1,6 @@
[project]
name = "litellm-enterprise"
version = "0.1.50"
version = "0.1.51"
description = "Package for LiteLLM Enterprise features"
readme = "README.md"
requires-python = ">=3.9"
@ -26,7 +26,7 @@ required-version = ">=0.10.9"
module-root = ""
[tool.commitizen]
version = "0.1.50"
version = "0.1.51"
version_files = [
"pyproject.toml:^version",
"../pyproject.toml:litellm-enterprise==",

View file

@ -46,4 +46,9 @@ Reminders:
- gateway.config.proxy_config (rendered into a ConfigMap and mounted at
/app/config/config.yaml; gateway reads it via
CONFIG_FILE_PATH)
- {component}.pdb.{enabled,minAvailable,maxUnavailable} (per-component PodDisruptionBudget; disabled by
default — with hpa.minReplicas of 1, minAvailable: 1
would block node drains)
- {component}.topologySpreadConstraints (standard k8s list, e.g. spread replicas across
topology.kubernetes.io/zone)
- Enable ingress.enabled=true to dispatch / → ui, gateway data-plane prefixes → gateway, and the catch-all → backend.

View file

@ -295,6 +295,52 @@ harmless no-op for the Job and authoritative for the app pods.
{{- end }}
{{- end -}}
{{/*
PodDisruptionBudget shared by gateway, backend, and ui.
Invoke with a dict:
(dict "root" $ "component" .Values.gateway "componentName" "gateway"
"fullname" (include "litellm.gateway.fullname" .)
"selectorLabels" (include "litellm.gateway.selectorLabels" .))
Renders nothing unless both the component and its `pdb.enabled` are on.
Only one of minAvailable / maxUnavailable should be set; if both are,
minAvailable wins. If neither is set, falls back to `maxUnavailable: 1` so
an enabled-but-unconfigured PDB still permits node drains.
"Set" means non-nil and non-empty-string, so an explicit 0 (e.g.
`maxUnavailable: 0` to forbid all voluntary disruptions) is honored rather
than silently replaced by the fallback.
*/}}
{{- define "litellm.pdb" -}}
{{- $root := .root -}}
{{- $component := .component -}}
{{- $min := $component.pdb.minAvailable -}}
{{- $max := $component.pdb.maxUnavailable -}}
{{- $minSet := not (or (kindIs "invalid" $min) (eq (printf "%v" $min) "")) -}}
{{- $maxSet := not (or (kindIs "invalid" $max) (eq (printf "%v" $max) "")) -}}
{{- if and $component.enabled $component.pdb $component.pdb.enabled }}
apiVersion: policy/v1
kind: PodDisruptionBudget
metadata:
name: {{ .fullname }}
labels:
{{- include "litellm.commonLabels" $root | nindent 4 }}
app.kubernetes.io/component: {{ .componentName }}
spec:
selector:
matchLabels:
{{- .selectorLabels | nindent 6 }}
{{- if $minSet }}
minAvailable: {{ $min }}
{{- else if $maxSet }}
maxUnavailable: {{ $max }}
{{- else }}
maxUnavailable: 1
{{- end }}
{{- end }}
{{- end -}}
{{/*
Renders `envFrom:` block for a component's `envConfigMaps` / `envSecrets`
lists. Each entry is a resource name; the chart wires the whole ConfigMap /

View file

@ -98,4 +98,8 @@ spec:
tolerations:
{{- toYaml . | nindent 8 }}
{{- end }}
{{- with .Values.backend.topologySpreadConstraints }}
topologySpreadConstraints:
{{- toYaml . | nindent 8 }}
{{- end }}
{{- end }}

View file

@ -0,0 +1,6 @@
{{- include "litellm.pdb" (dict
"root" $
"component" .Values.backend
"componentName" "backend"
"fullname" (include "litellm.backend.fullname" .)
"selectorLabels" (include "litellm.backend.selectorLabels" .)) }}

View file

@ -100,4 +100,8 @@ spec:
tolerations:
{{- toYaml . | nindent 8 }}
{{- end }}
{{- with .Values.gateway.topologySpreadConstraints }}
topologySpreadConstraints:
{{- toYaml . | nindent 8 }}
{{- end }}
{{- end }}

View file

@ -0,0 +1,6 @@
{{- include "litellm.pdb" (dict
"root" $
"component" .Values.gateway
"componentName" "gateway"
"fullname" (include "litellm.gateway.fullname" .)
"selectorLabels" (include "litellm.gateway.selectorLabels" .)) }}

View file

@ -76,4 +76,8 @@ spec:
tolerations:
{{- toYaml . | nindent 8 }}
{{- end }}
{{- with .Values.ui.topologySpreadConstraints }}
topologySpreadConstraints:
{{- toYaml . | nindent 8 }}
{{- end }}
{{- end }}

View file

@ -0,0 +1,6 @@
{{- include "litellm.pdb" (dict
"root" $
"component" .Values.ui
"componentName" "ui"
"fullname" (include "litellm.ui.fullname" .)
"selectorLabels" (include "litellm.ui.selectorLabels" .)) }}

View file

@ -0,0 +1,188 @@
suite: test pod disruption budgets and topology spread constraints
templates:
- gateway/poddisruptionbudget.yaml
- backend/poddisruptionbudget.yaml
- ui/poddisruptionbudget.yaml
- gateway/deployment.yaml
- gateway/configmap.yaml
- backend/deployment.yaml
- ui/deployment.yaml
values:
- ./values/required.yaml
tests:
- it: renders no PDB by default
templates:
- gateway/poddisruptionbudget.yaml
- backend/poddisruptionbudget.yaml
- ui/poddisruptionbudget.yaml
asserts:
- hasDocuments:
count: 0
- it: gateway PDB uses minAvailable and matches the gateway selector labels
template: gateway/poddisruptionbudget.yaml
set:
gateway.pdb.enabled: true
gateway.pdb.minAvailable: 1
asserts:
- isKind:
of: PodDisruptionBudget
- equal:
path: apiVersion
value: policy/v1
- equal:
path: metadata.name
value: RELEASE-NAME-litellm-gateway
- equal:
path: spec.minAvailable
value: 1
- notExists:
path: spec.maxUnavailable
- equal:
path: spec.selector.matchLabels
value:
app.kubernetes.io/name: litellm
app.kubernetes.io/instance: RELEASE-NAME
app.kubernetes.io/component: gateway
- it: backend PDB uses maxUnavailable when minAvailable is unset
template: backend/poddisruptionbudget.yaml
set:
backend.pdb.enabled: true
backend.pdb.maxUnavailable: 25%
asserts:
- equal:
path: spec.maxUnavailable
value: 25%
- notExists:
path: spec.minAvailable
- equal:
path: spec.selector.matchLabels
value:
app.kubernetes.io/name: litellm
app.kubernetes.io/instance: RELEASE-NAME
app.kubernetes.io/component: backend
- it: minAvailable wins when both minAvailable and maxUnavailable are set
template: gateway/poddisruptionbudget.yaml
set:
gateway.pdb.enabled: true
gateway.pdb.minAvailable: 2
gateway.pdb.maxUnavailable: 1
asserts:
- equal:
path: spec.minAvailable
value: 2
- notExists:
path: spec.maxUnavailable
- it: an explicit maxUnavailable 0 is honored instead of the fallback
template: backend/poddisruptionbudget.yaml
set:
backend.pdb.enabled: true
backend.pdb.maxUnavailable: 0
asserts:
- equal:
path: spec.maxUnavailable
value: 0
- notExists:
path: spec.minAvailable
- it: an explicit minAvailable 0 is honored and beats a set maxUnavailable
template: gateway/poddisruptionbudget.yaml
set:
gateway.pdb.enabled: true
gateway.pdb.minAvailable: 0
gateway.pdb.maxUnavailable: 1
asserts:
- equal:
path: spec.minAvailable
value: 0
- notExists:
path: spec.maxUnavailable
- it: enabled PDB with neither knob set falls back to maxUnavailable 1
template: ui/poddisruptionbudget.yaml
set:
ui.pdb.enabled: true
asserts:
- equal:
path: spec.maxUnavailable
value: 1
- notExists:
path: spec.minAvailable
- equal:
path: spec.selector.matchLabels
value:
app.kubernetes.io/name: litellm
app.kubernetes.io/instance: RELEASE-NAME
app.kubernetes.io/component: ui
- it: renders no PDB for a disabled component even when its pdb is enabled
template: gateway/poddisruptionbudget.yaml
set:
gateway.enabled: false
gateway.pdb.enabled: true
asserts:
- hasDocuments:
count: 0
- it: deployments omit topologySpreadConstraints by default
templates:
- gateway/deployment.yaml
- backend/deployment.yaml
- ui/deployment.yaml
asserts:
- notExists:
path: spec.template.spec.topologySpreadConstraints
- it: gateway deployment renders configured topologySpreadConstraints
template: gateway/deployment.yaml
set:
gateway.topologySpreadConstraints:
- maxSkew: 1
topologyKey: topology.kubernetes.io/zone
whenUnsatisfiable: ScheduleAnyway
labelSelector:
matchLabels:
app.kubernetes.io/component: gateway
asserts:
- equal:
path: spec.template.spec.topologySpreadConstraints
value:
- maxSkew: 1
topologyKey: topology.kubernetes.io/zone
whenUnsatisfiable: ScheduleAnyway
labelSelector:
matchLabels:
app.kubernetes.io/component: gateway
- it: backend deployment renders configured topologySpreadConstraints
template: backend/deployment.yaml
set:
backend.topologySpreadConstraints:
- maxSkew: 1
topologyKey: kubernetes.io/hostname
whenUnsatisfiable: DoNotSchedule
labelSelector:
matchLabels:
app.kubernetes.io/component: backend
asserts:
- equal:
path: spec.template.spec.topologySpreadConstraints[0].topologyKey
value: kubernetes.io/hostname
- equal:
path: spec.template.spec.topologySpreadConstraints[0].whenUnsatisfiable
value: DoNotSchedule
- it: ui deployment renders configured topologySpreadConstraints
template: ui/deployment.yaml
set:
ui.topologySpreadConstraints:
- maxSkew: 1
topologyKey: topology.kubernetes.io/zone
whenUnsatisfiable: ScheduleAnyway
asserts:
- equal:
path: spec.template.spec.topologySpreadConstraints[0].topologyKey
value: topology.kubernetes.io/zone

View file

@ -190,10 +190,28 @@ gateway:
maxReplicas: 10
targetCPUUtilizationPercentage: 70
targetMemoryUtilizationPercentage: 80
# PodDisruptionBudget for the gateway pods. Set exactly one of
# `minAvailable` / `maxUnavailable` (minAvailable wins if both are set;
# enabling without either falls back to `maxUnavailable: 1`). Disabled by
# default: with the default hpa.minReplicas of 1, a `minAvailable: 1` PDB
# would block node drains entirely.
pdb:
enabled: false
minAvailable: ""
maxUnavailable: ""
podAnnotations: {}
nodeSelector: {}
tolerations: []
affinity: {}
# Standard k8s topologySpreadConstraints for the gateway pods, e.g. to
# spread replicas across zones:
# - maxSkew: 1
# topologyKey: topology.kubernetes.io/zone
# whenUnsatisfiable: ScheduleAnyway
# labelSelector:
# matchLabels:
# app.kubernetes.io/component: gateway
topologySpreadConstraints: []
# ---------- backend (UI / management API) ----------
backend:
@ -233,10 +251,17 @@ backend:
minReplicas: 1
maxReplicas: 4
targetCPUUtilizationPercentage: 70
# Same shape as gateway.pdb.
pdb:
enabled: false
minAvailable: ""
maxUnavailable: ""
podAnnotations: {}
nodeSelector: {}
tolerations: []
affinity: {}
# Same shape as gateway.topologySpreadConstraints.
topologySpreadConstraints: []
# ---------- ui (Next.js static dashboard) ----------
ui:
@ -279,7 +304,14 @@ ui:
minReplicas: 1
maxReplicas: 3
targetCPUUtilizationPercentage: 80
# Same shape as gateway.pdb.
pdb:
enabled: false
minAvailable: ""
maxUnavailable: ""
podAnnotations: {}
nodeSelector: {}
tolerations: []
affinity: {}
# Same shape as gateway.topologySpreadConstraints.
topologySpreadConstraints: []

View file

@ -0,0 +1,2 @@
-- AlterTable
ALTER TABLE "LiteLLM_MCPServerTable" ADD COLUMN IF NOT EXISTS "issuer" TEXT;

View file

@ -325,6 +325,7 @@ model LiteLLM_MCPServerTable {
command String?
args String[] @default([])
env Json? @default("{}")
issuer String?
authorization_url String?
token_url String?
registration_url String?

View file

@ -1,6 +1,6 @@
[project]
name = "litellm-proxy-extras"
version = "0.4.77"
version = "0.4.78"
description = "Additional files for the LiteLLM Proxy. Reduces the size of the main litellm package."
readme = "README.md"
requires-python = ">=3.9"
@ -26,7 +26,7 @@ required-version = ">=0.10.9"
module-root = ""
[tool.commitizen]
version = "0.4.77"
version = "0.4.78"
version_files = [
"pyproject.toml:^version",
"../pyproject.toml:litellm-proxy-extras==",

View file

@ -688,10 +688,8 @@ def get_redis_connection_pool(
elif redis_connect_func and hasattr(redis_connect_func, "_gcp_service_account"):
redis_kwargs["credential_provider"] = GCPIAMCredentialProvider(redis_connect_func._gcp_service_account)
connection_class = async_redis.Connection
if redis_kwargs.pop("ssl", False):
connection_class = async_redis.SSLConnection
redis_kwargs["connection_class"] = connection_class
if redis_kwargs.pop("ssl", None):
redis_kwargs["connection_class"] = async_redis.SSLConnection
return async_redis.BlockingConnectionPool(timeout=REDIS_CONNECTION_POOL_TIMEOUT, **redis_kwargs)

View file

@ -103,6 +103,18 @@ class DualCache(BaseCache):
if default_redis_ttl is not None:
self.default_redis_ttl = default_redis_ttl
def _backfill_kwargs(self, kwargs: "dict[str, object]") -> "dict[str, object]":
"""
Kwargs for writing a Redis read result into the in-memory tier.
Applies ``default_in_memory_ttl`` exactly like the write paths do;
without it, backfilled entries fall to ``InMemoryCache``'s own default
TTL and can outlive the TTL this cache was configured with.
"""
if "ttl" not in kwargs and self.default_in_memory_ttl is not None:
return {**kwargs, "ttl": self.default_in_memory_ttl}
return kwargs
def set_cache(self, key, value, local_only: bool = False, **kwargs):
# Update both Redis and in-memory cache
try:
@ -160,7 +172,7 @@ class DualCache(BaseCache):
if redis_result is not None:
# Update in-memory cache with the value from Redis
self.in_memory_cache.set_cache(key, redis_result, **kwargs)
self.in_memory_cache.set_cache(key, redis_result, **self._backfill_kwargs(kwargs))
result = redis_result
@ -226,7 +238,7 @@ class DualCache(BaseCache):
if redis_result is not None:
# Update in-memory cache with the value from Redis
await self.in_memory_cache.async_set_cache(key, redis_result, **kwargs)
await self.in_memory_cache.async_set_cache(key, redis_result, **self._backfill_kwargs(kwargs))
result = redis_result
@ -318,7 +330,7 @@ class DualCache(BaseCache):
result[key_to_index[key]] = value
if value is not None and self.in_memory_cache is not None:
await self.in_memory_cache.async_set_cache(key, value, **kwargs)
await self.in_memory_cache.async_set_cache(key, value, **self._backfill_kwargs(kwargs))
return result
except Exception:

View file

@ -966,11 +966,15 @@ class BudgetExceededError(Exception):
max_budget: float,
message: Optional[str] = None,
llm_provider: Optional[str] = None,
entity_type: Optional[str] = None,
entity_id: Optional[str] = None,
):
self.current_cost = current_cost
self.max_budget = max_budget
self.status_code = 429
self.llm_provider = llm_provider or ""
self.entity_type = entity_type
self.entity_id = entity_id
# Surface unified rate-limit fields without joining the RateLimitError
# hierarchy so existing `except BudgetExceededError:` handlers keep
# working; custom callbacks reading StandardLoggingPayload pick these

View file

@ -267,29 +267,10 @@ class LangfuseOtelLogger(OpenTelemetry):
# If no keys, return default from env (likely logging to console or something else)
return OpenTelemetryConfig.from_env()
# Determine endpoint - default to US cloud
langfuse_host = LangfuseOtelLogger._get_langfuse_otel_host()
if langfuse_host:
# If LANGFUSE_HOST is provided, construct OTEL endpoint from it
if not langfuse_host.startswith("http"):
langfuse_host = "https://" + langfuse_host
endpoint = f"{langfuse_host.rstrip('/')}/api/public/otel"
verbose_logger.debug(f"Using Langfuse OTEL endpoint from host: {endpoint}")
else:
# Default to US cloud endpoint
endpoint = LANGFUSE_CLOUD_US_ENDPOINT
verbose_logger.debug(f"Using Langfuse US cloud endpoint: {endpoint}")
auth_header = LangfuseOtelLogger._get_langfuse_authorization_header(
public_key=public_key, secret_key=secret_key
)
otlp_auth_headers = f"Authorization={auth_header}"
return OpenTelemetryConfig(
exporter="otlp_http",
endpoint=endpoint,
headers=otlp_auth_headers,
return LangfuseOtelLogger._build_langfuse_otel_config(
public_key=public_key,
secret_key=secret_key,
langfuse_host=LangfuseOtelLogger._get_langfuse_otel_host(),
)
@staticmethod
@ -316,33 +297,36 @@ class LangfuseOtelLogger(OpenTelemetry):
"LANGFUSE_PUBLIC_KEY and LANGFUSE_SECRET_KEY must be set for Langfuse OpenTelemetry integration."
)
# Determine endpoint - default to US cloud
langfuse_host = LangfuseOtelLogger._get_langfuse_otel_host()
return LangfuseOtelLogger._build_langfuse_otel_config(
public_key=public_key,
secret_key=secret_key,
langfuse_host=LangfuseOtelLogger._get_langfuse_otel_host(),
)
@staticmethod
def _build_langfuse_otel_config(
public_key: str, secret_key: str, langfuse_host: Optional[str]
) -> "OpenTelemetryConfig":
"""
Builds an OTLP HTTP config pointing at the Langfuse OTEL endpoint for the
given host (US cloud when no host is provided), authorized with the given keys.
"""
if langfuse_host:
# If LANGFUSE_HOST is provided, construct OTEL endpoint from it
if not langfuse_host.startswith("http"):
langfuse_host = "https://" + langfuse_host
endpoint = f"{langfuse_host.rstrip('/')}/api/public/otel"
normalized_host = langfuse_host if langfuse_host.startswith("http") else f"https://{langfuse_host}"
endpoint = f"{normalized_host.rstrip('/')}/api/public/otel"
verbose_logger.debug(f"Using Langfuse OTEL endpoint from host: {endpoint}")
else:
# Default to US cloud endpoint
endpoint = LANGFUSE_CLOUD_US_ENDPOINT
verbose_logger.debug(f"Using Langfuse US cloud endpoint: {endpoint}")
auth_header = LangfuseOtelLogger._get_langfuse_authorization_header(
public_key=public_key, secret_key=secret_key
)
otlp_auth_headers = f"Authorization={auth_header}"
# Prevent modification of global env vars which causes leakage
# os.environ["OTEL_EXPORTER_OTLP_ENDPOINT"] = endpoint
# os.environ["OTEL_EXPORTER_OTLP_HEADERS"] = otlp_auth_headers
return OpenTelemetryConfig(
exporter="otlp_http",
endpoint=endpoint,
headers=otlp_auth_headers,
headers=f"Authorization={auth_header}",
)
@staticmethod
@ -378,6 +362,29 @@ class LangfuseOtelLogger(OpenTelemetry):
return dynamic_headers
def construct_dynamic_otel_config(
self, standard_callback_dynamic_params: StandardCallbackDynamicParams
) -> Optional["OpenTelemetryConfig"]:
"""
Build a full per-request OTLP config from team/key dynamic Langfuse credentials.
Key-scoped credentials must define the export target, not just the auth
headers: without this, a proxy with no global LANGFUSE_* env vars keeps its
init-time fallback exporter (console), so key-level langfuse_otel silently
never reaches Langfuse.
"""
public_key = standard_callback_dynamic_params.get("langfuse_public_key")
secret_key = standard_callback_dynamic_params.get("langfuse_secret_key")
if not public_key or not secret_key:
return None
langfuse_host = standard_callback_dynamic_params.get("langfuse_host") or self._get_langfuse_otel_host()
return LangfuseOtelLogger._build_langfuse_otel_config(
public_key=public_key,
secret_key=secret_key,
langfuse_host=langfuse_host,
)
def create_litellm_proxy_request_started_span(
self,
start_time: datetime,

View file

@ -28,6 +28,7 @@ from litellm.integrations.opentelemetry_utils.gen_ai_semconv import (
parse_semconv_opt_in,
)
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.litellm_core_utils.secret_redaction import redact_string
from litellm.secret_managers.main import get_secret_bool, str_to_bool
from litellm.types.services import ServiceLoggerPayload
from litellm.types.utils import (
@ -948,12 +949,22 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
Returns:
Tracer: The tracer to use for this request
"""
dynamic_config = self._get_dynamic_otel_config_from_kwargs(kwargs)
if dynamic_config is not None:
verbose_logger.debug(
"[OTEL DEBUG] Using DYNAMIC config tracer with endpoint: %s",
dynamic_config.endpoint,
)
return self._get_tracer_with_dynamic_config(dynamic_config)
dynamic_headers = self._get_dynamic_otel_headers_from_kwargs(kwargs)
if dynamic_headers is not None:
# Create spans using a temporary tracer with dynamic headers
tracer_to_use = self._get_tracer_with_dynamic_headers(dynamic_headers)
verbose_logger.debug("[OTEL DEBUG] Using DYNAMIC tracer with headers: %s", dynamic_headers)
verbose_logger.debug(
"[OTEL DEBUG] Using DYNAMIC tracer with headers: %s", redact_string(str(dynamic_headers))
)
else:
# For langfuse_otel without dynamic headers, create a provider with env var credentials
if hasattr(self, "callback_name") and self.callback_name == "langfuse_otel":
@ -989,6 +1000,32 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
return dynamic_headers if dynamic_headers else None
def _get_dynamic_otel_config_from_kwargs(self, kwargs: dict) -> Optional[OpenTelemetryConfig]:
"""Extract a full dynamic exporter config from kwargs if available."""
standard_callback_dynamic_params: Optional[StandardCallbackDynamicParams] = kwargs.get(
"standard_callback_dynamic_params"
)
if not standard_callback_dynamic_params:
return None
return self.construct_dynamic_otel_config(standard_callback_dynamic_params=standard_callback_dynamic_params)
def _get_tracer_with_dynamic_config(self, dynamic_config: OpenTelemetryConfig):
"""Create (or reuse) a tracer whose exporter target comes from a per-request config."""
from opentelemetry.sdk.trace import TracerProvider
cache_key = f"dynamic_config:{dynamic_config.exporter}:{dynamic_config.endpoint}:{dynamic_config.headers}"
if cache_key in self._tracer_provider_cache:
return self._tracer_provider_cache[cache_key].get_tracer(LITELLM_TRACER_NAME)
temp_provider = TracerProvider(resource=self._get_litellm_resource(self.config))
temp_provider.add_span_processor(self._get_span_processor(config_override=dynamic_config))
self._tracer_provider_cache[cache_key] = temp_provider
return temp_provider.get_tracer(LITELLM_TRACER_NAME)
def _get_tracer_with_dynamic_headers(self, dynamic_headers: dict):
"""Create a temporary tracer with dynamic headers for this request only."""
from opentelemetry.sdk.trace import TracerProvider
@ -1020,6 +1057,19 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
"""
return None
def construct_dynamic_otel_config(
self, standard_callback_dynamic_params: StandardCallbackDynamicParams
) -> Optional[OpenTelemetryConfig]:
"""
Construct a full exporter config from standard callback dynamic params.
Override this when team/key dynamic params must control the export
target (exporter kind + endpoint), not just the request headers. When
this returns a config, it takes precedence over
construct_dynamic_otel_headers for the request.
"""
return None
#########################################################
# End of Team/Key Based Logging Control Flow
#########################################################
@ -2747,7 +2797,11 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
verbose_logger.debug("OpenTelemetry: No parent context found, creating root span")
return None, None
def _get_span_processor(self, dynamic_headers: Optional[dict] = None):
def _get_span_processor(
self,
dynamic_headers: Optional[dict] = None,
config_override: Optional[OpenTelemetryConfig] = None,
):
from opentelemetry.sdk.trace.export import (
BatchSpanProcessor,
ConsoleSpanExporter,
@ -2755,40 +2809,45 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
SpanExporter,
)
otel_exporter = config_override.exporter if config_override else self.OTEL_EXPORTER
otel_endpoint = config_override.endpoint if config_override else self.OTEL_ENDPOINT
otel_headers = config_override.headers if config_override else self.OTEL_HEADERS
verbose_logger.debug(
"OpenTelemetry Logger, initializing span processor \nself.OTEL_EXPORTER: %s\nself.OTEL_ENDPOINT: %s\nself.OTEL_HEADERS: %s",
self.OTEL_EXPORTER,
self.OTEL_ENDPOINT,
self.OTEL_HEADERS,
"OpenTelemetry Logger, initializing span processor \nexporter: %s\nendpoint: %s\nheaders: %s",
otel_exporter,
otel_endpoint,
redact_string(str(otel_headers)),
)
_split_otel_headers = OpenTelemetry._get_headers_dictionary(headers=dynamic_headers or self.OTEL_HEADERS)
_split_otel_headers = OpenTelemetry._get_headers_dictionary(headers=dynamic_headers or otel_headers)
if dynamic_headers:
verbose_logger.debug(
"[OTEL DEBUG] Creating span processor with DYNAMIC headers: %s",
{k: v[:20] + "..." if len(str(v)) > 20 else v for k, v in _split_otel_headers.items()},
redact_string(str(_split_otel_headers)),
)
elif config_override:
verbose_logger.debug(
"[OTEL DEBUG] Creating span processor with DYNAMIC config, endpoint: %s",
otel_endpoint,
)
else:
verbose_logger.debug("[OTEL DEBUG] Creating span processor with GLOBAL headers")
if hasattr(self.OTEL_EXPORTER, "export"): # Check if it has the export method that SpanExporter requires
if hasattr(otel_exporter, "export"): # Check if it has the export method that SpanExporter requires
verbose_logger.debug(
"OpenTelemetry: intiializing SpanExporter. Value of OTEL_EXPORTER: %s",
self.OTEL_EXPORTER,
otel_exporter,
)
return SimpleSpanProcessor(cast(SpanExporter, self.OTEL_EXPORTER))
return SimpleSpanProcessor(cast(SpanExporter, otel_exporter))
if self.OTEL_EXPORTER == "console":
if otel_exporter == "console":
verbose_logger.debug(
"OpenTelemetry: intiializing console exporter. Value of OTEL_EXPORTER: %s",
self.OTEL_EXPORTER,
otel_exporter,
)
return BatchSpanProcessor(ConsoleSpanExporter())
elif (
self.OTEL_EXPORTER == "otlp_http"
or self.OTEL_EXPORTER == "http/protobuf"
or self.OTEL_EXPORTER == "http/json"
):
elif otel_exporter == "otlp_http" or otel_exporter == "http/protobuf" or otel_exporter == "http/json":
try:
from opentelemetry.exporter.otlp.proto.http.trace_exporter import (
OTLPSpanExporter as OTLPSpanExporterHTTP,
@ -2801,13 +2860,13 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
verbose_logger.debug(
"OpenTelemetry: intiializing http exporter. Value of OTEL_EXPORTER: %s",
self.OTEL_EXPORTER,
otel_exporter,
)
normalized_endpoint = self._normalize_otel_endpoint(self.OTEL_ENDPOINT, "traces")
normalized_endpoint = self._normalize_otel_endpoint(otel_endpoint, "traces")
return BatchSpanProcessor(
OTLPSpanExporterHTTP(endpoint=normalized_endpoint, headers=_split_otel_headers),
)
elif self.OTEL_EXPORTER == "otlp_grpc" or self.OTEL_EXPORTER == "grpc":
elif otel_exporter == "otlp_grpc" or otel_exporter == "grpc":
try:
from opentelemetry.exporter.otlp.proto.grpc.trace_exporter import (
OTLPSpanExporter as OTLPSpanExporterGRPC,
@ -2820,16 +2879,16 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
verbose_logger.debug(
"OpenTelemetry: intiializing grpc exporter. Value of OTEL_EXPORTER: %s",
self.OTEL_EXPORTER,
otel_exporter,
)
normalized_endpoint = self._normalize_otel_endpoint(self.OTEL_ENDPOINT, "traces")
normalized_endpoint = self._normalize_otel_endpoint(otel_endpoint, "traces")
return BatchSpanProcessor(
OTLPSpanExporterGRPC(endpoint=normalized_endpoint, headers=_split_otel_headers),
)
else:
verbose_logger.debug(
"OpenTelemetry: intiializing console exporter. Value of OTEL_EXPORTER: %s",
self.OTEL_EXPORTER,
otel_exporter,
)
return BatchSpanProcessor(ConsoleSpanExporter())
@ -2841,7 +2900,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
"OpenTelemetry Logger, initializing log exporter \nself.OTEL_EXPORTER: %s\nself.OTEL_ENDPOINT: %s\nself.OTEL_HEADERS: %s",
self.OTEL_EXPORTER,
self.OTEL_ENDPOINT,
self.OTEL_HEADERS,
redact_string(str(self.OTEL_HEADERS)),
)
_split_otel_headers = OpenTelemetry._get_headers_dictionary(self.OTEL_HEADERS)
@ -2928,7 +2987,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
"OpenTelemetry Logger, initializing metric reader\nself.OTEL_EXPORTER: %s\nself.OTEL_ENDPOINT: %s\nself.OTEL_HEADERS: %s",
self.OTEL_EXPORTER,
self.OTEL_ENDPOINT,
self.OTEL_HEADERS,
redact_string(str(self.OTEL_HEADERS)),
)
_split_otel_headers = OpenTelemetry._get_headers_dictionary(self.OTEL_HEADERS)

View file

@ -38,6 +38,7 @@ from litellm import (
)
from litellm._logging import _is_debugging_on, _redact_string, verbose_logger
from litellm.exceptions import (
BudgetExceededError,
validate_rate_limit_category,
validate_rate_limit_type,
)
@ -4597,6 +4598,10 @@ class StandardLoggingPayloadSetup:
user_api_key_spend=None,
user_api_key_max_budget=None,
user_api_key_budget_reset_at=None,
user_api_key_user_spend=None,
user_api_key_user_max_budget=None,
user_api_key_team_spend=None,
user_api_key_team_max_budget=None,
user_api_key_team_id=None,
user_api_key_org_id=None,
user_api_key_org_alias=None,
@ -4943,6 +4948,7 @@ class StandardLoggingPayloadSetup:
rate_limit_category = validate_rate_limit_category(getattr(original_exception, "category", None))
rate_limit_type = validate_rate_limit_type(getattr(original_exception, "rate_limit_type", None))
budget_error = original_exception if isinstance(original_exception, BudgetExceededError) else None
return StandardLoggingPayloadErrorInformation(
error_code=error_status,
@ -4952,6 +4958,10 @@ class StandardLoggingPayloadSetup:
error_message=error_message,
error_rate_limit_category=rate_limit_category,
error_rate_limit_type=rate_limit_type,
error_budget_entity_type=budget_error.entity_type if budget_error else None,
error_budget_entity_id=budget_error.entity_id if budget_error else None,
error_budget_limit=budget_error.max_budget if budget_error else None,
error_budget_spend=budget_error.current_cost if budget_error else None,
)
@staticmethod
@ -5428,6 +5438,10 @@ def get_standard_logging_metadata(
user_api_key_spend=None,
user_api_key_max_budget=None,
user_api_key_budget_reset_at=None,
user_api_key_user_spend=None,
user_api_key_user_max_budget=None,
user_api_key_team_spend=None,
user_api_key_team_max_budget=None,
user_api_key_team_id=None,
user_api_key_org_id=None,
user_api_key_org_alias=None,
@ -5527,6 +5541,10 @@ def create_dummy_standard_logging_payload() -> StandardLoggingPayload:
user_api_key_team_id=str("test_team"),
user_api_key_user_id=str("test_user"),
user_api_key_team_alias=str("test_team_alias"),
user_api_key_user_spend=None,
user_api_key_user_max_budget=None,
user_api_key_team_spend=None,
user_api_key_team_max_budget=None,
user_api_key_org_id=None,
spend_logs_metadata=None,
requester_ip_address=str("127.0.0.1"),

View file

@ -467,6 +467,7 @@ class ChunkProcessor:
cache_read_input_tokens: Optional[int] = None
completion_tokens_details: Optional[CompletionTokensDetails] = None
prompt_tokens_details: Optional[PromptTokensDetailsWrapper] = None
cost: Optional[float] = None
if "prompt_tokens" in usage_chunk:
prompt_tokens = usage_chunk.get("prompt_tokens", 0) or 0
@ -476,6 +477,8 @@ class ChunkProcessor:
cache_creation_input_tokens = usage_chunk.get("cache_creation_input_tokens")
if "cache_read_input_tokens" in usage_chunk:
cache_read_input_tokens = usage_chunk.get("cache_read_input_tokens")
if "cost" in usage_chunk:
cost = usage_chunk.get("cost")
if hasattr(usage_chunk, "completion_tokens_details"):
if isinstance(usage_chunk.completion_tokens_details, dict):
completion_tokens_details = CompletionTokensDetails(**usage_chunk.completion_tokens_details)
@ -494,6 +497,7 @@ class ChunkProcessor:
"cache_read_input_tokens": cache_read_input_tokens,
"completion_tokens_details": completion_tokens_details,
"prompt_tokens_details": prompt_tokens_details,
"cost": cost,
}
def count_reasoning_tokens(self, response: ModelResponse) -> Optional[int]:
@ -512,6 +516,22 @@ class ChunkProcessor:
return reasoning_tokens
@staticmethod
def _extract_usage_chunk(chunk: dict[str, Any] | ModelResponse | ModelResponseStream) -> Usage | None:
usage_chunk: Usage | dict[str, Any] | None = None
if hasattr(chunk, "usage") and chunk.usage is not None:
usage_chunk = chunk.usage
elif "usage" in chunk:
usage_chunk = chunk["usage"]
elif (isinstance(chunk, ModelResponse) or isinstance(chunk, ModelResponseStream)) and hasattr(
chunk, "_hidden_params"
):
usage_chunk = chunk._hidden_params.get("usage", None)
if isinstance(usage_chunk, dict):
return Usage(**usage_chunk)
return usage_chunk
def _calculate_usage_per_chunk(
self,
chunks: List[Union[Dict[str, Any], ModelResponse]],
@ -548,18 +568,12 @@ class ChunkProcessor:
# is last-wins, so without preserving this separately the 1h breakdown is
# lost and 1h cache writes get billed at the 5m rate.
cache_creation_token_details: Optional[CacheCreationTokenDetails] = None
cost: Optional[float] = None
for chunk in chunks:
usage_chunk: Optional[Usage] = None
if "usage" in chunk:
usage_chunk = chunk["usage"]
elif (isinstance(chunk, ModelResponse) or isinstance(chunk, ModelResponseStream)) and hasattr(
chunk, "_hidden_params"
):
usage_chunk = chunk._hidden_params.get("usage", None)
usage_chunk = self._extract_usage_chunk(chunk)
if usage_chunk is not None:
if isinstance(usage_chunk, dict):
usage_chunk = Usage(**usage_chunk)
usage_chunk_dict = self._usage_chunk_calculation_helper(usage_chunk)
if usage_chunk_dict["prompt_tokens"] is not None and usage_chunk_dict["prompt_tokens"] > 0:
prompt_tokens = usage_chunk_dict["prompt_tokens"]
@ -610,6 +624,9 @@ class ChunkProcessor:
prompt_tokens_details, cache_creation_token_details
)
if usage_chunk_dict["cost"] is not None:
cost = usage_chunk_dict["cost"]
prompt_tokens_details = self._attach_cache_creation_token_details(
prompt_tokens_details, cache_creation_token_details
)
@ -629,6 +646,7 @@ class ChunkProcessor:
web_search_requests=web_search_requests,
completion_tokens_details=completion_tokens_details,
prompt_tokens_details=prompt_tokens_details,
cost=cost,
)
@staticmethod
@ -727,6 +745,7 @@ class ChunkProcessor:
prompt_tokens_details: Optional[PromptTokensDetailsWrapper] = calculated_usage_per_chunk[
"prompt_tokens_details"
]
cost: Optional[float] = calculated_usage_per_chunk["cost"]
try:
returned_usage.prompt_tokens = prompt_tokens or token_counter(model=model, messages=messages)
@ -784,6 +803,9 @@ class ChunkProcessor:
else:
returned_usage.prompt_tokens_details.web_search_requests = web_search_requests
if cost is not None:
setattr(returned_usage, "cost", cost)
# Return a new usage object with the new values
returned_usage = Usage(**returned_usage.model_dump())

View file

@ -962,10 +962,11 @@ class CustomStreamWrapper:
if self.custom_llm_provider == "bedrock" and "trace" in model_response:
return model_response
# Default - return StopIteration
if hasattr(model_response, "usage"):
self.chunks.append(model_response)
raise StopIteration
# Don't raise StopIteration here - some providers (like OpenRouter)
# send usage/cost data in chunks after the finish_reason chunk
if hasattr(model_response, "usage") and model_response.usage is not None:
return model_response
return
# flush any remaining holding chunk
if len(self.holding_chunk) > 0:
if model_response.choices[0].delta.content is None:
@ -1474,12 +1475,16 @@ class CustomStreamWrapper:
self.tool_call = True
if hasattr(chunk, "usage") and chunk.usage is not None:
model_response.usage = chunk.usage
## RETURN ARG
return self.return_processed_chunk_logic(
result = self.return_processed_chunk_logic(
completion_obj=completion_obj,
model_response=model_response, # type: ignore
response_obj=response_obj,
)
return result
except StopIteration:
raise StopIteration
@ -1686,6 +1691,21 @@ class CustomStreamWrapper:
model_response.choices[0].finish_reason = "tool_calls"
return model_response
@staticmethod
def _propagate_usage_cost_to_hidden_params(
response: "ModelResponse",
) -> None:
"""
If the assembled response carries a provider-reported cost on
usage.cost, copy it into _hidden_params so litellm's cost
calculator uses it instead of a token-based estimate.
"""
_usage = getattr(response, "usage", None)
if _usage is not None and hasattr(_usage, "cost") and _usage.cost is not None:
if "additional_headers" not in response._hidden_params:
response._hidden_params["additional_headers"] = {}
response._hidden_params["additional_headers"]["llm_provider-x-litellm-response-cost"] = float(_usage.cost)
def __next__(self) -> "ModelResponseStream":
cache_hit = False
if self.custom_llm_provider is not None and self.custom_llm_provider == "cached_response":
@ -1741,6 +1761,10 @@ class CustomStreamWrapper:
# hasattr(response, "usage") is always True — must check
# `is not None` to avoid running this path on every chunk.
if getattr(response, "usage", None) is not None:
usage_to_preserve = response.usage
if usage_to_preserve:
response._hidden_params["usage"] = usage_to_preserve
obj_dict = response.model_dump()
if "usage" in obj_dict:
@ -1789,6 +1813,8 @@ class CustomStreamWrapper:
response = self.model_response_creator()
if complete_streaming_response is not None:
self._propagate_usage_cost_to_hidden_params(complete_streaming_response)
setattr(
response,
"usage",
@ -1974,97 +2000,7 @@ class CustomStreamWrapper:
self.chunks.append(processed_chunk)
return processed_chunk
except (StopAsyncIteration, StopIteration):
if self.sent_last_chunk is True:
# log the final chunk with accurate streaming values
try:
complete_streaming_response = litellm.stream_chunk_builder(
chunks=self.chunks,
messages=self.messages,
logging_obj=self.logging_obj,
)
except Exception as e:
# see sync __next__: a raise from stream_chunk_builder inside this
# except handler escapes __anext__ and drops the request from SpendLogs.
# Recover best-effort usage from the raw chunks so cost is still tracked
verbose_logger.warning(
"stream_chunk_builder raised at end-of-stream (%s); logging best-effort usage from chunks.",
str(e),
)
try:
complete_streaming_response = self.model_response_creator(
chunk={"usage": calculate_total_usage(chunks=self.chunks)}
)
except Exception:
complete_streaming_response = None
response = self.model_response_creator()
if complete_streaming_response is not None:
setattr(
response,
"usage",
getattr(complete_streaming_response, "usage"),
)
try:
_copy = complete_streaming_response.model_copy(deep=True)
except RuntimeError:
_copy = complete_streaming_response.model_copy()
asyncio.create_task(
self.async_cache_streaming_response(
processed_chunk=_copy,
cache_hit=cache_hit,
)
)
# Update hidden_params with final usage from
# stream_chunk_builder (see sync __next__ for full comment).
if (
self.stream_options is None
and complete_streaming_response is not None
and self._last_returned_hidden_params is not None
):
final_usage = getattr(complete_streaming_response, "usage", None)
if final_usage is not None:
self._last_returned_hidden_params["usage"] = final_usage
if self.sent_stream_usage is False and self.send_stream_usage is True:
self.sent_stream_usage = True
return response
_deferred_cb = getattr(
self.logging_obj,
"_on_deferred_stream_complete",
None,
)
if _deferred_cb is not None:
# Proxy has post-call guardrails. Store the assembled
# response so the outer streaming consumer
# (ProxyLogging.async_post_call_streaming_iterator_hook)
# can fire the deferred callback AFTER all guardrail
# end-of-stream blocks complete. Scheduling here via
# create_task would race with unified_guardrail's
# end-of-stream block for short-stream providers.
self.logging_obj._deferred_stream_complete_args = ( # type: ignore[attr-defined]
complete_streaming_response,
cache_hit,
)
else:
# prefer_async_handlers routes CustomLogger to async_success_handler
# when consumers use ``async for`` on sync-SDK streams. Legacy string
# callbacks still run via executor.submit inside dispatch_success_handlers.
asyncio.create_task(
self.logging_obj.dispatch_success_handlers(
complete_streaming_response,
cache_hit=cache_hit,
start_time=None,
end_time=None,
prefer_async_handlers=True,
)
)
raise StopAsyncIteration # Re-raise StopIteration
else:
self.sent_last_chunk = True
processed_chunk = self.finish_reason_handler()
return processed_chunk
return await self._finalize_completed_stream(cache_hit=cache_hit)
except httpx.TimeoutException as e: # if httpx read timeout error occues
traceback_exception = traceback.format_exc()
## ADD DEBUG INFORMATION - E.G. LITELLM REQUEST TIMEOUT
@ -2079,20 +2015,122 @@ class CustomStreamWrapper:
# Handle any exceptions that might occur during streaming
asyncio.create_task(self.logging_obj.async_failure_handler(e, traceback_exception))
self._handle_stream_fallback_error(e)
except (httpx.ReadError, httpx.RemoteProtocolError) as e:
if self.received_finish_reason is None:
self._log_stream_failure_and_raise(e)
return await self._finalize_completed_stream(cache_hit=cache_hit)
except Exception as e:
traceback_exception = traceback.format_exc()
if self.logging_obj is not None:
self._record_partial_usage_for_failure()
## LOGGING
threading.Thread(
target=self.logging_obj.failure_handler,
args=(e, traceback_exception),
).start() # log response
# Handle any exceptions that might occur during streaming
asyncio.create_task(
self.logging_obj.async_failure_handler(e, traceback_exception) # type: ignore
self._log_stream_failure_and_raise(e)
async def _finalize_completed_stream(self, cache_hit: bool) -> "ModelResponseStream":
if self.sent_last_chunk is True:
# log the final chunk with accurate streaming values
try:
complete_streaming_response = litellm.stream_chunk_builder(
chunks=self.chunks,
messages=self.messages,
logging_obj=self.logging_obj,
)
self._handle_stream_fallback_error(e)
except Exception as e:
# see sync __next__: a raise from stream_chunk_builder inside this
# except handler escapes __anext__ and drops the request from SpendLogs.
# Recover best-effort usage from the raw chunks so cost is still tracked
verbose_logger.warning(
"stream_chunk_builder raised at end-of-stream (%s); logging best-effort usage from chunks.",
str(e),
)
try:
complete_streaming_response = self.model_response_creator(
chunk={"usage": calculate_total_usage(chunks=self.chunks)}
)
except Exception:
complete_streaming_response = None
response = self.model_response_creator()
if complete_streaming_response is not None:
self._propagate_usage_cost_to_hidden_params(complete_streaming_response)
setattr(
response,
"usage",
getattr(complete_streaming_response, "usage"),
)
try:
_copy = complete_streaming_response.model_copy(deep=True)
except RuntimeError:
_copy = complete_streaming_response.model_copy()
asyncio.create_task(
self.async_cache_streaming_response(
processed_chunk=_copy,
cache_hit=cache_hit,
)
)
# Update hidden_params with final usage from
# stream_chunk_builder (see sync __next__ for full comment).
if (
self.stream_options is None
and complete_streaming_response is not None
and self._last_returned_hidden_params is not None
):
final_usage = getattr(complete_streaming_response, "usage", None)
if final_usage is not None:
self._last_returned_hidden_params["usage"] = final_usage
if self.sent_stream_usage is False and self.send_stream_usage is True:
self.sent_stream_usage = True
return response
_deferred_cb = getattr(
self.logging_obj,
"_on_deferred_stream_complete",
None,
)
if _deferred_cb is not None:
# Proxy has post-call guardrails. Store the assembled
# response so the outer streaming consumer
# (ProxyLogging.async_post_call_streaming_iterator_hook)
# can fire the deferred callback AFTER all guardrail
# end-of-stream blocks complete. Scheduling here via
# create_task would race with unified_guardrail's
# end-of-stream block for short-stream providers.
self.logging_obj._deferred_stream_complete_args = ( # type: ignore[attr-defined]
complete_streaming_response,
cache_hit,
)
else:
# prefer_async_handlers routes CustomLogger to async_success_handler
# when consumers use ``async for`` on sync-SDK streams. Legacy string
# callbacks still run via executor.submit inside dispatch_success_handlers.
asyncio.create_task(
self.logging_obj.dispatch_success_handlers(
complete_streaming_response,
cache_hit=cache_hit,
start_time=None,
end_time=None,
prefer_async_handlers=True,
)
)
raise StopAsyncIteration # Re-raise StopIteration
else:
self.sent_last_chunk = True
processed_chunk = self.finish_reason_handler()
return processed_chunk
def _log_stream_failure_and_raise(self, e: Exception) -> NoReturn:
traceback_exception = traceback.format_exc()
if self.logging_obj is not None:
self._record_partial_usage_for_failure()
## LOGGING
threading.Thread(
target=self.logging_obj.failure_handler,
args=(e, traceback_exception),
).start() # log response
# Handle any exceptions that might occur during streaming
asyncio.create_task(
self.logging_obj.async_failure_handler(e, traceback_exception) # type: ignore
)
self._handle_stream_fallback_error(e)
def _record_partial_usage_for_failure(self) -> None:
"""
@ -2228,12 +2266,16 @@ def calculate_total_usage(chunks: List[ModelResponse]) -> Usage:
"""Assume most recent usage chunk has total usage uptil then."""
prompt_tokens: int = 0
completion_tokens: int = 0
latest_usage_chunk = None
for chunk in chunks:
if "usage" in chunk and chunk["usage"] is not None:
if "prompt_tokens" in chunk["usage"]:
prompt_tokens = chunk["usage"].get("prompt_tokens", 0) or 0
if "completion_tokens" in chunk["usage"]:
completion_tokens = chunk["usage"].get("completion_tokens", 0) or 0
usage = chunk["usage"]
latest_usage_chunk = usage
if "prompt_tokens" in usage:
prompt_tokens = usage.get("prompt_tokens", 0) or 0
if "completion_tokens" in usage:
completion_tokens = usage.get("completion_tokens", 0) or 0
returned_usage_chunk = Usage(
prompt_tokens=prompt_tokens,
@ -2241,6 +2283,15 @@ def calculate_total_usage(chunks: List[ModelResponse]) -> Usage:
total_tokens=prompt_tokens + completion_tokens,
)
if latest_usage_chunk is not None:
latest_cost = (
latest_usage_chunk.get("cost")
if isinstance(latest_usage_chunk, dict)
else getattr(latest_usage_chunk, "cost", None)
)
if latest_cost is not None:
returned_usage_chunk.cost = latest_cost
return returned_usage_chunk

View file

@ -1403,11 +1403,6 @@ class LiteLLMAnthropicMessagesAdapter:
assert isinstance(thinking, str)
assert isinstance(signature, str)
if thinking and signature:
raise ValueError(
"Both `thinking` and `signature` in a single streaming chunk isn't supported."
)
return "thinking", ChatCompletionThinkingBlock(
type="thinking", thinking=thinking, signature=signature
)
@ -1463,17 +1458,14 @@ class LiteLLMAnthropicMessagesAdapter:
if choice.delta.reasoning_content is not None:
reasoning_content += choice.delta.reasoning_content
if reasoning_content and reasoning_signature:
raise ValueError("Both `reasoning` and `signature` in a single streaming chunk isn't supported.")
if partial_json is not None:
return "input_json_delta", ContentJsonBlockDelta(type="input_json_delta", partial_json=partial_json)
elif reasoning_content:
return "thinking_delta", ContentThinkingBlockDelta(type="thinking_delta", thinking=reasoning_content)
elif reasoning_signature:
return "signature_delta", ContentThinkingSignatureBlockDelta(
type="signature_delta", signature=reasoning_signature
)
elif reasoning_content:
return "thinking_delta", ContentThinkingBlockDelta(type="thinking_delta", thinking=reasoning_content)
else:
return "text_delta", ContentTextBlockDelta(type="text_delta", text=text)

View file

@ -85,29 +85,12 @@ class AiohttpResponseStream(httpx.AsyncByteStream):
try:
async for chunk in self._aiohttp_response.content.iter_chunked(self.CHUNK_SIZE):
yield chunk
except (
aiohttp.ClientPayloadError,
aiohttp.client_exceptions.ClientPayloadError,
) as e:
# Handle incomplete transfers more gracefully
# Log the error but don't re-raise if we've already yielded some data
verbose_logger.debug(f"Transfer incomplete, but continuing: {e}")
# If the error is due to incomplete transfer encoding, we can still
# return what we've received so far, similar to how httpx handles it
return
except RuntimeError as e:
# Some providers (e.g., SSE streams) may close the connection
# causing aiohttp StreamReader to raise a generic RuntimeError
# with message "Connection closed.". Treat this as a graceful
# end-of-stream so downstream consumers don't error.
if "Connection closed" in str(e):
verbose_logger.debug("Upstream closed streaming connection; ending iterator gracefully")
return
raise
if "Connection closed" not in str(e):
raise
raise httpx.ReadError(str(e)) from e
except aiohttp.http_exceptions.TransferEncodingError as e:
# Handle transfer encoding errors gracefully
verbose_logger.debug(f"Transfer encoding error, but continuing: {e}")
return
raise httpx.ReadError(str(e)) from e
except Exception:
# For other exceptions, use the normal mapping
with map_aiohttp_exceptions():

View file

@ -79,6 +79,7 @@ class LiteLLM_MCPServerTable(LiteLLMPydanticObjectBase):
command: Optional[str] = None
args: List[str] = Field(default_factory=list)
env: Dict[str, str] = Field(default_factory=dict)
issuer: Optional[str] = None
authorization_url: Optional[str] = None
token_url: Optional[str] = None
registration_url: Optional[str] = None

View file

@ -48,6 +48,7 @@ if TYPE_CHECKING:
_AUTH_FLOW_SCOPED_FIELDS: frozenset = frozenset(
{
"issuer",
"authorization_url",
"token_url",
"registration_url",
@ -60,6 +61,13 @@ _AUTH_FLOW_SCOPED_FIELDS: frozenset = frozenset(
}
)
def _blank_to_none(value: Optional[str]) -> Optional[str]:
if not isinstance(value, str):
return None
return value.strip() or None
# Token-exchange settings with dedicated columns that also exist on
# ``MCPCredentials`` as a legacy shape (rows and REST callers that predate the
# columns). Every write lifts blob values into the columns and strips them from
@ -697,13 +705,15 @@ async def update_mcp_server(
# of being reset to a schema default (transport=sse, allow_all_keys=False...).
data_dict = _prepare_mcp_server_data(data, exclude_unset=True, fields_set=fields_set)
# Pre-fetch existing record once if we need it for auth_type or credential logic
# Pre-fetch existing record once if we need it for auth_type, url, or credential logic
existing = None
has_credentials = "credentials" in data_dict and data_dict["credentials"] is not None
# An explicit token-exchange column write (set or clear) also migrates the
# legacy blob copies below, so the existing row is needed for those updates.
explicit_te_write = bool(_TOKEN_EXCHANGE_COLUMN_FIELDS & data_dict.keys())
if data.auth_type or has_credentials or explicit_te_write:
url_provided = "url" in data_dict and data_dict["url"] is not None
issuer_provided = "issuer" in data_dict
if data.auth_type or has_credentials or explicit_te_write or url_provided or issuer_provided:
existing = await MCPServerRepository(prisma_client).table.find_unique(where={"server_id": data.server_id})
auth_type_changed = bool(
@ -711,13 +721,30 @@ async def update_mcp_server(
and existing
and _credential_auth_class(existing.auth_type) != _credential_auth_class(data.auth_type)
)
# A url change re-points the server at a potentially different upstream, so any discovered or
# trust-on-first-use OAuth endpoints/issuer belong to the old upstream and must re-discover.
url_changed = bool(url_provided and existing and existing.url != data_dict["url"])
old_issuer = _blank_to_none(getattr(existing, "issuer", None)) if existing else None
issuer_changed = bool(
issuer_provided and old_issuer is not None and _blank_to_none(data_dict.get("issuer")) != old_issuer
)
# Clear stale credentials when auth_type changes but no new credentials provided
if auth_type_changed and "credentials" not in data_dict:
data_dict["credentials"] = None
if auth_type_changed:
data_dict.update({field: None for field in _AUTH_FLOW_SCOPED_FIELDS if field not in data_dict})
if auth_type_changed or url_changed or issuer_changed:
# Clear each auth-flow-scoped field that the caller either omitted (partial update) or
# resubmitted unchanged. The edit form re-sends every field, so a stale issuer/endpoint
# belonging to the old upstream would otherwise survive a url/auth_type change and win in the
# resolution merge; only a genuinely new submitted value is kept.
data_dict.update(
{
field: None
for field in _AUTH_FLOW_SCOPED_FIELDS
if field not in data_dict or data_dict[field] == getattr(existing, field, None)
}
)
# An explicit column write that does not touch credentials must still migrate
# the row's legacy blob copies: lift values for columns the caller left
@ -1181,6 +1208,7 @@ def mcp_oauth_token_identity(server: object) -> tuple[object, ...]:
getattr(server, "spec_path", None),
getattr(server, "auth_type", None),
getattr(server, "oauth2_flow", None),
getattr(server, "issuer", None),
getattr(server, "authorization_url", None),
getattr(server, "token_url", None),
getattr(server, "registration_url", None),

View file

@ -201,6 +201,38 @@ def _blank_to_none(value: str | None) -> str | None:
return value.strip() or None
def _uses_issuer_anchor(manual_issuer: str | None, is_discovery_auth_type: bool) -> bool:
"""Whether the endpoints are authoritatively anchored to an admin-pinned issuer (RFC 8414 §3.3).
This is the trust/provenance property, distinct from whether the ``issuer`` field is merely
populated: a trust-on-first-use discovered issuer sets ``issuer`` for token identity but is NOT
anchored, so its endpoints stay resource-rooted. Anchoring holds only when the issuer was pinned
(present on the row/config) on a discovery auth type. Every consumer of "is this anchored" reads
this one definition, so the answer cannot diverge across build paths.
"""
return _blank_to_none(manual_issuer) is not None and is_discovery_auth_type
def _endpoints_yield_to_issuer(
issuer: str | None,
is_discovery_auth_type: bool,
authorization_url: str | None,
token_url: str | None,
registration_url: str | None,
) -> tuple[str | None, str | None, str | None]:
"""The single rule that makes an admin-configured ``issuer`` the sole authoritative endpoint
source (RFC 8414 §3.3): when it is set for a discovery auth type, the stored/manual
``authorization_url``/``token_url``/``registration_url`` do not apply. They neither anchor nor
short-circuit discovery, never override the issuer document in the merge, and never substitute for
it when the issuer fetch fails (fail-closed). Returns the endpoint values that remain in force,
i.e. all ``None`` when issuer-anchored, else the inputs unchanged. Called at every resolution site
so the invariant holds in one place instead of being re-derived per merge.
"""
if issuer is not None and is_discovery_auth_type:
return None, None, None
return authorization_url, token_url, registration_url
def _normalized_authorize_endpoint(url: str) -> str:
"""Compare authorize endpoints on scheme, host, and path only. The default port is elided and
the host is lowercased so ``https://IDP.example.com:443/authorize/`` and
@ -217,6 +249,17 @@ def _normalized_authorize_endpoint(url: str) -> str:
return f"{scheme}://{authority}{parsed.path.rstrip('/')}"
def _issuer_matches(claimed_issuer: object, configured_issuer: str) -> bool:
"""RFC 8414 §3.3 issuer equality between the metadata document's self-attested ``issuer`` and the
admin-configured issuer, tolerant only of URL-insignificant differences (scheme/host case, the
default port, a trailing slash). A non-string or empty claimed issuer never matches, so a
document that omits ``issuer`` fails closed under issuer-anchored discovery.
"""
if not isinstance(claimed_issuer, str) or not claimed_issuer:
return False
return _normalized_authorize_endpoint(claimed_issuer) == _normalized_authorize_endpoint(configured_issuer)
def _endpoints_corroborate_authorization_url(
source_authorization_url: str | None,
trusted_authorization_url: str | None,
@ -260,11 +303,27 @@ def _carry_forward_resolved_oauth_endpoints(new_server: MCPServer, previous_serv
incoming build has no pinned authorize endpoint (``None`` -> we adopt the previous one too, a
consistent group) or pins the same one. An admin re-pointing ``authorization_url`` to a different
server must not keep serving the old server's token endpoint or granted scopes.
When the server is issuer-anchored (``issuer_is_anchored`` -- a pinned issuer on a discovery auth
type), the endpoints come solely from the §3.3-validated issuer document, so carry-forward is
skipped entirely for its endpoints: a failed issuer fetch leaves them ``None`` and must stay
``None`` (fail-closed), never resurrected from the previous registry entry. A merely discovered
(trust-on-first-use) issuer is NOT anchored -- ``issuer`` is set for token identity but the
endpoints are resource-rooted, so they still carry forward as last-known-good, gated by the
corroboration check below like any other resource-rooted server. Scopes stay resource-driven and
can carry either way.
"""
if previous_server is None:
return
if previous_server.url != new_server.url or previous_server.auth_type != new_server.auth_type:
return
if new_server.issuer_is_anchored:
# Endpoints come solely from the §3.3-validated issuer document; a failed fetch stays
# fail-closed and must not be resurrected from the previous entry. Only the resource-driven
# scopes carry as last-known-good.
if not new_server.scopes and previous_server.scopes:
new_server.scopes = previous_server.scopes
return
may_carry = _endpoints_corroborate_authorization_url(
previous_server.authorization_url, new_server.authorization_url
)
@ -1137,34 +1196,48 @@ class MCPServerManager:
)
auth_type = server_config.get("auth_type", None)
manual_issuer = _blank_to_none(server_config.get("issuer"))
manual_authorization_url = _blank_to_none(server_config.get("authorization_url"))
manual_token_url = _blank_to_none(server_config.get("token_url"))
manual_registration_url = _blank_to_none(server_config.get("registration_url"))
if server_url and (
auth_type in _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES
is_discovery_auth_type = auth_type in _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES
use_issuer_anchor = _uses_issuer_anchor(manual_issuer, is_discovery_auth_type)
manual_authorization_url, manual_token_url, manual_registration_url = _endpoints_yield_to_issuer(
manual_issuer,
is_discovery_auth_type,
manual_authorization_url,
manual_token_url,
manual_registration_url,
)
should_discover = bool(server_url) and (
is_discovery_auth_type
or self._obo_needs_endpoint_discovery(
auth_type,
server_config.get("token_exchange_endpoint"),
manual_token_url,
)
):
)
if not should_discover:
mcp_oauth_metadata = None
elif manual_issuer is not None and is_discovery_auth_type:
mcp_oauth_metadata = await self._fetch_issuer_anchored_oauth_metadata(manual_issuer, server_url)
else:
mcp_oauth_metadata = await self._descovery_metadata(
server_url=server_url,
allow_origin_fallback=auth_type in _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES,
allow_origin_fallback=is_discovery_auth_type,
)
else:
mcp_oauth_metadata = None
gated_oauth_metadata = (
_restrict_discovery_to_corroborated_authorization_server(
if use_issuer_anchor:
gated_oauth_metadata = mcp_oauth_metadata
elif is_discovery_auth_type:
gated_oauth_metadata = _restrict_discovery_to_corroborated_authorization_server(
mcp_oauth_metadata,
manual_authorization_url,
server_name or server_id,
bool(server_config.get("dcr_bridge")),
)
if auth_type in _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES
else mcp_oauth_metadata
)
else:
gated_oauth_metadata = mcp_oauth_metadata
# Filter blank scopes (e.g. YAML ``scopes: [""]``) the same way the DB-build path does, so
# an all-blank list normalizes to None rather than a ``("",)`` tuple that skips the
@ -1179,6 +1252,12 @@ class MCPServerManager:
resolved_registration_url = manual_registration_url or (
gated_oauth_metadata.registration_url if gated_oauth_metadata else None
)
discovered_issuer = (
gated_oauth_metadata.discovered_issuer
if gated_oauth_metadata and not gated_oauth_metadata.from_origin_fallback
else None
)
effective_issuer = manual_issuer or discovered_issuer
config_oauth2_flow = server_config.get("oauth2_flow", None)
if auth_type == MCPAuth.oauth2 and config_oauth2_flow not in (
@ -1227,6 +1306,8 @@ class MCPServerManager:
client_secret=server_config.get("client_secret", None),
oauth2_flow=self._explicit_oauth2_flow(config_oauth2_flow),
scopes=resolved_scopes,
issuer=effective_issuer,
issuer_is_anchored=use_issuer_anchor,
authorization_url=resolved_authorization_url,
token_url=resolved_token_url,
registration_url=resolved_registration_url,
@ -1487,6 +1568,52 @@ class MCPServerManager:
decrypt_global_env_var_values(env_vars_list)
return env_vars_list
async def _resolve_table_oauth_metadata(
self,
*,
mcp_server: LiteLLM_MCPServerTable,
auth_type: MCPAuthType,
server_url: Optional[str],
manual_issuer: Optional[str],
manual_authorization_url: Optional[str],
manual_token_url: Optional[str],
is_discovery_auth_type: bool,
use_issuer_anchor: bool,
scopes: Optional[list[str]],
token_exchange_endpoint: Optional[str],
) -> Optional[MCPOAuthMetadata]:
has_all_upstream_oauth_fields = bool(manual_authorization_url and manual_token_url and scopes)
needs_discovery = bool(server_url) and (
(is_discovery_auth_type and not has_all_upstream_oauth_fields)
or self._obo_needs_endpoint_discovery(auth_type, token_exchange_endpoint, manual_token_url)
)
if not needs_discovery:
mcp_oauth_metadata: Optional[MCPOAuthMetadata] = None
elif use_issuer_anchor and manual_issuer is not None:
mcp_oauth_metadata = await self._fetch_issuer_anchored_oauth_metadata(manual_issuer, server_url)
else:
mcp_oauth_metadata = await self._descovery_metadata(
server_url=server_url, # type: ignore[arg-type]
allow_origin_fallback=is_discovery_auth_type,
)
if needs_discovery and not use_issuer_anchor and mcp_oauth_metadata is None:
verbose_logger.warning(
"MCP OAuth discovery yielded no metadata for server %s (%s); "
"OAuth endpoints/scopes stay unresolved until a rebuild succeeds",
mcp_server.server_id,
server_url,
)
if use_issuer_anchor:
return mcp_oauth_metadata
if is_discovery_auth_type:
return _restrict_discovery_to_corroborated_authorization_server(
mcp_oauth_metadata,
manual_authorization_url,
mcp_server.server_id,
bool(getattr(mcp_server, "dcr_bridge", None)),
)
return mcp_oauth_metadata
async def build_mcp_server_from_table(
self,
mcp_server: LiteLLM_MCPServerTable,
@ -1570,46 +1697,38 @@ class MCPServerManager:
auth_type = cast(MCPAuthType, mcp_server.auth_type)
server_url = mcp_server.url
manual_issuer = _blank_to_none(mcp_server.issuer)
manual_authorization_url = _blank_to_none(mcp_server.authorization_url)
manual_token_url = _blank_to_none(mcp_server.token_url)
manual_registration_url = _blank_to_none(mcp_server.registration_url)
has_all_upstream_oauth_fields = bool(manual_authorization_url and manual_token_url and scopes)
needs_discovery = bool(server_url) and (
(auth_type in _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES and not has_all_upstream_oauth_fields)
or self._obo_needs_endpoint_discovery(
auth_type,
mcp_server.token_exchange_endpoint
or (credentials_dict.get("token_exchange_endpoint") if credentials_dict else None),
manual_token_url,
)
is_discovery_auth_type = auth_type in _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES
use_issuer_anchor = _uses_issuer_anchor(manual_issuer, is_discovery_auth_type)
manual_authorization_url, manual_token_url, manual_registration_url = _endpoints_yield_to_issuer(
manual_issuer, is_discovery_auth_type, manual_authorization_url, manual_token_url, manual_registration_url
)
mcp_oauth_metadata = (
await self._descovery_metadata(
server_url=server_url, # type: ignore[arg-type]
allow_origin_fallback=auth_type in _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES,
)
if needs_discovery
else None
token_exchange_endpoint = mcp_server.token_exchange_endpoint or (
credentials_dict.get("token_exchange_endpoint") if credentials_dict else None
)
if needs_discovery and mcp_oauth_metadata is None:
verbose_logger.warning(
"MCP OAuth discovery yielded no metadata for server %s (%s); "
"OAuth endpoints/scopes stay unresolved until a rebuild succeeds",
mcp_server.server_id,
server_url,
)
gated_oauth_metadata = (
_restrict_discovery_to_corroborated_authorization_server(
mcp_oauth_metadata,
manual_authorization_url,
mcp_server.server_id,
bool(getattr(mcp_server, "dcr_bridge", None)),
)
if auth_type in _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES
else mcp_oauth_metadata
gated_oauth_metadata = await self._resolve_table_oauth_metadata(
mcp_server=mcp_server,
auth_type=auth_type,
server_url=server_url,
manual_issuer=manual_issuer,
manual_authorization_url=manual_authorization_url,
manual_token_url=manual_token_url,
is_discovery_auth_type=is_discovery_auth_type,
use_issuer_anchor=use_issuer_anchor,
scopes=scopes,
token_exchange_endpoint=token_exchange_endpoint,
)
resolved_scopes = scopes or (gated_oauth_metadata.scopes if gated_oauth_metadata else None)
discovered_issuer = (
gated_oauth_metadata.discovered_issuer
if gated_oauth_metadata and not gated_oauth_metadata.from_origin_fallback
else None
)
effective_issuer = manual_issuer or discovered_issuer
new_server = MCPServer(
server_id=mcp_server.server_id,
@ -1629,6 +1748,8 @@ class MCPServerManager:
client_secret=client_secret_value or getattr(mcp_server, "client_secret", None),
oauth2_flow=self._explicit_oauth2_flow(getattr(mcp_server, "oauth2_flow", None)),
scopes=resolved_scopes,
issuer=effective_issuer,
issuer_is_anchored=use_issuer_anchor,
authorization_url=manual_authorization_url or getattr(gated_oauth_metadata, "authorization_url", None),
token_url=manual_token_url or getattr(gated_oauth_metadata, "token_url", None),
registration_url=manual_registration_url or getattr(gated_oauth_metadata, "registration_url", None),
@ -1688,10 +1809,12 @@ class MCPServerManager:
await self._persist_discovered_oauth_endpoints(
server_id=mcp_server.server_id,
auth_type=auth_type,
existing_issuer=manual_issuer,
existing_authorization_url=manual_authorization_url,
existing_token_url=manual_token_url,
existing_scopes=scopes,
metadata=gated_oauth_metadata,
is_issuer_anchored=use_issuer_anchor,
)
return new_server
@ -1735,10 +1858,12 @@ class MCPServerManager:
*,
server_id: str,
auth_type: MCPAuthType | None,
existing_issuer: str | None,
existing_authorization_url: str | None,
existing_token_url: str | None,
existing_scopes: list[str] | None,
metadata: MCPOAuthMetadata | None,
is_issuer_anchored: bool = False,
) -> None:
"""Write freshly discovered OAuth endpoints back onto the DB row.
@ -1752,19 +1877,37 @@ class MCPServerManager:
because ``_dcr_bridge_relays_client_registration`` keys off that column. Best-effort: a
failed write re-discovers on the next build. Scopes go through ``update_mcp_server`` so
they merge into the credentials blob without touching the stored client credentials.
For an issuer-anchored server (``is_issuer_anchored``) the endpoints are re-derived from the
§3.3-validated issuer document on every build, so they are NOT persisted into the endpoint
columns: persisting them would make the next build see populated endpoints and treat them as
authoritative stored values, defeating the "endpoints come solely from the issuer" invariant.
Only the resource-driven scopes are persisted for such servers.
"""
if auth_type not in _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES:
return
if metadata is None or metadata.from_origin_fallback:
return
issuer_update = (
{"issuer": metadata.discovered_issuer} if metadata.discovered_issuer and not existing_issuer else {}
)
authorization_url_update = (
{"authorization_url": metadata.authorization_url}
if metadata.authorization_url and not existing_authorization_url
if metadata.authorization_url and not existing_authorization_url and not is_issuer_anchored
else {}
)
token_url_update = (
{"token_url": metadata.token_url}
if metadata.token_url and not existing_token_url and not is_issuer_anchored
else {}
)
token_url_update = {"token_url": metadata.token_url} if metadata.token_url and not existing_token_url else {}
scopes_update = {"credentials": {"scopes": metadata.scopes}} if metadata.scopes and not existing_scopes else {}
updates: dict[str, object] = {**authorization_url_update, **token_url_update, **scopes_update}
updates: dict[str, object] = {
**issuer_update,
**authorization_url_update,
**token_url_update,
**scopes_update,
}
if not updates:
return
from litellm.proxy._experimental.mcp_server.db import ( # noqa: PLC0415 # db.py imports this module at load
@ -3337,8 +3480,41 @@ class MCPServerManager:
return metadata
return None
async def _fetch_issuer_anchored_oauth_metadata(
self, issuer: str, server_url: Optional[str]
) -> Optional[MCPOAuthMetadata]:
"""RFC 8414 issuer-anchored discovery for the OAuth endpoints, with resource-driven scopes.
Fetch authorization-server metadata from the admin-configured issuer's own origin and adopt
its ``token_endpoint``/``registration_endpoint`` only when the document self-attests that same
issuer (RFC 8414 §3.3). Because the trust anchor is the pinned issuer rather than anything the
MCP resource advertises, the endpoints are authoritative for that issuer and cannot be
substituted by a compromised resource. Fails closed (returns None) on a §3.3 mismatch or a
fetch failure. The issuer is passed as its own ``server_url`` so the endpoint fetch is treated
as same-authority and is not subject to the resource-scoped SSRF shortcut.
Scopes are NOT taken from the issuer document. Per the MCP authorization spec Scope Selection
Strategy and RFC 9728, the scopes a client requests are resource-driven (the WWW-Authenticate
challenge or the protected-resource ``scopes_supported``), so the resource's advertised scopes
are fetched separately and used; the resource can influence only the requested scope, which
the authorization server and user consent bound (RFC 6749 §3.3), never the token endpoint.
"""
metadata = await self._fetch_single_authorization_server_metadata(issuer, issuer, require_issuer=issuer)
if metadata is None:
verbose_logger.warning(
"MCP OAuth issuer-anchored discovery for issuer %s yielded no metadata whose issuer "
"matched (RFC 8414 §3.3); OAuth endpoints stay unresolved until a rebuild succeeds",
issuer,
)
return None
resource_metadata = (
await self._descovery_metadata(server_url, allow_origin_fallback=False) if server_url else None
)
resource_scopes = resource_metadata.scopes if resource_metadata else None
return metadata.model_copy(update={"scopes": resource_scopes})
async def _fetch_single_authorization_server_metadata(
self, issuer_url: str, server_url: str
self, issuer_url: str, server_url: str, require_issuer: Optional[str] = None
) -> Optional[MCPOAuthMetadata]:
try:
parsed = urlparse(issuer_url)
@ -3382,20 +3558,33 @@ class MCPServerManager:
)
continue
scopes = self._extract_scopes(data.get("scopes_supported"))
claimed_issuer = data.get("issuer")
verbose_logger.debug(
"Authorization server metadata from %s: issuer=%s grant_types_supported=%s "
"token_endpoint_auth_methods_supported=%s",
url,
data.get("issuer"),
claimed_issuer,
data.get("grant_types_supported"),
data.get("token_endpoint_auth_methods_supported"),
)
if require_issuer is not None and not _issuer_matches(claimed_issuer, require_issuer):
verbose_logger.warning(
"MCP OAuth issuer-anchored discovery: metadata at %s self-attests issuer %r, which "
"does not match the configured issuer %r (RFC 8414 §3.3); rejecting so a compromised "
"resource cannot substitute an attacker authorization server",
url,
claimed_issuer,
require_issuer,
)
continue
scopes = self._extract_scopes(data.get("scopes_supported"))
metadata = MCPOAuthMetadata(
scopes=scopes,
authorization_url=data.get("authorization_endpoint"),
token_url=data.get("token_endpoint"),
registration_url=data.get("registration_endpoint"),
discovered_issuer=claimed_issuer if isinstance(claimed_issuer, str) and claimed_issuer else None,
)
if any(
@ -5116,6 +5305,7 @@ class MCPServerManager:
command=getattr(server, "command", None),
args=getattr(server, "args", None) or [],
env=getattr(server, "env", None) or {},
issuer=server.issuer,
authorization_url=server.authorization_url,
token_url=server.token_url,
registration_url=server.registration_url,
@ -5225,6 +5415,7 @@ class MCPServerManager:
command=getattr(server, "command", None),
args=getattr(server, "args", None) or [],
env=getattr(server, "env", None) or {},
issuer=server.issuer,
authorization_url=server.authorization_url,
token_url=server.token_url,
registration_url=server.registration_url,

View file

@ -1138,6 +1138,7 @@ if MCP_AVAILABLE:
static_headers=request.static_headers,
client_id=client_id,
client_secret=client_secret,
issuer=request.issuer,
token_url=request.token_url,
scopes=scopes,
authorization_url=request.authorization_url,

View file

@ -4,6 +4,7 @@ Semantic MCP Tool Filtering using semantic-router
Filters MCP tools semantically for /chat/completions and /responses endpoints.
"""
import asyncio
from typing import TYPE_CHECKING, Any, Dict, List, Optional
from litellm._logging import verbose_logger
@ -76,6 +77,7 @@ class SemanticMCPToolFilter:
self.tool_router: Optional["SemanticRouter"] = None
self.context_window_error: Optional[str] = None
self._tool_map: Dict[str, Any] = {} # MCPTool objects or OpenAI function dicts
self._index_sync_lock = asyncio.Lock()
async def build_router_from_mcp_registry(self) -> None:
"""Build semantic router from all MCP tools in the registry (no auth checks)."""
@ -182,6 +184,81 @@ class SemanticMCPToolFilter:
return
raise
def _has_tools_missing_from_index(self, tools: list[Any]) -> bool:
"""Allocation-free check for any named tool not yet in the semantic index."""
return any(name and name not in self._tool_map for name in (self._extract_tool_info(t)[0] for t in tools))
def _tools_missing_from_index(self, tools: list[Any]) -> dict[str, Any]:
"""Map name -> tool for every named tool not yet in the semantic index."""
return {
name: tool
for name, tool in ((self._extract_tool_info(t)[0], t) for t in tools)
if name and name not in self._tool_map
}
async def _ensure_tools_indexed(self, available_tools: list[Any]) -> None:
"""
Index request-time tools the startup build never saw.
The startup index lists every registered MCP server WITHOUT per-user
credentials, so servers requiring per-user auth (interactive OAuth
tokens, user-scoped env vars) contribute zero routes. Tools reaching
the filter came through an authenticated expansion; without indexing
them here they can never be selected, so requests either bypass
filtering entirely (N->N) or lose every tool to unrelated matches.
Runs async-only (no synchronous embedding on the request path) and
never writes shared error state: an embedding failure here raises and
is scoped to the requesting call, so one request's oversized tool
description cannot poison the filter for other users on the worker.
"""
from semantic_router.routers import SemanticRouter
from semantic_router.routers.base import Route
from litellm.router_strategy.auto_router.litellm_encoder import (
LiteLLMRouterEncoder,
)
if not self._has_tools_missing_from_index(available_tools):
return
async with self._index_sync_lock:
missing = self._tools_missing_from_index(available_tools)
if not missing:
return
descriptions = {name: self._extract_tool_info(tool)[1] for name, tool in missing.items()}
routes = [
Route(
name=name,
description=description,
utterances=[description],
score_threshold=self.similarity_threshold,
)
for name, description in descriptions.items()
]
if self.tool_router is None:
router = SemanticRouter(
routes=[],
encoder=LiteLLMRouterEncoder(
litellm_router_instance=self.router_instance,
model_name=self.embedding_model,
score_threshold=self.similarity_threshold,
),
auto_sync="local",
top_k=self.top_k,
)
await router.aadd(routes)
self.tool_router = router
else:
await self.tool_router.aadd(routes)
self._tool_map.update(missing)
verbose_logger.info(
f"Semantic tool filter indexed {len(routes)} request-time tools missing from the startup index"
)
async def filter_tools(
self,
query: str,
@ -216,22 +293,34 @@ class SemanticMCPToolFilter:
if not query or not query.strip():
return available_tools
# Router should be built on startup - if not, something went wrong
if self.tool_router is None:
verbose_logger.warning("Router not initialized - was build_router_from_mcp_registry() called on startup?")
return available_tools
# Run semantic filtering
try:
await self._ensure_tools_indexed(available_tools)
if self.tool_router is None:
verbose_logger.warning("Semantic router could not be built from the request's tools")
return available_tools
available_names = [name for name in (self._extract_tool_info(t)[0] for t in available_tools) if name]
if not available_names:
return available_tools
limit = top_k or self.top_k
matches = self.tool_router(text=query, limit=limit)
if self.tool_router.top_k < limit:
self.tool_router.top_k = limit
matches = self.tool_router(text=query, limit=limit, route_filter=available_names)
matched_tool_names = self._extract_tool_names_from_matches(matches)
if not matched_tool_names:
return available_tools
return self._get_tools_by_names(matched_tool_names, available_tools)
filtered_tools = self._get_tools_by_names(matched_tool_names, available_tools)
if not filtered_tools:
return available_tools
return filtered_tools
except SemanticToolFilterContextWindowError:
raise
except Exception as e:
if _is_context_window_error(e):
verbose_logger.error(
@ -240,7 +329,7 @@ class SemanticMCPToolFilter:
)
raise SemanticToolFilterContextWindowError(
embedding_model=self.embedding_model,
stage="the user query",
stage="the user query or the MCP tool descriptions being indexed",
original_error=str(e),
) from e
verbose_logger.error(f"Semantic tool filter failed: {e}", exc_info=True)

View file

@ -1265,6 +1265,7 @@ class NewMCPServerRequest(LiteLLMPydanticObjectBase):
command: Optional[str] = None
args: List[str] = Field(default_factory=list)
env: Dict[str, str] = Field(default_factory=dict)
issuer: Optional[str] = None
authorization_url: Optional[str] = None
token_url: Optional[str] = None
registration_url: Optional[str] = None
@ -1370,6 +1371,7 @@ class UpdateMCPServerRequest(LiteLLMPydanticObjectBase):
command: Optional[str] = None
args: List[str] = Field(default_factory=list)
env: Dict[str, str] = Field(default_factory=dict)
issuer: Optional[str] = None
authorization_url: Optional[str] = None
token_url: Optional[str] = None
registration_url: Optional[str] = None
@ -2369,6 +2371,10 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
None,
description="If True, stores request messages and responses in spend logs. Default is False.",
)
disable_auto_add_proxy_admin_to_teams: bool | None = Field(
None,
description="By default, the user calling /team/new is automatically added to the new team as a team admin. If True, proxy admins are no longer auto-added; members explicitly listed in members_with_roles are unaffected. Default is False.",
)
maximum_spend_logs_retention_period: Optional[str] = Field(
None,
description="Maximum retention period for spend logs (e.g., '7d' for 7 days). Logs older than this will be deleted.",

View file

@ -361,7 +361,11 @@ def _global_proxy_budget_check(global_proxy_spend: Optional[float], skip_budget_
and route != "/models"
):
if math.isfinite(litellm.max_budget) and global_proxy_spend > litellm.max_budget:
raise litellm.BudgetExceededError(current_cost=global_proxy_spend, max_budget=litellm.max_budget)
raise litellm.BudgetExceededError(
current_cost=global_proxy_spend,
max_budget=litellm.max_budget,
entity_type=Litellm_EntityType.PROXY.value,
)
_GUARDRAIL_MODIFICATION_KEYS: tuple = (
@ -648,6 +652,8 @@ async def common_checks(
current_cost=user_spend,
max_budget=user_budget,
message=f"ExceededBudget: User={user_object.user_id} over budget. Spend={user_spend}, Budget={user_budget}",
entity_type=Litellm_EntityType.USER.value,
entity_id=user_object.user_id,
)
# Each scope reads a distinct counter key with no cross-scope ordering
@ -1093,6 +1099,8 @@ async def _check_end_user_budget(
current_cost=end_user_spend,
max_budget=end_user_budget,
message=f"ExceededBudget: End User={end_user_obj.user_id} over budget. Spend={end_user_spend}, Budget={end_user_budget}",
entity_type=Litellm_EntityType.END_USER.value,
entity_id=end_user_obj.user_id,
)
@ -3552,6 +3560,8 @@ async def _virtual_key_max_budget_check(
current_cost=spend,
max_budget=valid_token.max_budget,
message=f"Budget has been exceeded! Key={key_descriptor} Current cost: {spend}, Max budget: {valid_token.max_budget}",
entity_type=Litellm_EntityType.KEY.value,
entity_id=valid_token.token,
)
@ -3593,6 +3603,8 @@ async def _virtual_key_multi_budget_check(
f"ExceededBudget: Key over {w['budget_duration']} budget. "
f"Spend=${window_spend:.4f}, Limit=${w['max_budget']:.2f}"
),
entity_type=Litellm_EntityType.KEY.value,
entity_id=valid_token.token,
)
@ -3824,6 +3836,8 @@ async def _check_team_member_budget(
current_cost=team_member_spend,
max_budget=team_member_budget,
message=f"Budget has been exceeded! User={valid_token.user_id} in Team={team_object.team_id} Current cost: {team_member_spend}, Max budget: {team_member_budget}",
entity_type=Litellm_EntityType.TEAM_MEMBER.value,
entity_id=f"{valid_token.user_id}:{team_object.team_id}",
)
@ -3923,6 +3937,8 @@ async def _team_max_budget_check(
current_cost=spend,
max_budget=team_object.max_budget,
message=f"Budget has been exceeded! Team={team_object.team_id} Current cost: {spend}, Max budget: {team_object.max_budget}",
entity_type=Litellm_EntityType.TEAM.value,
entity_id=team_object.team_id,
)
@ -3960,6 +3976,8 @@ async def _team_multi_budget_check(
f"ExceededBudget: Team={team_object.team_id} over {w['budget_duration']} budget. "
f"Spend=${window_spend:.4f}, Limit=${w['max_budget']:.2f}"
),
entity_type=Litellm_EntityType.TEAM.value,
entity_id=team_object.team_id,
)
@ -4081,6 +4099,8 @@ async def _project_max_budget_check(
current_cost=project_object.spend,
max_budget=max_budget,
message=f"Budget has been exceeded! Project={project_object.project_id} Current cost: {project_object.spend}, Max budget: {max_budget}",
entity_type=Litellm_EntityType.PROJECT.value,
entity_id=project_object.project_id,
)
@ -4269,6 +4289,8 @@ async def _organization_max_budget_check(
current_cost=org_spend,
max_budget=org_max_budget,
message=f"Budget has been exceeded! Organization={org_id} Current cost: {org_spend}, Max budget: {org_max_budget}",
entity_type=Litellm_EntityType.ORGANIZATION.value,
entity_id=org_id,
)
@ -4326,6 +4348,8 @@ async def _tag_max_budget_check(
current_cost=tag_spend,
max_budget=tag_object.litellm_budget_table.max_budget,
message=f"Budget has been exceeded! Tag={tag_name} Current cost: {tag_spend}, Max budget: {tag_object.litellm_budget_table.max_budget}",
entity_type=Litellm_EntityType.TAG.value,
entity_id=tag_name,
)

View file

@ -1797,6 +1797,8 @@ async def _user_api_key_auth_builder(
raise litellm.BudgetExceededError(
current_cost=team_member_spend,
max_budget=team_member_budget,
entity_type=Litellm_EntityType.TEAM_MEMBER.value,
entity_id=f"{valid_token.user_id}:{valid_token.team_id}",
)
# Check 3. If token is expired
@ -1994,16 +1996,6 @@ async def _user_api_key_auth_builder(
raise HTTPException(401, detail="Invalid API key, no token associated")
api_key = valid_token.token
# Add hashed token to cache
asyncio.create_task(
_cache_key_object(
hashed_token=api_key,
user_api_key_obj=valid_token,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
)
valid_token_dict = valid_token.model_dump(exclude_none=True)
valid_token_dict.pop("token", None)
# budget_throttle_pct is excluded from model_dump (it must not leak

View file

@ -495,19 +495,21 @@ Cursor is not supported: it has no equivalent file-based config to hot-patch thi
#### Install the CLI
If you don't already have the `lite` command, install it with a single curl command -- no existing Python tooling required, `uv` is bootstrapped automatically if missing:
`lite autoroute up` builds and runs a throwaway litellm proxy locally, so unlike the rest of this CLI it needs the proxy server runtime, not just the thin `litellm[cli]` client. Install `litellm[proxy]` (which ships the `lite` command too) with a single curl command -- no existing Python tooling required, `uv` is bootstrapped automatically if missing:
```bash
curl -fsSL https://raw.githubusercontent.com/BerriAI/litellm/main/scripts/install-cli.sh | sh
curl -fsSL https://raw.githubusercontent.com/BerriAI/litellm/main/scripts/install.sh | sh
```
This installs only `litellm[cli]`, the thin client (`lite`), not the full proxy server. To try an unreleased branch or commit instead of the latest PyPI release, set `LITELLM_CLI_REF`:
To QA an unreleased branch or commit instead of the latest PyPI release, set `LITELLM_CLI_REF`:
```bash
curl -fsSL https://raw.githubusercontent.com/BerriAI/litellm/<branch-or-commit>/scripts/install-cli.sh | \
curl -fsSL https://raw.githubusercontent.com/BerriAI/litellm/<branch-or-commit>/scripts/install.sh | \
LITELLM_CLI_REF=<branch-or-commit> sh
```
The thin `scripts/install-cli.sh` installs only `litellm[cli]`, which is enough for `lite login`, `lite claude`, and `lite up`, but not for `lite autoroute up`; running it against a `litellm[cli]` install fails fast with a message telling you to install the proxy runtime.
Point the CLI at your real proxy and key before running any `lite model-groups` or `lite autoroute` command -- like every other command in this CLI, they read `LITELLM_PROXY_URL`/`LITELLM_PROXY_API_KEY` (or `--base-url`/`--api-key`), no `lite login` required:
```bash

View file

@ -72,7 +72,7 @@ def display_teams_table(teams: List[Dict[str, Any]]) -> None:
console = Console()
if not teams:
console.print("❌ No teams found for your user.")
console.print("No teams found for your user.")
return
table = Table(title="Available Teams")
@ -162,7 +162,7 @@ def display_interactive_team_selection(teams: List[Dict[str, Any]], selected_ind
# Clear the screen using Rich's method
console.clear()
console.print("🎯 Select a Team (Use ↑↓ arrows, Enter to select, 'q' to skip):\n")
console.print("Select a Team (Use up/down arrows, Enter to select, 'q' to skip):\n")
for i, team in enumerate(teams):
team_alias = team.get("team_alias") or "N/A"
@ -184,7 +184,7 @@ def display_interactive_team_selection(teams: List[Dict[str, Any]], selected_ind
# Highlight the selected item
if i == selected_index:
console.print(f"➤ [bold cyan]{team_alias}[/bold cyan] ({team_id})")
console.print(f"> [bold cyan]{team_alias}[/bold cyan] ({team_id})")
console.print(f" Models: [yellow]{models_str}[/yellow]")
console.print(f" Budget: [blue]{budget_str}[/blue]\n")
else:
@ -220,15 +220,13 @@ def prompt_team_selection(teams: List[Dict[str, Any]]) -> Optional[Dict[str, Any
# Clear screen and show selection
console = Console()
console.clear()
click.echo(
f"✅ Selected team: {selected_team.get('team_alias', 'N/A')} ({selected_team.get('team_id')})"
)
click.echo(f"Selected team: {selected_team.get('team_alias', 'N/A')} ({selected_team.get('team_id')})")
return selected_team
elif key == "quit" or key == "escape":
# Clear screen
console = Console()
console.clear()
click.echo("ℹ️ Team selection skipped.")
click.echo("Team selection skipped.")
return None
elif key is None:
# If we can't get key input, fall back to simple selection
@ -237,7 +235,7 @@ def prompt_team_selection(teams: List[Dict[str, Any]]) -> Optional[Dict[str, Any
except KeyboardInterrupt:
console = Console()
console.clear()
click.echo("\n❌ Team selection cancelled.")
click.echo("\nTeam selection cancelled.")
return None
except Exception:
# If interactive mode fails, fall back to simple selection
@ -265,15 +263,15 @@ def prompt_team_selection_fallback(
if 0 <= index < len(teams):
selected_team = teams[index]
click.echo(
f"\n✅ Selected team: {selected_team.get('team_alias', 'N/A')} ({selected_team.get('team_id')})"
f"\nSelected team: {selected_team.get('team_alias', 'N/A')} ({selected_team.get('team_id')})"
)
return selected_team
else:
click.echo(f"❌ Invalid selection. Please enter a number between 1 and {len(teams)}")
click.echo(f"Invalid selection. Please enter a number between 1 and {len(teams)}")
except ValueError:
click.echo("❌ Invalid input. Please enter a number or 'skip'")
click.echo("Invalid input. Please enter a number or 'skip'")
except KeyboardInterrupt:
click.echo("\n❌ Team selection cancelled.")
click.echo("\nTeam selection cancelled.")
return None
@ -437,7 +435,7 @@ def _poll_for_authentication(base_url: str, key_id: str, poll_secret: str) -> Op
user_id = data.get("user_id")
normalized_teams: List[Dict[str, Any]] = _normalize_teams(teams, team_details)
if not normalized_teams:
click.echo("⚠️ No teams available for selection.")
click.echo("Warning: No teams available for selection.")
return None
# User has multiple teams - let them select
@ -457,7 +455,7 @@ def _poll_for_authentication(base_url: str, key_id: str, poll_secret: str) -> Op
"team_id": None, # Set by server in JWT
}
click.echo("❌ Team selection cancelled or JWT generation failed.")
click.echo("Team selection cancelled or JWT generation failed.")
return None
# JWT is ready (single team or team already selected)
@ -468,7 +466,7 @@ def _poll_for_authentication(base_url: str, key_id: str, poll_secret: str) -> Op
# Show which team was assigned
if team_id and len(teams) == 1:
click.echo(f"\n✅ Automatically assigned to team: {team_id}")
click.echo(f"\nAutomatically assigned to team: {team_id}")
if api_key:
return {
@ -494,19 +492,19 @@ def _handle_team_selection_during_polling(
The JWT token with the selected team, or None if selection was skipped
"""
if not teams:
click.echo("ℹ️ No teams found. You can create or join teams using the web interface.")
click.echo("No teams found. You can create or join teams using the web interface.")
return None
click.echo("\n" + "=" * 60)
click.echo("📋 Select a team for your CLI session...")
click.echo("Select a team for your CLI session...")
team_id = _render_and_prompt_for_team_selection(teams)
if not team_id:
click.echo("ℹ️ No team selected.")
click.echo("No team selected.")
return None
click.echo(f"\n🔄 Generating JWT for team: {team_id}")
click.echo(f"\nGenerating JWT for team: {team_id}")
poll_url = f"{base_url}/sso/cli/poll/{key_id}?team_id={team_id}"
data = _poll_for_ready_data(
@ -520,7 +518,7 @@ def _handle_team_selection_during_polling(
return None
jwt_token = data.get("key")
if jwt_token:
click.echo(f"✅ Successfully generated JWT for team: {team_id}")
click.echo(f"Successfully generated JWT for team: {team_id}")
return jwt_token
return None
@ -568,14 +566,14 @@ def _render_and_prompt_for_team_selection(teams: List[Dict[str, Any]]) -> Option
selected_team = teams[index]
team_id = str(selected_team.get("team_id"))
team_alias = selected_team.get("team_alias") or team_id
click.echo(f"\n✅ Selected team: {team_alias} ({team_id})")
click.echo(f"\nSelected team: {team_alias} ({team_id})")
return team_id
click.echo(f"❌ Invalid selection. Please enter a number between 1 and {len(teams)}")
click.echo(f"Invalid selection. Please enter a number between 1 and {len(teams)}")
except ValueError:
click.echo("❌ Invalid input. Please enter a number or 'skip'")
click.echo("Invalid input. Please enter a number or 'skip'")
except KeyboardInterrupt:
click.echo("\n❌ Team selection cancelled.")
click.echo("\nTeam selection cancelled.")
return None
@ -628,7 +626,7 @@ def login(ctx: click.Context):
}
)
click.echo("\n✅ Login successful!")
click.echo("\nLogin successful!")
click.echo(f"JWT Token: {api_key[:20]}...")
click.echo("You can now use the CLI without specifying --api-key")
@ -637,7 +635,7 @@ def login(ctx: click.Context):
show_commands()
return
else:
click.echo("❌ Authentication timed out. Please try again.")
click.echo("Authentication timed out. Please try again.")
click.echo(
"The proxy never reported the browser sign-in as finished. If you did complete it, "
"check the proxy logs for /sso/callback errors and confirm SSO is configured on the proxy."
@ -645,10 +643,10 @@ def login(ctx: click.Context):
return
except KeyboardInterrupt:
click.echo("\n❌ Authentication cancelled by user.")
click.echo("\nAuthentication cancelled by user.")
return
except Exception as e:
click.echo(f"❌ Authentication failed: {e}")
click.echo(f"Authentication failed: {e}")
return
@ -656,7 +654,7 @@ def login(ctx: click.Context):
def logout():
"""Logout and clear stored authentication"""
clear_token()
click.echo("✅ Logged out successfully. Authentication token cleared.")
click.echo("Logged out successfully. Authentication token cleared.")
@click.command(name="print-token")
@ -703,10 +701,10 @@ def whoami():
token_data = load_token()
if not token_data:
click.echo("❌ Not authenticated. Run 'lite login' to authenticate.")
click.echo("Not authenticated. Run 'lite login' to authenticate.")
return
click.echo("✅ Authenticated")
click.echo("Authenticated")
click.echo(f"User Email: {token_data.get('user_email', 'Unknown')}")
click.echo(f"User ID: {token_data.get('user_id', 'Unknown')}")
click.echo(f"User Role: {token_data.get('user_role', 'Unknown')}")
@ -717,7 +715,7 @@ def whoami():
click.echo(f"Token age: {age_hours:.1f} hours")
if age_hours > CLI_JWT_EXPIRATION_HOURS:
click.echo(f"⚠️ Warning: Token is more than {CLI_JWT_EXPIRATION_HOURS} hours old and may have expired.")
click.echo(f"Warning: Token is more than {CLI_JWT_EXPIRATION_HOURS} hours old and may have expired.")
@click.group(name="auth")

View file

@ -21,6 +21,7 @@ from .process import (
clear_pid_record,
is_running,
launch_proxy,
missing_proxy_runtime_modules,
poll_liveliness,
read_pid_record,
secure_create,
@ -81,6 +82,16 @@ def up() -> None:
if not CONFIG_PATH.exists():
raise click.ClickException("No config found. Run `lite autoroute configure` first.")
missing = missing_proxy_runtime_modules()
if missing:
raise click.ClickException(
"lite autoroute up launches a local litellm proxy, which needs the proxy runtime that the "
f"thin `litellm[cli]` install does not include (missing: {', '.join(missing)}). Install the "
"proxy runtime with `uv tool install --force 'litellm[proxy]'`, or to QA a branch, "
"`curl -fsSL https://raw.githubusercontent.com/BerriAI/litellm/<branch>/scripts/install.sh | "
"LITELLM_CLI_REF=<branch> sh`."
)
try:
existing_pid = read_pid_record()
except UpError as e:

View file

@ -80,24 +80,32 @@ class NoSemanticMatching(BaseModel):
kind: Literal["none"] = "none"
class KeywordTierRule(BaseModel):
model_config = ConfigDict(frozen=True)
keywords: tuple[str, ...]
tier: str
# Satisfies complexity_router's "semantic matching requires non-empty keyword_tier_rules"
# invariant with a sane starting point; the wizard lets the user override these per tier.
DEFAULT_KEYWORD_TIER_RULES: tuple[KeywordTierRule, ...] = (
KeywordTierRule(keywords=("hi", "hello", "thanks"), tier="SIMPLE"),
KeywordTierRule(keywords=("explain", "how does"), tier="MEDIUM"),
KeywordTierRule(keywords=("refactor", "implement", "debug"), tier="COMPLEX"),
KeywordTierRule(keywords=("step by step", "think through", "prove"), tier="REASONING"),
)
class SemanticMatching(BaseModel):
model_config = ConfigDict(frozen=True)
kind: Literal["semantic"] = "semantic"
embedding_model: str
match_threshold: float = 0.5
keyword_tier_rules: tuple[KeywordTierRule, ...] = DEFAULT_KEYWORD_TIER_RULES
SemanticMatchingChoice = Union[NoSemanticMatching, SemanticMatching]
# Satisfies complexity_router's "semantic matching requires non-empty keyword_tier_rules"
# invariant with a sane starting point; the generated config.yaml can be hand-edited afterward.
_DEFAULT_KEYWORD_TIER_RULES: tuple[dict[str, JsonValue], ...] = (
{"keywords": ["hi", "hello", "thanks"], "tier": "SIMPLE"},
{"keywords": ["explain", "how does"], "tier": "MEDIUM"},
{"keywords": ["refactor", "implement", "debug"], "tier": "COMPLEX"},
{"keywords": ["step by step", "think through", "prove"], "tier": "REASONING"},
)
class AutorouteConfig(BaseModel):
model_config = ConfigDict(frozen=True)
@ -181,7 +189,9 @@ def build_generated_model_list(config: AutorouteConfig) -> list[JsonValue]:
complexity_router_config["semantic_keyword_matching"] = True
complexity_router_config["embedding_model"] = config.semantic_matching.embedding_model
complexity_router_config["match_threshold"] = config.semantic_matching.match_threshold
complexity_router_config["keyword_tier_rules"] = list(_DEFAULT_KEYWORD_TIER_RULES)
complexity_router_config["keyword_tier_rules"] = [
{"keywords": list(rule.keywords), "tier": rule.tier} for rule in config.semantic_matching.keyword_tier_rules
]
if config.adaptive:
complexity_router_config["adaptive"] = True
@ -222,8 +232,10 @@ __all__ = [
"AutorouteConfig",
"ClassifierChoice",
"ConfigGenerationError",
"DEFAULT_KEYWORD_TIER_RULES",
"DiscoveredModel",
"HeuristicClassifier",
"KeywordTierRule",
"LLMClassifier",
"NoSemanticMatching",
"SemanticMatching",

View file

@ -1,4 +1,5 @@
import contextlib
import importlib.util
import json
import os
import signal
@ -37,6 +38,20 @@ class PidRecord:
_PID_RECORD_ADAPTER = TypeAdapter(PidRecord)
_PROXY_RUNTIME_MODULES: tuple[str, ...] = ("fastapi", "uvicorn", "backoff", "orjson", "websockets", "apscheduler")
def missing_proxy_runtime_modules() -> tuple[str, ...]:
"""Proxy-server modules that ``lite autoroute up`` needs but the thin CLI install lacks.
``launch_proxy`` runs the full ``litellm.proxy.proxy_cli`` server, whose dependencies live in
the ``proxy`` extra, not the ``cli`` extra that installs the ``lite`` command. On a thin
``litellm[cli]`` install the subprocess dies with a bare ``ModuleNotFoundError``; detecting the
gap here lets ``up`` fail with an actionable message instead.
"""
return tuple(name for name in _PROXY_RUNTIME_MODULES if importlib.util.find_spec(name) is None)
def allocate_free_port() -> int:
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as sock:
sock.bind(("127.0.0.1", 0))
@ -165,6 +180,7 @@ __all__ = [
"clear_pid_record",
"is_running",
"launch_proxy",
"missing_proxy_runtime_modules",
"poll_liveliness",
"read_pid_record",
"secure_create",

View file

@ -8,11 +8,13 @@ from InquirerPy.base.control import Choice
from .... import Client
from .config import (
DEFAULT_KEYWORD_TIER_RULES,
TIER_NAMES,
AutorouteConfig,
ConfigGenerationError,
DiscoveredModel,
HeuristicClassifier,
KeywordTierRule,
LLMClassifier,
NoSemanticMatching,
SemanticMatching,
@ -63,6 +65,25 @@ def _render_and_prompt_for_models(models: tuple[DiscoveredModel, ...], prompt_la
return tuple(_fuzzy_pick(models, prompt_label, multiselect=True))
def _parse_keywords(raw: str) -> tuple[str, ...]:
return tuple(keyword.strip() for keyword in raw.split(",") if keyword.strip())
def _prompt_for_keyword_tier_rules() -> tuple[KeywordTierRule, ...]:
"""Let the user supply the semantic-matching keywords per tier, since matching those
keywords against the request is the whole point of enabling it. Each prompt is prefilled
with the built-in default, so pressing enter keeps it."""
click.echo("\nEnter example keywords/phrases per tier (comma-separated); press enter to keep the default:")
defaults = {rule.tier: rule.keywords for rule in DEFAULT_KEYWORD_TIER_RULES}
def _rule_for(tier: str) -> KeywordTierRule:
default_keywords = defaults.get(tier, ())
raw = click.prompt(f" {tier} keywords", default=", ".join(default_keywords), show_default=True)
return KeywordTierRule(keywords=_parse_keywords(raw) or default_keywords, tier=tier)
return tuple(_rule_for(tier) for tier in TIER_NAMES)
def run_configure_wizard(ctx: click.Context) -> Path:
"""Discover the caller's accessible models, walk them through tier assignment, write config."""
base_url = ctx.obj["base_url"]
@ -96,7 +117,8 @@ def run_configure_wizard(ctx: click.Context) -> Path:
semantic_matching = NoSemanticMatching()
if embedding_pool and click.confirm("\nEnable semantic keyword matching?", default=False):
embedding_model = _render_and_prompt_for_model(embedding_pool, "semantic embeddings")
semantic_matching = SemanticMatching(embedding_model=embedding_model)
keyword_tier_rules = _prompt_for_keyword_tier_rules()
semantic_matching = SemanticMatching(embedding_model=embedding_model, keyword_tier_rules=keyword_tier_rules)
adaptive = click.confirm("\nEnable adaptive (bandit) selection on top of tiering?", default=False)

View file

@ -150,7 +150,7 @@ def chat(
f"Max Tokens: [yellow]{max_tokens or 'unlimited'}[/yellow]\n\n"
f"Type your messages and press Enter. Type '/quit' or '/exit' to end the session.\n"
f"Type '/help' for more commands.",
title="🤖 Chat Session",
title="Chat Session",
)
)

View file

@ -32,7 +32,7 @@ def migrate(ctx: click.Context, check_only: bool, dry_run: bool):
Requires the proxy to be started with
``general_settings.encryption_algorithm: aes-256-gcm``. Idempotent and
resumable — safe to re-run after an interruption.
resumable; safe to re-run after an interruption.
Examples:
litellm-proxy encryption migrate --check # attestation scan, no writes

View file

@ -309,12 +309,12 @@ def _import_keys_to_destination(
imported_count += 1
key_alias = key.get("key_alias", "N/A")
click.echo(f"✓ Imported key: {key_alias}")
click.echo(f"Imported key: {key_alias}")
except Exception as e:
failed_count += 1
key_alias = key.get("key_alias", "N/A")
click.echo(f"✗ Failed to import key {key_alias}: {str(e)}", err=True)
click.echo(f"Failed to import key {key_alias}: {str(e)}", err=True)
return imported_count, failed_count

View file

@ -21,7 +21,7 @@ def display_teams_table(teams: List[Dict[str, Any]]) -> None:
console = Console()
if not teams:
console.print("❌ No teams found for your user.")
console.print("No teams found for your user.")
return
table = Table(title="Available Teams")
@ -91,10 +91,10 @@ def available(ctx: click.Context):
teams = client.teams.get_available()
if teams:
console = Console()
console.print("\n🎯 Available Teams to Join:")
console.print("\nAvailable Teams to Join:")
display_teams_table(teams)
else:
click.echo("ℹ️ No available teams to join.")
click.echo("No available teams to join.")
except requests.exceptions.HTTPError as e:
click.echo(f"Error: HTTP {e.response.status_code}", err=True)
error_body = e.response.json()
@ -113,7 +113,7 @@ def assign_key(ctx: click.Context, team_id: Optional[str]):
api_key = ctx.obj["api_key"]
if not api_key:
click.echo("❌ No API key found. Please login first using 'litellm login'")
click.echo("No API key found. Please login first using 'litellm login'")
raise click.Abort()
try:
@ -122,7 +122,7 @@ def assign_key(ctx: click.Context, team_id: Optional[str]):
teams = client.teams.list()
if not teams:
click.echo("❌ No teams found for your user.")
click.echo("No teams found for your user.")
return
# Use interactive selection from auth module
@ -133,14 +133,14 @@ def assign_key(ctx: click.Context, team_id: Optional[str]):
if selected_team:
team_id = selected_team.get("team_id")
else:
click.echo("❌ Operation cancelled.")
click.echo("Operation cancelled.")
return
# Update the key with the selected team
if team_id:
click.echo(f"\n🔄 Assigning your key to team: {team_id}")
click.echo(f"\nAssigning your key to team: {team_id}")
client.keys.update(key=api_key, team_id=team_id)
click.echo(f"✅ Successfully assigned key to team: {team_id}")
click.echo(f"Successfully assigned key to team: {team_id}")
# Show team details if available
teams = client.teams.list()
@ -148,9 +148,9 @@ def assign_key(ctx: click.Context, team_id: Optional[str]):
if team.get("team_id") == team_id:
models = team.get("models", [])
if models:
click.echo(f"🎯 You can now access models: {', '.join(models)}")
click.echo(f"You can now access models: {', '.join(models)}")
else:
click.echo("🎯 You can now access all available models")
click.echo("You can now access all available models")
break
except requests.exceptions.HTTPError as e:

View file

@ -27,13 +27,13 @@ def styled_prompt():
verbose_logger.debug(f"Error getting terminal size: {e}")
click.echo("\n" * 3)
# Unicode box drawing characters
top_left = "┌"
top_right = "┐"
bottom_left = "└"
bottom_right = "┘"
horizontal = "─"
vertical = "│"
# ASCII box drawing characters
top_left = "+"
top_right = "+"
bottom_left = "+"
bottom_right = "+"
horizontal = "-"
vertical = "|"
# Create the box with increased width
width = 80

View file

@ -27,7 +27,7 @@ from typing import (
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.caching import DualCache, RedisCache
from litellm.caching import RedisCache
from litellm.constants import (
DB_SPEND_UPDATE_JOB_NAME,
DB_DAILY_TAG_SPEND_UPDATE_JOB_NAME,
@ -44,7 +44,6 @@ from litellm.proxy._types import (
DailyUserSpendTransaction,
DBSpendUpdateTransactions,
Litellm_EntityType,
LiteLLM_UserTable,
SpendLogsMetadata,
SpendLogsPayload,
SpendUpdateQueueItem,
@ -137,7 +136,6 @@ class DBSpendUpdateWriter:
disable_spend_logs,
litellm_proxy_budget_name,
prisma_client,
user_api_key_cache,
)
from litellm.proxy.utils import ProxyUpdateSpend, hash_token
@ -195,7 +193,6 @@ class DBSpendUpdateWriter:
org_id=org_id,
end_user_id=end_user_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
litellm_proxy_budget_name=litellm_proxy_budget_name,
payload=payload,
)
@ -326,7 +323,6 @@ class DBSpendUpdateWriter:
org_id: Optional[str],
end_user_id: Optional[str],
prisma_client: Optional[PrismaClient],
user_api_key_cache: DualCache,
litellm_proxy_budget_name: Optional[str],
payload: SpendLogsPayload,
):
@ -345,7 +341,6 @@ class DBSpendUpdateWriter:
response_cost=response_cost,
user_id=user_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
litellm_proxy_budget_name=litellm_proxy_budget_name,
end_user_id=end_user_id,
)
@ -510,7 +505,6 @@ class DBSpendUpdateWriter:
response_cost: Optional[float],
user_id: Optional[str],
prisma_client: Optional[PrismaClient],
user_api_key_cache: DualCache,
litellm_proxy_budget_name: Optional[str],
end_user_id: Optional[str] = None,
):
@ -518,10 +512,6 @@ class DBSpendUpdateWriter:
- Update that user's row
- Update litellm-proxy-budget row (global proxy spend)
"""
## if an end-user is passed in, do an upsert - we can't guarantee they already exist in db
existing_user_obj = await user_api_key_cache.async_get_cache(key=user_id)
if existing_user_obj is not None and isinstance(existing_user_obj, dict):
existing_user_obj = LiteLLM_UserTable(**existing_user_obj)
try:
if prisma_client is not None: # update
user_ids = [user_id]

View file

@ -0,0 +1,212 @@
"""Supervisor-side reaper for orphaned Prisma query-engine processes.
Each proxy worker owns a Prisma query-engine subprocess whose only cleanup
hook is an in-process ``atexit`` handler. When a multi-worker supervisor
(uvicorn's multiprocess manager, the gunicorn arbiter) force-kills a hung or
crashed worker, that handler never runs: the engine reparents to the nearest
subreaper (PID 1 in a container, which is the supervisor itself under the
standard docker entrypoint) and keeps its database connection pool
established forever, while the replacement worker opens a fresh pool. Over
repeated worker deaths the active database connections grow without bound.
The reaper runs only in the supervisor process, where a query-engine process
can never be a legitimate direct child: workers own their engines, and the
supervisor never starts one. Any direct child whose command name begins with
``query-engine`` is therefore an adopted orphan and is terminated
(SIGTERM, bounded grace, SIGKILL) and reaped. On Linux the supervisor also
marks itself a child subreaper so orphans reparent to it even when it is not
PID 1.
Linux-only by construction (``/proc`` scan, ``prctl``); a no-op elsewhere.
"""
import ctypes
import os
import signal
import sys
import threading
import time
from typing import Optional
from litellm._logging import verbose_proxy_logger
QUERY_ENGINE_COMM_PREFIX = "query-engine"
REAPER_SCAN_INTERVAL_SECONDS = 5.0
SIGTERM_GRACE_SECONDS = 10.0
PR_SET_CHILD_SUBREAPER = 36
def set_child_subreaper() -> bool:
"""Mark this process as a child subreaper so orphaned descendants
reparent to it instead of PID 1. Best-effort: when it fails (or on
non-Linux) the reaper still covers the containerized case where the
supervisor already is PID 1."""
if not sys.platform.startswith("linux"):
return False
try:
libc = ctypes.CDLL(None, use_errno=True)
result: int = libc.prctl( # pyright: ignore[reportAny] # ctypes types foreign calls as Any; default restype is c_int
PR_SET_CHILD_SUBREAPER, 1, 0, 0, 0
)
return result == 0
except (OSError, AttributeError):
return False
def _read_comm_and_ppid(pid: int, proc_root: str) -> Optional[tuple[str, int]]:
try:
with open(f"{proc_root}/{pid}/stat", encoding="ascii", errors="replace") as stat_file:
data = stat_file.read()
except (FileNotFoundError, ProcessLookupError, PermissionError, OSError):
return None
lparen = data.find("(")
rparen = data.rfind(")")
if lparen == -1 or rparen == -1 or rparen < lparen:
return None
comm = data[lparen + 1 : rparen]
fields = data[rparen + 2 :].split()
if len(fields) < 2:
return None
try:
ppid = int(fields[1])
except ValueError:
return None
return comm, ppid
def list_orphaned_engine_pids(parent_pid: int, proc_root: str = "/proc") -> tuple[int, ...]:
"""PIDs of direct children of ``parent_pid`` whose command name marks
them as Prisma query engines. In the supervisor these are always
adopted orphans: live engines are children of workers, not of the
supervisor."""
try:
entries = os.listdir(proc_root)
except (FileNotFoundError, OSError):
return ()
candidate_pids = (int(entry) for entry in entries if entry.isdigit())
return tuple(
pid
for pid in candidate_pids
if (info := _read_comm_and_ppid(pid, proc_root)) is not None
and info[1] == parent_pid
and info[0].startswith(QUERY_ENGINE_COMM_PREFIX)
)
def _try_reap(pid: int) -> bool:
try:
reaped_pid, _ = os.waitpid(pid, os.WNOHANG)
except ChildProcessError:
return True
except OSError:
return True
return reaped_pid == pid
def _send_signal(pid: int, signum: int) -> None:
try:
os.kill(pid, signum)
except (ProcessLookupError, PermissionError, OSError):
pass
def _await_reaped(pids: tuple[int, ...], timeout_seconds: float) -> tuple[int, ...]:
"""Poll until every PID is reaped or the shared deadline passes.
Returns the PIDs still alive at the deadline."""
deadline = time.monotonic() + timeout_seconds
remaining = pids
while remaining and time.monotonic() < deadline:
remaining = tuple(pid for pid in remaining if not _try_reap(pid))
if remaining:
time.sleep(0.2)
return remaining
def terminate_and_reap(pid: int, grace_seconds: float = SIGTERM_GRACE_SECONDS) -> None:
"""SIGTERM the orphaned engine, escalate to SIGKILL after the grace
period, and reap it so it does not linger as a zombie."""
terminate_and_reap_all((pid,), grace_seconds=grace_seconds)
def terminate_and_reap_all(
pids: tuple[int, ...],
grace_seconds: float = SIGTERM_GRACE_SECONDS,
) -> None:
"""Terminate a batch of orphaned engines concurrently: SIGTERM all of
them, share one grace period, SIGKILL the stragglers, and reap. The
shared deadline keeps cleanup time bounded when several workers die
at once instead of paying the grace period once per orphan."""
for pid in pids:
verbose_proxy_logger.warning(
"Reaping orphaned prisma query-engine PID %s (its worker process exited without cleanup).",
pid,
)
_send_signal(pid, signal.SIGTERM)
survivors = _await_reaped(pids, grace_seconds)
if not survivors:
return
for pid in survivors:
verbose_proxy_logger.warning(
"Orphaned prisma query-engine PID %s did not exit within %.1fs of SIGTERM; sending SIGKILL.",
pid,
grace_seconds,
)
_send_signal(pid, signal.SIGKILL)
unkillable = _await_reaped(survivors, 5.0)
for pid in unkillable:
verbose_proxy_logger.error(
"Orphaned prisma query-engine PID %s survived SIGKILL; will retry on the next scan.",
pid,
)
def reap_orphaned_engines(parent_pid: int, proc_root: str = "/proc") -> tuple[int, ...]:
"""One scan-and-reap pass. Returns the PIDs it acted on."""
orphaned_pids = list_orphaned_engine_pids(parent_pid, proc_root=proc_root)
if orphaned_pids:
terminate_and_reap_all(orphaned_pids)
return orphaned_pids
def _reaper_loop(parent_pid: int) -> None:
while True:
try:
reap_orphaned_engines(parent_pid)
except Exception as scan_error: # noqa: BLE001 # reaper thread must survive any scan failure
verbose_proxy_logger.debug("Orphaned query-engine scan failed: %s", scan_error)
time.sleep(REAPER_SCAN_INTERVAL_SECONDS)
REAPER_THREAD_NAME = "litellm-orphan-query-engine-reaper"
def start_query_engine_reaper() -> Optional[threading.Thread]:
"""Start the reaper daemon thread in the supervisor process.
Must only be called from a process that never hosts the proxy app
itself (uvicorn with ``workers > 1``, the gunicorn arbiter): with a
single in-process uvicorn worker the query engine is a legitimate
direct child and must not be touched. Idempotent: a reaper already
running in this process is returned instead of starting a second one.
"""
if not sys.platform.startswith("linux"):
return None
existing = next(
(thread for thread in threading.enumerate() if thread.name == REAPER_THREAD_NAME),
None,
)
if existing is not None:
return existing
set_child_subreaper()
reaper_thread = threading.Thread(
target=_reaper_loop,
args=(os.getpid(),),
daemon=True,
name=REAPER_THREAD_NAME,
)
reaper_thread.start()
verbose_proxy_logger.info(
"Started orphaned prisma query-engine reaper in supervisor process %s.",
os.getpid(),
)
return reaper_thread

View file

@ -26,6 +26,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
mask_request_content=litellm_params.mask_request_content,
mask_response_content=litellm_params.mask_response_content,
fail_on_error=litellm_params.fail_on_error,
skip_unscannable_attachments=litellm_params.skip_unscannable_attachments,
)
litellm.logging_callback_manager.add_litellm_callback(_model_armor_callback)

View file

@ -25,10 +25,6 @@ from litellm.types.llms.openai import AllMessageValues
MODEL_ARMOR_MAX_FILE_SIZE_BYTES = 4 * 1024 * 1024
# Hard cap on how many attachments a single request may submit to Model Armor, to bound
# per-request fan-out (latency and quota).
MAX_FILE_ATTACHMENTS_PER_REQUEST = 10
_REMOTE_URI_SCHEMES = ("gs://", "http://", "https://")
ModelArmorByteDataType = Literal["PDF", "WORD_DOCUMENT", "EXCEL_DOCUMENT", "POWERPOINT_DOCUMENT", "CSV", "TXT"]

View file

@ -32,7 +32,6 @@ from litellm.llms.custom_httpx.http_handler import (
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.guardrails.guardrail_hooks.model_armor.file_scanning import (
MAX_FILE_ATTACHMENTS_PER_REQUEST,
MODEL_ARMOR_MAX_FILE_SIZE_BYTES,
plan_file_scans,
)
@ -383,10 +382,14 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
Each attachment is sent through the byte API and a MATCH_FOUND raises a 400 before the
request reaches the LLM. File scanning does not support masking (Model Armor returns
findings, not a sanitized document), so it only blocks. Anything the guardrail cannot
scan - a file_id or remote URL reference with no inline bytes, a document over the 4 MB
byte limit, or more attachments than the per-request cap - is a guardrail failure and
blocks unless the operator has opted into fail-open via fail_on_error=False.
findings, not a sanitized document), so it only blocks. A file_id or remote URL reference
with no inline bytes and a document over the 4 MB byte limit are guardrail failures that
block unless the operator has opted into fail-open via fail_on_error=False.
skip_unscannable_attachments decouples reference-only attachments from fail_on_error: when
enabled, attachments Model Armor cannot scan (file_id, gs://, or http(s) references with no
inline bytes, and inline content whose base64 will not decode) pass through instead of
blocking, while fail_on_error still governs real Model Armor API errors.
"""
from litellm.proxy.common_utils.callback_utils import (
_get_or_create_proxy_metadata_bucket,
@ -395,7 +398,14 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
plan = plan_file_scans(messages)
attachments = plan.attachments
unscannable_references = plan.unscannable_count
skip_unscannable = bool(self.optional_params.get("skip_unscannable_attachments", False))
if skip_unscannable and plan.unscannable_count > 0:
verbose_proxy_logger.warning(
"Model Armor: allowing %d unscannable attachment(s) through because "
"skip_unscannable_attachments is enabled",
plan.unscannable_count,
)
unscannable_references = 0 if skip_unscannable else plan.unscannable_count
if not attachments and unscannable_references == 0:
return
@ -415,14 +425,6 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
metadata["_model_armor_status"] = "blocked"
raise self._unscannable_block_error(reason)
if len(attachments) > MAX_FILE_ATTACHMENTS_PER_REQUEST:
reason = f"{len(attachments)} attachments exceed the per-request scan limit of {MAX_FILE_ATTACHMENTS_PER_REQUEST}"
verbose_proxy_logger.warning("Model Armor: %s", reason)
if fail_on_error:
metadata["_model_armor_status"] = "blocked"
raise self._unscannable_block_error(reason)
attachments = attachments[:MAX_FILE_ATTACHMENTS_PER_REQUEST]
for attachment in attachments:
if len(attachment.file_bytes) > MODEL_ARMOR_MAX_FILE_SIZE_BYTES:
reason = (

View file

@ -494,6 +494,7 @@ class InMemoryGuardrailHandler:
guardrail_id=guardrail.get("guardrail_id"),
guardrail_name=guardrail["guardrail_name"],
litellm_params=litellm_params,
guardrail_info=guardrail.get("guardrail_info"),
)
# store references to the guardrail in memory
@ -612,6 +613,27 @@ class InMemoryGuardrailHandler:
"""
return self._sources.get(guardrail_id)
def list_config_guardrails(self) -> List[Guardrail]:
"""
List in-memory guardrails owned by config.yaml.
DB-sourced entries are excluded: a read surface that also queries the DB
would double-count live ones, and a DB-sourced entry that's missing from
the DB is stale (deleted on another pod, awaiting reconciliation here).
"""
return [g for gid, g in self.IN_MEMORY_GUARDRAILS.items() if self._sources.get(gid) == "config"]
def get_config_guardrail_by_id(self, guardrail_id: str) -> Optional[Guardrail]:
"""
Get a config-owned in-memory guardrail by its ID, or None.
Mirrors the fallback in get_guardrail_info: a DB-sourced in-memory entry
that missed the DB lookup is stale and must not be surfaced.
"""
if self._sources.get(guardrail_id) != "config":
return None
return self.IN_MEMORY_GUARDRAILS.get(guardrail_id)
def reconcile_db_guardrails(self, db_guardrail_ids: Set[str]) -> List[str]:
"""
Drop in-memory entries that originated from the DB but are no longer

View file

@ -137,10 +137,26 @@ def _chart_from_metrics(metrics: Any) -> List[Dict[str, Any]]:
return [{"date": d, "passed": v["passed"], "blocked": v["blocked"]} for d, v in sorted(chart_by_date.items())]
def _get_guardrail_field(g: Any, field: str) -> Any:
"""Read `field` off a guardrail whether it's a Prisma row (attr) or a dict/TypedDict (key)."""
if isinstance(g, dict):
return g.get(field)
return getattr(g, field, None)
def _to_dict(value: Any) -> Dict[str, Any]:
"""Coerce a pydantic model (e.g. LitellmParams) / dict value into a plain dict."""
if isinstance(value, BaseModel):
return value.model_dump(exclude_none=True)
if isinstance(value, dict):
return value
return {}
def _get_guardrail_attrs(g: Any) -> tuple[Any, str]:
"""Get (guardrail_id, display_name) from guardrail - handles Prisma model or dict."""
gid = getattr(g, "guardrail_id", None) or (g.get("guardrail_id") if isinstance(g, dict) else None)
name = getattr(g, "guardrail_name", None) or (g.get("guardrail_name") if isinstance(g, dict) else None)
gid = _get_guardrail_field(g, "guardrail_id")
name = _get_guardrail_field(g, "guardrail_name")
return gid, (name or gid or "")
@ -163,9 +179,9 @@ def _guardrail_overview_rows(
break
req, blocked = a["requests"], a["blocked"]
fail_rate = (100.0 * blocked / req) if req else 0.0
litellm_params = (g.litellm_params or {}) if isinstance(g.litellm_params, dict) else {}
litellm_params = _to_dict(_get_guardrail_field(g, "litellm_params"))
provider = str(litellm_params.get("guardrail", "Unknown"))
guardrail_info = (g.guardrail_info or {}) if isinstance(g.guardrail_info, dict) else {}
guardrail_info = _to_dict(_get_guardrail_field(g, "guardrail_info"))
gtype = str(guardrail_info.get("type", "Guardrail"))
prev_fail = 0.0
for k in lookup_keys:
@ -262,9 +278,15 @@ async def guardrails_usage_overview(
end = end_date or now.strftime("%Y-%m-%d")
start = start_date or (now - timedelta(days=7)).strftime("%Y-%m-%d")
from litellm.proxy.guardrails.guardrail_registry import IN_MEMORY_GUARDRAIL_HANDLER
try:
# Guardrails from DB
guardrails = await GuardrailsRepository(prisma_client).table.find_many()
db_guardrails = await GuardrailsRepository(prisma_client).table.find_many()
seen_ids = {gid for g in db_guardrails if (gid := _get_guardrail_field(g, "guardrail_id")) is not None}
config_guardrails = [
g for g in IN_MEMORY_GUARDRAIL_HANDLER.list_config_guardrails() if g.get("guardrail_id") not in seen_ids
]
guardrails: List[Any] = [*db_guardrails, *config_guardrails]
# Daily metrics in range
metrics = await DailyGuardrailMetricsRepository(prisma_client).table.find_many(
@ -321,16 +343,18 @@ async def guardrails_usage_detail(
end = end_date or now.strftime("%Y-%m-%d")
start = start_date or (now - timedelta(days=7)).strftime("%Y-%m-%d")
guardrail = await GuardrailsRepository(prisma_client).table.find_unique(where={"guardrail_id": guardrail_id})
if not guardrail:
from litellm.proxy.guardrails.guardrail_registry import IN_MEMORY_GUARDRAIL_HANDLER
guardrail: Any = await GuardrailsRepository(prisma_client).table.find_unique(where={"guardrail_id": guardrail_id})
if guardrail is None:
guardrail = IN_MEMORY_GUARDRAIL_HANDLER.get_config_guardrail_by_id(guardrail_id=guardrail_id)
if guardrail is None:
from fastapi import HTTPException
raise HTTPException(status_code=404, detail="Guardrail not found")
# Metrics are keyed by logical name (from spend log metadata), not UUID
logical_id = getattr(guardrail, "guardrail_name", None) or (
guardrail.get("guardrail_name") if isinstance(guardrail, dict) else None
)
logical_id = _get_guardrail_field(guardrail, "guardrail_name")
metric_ids = [i for i in (logical_id, guardrail_id) if i]
metrics = await DailyGuardrailMetricsRepository(prisma_client).table.find_many(
@ -367,17 +391,9 @@ async def guardrails_usage_detail(
{"date": d, "passed": v["passed"], "blocked": v["blocked"], "score": None}
for d, v in sorted(ts_by_date.items())
]
_litellm_params = getattr(guardrail, "litellm_params", None) or (
guardrail.get("litellm_params") if isinstance(guardrail, dict) else None
)
litellm_params = _litellm_params if isinstance(_litellm_params, dict) else {}
_guardrail_info = getattr(guardrail, "guardrail_info", None) or (
guardrail.get("guardrail_info") if isinstance(guardrail, dict) else None
)
guardrail_info = _guardrail_info if isinstance(_guardrail_info, dict) else {}
_guardrail_name = getattr(guardrail, "guardrail_name", None) or (
guardrail.get("guardrail_name") if isinstance(guardrail, dict) else None
)
litellm_params = _to_dict(_get_guardrail_field(guardrail, "litellm_params"))
guardrail_info = _to_dict(_get_guardrail_field(guardrail, "guardrail_info"))
_guardrail_name = _get_guardrail_field(guardrail, "guardrail_name")
return UsageDetailResponse(
guardrail_id=guardrail_id,
@ -548,11 +564,15 @@ async def guardrails_usage_logs(
# Query by both so we match regardless of which was written.
effective_guardrail_ids: List[str] = [guardrail_id] if guardrail_id else []
if guardrail_id:
guardrail = await GuardrailsRepository(prisma_client).table.find_unique(
from litellm.proxy.guardrails.guardrail_registry import IN_MEMORY_GUARDRAIL_HANDLER
guardrail: Any = await GuardrailsRepository(prisma_client).table.find_unique(
where={"guardrail_id": guardrail_id}
)
if guardrail is None:
guardrail = IN_MEMORY_GUARDRAIL_HANDLER.get_config_guardrail_by_id(guardrail_id=guardrail_id)
if guardrail:
logical_name = getattr(guardrail, "guardrail_name", None)
logical_name = _get_guardrail_field(guardrail, "guardrail_name")
if logical_name and logical_name not in effective_guardrail_ids:
effective_guardrail_ids.append(logical_name)

View file

@ -155,6 +155,40 @@ class SemanticToolFilterHook(CustomLogger):
return await self.filter.filter_tools(query=user_query, available_tools=expanded_tools)
def _selected_tool_names(self, filtered_tools: list[dict[str, Any]]) -> list[str]:
"""Names of the semantically selected tools, as produced by the MCP expansion."""
names = (self.filter._extract_tool_info(tool)[0] for tool in filtered_tools)
return [name for name in names if name]
@staticmethod
def _narrow_mcp_references(tools: list[Any], selected_tool_names: list[str]) -> list[Any]:
"""
Restrict each litellm_proxy MCP reference to the semantically selected tools.
The reference block is preserved rather than replaced with expanded tools, so the
MCP gateway still performs the expansion. That keeps the per-endpoint tool shape
and tool auto-execution intact. Expansion already applied any caller-supplied
allowed_tools, so this selection can only narrow a block further.
Whether an undecidable selection exposes every tool or none is owned by
SemanticMCPToolFilter.filter_tools, which returns the full set when nothing
matches; the same policy therefore governs references and plain tools. Passing an
empty selection through is safe rather than a hidden allow-all: the gateway reads
the union of every reference's allowed_tools and treats an empty union as unset.
"""
from litellm.responses.mcp.litellm_proxy_mcp_handler import (
LiteLLM_Proxy_MCP_Handler,
)
return [
(
{**tool, "allowed_tools": selected_tool_names}
if isinstance(tool, dict) and LiteLLM_Proxy_MCP_Handler._should_use_litellm_mcp_gateway([tool])
else tool
)
for tool in tools
]
def _is_mcp_tool(self, tool: object) -> bool:
"""
Check whether *tool* is registered in the MCP semantic router.
@ -261,36 +295,30 @@ class SemanticToolFilterHook(CustomLogger):
if self._should_expand_mcp_tools(tools):
verbose_proxy_logger.debug("Detected litellm_proxy MCP references, expanding before semantic filtering")
if not self.filter.enabled:
verbose_proxy_logger.debug("Semantic filter disabled, leaving MCP references untouched")
return None
try:
native_tools_before_expand = [t for t in tools if not (isinstance(t, dict) and t.get("type") == "mcp")]
expanded_tools = await self._expand_mcp_tools(tools, user_api_key_dict)
if not expanded_tools:
if native_tools_before_expand:
data["tools"] = native_tools_before_expand
verbose_proxy_logger.warning(
f"No MCP tools expanded, preserving {len(native_tools_before_expand)} native tools"
)
return data
verbose_proxy_logger.warning("No tools expanded from MCP references")
return None
if not self.filter.enabled:
data["tools"] = native_tools_before_expand + expanded_tools
verbose_proxy_logger.debug("Semantic filter disabled, forwarding expanded MCP tools unfiltered")
return data
filtered_expanded_tools = await self._filter_expanded_tools(data=data, expanded_tools=expanded_tools)
combined_tools = native_tools_before_expand + filtered_expanded_tools
data["tools"] = combined_tools
selected_tool_names = self._selected_tool_names(filtered_expanded_tools)
narrowed_tools = self._narrow_mcp_references(tools, selected_tool_names)
data["tools"] = narrowed_tools
self._emit_filter_metadata_safe(
data=data,
mcp_tools=expanded_tools,
filtered_mcp_tools=filtered_expanded_tools,
native_tools=native_tools_before_expand,
filtered_tools=combined_tools,
filtered_tools=narrowed_tools,
)
verbose_proxy_logger.info(
f"Expanded MCP references to {len(expanded_tools)} tools "

View file

@ -5,7 +5,7 @@ import litellm
from litellm._logging import verbose_proxy_logger
from litellm.caching.caching import DualCache
from litellm.integrations.custom_logger import Span
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy._types import Litellm_EntityType, UserAPIKeyAuth
from litellm.router_strategy.budget_limiter import RouterBudgetLimiting
from litellm.types.llms.openai import AllMessageValues
from litellm.types.utils import (
@ -76,6 +76,8 @@ class _PROXY_VirtualKeyModelMaxBudgetLimiter(RouterBudgetLimiting):
message=f"LiteLLM Virtual Key: {user_api_key_dict.token}, key_alias: {user_api_key_dict.key_alias}, exceeded budget for model={model}",
current_cost=_current_spend,
max_budget=_current_model_budget_info.max_budget,
entity_type=Litellm_EntityType.KEY.value,
entity_id=user_api_key_dict.token,
)
return True
@ -140,6 +142,8 @@ class _PROXY_VirtualKeyModelMaxBudgetLimiter(RouterBudgetLimiting):
message=f"LiteLLM End User: {end_user_id}, exceeded budget for model={model}",
current_cost=_current_spend,
max_budget=_current_model_budget_info.max_budget,
entity_type=Litellm_EntityType.END_USER.value,
entity_id=end_user_id,
)
return True

View file

@ -949,6 +949,10 @@ class LiteLLMProxyRequestSetup:
user_api_key_alias=user_api_key_dict.key_alias,
user_api_key_spend=user_api_key_dict.spend,
user_api_key_max_budget=user_api_key_dict.max_budget,
user_api_key_user_spend=user_api_key_dict.user_spend,
user_api_key_user_max_budget=user_api_key_dict.user_max_budget,
user_api_key_team_spend=user_api_key_dict.team_spend,
user_api_key_team_max_budget=user_api_key_dict.team_max_budget,
user_api_key_team_id=user_api_key_dict.team_id,
user_api_key_project_id=user_api_key_dict.project_id,
user_api_key_project_alias=user_api_key_dict.project_alias,

View file

@ -61,6 +61,8 @@ from litellm.types.proxy.management_endpoints.common_daily_activity import (
)
from litellm.types.proxy.management_endpoints.scim_v2 import (
SCIM_ENTERPRISE_METADATA_KEY,
SCIM_ENTITLEMENTS_METADATA_KEY,
SCIM_ROLES_METADATA_KEY,
)
from litellm.types.proxy.management_endpoints.internal_user_endpoints import (
BulkUpdateUserRequest,
@ -690,15 +692,21 @@ async def _get_user_info_teams(
return team_list, teams_1
_SCIM_DIRECTORY_METADATA_KEYS = frozenset(
{SCIM_ENTERPRISE_METADATA_KEY, SCIM_ENTITLEMENTS_METADATA_KEY, SCIM_ROLES_METADATA_KEY}
)
def _redact_scim_enterprise_metadata(
metadata: Optional[Dict[str, Any]],
) -> Optional[Dict[str, Any]]:
"""SCIM enterprise attributes are persisted in user metadata so reporting can
group on them, but they are directory-only fields that generic user-info
endpoints must not surface; SCIM clients read them through the SCIM endpoints."""
if not isinstance(metadata, dict) or SCIM_ENTERPRISE_METADATA_KEY not in metadata:
"""SCIM enterprise attributes, entitlements, and roles are persisted in user
metadata so reporting can group on them, but they are directory-only fields
that generic user-info endpoints must not surface; SCIM clients read them
through the SCIM endpoints."""
if not isinstance(metadata, dict) or not _SCIM_DIRECTORY_METADATA_KEYS.intersection(metadata):
return metadata
return {k: v for k, v in metadata.items() if k != SCIM_ENTERPRISE_METADATA_KEY}
return {k: v for k, v in metadata.items() if k not in _SCIM_DIRECTORY_METADATA_KEYS}
def _build_user_info_response(

View file

@ -536,6 +536,7 @@ if MCP_AVAILABLE:
sanitized.env = {}
sanitized.command = None
sanitized.args = []
sanitized.issuer = None
sanitized.authorization_url = None
sanitized.token_url = None
sanitized.registration_url = None
@ -581,6 +582,7 @@ if MCP_AVAILABLE:
sanitized.teams = []
sanitized.env_vars = None
sanitized.issuer = None
sanitized.authorization_url = None
sanitized.token_url = None
sanitized.registration_url = None
@ -686,6 +688,7 @@ if MCP_AVAILABLE:
command=payload.command,
args=payload.args,
env=payload.env,
issuer=payload.issuer,
authorization_url=payload.authorization_url,
token_url=payload.token_url,
registration_url=payload.registration_url,

View file

@ -1,5 +1,8 @@
from typing import List, Union
from typing import Callable, List, TypeVar, Union
from pydantic import ValidationError
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import (
LiteLLM_TeamTable,
LiteLLM_UserTable,
@ -9,6 +12,8 @@ from litellm.proxy._types import (
from litellm.repositories.team_repository import TeamRepository
from litellm.types.proxy.management_endpoints.scim_v2 import *
T = TypeVar("T")
class ScimTransformations:
DEFAULT_SCIM_NAME = "Unknown User"
@ -47,11 +52,19 @@ class ScimTransformations:
active = True if scim_active is None else bool(scim_active)
schemas = ["urn:ietf:params:scim:schemas:core:2.0:User"]
enterprise_user = None
if metadata.get(SCIM_ENTERPRISE_METADATA_KEY):
enterprise_user = SCIMEnterpriseUser.model_validate(metadata[SCIM_ENTERPRISE_METADATA_KEY])
enterprise_user = ScimTransformations._parse_directory_metadata(
user, SCIM_ENTERPRISE_METADATA_KEY, SCIMEnterpriseUser.model_validate
)
if enterprise_user is not None:
schemas.append(SCIM_ENTERPRISE_USER_SCHEMA)
entitlements = ScimTransformations._parse_directory_metadata(
user, SCIM_ENTITLEMENTS_METADATA_KEY, SCIM_MULTI_VALUED_LIST_ADAPTER.validate_python
)
roles = ScimTransformations._parse_directory_metadata(
user, SCIM_ROLES_METADATA_KEY, SCIM_MULTI_VALUED_LIST_ADAPTER.validate_python
)
return SCIMUser(
schemas=schemas,
id=user.user_id,
@ -64,6 +77,8 @@ class ScimTransformations:
emails=emails,
groups=groups,
active=active,
entitlements=entitlements,
roles=roles,
enterprise_user=enterprise_user,
meta={
"resourceType": "User",
@ -72,6 +87,31 @@ class ScimTransformations:
},
)
@staticmethod
def _parse_directory_metadata(
user: Union[LiteLLM_UserTable, NewUserResponse],
key: str,
validate: Callable[[object], T],
) -> T | None:
"""A SCIM directory attribute parsed from user metadata, or None when absent or malformed.
Metadata is writable outside the SCIM surface, so a malformed value on one user must not
fail the whole directory response; the attribute is omitted and the corruption logged.
"""
metadata = user.metadata or {}
raw = metadata.get(key)
if not raw:
return None
try:
return validate(raw)
except ValidationError:
verbose_proxy_logger.warning(
"Skipping malformed %s metadata on user %s in SCIM response",
key,
user.user_id,
)
return None
@staticmethod
def _get_scim_user_name(user: Union[LiteLLM_UserTable, NewUserResponse]) -> str:
"""

View file

@ -17,7 +17,7 @@ from fastapi import (
Request,
Response,
)
from pydantic import BaseModel
from pydantic import BaseModel, ValidationError
from typing_extensions import TypedDict
import litellm
@ -125,6 +125,8 @@ class ScimUserData(TypedDict):
family_name: Optional[str]
active: Optional[bool]
enterprise: Optional[SCIMEnterpriseUser]
entitlements: list[SCIMMultiValuedAttribute] | None
roles: list[SCIMMultiValuedAttribute] | None
class GroupMemberExtractionResult(BaseModel):
@ -199,6 +201,8 @@ def _extract_scim_user_data(user: SCIMUser) -> ScimUserData:
"family_name": user.name.familyName if user.name else None,
"active": user.active,
"enterprise": user.enterprise_user,
"entitlements": user.entitlements,
"roles": user.roles,
}
@ -207,6 +211,8 @@ def _build_scim_metadata(
family_name: Optional[str],
active: Optional[bool] = None,
enterprise: Optional[SCIMEnterpriseUser] = None,
entitlements: list[SCIMMultiValuedAttribute] | None = None,
roles: list[SCIMMultiValuedAttribute] | None = None,
) -> Dict[str, Any]:
"""Build metadata dictionary with SCIM data."""
metadata: Dict[str, Any] = {
@ -222,6 +228,12 @@ def _build_scim_metadata(
if enterprise is not None:
metadata[SCIM_ENTERPRISE_METADATA_KEY] = enterprise.model_dump(by_alias=True, exclude_none=True)
if entitlements is not None:
metadata[SCIM_ENTITLEMENTS_METADATA_KEY] = [e.model_dump(exclude_none=True) for e in entitlements]
if roles is not None:
metadata[SCIM_ROLES_METADATA_KEY] = [r.model_dump(exclude_none=True) for r in roles]
return metadata
@ -739,6 +751,62 @@ def _get_schemas() -> list:
),
],
),
SCIMSchemaAttribute(
name="entitlements",
type="complex",
multiValued=True,
description="A list of entitlements for the user.",
subAttributes=[
SCIMSchemaAttribute(
name="value",
type="string",
description="The value of an entitlement.",
),
SCIMSchemaAttribute(
name="display",
type="string",
description="A human-readable name for the entitlement.",
),
SCIMSchemaAttribute(
name="type",
type="string",
description="A label indicating the entitlement's function.",
),
SCIMSchemaAttribute(
name="primary",
type="boolean",
description="Whether this is the primary entitlement.",
),
],
),
SCIMSchemaAttribute(
name="roles",
type="complex",
multiValued=True,
description="A list of roles for the user.",
subAttributes=[
SCIMSchemaAttribute(
name="value",
type="string",
description="The value of a role.",
),
SCIMSchemaAttribute(
name="display",
type="string",
description="A human-readable name for the role.",
),
SCIMSchemaAttribute(
name="type",
type="string",
description="A label indicating the role's function.",
),
SCIMSchemaAttribute(
name="primary",
type="boolean",
description="Whether this is the primary role.",
),
],
),
],
meta={
"location": "/scim/v2/Schemas/urn:ietf:params:scim:schemas:core:2.0:User",
@ -1074,6 +1142,8 @@ async def create_user(
user_data["given_name"],
user_data["family_name"],
enterprise=user_data["enterprise"],
entitlements=user_data["entitlements"],
roles=user_data["roles"],
)
default_role = _default_scim_user_role()
@ -1152,6 +1222,8 @@ async def update_user(
user_data["family_name"],
scim_active_for_metadata,
enterprise=user_data["enterprise"],
entitlements=user_data["entitlements"],
roles=user_data["roles"],
)
await _handle_team_membership_changes(
@ -1311,6 +1383,48 @@ def _handle_group_operations(op_type: str, value: Any, teams_set: Set[str]) -> O
return None
def _multi_valued_attribute_base(path: str) -> str:
"""The attribute name a SCIM path targets, stripped of any value filter or sub-attribute."""
return path.split("[", 1)[0].split(".", 1)[0]
def _handle_multi_valued_attribute_update(path: str, op_type: str, value: Any, metadata: dict[str, Any]) -> None:
"""Handle add/replace/remove for the entitlements and roles multi-valued attributes."""
base = _multi_valued_attribute_base(path)
metadata_key = SCIM_MULTI_VALUED_ATTRIBUTE_METADATA_KEYS[base]
if path != base:
raise HTTPException(
status_code=400,
detail={"error": f"Filtered or sub-attribute paths are not supported for {base}; PATCH the full attribute"},
)
if op_type == "remove":
metadata.pop(metadata_key, None)
return
if value is None:
raise HTTPException(
status_code=400,
detail={"error": f"The {op_type} operation on {base} requires a 'value' member (RFC 7644 Section 3.5.2)"},
)
normalized = value if isinstance(value, list) else [value]
try:
attrs = SCIM_MULTI_VALUED_LIST_ADAPTER.validate_python(normalized)
except ValidationError:
raise HTTPException(
status_code=400,
detail={"error": f"Invalid value for {base}: expected a list of objects with a 'value' sub-attribute"},
)
dumped = [attr.model_dump(exclude_none=True) for attr in attrs]
existing = metadata.get(metadata_key)
if op_type == "add" and isinstance(existing, list):
metadata[metadata_key] = existing + dumped
return
metadata[metadata_key] = dumped
def _handle_generic_metadata(path: str, op_type: str, value: Any, metadata: Dict[str, Any]) -> None:
"""Handle generic metadata operations for unknown paths."""
if op_type == "remove":
@ -1346,6 +1460,8 @@ def _apply_patch_ops(
_handle_displayname_update(op_type, val, update_data)
elif key_lower == "externalid":
_handle_externalid_update(op_type, val, update_data)
elif key_lower in SCIM_MULTI_VALUED_ATTRIBUTE_METADATA_KEYS:
_handle_multi_valued_attribute_update(key_lower, op_type, val, metadata)
elif key_lower == "name" and isinstance(val, dict):
for name_key, name_val in val.items():
name_key_lower = name_key.lower()
@ -1366,6 +1482,8 @@ def _apply_patch_ops(
_handle_active_update(op_type, value, metadata)
elif path in ("name.givenname", "name.familyname"):
_handle_name_update(path, op_type, value, scim_metadata)
elif _multi_valued_attribute_base(path) in SCIM_MULTI_VALUED_ATTRIBUTE_METADATA_KEYS:
_handle_multi_valued_attribute_update(path, op_type, value, metadata)
elif path.startswith("groups"):
new_replace_set = _handle_group_operations(op_type, value, teams_set)
if new_replace_set is not None:

View file

@ -14,7 +14,7 @@ import json
import math
import traceback
from datetime import datetime, timezone
from typing import Annotated, Any, Dict, List, Optional, Sequence, Tuple, Union, cast
from typing import Annotated, Any, Dict, List, Mapping, Optional, Sequence, Tuple, Union, cast
import fastapi
from fastapi import APIRouter, Depends, Header, HTTPException, Request, status
@ -907,6 +907,17 @@ def _check_team_budget_update_authority(
)
def _should_auto_add_team_creator(
user_api_key_dict: UserAPIKeyAuth,
general_settings: Mapping[str, object],
) -> bool:
if user_api_key_dict.user_id is None:
return False
if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN:
return True
return general_settings.get("disable_auto_add_proxy_admin_to_teams") is not True
#### TEAM MANAGEMENT ####
@router.post(
"/team/new",
@ -1010,6 +1021,7 @@ async def new_team(
from litellm.proxy.proxy_server import (
_license_check,
create_audit_log_for_update,
general_settings,
litellm_proxy_admin_name,
prisma_client,
user_api_key_cache,
@ -1136,13 +1148,11 @@ async def new_team(
user_api_key_cache=user_api_key_cache,
)
if user_api_key_dict.user_id is not None:
creating_user_in_list = False
for member in data.members_with_roles:
if member.user_id == user_api_key_dict.user_id:
creating_user_in_list = True
if creating_user_in_list is False:
if _should_auto_add_team_creator(user_api_key_dict, general_settings):
creating_user_in_list = any(
member.user_id == user_api_key_dict.user_id for member in data.members_with_roles
)
if not creating_user_in_list:
data.members_with_roles.append(Member(role="admin", user_id=user_api_key_dict.user_id))
_check_passthrough_routes_caller_permission(data, user_api_key_dict, entity="team")

View file

@ -15,6 +15,7 @@ from dotenv import load_dotenv
import litellm
from litellm.constants import DEFAULT_NUM_WORKERS_LITELLM_PROXY
from litellm.proxy.db.query_engine_reaper import start_query_engine_reaper
from litellm.secret_managers.main import get_secret_bool
if TYPE_CHECKING:
@ -495,6 +496,7 @@ class ProxyInitializationHelpers:
gunicorn_options["certfile"] = ssl_certfile_path
gunicorn_options["keyfile"] = ssl_keyfile_path
start_query_engine_reaper()
StandaloneApplication(app=app, options=gunicorn_options).run() # Run gunicorn
@staticmethod
@ -1261,6 +1263,8 @@ def run_server(
if reload:
ProxyInitializationHelpers._configure_dev_reload(uvicorn_args, config)
if num_workers > 1:
start_query_engine_reaper()
uvicorn.run(
**uvicorn_args,
workers=num_workers,

View file

@ -2814,21 +2814,6 @@ async def update_cache(
)
# set cooldown on alert
if existing_spend_obj is not None and getattr(existing_spend_obj, "team_spend", None) is not None:
existing_team_spend = existing_spend_obj.team_spend or 0
# Calculate the new cost by adding the existing cost and response_cost
existing_spend_obj.team_spend = existing_team_spend + response_cost
if existing_spend_obj is not None and getattr(existing_spend_obj, "team_member_spend", None) is not None:
existing_team_member_spend = existing_spend_obj.team_member_spend or 0
# Calculate the new cost by adding the existing cost and response_cost
existing_spend_obj.team_member_spend = existing_team_member_spend + response_cost
# Existing spend_obj is mutated; UserApiKeyCache.async_set_cache_pipeline turns
# BaseModel values into dicts for Redis (same Codec path as async_set_cache).
existing_spend_obj.spend = new_spend
values_to_update_in_cache.append((hashed_token, existing_spend_obj))
### UPDATE USER SPEND ###
async def _update_user_cache():
## UPDATE CACHE FOR USER ID + GLOBAL PROXY
@ -3032,13 +3017,27 @@ async def update_cache(
if tags is not None:
await _update_tag_cache()
asyncio.create_task(
user_api_key_cache.async_set_cache_pipeline(
cache_list=values_to_update_in_cache,
ttl=get_management_object_ttl(user_api_key_cache),
litellm_parent_otel_span=parent_otel_span,
global_proxy_spend_key = "{}:spend".format(litellm_proxy_admin_name)
local_object_updates = tuple((k, v) for k, v in values_to_update_in_cache if k != global_proxy_spend_key)
shared_scalar_updates = tuple((k, v) for k, v in values_to_update_in_cache if k == global_proxy_spend_key)
if local_object_updates:
asyncio.create_task(
user_api_key_cache.async_set_cache_pipeline(
cache_list=list(local_object_updates),
ttl=get_management_object_ttl(user_api_key_cache),
litellm_parent_otel_span=parent_otel_span,
local_only=True,
)
)
if shared_scalar_updates:
asyncio.create_task(
user_api_key_cache.async_set_cache_pipeline(
cache_list=list(shared_scalar_updates),
ttl=get_management_object_ttl(user_api_key_cache),
litellm_parent_otel_span=parent_otel_span,
)
)
)
def run_ollama_serve():
@ -4425,6 +4424,15 @@ class ProxyConfig:
litellm.default_max_internal_user_budget = float(value)
if litellm.max_internal_user_budget is None:
litellm.max_internal_user_budget = litellm.default_max_internal_user_budget
elif key == "default_internal_user_params" and isinstance(value, dict):
litellm.default_internal_user_params = (
{**value, "max_budget": float(value["max_budget"])}
if value.get("max_budget") is not None
else value
)
verbose_proxy_logger.debug(
f"{blue_color_code} setting litellm.{key}={_redact_general_setting_value(key, litellm.default_internal_user_params, is_full_admin=False)}{reset_color_code}"
)
elif key == "custom_provider_map":
from litellm.utils import custom_llm_setup
@ -5785,6 +5793,13 @@ class ProxyConfig:
# For other types, convert to bool
general_settings["store_prompts_in_spend_logs"] = bool(value)
if "disable_auto_add_proxy_admin_to_teams" in _general_settings:
value = _general_settings["disable_auto_add_proxy_admin_to_teams"]
if isinstance(value, str):
general_settings["disable_auto_add_proxy_admin_to_teams"] = value.lower() == "true"
else:
general_settings["disable_auto_add_proxy_admin_to_teams"] = value if value is None else bool(value)
## STORE MODEL IN DB ##
if "store_model_in_db" in _general_settings:
value = _general_settings["store_model_in_db"]
@ -14898,6 +14913,7 @@ async def get_config_list(
"mcp_required_fields": {"type": "List"},
"cancel_on_disconnect": {"type": "Boolean"},
"skip_user_budget_on_team_key": {"type": "Boolean"},
"disable_auto_add_proxy_admin_to_teams": {"type": "Boolean"},
}
return_val = []

View file

@ -362,6 +362,8 @@ async def route_request(
for _key in _MOCK_TESTING_KWARG_NAMES:
data.pop(_key, None)
data.pop("enable_tag_filtering", None)
team_id = get_team_id_from_data(data)
router_model_names = llm_router.model_names if llm_router is not None else []
is_proxy_admin_without_team = team_id is None and _is_proxy_admin_request(data)
@ -409,6 +411,8 @@ async def route_request(
"num_retries",
"timeout",
"model_group_retry_policy",
"routing_strategy",
"enable_tag_filtering",
]
# Merge override settings into data (only if not already set in request)

View file

@ -325,6 +325,7 @@ model LiteLLM_MCPServerTable {
command String?
args String[] @default([])
env Json? @default("{}")
issuer String?
authorization_url String?
token_url String?
registration_url String?

View file

@ -4,7 +4,7 @@ import asyncio
import json
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
from typing import Any, Dict, List, Optional, Sequence, cast
from typing import Any, Dict, List, Mapping, Optional, Sequence, cast
import litellm
from litellm._logging import verbose_proxy_logger
@ -12,6 +12,7 @@ from litellm.caching import DualCache
from litellm.litellm_core_utils.duration_parser import duration_in_seconds
from litellm.litellm_core_utils.llm_cost_calc.tiered_pricing import select_tier_for_input, tier_rate
from litellm.proxy._types import (
Litellm_EntityType,
LiteLLM_TeamMembership,
LiteLLM_TeamTable,
LiteLLM_UserTable,
@ -36,6 +37,17 @@ class _BudgetCounter:
window_start: Optional[datetime] = None
_COUNTER_ENTITY_TYPES: Mapping[str, str] = {
"Key": Litellm_EntityType.KEY.value,
"Team": Litellm_EntityType.TEAM.value,
"TeamMember": Litellm_EntityType.TEAM_MEMBER.value,
"User": Litellm_EntityType.USER.value,
"EndUser": Litellm_EntityType.END_USER.value,
"Tag": Litellm_EntityType.TAG.value,
"Organization": Litellm_EntityType.ORGANIZATION.value,
}
class _CounterReservationUnavailable(Exception):
def __init__(
self,
@ -108,6 +120,8 @@ async def _apply_over_budget_reservation_policy(
f"Current cost: {current_spend}, "
f"Max budget: {counter.max_budget}"
),
entity_type=_COUNTER_ENTITY_TYPES.get(counter.entity_type),
entity_id=counter.spend_log_entity_id or counter.entity_id,
)

View file

@ -718,6 +718,12 @@ class ResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
except StopAsyncIteration:
# Normal end of stream - don't log as failure
raise
except (httpx.ReadError, httpx.RemoteProtocolError) as e:
self.finished = True
if self.completed_response is None:
self._handle_failure(e)
raise
raise StopAsyncIteration from e
except httpx.HTTPError as e:
# Handle HTTP errors
self.finished = True
@ -794,6 +800,12 @@ class SyncResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
except StopIteration:
# Normal end of stream - don't log as failure
raise
except (httpx.ReadError, httpx.RemoteProtocolError) as e:
self.finished = True
if self.completed_response is None:
self._handle_failure(e)
raise
raise StopIteration from e
except httpx.HTTPError as e:
# Handle HTTP errors
self.finished = True

View file

@ -251,6 +251,15 @@ else:
PreRoutingHookResponse = Any
def _cost_value_as_float(value: Union[str, int, float, None]) -> Optional[float]:
if value is None:
return None
try:
return float(value)
except (TypeError, ValueError):
return None
class RoutingArgs(enum.Enum):
ttl = 60 # 1min (RPM/TPM expire key)
@ -622,6 +631,8 @@ class Router:
routing_strategy_args=routing_strategy_args,
)
self._init_routing_groups(self._routing_groups_input)
self._override_selectors: dict[str, Any] = {}
self._override_selectors_lock = threading.Lock()
self.access_groups = None
## USAGE TRACKING ##
if isinstance(litellm._async_success_callback, list):
@ -893,7 +904,9 @@ class Router:
self._unregister_router_selectors(
[getattr(self, attr, None) for attr in self._DEFAULT_SELECTOR_ATTR_BY_STRATEGY.values()]
+ list(getattr(self, "_override_selectors", {}).values())
)
self._override_selectors = {}
self.leastbusy_logger: Optional[LeastBusyLoggingHandler] = None
self.lowesttpm_logger: Optional[LowestTPMLoggingHandler] = None
@ -983,12 +996,67 @@ class Router:
{strategy_value: group_selector} if group_selector is not None else {}
)
def _get_routing_context(self, model: str) -> Tuple[Optional[str], Optional[Any]]:
_OVERRIDABLE_ROUTING_STRATEGIES: frozenset[str] = frozenset({"simple-shuffle", *_DEFAULT_SELECTOR_ATTR_BY_STRATEGY})
def _get_request_routing_strategy_override(self, request_kwargs: Optional[dict]) -> Optional[str]:
"""
Reads a per-request `routing_strategy` override (forwarded by the proxy
from key/team `router_settings`) out of the request kwargs.
Only strategies with a per-request-capable selector are honored;
anything else (unknown strings, `lar1`, `provider-budget-routing`) is
ignored with a warning so a bad value stored on a key or team can
never take down that caller's traffic.
"""
if not request_kwargs:
return None
raw_strategy = request_kwargs.get("routing_strategy")
if raw_strategy is None:
return None
strategy = self._normalize_strategy(raw_strategy) if isinstance(raw_strategy, (str, RoutingStrategy)) else None
if not isinstance(strategy, str) or strategy not in self._OVERRIDABLE_ROUTING_STRATEGIES:
verbose_router_logger.warning(
"Ignoring per-request routing_strategy override '%s'; supported overrides: %s.",
raw_strategy,
sorted(self._OVERRIDABLE_ROUTING_STRATEGIES),
)
return None
return strategy
def _get_override_strategy_selector(self, strategy: str) -> Optional[Any]:
"""
Returns the selector for a per-request strategy override.
Reuses the default group's selector when the override matches the
router's configured strategy (so shared state keeps accumulating in
one place); otherwise lazily builds one selector per strategy and
caches it for the router's lifetime so its usage/latency state
persists across requests.
"""
if strategy == self._normalize_strategy(self.routing_strategy):
attr = self._DEFAULT_SELECTOR_ATTR_BY_STRATEGY.get(strategy)
return getattr(self, attr, None) if attr is not None else None
with self._override_selectors_lock:
if strategy not in self._override_selectors:
self._override_selectors[strategy] = self._build_strategy_selector(
strategy=strategy,
routing_strategy_args={},
)
return self._override_selectors[strategy]
def _get_routing_context(
self, model: str, request_kwargs: Optional[dict] = None
) -> tuple[Optional[str], Optional[Any]]:
"""
Resolves the routing strategy and selector to use for the given model.
Every model belongs to exactly one group: an explicit entry from
`routing_groups`, or the implicit `"default"` group driven by the
A per-request `routing_strategy` in `request_kwargs` (forwarded by the
proxy from key/team `router_settings`) takes precedence over both the
model's routing group and the router's top-level strategy, since it is
the most specific expression of caller intent.
Otherwise every model belongs to exactly one group: an explicit entry
from `routing_groups`, or the implicit `"default"` group driven by the
router's top-level `routing_strategy` / `routing_strategy_args`.
`self.routing_strategy` may be either a string or a `RoutingStrategy`
@ -996,6 +1064,11 @@ class Router:
string here. Downstream call sites and `_select_deployment_*` arms
compare against string literals.
"""
override = self._get_request_routing_strategy_override(request_kwargs)
if override is not None:
verbose_router_logger.debug("routing_group=request-override model=%s strategy=%s", model, override)
return override, self._get_override_strategy_selector(override)
group_name = self._model_to_group.get(model)
if group_name is None:
strategy = self._normalize_strategy(self.routing_strategy)
@ -5936,7 +6009,7 @@ class Router:
input_kwargs: dict,
) -> Optional[Any]:
"""Same-model-group retry after a failed deployment; returns None if not applicable."""
strategy, _ = self._get_routing_context(original_model_group)
strategy, _ = self._get_routing_context(original_model_group, kwargs)
if strategy != "simple-shuffle":
return None
@ -8750,8 +8823,8 @@ class Router:
# Get mode from database model_info if available, otherwise default to "chat"
db_model_info = model.get("model_info", {})
mode = db_model_info.get("mode", "chat")
input_cost_per_token = db_model_info.get("input_cost_per_token")
output_cost_per_token = db_model_info.get("output_cost_per_token")
input_cost_per_token = _cost_value_as_float(db_model_info.get("input_cost_per_token"))
output_cost_per_token = _cost_value_as_float(db_model_info.get("output_cost_per_token"))
model_info = ModelMapInfo(
key=model_group,
@ -8802,16 +8875,18 @@ class Router:
)
):
model_group_info.max_output_tokens = model_info["max_output_tokens"]
if model_info.get("input_cost_per_token", None) is not None and (
_input_cost_per_token = _cost_value_as_float(model_info.get("input_cost_per_token"))
if _input_cost_per_token is not None and (
model_group_info.input_cost_per_token is None
or (model_info["input_cost_per_token"] or 0.0) > (model_group_info.input_cost_per_token or 0.0)
or _input_cost_per_token > (model_group_info.input_cost_per_token or 0.0)
):
model_group_info.input_cost_per_token = model_info["input_cost_per_token"]
if model_info.get("output_cost_per_token", None) is not None and (
model_group_info.input_cost_per_token = _input_cost_per_token
_output_cost_per_token = _cost_value_as_float(model_info.get("output_cost_per_token"))
if _output_cost_per_token is not None and (
model_group_info.output_cost_per_token is None
or (model_info["output_cost_per_token"] or 0.0) > (model_group_info.output_cost_per_token or 0.0)
or _output_cost_per_token > (model_group_info.output_cost_per_token or 0.0)
):
model_group_info.output_cost_per_token = model_info["output_cost_per_token"]
model_group_info.output_cost_per_token = _output_cost_per_token
if (
model_info.get("supports_parallel_function_calling", None) is not None
and model_info["supports_parallel_function_calling"] is True # type: ignore
@ -9703,6 +9778,7 @@ class Router:
"retry_policy",
"model_group_alias",
"enable_weighted_failover",
"enable_tag_filtering",
]
for var in vars_to_include:
@ -9739,6 +9815,7 @@ class Router:
"model_group_retry_policy",
"model_group_alias",
"enable_weighted_failover",
"enable_tag_filtering",
]
_int_settings = [
@ -10462,7 +10539,7 @@ class Router:
# Resolve the strategy and logger AFTER the pre-routing hook, since
# the hook can replace `model` and routing-group lookup must key
# off the final model name.
strategy, strategy_selector = self._get_routing_context(model)
strategy, strategy_selector = self._get_routing_context(model, request_kwargs)
healthy_deployments = await self.async_get_healthy_deployments(
model=model,
@ -10606,7 +10683,7 @@ class Router:
# 5. Apply load balancing strategy
start_time = time.perf_counter()
strategy, strategy_selector = self._get_routing_context(model)
strategy, strategy_selector = self._get_routing_context(model, request_kwargs)
if strategy == "simple-shuffle":
return simple_shuffle(
llm_router_instance=self,
@ -10893,7 +10970,7 @@ class Router:
cooldown_list=_cooldown_list,
)
strategy, strategy_selector = self._get_routing_context(model)
strategy, strategy_selector = self._get_routing_context(model, request_kwargs)
if strategy == "simple-shuffle":
# if users pass rpm or tpm, we do a random weighted pick - based on rpm/tpm
############## Check 'weight' param set for weighted pick #################
@ -11033,7 +11110,7 @@ class Router:
)
# 6. Apply load balancing strategy
strategy, strategy_selector = self._get_routing_context(model)
strategy, strategy_selector = self._get_routing_context(model, request_kwargs)
if strategy == "simple-shuffle":
return simple_shuffle(
llm_router_instance=self,

View file

@ -160,8 +160,14 @@ async def get_deployments_for_tag(
Returns a list of deployments that match the requested model and tags in the request.
Executes tag based filtering based on the tags in request metadata and the tags on the deployments
Runs when the router-level `enable_tag_filtering` is True or the request carries
`enable_tag_filtering=True` (set from key/team router_settings by the proxy).
A request-level False never disables a router-level True, so per-request settings
cannot escape an operator's global tag-routing policy.
"""
if llm_router_instance.enable_tag_filtering is not True:
request_enable_tag_filtering = request_kwargs.get("enable_tag_filtering") if request_kwargs else None
if request_enable_tag_filtering is not True and llm_router_instance.enable_tag_filtering is not True:
return healthy_deployments
if request_kwargs is None:

View file

@ -800,6 +800,14 @@ class BaseLitellmParams(ContentFilterConfigModel): # works for new and patch up
"so only a valid guardrail response can block or modify it."
),
)
skip_unscannable_attachments: Optional[bool] = Field(
default=False,
description=(
"Implemented by guardrail='model_armor'. When True, attachment references that carry no "
"inline bytes (file_id, gs://, or http(s) URLs) pass through unscanned instead of blocking, "
"while fail_on_error still governs real Model Armor API errors. Default False blocks them."
),
)
additional_provider_specific_params: Optional[Dict[str, Any]] = Field(
default=None,

View file

@ -14,3 +14,4 @@ class UsagePerChunk(TypedDict):
web_search_requests: Optional[int]
completion_tokens_details: Optional[CompletionTokensDetails]
prompt_tokens_details: Optional[PromptTokensDetailsWrapper]
cost: Optional[float]

View file

@ -26,6 +26,11 @@ class MCPOAuthMetadata(BaseModel):
authorization_url: Optional[str] = None
token_url: Optional[str] = None
registration_url: Optional[str] = None
discovered_issuer: Optional[str] = None
"""The ``issuer`` the authorization-server metadata document self-attests (RFC 8414). Persisted
trust-on-first-use as the server's ``issuer`` when none is configured, so that later rebuilds
anchor discovery on it (RFC 8414 §3.3) and a subsequently compromised resource cannot re-point
it. Never overwrites an admin-configured issuer."""
from_origin_fallback: bool = False
"""True when the metadata came from guessing the resource origin as its authorization
server rather than from an RFC 9728/8414-advertised document. Guessed endpoints are
@ -60,6 +65,8 @@ class MCPServer(BaseModel):
# OAuth-specific fields
client_id: Optional[str] = None
client_secret: Optional[str] = None
issuer: Optional[str] = None
issuer_is_anchored: bool = False
scopes: Optional[List[str]] = None
authorization_url: Optional[str] = None
token_url: Optional[str] = None

View file

@ -6,13 +6,17 @@ from pydantic import (
ConfigDict,
EmailStr,
Field,
TypeAdapter,
field_validator,
model_serializer,
model_validator,
)
from pydantic_core.core_schema import SerializerFunctionWrapHandler
SCIM_ENTERPRISE_USER_SCHEMA = "urn:ietf:params:scim:schemas:extension:enterprise:2.0:User"
SCIM_ENTERPRISE_METADATA_KEY = "scim_enterprise"
SCIM_ENTITLEMENTS_METADATA_KEY = "scim_entitlements"
SCIM_ROLES_METADATA_KEY = "scim_roles"
class LiteLLM_UserScimMetadata(BaseModel):
@ -53,6 +57,28 @@ class SCIMUserGroup(BaseModel):
type: Optional[str] = "direct" # direct or indirect
class SCIMMultiValuedAttribute(BaseModel):
value: str
display: Optional[str] = None
type: Optional[str] = None
primary: Optional[bool] = None
@model_validator(mode="before")
@classmethod
def coerce_bare_string(cls, data: object) -> object:
if isinstance(data, str):
return {"value": data}
return data
SCIM_MULTI_VALUED_LIST_ADAPTER = TypeAdapter(List[SCIMMultiValuedAttribute])
SCIM_MULTI_VALUED_ATTRIBUTE_METADATA_KEYS = {
"entitlements": SCIM_ENTITLEMENTS_METADATA_KEY,
"roles": SCIM_ROLES_METADATA_KEY,
}
class SCIMUserManager(BaseModel):
model_config = ConfigDict(populate_by_name=True)
@ -81,6 +107,8 @@ class SCIMUser(SCIMResource):
active: bool = True
emails: Optional[List[SCIMUserEmail]] = None
groups: Optional[List[SCIMUserGroup]] = None
entitlements: Optional[List[SCIMMultiValuedAttribute]] = None
roles: Optional[List[SCIMMultiValuedAttribute]] = None
enterprise_user: Optional[SCIMEnterpriseUser] = Field(
default=None,
alias=SCIM_ENTERPRISE_USER_SCHEMA,
@ -88,11 +116,15 @@ class SCIMUser(SCIMResource):
)
@model_serializer(mode="wrap")
def _omit_absent_enterprise(self, handler: SerializerFunctionWrapHandler) -> Dict[str, Any]:
def _omit_absent_optional_blocks(self, handler: SerializerFunctionWrapHandler) -> Dict[str, Any]:
dumped = handler(self)
if self.enterprise_user is None:
dumped.pop(SCIM_ENTERPRISE_USER_SCHEMA, None)
dumped.pop("enterprise_user", None)
if self.entitlements is None:
dumped.pop("entitlements", None)
if self.roles is None:
dumped.pop("roles", None)
return dumped

View file

@ -117,6 +117,7 @@ class UpdateRouterConfig(BaseModel):
fallbacks: Optional[List[dict]] = None
context_window_fallbacks: Optional[List[dict]] = None
model_group_alias: Optional[Dict[str, Union[str, Dict]]] = {}
enable_tag_filtering: Optional[bool] = None
model_config = ConfigDict(protected_namespaces=())

View file

@ -1798,14 +1798,17 @@ class ModelResponseStream(ModelResponseBase):
else:
created = created
usage_to_set = None
if "usage" in kwargs and kwargs["usage"] is not None:
if isinstance(kwargs["usage"], dict):
kwargs["usage"] = Usage(**kwargs["usage"])
usage_to_set = Usage(**kwargs["usage"])
kwargs["usage"] = usage_to_set
elif isinstance(kwargs["usage"], BaseModel):
dump = (
kwargs["usage"].model_dump() if hasattr(kwargs["usage"], "model_dump") else kwargs["usage"].dict()
)
kwargs["usage"] = Usage(**dump)
usage_to_set = Usage(**dump)
kwargs["usage"] = usage_to_set
kwargs["id"] = id
kwargs["created"] = created
@ -1814,6 +1817,9 @@ class ModelResponseStream(ModelResponseBase):
super().__init__(**kwargs)
if usage_to_set is not None:
self.usage = usage_to_set
def __contains__(self, key):
# Define custom behavior for the 'in' operator
return hasattr(self, key)
@ -2469,6 +2475,10 @@ class StandardLoggingUserAPIKeyMetadata(TypedDict):
user_api_key_spend: Optional[float]
user_api_key_max_budget: Optional[float]
user_api_key_budget_reset_at: Optional[str]
user_api_key_user_spend: Optional[float]
user_api_key_user_max_budget: Optional[float]
user_api_key_team_spend: Optional[float]
user_api_key_team_max_budget: Optional[float]
user_api_key_org_id: Optional[str]
user_api_key_org_alias: Optional[str]
user_api_key_team_id: Optional[str]
@ -2687,6 +2697,10 @@ class StandardLoggingPayloadErrorInformation(TypedDict, total=False):
# hints). Lets dashboards split rate-limit failures by cause without
# parsing free-text error messages.
error_rate_limit_type: Optional[str]
error_budget_entity_type: Optional[str]
error_budget_entity_id: Optional[str]
error_budget_limit: Optional[float]
error_budget_spend: Optional[float]
class GuardrailMode(TypedDict, total=False):
@ -3136,6 +3150,7 @@ all_litellm_params = (
"use_client",
"id",
"fallbacks",
"routing_strategy",
"azure",
"headers",
"model_list",
@ -3207,6 +3222,7 @@ all_litellm_params = (
"shared_session",
"search_tool_name",
"order",
"enable_tag_filtering",
"enable_json_schema_validation",
"use_xai_oauth",
"_litellm_rate_limit_descriptors",

View file

@ -62,8 +62,8 @@ proxy = [
"azure-identity>=1.25.2,<2.0",
"azure-storage-blob>=12.28.0,<13.0",
"mcp>=1.26.0,<2.0",
"litellm-proxy-extras==0.4.77",
"litellm-enterprise==0.1.50",
"litellm-proxy-extras==0.4.78",
"litellm-enterprise==0.1.51",
"RestrictedPython>=8.1,<9.0",
"rich>=13.9.4,<14.0",
"InquirerPy>=0.3.4,<1.0",

View file

@ -325,6 +325,7 @@ model LiteLLM_MCPServerTable {
command String?
args String[] @default([])
env Json? @default("{}")
issuer String?
authorization_url String?
token_url String?
registration_url String?

View file

@ -5,12 +5,24 @@
# Needs only curl: uv is bootstrapped if missing, and uv provisions a compatible
# Python itself (reusing a suitable system one, else downloading a managed build).
#
# To install from an unreleased branch, tag, or commit instead of the latest PyPI
# release, set LITELLM_CLI_REF:
# curl -fsSL https://raw.githubusercontent.com/BerriAI/litellm/<branch>/scripts/install.sh | \
# LITELLM_CLI_REF=<branch> sh
#
# NOTE: set -e without pipefail for POSIX sh compatibility (dash on Ubuntu/Debian
# ignores the shebang when invoked as `sh` and does not support `pipefail`).
set -eu
# NOTE: before merging, this must stay as "litellm[proxy]" to install from PyPI.
LITELLM_PACKAGE="litellm[proxy]"
# LITELLM_CLI_REF opts into installing from a branch, tag, or commit instead (for
# example, to QA lite autoroute against an unreleased branch, which needs this proxy
# runtime, not the thin litellm[cli] install).
if [ -n "${LITELLM_CLI_REF:-}" ]; then
LITELLM_PACKAGE="litellm[proxy] @ git+https://github.com/BerriAI/litellm.git@${LITELLM_CLI_REF}"
else
LITELLM_PACKAGE="litellm[proxy]"
fi
UV_VERSION="0.10.9"
# ── colours ────────────────────────────────────────────────────────────────
@ -81,7 +93,11 @@ fi
# ── install ────────────────────────────────────────────────────────────────
echo ""
header "Installing litellm[proxy]…"
if [ -n "${LITELLM_CLI_REF:-}" ]; then
header "Installing litellm[proxy] from ${LITELLM_CLI_REF}…"
else
header "Installing litellm[proxy]…"
fi
echo ""
# --python-preference system: reuse a compatible system Python when present,

View file

@ -17,7 +17,7 @@ Each subdirectory under `tests/e2e/` is one suite, scoped to an endpoint family
- `security/` - secret handling and log-leak protection
- `router/` - routing and reliability behavior (fallbacks, cooldowns)
- `gateway/` - proxy configuration only (`litellm-config.yml`); no tests
- `claude_code/` - the Claude Code compatibility matrix: drives the real `claude` CLI (and HTTP probes) against a proxy for each feature x provider cell, reporting tagged-union outcomes via the `compat_result` fixture; ships its own driver/builder/publisher plus `_*_unit_tests/` trees, and does not use the shared transport harness
- `claude_code/` - the Claude Code compatibility matrix: drives the real `claude` CLI (and HTTP probes) against a proxy for each feature x provider cell, reporting tagged-union outcomes via the `compat_result` fixture; ships its own driver/builder/publisher and does not use the shared transport harness
## Lay the pattern down in a class
@ -53,7 +53,7 @@ Each suite provides its own `client` fixture (see `llm_translation/passthrough_c
Request and response bodies are typed pydantic models in `models.py`; only the fields a test reads are modelled, and nothing passes raw dicts. Outcomes come back as a `Result[R]` tagged union (`Success`, `NetworkError`, `UnauthorizedError`, `RateLimitedError`, `ValidationError`, `UnknownApiError`). Handle them with `match`, or call `unwrap(...)` when a non-success should fail the test. The skip-vs-fail split is deliberate: a test marked `e2e` skips when no proxy answers its liveness probe, but once a request reaches the proxy any wrong behavior is a hard failure, never a skip
Mark live tests with `@pytest.mark.e2e` (on the class or the module). Pure coverage of the harness itself carries no marker and runs regardless. Use `scoped_key` for a fresh all-models key that auto-deletes, `resources` when you need to create and tear down more than a key, and `unique_marker()` from `e2e_config` to keep prompts, tags, and customer ids from colliding across concurrent runs and the shared response cache
Mark live tests with `@pytest.mark.e2e` (on the class or the module). `tests/e2e/` is for live proxy suites only; do not put unit tests here. Use `scoped_key` for a fresh all-models key that auto-deletes, `resources` when you need to create and tear down more than a key, and `unique_marker()` from `e2e_config` to keep prompts, tags, and customer ids from colliding across concurrent runs and the shared response cache
## Typing
@ -171,3 +171,16 @@ other.<area>.<case>.<assertion>
e.g. other.auth.jwt.valid_token_allows
other.lifecycle.readiness.reports_db
```
## Hard Rules
- no monkeypatching, mock tests or unit tests of any kind. if a contributor asks you to write an end to end test, do NOT stage a unit test with it. if you find a product gap, call it out in the PR description
- use model management endpoints to create new models for a test. this could be in a conftest / inline for each test. ask the user what they want.
- do not overengineer a test, i need you to write readable, clean code of what would look like a natural user scenario
- when it comes to typing an input schema for an api endpoint, have it type X = A | B | C ... where X = exhaustive union of all supported input schemas and A, B, C typically are composed by a base type. types are only pretty for a api request / response body. make sure to compose types instead of repeating the same base attributes over and over again.
- use the docker-compose to your advantage and spin up a local proxy, make sure all tests pass. if a test fails due to an internally found issue, let users know to create a linear ticket for it.
- do not use xfail markers, tests should be written in a form that the end user expects it to pass

View file

@ -27,19 +27,20 @@ collecting this module as a test file.
from __future__ import annotations
import os
from typing import Any, Mapping, Sequence
from typing import Any, Callable, Mapping, Sequence
import pytest
from claude_code._env import require_proxy
from claude_code.cli_driver import (
ClaudeCLIError,
DriverResult,
failure_diagnostic,
run_claude_models_parallel,
)
PROXY_BASE_URL_ENV = "LITELLM_PROXY_BASE_URL"
PROXY_API_KEY_ENV = "LITELLM_PROXY_API_KEY"
ClaudeRunner = Callable[..., Mapping[str, DriverResult | ClaudeCLIError]]
# Floor on the number of `stream_event` records (with delta payloads)
# we expect to see when the proxy actually streams. With
@ -79,6 +80,8 @@ def run_basic_messaging_cell(
models: Sequence[str],
prompt: str,
verify_streaming: bool = False,
env: Mapping[str, str] | None = None,
runner: ClaudeRunner = run_claude_models_parallel,
) -> None:
"""Run the shared `basic_messaging_*` × <provider> cell body.
@ -99,28 +102,13 @@ def run_basic_messaging_cell(
streamed reply to a single ``assistant`` event in
``--print --output-format stream-json`` mode).
"""
base_url = os.environ.get(PROXY_BASE_URL_ENV)
api_key = os.environ.get(PROXY_API_KEY_ENV)
if not base_url or not api_key:
compat_result.set(
{
"status": "fail",
"error": (
f"missing required env: set {PROXY_BASE_URL_ENV} and "
f"{PROXY_API_KEY_ENV} to point at a running LiteLLM proxy"
),
}
)
pytest.fail(
f"{PROXY_BASE_URL_ENV} / {PROXY_API_KEY_ENV} not configured",
pytrace=False,
)
base_url, api_key = require_proxy(compat_result, env=env)
extra_args: Sequence[str] = (
("--include-partial-messages",) if verify_streaming else ()
)
outcomes = run_claude_models_parallel(
outcomes = runner(
models=models,
prompt=prompt,
base_url=base_url,

View file

@ -1,38 +0,0 @@
{
"schema_version": "1",
"generated_at": "2026-04-25T00:00:00Z",
"litellm_version": "v1.83.0-stable",
"claude_code_version": "2.1.120",
"providers": [
"anthropic",
"bedrock_invoke"
],
"features": [
{
"id": "basic_messaging_non_streaming",
"name": "Basic messaging (non-streaming)",
"providers": {
"anthropic": {
"status": "pass"
},
"bedrock_invoke": {
"status": "not_tested"
}
}
},
{
"id": "tool_use",
"name": "Tool use",
"providers": {
"anthropic": {
"status": "fail",
"error": "[claude-sonnet-4-6] tool call dropped"
},
"bedrock_invoke": {
"status": "not_applicable",
"reason": "tool use not yet wired up for Bedrock Invoke"
}
}
}
]
}

View file

@ -1,9 +0,0 @@
schema_version: "1"
providers:
- anthropic
- bedrock_invoke
features:
- id: basic_messaging_non_streaming
name: Basic messaging (non-streaming)
- id: tool_use
name: Tool use

View file

@ -1,41 +0,0 @@
{
"schema_version": "1",
"results": [
{
"feature_id": "basic_messaging_non_streaming",
"provider": "anthropic",
"nodeid": "tests/e2e/claude_code/basic_messaging_non_streaming/test_anthropic.py::test_basic_messaging_non_streaming_anthropic[claude-haiku-4-5]",
"result": {"status": "pass"}
},
{
"feature_id": "basic_messaging_non_streaming",
"provider": "anthropic",
"nodeid": "tests/e2e/claude_code/basic_messaging_non_streaming/test_anthropic.py::test_basic_messaging_non_streaming_anthropic[claude-sonnet-4-6]",
"result": {"status": "pass"}
},
{
"feature_id": "basic_messaging_non_streaming",
"provider": "anthropic",
"nodeid": "tests/e2e/claude_code/basic_messaging_non_streaming/test_anthropic.py::test_basic_messaging_non_streaming_anthropic[claude-opus-4-7]",
"result": {"status": "pass"}
},
{
"feature_id": "tool_use",
"provider": "anthropic",
"nodeid": "tests/e2e/claude_code/tool_use/test_anthropic.py::test_x[claude-haiku-4-5]",
"result": {"status": "pass"}
},
{
"feature_id": "tool_use",
"provider": "anthropic",
"nodeid": "tests/e2e/claude_code/tool_use/test_anthropic.py::test_x[claude-sonnet-4-6]",
"result": {"status": "fail", "error": "[claude-sonnet-4-6] tool call dropped"}
},
{
"feature_id": "tool_use",
"provider": "bedrock_invoke",
"nodeid": "tests/e2e/claude_code/tool_use/test_bedrock_invoke.py::test_x[claude-haiku-4-5]",
"result": {"status": "not_applicable", "reason": "tool use not yet wired up for Bedrock Invoke"}
}
]
}

View file

@ -1,479 +0,0 @@
"""Golden-file tests for the Matrix JSON Builder.
These tests fix the published JSON schema. The builder is a pure function
from (manifest, results, metadata) → matrix dict, so we feed it a fixture
input set and compare the produced dict to a checked-in expected output.
Any schema drift — intentional or accidental — surfaces as a diff in PR
review.
"""
from __future__ import annotations
import json
from pathlib import Path
import pytest
from claude_code.matrix_builder import (
ManifestError,
ResultsError,
build_from_paths,
build_matrix,
load_manifest,
load_results,
)
FIXTURES = Path(__file__).parent / "fixtures"
def test_build_matrix_matches_golden_file(tmp_path):
manifest = load_manifest(FIXTURES / "manifest.yaml")
results = load_results(FIXTURES / "results.json")
matrix = build_matrix(
manifest=manifest,
results=results,
litellm_version="v1.83.0-stable",
claude_code_version="2.1.120",
generated_at="2026-04-25T00:00:00Z",
)
expected = json.loads((FIXTURES / "expected_matrix.json").read_text())
assert matrix == expected
def test_build_matrix_pass_requires_all_models_pass():
"""Multiple results in one cell must all be pass for the cell to be pass."""
manifest = {
"schema_version": "1",
"providers": ["anthropic"],
"features": [{"id": "f", "name": "F"}],
}
results = [
{"feature_id": "f", "provider": "anthropic", "result": {"status": "pass"}},
{"feature_id": "f", "provider": "anthropic", "result": {"status": "pass"}},
{"feature_id": "f", "provider": "anthropic", "result": {"status": "pass"}},
]
matrix = build_matrix(
manifest=manifest,
results=results,
litellm_version="v",
claude_code_version="c",
generated_at="t",
)
assert matrix["features"][0]["providers"]["anthropic"] == {"status": "pass"}
def test_build_matrix_any_fail_makes_cell_fail():
manifest = {
"schema_version": "1",
"providers": ["anthropic"],
"features": [{"id": "f", "name": "F"}],
}
results = [
{"feature_id": "f", "provider": "anthropic", "result": {"status": "pass"}},
{
"feature_id": "f",
"provider": "anthropic",
"result": {"status": "fail", "error": "[claude-opus-4-7] timeout"},
},
{"feature_id": "f", "provider": "anthropic", "result": {"status": "pass"}},
]
matrix = build_matrix(
manifest=manifest,
results=results,
litellm_version="v",
claude_code_version="c",
generated_at="t",
)
cell = matrix["features"][0]["providers"]["anthropic"]
assert cell["status"] == "fail"
assert cell["error"] == "[claude-opus-4-7] timeout"
def test_build_matrix_joins_all_failure_errors_in_one_cell():
"""When multiple tiers fail for different reasons within the same cell,
every failure's error must appear in the published cell so triage
isn't reduced to a single tier's diagnostic.
"""
manifest = {
"schema_version": "1",
"providers": ["anthropic"],
"features": [{"id": "f", "name": "F"}],
}
results = [
{
"feature_id": "f",
"provider": "anthropic",
"result": {"status": "fail", "error": "[claude-haiku-4-5] 429"},
},
{"feature_id": "f", "provider": "anthropic", "result": {"status": "pass"}},
{
"feature_id": "f",
"provider": "anthropic",
"result": {"status": "fail", "error": "[claude-opus-4-7] timeout"},
},
]
matrix = build_matrix(
manifest=manifest,
results=results,
litellm_version="v",
claude_code_version="c",
generated_at="t",
)
cell = matrix["features"][0]["providers"]["anthropic"]
assert cell["status"] == "fail"
assert "[claude-haiku-4-5] 429" in cell["error"]
assert "[claude-opus-4-7] timeout" in cell["error"]
def test_build_matrix_mixed_pass_and_not_tested_surfaces_pass():
"""A `not_tested` row mixed with `pass` rows must not silently demote
the cell to `not_tested` — `not_tested` is "absent data", not a
negative signal. Otherwise a partial crash mid-test, or a test that
explicitly recorded "tier didn't run", would discard real passing
results from the published cell.
"""
manifest = {
"schema_version": "1",
"providers": ["anthropic"],
"features": [{"id": "f", "name": "F"}],
}
results = [
{"feature_id": "f", "provider": "anthropic", "result": {"status": "pass"}},
{
"feature_id": "f",
"provider": "anthropic",
"result": {"status": "not_tested"},
},
{"feature_id": "f", "provider": "anthropic", "result": {"status": "pass"}},
]
matrix = build_matrix(
manifest=manifest,
results=results,
litellm_version="v",
claude_code_version="c",
generated_at="t",
)
assert matrix["features"][0]["providers"]["anthropic"] == {"status": "pass"}
def test_build_matrix_all_not_tested_stays_not_tested():
"""A cell whose every row is `not_tested` (or empty) must remain
`not_tested` — the absent-data rule only drops `not_tested` rows
when there's other signal to surface.
"""
manifest = {
"schema_version": "1",
"providers": ["anthropic"],
"features": [{"id": "f", "name": "F"}],
}
results = [
{
"feature_id": "f",
"provider": "anthropic",
"result": {"status": "not_tested"},
},
{
"feature_id": "f",
"provider": "anthropic",
"result": {"status": "not_tested"},
},
]
matrix = build_matrix(
manifest=manifest,
results=results,
litellm_version="v",
claude_code_version="c",
generated_at="t",
)
assert matrix["features"][0]["providers"]["anthropic"] == {"status": "not_tested"}
def test_build_matrix_mixed_pass_and_not_applicable_surfaces_pass():
"""A `not_applicable` row mixed with `pass` rows must surface as
`pass`, not `not_applicable`. The published cell answers "does this
feature work on this provider?"; if any tier passes, the feature
works there. Discarding passing tiers because one tier is NA would
misrepresent the cell as unsupported.
"""
manifest = {
"schema_version": "1",
"providers": ["anthropic"],
"features": [{"id": "f", "name": "F"}],
}
results = [
{"feature_id": "f", "provider": "anthropic", "result": {"status": "pass"}},
{
"feature_id": "f",
"provider": "anthropic",
"result": {
"status": "not_applicable",
"reason": "haiku does not support extended thinking",
},
},
{"feature_id": "f", "provider": "anthropic", "result": {"status": "pass"}},
]
matrix = build_matrix(
manifest=manifest,
results=results,
litellm_version="v",
claude_code_version="c",
generated_at="t",
)
assert matrix["features"][0]["providers"]["anthropic"] == {"status": "pass"}
def test_build_matrix_all_not_applicable_stays_not_applicable():
"""When every observed row is `not_applicable`, the cell remains
`not_applicable` and the first row's reason carries through to the
published matrix.
"""
manifest = {
"schema_version": "1",
"providers": ["anthropic"],
"features": [{"id": "f", "name": "F"}],
}
results = [
{
"feature_id": "f",
"provider": "anthropic",
"result": {
"status": "not_applicable",
"reason": "feature unsupported on this provider",
},
},
{
"feature_id": "f",
"provider": "anthropic",
"result": {"status": "not_applicable", "reason": "ditto"},
},
]
matrix = build_matrix(
manifest=manifest,
results=results,
litellm_version="v",
claude_code_version="c",
generated_at="t",
)
assert matrix["features"][0]["providers"]["anthropic"] == {
"status": "not_applicable",
"reason": "feature unsupported on this provider",
}
def test_build_matrix_fills_not_tested_for_missing_cells():
manifest = {
"schema_version": "1",
"providers": ["anthropic", "azure"],
"features": [{"id": "f", "name": "F"}],
}
results = [
{"feature_id": "f", "provider": "anthropic", "result": {"status": "pass"}},
]
matrix = build_matrix(
manifest=manifest,
results=results,
litellm_version="v",
claude_code_version="c",
generated_at="t",
)
cells = matrix["features"][0]["providers"]
assert cells["anthropic"] == {"status": "pass"}
assert cells["azure"] == {"status": "not_tested"}
def test_build_matrix_preserves_provider_and_feature_order():
manifest = {
"schema_version": "1",
"providers": ["azure", "anthropic", "vertex_ai"],
"features": [
{"id": "z", "name": "Z"},
{"id": "a", "name": "A"},
],
}
matrix = build_matrix(
manifest=manifest,
results=[],
litellm_version="v",
claude_code_version="c",
generated_at="t",
)
assert matrix["providers"] == ["azure", "anthropic", "vertex_ai"]
assert [f["id"] for f in matrix["features"]] == ["z", "a"]
assert list(matrix["features"][0]["providers"].keys()) == [
"azure",
"anthropic",
"vertex_ai",
]
def test_build_matrix_emits_schema_version_one():
manifest = {
"schema_version": "1",
"providers": ["anthropic"],
"features": [{"id": "f", "name": "F"}],
}
matrix = build_matrix(
manifest=manifest,
results=[],
litellm_version="v",
claude_code_version="c",
generated_at="t",
)
assert matrix["schema_version"] == "1"
def test_load_manifest_rejects_wrong_schema_version(tmp_path):
bad = tmp_path / "manifest.yaml"
bad.write_text(
'schema_version: "2"\nproviders: [anthropic]\nfeatures:\n - id: f\n name: F\n'
)
with pytest.raises(ManifestError, match="schema_version"):
load_manifest(bad)
def test_load_manifest_rejects_empty_features(tmp_path):
bad = tmp_path / "manifest.yaml"
bad.write_text('schema_version: "1"\nproviders: [anthropic]\nfeatures: []\n')
with pytest.raises(ManifestError):
load_manifest(bad)
def test_load_results_rejects_missing_results_key(tmp_path):
bad = tmp_path / "results.json"
bad.write_text(json.dumps({"schema_version": "1"}))
with pytest.raises(ResultsError):
load_results(bad)
def test_build_matrix_6x5_grid_matches_published_sample():
"""Slice 5 acceptance: feeding the per-model results the full v0
row set produces reproduces the hand-authored 6x5 sample that the
docs page renders.
Inputs mirror the structure of `compat-results.json` after a real
run with the proxy configured for all five columns and all six
feature directories: every (feature, provider, model) cell yields a
`pass`. Anthropic announced Claude in Microsoft Foundry on
2025-11-18, so the Azure column is now exercised end-to-end like
the others rather than reporting `not_applicable`.
The aggregated matrix must equal the checked-in
`sample_compatibility-matrix.json` byte-for-byte (after JSON load),
so any future schema drift surfaces here in review.
"""
repo_root = Path(__file__).resolve().parents[1]
full_manifest = load_manifest(repo_root / "manifest.yaml")
# The v0 sample matrix is a frozen baseline: it covers exactly the
# six features the PRD shipped with, in their canonical order. The
# live manifest may carry additional rows (extensions added after
# v0 shipped), but the sample is derived only from the v0 slice so
# this test stays a meaningful regression gate for the v0 cell
# shape rather than chasing every new row added downstream.
v0_feature_ids = [
"basic_messaging_non_streaming",
"basic_messaging_streaming",
"tool_use",
"prompt_caching_5m",
"vision",
# Row 6 of the v0 PRD; originally shipped as `extended_thinking`.
# The id was renamed in-place to `thinking` to match Anthropic's
# current docs (which reserve "extended thinking" for the
# deprecated manual mode only). The row's *position* in v0 is
# the load-bearing invariant, not the id string.
"thinking",
]
v0_features = [
feature
for feature in full_manifest["features"]
if feature["id"] in v0_feature_ids
]
manifest = {**full_manifest, "features": v0_features}
feature_ids = [feature["id"] for feature in manifest["features"]]
providers = manifest["providers"]
models = ["claude-haiku-4-5", "claude-sonnet-4-6", "claude-opus-4-7"]
results = []
for feature_id in feature_ids:
for provider in providers:
for model in models:
results.append(
{
"feature_id": feature_id,
"provider": provider,
"nodeid": (
f"tests/e2e/claude_code/{feature_id}/test_{provider}.py"
f"::test[{model}]"
),
"result": {"status": "pass"},
}
)
matrix = build_matrix(
manifest=manifest,
results=results,
litellm_version="v1.83.0-stable",
claude_code_version="2.1.120",
generated_at="2026-04-25T00:00:00Z",
)
expected = json.loads((repo_root / "sample_compatibility-matrix.json").read_text())
assert matrix == expected
def test_build_matrix_1x5_grid_one_failing_model_breaks_cell():
"""If even one of three models fails on a provider, that cell is fail
and the error string carries the failing model id so the docs
tooltip can name the outlier."""
repo_root = Path(__file__).resolve().parents[1]
manifest = load_manifest(repo_root / "manifest.yaml")
results = [
{
"feature_id": "basic_messaging_non_streaming",
"provider": "bedrock_invoke",
"result": {"status": "pass"},
},
{
"feature_id": "basic_messaging_non_streaming",
"provider": "bedrock_invoke",
"result": {
"status": "fail",
"error": "[claude-opus-4-7-bedrock-invoke] claude CLI exited 1: throttled",
},
},
{
"feature_id": "basic_messaging_non_streaming",
"provider": "bedrock_invoke",
"result": {"status": "pass"},
},
]
matrix = build_matrix(
manifest=manifest,
results=results,
litellm_version="v",
claude_code_version="c",
generated_at="t",
)
cell = matrix["features"][0]["providers"]["bedrock_invoke"]
assert cell["status"] == "fail"
assert "claude-opus-4-7-bedrock-invoke" in cell["error"]
def test_build_from_paths_writes_output(tmp_path):
out = tmp_path / "compatibility-matrix.json"
matrix = build_from_paths(
manifest_path=FIXTURES / "manifest.yaml",
results_path=FIXTURES / "results.json",
litellm_version="v1.83.0-stable",
claude_code_version="2.1.120",
generated_at="2026-04-25T00:00:00Z",
output_path=out,
)
assert out.exists()
on_disk = json.loads(out.read_text())
assert on_disk == matrix
expected = json.loads((FIXTURES / "expected_matrix.json").read_text())
assert on_disk == expected

View file

@ -1,183 +0,0 @@
"""Structural tests for the full v0 6x5 matrix layout.
These tests don't run the `claude` CLI — they only verify that the
shape of the test suite on disk matches what the PRD declares: six
features in the prescribed order, and for each feature a directory
with one test file per provider column.
Catching layout drift here means the daily-cron VM and the PR gate
both see the same row set the docs page declares.
"""
from __future__ import annotations
from pathlib import Path
import pytest
import yaml
SUITE_ROOT = Path(__file__).resolve().parents[1]
MANIFEST_PATH = SUITE_ROOT / "manifest.yaml"
# The PRD's "Features in v0" section, in row order.
EXPECTED_FEATURE_IDS = [
"basic_messaging_non_streaming",
"basic_messaging_streaming",
"tool_use",
"prompt_caching_5m",
"vision",
# v0 originally shipped this row as `extended_thinking`. It was
# renamed in-place to `thinking` because Anthropic's docs reserve
# "extended thinking" for the deprecated manual API mode only; the
# single row exercises both manual and adaptive shapes since Claude
# Code picks per model. The PRD's "v0" identity is the *position*
# (row 6, 0-indexed 5), not the id string.
"thinking",
]
# The PRD's column order. Every feature directory must have one
# `test_<provider>.py` for each of these.
EXPECTED_PROVIDERS = [
"anthropic",
"bedrock_invoke",
"bedrock_converse",
"vertex_ai",
"azure",
]
def _all_manifest_feature_ids() -> list[str]:
"""Every feature_id currently declared in `manifest.yaml`.
Evaluated at import time so the result can drive parametrized
structural tests below. Used to catch layout drift on post-v0
feature rows added after the matrix shipped — the v0 anchor
constants above only validate the original six rows by design.
"""
return [
feature["id"]
for feature in yaml.safe_load(MANIFEST_PATH.read_text())["features"]
]
ALL_FEATURE_IDS = _all_manifest_feature_ids()
@pytest.fixture(scope="module")
def manifest() -> dict:
return yaml.safe_load(MANIFEST_PATH.read_text())
def test_manifest_lists_all_six_v0_features_in_order(manifest):
"""The PRD's v0 row set must appear at the top of the manifest in
order. Features beyond v0 (extensions added after the matrix
shipped) are allowed but must not reorder or displace the v0
rows — the docs page anchors row links by index, so v0 stays
pinned at positions [0:6] for the lifetime of the schema.
"""
ids = [feature["id"] for feature in manifest["features"]]
assert ids[: len(EXPECTED_FEATURE_IDS)] == EXPECTED_FEATURE_IDS
def test_manifest_lists_all_five_v0_providers_in_order(manifest):
assert manifest["providers"] == EXPECTED_PROVIDERS
def test_manifest_every_feature_has_human_readable_name(manifest):
for feature in manifest["features"]:
assert isinstance(feature["name"], str) and feature["name"].strip()
@pytest.mark.parametrize("feature_id", EXPECTED_FEATURE_IDS)
def test_feature_directory_exists(feature_id):
feature_dir = SUITE_ROOT / feature_id
assert feature_dir.is_dir(), f"missing feature directory: {feature_dir}"
@pytest.mark.parametrize("feature_id", EXPECTED_FEATURE_IDS)
@pytest.mark.parametrize("provider", EXPECTED_PROVIDERS)
def test_per_provider_test_file_exists(feature_id, provider):
test_file = SUITE_ROOT / feature_id / f"test_{provider}.py"
assert test_file.is_file(), f"missing per-provider test file: {test_file}"
@pytest.mark.parametrize("feature_id", EXPECTED_FEATURE_IDS)
def test_feature_directory_has_init_file(feature_id):
"""Each feature directory needs an __init__.py so pytest collects
the per-provider test files as a package — matches the layout
established by `basic_messaging_non_streaming/`."""
init_file = SUITE_ROOT / feature_id / "__init__.py"
assert init_file.is_file(), f"missing __init__.py: {init_file}"
# Manifest-driven structural tests: every feature in `manifest.yaml`
# (v0 and post-v0 alike) must have the expected on-disk layout. The
# v0-only tests above pin the position of the original six rows; these
# extend the same structural guarantees to any row added afterward so
# a broken post-v0 directory still fails CI.
@pytest.mark.parametrize("feature_id", ALL_FEATURE_IDS)
def test_every_manifest_feature_has_directory(feature_id):
feature_dir = SUITE_ROOT / feature_id
assert feature_dir.is_dir(), (
f"manifest declares {feature_id!r} but {feature_dir} is missing — "
"feature_id MUST match its on-disk directory (see manifest.yaml header)."
)
@pytest.mark.parametrize("feature_id", ALL_FEATURE_IDS)
def test_every_manifest_feature_has_init_file(feature_id):
init_file = SUITE_ROOT / feature_id / "__init__.py"
assert init_file.is_file(), f"missing __init__.py: {init_file}"
@pytest.mark.parametrize("feature_id", ALL_FEATURE_IDS)
@pytest.mark.parametrize("provider", EXPECTED_PROVIDERS)
def test_every_manifest_feature_has_per_provider_test_file(feature_id, provider):
"""Every (feature, provider) cell in the rendered matrix must be
backed by a per-provider test file. Without this check, a missing
file silently becomes a `not_tested` cell in the published matrix
rather than a CI failure surfacing the layout drift."""
test_file = SUITE_ROOT / feature_id / f"test_{provider}.py"
assert test_file.is_file(), f"missing per-provider test file: {test_file}"
@pytest.mark.parametrize("feature_id", EXPECTED_FEATURE_IDS)
@pytest.mark.parametrize("provider", EXPECTED_PROVIDERS)
def test_per_provider_test_file_imports_and_parametrizes_three_models(
feature_id, provider
):
"""Every test file must reference the three Claude tiers required
by the PRD: Haiku 4.5, Sonnet 4.6, Opus 4.7. Implementations may
use plain aliases or per-provider-suffixed aliases (e.g.
`claude-opus-4-7-bedrock-invoke`), so we check for the tier
substrings rather than exact alias names."""
text = (SUITE_ROOT / feature_id / f"test_{provider}.py").read_text()
for tier in ("haiku-4-5", "sonnet-4-6", "opus-4-7"):
assert (
tier in text
), f"{feature_id}/test_{provider}.py does not reference {tier}"
@pytest.mark.parametrize("feature_id", EXPECTED_FEATURE_IDS)
def test_azure_test_file_drives_the_proxy(feature_id):
"""Azure (Microsoft Foundry) hosts Anthropic Claude as of 2025-11-18,
so every Azure cell in the v0 matrix exercises a real route through
the LiteLLM proxy — same shape as the other provider columns. Pin
that here so a future regression doesn't silently revert these
cells to the old `not_applicable` boilerplate.
We accept either the direct `run_claude(...)` family of entrypoints
or a per-feature shared helper (e.g. `run_basic_messaging_cell`)
that wraps them — both shapes drive the proxy, and we don't want
this layout pin to block legitimate de-duplication of test bodies.
"""
text = (SUITE_ROOT / feature_id / "test_azure.py").read_text()
assert "run_claude" in text or "run_basic_messaging_cell" in text, (
f"{feature_id}/test_azure.py must drive the claude CLI via run_claude() "
"or a shared helper that wraps it; the not_applicable stub was removed "
"when Foundry started hosting Claude."
)
assert '"status": "not_applicable"' not in text, (
f"{feature_id}/test_azure.py still reports not_applicable; Microsoft Foundry "
"now hosts Claude (Haiku 4.5, Sonnet 4.6, Opus 4.7), so this row must run."
)

View file

@ -0,0 +1,86 @@
"""Load the claude_code compat matrix's deployment list from
``test_config.yaml``.
``test_config.yaml`` is the ground-truth config the stage deployment
uses; parsing it at fixture time means a change there (new tier, tier
retirement, provider swap, endpoint rename) reaches the fixture with
no extra edit. A drift-check test asserts every ``*_MODELS`` list
referenced by the compat cells is covered by the yaml, so a cell that
adds a probe for a name the yaml doesn't know about fails loudly at
collection instead of at 400-time.
"""
from __future__ import annotations
from dataclasses import dataclass
from pathlib import Path
from typing import Callable, Mapping
import yaml
from models import LiteLLMParamsBody
CONFIG_PATH = Path(__file__).resolve().parent / "test_config.yaml"
@dataclass(frozen=True, slots=True)
class CompatDeployment:
model_name: str
litellm_params: LiteLLMParamsBody
# The yaml uses ``vertex_ai_*`` for the vertex project/location fields
# (that is the spelling the proxy config file historically standardized
# on), while ``LiteLLMParamsBody`` names them without the ``_ai`` infix
# (matching the proxy's DB column). Both spellings resolve at call time
# on the proxy side, but pydantic silently drops unknown fields, so a
# raw ``LiteLLMParamsBody(**entry)`` would produce a body with the
# vertex project stripped - the resulting deployment 400s at
# ``/v1/messages`` with "Invalid model name". Normalize the yaml keys
# to the pydantic names in one place.
_YAML_TO_PYDANTIC_ALIASES = {
"vertex_ai_project": "vertex_project",
"vertex_ai_location": "vertex_location",
"vertex_ai_credentials": "vertex_credentials",
}
def _normalize_params(raw: Mapping[str, object]) -> dict[str, object]:
return {_YAML_TO_PYDANTIC_ALIASES.get(k, k): v for k, v in raw.items()}
ConfigReader = Callable[[Path], str]
def _default_reader(path: Path) -> str:
return path.read_text()
def load_all_deployments(
config_path: Path = CONFIG_PATH,
reader: ConfigReader = _default_reader,
) -> tuple[CompatDeployment, ...]:
"""Every deployment declared in the yaml, in file order."""
doc = yaml.safe_load(reader(config_path))
model_list = doc.get("model_list") or []
return tuple(
CompatDeployment(
model_name=entry["model_name"],
litellm_params=LiteLLMParamsBody(
**_normalize_params(entry["litellm_params"])
),
)
for entry in model_list
)
def all_expected_model_names(
*,
config_path: Path = CONFIG_PATH,
reader: ConfigReader = _default_reader,
) -> frozenset[str]:
"""Every virtual name the compat matrix declares - the ground truth
the cells are supposed to probe. Used by the drift-check test."""
return frozenset(
d.model_name for d in load_all_deployments(config_path, reader)
)

View file

@ -1,32 +0,0 @@
"""Local conftest for the driver unit tests.
Installs a hermetic, no-op rate limiter for every test in this
subdirectory. Without this, importing `cli_driver` and calling
`run_claude(..., runner=fake)` would silently consume tokens from the
shared default limiter (which writes to `$TMPDIR/...`), polluting the
on-disk state another test run might rely on and adding flakiness if
the env vars say "rate=0.1/s".
A no-op limiter (rate=0 for every provider) returns immediately from
`acquire(...)`, so unit tests behave exactly as they did before the
limiter was added.
"""
from __future__ import annotations
import pytest
from claude_code.rate_limiter import (
ALL_PROVIDERS,
ProviderConfig,
RateLimiter,
use_limiter,
)
@pytest.fixture(autouse=True)
def _hermetic_rate_limiter(tmp_path):
config = {p: ProviderConfig(rate_per_sec=0.0, burst=0.0) for p in ALL_PROVIDERS}
limiter = RateLimiter(config=config, state_dir=tmp_path)
with use_limiter(limiter):
yield

View file

@ -1,201 +0,0 @@
"""Unit tests for the shared `run_basic_messaging_cell` helper.
These tests mock `run_claude_models_parallel` so they exercise the
helper's branching (env-missing guard, per-model pass/fail/empty-text,
streaming wire check) without spawning the real CLI. The streaming
check is the regression we care about: a proxy that buffers the
upstream stream must turn the cell red, not green.
"""
from __future__ import annotations
import os
from typing import Any, Dict, List, Mapping, Optional, Sequence
import pytest
from claude_code import _basic_messaging
from claude_code._basic_messaging import (
MIN_STREAM_DELTA_EVENTS,
_count_stream_event_deltas,
run_basic_messaging_cell,
)
from claude_code.cli_driver import DriverResult
class _FakeResult:
"""Stand-in for the test's `compat_result` fixture.
Records every `set` / `add` payload so assertions can inspect what
the cell reported, in order, without needing the real
`pytest_runtest_logreport` plumbing from `conftest.py`.
"""
def __init__(self) -> None:
self.rows: List[Dict[str, Any]] = []
self.single: Optional[Dict[str, Any]] = None
def set(self, payload: Mapping[str, Any]) -> None:
self.single = dict(payload)
def add(self, payload: Mapping[str, Any]) -> None:
self.rows.append(dict(payload))
def _streamed_events(n_deltas: int = 5) -> List[Dict[str, Any]]:
"""Build a stream-json event list that *looks* streamed.
Includes `n_deltas` `stream_event` records (matching what
`--include-partial-messages` produces) plus the usual
`system`/`assistant`/`result` boilerplate the CLI always emits.
"""
events: List[Dict[str, Any]] = [{"type": "system", "subtype": "init"}]
for i in range(n_deltas):
events.append(
{
"type": "stream_event",
"event": {
"type": "content_block_delta",
"delta": {"type": "text_delta", "text": str(i)},
},
}
)
events.append(
{
"type": "assistant",
"message": {"content": [{"type": "text", "text": "1\n2\n3"}]},
}
)
events.append({"type": "result"})
return events
def _buffered_events() -> List[Dict[str, Any]]:
"""Event list a buffering proxy would produce: zero `stream_event`s."""
return [
{"type": "system", "subtype": "init"},
{
"type": "assistant",
"message": {"content": [{"type": "text", "text": "1\n2\n3"}]},
},
{"type": "result"},
]
def _install_fake_runner(monkeypatch, *, outcomes_by_model):
"""Patch `run_claude_models_parallel` to return canned outcomes.
Captures the kwargs the cell passed in so tests can assert on
`extra_args` (which is how the streaming variant opts into
`--include-partial-messages`).
"""
captured: Dict[str, Any] = {}
def fake(*, models, prompt, base_url, api_key, extra_args=None, **_kwargs):
captured["models"] = list(models)
captured["prompt"] = prompt
captured["base_url"] = base_url
captured["api_key"] = api_key
captured["extra_args"] = list(extra_args) if extra_args else []
return {model: outcomes_by_model[model] for model in models}
monkeypatch.setattr(_basic_messaging, "run_claude_models_parallel", fake)
return captured
@pytest.fixture(autouse=True)
def _proxy_env(monkeypatch):
monkeypatch.setenv("LITELLM_PROXY_BASE_URL", "http://localhost:4000")
monkeypatch.setenv("LITELLM_PROXY_API_KEY", "sk-test")
def test_count_stream_event_deltas_only_counts_records_with_event_payload():
events = [
{"type": "system"},
{"type": "stream_event", "event": {"type": "message_start"}},
{"type": "stream_event", "event": {"type": "content_block_delta"}},
{"type": "stream_event"},
{"type": "stream_event", "event": None},
{"type": "stream_event", "event": "not-a-dict"},
{"type": "assistant"},
{"type": "result"},
]
assert _count_stream_event_deltas(events) == 2
def test_verify_streaming_passes_when_proxy_streams(monkeypatch):
fake_result = _FakeResult()
model = "claude-haiku-4-5"
outcome = DriverResult(text="1\n2\n3", events=_streamed_events(n_deltas=5))
captured = _install_fake_runner(monkeypatch, outcomes_by_model={model: outcome})
run_basic_messaging_cell(
compat_result=fake_result,
models=[model],
prompt="Count from 1 to 5, one number per line.",
verify_streaming=True,
)
assert captured["extra_args"] == ["--include-partial-messages"]
assert fake_result.rows == [{"status": "pass"}]
def test_verify_streaming_fails_when_proxy_buffers(monkeypatch):
fake_result = _FakeResult()
model = "claude-haiku-4-5"
outcome = DriverResult(text="1\n2\n3", events=_buffered_events())
_install_fake_runner(monkeypatch, outcomes_by_model={model: outcome})
with pytest.raises(pytest.fail.Exception):
run_basic_messaging_cell(
compat_result=fake_result,
models=[model],
prompt="Count from 1 to 5, one number per line.",
verify_streaming=True,
)
assert len(fake_result.rows) == 1
row = fake_result.rows[0]
assert row["status"] == "fail"
assert "stream_event" in row["error"]
assert f"< {MIN_STREAM_DELTA_EVENTS}" in row["error"]
def test_non_streaming_variant_omits_partial_messages_flag(monkeypatch):
"""Default `verify_streaming=False` keeps the non-streaming wire identical."""
fake_result = _FakeResult()
model = "claude-haiku-4-5"
outcome = DriverResult(text="pong", events=_buffered_events())
captured = _install_fake_runner(monkeypatch, outcomes_by_model={model: outcome})
run_basic_messaging_cell(
compat_result=fake_result,
models=[model],
prompt="Reply with the single word 'pong' and nothing else.",
)
assert captured["extra_args"] == []
assert fake_result.rows == [{"status": "pass"}]
def test_verify_streaming_requires_all_models_to_stream(monkeypatch):
"""If any one tier buffers, the cell fails — same all-must-pass shape as
the non-streaming check."""
fake_result = _FakeResult()
outcomes = {
"claude-haiku-4-5": DriverResult(text="ok", events=_streamed_events(5)),
"claude-sonnet-4-6": DriverResult(text="ok", events=_buffered_events()),
"claude-opus-4-7": DriverResult(text="ok", events=_streamed_events(5)),
}
_install_fake_runner(monkeypatch, outcomes_by_model=outcomes)
with pytest.raises(pytest.fail.Exception):
run_basic_messaging_cell(
compat_result=fake_result,
models=list(outcomes.keys()),
prompt="Count from 1 to 5, one number per line.",
verify_streaming=True,
)
statuses = [row["status"] for row in fake_result.rows]
assert statuses == ["pass", "fail", "pass"]

View file

@ -1,992 +0,0 @@
"""Unit tests for the Claude Code CLI Driver.
These tests mock the subprocess so they run anywhere — no network, no
`claude` install, no API keys. They cover the behavior contract:
argument assembly, environment overlay, stream-JSON parsing, exit-code
plumbing, and the structured failure modes (CLI not found, timeout).
"""
from __future__ import annotations
import json
import subprocess
from dataclasses import dataclass
from typing import List, Optional
import pytest
from claude_code.cli_driver import (
ClaudeCLIError,
DriverResult,
failure_diagnostic,
is_rate_limit_shaped,
run_claude,
run_claude_models_parallel,
)
@dataclass
class _Completed:
returncode: int = 0
stdout: str = ""
stderr: str = ""
def _make_runner(*, stdout: str = "", returncode: int = 0, stderr: str = ""):
captured = {}
def runner(cmd, env, capture_output, text, timeout, check, input=None):
captured["cmd"] = cmd
captured["env"] = env
captured["timeout"] = timeout
captured["input"] = input
return _Completed(returncode=returncode, stdout=stdout, stderr=stderr)
return runner, captured
def test_run_claude_assembles_command_correctly():
runner, captured = _make_runner(
stdout='{"type":"assistant","message":{"content":[{"type":"text","text":"ok"}]}}\n'
)
run_claude(
prompt="hello",
model="claude-haiku-4-5",
base_url="http://localhost:4000",
api_key="sk-test",
runner=runner,
)
cmd = captured["cmd"]
assert cmd[0] == "claude"
assert "--print" in cmd
assert "--output-format" in cmd
assert "stream-json" in cmd
assert "--model" in cmd
assert "claude-haiku-4-5" in cmd
# prompt is the last positional after the `--` end-of-options marker.
assert cmd[-2:] == ["--", "hello"]
def test_run_claude_places_extra_args_before_prompt():
"""`claude --print` expects the prompt as the final positional arg.
Flags appearing after the prompt are ignored or eaten by the prompt
parser (especially variadic flags like `--allowed-tools <tools...>`),
which silently broke the tool_use, vision, and web_search cells
before the fix. Pin the ordering: every flag (including
caller-supplied `extra_args`) must precede the `--` end-of-options
marker, which itself precedes the prompt.
"""
runner, captured = _make_runner(stdout="")
run_claude(
prompt="say hi",
model="claude-haiku-4-5",
base_url="http://localhost:4000",
api_key="sk-test",
extra_args=["--allowed-tools", "Bash"],
runner=runner,
)
cmd = captured["cmd"]
# Prompt is last, `--` immediately precedes it, and the caller's
# extra_args sit somewhere earlier in the command.
assert cmd[-2:] == ["--", "say hi"]
assert "--allowed-tools" in cmd
assert cmd.index("--allowed-tools") < cmd.index("--")
def test_run_claude_overlays_proxy_env():
runner, captured = _make_runner(stdout="")
run_claude(
prompt="hi",
model="claude-opus-4-7",
base_url="http://proxy.example:4000",
api_key="sk-abc",
runner=runner,
)
env = captured["env"]
assert env["ANTHROPIC_BASE_URL"] == "http://proxy.example:4000"
assert env["ANTHROPIC_AUTH_TOKEN"] == "sk-abc"
def test_run_claude_extra_env_is_added_to_subprocess_env():
"""Caller-supplied extra_env entries land on the subprocess env."""
runner, captured = _make_runner(stdout="")
run_claude(
prompt="hi",
model="claude-opus-4-7",
base_url="http://localhost",
api_key="sk-abc",
extra_env={"MAX_THINKING_TOKENS": "4096"},
runner=runner,
)
assert captured["env"]["MAX_THINKING_TOKENS"] == "4096"
def test_run_claude_inherits_only_allowlisted_os_environ(monkeypatch):
"""Process-runtime vars (PATH) flow through; credentials don't.
The `claude` CLI is a Node binary installed dynamically from npm in
CI. If the package were ever compromised, inheriting the entire
parent environment would hand it every credential the surrounding
proxy job loads (AWS keys, Azure Foundry key, GitHub token, etc.).
Pin the contract: only the small allowlist of runtime vars is
inherited; everything else is dropped unless the caller passes it
explicitly via extra_env.
`HOME` is *not* on the allowlist anymore — see the dedicated
isolated-HOME test below for the reason.
"""
monkeypatch.setenv("PATH", "/usr/bin:/usr/local/bin")
monkeypatch.setenv("HOME", "/home/runner")
monkeypatch.setenv("AWS_SECRET_ACCESS_KEY", "totally-secret")
monkeypatch.setenv("ANTHROPIC_API_KEY", "sk-proxy-only")
monkeypatch.setenv("AZURE_FOUNDRY_API_KEY", "azure-secret")
monkeypatch.setenv("VERTEXAI_CREDENTIALS", '{"private_key": "leak"}')
monkeypatch.setenv("GITHUB_TOKEN", "ghs_xxx")
runner, captured = _make_runner(stdout="")
run_claude(
prompt="hi",
model="claude-opus-4-7",
base_url="http://localhost",
api_key="sk-abc",
runner=runner,
)
env = captured["env"]
assert env["PATH"] == "/usr/bin:/usr/local/bin"
assert "AWS_SECRET_ACCESS_KEY" not in env
assert "ANTHROPIC_API_KEY" not in env
assert "AZURE_FOUNDRY_API_KEY" not in env
assert "VERTEXAI_CREDENTIALS" not in env
assert "GITHUB_TOKEN" not in env
def test_run_claude_uses_isolated_per_invocation_home(monkeypatch, tmp_path):
"""`claude` subprocess never sees the runtime user's real $HOME.
The CLI needs *a* HOME (it caches per-session state under
`$HOME/.claude/projects/<sha>/`), but it has no business reading
the runtime user's real one. On the cron VM the runtime user is a
real interactive account with a populated home directory
(~/.config/gh/hosts.yml carrying a GitHub token, ~/.ssh/, etc.);
handing /home/mateo to a compromised npm package — or to a
model-directed `Read` tool call during the PDF/vision cells —
would let it exfiltrate those files. We hand the CLI a fresh
empty per-invocation tmpdir instead.
"""
monkeypatch.setenv("HOME", "/home/runner")
runner, captured = _make_runner(stdout="")
run_claude(
prompt="hi",
model="claude-opus-4-7",
base_url="http://localhost",
api_key="sk-abc",
runner=runner,
)
env = captured["env"]
assert "HOME" in env, "claude CLI needs HOME to find ~/.claude session dir"
assert (
env["HOME"] != "/home/runner"
), "HOME must not leak the parent process's HOME to claude"
# The isolated HOME is a fresh tmpdir prefixed `claude-cli-home-`;
# see `_make_isolated_home` in cli_driver.py. It exists during the
# subprocess call and is removed afterwards (cleanup runs in a
# `finally`, so by the time this assertion runs the dir is gone —
# we only check the *prefix* of the path string we captured).
assert "claude-cli-home-" in env["HOME"]
def test_run_claude_isolated_home_is_distinct_per_invocation(monkeypatch):
"""Two consecutive calls get two different isolated HOMEs.
Reusing a single tmpdir across calls would defeat the isolation
in the parallel matrix run (a compromised CLI could plant a file
in HOME on one model's run and read it on the next). Pin: each
`run_claude` invocation gets its own freshly-created HOME.
"""
monkeypatch.setenv("HOME", "/home/runner")
runner, captured = _make_runner(stdout="")
run_claude(
prompt="hi",
model="claude-opus-4-7",
base_url="http://localhost",
api_key="sk-abc",
runner=runner,
)
home_a = captured["env"]["HOME"]
runner, captured = _make_runner(stdout="")
run_claude(
prompt="hi",
model="claude-opus-4-7",
base_url="http://localhost",
api_key="sk-abc",
runner=runner,
)
home_b = captured["env"]["HOME"]
assert home_a != home_b
def test_run_claude_isolated_home_cleaned_up_after_run(monkeypatch):
"""The per-invocation HOME tmpdir is rm-rf'd when run_claude returns.
Without cleanup, a long matrix run would accumulate one tmpdir
per cell × per model × per CLI call (~75 dirs per cron run,
growing without bound across days).
"""
import os as _os
monkeypatch.setenv("HOME", "/home/runner")
runner, captured = _make_runner(stdout="")
run_claude(
prompt="hi",
model="claude-opus-4-7",
base_url="http://localhost",
api_key="sk-abc",
runner=runner,
)
isolated_home = captured["env"]["HOME"]
assert not _os.path.exists(
isolated_home
), f"isolated HOME {isolated_home!r} should be removed after run_claude returns"
def test_run_claude_isolated_home_cleaned_up_on_subprocess_failure(monkeypatch):
"""Cleanup runs even when the CLI subprocess raises.
If the CLI is missing or times out, `run_claude` raises
`ClaudeCLIError` — but the per-invocation HOME tmpdir must still
be removed (the `finally` clause), otherwise long failure-prone
runs leak tmpdirs.
"""
import os as _os
monkeypatch.setenv("HOME", "/home/runner")
captured: dict = {}
def runner(cmd, env, capture_output, text, timeout, check, input=None):
captured["env"] = env
raise subprocess.TimeoutExpired(cmd=cmd, timeout=timeout)
with pytest.raises(ClaudeCLIError):
run_claude(
prompt="hi",
model="claude-opus-4-7",
base_url="http://localhost",
api_key="sk-abc",
runner=runner,
)
isolated_home = captured["env"]["HOME"]
assert not _os.path.exists(
isolated_home
), f"isolated HOME {isolated_home!r} should be removed even on timeout"
def test_run_claude_extra_env_can_pass_through_otherwise_blocked_var(monkeypatch):
"""The allowlist applies to inherited os.environ; extra_env is the
sanctioned way for a test to opt-in to passing something extra."""
monkeypatch.setenv("ANTHROPIC_API_KEY", "from-os")
runner, captured = _make_runner(stdout="")
run_claude(
prompt="hi",
model="claude-opus-4-7",
base_url="http://localhost",
api_key="sk-abc",
extra_env={"ANTHROPIC_API_KEY": "from-arg"},
runner=runner,
)
assert captured["env"]["ANTHROPIC_API_KEY"] == "from-arg"
def test_run_claude_parses_stream_json_assistant_text():
events = [
{"type": "system", "session_id": "abc"},
{
"type": "assistant",
"message": {
"content": [
{"type": "text", "text": "Hello "},
{"type": "text", "text": "world"},
]
},
},
{"type": "result", "usage": {"input_tokens": 10, "output_tokens": 2}},
]
stdout = "\n".join(json.dumps(e) for e in events) + "\n"
runner, _ = _make_runner(stdout=stdout)
result = run_claude(
prompt="hi",
model="claude-haiku-4-5",
base_url="http://localhost",
api_key="sk-abc",
runner=runner,
)
assert isinstance(result, DriverResult)
assert result.text == "Hello world"
assert len(result.events) == 3
assert result.usage == {"input_tokens": 10, "output_tokens": 2}
assert result.exit_code == 0
def test_run_claude_handles_string_message_content():
"""Some CLI versions emit `message.content` as a plain string."""
stdout = (
json.dumps({"type": "assistant", "message": {"content": "bare text"}}) + "\n"
)
runner, _ = _make_runner(stdout=stdout)
result = run_claude(
prompt="hi",
model="m",
base_url="http://x",
api_key="k",
runner=runner,
)
assert result.text == "bare text"
def test_run_claude_skips_malformed_lines():
stdout = (
"not-json\n"
+ json.dumps(
{
"type": "assistant",
"message": {"content": [{"type": "text", "text": "x"}]},
}
)
+ "\n"
+ "{also-bad\n"
)
runner, _ = _make_runner(stdout=stdout)
result = run_claude(
prompt="hi",
model="m",
base_url="http://x",
api_key="k",
runner=runner,
)
assert result.text == "x"
assert len(result.events) == 1
def test_run_claude_propagates_nonzero_exit_code():
runner, _ = _make_runner(stdout="", returncode=2, stderr="auth failed")
result = run_claude(
prompt="hi",
model="m",
base_url="http://x",
api_key="k",
runner=runner,
)
assert result.exit_code == 2
assert result.stderr == "auth failed"
assert result.text == ""
def test_run_claude_raises_on_missing_cli():
def runner(*args, **kwargs):
raise FileNotFoundError(2, "no such file", "claude")
with pytest.raises(ClaudeCLIError, match="claude CLI not found"):
run_claude(
prompt="hi",
model="m",
base_url="http://x",
api_key="k",
runner=runner,
)
def test_run_claude_raises_on_timeout():
def runner(*args, **kwargs):
raise subprocess.TimeoutExpired(cmd="claude", timeout=1)
with pytest.raises(ClaudeCLIError, match="timed out"):
run_claude(
prompt="hi",
model="m",
base_url="http://x",
api_key="k",
timeout=1,
runner=runner,
)
def test_run_claude_validates_required_params():
runner, _ = _make_runner()
with pytest.raises(ValueError, match="prompt"):
run_claude(
prompt="",
model="m",
base_url="http://x",
api_key="k",
runner=runner,
)
with pytest.raises(ValueError, match="stdin_input"):
run_claude(
prompt=None,
stdin_input="",
model="m",
base_url="http://x",
api_key="k",
runner=runner,
)
with pytest.raises(ValueError, match="model"):
run_claude(
prompt="hi",
model="",
base_url="http://x",
api_key="k",
runner=runner,
)
with pytest.raises(ValueError, match="base_url"):
run_claude(
prompt="hi",
model="m",
base_url="",
api_key="k",
runner=runner,
)
with pytest.raises(ValueError, match="api_key"):
run_claude(
prompt="hi",
model="m",
base_url="http://x",
api_key="",
runner=runner,
)
# ---------------------------------------------------------------------------
# failure_diagnostic
#
# Regression coverage for the bring-up incident where the proxy was started
# with the wrong config and tests reported only `claude CLI exited 1` while
# the actual 400 from LiteLLM was sitting in stdout. The helper must surface
# api_status, the assistant text (where API errors land), stderr, and the
# exit code together — and gracefully degrade when individual pieces are
# missing.
# ---------------------------------------------------------------------------
def test_failure_diagnostic_surfaces_api_error_text_from_stdout():
"""The CLI hides 4xx/5xx from the proxy in `assistant.message.content` text."""
api_error_text = (
'API Error: 400 {"error":{"message":"litellm.BadRequestError: '
"You passed in model=claude-haiku-4-5. There are no healthy "
'deployments..."}}'
)
result = DriverResult(
text=api_error_text,
events=[
{
"type": "assistant",
"message": {"content": [{"type": "text", "text": api_error_text}]},
},
{
"type": "result",
"is_error": True,
"api_error_status": 400,
"result": api_error_text,
},
],
exit_code=1,
stderr="",
)
diag = failure_diagnostic(result)
assert "exit=1" in diag
assert "api_status=400" in diag
assert "There are no healthy deployments" in diag
def test_failure_diagnostic_falls_back_to_stderr_when_no_text():
result = DriverResult(text="", events=[], exit_code=2, stderr="boom\n")
diag = failure_diagnostic(result)
assert "exit=2" in diag
assert "stderr=boom" in diag
def test_failure_diagnostic_handles_completely_empty_result():
"""A run that produced literally nothing should still yield a useful string."""
result = DriverResult(text="", events=[], exit_code=137, stderr="")
diag = failure_diagnostic(result)
assert "exit=137" in diag
assert "no diagnostic output" in diag
def test_failure_diagnostic_truncates_long_text():
"""Don't let a 5MB HTML 502 page from a load balancer wreck the matrix JSON."""
huge = "x" * 5000
result = DriverResult(text=huge, events=[], exit_code=1, stderr="")
diag = failure_diagnostic(result, max_len=100)
assert "truncated" in diag
# Allow some slack for the prefix/suffix/separator characters.
assert len(diag) < 300
def test_failure_diagnostic_ignores_non_int_api_error_status():
"""The CLI sometimes emits api_error_status as a string; don't crash."""
result = DriverResult(
text="oops",
events=[{"type": "result", "api_error_status": "n/a"}],
exit_code=1,
stderr="",
)
diag = failure_diagnostic(result)
assert "api_status" not in diag
assert "text=oops" in diag
# ---------------------------------------------------------------------------
# run_claude_models_parallel
#
# The matrix runs three Claude tiers per cell, so the parallel helper has to
# (a) invoke `run_claude` once per model, (b) preserve each model's outcome
# separately, and (c) return errors as values rather than raising — callers
# need both the failed and the succeeded model results to report per-cell
# rows accurately.
# ---------------------------------------------------------------------------
def test_run_claude_models_parallel_returns_one_result_per_model():
"""Each model gets its own DriverResult keyed under the helper's dict."""
seen_models: List[str] = []
def runner(cmd, env, capture_output, text, timeout, check, input=None):
# The model id is two slots after `--model` in the assembled command.
idx = cmd.index("--model")
model = cmd[idx + 1]
seen_models.append(model)
return _Completed(
returncode=0,
stdout=json.dumps(
{
"type": "assistant",
"message": {
"content": [{"type": "text", "text": f"reply-{model}"}]
},
}
)
+ "\n",
)
outcomes = run_claude_models_parallel(
models=["a", "b", "c"],
prompt="hi",
base_url="http://x",
api_key="k",
runner=runner,
)
assert set(outcomes.keys()) == {"a", "b", "c"}
for model in ("a", "b", "c"):
result = outcomes[model]
assert isinstance(result, DriverResult)
assert result.text == f"reply-{model}"
assert result.exit_code == 0
assert sorted(seen_models) == ["a", "b", "c"]
def test_run_claude_models_parallel_returns_errors_as_values():
"""A model whose CLI is missing surfaces as a ClaudeCLIError, not a raise."""
def runner(cmd, env, capture_output, text, timeout, check, input=None):
idx = cmd.index("--model")
model = cmd[idx + 1]
if model == "boom":
raise FileNotFoundError(2, "no such file", "claude")
return _Completed(
returncode=0,
stdout=json.dumps(
{
"type": "assistant",
"message": {"content": [{"type": "text", "text": "ok"}]},
}
)
+ "\n",
)
outcomes = run_claude_models_parallel(
models=["ok-model", "boom"],
prompt="hi",
base_url="http://x",
api_key="k",
runner=runner,
)
assert isinstance(outcomes["ok-model"], DriverResult)
assert outcomes["ok-model"].text == "ok"
assert isinstance(outcomes["boom"], ClaudeCLIError)
assert "claude CLI not found" in str(outcomes["boom"])
def test_run_claude_models_parallel_preserves_nonzero_exit_codes():
"""Mixed success/failure on exit code should not collapse into one verdict."""
def runner(cmd, env, capture_output, text, timeout, check, input=None):
idx = cmd.index("--model")
model = cmd[idx + 1]
if model == "fail":
return _Completed(returncode=2, stdout="", stderr="auth failed")
return _Completed(
returncode=0,
stdout=json.dumps(
{
"type": "assistant",
"message": {"content": [{"type": "text", "text": "ok"}]},
}
)
+ "\n",
)
outcomes = run_claude_models_parallel(
models=["ok-model", "fail"],
prompt="hi",
base_url="http://x",
api_key="k",
runner=runner,
)
assert outcomes["ok-model"].exit_code == 0
assert outcomes["fail"].exit_code == 2
assert outcomes["fail"].stderr == "auth failed"
def test_run_claude_models_parallel_rejects_empty_models():
with pytest.raises(ValueError, match="non-empty"):
run_claude_models_parallel(
models=[],
prompt="hi",
base_url="http://x",
api_key="k",
)
def test_run_claude_models_parallel_stamps_duration_on_each_result():
"""Each DriverResult carries the per-model wall time so callers can
attribute slow cells without re-timing the work themselves.
The fake runner sleeps for very different durations per model so
we can prove each result is timing its own work (not the batch
wall time). We use generous absolute bounds because thread-pool
scheduling on a loaded CI box adds noise on the order of tens of
milliseconds.
"""
import time
def runner(cmd, env, capture_output, text, timeout, check, input=None):
idx = cmd.index("--model")
model = cmd[idx + 1]
time.sleep(0.05 if model == "fast" else 0.40)
return _Completed(returncode=0, stdout="")
outcomes = run_claude_models_parallel(
models=["fast", "slow"],
prompt="hi",
base_url="http://x",
api_key="k",
runner=runner,
)
fast_ms = outcomes["fast"].duration_ms
slow_ms = outcomes["slow"].duration_ms
assert fast_ms is not None and slow_ms is not None
# 50ms sleep ⇒ ~50–250ms after scheduling overhead; 400ms sleep ⇒
# 400–700ms. We just need the two distributions to be non-overlapping
# so we know each row's duration is its own work, not the batch's.
assert fast_ms < 300, fast_ms
assert slow_ms >= 350, slow_ms
assert slow_ms > fast_ms
def test_run_claude_models_parallel_breakdown_logs_to_stderr(capsys):
"""The breakdown helper must emit a per-model timing block so users
can answer "why didn't parallel help?" without re-instrumenting."""
def runner(cmd, env, capture_output, text, timeout, check, input=None):
return _Completed(returncode=0, stdout="")
run_claude_models_parallel(
models=["model-x", "model-y"],
prompt="hi",
base_url="http://x",
api_key="k",
runner=runner,
)
captured = capsys.readouterr()
assert "[parallel] per-model wall time:" in captured.err
assert "model-x" in captured.err
assert "model-y" in captured.err
assert "speedup=" in captured.err
assert "slowest=" in captured.err
def test_run_claude_models_parallel_breakdown_marks_cli_errors(capsys):
"""When a model raises ClaudeCLIError, the breakdown should still
show its row tagged as `cli-error` rather than crashing or omitting it."""
def runner(cmd, env, capture_output, text, timeout, check, input=None):
idx = cmd.index("--model")
if cmd[idx + 1] == "boom":
raise FileNotFoundError(2, "no such file", "claude")
return _Completed(returncode=0, stdout="")
run_claude_models_parallel(
models=["ok-model", "boom"],
prompt="hi",
base_url="http://x",
api_key="k",
runner=runner,
)
captured = capsys.readouterr()
assert "ok-model" in captured.err
assert "boom" in captured.err
assert "cli-error" in captured.err
def test_run_claude_models_parallel_forwards_extra_args_and_env():
"""Shared kwargs must reach every per-model invocation unchanged."""
captured_envs: List[dict] = []
captured_cmds: List[List[str]] = []
def runner(cmd, env, capture_output, text, timeout, check, input=None):
captured_envs.append(env)
captured_cmds.append(cmd)
return _Completed(returncode=0, stdout="")
run_claude_models_parallel(
models=["a", "b"],
prompt="hi",
base_url="http://x",
api_key="k",
extra_env={"MAX_THINKING_TOKENS": "4096"},
extra_args=["--allowed-tools", "Bash"],
runner=runner,
)
assert all(env["MAX_THINKING_TOKENS"] == "4096" for env in captured_envs)
for cmd in captured_cmds:
assert "--allowed-tools" in cmd
assert "Bash" in cmd
def test_failure_diagnostic_uses_last_result_event_status():
"""If multiple `result` events appear, the most recent status wins."""
result = DriverResult(
text="",
events=[
{"type": "result", "api_error_status": 500},
{"type": "assistant", "message": {"content": []}},
{"type": "result", "api_error_status": 429},
],
exit_code=1,
stderr="",
)
diag = failure_diagnostic(result)
assert "api_status=429" in diag
assert "500" not in diag
_RATE_LIMITED_STDOUT = (
json.dumps(
{
"type": "assistant",
"message": {
"content": [
{"type": "text", "text": "API Error: 429 Too Many Requests"}
]
},
}
)
+ "\n"
+ json.dumps({"type": "result", "api_error_status": 429})
+ "\n"
)
_OK_STDOUT = (
json.dumps(
{
"type": "assistant",
"message": {"content": [{"type": "text", "text": "pong"}]},
}
)
+ "\n"
)
class _FlakyRunner:
"""Fake runner that rate-limits each model N times before succeeding.
Keeps a per-model call count so tests can assert exactly how many
attempts the retry loop made — the load-bearing detail a canned
single-response runner can't express.
"""
def __init__(self, failures_before_success: dict):
self.failures_before_success = dict(failures_before_success)
self.calls: dict = {}
def __call__(self, cmd, env, capture_output, text, timeout, check, input=None):
model = cmd[cmd.index("--model") + 1]
self.calls[model] = self.calls.get(model, 0) + 1
if self.calls[model] <= self.failures_before_success.get(model, 0):
return _Completed(returncode=1, stdout=_RATE_LIMITED_STDOUT)
return _Completed(returncode=0, stdout=_OK_STDOUT)
@pytest.mark.parametrize(
"outcome,expected",
[
(ClaudeCLIError("claude CLI timed out after 120.0s"), True),
(ClaudeCLIError("claude CLI not found at 'claude'"), False),
(
DriverResult(
text="",
events=[{"type": "result", "api_error_status": 429}],
exit_code=1,
),
True,
),
(DriverResult(text="Too Many Requests", exit_code=1), True),
(DriverResult(text="", stderr="throttled by upstream", exit_code=1), True),
(DriverResult(text="rate limit exceeded", exit_code=0), False),
(DriverResult(text="", stderr="auth failed", exit_code=2), False),
],
)
def test_is_rate_limit_shaped_classification(outcome, expected):
"""The retry trigger must match 429/throttle/timeout markers on
failures only — a passing result mentioning '429' in its reply text
must never be classified as retryable."""
assert is_rate_limit_shaped(outcome) is expected
def test_run_claude_models_parallel_retries_rate_limited_model_until_success():
"""A model that 429s once must be retried after the backoff sleep and
end up green, while an untroubled sibling model runs exactly once."""
runner = _FlakyRunner({"flaky": 1})
sleeps: List[float] = []
outcomes = run_claude_models_parallel(
models=["flaky", "steady"],
prompt="hi",
base_url="http://x",
api_key="k",
runner=runner,
rate_limit_retries=2,
rate_limit_backoff_seconds=0.5,
sleep=sleeps.append,
)
assert isinstance(outcomes["flaky"], DriverResult)
assert outcomes["flaky"].exit_code == 0
assert outcomes["flaky"].text == "pong"
assert runner.calls == {"flaky": 2, "steady": 1}
assert sleeps == [0.5]
def test_run_claude_models_parallel_does_not_retry_non_rate_limit_failures():
"""A deterministic failure (bad auth) must fail fast: no sleeps, one
attempt — retrying it would just triple the matrix wall time."""
def runner(cmd, env, capture_output, text, timeout, check, input=None):
return _Completed(returncode=2, stdout="", stderr="auth failed")
sleeps: List[float] = []
outcomes = run_claude_models_parallel(
models=["a"],
prompt="hi",
base_url="http://x",
api_key="k",
runner=runner,
rate_limit_retries=2,
rate_limit_backoff_seconds=0.5,
sleep=sleeps.append,
)
assert outcomes["a"].exit_code == 2
assert sleeps == []
def test_run_claude_models_parallel_returns_last_failure_when_retries_exhausted():
"""A persistently rate-limited model exhausts its budget (initial
attempt + N retries, each preceded by one backoff sleep) and still
surfaces the 429 diagnostic instead of masking it."""
runner = _FlakyRunner({"stuck": 99})
sleeps: List[float] = []
outcomes = run_claude_models_parallel(
models=["stuck"],
prompt="hi",
base_url="http://x",
api_key="k",
runner=runner,
rate_limit_retries=2,
rate_limit_backoff_seconds=0.25,
sleep=sleeps.append,
)
assert runner.calls == {"stuck": 3}
assert sleeps == [0.25, 0.25]
assert outcomes["stuck"].exit_code == 1
assert "429" in failure_diagnostic(outcomes["stuck"])
def test_run_claude_models_parallel_retries_timeout_shaped_cli_errors():
"""CLI timeouts are how saturated upstreams usually present (the CLI
retries 429s internally until the harness kills it), so a timeout
must be retried like an explicit 429."""
calls: List[int] = []
def runner(cmd, env, capture_output, text, timeout, check, input=None):
calls.append(1)
if len(calls) == 1:
raise subprocess.TimeoutExpired(cmd="claude", timeout=1)
return _Completed(returncode=0, stdout=_OK_STDOUT)
sleeps: List[float] = []
outcomes = run_claude_models_parallel(
models=["a"],
prompt="hi",
base_url="http://x",
api_key="k",
runner=runner,
rate_limit_retries=1,
rate_limit_backoff_seconds=0.5,
sleep=sleeps.append,
)
assert isinstance(outcomes["a"], DriverResult)
assert outcomes["a"].text == "pong"
assert len(calls) == 2
assert sleeps == [0.5]
def test_run_claude_models_parallel_zero_retries_disables_backoff():
"""`rate_limit_retries=0` must restore the old single-attempt
behavior exactly: one call, no sleeps, failure returned as-is."""
runner = _FlakyRunner({"stuck": 99})
sleeps: List[float] = []
outcomes = run_claude_models_parallel(
models=["stuck"],
prompt="hi",
base_url="http://x",
api_key="k",
runner=runner,
rate_limit_retries=0,
rate_limit_backoff_seconds=0.5,
sleep=sleeps.append,
)
assert runner.calls == {"stuck": 1}
assert sleeps == []
assert outcomes["stuck"].exit_code == 1

View file

@ -1,138 +0,0 @@
"""Tests for the `compat_result` fixture's tagged-union validation.
The conftest's `pytest_runtest_makereport` hook is exercised end-to-end by
the matrix-builder golden-file tests (which consume a results.json that
the harness would produce). Here we just test the input-validation
contract on `CompatResult.set()`.
"""
from __future__ import annotations
import pytest
from claude_code.conftest import CompatResult
def test_set_pass_is_accepted():
r = CompatResult()
r.set({"status": "pass"})
assert r.value == {"status": "pass"}
def test_set_fail_requires_error():
r = CompatResult()
with pytest.raises(ValueError, match="requires 'error'"):
r.set({"status": "fail"})
def test_set_fail_with_error_is_accepted():
r = CompatResult()
r.set({"status": "fail", "error": "boom"})
assert r.value == {"status": "fail", "error": "boom"}
def test_set_not_applicable_requires_reason():
r = CompatResult()
with pytest.raises(ValueError, match="requires 'reason'"):
r.set({"status": "not_applicable"})
def test_set_not_applicable_with_reason_is_accepted():
r = CompatResult()
r.set({"status": "not_applicable", "reason": "Bedrock has no /thinking"})
assert r.value == {"status": "not_applicable", "reason": "Bedrock has no /thinking"}
def test_set_not_tested_is_accepted():
r = CompatResult()
r.set({"status": "not_tested"})
assert r.value == {"status": "not_tested"}
def test_set_rejects_unknown_status():
r = CompatResult()
with pytest.raises(ValueError, match="status must be one of"):
r.set({"status": "maybe"})
def test_set_rejects_non_dict():
r = CompatResult()
with pytest.raises(TypeError):
r.set("pass") # type: ignore[arg-type]
def test_set_copies_input():
"""Mutating the dict after set() must not change the stored value."""
r = CompatResult()
payload = {"status": "fail", "error": "x"}
r.set(payload)
payload["error"] = "mutated"
assert r.value["error"] == "x"
# ---------------------------------------------------------------------------
# add() / collected()
#
# When a single test exercises three Claude tiers in parallel, each tier
# needs its own row in the results artifact so the matrix builder can
# apply its "all three must pass" aggregation. `add()` is the per-tier
# recorder; `collected()` is what the conftest hook reads.
# ---------------------------------------------------------------------------
def test_add_appends_each_call_to_values():
r = CompatResult()
r.add({"status": "pass"})
r.add({"status": "fail", "error": "bad"})
assert r.values == [
{"status": "pass"},
{"status": "fail", "error": "bad"},
]
def test_add_validates_like_set():
"""The add() and set() validators are the same; both must reject bad payloads."""
r = CompatResult()
with pytest.raises(ValueError, match="requires 'error'"):
r.add({"status": "fail"})
with pytest.raises(ValueError, match="requires 'reason'"):
r.add({"status": "not_applicable"})
with pytest.raises(ValueError, match="status must be one of"):
r.add({"status": "maybe"})
with pytest.raises(TypeError):
r.add("pass") # type: ignore[arg-type]
def test_add_copies_input():
"""Same defensive copy contract as set()."""
r = CompatResult()
payload = {"status": "fail", "error": "x"}
r.add(payload)
payload["error"] = "mutated"
assert r.values[0]["error"] == "x"
def test_collected_returns_values_when_added():
r = CompatResult()
r.add({"status": "pass"})
r.add({"status": "pass"})
assert r.collected() == [{"status": "pass"}, {"status": "pass"}]
def test_collected_returns_single_value_when_only_set_called():
"""Legacy single-result tests should still surface their one outcome."""
r = CompatResult()
r.set({"status": "pass"})
assert r.collected() == [{"status": "pass"}]
def test_collected_prefers_added_values_over_set_value():
"""If both are populated, the per-tier list wins — that's the multi-model shape."""
r = CompatResult()
r.set({"status": "pass"})
r.add({"status": "fail", "error": "tier-2 broke"})
assert r.collected() == [{"status": "fail", "error": "tier-2 broke"}]
def test_collected_returns_empty_when_nothing_reported():
assert CompatResult().collected() == []

View file

@ -1,195 +0,0 @@
"""Unit tests for the shared `run_passthrough_cell` helper.
These tests inject a fake `run_models` callable and an explicit `env`
mapping (both are first-class parameters, no monkeypatching), so they
exercise the helper's branching -- env-missing guard, base-URL
assembly, extra-env forwarding, per-model pass/fail -- without
spawning the real CLI.
The env-builder tests pin the provider-mode contract itself: the
CLAUDE_CODE_USE_* / CLAUDE_CODE_SKIP_*_AUTH flags and the passthrough
route each mode must target. Those values are the feature -- e.g.
dropping the `/v1` from the vertex base URL produces a request Google
404s on -- so a mutation to any of them must fail here before it burns
a live matrix run.
"""
from __future__ import annotations
from typing import Any, Dict, List, Mapping, Optional
import pytest
from claude_code._passthrough import (
ANTHROPIC_PASSTHROUGH_BASE_PATH,
CLIENT_SIDE_AWS_REGION,
VERTEX_PLACEHOLDER_PROJECT,
VERTEX_PLACEHOLDER_REGION,
bedrock_extra_env,
foundry_extra_env,
run_passthrough_cell,
vertex_extra_env,
)
from claude_code.cli_driver import ClaudeCLIError, DriverResult
PROXY_ENV = {
"LITELLM_PROXY_BASE_URL": "http://localhost:4000",
"LITELLM_PROXY_API_KEY": "sk-test",
}
class _FakeResult:
def __init__(self) -> None:
self.rows: List[Dict[str, Any]] = []
self.single: Optional[Dict[str, Any]] = None
def set(self, payload: Mapping[str, Any]) -> None:
self.single = dict(payload)
def add(self, payload: Mapping[str, Any]) -> None:
self.rows.append(dict(payload))
def _fake_run_models(outcomes_by_model, captured: Dict[str, Any]):
def fake(*, models, prompt, base_url, api_key, extra_env=None, **_kwargs):
captured["models"] = list(models)
captured["prompt"] = prompt
captured["base_url"] = base_url
captured["api_key"] = api_key
captured["extra_env"] = dict(extra_env) if extra_env is not None else None
return {model: outcomes_by_model[model] for model in models}
return fake
def test_env_missing_guard_reports_fail_and_aborts():
fake_result = _FakeResult()
with pytest.raises(pytest.fail.Exception):
run_passthrough_cell(
compat_result=fake_result,
models=["claude-haiku-4-5"],
prompt="ping",
env={},
)
assert fake_result.single is not None
assert fake_result.single["status"] == "fail"
assert "LITELLM_PROXY_BASE_URL" in fake_result.single["error"]
def test_anthropic_base_path_appended_to_normalized_proxy_url():
fake_result = _FakeResult()
captured: Dict[str, Any] = {}
outcome = DriverResult(text="pong")
run_passthrough_cell(
compat_result=fake_result,
models=["claude-haiku-4-5"],
prompt="ping",
passthrough_base_path=ANTHROPIC_PASSTHROUGH_BASE_PATH,
run_models=_fake_run_models({"claude-haiku-4-5": outcome}, captured),
env={**PROXY_ENV, "LITELLM_PROXY_BASE_URL": "http://localhost:4000/"},
)
assert captured["base_url"] == "http://localhost:4000/anthropic"
assert captured["extra_env"] is None
assert fake_result.rows == [{"status": "pass"}]
def test_extra_env_builder_receives_normalized_base_and_is_forwarded():
fake_result = _FakeResult()
captured: Dict[str, Any] = {}
outcome = DriverResult(text="pong")
seen_bases: List[str] = []
def build(proxy_base: str) -> Dict[str, str]:
seen_bases.append(proxy_base)
return {"SOME_FLAG": "1"}
run_passthrough_cell(
compat_result=fake_result,
models=["claude-haiku-4-5"],
prompt="ping",
build_extra_env=build,
run_models=_fake_run_models({"claude-haiku-4-5": outcome}, captured),
env={**PROXY_ENV, "LITELLM_PROXY_BASE_URL": "http://localhost:4000/"},
)
assert seen_bases == ["http://localhost:4000"]
assert captured["extra_env"] == {"SOME_FLAG": "1"}
assert captured["base_url"] == "http://localhost:4000"
def test_per_model_failures_reported_individually():
fake_result = _FakeResult()
captured: Dict[str, Any] = {}
outcomes = {
"claude-haiku-4-5": DriverResult(text="pong"),
"claude-sonnet-4-6": ClaudeCLIError("claude CLI timed out after 120s"),
"claude-opus-4-7": DriverResult(text="", exit_code=1),
}
with pytest.raises(pytest.fail.Exception):
run_passthrough_cell(
compat_result=fake_result,
models=list(outcomes.keys()),
prompt="ping",
run_models=_fake_run_models(outcomes, captured),
env=PROXY_ENV,
)
statuses = [row["status"] for row in fake_result.rows]
assert statuses == ["pass", "fail", "fail"]
assert "timed out" in fake_result.rows[1]["error"]
assert "claude CLI failed" in fake_result.rows[2]["error"]
def test_empty_assistant_text_is_a_fail():
fake_result = _FakeResult()
captured: Dict[str, Any] = {}
outcomes = {"claude-haiku-4-5": DriverResult(text=" ")}
with pytest.raises(pytest.fail.Exception):
run_passthrough_cell(
compat_result=fake_result,
models=["claude-haiku-4-5"],
prompt="ping",
run_models=_fake_run_models(outcomes, captured),
env=PROXY_ENV,
)
assert fake_result.rows == [
{
"status": "fail",
"error": "[claude-haiku-4-5] claude returned empty assistant text",
}
]
def test_bedrock_extra_env_targets_proxy_bedrock_route():
env = bedrock_extra_env("http://localhost:4000")
assert env == {
"CLAUDE_CODE_USE_BEDROCK": "1",
"CLAUDE_CODE_SKIP_BEDROCK_AUTH": "1",
"ANTHROPIC_BEDROCK_BASE_URL": "http://localhost:4000/bedrock",
"AWS_REGION": CLIENT_SIDE_AWS_REGION,
}
def test_vertex_extra_env_keeps_the_api_version_in_the_base_url():
env = vertex_extra_env("http://localhost:4000")
assert env == {
"CLAUDE_CODE_USE_VERTEX": "1",
"CLAUDE_CODE_SKIP_VERTEX_AUTH": "1",
"ANTHROPIC_VERTEX_BASE_URL": "http://localhost:4000/vertex_ai/v1",
"ANTHROPIC_VERTEX_PROJECT_ID": VERTEX_PLACEHOLDER_PROJECT,
"CLOUD_ML_REGION": VERTEX_PLACEHOLDER_REGION,
}
def test_foundry_extra_env_targets_proxy_azure_route():
env = foundry_extra_env("http://localhost:4000")
assert env == {
"CLAUDE_CODE_USE_FOUNDRY": "1",
"CLAUDE_CODE_SKIP_FOUNDRY_AUTH": "1",
"ANTHROPIC_FOUNDRY_BASE_URL": "http://localhost:4000/azure",
}

View file

@ -1,329 +0,0 @@
"""Unit tests for the cross-process token-bucket rate limiter.
The tests cover three layers:
1. Provider inference from model alias — the matrix-column mapping the
live tests rely on (`-bedrock-converse` vs `-bedrock-invoke` vs
`-azure` vs `-vertex` vs bare = anthropic).
2. Config parsing — env-var precedence, fallback to default, malformed
input handling, burst override semantics. These run against
`os.environ`-shaped dicts so we don't have to monkeypatch globals.
3. Token-bucket behavior — enforcing rate, accumulating burst, never
over-spending across a fake clock. Filesystem state is exercised
with a real `tmp_path` because the persistence is the whole point;
the only injected seam is `clock` (and `sleep`, so tests don't
actually wait on wall time).
The cross-process flock semantics are exercised indirectly: every
test creates a fresh `RateLimiter` rooted at `tmp_path`, so the same
file lock that protects production is exercised here too. We don't
fork to test multi-process behavior in this file because pytest
fixtures + xdist already do that for the integration suite.
"""
from __future__ import annotations
import json
import time
from pathlib import Path
from typing import List
import pytest
from claude_code.rate_limiter import (
ALL_PROVIDERS,
BURST_ENV,
DEFAULT_RATE,
PROVIDER_ANTHROPIC,
PROVIDER_AZURE,
PROVIDER_BEDROCK_CONVERSE,
PROVIDER_BEDROCK_INVOKE,
PROVIDER_VERTEX_AI,
ProviderConfig,
RateLimiter,
get_default_limiter,
infer_provider,
load_config,
reset_default_limiter,
use_limiter,
)
# ---------------------------------------------------------------------------
# Provider inference
# ---------------------------------------------------------------------------
@pytest.mark.parametrize(
"model, expected",
[
("claude-haiku-4-5", PROVIDER_ANTHROPIC),
("claude-sonnet-4-6", PROVIDER_ANTHROPIC),
("claude-opus-4-7", PROVIDER_ANTHROPIC),
("claude-haiku-4-5-azure", PROVIDER_AZURE),
("claude-sonnet-4-6-azure", PROVIDER_AZURE),
("claude-opus-4-7-vertex", PROVIDER_VERTEX_AI),
("claude-haiku-4-5-bedrock-converse", PROVIDER_BEDROCK_CONVERSE),
("claude-haiku-4-5-bedrock-invoke", PROVIDER_BEDROCK_INVOKE),
],
)
def test_infer_provider_maps_alias_suffix_to_column(model, expected):
assert infer_provider(model) == expected
def test_infer_provider_bedrock_converse_beats_bedrock_invoke_lookup_order():
"""Both bedrock suffixes contain `bedrock`; the more-specific suffix wins."""
assert infer_provider("claude-foo-bedrock-converse") == PROVIDER_BEDROCK_CONVERSE
assert infer_provider("claude-foo-bedrock-invoke") == PROVIDER_BEDROCK_INVOKE
def test_infer_provider_rejects_empty_string():
with pytest.raises(ValueError, match="non-empty"):
infer_provider("")
def test_infer_provider_is_case_insensitive():
"""Aliases in the proxy config sometimes drift between cases; we
should still route them to the right column."""
assert infer_provider("CLAUDE-OPUS-4-7-AZURE") == PROVIDER_AZURE
# ---------------------------------------------------------------------------
# Config parsing
# ---------------------------------------------------------------------------
def test_load_config_uses_default_rate_when_env_absent():
cfg = load_config(env={})
for provider in ALL_PROVIDERS:
assert cfg[provider].rate_per_sec == DEFAULT_RATE
assert cfg[provider].burst == DEFAULT_RATE
def test_load_config_reads_per_provider_rate():
cfg = load_config(
env={
"LITELLM_COMPAT_RATE_ANTHROPIC": "10",
"LITELLM_COMPAT_RATE_AZURE": "0.5",
}
)
assert cfg[PROVIDER_ANTHROPIC].rate_per_sec == 10.0
assert cfg[PROVIDER_AZURE].rate_per_sec == 0.5
assert cfg[PROVIDER_VERTEX_AI].rate_per_sec == DEFAULT_RATE
def test_load_config_zero_rate_disables_provider():
cfg = load_config(env={"LITELLM_COMPAT_RATE_BEDROCK_INVOKE": "0"})
assert cfg[PROVIDER_BEDROCK_INVOKE].enabled is False
def test_load_config_burst_override_applies_to_every_provider():
cfg = load_config(
env={
"LITELLM_COMPAT_RATE_ANTHROPIC": "5",
BURST_ENV: "20",
}
)
for provider in ALL_PROVIDERS:
assert cfg[provider].burst == 20.0
def test_load_config_falls_back_on_malformed_value():
cfg = load_config(env={"LITELLM_COMPAT_RATE_ANTHROPIC": "not-a-number"})
assert cfg[PROVIDER_ANTHROPIC].rate_per_sec == DEFAULT_RATE
def test_load_config_burst_floors_at_one_when_rate_is_low():
"""A 0.5/s rate with no burst override must still allow at least
one immediate request — otherwise the very first call would block."""
cfg = load_config(env={"LITELLM_COMPAT_RATE_ANTHROPIC": "0.5"})
assert cfg[PROVIDER_ANTHROPIC].burst == 1.0
# ---------------------------------------------------------------------------
# Token bucket
# ---------------------------------------------------------------------------
@pytest.fixture
def fake_clock():
"""A controllable monotonic clock + sleep for the limiter under test.
Tests advance `clock.now` to simulate elapsed wall time. `sleep`
adds the requested duration to `clock.now` instead of actually
sleeping, so a "wait 200ms" code path runs in microseconds and
is deterministic.
"""
class Clock:
def __init__(self):
self.now = 1_000.0
self.sleeps: List[float] = []
def __call__(self):
return self.now
def sleep(self, seconds: float) -> None:
self.sleeps.append(seconds)
self.now += seconds
return Clock()
def _make_limiter(tmp_path: Path, fake_clock, *, rate=10.0, burst=None):
cfg = {
p: ProviderConfig(rate_per_sec=rate, burst=burst if burst is not None else rate)
for p in ALL_PROVIDERS
}
return RateLimiter(
config=cfg,
state_dir=tmp_path,
clock=fake_clock,
sleep=fake_clock.sleep,
)
def test_acquire_first_call_does_not_wait(tmp_path, fake_clock):
"""A freshly-initialized bucket starts full; the first acquire is free."""
limiter = _make_limiter(tmp_path, fake_clock, rate=10.0, burst=10.0)
waited = limiter.acquire(PROVIDER_ANTHROPIC)
assert waited == 0.0
assert fake_clock.sleeps == []
def test_acquire_disabled_provider_returns_immediately(tmp_path, fake_clock):
"""rate=0 ⇒ no throttling, even if every other provider is throttled."""
cfg = {p: ProviderConfig(rate_per_sec=0.0, burst=0.0) for p in ALL_PROVIDERS}
limiter = RateLimiter(
config=cfg, state_dir=tmp_path, clock=fake_clock, sleep=fake_clock.sleep
)
for _ in range(100):
assert limiter.acquire(PROVIDER_ANTHROPIC) == 0.0
assert fake_clock.sleeps == []
def test_acquire_burns_through_burst_then_throttles(tmp_path, fake_clock):
"""`burst` immediate requests succeed; the next one waits 1/rate seconds."""
limiter = _make_limiter(tmp_path, fake_clock, rate=2.0, burst=3.0)
for _ in range(3):
assert limiter.acquire(PROVIDER_ANTHROPIC) == 0.0
# Bucket is empty; next call must sleep ~0.5s to earn one token at 2/s.
waited = limiter.acquire(PROVIDER_ANTHROPIC)
assert waited == pytest.approx(0.5, abs=0.01)
def test_acquire_refills_with_elapsed_time(tmp_path, fake_clock):
"""Advancing the clock between calls credits tokens at the configured rate."""
limiter = _make_limiter(tmp_path, fake_clock, rate=4.0, burst=1.0)
assert limiter.acquire(PROVIDER_ANTHROPIC) == 0.0 # consumes the 1-token burst
fake_clock.now += 0.25 # 0.25s × 4/s = 1 token earned
assert limiter.acquire(PROVIDER_ANTHROPIC) == 0.0
def test_acquire_caps_refill_at_burst(tmp_path, fake_clock):
"""A long quiet period must not let the bucket grow past `burst`."""
limiter = _make_limiter(tmp_path, fake_clock, rate=10.0, burst=2.0)
fake_clock.now += 1_000 # would earn 10_000 tokens uncapped
# Only `burst` (=2) immediate calls should succeed before throttling.
assert limiter.acquire(PROVIDER_ANTHROPIC) == 0.0
assert limiter.acquire(PROVIDER_ANTHROPIC) == 0.0
waited = limiter.acquire(PROVIDER_ANTHROPIC)
assert waited > 0
def test_acquire_independent_buckets_per_provider(tmp_path, fake_clock):
"""Anthropic exhaustion must not throttle Azure (each column has its own bucket)."""
limiter = _make_limiter(tmp_path, fake_clock, rate=2.0, burst=1.0)
assert limiter.acquire(PROVIDER_ANTHROPIC) == 0.0
# Anthropic bucket is now empty; Azure is untouched.
assert limiter.acquire(PROVIDER_AZURE) == 0.0
def test_acquire_persists_state_across_limiter_instances(tmp_path):
"""A fresh RateLimiter must read the on-disk state, not start fresh.
This is the property that makes the limiter cross-process: an
xdist worker created mid-run sees the credit consumed by other
workers, instead of getting its own private bucket.
"""
cfg = {p: ProviderConfig(rate_per_sec=10.0, burst=2.0) for p in ALL_PROVIDERS}
state = {"now": 1_000.0, "sleeps": []}
def clock():
return state["now"]
def sleep(seconds):
state["sleeps"].append(seconds)
state["now"] += seconds
first = RateLimiter(config=cfg, state_dir=tmp_path, clock=clock, sleep=sleep)
first.acquire(PROVIDER_ANTHROPIC)
first.acquire(PROVIDER_ANTHROPIC)
# bucket is now empty
second = RateLimiter(config=cfg, state_dir=tmp_path, clock=clock, sleep=sleep)
waited = second.acquire(PROVIDER_ANTHROPIC)
assert waited > 0 # had to wait, didn't see a fresh full bucket
def test_acquire_recovers_from_corrupt_state_file(tmp_path, fake_clock):
"""A truncated/garbage state file must not crash the test session."""
state_file = tmp_path / f"{PROVIDER_ANTHROPIC}.json"
state_file.write_text("not-json {{")
limiter = _make_limiter(tmp_path, fake_clock, rate=5.0, burst=5.0)
assert limiter.acquire(PROVIDER_ANTHROPIC) == 0.0
def test_acquire_handles_clock_going_backward(tmp_path, fake_clock):
"""Across a host suspend/resume the monotonic clock can briefly
go backward; we must not interpret that as removing tokens."""
limiter = _make_limiter(tmp_path, fake_clock, rate=1.0, burst=2.0)
limiter.acquire(PROVIDER_ANTHROPIC)
fake_clock.now -= 10 # clock moved backward
# Bucket should still have ~1 token left from the burst, not -9.
assert limiter.acquire(PROVIDER_ANTHROPIC) == 0.0
# ---------------------------------------------------------------------------
# Process-default singleton
# ---------------------------------------------------------------------------
def test_use_limiter_swaps_default_for_block(tmp_path):
sentinel_cfg = {
p: ProviderConfig(rate_per_sec=0.0, burst=0.0) for p in ALL_PROVIDERS
}
sentinel = RateLimiter(config=sentinel_cfg, state_dir=tmp_path)
reset_default_limiter()
try:
with use_limiter(sentinel):
assert get_default_limiter() is sentinel
# After the context exits, the default goes back to whatever it
# was — in this test that's "rebuilt on next access" because we
# called reset_default_limiter() above.
assert get_default_limiter() is not sentinel
finally:
reset_default_limiter()
# ---------------------------------------------------------------------------
# Persistence shape
# ---------------------------------------------------------------------------
def test_state_file_is_json_after_acquire(tmp_path, fake_clock):
limiter = _make_limiter(tmp_path, fake_clock, rate=5.0, burst=5.0)
limiter.acquire(PROVIDER_ANTHROPIC)
state_file = tmp_path / f"{PROVIDER_ANTHROPIC}.json"
payload = json.loads(state_file.read_text())
assert "tokens" in payload
assert "last_refill" in payload
assert payload["tokens"] == pytest.approx(4.0)

View file

@ -0,0 +1,74 @@
"""Proxy-env resolution for the claude_code compat cells.
Uses the same ``LITELLM_PROXY_URL`` / ``LITELLM_MASTER_KEY`` names as
``e2e_config.py`` and the rest of ``tests/e2e/``. Everything under
``claude_code/`` goes through ``resolve_proxy`` / ``require_proxy`` here
so the naming lives in one place.
"""
from __future__ import annotations
import os
from typing import Mapping, NamedTuple
import pytest
class ProxyConfig(NamedTuple):
base_url: str
api_key: str
PRIMARY_BASE_URL_ENV = "LITELLM_PROXY_URL"
PRIMARY_API_KEY_ENV = "LITELLM_MASTER_KEY"
def resolve_proxy_from(mapping: Mapping[str, str]) -> ProxyConfig | None:
"""Pure resolver: takes an env mapping, returns a ProxyConfig if
both a base URL and an API key are present, else None. Extracted so
tests can exercise it without mutating ``os.environ``."""
base_url = mapping.get(PRIMARY_BASE_URL_ENV) or None
api_key = mapping.get(PRIMARY_API_KEY_ENV) or None
if not base_url or not api_key:
return None
return ProxyConfig(base_url=base_url, api_key=api_key)
def resolve_proxy(env: Mapping[str, str] | None = None) -> ProxyConfig | None:
"""Convenience wrapper that defaults to ``os.environ``. Prefer
calling ``resolve_proxy_from(env)`` from tests so nothing has to
reach into the process environment."""
return resolve_proxy_from(os.environ if env is None else env)
def _fail_missing_proxy_env(compat_result) -> None:
compat_result.set(
{
"status": "fail",
"error": (
f"missing required env: set {PRIMARY_BASE_URL_ENV} and "
f"{PRIMARY_API_KEY_ENV} to point at a running LiteLLM proxy"
),
}
)
pytest.fail(
f"{PRIMARY_BASE_URL_ENV} / {PRIMARY_API_KEY_ENV} not configured",
pytrace=False,
)
def require_proxy(
compat_result,
*,
env: Mapping[str, str] | None = None,
) -> ProxyConfig:
"""Return the proxy config (base URL + master key), or hard-fail
the test.
``env`` is injected for tests; production callers pass nothing and
the process env is used. This keeps tests off ``monkeypatch.setenv``
for a check that is a pure function of its inputs."""
cfg = resolve_proxy(env)
if cfg is None:
_fail_missing_proxy_env(compat_result)
return cfg

View file

@ -54,20 +54,17 @@ them unset.
from __future__ import annotations
import os
from typing import Any, Callable, Dict, Mapping, Optional, Sequence
import pytest
from claude_code._env import require_proxy
from claude_code.cli_driver import (
ClaudeCLIError,
failure_diagnostic,
run_claude_models_parallel,
)
PROXY_BASE_URL_ENV = "LITELLM_PROXY_BASE_URL"
PROXY_API_KEY_ENV = "LITELLM_PROXY_API_KEY"
ANTHROPIC_PASSTHROUGH_BASE_PATH = "/anthropic"
CLIENT_SIDE_AWS_REGION = "us-east-1"
@ -140,32 +137,15 @@ def run_passthrough_cell(
the trailing-slash-normalized proxy base URL and returns the
provider-mode env for the CLI subprocess.
"""
environ = env if env is not None else os.environ
base_url = environ.get(PROXY_BASE_URL_ENV)
api_key = environ.get(PROXY_API_KEY_ENV)
if not base_url or not api_key:
compat_result.set(
{
"status": "fail",
"error": (
f"missing required env: set {PROXY_BASE_URL_ENV} and "
f"{PROXY_API_KEY_ENV} to point at a running LiteLLM proxy"
),
}
)
pytest.fail(
f"{PROXY_BASE_URL_ENV} / {PROXY_API_KEY_ENV} not configured",
pytrace=False,
)
proxy_base = base_url.rstrip("/")
proxy = require_proxy(compat_result, env=env)
proxy_base = proxy.base_url.rstrip("/")
extra_env = dict(build_extra_env(proxy_base)) if build_extra_env else None
outcomes = run_models(
models=models,
prompt=prompt,
base_url=proxy_base + passthrough_base_path,
api_key=api_key,
api_key=proxy.api_key,
extra_env=extra_env,
)

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