mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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:
commit
a3768b489f
356 changed files with 12582 additions and 10825 deletions
2
.github/workflows/test-unit-proxy-db.yml
vendored
2
.github/workflows/test-unit-proxy-db.yml
vendored
|
|
@ -5,6 +5,8 @@ on:
|
|||
branches:
|
||||
- main
|
||||
- litellm_internal_staging
|
||||
- litellm_oss_staging
|
||||
- "litellm_**"
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
|
|
|
|||
6
.github/workflows/zizmor.yml
vendored
6
.github/workflows/zizmor.yml
vendored
|
|
@ -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 }}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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==",
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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 /
|
||||
|
|
|
|||
|
|
@ -98,4 +98,8 @@ spec:
|
|||
tolerations:
|
||||
{{- toYaml . | nindent 8 }}
|
||||
{{- end }}
|
||||
{{- with .Values.backend.topologySpreadConstraints }}
|
||||
topologySpreadConstraints:
|
||||
{{- toYaml . | nindent 8 }}
|
||||
{{- end }}
|
||||
{{- end }}
|
||||
|
|
|
|||
6
helm/litellm/templates/backend/poddisruptionbudget.yaml
Normal file
6
helm/litellm/templates/backend/poddisruptionbudget.yaml
Normal 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" .)) }}
|
||||
|
|
@ -100,4 +100,8 @@ spec:
|
|||
tolerations:
|
||||
{{- toYaml . | nindent 8 }}
|
||||
{{- end }}
|
||||
{{- with .Values.gateway.topologySpreadConstraints }}
|
||||
topologySpreadConstraints:
|
||||
{{- toYaml . | nindent 8 }}
|
||||
{{- end }}
|
||||
{{- end }}
|
||||
|
|
|
|||
6
helm/litellm/templates/gateway/poddisruptionbudget.yaml
Normal file
6
helm/litellm/templates/gateway/poddisruptionbudget.yaml
Normal 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" .)) }}
|
||||
|
|
@ -76,4 +76,8 @@ spec:
|
|||
tolerations:
|
||||
{{- toYaml . | nindent 8 }}
|
||||
{{- end }}
|
||||
{{- with .Values.ui.topologySpreadConstraints }}
|
||||
topologySpreadConstraints:
|
||||
{{- toYaml . | nindent 8 }}
|
||||
{{- end }}
|
||||
{{- end }}
|
||||
|
|
|
|||
6
helm/litellm/templates/ui/poddisruptionbudget.yaml
Normal file
6
helm/litellm/templates/ui/poddisruptionbudget.yaml
Normal 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" .)) }}
|
||||
188
helm/litellm/tests/pdb_topology_spread_tests.yaml
Normal file
188
helm/litellm/tests/pdb_topology_spread_tests.yaml
Normal 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
|
||||
|
|
@ -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: []
|
||||
|
|
|
|||
|
|
@ -0,0 +1,2 @@
|
|||
-- AlterTable
|
||||
ALTER TABLE "LiteLLM_MCPServerTable" ADD COLUMN IF NOT EXISTS "issuer" TEXT;
|
||||
|
|
@ -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?
|
||||
|
|
|
|||
|
|
@ -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==",
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"),
|
||||
|
|
|
|||
|
|
@ -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())
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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.",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
)
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
212
litellm/proxy/db/query_engine_reaper.py
Normal file
212
litellm/proxy/db/query_engine_reaper.py
Normal 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
|
||||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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 = (
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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 "
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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 = []
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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?
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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=())
|
||||
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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?
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
}
|
||||
}
|
||||
}
|
||||
]
|
||||
}
|
||||
|
|
@ -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
|
||||
|
|
@ -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"}
|
||||
}
|
||||
]
|
||||
}
|
||||
|
|
@ -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
|
||||
|
|
@ -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."
|
||||
)
|
||||
86
tests/e2e/claude_code/_compat_models.py
Normal file
86
tests/e2e/claude_code/_compat_models.py
Normal 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)
|
||||
)
|
||||
|
|
@ -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
|
||||
|
|
@ -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"]
|
||||
|
|
@ -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
|
||||
|
|
@ -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() == []
|
||||
|
|
@ -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",
|
||||
}
|
||||
|
|
@ -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)
|
||||
74
tests/e2e/claude_code/_env.py
Normal file
74
tests/e2e/claude_code/_env.py
Normal 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
|
||||
|
|
@ -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
Loading…
Add table
Reference in a new issue