diff --git a/.github/workflows/test-unit-proxy-db.yml b/.github/workflows/test-unit-proxy-db.yml index 2ac9a3b7c1c..b0ee56f5a5c 100644 --- a/.github/workflows/test-unit-proxy-db.yml +++ b/.github/workflows/test-unit-proxy-db.yml @@ -5,6 +5,8 @@ on: branches: - main - litellm_internal_staging + - litellm_oss_staging + - "litellm_**" permissions: contents: read diff --git a/.github/workflows/zizmor.yml b/.github/workflows/zizmor.yml index db79fe43038..df242e5a3b6 100644 --- a/.github/workflows/zizmor.yml +++ b/.github/workflows/zizmor.yml @@ -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 }} diff --git a/Dockerfile b/Dockerfile index bc0e6a5ca6f..581d1808f0a 100644 --- a/Dockerfile +++ b/Dockerfile @@ -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 diff --git a/docker/Dockerfile.database b/docker/Dockerfile.database index 4564ee403fe..868b6682276 100644 --- a/docker/Dockerfile.database +++ b/docker/Dockerfile.database @@ -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 diff --git a/docker/Dockerfile.non_root b/docker/Dockerfile.non_root index 1883e87be60..839f5da565c 100644 --- a/docker/Dockerfile.non_root +++ b/docker/Dockerfile.non_root @@ -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 diff --git a/enterprise/litellm_enterprise/enterprise_callbacks/pagerduty/pagerduty.py b/enterprise/litellm_enterprise/enterprise_callbacks/pagerduty/pagerduty.py index 12fdaeb6a81..f920aa7ac13 100644 --- a/enterprise/litellm_enterprise/enterprise_callbacks/pagerduty/pagerduty.py +++ b/enterprise/litellm_enterprise/enterprise_callbacks/pagerduty/pagerduty.py @@ -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, diff --git a/enterprise/pyproject.toml b/enterprise/pyproject.toml index 97571a4576d..04643b1ec33 100644 --- a/enterprise/pyproject.toml +++ b/enterprise/pyproject.toml @@ -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==", diff --git a/helm/litellm/templates/NOTES.txt b/helm/litellm/templates/NOTES.txt index 5b939fe480a..468cf621b32 100644 --- a/helm/litellm/templates/NOTES.txt +++ b/helm/litellm/templates/NOTES.txt @@ -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. diff --git a/helm/litellm/templates/_helpers.tpl b/helm/litellm/templates/_helpers.tpl index 043a0afc173..a0205c0a3a2 100644 --- a/helm/litellm/templates/_helpers.tpl +++ b/helm/litellm/templates/_helpers.tpl @@ -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 / diff --git a/helm/litellm/templates/backend/deployment.yaml b/helm/litellm/templates/backend/deployment.yaml index 9d056167fe1..892b84ff7d5 100644 --- a/helm/litellm/templates/backend/deployment.yaml +++ b/helm/litellm/templates/backend/deployment.yaml @@ -98,4 +98,8 @@ spec: tolerations: {{- toYaml . | nindent 8 }} {{- end }} + {{- with .Values.backend.topologySpreadConstraints }} + topologySpreadConstraints: + {{- toYaml . | nindent 8 }} + {{- end }} {{- end }} diff --git a/helm/litellm/templates/backend/poddisruptionbudget.yaml b/helm/litellm/templates/backend/poddisruptionbudget.yaml new file mode 100644 index 00000000000..02853ac879c --- /dev/null +++ b/helm/litellm/templates/backend/poddisruptionbudget.yaml @@ -0,0 +1,6 @@ +{{- include "litellm.pdb" (dict + "root" $ + "component" .Values.backend + "componentName" "backend" + "fullname" (include "litellm.backend.fullname" .) + "selectorLabels" (include "litellm.backend.selectorLabels" .)) }} diff --git a/helm/litellm/templates/gateway/deployment.yaml b/helm/litellm/templates/gateway/deployment.yaml index 4c80d784156..b2e22612905 100644 --- a/helm/litellm/templates/gateway/deployment.yaml +++ b/helm/litellm/templates/gateway/deployment.yaml @@ -100,4 +100,8 @@ spec: tolerations: {{- toYaml . | nindent 8 }} {{- end }} + {{- with .Values.gateway.topologySpreadConstraints }} + topologySpreadConstraints: + {{- toYaml . | nindent 8 }} + {{- end }} {{- end }} diff --git a/helm/litellm/templates/gateway/poddisruptionbudget.yaml b/helm/litellm/templates/gateway/poddisruptionbudget.yaml new file mode 100644 index 00000000000..15e89af17d7 --- /dev/null +++ b/helm/litellm/templates/gateway/poddisruptionbudget.yaml @@ -0,0 +1,6 @@ +{{- include "litellm.pdb" (dict + "root" $ + "component" .Values.gateway + "componentName" "gateway" + "fullname" (include "litellm.gateway.fullname" .) + "selectorLabels" (include "litellm.gateway.selectorLabels" .)) }} diff --git a/helm/litellm/templates/ui/deployment.yaml b/helm/litellm/templates/ui/deployment.yaml index 79e9a3e43bb..cd1f8c08fd4 100644 --- a/helm/litellm/templates/ui/deployment.yaml +++ b/helm/litellm/templates/ui/deployment.yaml @@ -76,4 +76,8 @@ spec: tolerations: {{- toYaml . | nindent 8 }} {{- end }} + {{- with .Values.ui.topologySpreadConstraints }} + topologySpreadConstraints: + {{- toYaml . | nindent 8 }} + {{- end }} {{- end }} diff --git a/helm/litellm/templates/ui/poddisruptionbudget.yaml b/helm/litellm/templates/ui/poddisruptionbudget.yaml new file mode 100644 index 00000000000..f7a3a694e9c --- /dev/null +++ b/helm/litellm/templates/ui/poddisruptionbudget.yaml @@ -0,0 +1,6 @@ +{{- include "litellm.pdb" (dict + "root" $ + "component" .Values.ui + "componentName" "ui" + "fullname" (include "litellm.ui.fullname" .) + "selectorLabels" (include "litellm.ui.selectorLabels" .)) }} diff --git a/helm/litellm/tests/pdb_topology_spread_tests.yaml b/helm/litellm/tests/pdb_topology_spread_tests.yaml new file mode 100644 index 00000000000..8aa05f3a969 --- /dev/null +++ b/helm/litellm/tests/pdb_topology_spread_tests.yaml @@ -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 diff --git a/helm/litellm/values.yaml b/helm/litellm/values.yaml index 74d02a25b7a..461935b2f50 100644 --- a/helm/litellm/values.yaml +++ b/helm/litellm/values.yaml @@ -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: [] diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260715000000_add_issuer_to_mcp_server_table/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260715000000_add_issuer_to_mcp_server_table/migration.sql new file mode 100644 index 00000000000..f7f23e6a55e --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260715000000_add_issuer_to_mcp_server_table/migration.sql @@ -0,0 +1,2 @@ +-- AlterTable +ALTER TABLE "LiteLLM_MCPServerTable" ADD COLUMN IF NOT EXISTS "issuer" TEXT; diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index a23cecc3911..f842bf13da9 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -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? diff --git a/litellm-proxy-extras/pyproject.toml b/litellm-proxy-extras/pyproject.toml index b67d9d8570a..cbb4109a652 100644 --- a/litellm-proxy-extras/pyproject.toml +++ b/litellm-proxy-extras/pyproject.toml @@ -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==", diff --git a/litellm/_redis.py b/litellm/_redis.py index 0b91cdabffc..fe5c5cdabe9 100644 --- a/litellm/_redis.py +++ b/litellm/_redis.py @@ -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) diff --git a/litellm/caching/dual_cache.py b/litellm/caching/dual_cache.py index be618815a53..0e3c93946fd 100644 --- a/litellm/caching/dual_cache.py +++ b/litellm/caching/dual_cache.py @@ -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: diff --git a/litellm/exceptions.py b/litellm/exceptions.py index aca3fb551cc..fd0a2afb3e8 100644 --- a/litellm/exceptions.py +++ b/litellm/exceptions.py @@ -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 diff --git a/litellm/integrations/langfuse/langfuse_otel.py b/litellm/integrations/langfuse/langfuse_otel.py index fc7c1b211c0..449457bd123 100644 --- a/litellm/integrations/langfuse/langfuse_otel.py +++ b/litellm/integrations/langfuse/langfuse_otel.py @@ -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, diff --git a/litellm/integrations/opentelemetry.py b/litellm/integrations/opentelemetry.py index de543fa042b..fea55cd1db4 100644 --- a/litellm/integrations/opentelemetry.py +++ b/litellm/integrations/opentelemetry.py @@ -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) diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index 268dcc3df78..9a0b4937fdb 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -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"), diff --git a/litellm/litellm_core_utils/streaming_chunk_builder_utils.py b/litellm/litellm_core_utils/streaming_chunk_builder_utils.py index 38bc68f2f78..d52d9849310 100644 --- a/litellm/litellm_core_utils/streaming_chunk_builder_utils.py +++ b/litellm/litellm_core_utils/streaming_chunk_builder_utils.py @@ -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()) diff --git a/litellm/litellm_core_utils/streaming_handler.py b/litellm/litellm_core_utils/streaming_handler.py index 128ba0bf3ab..f518cbaadea 100644 --- a/litellm/litellm_core_utils/streaming_handler.py +++ b/litellm/litellm_core_utils/streaming_handler.py @@ -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 diff --git a/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py b/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py index cd75eed2e6e..4b6617fbeac 100644 --- a/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py +++ b/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py @@ -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) diff --git a/litellm/llms/custom_httpx/aiohttp_transport.py b/litellm/llms/custom_httpx/aiohttp_transport.py index 3172d3667e1..adac6a1b276 100644 --- a/litellm/llms/custom_httpx/aiohttp_transport.py +++ b/litellm/llms/custom_httpx/aiohttp_transport.py @@ -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(): diff --git a/litellm/models/mcp_server.py b/litellm/models/mcp_server.py index af2efa822b0..23b26bd8e89 100644 --- a/litellm/models/mcp_server.py +++ b/litellm/models/mcp_server.py @@ -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 diff --git a/litellm/proxy/_experimental/mcp_server/db.py b/litellm/proxy/_experimental/mcp_server/db.py index 97cefb3f2cb..d55eb3ac014 100644 --- a/litellm/proxy/_experimental/mcp_server/db.py +++ b/litellm/proxy/_experimental/mcp_server/db.py @@ -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), diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 8b1d00c2855..115ff2e492c 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -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, diff --git a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py index 111fde86ea0..7ca4923337c 100644 --- a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py @@ -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, diff --git a/litellm/proxy/_experimental/mcp_server/semantic_tool_filter.py b/litellm/proxy/_experimental/mcp_server/semantic_tool_filter.py index e12c6cdbd56..37d3e0aedea 100644 --- a/litellm/proxy/_experimental/mcp_server/semantic_tool_filter.py +++ b/litellm/proxy/_experimental/mcp_server/semantic_tool_filter.py @@ -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) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 44550c80aa9..a95076d87c3 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -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.", diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 00f6d44e25a..fa354b8cccb 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -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, ) diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index b402212fb2e..4d07d4c043c 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -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 diff --git a/litellm/proxy/client/cli/README.md b/litellm/proxy/client/cli/README.md index be8905e11ca..17041751d15 100644 --- a/litellm/proxy/client/cli/README.md +++ b/litellm/proxy/client/cli/README.md @@ -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//scripts/install-cli.sh | \ +curl -fsSL https://raw.githubusercontent.com/BerriAI/litellm//scripts/install.sh | \ LITELLM_CLI_REF= 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 diff --git a/litellm/proxy/client/cli/commands/auth.py b/litellm/proxy/client/cli/commands/auth.py index 785d3b1e37b..61495403407 100644 --- a/litellm/proxy/client/cli/commands/auth.py +++ b/litellm/proxy/client/cli/commands/auth.py @@ -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") diff --git a/litellm/proxy/client/cli/commands/autoroute/commands.py b/litellm/proxy/client/cli/commands/autoroute/commands.py index 428daeefa60..161907f5b27 100644 --- a/litellm/proxy/client/cli/commands/autoroute/commands.py +++ b/litellm/proxy/client/cli/commands/autoroute/commands.py @@ -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//scripts/install.sh | " + "LITELLM_CLI_REF= sh`." + ) + try: existing_pid = read_pid_record() except UpError as e: diff --git a/litellm/proxy/client/cli/commands/autoroute/config.py b/litellm/proxy/client/cli/commands/autoroute/config.py index fe2c4e467e7..2d760ef0f8a 100644 --- a/litellm/proxy/client/cli/commands/autoroute/config.py +++ b/litellm/proxy/client/cli/commands/autoroute/config.py @@ -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", diff --git a/litellm/proxy/client/cli/commands/autoroute/process.py b/litellm/proxy/client/cli/commands/autoroute/process.py index 86fa584a223..712f2eed2da 100644 --- a/litellm/proxy/client/cli/commands/autoroute/process.py +++ b/litellm/proxy/client/cli/commands/autoroute/process.py @@ -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", diff --git a/litellm/proxy/client/cli/commands/autoroute/wizard.py b/litellm/proxy/client/cli/commands/autoroute/wizard.py index 89b53d8f9fd..60696fb2e7e 100644 --- a/litellm/proxy/client/cli/commands/autoroute/wizard.py +++ b/litellm/proxy/client/cli/commands/autoroute/wizard.py @@ -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) diff --git a/litellm/proxy/client/cli/commands/chat.py b/litellm/proxy/client/cli/commands/chat.py index 9d1a1aa7a30..d78feb84bd6 100644 --- a/litellm/proxy/client/cli/commands/chat.py +++ b/litellm/proxy/client/cli/commands/chat.py @@ -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", ) ) diff --git a/litellm/proxy/client/cli/commands/encryption.py b/litellm/proxy/client/cli/commands/encryption.py index f67c9746fa9..4b460bac19c 100644 --- a/litellm/proxy/client/cli/commands/encryption.py +++ b/litellm/proxy/client/cli/commands/encryption.py @@ -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 diff --git a/litellm/proxy/client/cli/commands/keys.py b/litellm/proxy/client/cli/commands/keys.py index 45e27442708..afbaa3702c1 100644 --- a/litellm/proxy/client/cli/commands/keys.py +++ b/litellm/proxy/client/cli/commands/keys.py @@ -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 diff --git a/litellm/proxy/client/cli/commands/teams.py b/litellm/proxy/client/cli/commands/teams.py index c0d45544b11..b0ccdc8f9bf 100644 --- a/litellm/proxy/client/cli/commands/teams.py +++ b/litellm/proxy/client/cli/commands/teams.py @@ -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: diff --git a/litellm/proxy/client/cli/interface.py b/litellm/proxy/client/cli/interface.py index 33f1f4a4480..e953742f412 100644 --- a/litellm/proxy/client/cli/interface.py +++ b/litellm/proxy/client/cli/interface.py @@ -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 diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index ca6875ca800..54a4c2dad91 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -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] diff --git a/litellm/proxy/db/query_engine_reaper.py b/litellm/proxy/db/query_engine_reaper.py new file mode 100644 index 00000000000..0e5f0e68910 --- /dev/null +++ b/litellm/proxy/db/query_engine_reaper.py @@ -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 diff --git a/litellm/proxy/guardrails/guardrail_hooks/model_armor/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/model_armor/__init__.py index 7398e8defea..5e62ab96f0c 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/model_armor/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/model_armor/__init__.py @@ -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) diff --git a/litellm/proxy/guardrails/guardrail_hooks/model_armor/file_scanning.py b/litellm/proxy/guardrails/guardrail_hooks/model_armor/file_scanning.py index 0bc6e67eb35..b879f0d29c7 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/model_armor/file_scanning.py +++ b/litellm/proxy/guardrails/guardrail_hooks/model_armor/file_scanning.py @@ -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"] diff --git a/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py b/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py index 3ca63a1e287..32a3cebfca0 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py +++ b/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py @@ -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 = ( diff --git a/litellm/proxy/guardrails/guardrail_registry.py b/litellm/proxy/guardrails/guardrail_registry.py index e9eee3a1a8a..1f2c9e0c182 100644 --- a/litellm/proxy/guardrails/guardrail_registry.py +++ b/litellm/proxy/guardrails/guardrail_registry.py @@ -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 diff --git a/litellm/proxy/guardrails/usage_endpoints.py b/litellm/proxy/guardrails/usage_endpoints.py index e03bdbb95d2..f56b22ddd49 100644 --- a/litellm/proxy/guardrails/usage_endpoints.py +++ b/litellm/proxy/guardrails/usage_endpoints.py @@ -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) diff --git a/litellm/proxy/hooks/mcp_semantic_filter/hook.py b/litellm/proxy/hooks/mcp_semantic_filter/hook.py index bad6ef44ccd..5f1d061c7cb 100644 --- a/litellm/proxy/hooks/mcp_semantic_filter/hook.py +++ b/litellm/proxy/hooks/mcp_semantic_filter/hook.py @@ -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 " diff --git a/litellm/proxy/hooks/model_max_budget_limiter.py b/litellm/proxy/hooks/model_max_budget_limiter.py index 803fe64c193..e8cf5fbc718 100644 --- a/litellm/proxy/hooks/model_max_budget_limiter.py +++ b/litellm/proxy/hooks/model_max_budget_limiter.py @@ -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 diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index afd8a437cc1..15d1876e5a2 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -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, diff --git a/litellm/proxy/management_endpoints/internal_user_endpoints.py b/litellm/proxy/management_endpoints/internal_user_endpoints.py index ccd15a68437..f741783134e 100644 --- a/litellm/proxy/management_endpoints/internal_user_endpoints.py +++ b/litellm/proxy/management_endpoints/internal_user_endpoints.py @@ -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( diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index 288282dd08b..d920ee474cc 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -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, diff --git a/litellm/proxy/management_endpoints/scim/scim_transformations.py b/litellm/proxy/management_endpoints/scim/scim_transformations.py index cc3f18f593d..65651752944 100644 --- a/litellm/proxy/management_endpoints/scim/scim_transformations.py +++ b/litellm/proxy/management_endpoints/scim/scim_transformations.py @@ -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: """ diff --git a/litellm/proxy/management_endpoints/scim/scim_v2.py b/litellm/proxy/management_endpoints/scim/scim_v2.py index 808b80cd1ed..fa123b7d76c 100644 --- a/litellm/proxy/management_endpoints/scim/scim_v2.py +++ b/litellm/proxy/management_endpoints/scim/scim_v2.py @@ -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: diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index 74109ef8358..342c356561a 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -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") diff --git a/litellm/proxy/proxy_cli.py b/litellm/proxy/proxy_cli.py index 74ec0cc8700..9bed3657b20 100644 --- a/litellm/proxy/proxy_cli.py +++ b/litellm/proxy/proxy_cli.py @@ -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, diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 6661474d215..dcde9a27ec0 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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 = [] diff --git a/litellm/proxy/route_llm_request.py b/litellm/proxy/route_llm_request.py index 94871ff072d..25fa0819930 100644 --- a/litellm/proxy/route_llm_request.py +++ b/litellm/proxy/route_llm_request.py @@ -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) diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index a23cecc3911..f842bf13da9 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -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? diff --git a/litellm/proxy/spend_tracking/budget_reservation.py b/litellm/proxy/spend_tracking/budget_reservation.py index e1a093a4a48..80fd8a1594e 100644 --- a/litellm/proxy/spend_tracking/budget_reservation.py +++ b/litellm/proxy/spend_tracking/budget_reservation.py @@ -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, ) diff --git a/litellm/responses/streaming_iterator.py b/litellm/responses/streaming_iterator.py index eb78e6f9c8d..357a7ecefe6 100644 --- a/litellm/responses/streaming_iterator.py +++ b/litellm/responses/streaming_iterator.py @@ -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 diff --git a/litellm/router.py b/litellm/router.py index f408d030b8b..78e156801f8 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -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, diff --git a/litellm/router_strategy/tag_based_routing.py b/litellm/router_strategy/tag_based_routing.py index 6ca4e1de322..710c2199107 100644 --- a/litellm/router_strategy/tag_based_routing.py +++ b/litellm/router_strategy/tag_based_routing.py @@ -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: diff --git a/litellm/types/guardrails.py b/litellm/types/guardrails.py index f3410935ec7..3dda4e3990c 100644 --- a/litellm/types/guardrails.py +++ b/litellm/types/guardrails.py @@ -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, diff --git a/litellm/types/litellm_core_utils/streaming_chunk_builder_utils.py b/litellm/types/litellm_core_utils/streaming_chunk_builder_utils.py index a1f89dac5cf..c9f9d4e6baa 100644 --- a/litellm/types/litellm_core_utils/streaming_chunk_builder_utils.py +++ b/litellm/types/litellm_core_utils/streaming_chunk_builder_utils.py @@ -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] diff --git a/litellm/types/mcp_server/mcp_server_manager.py b/litellm/types/mcp_server/mcp_server_manager.py index 801436c774a..d0d8cc4cb28 100644 --- a/litellm/types/mcp_server/mcp_server_manager.py +++ b/litellm/types/mcp_server/mcp_server_manager.py @@ -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 diff --git a/litellm/types/proxy/management_endpoints/scim_v2.py b/litellm/types/proxy/management_endpoints/scim_v2.py index 3b1ea8f572e..8c434481975 100644 --- a/litellm/types/proxy/management_endpoints/scim_v2.py +++ b/litellm/types/proxy/management_endpoints/scim_v2.py @@ -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 diff --git a/litellm/types/router.py b/litellm/types/router.py index d62c613bf57..69a8ca9f19e 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -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=()) diff --git a/litellm/types/utils.py b/litellm/types/utils.py index 6487a8aa33f..88b3a39844f 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -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", diff --git a/pyproject.toml b/pyproject.toml index 2c19bf64b4f..45ae3c179d5 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -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", diff --git a/schema.prisma b/schema.prisma index a23cecc3911..f842bf13da9 100644 --- a/schema.prisma +++ b/schema.prisma @@ -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? diff --git a/scripts/install.sh b/scripts/install.sh index 06e6249c9ba..213f8a7b440 100755 --- a/scripts/install.sh +++ b/scripts/install.sh @@ -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//scripts/install.sh | \ +# LITELLM_CLI_REF= 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, diff --git a/tests/e2e/CLAUDE.md b/tests/e2e/CLAUDE.md index f0d283629b0..5d16761ac44 100644 --- a/tests/e2e/CLAUDE.md +++ b/tests/e2e/CLAUDE.md @@ -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... 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 diff --git a/tests/e2e/claude_code/_basic_messaging.py b/tests/e2e/claude_code/_basic_messaging.py index f6b82a38f6a..7c581cc5e38 100644 --- a/tests/e2e/claude_code/_basic_messaging.py +++ b/tests/e2e/claude_code/_basic_messaging.py @@ -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_*` × 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, diff --git a/tests/e2e/claude_code/_builder_unit_tests/__init__.py b/tests/e2e/claude_code/_builder_unit_tests/__init__.py deleted file mode 100644 index e69de29bb2d..00000000000 diff --git a/tests/e2e/claude_code/_builder_unit_tests/fixtures/expected_matrix.json b/tests/e2e/claude_code/_builder_unit_tests/fixtures/expected_matrix.json deleted file mode 100644 index d3aca0142dc..00000000000 --- a/tests/e2e/claude_code/_builder_unit_tests/fixtures/expected_matrix.json +++ /dev/null @@ -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" - } - } - } - ] -} diff --git a/tests/e2e/claude_code/_builder_unit_tests/fixtures/manifest.yaml b/tests/e2e/claude_code/_builder_unit_tests/fixtures/manifest.yaml deleted file mode 100644 index e88bdc6ddf5..00000000000 --- a/tests/e2e/claude_code/_builder_unit_tests/fixtures/manifest.yaml +++ /dev/null @@ -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 diff --git a/tests/e2e/claude_code/_builder_unit_tests/fixtures/results.json b/tests/e2e/claude_code/_builder_unit_tests/fixtures/results.json deleted file mode 100644 index a01540c394f..00000000000 --- a/tests/e2e/claude_code/_builder_unit_tests/fixtures/results.json +++ /dev/null @@ -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"} - } - ] -} diff --git a/tests/e2e/claude_code/_builder_unit_tests/test_matrix_builder.py b/tests/e2e/claude_code/_builder_unit_tests/test_matrix_builder.py deleted file mode 100644 index 9ddbdd29846..00000000000 --- a/tests/e2e/claude_code/_builder_unit_tests/test_matrix_builder.py +++ /dev/null @@ -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 diff --git a/tests/e2e/claude_code/_builder_unit_tests/test_v0_layout.py b/tests/e2e/claude_code/_builder_unit_tests/test_v0_layout.py deleted file mode 100644 index a3569ebdb49..00000000000 --- a/tests/e2e/claude_code/_builder_unit_tests/test_v0_layout.py +++ /dev/null @@ -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_.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." - ) diff --git a/tests/e2e/claude_code/_compat_models.py b/tests/e2e/claude_code/_compat_models.py new file mode 100644 index 00000000000..23c9c82e596 --- /dev/null +++ b/tests/e2e/claude_code/_compat_models.py @@ -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) + ) diff --git a/tests/e2e/claude_code/_driver_unit_tests/__init__.py b/tests/e2e/claude_code/_driver_unit_tests/__init__.py deleted file mode 100644 index e69de29bb2d..00000000000 diff --git a/tests/e2e/claude_code/_driver_unit_tests/conftest.py b/tests/e2e/claude_code/_driver_unit_tests/conftest.py deleted file mode 100644 index bfeaa57c736..00000000000 --- a/tests/e2e/claude_code/_driver_unit_tests/conftest.py +++ /dev/null @@ -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 diff --git a/tests/e2e/claude_code/_driver_unit_tests/test_basic_messaging.py b/tests/e2e/claude_code/_driver_unit_tests/test_basic_messaging.py deleted file mode 100644 index 018121a8e5c..00000000000 --- a/tests/e2e/claude_code/_driver_unit_tests/test_basic_messaging.py +++ /dev/null @@ -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"] diff --git a/tests/e2e/claude_code/_driver_unit_tests/test_cli_driver.py b/tests/e2e/claude_code/_driver_unit_tests/test_cli_driver.py deleted file mode 100644 index f1a0534906e..00000000000 --- a/tests/e2e/claude_code/_driver_unit_tests/test_cli_driver.py +++ /dev/null @@ -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 `), - 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//`), 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 diff --git a/tests/e2e/claude_code/_driver_unit_tests/test_compat_result.py b/tests/e2e/claude_code/_driver_unit_tests/test_compat_result.py deleted file mode 100644 index b3a904946d3..00000000000 --- a/tests/e2e/claude_code/_driver_unit_tests/test_compat_result.py +++ /dev/null @@ -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() == [] diff --git a/tests/e2e/claude_code/_driver_unit_tests/test_passthrough.py b/tests/e2e/claude_code/_driver_unit_tests/test_passthrough.py deleted file mode 100644 index 2d17a84d418..00000000000 --- a/tests/e2e/claude_code/_driver_unit_tests/test_passthrough.py +++ /dev/null @@ -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", - } diff --git a/tests/e2e/claude_code/_driver_unit_tests/test_rate_limiter.py b/tests/e2e/claude_code/_driver_unit_tests/test_rate_limiter.py deleted file mode 100644 index 92907eda3c4..00000000000 --- a/tests/e2e/claude_code/_driver_unit_tests/test_rate_limiter.py +++ /dev/null @@ -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) diff --git a/tests/e2e/claude_code/_env.py b/tests/e2e/claude_code/_env.py new file mode 100644 index 00000000000..4f93cb57fdc --- /dev/null +++ b/tests/e2e/claude_code/_env.py @@ -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 diff --git a/tests/e2e/claude_code/_passthrough.py b/tests/e2e/claude_code/_passthrough.py index 24b6d694d3f..be7a475dff7 100644 --- a/tests/e2e/claude_code/_passthrough.py +++ b/tests/e2e/claude_code/_passthrough.py @@ -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, ) diff --git a/tests/e2e/claude_code/_pr_gate_unit_tests/__init__.py b/tests/e2e/claude_code/_pr_gate_unit_tests/__init__.py deleted file mode 100644 index e69de29bb2d..00000000000 diff --git a/tests/e2e/claude_code/_pr_gate_unit_tests/test_bash_tool_restrictions.py b/tests/e2e/claude_code/_pr_gate_unit_tests/test_bash_tool_restrictions.py deleted file mode 100644 index d698131670a..00000000000 --- a/tests/e2e/claude_code/_pr_gate_unit_tests/test_bash_tool_restrictions.py +++ /dev/null @@ -1,147 +0,0 @@ -"""Pin tests for the `Bash`-using compat cells. - -Every cell that passes `--allowed-tools Bash` to the `claude` CLI is -giving a model-controlled response the ability to run host commands. -On the PR-gate CircleCI machine executor, those commands have access -to the Docker socket and can read `docker inspect compat-proxy` to -recover the provider credentials living inside the proxy container. - -To narrow that surface, every Bash-using cell must: - -1. Restrict the allow rule to the *exact* command `Bash(echo pong)` so - a compromised provider response cannot turn `Bash` into arbitrary - host execution by emitting a `tool_use` with a different command. - -2. Pair it with `--permission-mode dontAsk` so anything not matching - an allow rule is auto-denied instead of prompting (which would - abort the CLI in headless mode, but auto-denial is the explicit - contract). - -These restrictions are enforced by the `claude` CLI, not by the -model — see https://code.claude.com/docs/en/permissions for the -permission-rule precedence (`deny` → `ask` → `allow`). - -This test scans every cell under the three Bash-using feature -directories (`tool_use`, `tool_use_streaming`, `thinking_with_tool_use`) -and pins both requirements so a future test refactor cannot silently -revert any cell to the broad `Bash` allow that was originally -flagged by Veria. -""" - -from __future__ import annotations - -from pathlib import Path -from typing import Iterable - -import pytest - -REPO_ROOT = Path(__file__).resolve().parents[4] -CLAUDE_CODE_DIR = REPO_ROOT / "tests" / "e2e" / "claude_code" - -# Feature directories whose cells drive the `Bash` built-in tool. Add -# new entries here when a new Bash-using feature is added; the test -# fails loudly for any unhandled directory so we never miss one by -# silent omission. -BASH_FEATURE_DIRS = ( - "tool_use", - "tool_use_streaming", - "thinking_with_tool_use", -) - - -def _bash_cells() -> Iterable[Path]: - for feature in BASH_FEATURE_DIRS: - feature_dir = CLAUDE_CODE_DIR / feature - assert feature_dir.is_dir(), ( - f"{feature_dir} is missing — BASH_FEATURE_DIRS is out of sync " - f"with the layout under tests/e2e/claude_code/." - ) - for path in sorted(feature_dir.glob("test_*.py")): - yield path - - -def _has_bare_bash_token(text: str) -> bool: - """Return True if `text` contains a `"Bash"` token outside the - `"Bash(echo pong)"` allow rule. - - Extracted as a pure helper so the negative path can be unit-tested - directly. Without it, the previous structure of this assertion was - `'"Bash"' not in text or '"Bash(echo pong)"' in text`, which - short-circuits to True any time the allow rule is present and lets - a stray bare `"Bash"` slip through the security pin undetected. - """ - return '"Bash"' in text.replace('"Bash(echo pong)"', "") - - -@pytest.mark.parametrize( - "cell", list(_bash_cells()), ids=lambda p: str(p.relative_to(REPO_ROOT)) -) -def test_bash_allow_rule_is_pinned_to_exact_echo_pong(cell: Path) -> None: - """The cell must pass `Bash(echo pong)` as the allow rule, not the - unrestricted `Bash` value that was originally flagged.""" - text = cell.read_text() - assert '"Bash(echo pong)"' in text, ( - f"{cell.relative_to(REPO_ROOT)} must restrict `--allowed-tools` to " - f'`Bash(echo pong)` (exact-match pattern). Unrestricted `"Bash"` ' - f"grants arbitrary host command execution to model-controlled " - f"tool_use blocks, which can read `docker inspect compat-proxy` " - f"to exfiltrate provider credentials from the proxy container." - ) - # The only place `"Bash"` (the bare token, surrounded by quotes - # exactly as it would appear in `--allowed-tools` lists) is allowed - # to appear is *inside* the exact-match `"Bash(echo pong)"` rule. - # `_has_bare_bash_token` keeps that scan independent of the first - # assertion — otherwise `'"Bash"' not in text or '"Bash(echo pong)"' - # in text` short-circuits to True and lets a stray bare `"Bash"` - # slip through silently. - assert not _has_bare_bash_token(text), ( - f"{cell.relative_to(REPO_ROOT)} still references the unrestricted " - f'`"Bash"` value outside the `"Bash(echo pong)"` allow rule — ' - f"sweep it out before merging." - ) - - -def test_has_bare_bash_token_flags_unrestricted_value(): - """A file that allows the bare `"Bash"` token alongside the - exact-match rule must be flagged. Without this guard the security - pin reverts to the dead-code `or` it had originally, which let - arbitrary host commands through under the noise of a passing test. - """ - text = '--allowed-tools "Bash" "Bash(echo pong)"' - assert _has_bare_bash_token(text) - - -def test_has_bare_bash_token_accepts_only_exact_match(): - """The standard pattern — only the exact-match allow rule, no bare - `"Bash"` — must be accepted. This is the shape every Bash-using - cell in the suite is required to take. - """ - text = '--allowed-tools "Bash(echo pong)" --permission-mode "dontAsk"' - assert not _has_bare_bash_token(text) - - -def test_has_bare_bash_token_ignores_unrelated_substrings(): - """`Bash(echo pong)` is the only allowed shape; substrings like - `BashTool` or `Bashing` are unrelated identifiers and must not be - confused with the bare `"Bash"` token (i.e. the exact quoted - string `"Bash"`).""" - text = "BashTool helper used by the bashing harness" - assert not _has_bare_bash_token(text) - - -@pytest.mark.parametrize( - "cell", list(_bash_cells()), ids=lambda p: str(p.relative_to(REPO_ROOT)) -) -def test_bash_cell_uses_dontask_permission_mode(cell: Path) -> None: - """The cell must pair the allow rule with `--permission-mode dontAsk` - so tool calls that don't match the allow rule are auto-denied (as - opposed to defaulting to "ask", which in headless mode would - succeed without ever surfacing the security issue).""" - text = cell.read_text() - assert '"--permission-mode"' in text and '"dontAsk"' in text, ( - f"{cell.relative_to(REPO_ROOT)} must pass `--permission-mode dontAsk` " - f"alongside the `Bash(echo pong)` allow rule. Without dontAsk, " - f"commands outside the allow rule fall back to the default ask-" - f"mode behavior, which in `--print` (headless) mode is non-" - f"interactive — defeating the explicit-allow contract." - ) diff --git a/tests/e2e/claude_code/_pr_gate_unit_tests/test_pr_gate_version_resolver.py b/tests/e2e/claude_code/_pr_gate_unit_tests/test_pr_gate_version_resolver.py deleted file mode 100644 index 5c516da81c5..00000000000 --- a/tests/e2e/claude_code/_pr_gate_unit_tests/test_pr_gate_version_resolver.py +++ /dev/null @@ -1,164 +0,0 @@ -"""Unit tests for the Claude Code PR-Gate Version Resolver. - -The resolver picks the newest `@anthropic-ai/claude-code` version whose -publish timestamp is at least 3 days old. The 3-day window is a security -review buffer: a malicious or broken Claude Code release that slipped -through the npm publish process gets at least 72 hours to be detected -before it can land in the LiteLLM PR gate. - -The unit tests inject npm metadata directly (no network) and a fixed -`as_of` clock (no real time), so they run anywhere and never flake on -the wall clock or registry availability. -""" - -from __future__ import annotations - -from datetime import datetime, timedelta, timezone - -import pytest - -from claude_code.pr_gate_version_resolver import ( - NoEligibleVersionError, - resolve_pr_gate_version, -) - - -def _t(iso: str) -> str: - """Helper for readable ISO-8601 publish timestamps in fixtures.""" - return iso - - -# A clock fixed at a moment well after every fixture publish time below. -NOW = datetime(2026, 4, 25, 12, 0, 0, tzinfo=timezone.utc) - - -def _metadata_with_times(times: dict) -> dict: - """Shape an npm `packument`-like dict with the `time` field populated. - - The npm registry response includes `time.created` / `time.modified` - keys alongside per-version timestamps; the resolver must skip those. - """ - return { - "name": "@anthropic-ai/claude-code", - "time": { - "created": _t("2024-01-01T00:00:00.000Z"), - "modified": _t("2026-04-25T00:00:00.000Z"), - **times, - }, - } - - -def test_picks_newest_version_at_least_three_days_old(): - metadata = _metadata_with_times( - { - "2.1.118": _t("2026-04-15T10:00:00.000Z"), - "2.1.119": _t("2026-04-21T10:00:00.000Z"), # 4d 2h old - "2.1.120": _t("2026-04-23T10:00:00.000Z"), # 2d 2h old — too new - "2.1.121": _t("2026-04-25T11:00:00.000Z"), # 1h old — too new - } - ) - assert resolve_pr_gate_version(metadata=metadata, as_of=NOW) == "2.1.119" - - -def test_skips_created_and_modified_meta_keys(): - """`time` contains `created` / `modified` non-version entries — must be ignored.""" - metadata = { - "name": "@anthropic-ai/claude-code", - "time": { - "created": _t("2024-01-01T00:00:00.000Z"), - "modified": _t("2026-04-25T00:00:00.000Z"), - "2.0.0": _t("2026-04-10T00:00:00.000Z"), - }, - } - assert resolve_pr_gate_version(metadata=metadata, as_of=NOW) == "2.0.0" - - -def test_min_age_boundary_is_inclusive(): - """A version published exactly 3 days ago is eligible (>= cutoff).""" - three_days_ago = NOW - timedelta(days=3) - metadata = _metadata_with_times( - { - "2.1.0": three_days_ago.isoformat().replace("+00:00", "Z"), - } - ) - assert resolve_pr_gate_version(metadata=metadata, as_of=NOW) == "2.1.0" - - -def test_raises_when_every_version_is_too_new(): - metadata = _metadata_with_times( - { - "2.1.121": _t("2026-04-25T08:00:00.000Z"), # 4h old - "2.1.120": _t("2026-04-24T10:00:00.000Z"), # ~26h old - } - ) - with pytest.raises(NoEligibleVersionError): - resolve_pr_gate_version(metadata=metadata, as_of=NOW) - - -def test_raises_when_metadata_has_no_versions(): - metadata = {"name": "@anthropic-ai/claude-code", "time": {}} - with pytest.raises(NoEligibleVersionError): - resolve_pr_gate_version(metadata=metadata, as_of=NOW) - - -def test_picks_latest_publish_time_not_largest_semver(): - """If a patch is published to an old major after a newer release, - "newest" is by publish time, not semver string ordering.""" - metadata = _metadata_with_times( - { - "1.9.99": _t("2026-04-22T10:00:00.000Z"), # patched recently — wins - "2.0.0": _t("2026-03-01T10:00:00.000Z"), # older publish - } - ) - assert resolve_pr_gate_version(metadata=metadata, as_of=NOW) == "1.9.99" - - -def test_uses_custom_min_age(): - metadata = _metadata_with_times( - { - "1.0.0": _t("2026-04-23T10:00:00.000Z"), # 2d 2h old - "0.9.0": _t("2026-04-10T10:00:00.000Z"), # 15d old - } - ) - # min_age = 5 days disqualifies 1.0.0 - out = resolve_pr_gate_version( - metadata=metadata, as_of=NOW, min_age=timedelta(days=5) - ) - assert out == "0.9.0" - - -def test_excludes_prerelease_versions(): - """Pre-release tags (1.0.0-alpha.1, 2.0.0-rc.1, etc.) must never win, - even if their publish timestamp is the newest eligible one.""" - metadata = _metadata_with_times( - { - "2.1.119": _t("2026-04-21T10:00:00.000Z"), # stable, 4d old - "2.2.0-alpha.1": _t("2026-04-22T10:00:00.000Z"), # newer publish - "2.2.0-rc.1": _t("2026-04-22T11:00:00.000Z"), # newest publish - "3.0.0-beta": _t("2026-04-22T12:00:00.000Z"), # newest publish - } - ) - assert resolve_pr_gate_version(metadata=metadata, as_of=NOW) == "2.1.119" - - -def test_raises_when_only_prereleases_are_eligible(): - metadata = _metadata_with_times( - { - "2.2.0-alpha.1": _t("2026-04-22T10:00:00.000Z"), - "2.2.0-rc.1": _t("2026-04-22T11:00:00.000Z"), - } - ) - with pytest.raises(NoEligibleVersionError): - resolve_pr_gate_version(metadata=metadata, as_of=NOW) - - -def test_resolver_uses_fetcher_when_metadata_not_provided(): - captured = {} - - def fake_fetch(package_name: str) -> dict: - captured["package"] = package_name - return _metadata_with_times({"3.0.0": _t("2026-04-10T10:00:00.000Z")}) - - out = resolve_pr_gate_version(as_of=NOW, fetcher=fake_fetch) - assert out == "3.0.0" - assert captured["package"] == "@anthropic-ai/claude-code" diff --git a/tests/e2e/claude_code/_publisher_unit_tests/__init__.py b/tests/e2e/claude_code/_publisher_unit_tests/__init__.py deleted file mode 100644 index e69de29bb2d..00000000000 diff --git a/tests/e2e/claude_code/_publisher_unit_tests/test_run_daily_pytest_scrubs_env.py b/tests/e2e/claude_code/_publisher_unit_tests/test_run_daily_pytest_scrubs_env.py deleted file mode 100644 index 418d308674b..00000000000 --- a/tests/e2e/claude_code/_publisher_unit_tests/test_run_daily_pytest_scrubs_env.py +++ /dev/null @@ -1,95 +0,0 @@ -"""Pin: the cron `pytest` invocation must run under `env -i`. - -The systemd service `litellm-compat-matrix.service` loads provider -credentials (`ANTHROPIC_API_KEY`, `AWS_BEARER_TOKEN_BEDROCK`, -`AZURE_FOUNDRY_API_KEY`, `VERTEXAI_*`) and the agent-shin GitHub token -(`AGENT_SHIN_GITHUB_TOKEN`) into `run_daily.sh`'s environment from -`/etc/litellm-compat-matrix.env`. Pytest only needs to talk to the -loopback proxy at `127.0.0.1:${PROXY_PORT}` and has no legitimate reason -to see provider creds in its own `os.environ`. Leaving them in would -let a test under `tests/e2e/claude_code/` read them via `os.environ` and -exfiltrate them, and would also let a model-directed `Read` tool call -during a PDF/vision cell reach `/proc//environ`. The -PR-gate's pytest step in `.circleci/config.yml` already runs under -`env -i`; this pin enforces the same scrub on the cron path. -""" - -from __future__ import annotations - -from pathlib import Path - -REPO_ROOT = Path(__file__).resolve().parents[4] -RUN_DAILY = REPO_ROOT / "tests" / "e2e" / "claude_code" / "cron_vm" / "run_daily.sh" - - -def _pytest_invocation_block() -> str: - """Return only the executable lines around the pytest invocation. - - Comment text in run_daily.sh explains *why* certain credential - names must not appear, so a naïve substring scan over the whole - region would false-positive on the rationale itself. Strip lines - whose first non-space character is `#`. - """ - body = RUN_DAILY.read_text() - start = body.index('log "running pytest"') - end = body.index("PYTEST_EXIT=$?", start) - return "\n".join( - line for line in body[start:end].splitlines() - if line.lstrip()[:1] != "#" - ) - - -def test_pytest_invocation_wraps_in_env_i() -> None: - block = _pytest_invocation_block() - assert "env -i" in block, ( - "run_daily.sh: the pytest invocation must run under `env -i` so " - "PR-controlled test code under tests/e2e/claude_code/ cannot read " - "provider/agent-shin credentials out of the systemd service " - "environment, and so a model-directed `Read` tool call cannot " - "reach /proc//environ to pull them out." - ) - assert block.index("env -i") < block.index('"${WORKTREE_UV}" run pytest'), ( - "run_daily.sh: `env -i` must precede the pytest invocation; " - "otherwise pytest inherits the full credential-bearing env." - ) - - -def test_pytest_invocation_env_i_excludes_provider_secrets() -> None: - block = _pytest_invocation_block() - for forbidden in ( - "ANTHROPIC_API_KEY", - "AWS_BEARER_TOKEN_BEDROCK", - "AWS_ACCESS_KEY_ID", - "AWS_SECRET_ACCESS_KEY", - "VERTEXAI_CREDENTIALS", - "VERTEXAI_PROJECT", - "VERTEXAI_LOCATION", - "AZURE_FOUNDRY_API_KEY", - "AZURE_FOUNDRY_API_BASE", - "GITHUB_TOKEN", - "AGENT_SHIN_GITHUB_TOKEN", - ): - assert forbidden not in block, ( - f"run_daily.sh: the pytest-step `env -i` allowlist must not " - f"pass {forbidden} through. Found it inside the pytest " - f"invocation block." - ) - - -def test_pytest_invocation_passes_proxy_url_and_key_explicitly() -> None: - block = _pytest_invocation_block() - assert "LITELLM_PROXY_BASE_URL=" in block, ( - "run_daily.sh: the pytest `env -i` block must still pass " - "LITELLM_PROXY_BASE_URL so the test suite knows where to find " - "the loopback proxy." - ) - assert "LITELLM_PROXY_API_KEY=" in block, ( - "run_daily.sh: the pytest `env -i` block must still pass " - "LITELLM_PROXY_API_KEY so the test suite can authenticate to " - "the loopback proxy." - ) - assert "COMPAT_RESULTS_PATH=" in block, ( - "run_daily.sh: the pytest `env -i` block must still pass " - "COMPAT_RESULTS_PATH so the conftest writes the per-cell " - "tagged-union artifact to the script-managed path." - ) diff --git a/tests/e2e/claude_code/_publisher_unit_tests/test_run_daily_release_pagination.py b/tests/e2e/claude_code/_publisher_unit_tests/test_run_daily_release_pagination.py deleted file mode 100644 index fc733845ba3..00000000000 --- a/tests/e2e/claude_code/_publisher_unit_tests/test_run_daily_release_pagination.py +++ /dev/null @@ -1,289 +0,0 @@ -"""Regression tests for the GitHub release pagination in `run_daily.sh`. - -The cron job resolves "newest LiteLLM v*-stable" via the GitHub Releases -API. A previous version of the loop broke as soon as the current page -contained ANY v*-stable tag. The Releases endpoint orders by -`created_at`, NOT by semver, so a backport on an older series cut today -(e.g. v1.80.1-stable) can land on an earlier page than a higher-version -release cut two weeks ago (e.g. v1.83.0-stable). The early-break would -silently pin the cron to a stale tag because the higher-version release -on a later page never made it into the merged set the final `sort_by` -consumed. - -These tests pin two things: - - 1. The buggy early-break-on-first-stable pattern must not return. - 2. The loop still terminates early on the standard "empty page" guard - so a quiet release feed doesn't burn API quota. - -The shell loop itself is exercised end-to-end with a fake `curl` that -serves canned page JSON, demonstrating that the resolved tag is the -highest-semver stable across all pages even when the highest tag lives -on page 2+. -""" - -from __future__ import annotations - -import os -import shutil -import subprocess -import textwrap -from pathlib import Path - -import pytest - -REPO_ROOT = Path(__file__).resolve().parents[4] -RUN_DAILY = REPO_ROOT / "tests" / "e2e" / "claude_code" / "cron_vm" / "run_daily.sh" - -# The extracted snippet starts AFTER `log`/`die` are defined in run_daily.sh, -# so the test harness has to provide its own stubs. Without them, a failure -# inside the snippet (e.g. jq returning an empty LITELLM_VERSION) would crash -# with `bash: die: command not found` (exit 127) instead of the intended -# diagnostic, making test failures unnecessarily hard to debug. -_PREAMBLE = ( - "set -Eeuo pipefail\n" - "log() { printf '==> %s\\n' \"$*\" >&2; }\n" - "die() { printf 'ERROR: %s\\n' \"$*\" >&2; exit 1; }\n" -) - - -def test_run_daily_does_not_early_break_on_first_stable_page() -> None: - """The regex pattern `select(test("...stable$"))] | length > 0` followed - by `break` is exactly the buggy early-stop. If it ever returns the - cron will silently start testing against a stale stable tag. - """ - body = RUN_DAILY.read_text() - assert ( - "length > 0" not in body - or "break" not in body - or ( - # If both substrings exist, make sure they aren't both inside the - # same release-pagination loop. The current loop only contains - # a `break` for the empty-page guard, not for any "length > 0" - # condition. - not _shares_loop_body(body, "length > 0", "break") - ) - ), ( - "run_daily.sh contains the old early-break-on-stable pattern. The " - "Releases endpoint orders by created_at, not semver, so breaking " - "on first-stable-seen can miss higher-versioned releases sitting " - "on later pages." - ) - - -def _shares_loop_body(body: str, needle_a: str, needle_b: str) -> bool: - """Heuristic: do both needles live inside a `for page in ...; do ... done` - block? Used as a defensive guard for the static check above.""" - in_loop = False - saw_a = False - saw_b = False - for line in body.splitlines(): - stripped = line.strip() - if stripped.startswith("for page in"): - in_loop = True - saw_a = False - saw_b = False - continue - if in_loop and stripped == "done": - if saw_a and saw_b: - return True - in_loop = False - continue - if in_loop: - if needle_a in line: - saw_a = True - if needle_b in line: - saw_b = True - return False - - -def test_run_daily_keeps_empty_page_break_guard() -> None: - """The empty-page break is the only break that should remain in the - pagination loop — without it a quiet release feed wastes API quota - walking past the last real page.""" - body = RUN_DAILY.read_text() - assert "jq 'length' \"${PAGE_JSON}\"" in body, ( - "run_daily.sh must still detect empty pages via `jq 'length' " - "${PAGE_JSON}`; without this the loop walks the full 5-page cap " - "even when there are no more releases." - ) - assert ( - '== "0"' in body - ), 'The empty-page guard must compare jq\'s length output to "0".' - - -def _make_fake_curl(scratch: Path, pages: dict[int, str]) -> Path: - """Build a fake `curl` shim that serves the canned page JSON for - each `page=N` request and an empty array for any page past the - last canned one. - - The shim mimics just enough of curl's CLI surface for the cron - script: it accepts the headers + URL we pass, ignores everything - we don't need, and writes the canned body to either stdout or the - --output target if one is given. - """ - pages_dir = scratch / "pages" - pages_dir.mkdir() - for page_num, body in pages.items(): - (pages_dir / f"page{page_num}.json").write_text(body) - - curl_path = scratch / "curl" - curl_path.write_text( - textwrap.dedent( - f"""\ - #!/usr/bin/env bash - # Fake curl for run_daily.sh release pagination tests. Serves - # page JSON from {pages_dir} keyed by the `page=` query value, - # and returns "[]" for pages past the last canned one (which - # is exactly how the real GitHub API behaves past the end). - url="" - output="" - while [[ $# -gt 0 ]]; do - case "$1" in - -fsS|-fsSL|-H|-o|--output) - if [[ "$1" == "-o" || "$1" == "--output" ]]; then - output="$2"; shift 2 - elif [[ "$1" == "-H" ]]; then - shift 2 - else - shift - fi - ;; - http*) - url="$1"; shift - ;; - *) - shift - ;; - esac - done - page="$(printf '%s' "$url" | sed -n 's/.*[?&]page=\\([0-9]*\\).*/\\1/p')" - [[ -z "$page" ]] && page=1 - file="{pages_dir}/page${{page}}.json" - if [[ -f "$file" ]]; then - if [[ -n "$output" ]]; then cp "$file" "$output"; else cat "$file"; fi - else - if [[ -n "$output" ]]; then printf '[]' > "$output"; else printf '[]'; fi - fi - """ - ) - ) - curl_path.chmod(0o755) - return curl_path - - -def _extract_resolution_snippet() -> str: - """Pull the pagination + sort_by + assignment block out of run_daily.sh - so the test exercises the actual production code path (not a copy). - - The block is everything from the GH_AUTH_HEADER setup down through - the LITELLM_VERSION emission. - """ - body = RUN_DAILY.read_text() - start = body.index("GH_AUTH_HEADER=()") - end = body.index('log "resolved litellm:') - return body[start:end] - - -@pytest.mark.skipif(shutil.which("jq") is None, reason="jq not available") -def test_run_daily_resolves_highest_semver_across_pages(tmp_path: Path) -> None: - """End-to-end: drive the actual run_daily.sh pagination loop with a - fake curl whose page 1 contains a freshly-cut LOW-version backport - (v1.80.1-stable) and page 2 contains a two-weeks-old HIGH-version - release (v1.83.0-stable). The correct behavior is to resolve - v1.83.0-stable. The pre-fix behavior would resolve v1.80.1-stable - because the early-break consumed only page 1. - """ - pages = { - # Page 1: most-recently-created releases. The order here matches - # what /releases?page=1 returns: created-at descending. The - # freshly-cut v1.80.1-stable backport sits at the top, plus a - # bunch of non-stable releases. - 1: """[ - {"tag_name": "v1.84.0-nightly.1"}, - {"tag_name": "v1.80.1-stable"}, - {"tag_name": "v1.84.0-nightly.0"} - ]""", - # Page 2: older releases. The HIGHER-version stable lives here - # because it was cut two weeks ago, before the v1.80.1 backport. - 2: """[ - {"tag_name": "v1.83.0-rc.5"}, - {"tag_name": "v1.83.0-stable"}, - {"tag_name": "v1.82.4-stable"} - ]""", - # Page 3+: empty -> the loop's empty-page guard fires here. - } - fake_curl_dir = tmp_path / "shim" - fake_curl_dir.mkdir() - _make_fake_curl(fake_curl_dir, pages) - - workdir = tmp_path / "work" - workdir.mkdir() - - snippet = _extract_resolution_snippet() - script = ( - _PREAMBLE - + f"WORKDIR={workdir!s}\n" - + snippet - + 'printf "%s" "${LITELLM_VERSION}"\n' - ) - - env = { - **os.environ, - "PATH": f"{fake_curl_dir}:{os.environ.get('PATH', '')}", - } - # Make sure the loop hits the fake curl, not the system one. - env.pop("GITHUB_TOKEN", None) - result = subprocess.run( - ["bash", "-c", script], - capture_output=True, - text=True, - env=env, - check=True, - ) - assert result.stdout == "v1.83.0-stable", ( - f"Expected the highest-semver stable across pages 1-2, got " - f"{result.stdout!r}. stderr={result.stderr!r}" - ) - - -@pytest.mark.skipif(shutil.which("jq") is None, reason="jq not available") -def test_run_daily_terminates_on_empty_page(tmp_path: Path) -> None: - """The empty-page guard must fire so we don't always walk all 5 - pages. With a single populated page and an empty page 2 we should - stop after fetching page 2 (the first empty response).""" - pages = {1: '[{"tag_name": "v1.50.0-stable"}]'} - fake_curl_dir = tmp_path / "shim" - fake_curl_dir.mkdir() - _make_fake_curl(fake_curl_dir, pages) - - workdir = tmp_path / "work" - workdir.mkdir() - - snippet = _extract_resolution_snippet() - script = ( - _PREAMBLE - + f"WORKDIR={workdir!s}\n" - + snippet - + 'printf "%s" "${LITELLM_VERSION}"\n' - ) - - env = { - **os.environ, - "PATH": f"{tmp_path}/shim:{os.environ.get('PATH', '')}", - } - env.pop("GITHUB_TOKEN", None) - result = subprocess.run( - ["bash", "-c", script], - capture_output=True, - text=True, - env=env, - check=True, - ) - assert result.stdout == "v1.50.0-stable" - # Only pages 1 and 2 should have been fetched (2 is empty -> break). - assert (workdir / "releases.page2.json").exists() - assert not (workdir / "releases.page3.json").exists(), ( - "Empty-page guard didn't fire — the loop kept walking past the " - "first empty response." - ) diff --git a/tests/e2e/claude_code/_publisher_unit_tests/test_run_daily_version_probe_scrubs_env.py b/tests/e2e/claude_code/_publisher_unit_tests/test_run_daily_version_probe_scrubs_env.py deleted file mode 100644 index 1c3959764f3..00000000000 --- a/tests/e2e/claude_code/_publisher_unit_tests/test_run_daily_version_probe_scrubs_env.py +++ /dev/null @@ -1,92 +0,0 @@ -"""Pin: the cron `claude --version` probe must run under `env -i`. - -The systemd service `litellm-compat-matrix.service` loads provider -credentials (`ANTHROPIC_API_KEY`, `AWS_BEARER_TOKEN_BEDROCK`, -`AZURE_FOUNDRY_API_KEY`) and the agent-shin GitHub token -(`AGENT_SHIN_GITHUB_TOKEN`) into `run_daily.sh`'s environment from -`/etc/litellm-compat-matrix.env`. Running the npm-installed `claude` -binary directly there would hand that full env to package code, so a -compromised `@anthropic-ai/claude-code` release could read those -secrets out of `os.environ` before the proxy or test harness ever -starts. The version probe must be wrapped in `env -i` with a minimal -PATH/HOME/USER/TERM/LANG/LC_ALL/TMPDIR allowlist — matching the -PR-gate's resolver/npm-install/pytest scrubs. -""" - -from __future__ import annotations - -from pathlib import Path - -REPO_ROOT = Path(__file__).resolve().parents[4] -RUN_DAILY = REPO_ROOT / "tests" / "e2e" / "claude_code" / "cron_vm" / "run_daily.sh" - - -def _version_probe_block() -> str: - body = RUN_DAILY.read_text() - start = body.index("CLAUDE_CODE_VERSION=") - end = body.index('[[ -n "${CLAUDE_CODE_VERSION}" ]]', start) - return body[start:end] - - -def test_version_probe_wraps_claude_in_env_i() -> None: - block = _version_probe_block() - assert "env -i" in block, ( - "run_daily.sh: the `claude --version` probe must run under " - "`env -i` so a compromised @anthropic-ai/claude-code package " - "cannot read provider/GitHub credentials out of the systemd " - "service environment." - ) - assert block.index("env -i") < block.index("claude --version"), ( - "run_daily.sh: `env -i` must precede `claude --version`; " - "otherwise the binary inherits the full credential-bearing env." - ) - - -def test_version_probe_env_i_excludes_provider_secrets() -> None: - block = _version_probe_block() - for forbidden in ( - "ANTHROPIC_API_KEY", - "AWS_BEARER_TOKEN_BEDROCK", - "AWS_ACCESS_KEY_ID", - "AWS_SECRET_ACCESS_KEY", - "VERTEXAI_CREDENTIALS", - "AZURE_FOUNDRY_API_KEY", - "GITHUB_TOKEN", - "AGENT_SHIN_GITHUB_TOKEN", - ): - assert forbidden not in block, ( - f"run_daily.sh: the version-probe `env -i` allowlist must " - f"not pass {forbidden} through. Found it inside the probe " - f"block." - ) - - -def test_version_probe_uses_isolated_home_not_runtime_user_home() -> None: - """Pin: the `claude --version` probe runs under a fresh empty HOME. - - `ProtectHome=read-only` in the systemd unit allows reads of the - runtime user's real home directory. If the probe's `env -i` - block forwards `HOME=${HOME}`, a compromised `claude` package - can `os.path.expanduser("~/.config/gh/hosts.yml")` or - `os.path.expanduser("~/.ssh/...")` and exfiltrate the contents - before the proxy or test harness ever starts. The probe must - point HOME at a per-run tmpdir under `${WORKDIR}` so the CLI - sees an empty HOME instead. - """ - block = _version_probe_block() - body = RUN_DAILY.read_text() - - assert "CLAUDE_PROBE_HOME=" in body, ( - "run_daily.sh: must define a `CLAUDE_PROBE_HOME` per-run tmpdir " - "for the `claude --version` probe so the CLI never sees the " - "runtime user's real $HOME." - ) - assert 'HOME="${CLAUDE_PROBE_HOME}"' in block, ( - "run_daily.sh: the probe's `env -i` block must set HOME to " - "the per-run isolated tmpdir, not to the runtime user's $HOME." - ) - assert 'HOME="${HOME}"' not in block, ( - "run_daily.sh: the probe's `env -i` block must not forward the " - "runtime user's $HOME to `claude --version`. Use the isolated " - "$CLAUDE_PROBE_HOME tmpdir instead." - ) diff --git a/tests/e2e/claude_code/_publisher_unit_tests/test_systemd_unit_credential_isolation.py b/tests/e2e/claude_code/_publisher_unit_tests/test_systemd_unit_credential_isolation.py deleted file mode 100644 index 12edce3cb50..00000000000 --- a/tests/e2e/claude_code/_publisher_unit_tests/test_systemd_unit_credential_isolation.py +++ /dev/null @@ -1,104 +0,0 @@ -"""Pin: the cron systemd unit hides credential-bearing dotdirs. - -`ProtectHome=read-only` blocks writes to /home/mateo but still allows -reads. A model-directed `Read` tool call (the PDF cells pass -`--allowed-tools Read` to the `claude` CLI) or a compromised -`@anthropic-ai/claude-code` package can read absolute paths under -the runtime user's home and exfiltrate the contents — even with the -per-`claude`-invocation HOME isolation in place, because absolute -paths bypass `~`-expansion. - -This file pins the second line of defense: the systemd unit lists -the credential-bearing dotdirs (`~/.config/gh`, `~/.ssh`, `~/.aws`, -`~/.docker`, `~/.kube`, `~/.gnupg`) under `InaccessiblePaths=` so -the kernel hides them from every process in the unit's mount -namespace, including any child of `claude --version` or the pytest -run. It also pins that `~/.config/gh` is *not* in `ReadWritePaths=` -— we pass `GH_TOKEN` inline to every `gh` invocation in -`run_daily.sh`, so the host gh-cli config is unused. -""" - -from __future__ import annotations - -import re -from pathlib import Path - -REPO_ROOT = Path(__file__).resolve().parents[4] -SERVICE = ( - REPO_ROOT / "tests" / "e2e" / "claude_code" / "cron_vm" / "litellm-compat-matrix.service" -) - - -def _service_text() -> str: - return SERVICE.read_text() - - -def _directive(name: str) -> str: - """Return the value of a single-line systemd directive (or empty).""" - text = _service_text() - match = re.search(rf"^\s*{re.escape(name)}\s*=\s*(.*)$", text, re.MULTILINE) - return match.group(1).strip() if match else "" - - -def test_inaccessible_paths_hides_credential_dotdirs() -> None: - """Every credential-bearing dotdir must be under `InaccessiblePaths=`.""" - inaccessible = _directive("InaccessiblePaths") - assert inaccessible, ( - "litellm-compat-matrix.service: must declare `InaccessiblePaths=` " - "to hide credential dotdirs from the `claude` subprocess and the " - "model-directed Read tool. Without this, an absolute-path read " - "like `Read('/home/mateo/.config/gh/hosts.yml')` exfiltrates " - "the gh-cli token despite the per-invocation HOME isolation." - ) - for path in ( - "/home/mateo/.config/gh", - "/home/mateo/.ssh", - "/home/mateo/.aws", - "/home/mateo/.docker", - "/home/mateo/.kube", - "/home/mateo/.gnupg", - ): - # Tolerated `-` prefix means "ignore if missing on host". - assert path in inaccessible, ( - f"litellm-compat-matrix.service: `{path}` must appear in " - f"`InaccessiblePaths=` so the cron `claude` subprocess can " - f"never read it (even via an absolute path that bypasses " - f"the per-invocation HOME override)." - ) - - -def test_gh_config_is_not_writeable() -> None: - """`~/.config/gh` is not whitelisted under `ReadWritePaths=`. - - We pass `GH_TOKEN` inline to every `gh` invocation in - `run_daily.sh` (`gh repo clone`, `gh pr create`, `gh pr edit`). - The host `~/.config/gh/hosts.yml` is therefore never consulted - or written to. Keeping it out of `ReadWritePaths=` is the second - line of defense: a future regression that drops the inline-token - convention will fail loudly (gh writes a new login config and - hits a read-only filesystem) rather than silently re-introduce - the credential exfiltration surface that - `InaccessiblePaths=/home/mateo/.config/gh` is closing. - """ - rw = _directive("ReadWritePaths") - assert ".config/gh" not in rw, ( - "litellm-compat-matrix.service: `/home/mateo/.config/gh` must " - "*not* appear in `ReadWritePaths=`. We pass `GH_TOKEN` inline " - "to every `gh` invocation in run_daily.sh, so the host gh-cli " - "config is never consulted or written to. Keeping the path out " - "of ReadWritePaths means a future regression that drops the " - "inline-token convention will fail loudly instead of silently " - "re-opening the credential exfiltration surface that " - "`InaccessiblePaths=` is closing." - ) - - -def test_protect_home_is_read_only_or_stricter() -> None: - """`ProtectHome=` must be at least `read-only`.""" - value = _directive("ProtectHome") - assert value in ("read-only", "tmpfs", "yes", "true"), ( - f"litellm-compat-matrix.service: `ProtectHome=` must be `read-only`, " - f"`tmpfs`, or `yes`. Got: {value!r}. Without this, the unit can " - f"write anywhere under /home/mateo, including overwriting " - f"~/.config/gh/hosts.yml." - ) diff --git a/tests/e2e/claude_code/basic_messaging_non_streaming/test_anthropic.py b/tests/e2e/claude_code/basic_messaging_non_streaming/test_anthropic.py index c06fff28d2d..21383b85da5 100644 --- a/tests/e2e/claude_code/basic_messaging_non_streaming/test_anthropic.py +++ b/tests/e2e/claude_code/basic_messaging_non_streaming/test_anthropic.py @@ -20,6 +20,7 @@ the matrix builder still sees three rows for this (feature, provider). from __future__ import annotations +import pytest from claude_code._basic_messaging import run_basic_messaging_cell # Per the PRD: each cell is exercised against three Claude tiers via the @@ -27,11 +28,12 @@ from claude_code._basic_messaging import run_basic_messaging_cell # routing config; the driver only sends the alias. ANTHROPIC_MODELS = [ "claude-haiku-4-5", - "claude-sonnet-4-6", + "claude-sonnet-4-5", "claude-opus-4-7", ] +@pytest.mark.covers("llm.messages.anthropic.basic.nonstream.works") def test_basic_messaging_non_streaming_anthropic(compat_result): """Drive the `claude` CLI against the LiteLLM proxy and assert a reply. diff --git a/tests/e2e/claude_code/basic_messaging_non_streaming/test_azure.py b/tests/e2e/claude_code/basic_messaging_non_streaming/test_azure.py index 2a962b244a8..19e88dbe3cb 100644 --- a/tests/e2e/claude_code/basic_messaging_non_streaming/test_azure.py +++ b/tests/e2e/claude_code/basic_messaging_non_streaming/test_azure.py @@ -25,6 +25,7 @@ the matrix builder still sees three rows for this (feature, provider). from __future__ import annotations +import pytest from claude_code._basic_messaging import run_basic_messaging_cell # Per-model aliases registered in the LiteLLM proxy's routing config to @@ -33,11 +34,12 @@ from claude_code._basic_messaging import run_basic_messaging_cell # resource URL and API key. AZURE_MODELS = [ "claude-haiku-4-5-azure", - "claude-sonnet-4-6-azure", + "claude-sonnet-4-5-azure", "claude-opus-4-7-azure", ] +@pytest.mark.covers("llm.messages.azure_foundry.basic.nonstream.works") def test_basic_messaging_non_streaming_azure(compat_result): """Drive the `claude` CLI against the LiteLLM proxy and assert a reply. diff --git a/tests/e2e/claude_code/basic_messaging_non_streaming/test_bedrock_converse.py b/tests/e2e/claude_code/basic_messaging_non_streaming/test_bedrock_converse.py index 2245ed7417a..2b0f49bc205 100644 --- a/tests/e2e/claude_code/basic_messaging_non_streaming/test_bedrock_converse.py +++ b/tests/e2e/claude_code/basic_messaging_non_streaming/test_bedrock_converse.py @@ -20,6 +20,7 @@ the matrix builder still sees three rows for this (feature, provider). from __future__ import annotations +import pytest from claude_code._basic_messaging import run_basic_messaging_cell # Per-model aliases registered in the LiteLLM proxy's routing config to @@ -28,11 +29,12 @@ from claude_code._basic_messaging import run_basic_messaging_cell # strategy. BEDROCK_CONVERSE_MODELS = [ "claude-haiku-4-5-bedrock-converse", - "claude-sonnet-4-6-bedrock-converse", + "claude-sonnet-4-5-bedrock-converse", "claude-opus-4-7-bedrock-converse", ] +@pytest.mark.covers("llm.messages.bedrock_converse.basic.nonstream.works") def test_basic_messaging_non_streaming_bedrock_converse(compat_result): """Drive the `claude` CLI against the LiteLLM proxy and assert a reply.""" run_basic_messaging_cell( diff --git a/tests/e2e/claude_code/basic_messaging_non_streaming/test_bedrock_invoke.py b/tests/e2e/claude_code/basic_messaging_non_streaming/test_bedrock_invoke.py index e0a6e77f3c1..937ea5ee27e 100644 --- a/tests/e2e/claude_code/basic_messaging_non_streaming/test_bedrock_invoke.py +++ b/tests/e2e/claude_code/basic_messaging_non_streaming/test_bedrock_invoke.py @@ -20,6 +20,7 @@ the matrix builder still sees three rows for this (feature, provider). from __future__ import annotations +import pytest from claude_code._basic_messaging import run_basic_messaging_cell # Per-model aliases registered in the LiteLLM proxy's routing config to @@ -28,11 +29,12 @@ from claude_code._basic_messaging import run_basic_messaging_cell # routing strategy. BEDROCK_INVOKE_MODELS = [ "claude-haiku-4-5-bedrock-invoke", - "claude-sonnet-4-6-bedrock-invoke", + "claude-sonnet-4-5-bedrock-invoke", "claude-opus-4-7-bedrock-invoke", ] +@pytest.mark.covers("llm.messages.bedrock_invoke.basic.nonstream.works") def test_basic_messaging_non_streaming_bedrock_invoke(compat_result): """Drive the `claude` CLI against the LiteLLM proxy and assert a reply.""" run_basic_messaging_cell( diff --git a/tests/e2e/claude_code/basic_messaging_non_streaming/test_vertex_ai.py b/tests/e2e/claude_code/basic_messaging_non_streaming/test_vertex_ai.py index e4e2a39e6cd..c46e5a8f762 100644 --- a/tests/e2e/claude_code/basic_messaging_non_streaming/test_vertex_ai.py +++ b/tests/e2e/claude_code/basic_messaging_non_streaming/test_vertex_ai.py @@ -20,6 +20,7 @@ the matrix builder still sees three rows for this (feature, provider). from __future__ import annotations +import pytest from claude_code._basic_messaging import run_basic_messaging_cell # Per-model aliases registered in the LiteLLM proxy's routing config to @@ -28,11 +29,12 @@ from claude_code._basic_messaging import run_basic_messaging_cell # model id and the GCP region. VERTEX_AI_MODELS = [ "claude-haiku-4-5-vertex", - "claude-sonnet-4-6-vertex", + "claude-sonnet-4-5-vertex", "claude-opus-4-7-vertex", ] +@pytest.mark.covers("llm.messages.vertex.basic.nonstream.works") def test_basic_messaging_non_streaming_vertex_ai(compat_result): """Drive the `claude` CLI against the LiteLLM proxy and assert a reply.""" run_basic_messaging_cell( diff --git a/tests/e2e/claude_code/basic_messaging_streaming/test_anthropic.py b/tests/e2e/claude_code/basic_messaging_streaming/test_anthropic.py index 56e3fb6c181..ce453f3e523 100644 --- a/tests/e2e/claude_code/basic_messaging_streaming/test_anthropic.py +++ b/tests/e2e/claude_code/basic_messaging_streaming/test_anthropic.py @@ -25,15 +25,17 @@ sees three rows for this (feature, provider). from __future__ import annotations +import pytest from claude_code._basic_messaging import run_basic_messaging_cell ANTHROPIC_MODELS = [ "claude-haiku-4-5", - "claude-sonnet-4-6", + "claude-sonnet-4-5", "claude-opus-4-7", ] +@pytest.mark.covers("llm.messages.anthropic.basic.stream.works") def test_basic_messaging_streaming_anthropic(compat_result): """Drive the `claude` CLI against the LiteLLM proxy and assert a non-empty streamed reply (one row per Claude tier). diff --git a/tests/e2e/claude_code/basic_messaging_streaming/test_azure.py b/tests/e2e/claude_code/basic_messaging_streaming/test_azure.py index b6c002d0b27..3307194e862 100644 --- a/tests/e2e/claude_code/basic_messaging_streaming/test_azure.py +++ b/tests/e2e/claude_code/basic_messaging_streaming/test_azure.py @@ -19,15 +19,17 @@ The (feature, provider) for this cell is inferred from the file path by from __future__ import annotations +import pytest from claude_code._basic_messaging import run_basic_messaging_cell AZURE_MODELS = [ "claude-haiku-4-5-azure", - "claude-sonnet-4-6-azure", + "claude-sonnet-4-5-azure", "claude-opus-4-7-azure", ] +@pytest.mark.covers("llm.messages.azure_foundry.basic.stream.works") def test_basic_messaging_streaming_azure(compat_result): """Drive the `claude` CLI against the LiteLLM proxy and assert a non-empty streamed reply (one row per Claude tier). diff --git a/tests/e2e/claude_code/basic_messaging_streaming/test_bedrock_converse.py b/tests/e2e/claude_code/basic_messaging_streaming/test_bedrock_converse.py index 44ac54515f0..a8bc0b77a5d 100644 --- a/tests/e2e/claude_code/basic_messaging_streaming/test_bedrock_converse.py +++ b/tests/e2e/claude_code/basic_messaging_streaming/test_bedrock_converse.py @@ -15,15 +15,17 @@ The (feature, provider) for this cell is inferred from the file path by from __future__ import annotations +import pytest from claude_code._basic_messaging import run_basic_messaging_cell BEDROCK_CONVERSE_MODELS = [ "claude-haiku-4-5-bedrock-converse", - "claude-sonnet-4-6-bedrock-converse", + "claude-sonnet-4-5-bedrock-converse", "claude-opus-4-7-bedrock-converse", ] +@pytest.mark.covers("llm.messages.bedrock_converse.basic.stream.works") def test_basic_messaging_streaming_bedrock_converse(compat_result): """Drive the `claude` CLI against the LiteLLM proxy and assert a non-empty streamed reply (one row per Claude tier). diff --git a/tests/e2e/claude_code/basic_messaging_streaming/test_bedrock_invoke.py b/tests/e2e/claude_code/basic_messaging_streaming/test_bedrock_invoke.py index 1d59d16cdc1..c0ece0e0721 100644 --- a/tests/e2e/claude_code/basic_messaging_streaming/test_bedrock_invoke.py +++ b/tests/e2e/claude_code/basic_messaging_streaming/test_bedrock_invoke.py @@ -15,15 +15,17 @@ The (feature, provider) for this cell is inferred from the file path by from __future__ import annotations +import pytest from claude_code._basic_messaging import run_basic_messaging_cell BEDROCK_INVOKE_MODELS = [ "claude-haiku-4-5-bedrock-invoke", - "claude-sonnet-4-6-bedrock-invoke", + "claude-sonnet-4-5-bedrock-invoke", "claude-opus-4-7-bedrock-invoke", ] +@pytest.mark.covers("llm.messages.bedrock_invoke.basic.stream.works") def test_basic_messaging_streaming_bedrock_invoke(compat_result): """Drive the `claude` CLI against the LiteLLM proxy and assert a non-empty streamed reply (one row per Claude tier). diff --git a/tests/e2e/claude_code/basic_messaging_streaming/test_vertex_ai.py b/tests/e2e/claude_code/basic_messaging_streaming/test_vertex_ai.py index 014a31160a8..13f1a0abf40 100644 --- a/tests/e2e/claude_code/basic_messaging_streaming/test_vertex_ai.py +++ b/tests/e2e/claude_code/basic_messaging_streaming/test_vertex_ai.py @@ -15,15 +15,17 @@ The (feature, provider) for this cell is inferred from the file path by from __future__ import annotations +import pytest from claude_code._basic_messaging import run_basic_messaging_cell VERTEX_AI_MODELS = [ "claude-haiku-4-5-vertex", - "claude-sonnet-4-6-vertex", + "claude-sonnet-4-5-vertex", "claude-opus-4-7-vertex", ] +@pytest.mark.covers("llm.messages.vertex.basic.stream.works") def test_basic_messaging_streaming_vertex_ai(compat_result): """Drive the `claude` CLI against the LiteLLM proxy and assert a non-empty streamed reply (one row per Claude tier). diff --git a/tests/e2e/claude_code/conftest.py b/tests/e2e/claude_code/conftest.py index 6ee8b940648..8f5c09fa4e2 100644 --- a/tests/e2e/claude_code/conftest.py +++ b/tests/e2e/claude_code/conftest.py @@ -169,8 +169,9 @@ def _manifest_feature_ids() -> FrozenSet[str]: Used as a positive filter so only directories that correspond to a real matrix row contribute results — utility/support directories - (e.g. `cron_vm`, `_driver_unit_tests`) are dropped regardless of - naming convention, and the rate-limit summary stays clean. + (e.g. `_driver_unit_tests`, `_builder_unit_tests`) are dropped + regardless of naming convention, and the rate-limit summary stays + clean. Returns an empty set if the manifest is missing or malformed; the caller treats that as "no path is a feature path", which is the @@ -199,11 +200,10 @@ def _infer_feature_and_provider(node_path: Path) -> Optional[tuple]: Path shape: tests/e2e/claude_code//test_.py Returns None if the file is not a per-feature test (e.g. unit tests - under `_driver_unit_tests/` or support code under `cron_vm/`), so - those don't pollute the matrix artifact. We positively filter the - parent directory against `manifest.yaml` rather than relying on - naming conventions, because non-feature siblings don't all share - an underscore prefix. + under `_driver_unit_tests/`), so those don't pollute the matrix + artifact. We positively filter the parent directory against + `manifest.yaml` rather than relying on naming conventions, because + non-feature siblings don't all share an underscore prefix. """ name = node_path.name if not name.startswith("test_") or not name.endswith(".py"): @@ -548,3 +548,117 @@ def pytest_sessionfinish(session, exitstatus): ) summary_path.write_text(json.dumps(summary, indent=2, sort_keys=True)) _print_rate_limit_summary(summary) + + +# --------------------------------------------------------------------------- +# Session-scoped compat model registration. +# +# The compat cells probe hardcoded virtual names like ``claude-sonnet-4-5`` +# and ``claude-sonnet-4-5-bedrock-invoke``. On stage those live in the +# gateway's model_list at deploy time; locally the docker-config.yaml +# under tests/e2e/ only declares one of them, so every non-haiku cell +# 400s with ``Invalid model name``. The fixture here reconciles the two: +# it reads ``test_config.yaml`` (the ground-truth compat matrix config) +# and POSTs ``/model/new`` for the subset whose provider credentials are +# actually set in the current environment, then tears them all down at +# session end. +# +# Kept below the rest of the conftest so the compat-artifact hooks stay +# grouped up top. The fixture is opt-in via autouse=True on the session +# scope, so a cell that hits the proxy sees the deployment ready without +# any per-cell wiring, and pure unit tests that never reach the proxy +# pay only one skipped-liveness check. +# --------------------------------------------------------------------------- + +from claude_code._env import ProxyConfig, resolve_proxy # noqa: E402 +from claude_code._compat_models import ( # noqa: E402 + CompatDeployment, + load_all_deployments, +) + + +def _build_control_gateway(proxy: ProxyConfig): + """Local import of the shared harness so the pure-unit-test tree + under ``_driver_unit_tests/`` etc. never has to pull it in. The + control plane transport is what /model/new lives on; SplitTransport + routes it correctly for both monolithic and split deployments. + + The endpoints come from the *resolved* proxy, not from a second + independent env read, so registration and the cells always hit the + same host and key. Both planes get the one URL the cells use; the + deployment is fronted by a single address that routes management + and LLM paths itself.""" + from e2e_gateway import build_gateway + + return build_gateway( + base_url=proxy.base_url, + master_key=proxy.api_key, + control_plane_base_url=proxy.base_url, + ) + + +def _register_deployment(gateway, deployment: CompatDeployment) -> str: + """Register one deployment and return its proxy-assigned model_id + once it is servable on the data plane.""" + return gateway.create_model( + deployment.model_name, + deployment.litellm_params, + ) + + +@pytest.fixture(scope="session", autouse=True) +def _compat_models_registered() -> Any: + """Register every compat deployment against the running proxy, then + tear them all down on session exit. + + Skips silently if the proxy env is not configured (no + ``LITELLM_PROXY_URL``/``LITELLM_MASTER_KEY``) so unit-test runs + stay hermetic. + + Design note: we always attempt to register all 15 deployments, + regardless of what credentials are exported in the test-runner's + shell. The credentials live in the proxy container's environment + (via docker-compose ``env_file``), not the shell running pytest - + so gating on shell env would filter out deployments the proxy can + actually serve. Per-deployment ``/model/new`` failures are printed + but do not abort the session: the cells that need that specific + deployment will 400 with "Invalid model name" and fail loudly, + which is the right signal (missing cred on the proxy side).""" + proxy = resolve_proxy() + if proxy is None: + yield + return + + from requests import RequestException + + gateway = _build_control_gateway(proxy) + registered_ids: list[str] = [] + failures: list[tuple[str, str]] = [] + try: + for deployment in load_all_deployments(): + try: + model_id = _register_deployment(gateway, deployment) + registered_ids.append(model_id) + except (AssertionError, RequestException) as exc: + failures.append((deployment.model_name, str(exc))) + if failures: + summary = "\n".join( + f" - {name}: {reason}" for name, reason in failures + ) + print( + f"[compat fixture] {len(failures)} of " + f"{len(failures) + len(registered_ids)} deployments " + f"failed to register (proxy likely missing that provider's " + f"credentials); cells that target them will fail loudly:\n" + f"{summary}" + ) + yield + finally: + for model_id in registered_ids: + try: + gateway.delete_model(model_id) + except (AssertionError, RequestException): + # Best-effort — teardown surfaces via warnings inside + # ``delete_model`` already; swallowing here so one flaky + # delete does not mask real test failures. + pass diff --git a/tests/e2e/claude_code/count_tokens/test_anthropic.py b/tests/e2e/claude_code/count_tokens/test_anthropic.py index 3508063459c..2fbdd4212c4 100644 --- a/tests/e2e/claude_code/count_tokens/test_anthropic.py +++ b/tests/e2e/claude_code/count_tokens/test_anthropic.py @@ -37,44 +37,27 @@ CLI rows isn't useful here. from __future__ import annotations -import os - import pytest +from claude_code._env import require_proxy from claude_code.http_probe import ( assert_count_tokens_shape, probe_count_tokens, ) -PROXY_BASE_URL_ENV = "LITELLM_PROXY_BASE_URL" -PROXY_API_KEY_ENV = "LITELLM_PROXY_API_KEY" ANTHROPIC_MODELS = [ "claude-haiku-4-5", - "claude-sonnet-4-6", + "claude-sonnet-4-5", "claude-opus-4-7", ] +@pytest.mark.covers("llm.messages.anthropic.count_tokens.nonstream.works") def test_count_tokens_anthropic(compat_result): """Probe `/v1/messages/count_tokens` for each Anthropic tier and assert the response shape.""" - 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) failures = [] for model in ANTHROPIC_MODELS: diff --git a/tests/e2e/claude_code/count_tokens/test_azure.py b/tests/e2e/claude_code/count_tokens/test_azure.py index 2b8707b50b0..a9aa168ccea 100644 --- a/tests/e2e/claude_code/count_tokens/test_azure.py +++ b/tests/e2e/claude_code/count_tokens/test_azure.py @@ -37,44 +37,27 @@ CLI rows isn't useful here. from __future__ import annotations -import os - import pytest +from claude_code._env import require_proxy from claude_code.http_probe import ( assert_count_tokens_shape, probe_count_tokens, ) -PROXY_BASE_URL_ENV = "LITELLM_PROXY_BASE_URL" -PROXY_API_KEY_ENV = "LITELLM_PROXY_API_KEY" AZURE_MODELS = [ "claude-haiku-4-5-azure", - "claude-sonnet-4-6-azure", + "claude-sonnet-4-5-azure", "claude-opus-4-7-azure", ] +@pytest.mark.covers("llm.messages.azure_foundry.count_tokens.nonstream.works") def test_count_tokens_azure(compat_result): """Probe `/v1/messages/count_tokens` for each Azure (Microsoft Foundry) tier and assert the response shape.""" - 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) failures = [] for model in AZURE_MODELS: diff --git a/tests/e2e/claude_code/count_tokens/test_bedrock_converse.py b/tests/e2e/claude_code/count_tokens/test_bedrock_converse.py index 4221773ead2..6dcff3ecae3 100644 --- a/tests/e2e/claude_code/count_tokens/test_bedrock_converse.py +++ b/tests/e2e/claude_code/count_tokens/test_bedrock_converse.py @@ -37,44 +37,27 @@ CLI rows isn't useful here. from __future__ import annotations -import os - import pytest +from claude_code._env import require_proxy from claude_code.http_probe import ( assert_count_tokens_shape, probe_count_tokens, ) -PROXY_BASE_URL_ENV = "LITELLM_PROXY_BASE_URL" -PROXY_API_KEY_ENV = "LITELLM_PROXY_API_KEY" BEDROCK_CONVERSE_MODELS = [ "claude-haiku-4-5-bedrock-converse", - "claude-sonnet-4-6-bedrock-converse", + "claude-sonnet-4-5-bedrock-converse", "claude-opus-4-7-bedrock-converse", ] +@pytest.mark.covers("llm.messages.bedrock_converse.count_tokens.nonstream.works") def test_count_tokens_bedrock_converse(compat_result): """Probe `/v1/messages/count_tokens` for each Bedrock (Converse) tier and assert the response shape.""" - 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) failures = [] for model in BEDROCK_CONVERSE_MODELS: diff --git a/tests/e2e/claude_code/count_tokens/test_bedrock_invoke.py b/tests/e2e/claude_code/count_tokens/test_bedrock_invoke.py index cc70bf12392..ae89067dc00 100644 --- a/tests/e2e/claude_code/count_tokens/test_bedrock_invoke.py +++ b/tests/e2e/claude_code/count_tokens/test_bedrock_invoke.py @@ -37,44 +37,27 @@ CLI rows isn't useful here. from __future__ import annotations -import os - import pytest +from claude_code._env import require_proxy from claude_code.http_probe import ( assert_count_tokens_shape, probe_count_tokens, ) -PROXY_BASE_URL_ENV = "LITELLM_PROXY_BASE_URL" -PROXY_API_KEY_ENV = "LITELLM_PROXY_API_KEY" BEDROCK_INVOKE_MODELS = [ "claude-haiku-4-5-bedrock-invoke", - "claude-sonnet-4-6-bedrock-invoke", + "claude-sonnet-4-5-bedrock-invoke", "claude-opus-4-7-bedrock-invoke", ] +@pytest.mark.covers("llm.messages.bedrock_invoke.count_tokens.nonstream.works") def test_count_tokens_bedrock_invoke(compat_result): """Probe `/v1/messages/count_tokens` for each Bedrock (Invoke) tier and assert the response shape.""" - 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) failures = [] for model in BEDROCK_INVOKE_MODELS: diff --git a/tests/e2e/claude_code/count_tokens/test_vertex_ai.py b/tests/e2e/claude_code/count_tokens/test_vertex_ai.py index 8c2678f7010..0f952496566 100644 --- a/tests/e2e/claude_code/count_tokens/test_vertex_ai.py +++ b/tests/e2e/claude_code/count_tokens/test_vertex_ai.py @@ -37,44 +37,27 @@ CLI rows isn't useful here. from __future__ import annotations -import os - import pytest +from claude_code._env import require_proxy from claude_code.http_probe import ( assert_count_tokens_shape, probe_count_tokens, ) -PROXY_BASE_URL_ENV = "LITELLM_PROXY_BASE_URL" -PROXY_API_KEY_ENV = "LITELLM_PROXY_API_KEY" VERTEX_AI_MODELS = [ "claude-haiku-4-5-vertex", - "claude-sonnet-4-6-vertex", + "claude-sonnet-4-5-vertex", "claude-opus-4-7-vertex", ] +@pytest.mark.covers("llm.messages.vertex.count_tokens.nonstream.works") def test_count_tokens_vertex_ai(compat_result): """Probe `/v1/messages/count_tokens` for each Vertex AI tier and assert the response shape.""" - 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) failures = [] for model in VERTEX_AI_MODELS: diff --git a/tests/e2e/claude_code/cron_vm/build_matrix.py b/tests/e2e/claude_code/cron_vm/build_matrix.py deleted file mode 100644 index 128f041cced..00000000000 --- a/tests/e2e/claude_code/cron_vm/build_matrix.py +++ /dev/null @@ -1,50 +0,0 @@ -"""Tiny CLI wrapper around `claude_code.matrix_builder.build_from_paths`. - -Exists only so `run_daily.sh` can hand the version metadata + paths into -the matrix builder without re-implementing it in bash. All real logic -lives in `matrix_builder.py`, which has its own unit tests under -`_builder_unit_tests/`. - -Invoked from the cron worktree (where `uv sync` has installed pyyaml), -not the dev checkout — the bash script `cd`s into the worktree before -`uv run python`-ing this file. -""" - -from __future__ import annotations - -import argparse -import datetime -import sys -from pathlib import Path - -sys.path.insert(0, str(Path(__file__).resolve().parents[2])) - -from claude_code.matrix_builder import build_from_paths # noqa: E402 # import needs the sys.path bootstrap above - - -def main() -> int: - parser = argparse.ArgumentParser(description=__doc__) - parser.add_argument("--manifest", type=Path, required=True) - parser.add_argument("--results", type=Path, required=True) - parser.add_argument("--output", type=Path, required=True) - parser.add_argument("--litellm-version", required=True) - parser.add_argument("--claude-code-version", required=True) - args = parser.parse_args() - - generated_at = datetime.datetime.now(datetime.timezone.utc).strftime( - "%Y-%m-%dT%H:%M:%SZ" - ) - build_from_paths( - manifest_path=args.manifest, - results_path=args.results, - litellm_version=args.litellm_version, - claude_code_version=args.claude_code_version, - generated_at=generated_at, - output_path=args.output, - ) - print(f"wrote {args.output}") - return 0 - - -if __name__ == "__main__": - sys.exit(main()) diff --git a/tests/e2e/claude_code/cron_vm/litellm-compat-matrix.env.example b/tests/e2e/claude_code/cron_vm/litellm-compat-matrix.env.example deleted file mode 100644 index 5ca7937a426..00000000000 --- a/tests/e2e/claude_code/cron_vm/litellm-compat-matrix.env.example +++ /dev/null @@ -1,59 +0,0 @@ -# Environment file consumed by `litellm-compat-matrix.service`. -# -# Install at `/etc/litellm-compat-matrix.env` and chmod 0600. -# `EnvironmentFile=-` in the unit means the service is allowed to start -# even if this file is missing, but the populator will fail at the -# first provider request without these credentials. - -# Anthropic -ANTHROPIC_API_KEY= - -# Bedrock (invoke + converse columns). -# Use Anthropic's Bedrock API-key passthrough (long-lived bearer token). -# No AWS_ACCESS_KEY_ID/AWS_SECRET_ACCESS_KEY required for the matrix -- -# both the LiteLLM invoke and converse routes pick up -# AWS_BEARER_TOKEN_BEDROCK when present. -AWS_BEARER_TOKEN_BEDROCK= -AWS_REGION_NAME=us-east-1 - -# Vertex AI. -# On the GCP VM, the default service-account ADC from the metadata server -# is used -- no JSON key file is needed. If you ever need to run outside -# GCP, also export GOOGLE_APPLICATION_CREDENTIALS=/path/to/sa.json. -VERTEXAI_PROJECT= -VERTEXAI_LOCATION=global - -# Microsoft Foundry (Azure column) -AZURE_FOUNDRY_API_KEY= -AZURE_FOUNDRY_API_BASE= - -# Azure cell of the `passthrough` row. Foundry-mode Claude Code sends -# the model in the request body, so the proxy's /azure passthrough -# cannot resolve a router alias and falls back to these env vars. -# AZURE_API_BASE is the Foundry resource's Anthropic surface, i.e. -# https://.services.ai.azure.com/anthropic ; AZURE_API_KEY -# is the same key as AZURE_FOUNDRY_API_KEY. -AZURE_API_BASE= -AZURE_API_KEY= - -# REQUIRED for publishing: PAT for the `agent-shin` user, used to push -# the daily compat-matrix branch to its fork (agent-shin/litellm-docs) -# and open the cross-repo PR against BerriAI/litellm-docs. Scopes: -# classic `repo` + `workflow`, or fine-grained on agent-shin/litellm-docs -# with Contents:RW + Pull requests:RW + Workflows:RW. -# Skip by setting SKIP_PUBLISH=1 (publishes nothing; only writes the -# matrix JSON locally). -AGENT_SHIN_GITHUB_TOKEN= - -# Optional: lifts the unauthenticated rate limit on the GitHub Releases -# API used by `resolver.py`. Any token works (read-only). Not required. -# GITHUB_TOKEN= - -# Optional overrides; defaults are sensible for the cron VM. -# PROXY_PORT=4100 -# LITELLM_WORKTREE=/home/mateo/litellm-cron-worktree -# DOCS_REPO=BerriAI/litellm-docs -# DOCS_BRANCH=main -# DOCS_TARGET_PATH=src/data/compatibility-matrix.json -# FORK_OWNER=agent-shin -# FORK_REPO=agent-shin/litellm-docs diff --git a/tests/e2e/claude_code/cron_vm/litellm-compat-matrix.service b/tests/e2e/claude_code/cron_vm/litellm-compat-matrix.service deleted file mode 100644 index c05ece90f50..00000000000 --- a/tests/e2e/claude_code/cron_vm/litellm-compat-matrix.service +++ /dev/null @@ -1,141 +0,0 @@ -# systemd service for the Claude Code compatibility-matrix populator. -# -# Triggered by `litellm-compat-matrix.timer`; not started directly. The -# unit is a `Type=oneshot` so the timer's `OnCalendar=` semantics -# describe "run once per day" cleanly — there's no long-lived daemon to -# supervise; each invocation runs the populator end-to-end and exits. -# -# Install -# ------- -# -# sudo cp tests/e2e/claude_code/cron_vm/litellm-compat-matrix.service /etc/systemd/system/ -# sudo cp tests/e2e/claude_code/cron_vm/litellm-compat-matrix.timer /etc/systemd/system/ -# sudo systemctl daemon-reload -# sudo systemctl enable --now litellm-compat-matrix.timer -# -# Paths are hard-coded to /home/mateo rather than using systemd's %h -# specifier. Why: in *system* units (this one), %h is expanded at -# parse time against the *manager's* home -- which is /root for PID 1 -# -- and *not* against the User= directive. That mismatch makes -# ReadWritePaths point at /root/.cache (which doesn't exist), causing -# the namespace setup to fail with status=226/NAMESPACE before the -# script ever runs. The runtime user (`User=mateo`) must: -# -# * have a checkout of `BerriAI/litellm` at `~/litellm/litellm` so the -# publisher module is importable; -# * have a uv venv at `~/litellm/litellm/.venv` (created by -# `uv sync --frozen` inside that checkout once); -# * have `gh` already authenticated against an account with -# `pull-requests: write` on `BerriAI/litellm-docs`; -# * have provider credentials exported in `/etc/litellm-compat-matrix.env` -# (see `litellm-compat-matrix.env.example` in this directory). - -[Unit] -Description=Claude Code compatibility-matrix populator (oneshot) -Wants=network-online.target -After=network-online.target - -[Service] -Type=oneshot -User=mateo -Group=mateo - -# Provider credentials + any gh/PROXY_PORT overrides live here. Format -# is the standard `KEY=value` one line per env var. -EnvironmentFile=-/etc/litellm-compat-matrix.env - -# systemd starts with a minimal PATH (~/usr/local/bin:/usr/bin:/bin). -# `uv` and `claude` are installed under the runtime user's `~/.local/bin` -# so we have to prepend it explicitly; otherwise run_daily.sh fails at -# the up-front command-presence check. -Environment=PATH=/home/mateo/.local/bin:/usr/local/sbin:/usr/local/bin:/usr/sbin:/usr/bin:/sbin:/bin - -# `HOME` is auto-set to /home/mateo when User=mateo is honored, but be -# explicit so anything that reads $HOME (e.g. uv's cache lookup, the -# claude CLI's per-session dir) sees the right value even if a future -# refactor flips DynamicUser= or PrivateUsers= on. -Environment=HOME=/home/mateo - -WorkingDirectory=/home/mateo/litellm/litellm - -ExecStart=/home/mateo/litellm/litellm/tests/e2e/claude_code/cron_vm/run_daily.sh - -# 90 minutes is generous: cold runs do `git clone` + `uv sync` of a new -# tag's lockfile, which can take a couple of minutes on a 2-vCPU VM, -# plus 30 cells of pytest hitting four cloud providers. -TimeoutStartSec=90min - -# A failed run shouldn't restart automatically — the next timer fire is -# the right retry. Reruns of the same day's matrix are idempotent. -Restart=no - -# Security hardening: the populator only reads the litellm checkout and -# the env-file; everything else it writes lives in either the worktree -# (managed) or `/tmp` (cleaned up by tempfile). -# -# `ProtectHome=read-only` blocks writes to /home/mateo but still -# allows reads. That's safe for the trusted run_daily.sh script -# itself, but unsafe for any subprocess we don't control: a -# compromised npm-installed `claude` package, or a model-directed -# `Read` tool call during a PDF/vision cell, could read sensitive -# host files like `~/.config/gh/hosts.yml` (gh-host token), -# `~/.ssh/`, or `~/.bash_history`. We mitigate that at the call -# boundary: every `claude` subprocess (the up-front `claude --version` -# probe in run_daily.sh, plus every CLI invocation routed through -# tests/e2e/claude_code/cli_driver.py) runs with `HOME` pointed at a -# fresh empty per-invocation tmpdir, not at /home/mateo. The CLI -# never sees the runtime user's real dotfiles. `gh` invocations in -# run_daily.sh pass `GH_TOKEN` inline, so they never need to read -# ~/.config/gh either; that path is intentionally NOT in the -# whitelist below — keeping it out is the second line of defense if -# the inline-token convention is ever accidentally regressed. -# -# ReadWritePaths whitelist: -# * litellm-cron-worktree - the long-lived stable-tag checkout + -# its `.venv` (`uv sync` rewrites every -# run) + `.uv-bin` (pinned `uv` binary -# cache). -# * .cache - uv's wheel cache (~/.cache/uv) so we -# don't redownload pinned deps each -# run. Used only by the trusted `uv` -# process; not exposed to `claude`. -# * /tmp - mktemp -d workdir, proxy logs, and -# the per-`claude`-invocation isolated -# HOME tmpdirs. PrivateTmp=true below -# gives the service its own tmpfs view -# so these don't escape to the host. -NoNewPrivileges=true -ProtectSystem=strict -ProtectHome=read-only -ReadWritePaths=/home/mateo/litellm-cron-worktree /home/mateo/.cache /tmp -PrivateTmp=true - -# Filesystem-level hiding for credential-bearing dotdirs/files. Even -# though `ProtectHome=read-only` prevents writes, a model-directed -# `Read` tool call (the PDF cells pass `--allowed-tools Read`) or a -# compromised `claude` package can read absolute paths under -# /home/mateo and exfiltrate the contents. `InaccessiblePaths=` makes -# the listed paths look like empty/missing to every process in the -# unit's mount namespace -- including the trusted populator script, -# which is fine because it doesn't need any of these: -# -# * .config/gh - gh CLI host token; we pass GH_TOKEN inline to -# every `gh` invocation (clone/PR/reviewer) so the -# host config is never consulted. -# * .ssh - never used by the populator. -# * .aws - upstream AWS credentials are passed to the proxy -# via the EnvironmentFile (provider env vars), not -# via shared SDK config files. -# * .docker - the populator never talks to a docker socket. -# * .kube - the populator never talks to a k8s API. -# * .gnupg - no GPG signing on the bot's commits. -# -# Leading `-` makes systemd tolerant if a path doesn't exist on the -# host (the unit is portable across VMs that may not have all of -# them set up). Anything else under /home/mateo (the litellm -# checkout, the cron worktree, the uv cache, .local/bin for the -# claude/uv/gh binaries on PATH) stays read-accessible. -InaccessiblePaths=-/home/mateo/.config/gh -/home/mateo/.ssh -/home/mateo/.aws -/home/mateo/.docker -/home/mateo/.kube -/home/mateo/.gnupg - -[Install] -WantedBy=multi-user.target diff --git a/tests/e2e/claude_code/cron_vm/litellm-compat-matrix.timer b/tests/e2e/claude_code/cron_vm/litellm-compat-matrix.timer deleted file mode 100644 index ee22538c6ed..00000000000 --- a/tests/e2e/claude_code/cron_vm/litellm-compat-matrix.timer +++ /dev/null @@ -1,25 +0,0 @@ -# Daily timer for the compatibility-matrix populator. -# -# 06:00 UTC matches the original GitHub Actions cron schedule; chosen so -# operators in US/EU timezones see fresh PRs at the start of their work -# day. -# -# `Persistent=true` causes a missed run (VM was off / suspended) to -# fire the next time the timer is started, which is the property we -# want for a once-a-day job: the matrix should refresh as soon as the -# VM is reachable again, not wait another 24h. -# -# `RandomizedDelaySec=10min` smears load if multiple matrix-style -# pipelines are ever colocated on the same VM in the future. - -[Unit] -Description=Run the Claude Code compatibility-matrix populator daily - -[Timer] -OnCalendar=*-*-* 06:00:00 UTC -Persistent=true -RandomizedDelaySec=10min -Unit=litellm-compat-matrix.service - -[Install] -WantedBy=timers.target diff --git a/tests/e2e/claude_code/cron_vm/run_daily.sh b/tests/e2e/claude_code/cron_vm/run_daily.sh deleted file mode 100755 index ae2d67c070c..00000000000 --- a/tests/e2e/claude_code/cron_vm/run_daily.sh +++ /dev/null @@ -1,590 +0,0 @@ -#!/usr/bin/env bash -# Daily Claude Code compatibility-matrix populator. -# -# Runs from the GCP VM `litellm-compatibility-matrix-populator` via the -# systemd timer in this directory. The flow is: -# -# 1. Resolve the latest LiteLLM v*-stable tag from the GitHub Releases API. -# 2. Update a long-lived worktree at $WORKTREE to that tag and `uv sync` it. -# 3. Boot the proxy as a background subprocess on $PROXY_PORT (default -# 4100; a separate port from the human-tended :4000 proxy). -# 4. Run `pytest tests/e2e/claude_code/` against the proxy. Test failures -# become `fail` cells in the JSON, not script errors. -# 5. Hand the per-test results artifact + manifest to a small Python -# CLI (`build_matrix.py`) that wraps the existing -# `matrix_builder.build_from_paths` to produce the published -# compatibility-matrix.json. -# 6. `gh repo clone` litellm-docs, write the JSON to a deterministic -# branch (`compat-matrix/--`), commit, -# `git push --force`, and `gh pr create`. -# -# Same-day reruns land on the same branch so they update the existing PR -# rather than spawning a new one. If the JSON is byte-identical to the -# docs branch, we skip the push entirely. -# -# Required commands on $PATH: git, uv, gh, jq, curl, claude. -# Required state: ~/litellm/litellm checked out (this file lives in it), -# $WORKTREE is created on first run, gh is already authenticated. -# -# Override any default by setting the matching env var; see the systemd -# unit for the production wiring. - -set -Eeuo pipefail - -LITELLM_REPO="${LITELLM_REPO:-${HOME}/litellm/litellm}" -WORKTREE="${LITELLM_WORKTREE:-${HOME}/litellm-cron-worktree}" -PROXY_PORT="${PROXY_PORT:-4100}" -PROXY_API_KEY="${PROXY_API_KEY:-sk-cron-matrix}" -DOCS_REPO="${DOCS_REPO:-BerriAI/litellm-docs}" -DOCS_BRANCH="${DOCS_BRANCH:-main}" -DOCS_TARGET_PATH="${DOCS_TARGET_PATH:-src/data/compatibility-matrix.json}" -SKIP_PUBLISH="${SKIP_PUBLISH:-0}" -PYTEST_K="${PYTEST_K:-}" -# Comma-separated GitHub usernames to request a review from on every PR. -# Reviewers must have at least read access to ${DOCS_REPO}. PR-author -# (agent-shin) has implicit rights to request reviews from anyone with -# read access, so no extra token scope is needed. Set to empty to skip. -PR_REVIEWERS="${PR_REVIEWERS:-mateo-berri}" - -POPULATOR_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)" -WORKDIR="$(mktemp -d -t litellm-compat-matrix.XXXXXX)" -PROXY_PID_FILE="${WORKDIR}/proxy.pid" - -# Cleanup is intentionally aggressive: it can run on normal exit, on a -# signal received by the script, or after a partial failure where the -# proxy is up but ${PROXY_PID_FILE} is stale. We try four things in -# order and stop as soon as the proxy port is free: -# -# 1. SIGTERM the pid recorded in proxy.pid. -# 2. SIGKILL anything from `pgrep -f "litellm.*--port ${PROXY_PORT}"` -# that survived. This catches the common case where the recorded -# pid was the sh wrapper, not the long-lived python child. -# 3. ss -K on the port (kernel kills sockets but not processes; -# mostly useful for catching lingering CLOSE_WAITs). -# 4. wipe ${WORKDIR}. -cleanup() { - local rc=$? - set +e - local proxy_pid - if [[ -f "${PROXY_PID_FILE}" ]]; then - proxy_pid="$(cat "${PROXY_PID_FILE}")" - if [[ -n "${proxy_pid}" ]]; then - kill -TERM "-${proxy_pid}" 2>/dev/null || kill -TERM "${proxy_pid}" 2>/dev/null || true - for _ in 1 2 3 4 5; do - kill -0 "${proxy_pid}" 2>/dev/null || break - sleep 1 - done - fi - fi - # Belt-and-braces: any python or uv talking to ${PROXY_PORT} that - # survived the SIGTERM gets SIGKILL'd by name. - pgrep -f "litellm.*--port[ =]?${PROXY_PORT}([^0-9]|$)" 2>/dev/null \ - | xargs -r kill -KILL 2>/dev/null || true - pgrep -f "${WORKTREE}/.uv-bin/uv.*run litellm" 2>/dev/null \ - | xargs -r kill -KILL 2>/dev/null || true - rm -rf "${WORKDIR}" - exit "${rc}" -} -trap cleanup EXIT INT TERM - -log() { printf '==> %s\n' "$*" >&2; } -die() { printf 'ERROR: %s\n' "$*" >&2; exit 1; } - -for cmd in git uv gh jq curl claude; do - command -v "${cmd}" >/dev/null 2>&1 || die "missing required command: ${cmd}" -done - -# Publishing is from a fork (agent-shin/litellm-docs) so neither the cron -# host nor the bot identity needs write access to BerriAI/litellm-docs. We -# require the fork token up front -- failing 30 minutes into a run because -# the env file is missing one line is a waste of CI quota. -if [[ "${SKIP_PUBLISH}" != "1" ]]; then - [[ -n "${AGENT_SHIN_GITHUB_TOKEN:-}" ]] \ - || die "AGENT_SHIN_GITHUB_TOKEN required to open PRs from agent-shin/litellm-docs (or set SKIP_PUBLISH=1)" -fi - -# --------------------------------------------------------------------------- -# 1. Resolve versions -# --------------------------------------------------------------------------- - -# Newest v*-stable release on BerriAI/litellm. The `select(...)` filter -# drops drafts/non-stable, the version_key sort handles 1.10 > 1.9. -# -# Paginate through the releases endpoint instead of grabbing only page 1 -# (default page_size=30). LiteLLM ships multiple non-stable releases per -# day, so it's common to need to walk past 30+ entries before hitting -# the most recent v*-stable. We cap at 5 pages (500 releases) which is -# conservatively beyond the worst observed gap. -# -# We deliberately do NOT short-circuit on the first page that contains a -# v*-stable tag. The /releases endpoint orders by `created_at`, not by -# semver, so a backport on an older series (e.g. v1.80.1-stable cut -# today) can show up on an earlier page than a higher-versioned release -# (v1.83.0-stable cut two weeks ago). Breaking early on first-stable-seen -# would silently pin the cron to the stale tag because the -# higher-versioned release still on a later page would never make it -# into the merged set the `sort_by` below consumes. The only break we -# keep is the empty-page guard, which means a quiet period in the -# release feed doesn't waste API quota — we just always walk far enough -# to be confident we've seen the highest stable tag. -GH_AUTH_HEADER=() -if [[ -n "${GITHUB_TOKEN:-}" ]]; then - GH_AUTH_HEADER=(-H "Authorization: Bearer ${GITHUB_TOKEN}") -fi -RELEASES_JSON="${WORKDIR}/releases.json" -echo "[]" >"${RELEASES_JSON}" -for page in 1 2 3 4 5; do - PAGE_JSON="${WORKDIR}/releases.page${page}.json" - curl -fsS \ - -H 'Accept: application/vnd.github+json' \ - -H 'User-Agent: litellm-compat-matrix' \ - "${GH_AUTH_HEADER[@]}" \ - "https://api.github.com/repos/BerriAI/litellm/releases?per_page=100&page=${page}" \ - >"${PAGE_JSON}" - jq -s '.[0] + .[1]' "${RELEASES_JSON}" "${PAGE_JSON}" >"${RELEASES_JSON}.merged" - mv "${RELEASES_JSON}.merged" "${RELEASES_JSON}" - # No more pages? GitHub returns an empty array past the last page. - if [[ "$(jq 'length' "${PAGE_JSON}")" == "0" ]]; then - break - fi -done -LITELLM_VERSION="$( - jq -r ' - [ .[] | .tag_name // empty - | select(test("^v[0-9]+\\.[0-9]+\\.[0-9]+-stable$")) - ] - | sort_by( - capture("^v(?[0-9]+)\\.(?[0-9]+)\\.(?[0-9]+)-stable$") - | [(.a|tonumber), (.b|tonumber), (.c|tonumber)] - ) - | last // empty - ' "${RELEASES_JSON}" -)" -[[ -n "${LITELLM_VERSION}" ]] || die "could not resolve latest v*-stable tag in 5 pages of releases" -log "resolved litellm: ${LITELLM_VERSION}" - -# The systemd unit loads provider credentials and the agent-shin GitHub -# token from /etc/litellm-compat-matrix.env into this script's -# environment. Running the npm-installed `claude` binary directly here -# would hand that full env to package code -- a compromised -# @anthropic-ai/claude-code release could read ANTHROPIC_API_KEY / -# AWS_BEARER_TOKEN_BEDROCK / AZURE_FOUNDRY_API_KEY / -# AGENT_SHIN_GITHUB_TOKEN from os.environ and exfiltrate them before -# the proxy or test harness ever starts. Probe under `env -i` with the -# same minimal allowlist the PR-gate uses (the matrix run itself goes -# through cli_driver.py, which already scrubs the CLI env). -# -# The probe also runs under a fresh empty HOME instead of the runtime -# user's real $HOME. `ProtectHome=read-only` in the systemd unit -# blocks *writes* to /home/mateo but still allows reads, so a -# compromised claude package invoked here with HOME=/home/mateo could -# read ~/.config/gh/hosts.yml (the gh-host token), ~/.bash_history, -# or ~/.ssh/. Pointing HOME at a per-run dir under ${WORKDIR} hides -# those entirely from the subprocess; ${WORKDIR} is rm -rf'd by the -# script-wide cleanup() trap regardless of probe outcome. -CLAUDE_PROBE_HOME="${WORKDIR}/claude-probe-home" -mkdir -p "${CLAUDE_PROBE_HOME}" -CLAUDE_CODE_VERSION="$(env -i \ - PATH="${PATH}" \ - HOME="${CLAUDE_PROBE_HOME}" \ - USER="${USER:-mateo}" \ - TERM="${TERM:-dumb}" \ - LANG="${LANG:-C.UTF-8}" \ - LC_ALL="${LC_ALL:-}" \ - TMPDIR="${TMPDIR:-/tmp}" \ - claude --version 2>/dev/null \ - | grep -oE '[0-9]+\.[0-9]+\.[0-9]+([.-][A-Za-z0-9.-]+)?' \ - | head -n1 || true)" -# `|| true` above keeps `set -Eeuo pipefail` from aborting silently when -# `grep` finds no match (exit 1) — without it the assignment inherits the -# pipeline's non-zero exit, `set -e` kills the script, and the operator -# never sees the helpful diagnostic below. -[[ -n "${CLAUDE_CODE_VERSION}" ]] || die "could not parse semver from 'claude --version'" -log "local claude code: ${CLAUDE_CODE_VERSION}" - -# --------------------------------------------------------------------------- -# 2. Update the worktree to that tag -# --------------------------------------------------------------------------- - -if [[ ! -d "${WORKTREE}/.git" ]]; then - log "first run: cloning litellm into ${WORKTREE}" - mkdir -p "$(dirname "${WORKTREE}")" - git clone https://github.com/BerriAI/litellm.git "${WORKTREE}" -fi - -log "updating worktree to ${LITELLM_VERSION}" -git -C "${WORKTREE}" fetch --tags --force -git -C "${WORKTREE}" reset --hard -# Keep the venv and the .uv-bin cache around — uv sync will reconcile -# the venv on every run, and we don't want to re-download the pinned -# uv binary each time. Drop everything else (including any prior -# tests/e2e/claude_code/ shim) so each run starts clean before the shim -# below rewrites it from the dev checkout. -git -C "${WORKTREE}" clean -fdx -e .venv -e .uv-bin -git -C "${WORKTREE}" checkout --force "${LITELLM_VERSION}" - -# Always overwrite tests/e2e/claude_code/ in the worktree with the copy -# from the dev checkout, regardless of whether the resolved -# ${LITELLM_VERSION} tag already ships a tests/e2e/claude_code/ tree of -# its own. Rationale: the matrix populator's job is to exercise -# today's tests against the latest stable proxy. The dev checkout -# carries the most recent test fixes (e.g. the stream-json vision -# rewrite, the --effort thinking knob, the WebSearch tool_use -# assertion) that haven't yet rolled into a v*-stable, and we want -# every cron run to pick those up the moment they land on -# ${LITELLM_REPO}, not whenever the next stable release happens. -# -# Concretely this means a fresh `rm -rf` + `cp -r` every run so the -# tree is byte-identical to ${LITELLM_REPO}/tests/e2e/claude_code (no -# stale files left over from the tag's own checkout, no drift across -# runs). -if [[ ! -d "${LITELLM_REPO}/tests/e2e/claude_code" ]]; then - die "no shim source at ${LITELLM_REPO}/tests/e2e/claude_code" -fi -log "shimming tests/e2e/claude_code/ from ${LITELLM_REPO} (always-overwrite)" -rm -rf "${WORKTREE}/tests/e2e/claude_code" -mkdir -p "${WORKTREE}/tests/e2e" -cp -r "${LITELLM_REPO}/tests/e2e/claude_code" "${WORKTREE}/tests/e2e/" - -# litellm pins an exact uv version in pyproject.toml's [tool.uv] -# `required-version` field, so a system uv that's newer or older -# refuses to sync. We pin our own local copy at the version the -# checked-out tag asks for, cached under .uv-bin/ inside the worktree -# so subsequent runs skip the download. -PINNED_UV_VERSION="$( - awk -F'"' ' - /^required-version[[:space:]]*=/ { - # Field 2 is the value between the quotes, e.g. ">=0.10.9" or - # "0.10.9". Strip any leading specifier prefix so we end up with - # the bare version string, which is what /releases/download// - # expects. - v = $2 - sub(/^[[:space:]=<>!~]+/, "", v) - if (v != "") { print v; exit } - } - ' "${WORKTREE}/pyproject.toml" -)" -if [[ -z "${PINNED_UV_VERSION}" ]]; then - log "no uv version pin in pyproject.toml; using system uv" - WORKTREE_UV="$(command -v uv)" -else - WORKTREE_UV="${WORKTREE}/.uv-bin/uv-${PINNED_UV_VERSION}" - if [[ ! -x "${WORKTREE_UV}" ]]; then - log "downloading uv ${PINNED_UV_VERSION} for the worktree" - mkdir -p "${WORKTREE}/.uv-bin" - # Detect host arch so the same script works on x86_64 GCP VMs and on - # aarch64 hosts (Astral publishes both `uv-x86_64-unknown-linux-gnu` - # and `uv-aarch64-unknown-linux-gnu` tarballs under the same release - # tag, and `uname -m` already returns the exact token uv uses). - UV_ARCH="$(uname -m)" - UV_TRIPLE="uv-${UV_ARCH}-unknown-linux-gnu" - UV_TARBALL_NAME="${UV_TRIPLE}.tar.gz" - UV_DOWNLOAD_URL="https://github.com/astral-sh/uv/releases/download/${PINNED_UV_VERSION}/${UV_TARBALL_NAME}" - UV_TMPDIR="$(mktemp -d -t uv-download.XXXXXX)" - # Download the tarball and Astral's official .sha256 sidecar to disk - # and verify the digest before extracting/executing anything. This - # closes the supply-chain trust gap of piping a remote binary - # straight into `tar -xzO ... > file ; chmod +x` (see CLAUDE.md - # "CI Supply-Chain Safety"). - curl -fsSL --output "${UV_TMPDIR}/${UV_TARBALL_NAME}" "${UV_DOWNLOAD_URL}" - curl -fsSL --output "${UV_TMPDIR}/${UV_TARBALL_NAME}.sha256" "${UV_DOWNLOAD_URL}.sha256" - (cd "${UV_TMPDIR}" && sha256sum -c "${UV_TARBALL_NAME}.sha256") \ - || { rm -rf "${UV_TMPDIR}"; die "uv ${PINNED_UV_VERSION} sha256 mismatch — refusing to install"; } - tar -xzf "${UV_TMPDIR}/${UV_TARBALL_NAME}" -C "${UV_TMPDIR}" "${UV_TRIPLE}/uv" - mv "${UV_TMPDIR}/${UV_TRIPLE}/uv" "${WORKTREE_UV}.tmp" - chmod +x "${WORKTREE_UV}.tmp" - mv "${WORKTREE_UV}.tmp" "${WORKTREE_UV}" - rm -rf "${UV_TMPDIR}" - fi -fi -# `--extra proxy` pulls fastapi/uvicorn/etc. so `uv run litellm` can -# actually serve. `--group proxy-dev` brings in pytest and the rest of -# what tests/e2e/claude_code/ needs. -log "uv sync --frozen --group proxy-dev --extra proxy (uv ${PINNED_UV_VERSION:-system})" -(cd "${WORKTREE}" && "${WORKTREE_UV}" sync --frozen --group proxy-dev --extra proxy) - -PROXY_CONFIG="${WORKTREE}/tests/e2e/claude_code/test_config.yaml" -[[ -f "${PROXY_CONFIG}" ]] || die "proxy config not found at ${PROXY_CONFIG} (does ${LITELLM_VERSION} predate the compat matrix work?)" - -# --------------------------------------------------------------------------- -# 3. Boot the proxy -# --------------------------------------------------------------------------- - -log "starting proxy on 127.0.0.1:${PROXY_PORT}" -# Bind the proxy to loopback only. The populator proxy is talked to -# exclusively by the pytest run on the same host (the health check and -# the test env set `LITELLM_PROXY_BASE_URL=http://127.0.0.1:...`), -# so there's no reason to expose it on the VM's external interfaces. -# Without `--host`, `litellm` defaults to 0.0.0.0, which combined with -# the predictable default `LITELLM_MASTER_KEY=sk-cron-matrix` would -# allow anything that can reach :${PROXY_PORT} on the VM to authenticate -# and burn upstream provider credentials. -# -# `setsid` puts the proxy in its own session+pgroup so cleanup() can -# SIGTERM the whole tree by passing the pgid as a negative pid. We -# write that pid to a file so cleanup() doesn't need to remember a -# variable that might be stale by the time the trap fires. -# -# Pass the master key as a shell-prefix assignment on `setsid` (inherited -# via the environment) rather than as `env KEY=VAL ...` argv. The argv -# form would land the literal key in /proc//cmdline, where -# any local reader (a model-directed `Read` tool call, another user on -# the VM, a crash dump) could pick it up before the process execs into -# the litellm child. The shell-prefix form keeps the key out of argv at -# every layer (setsid → bash → uv → litellm). -LITELLM_MASTER_KEY="${PROXY_API_KEY}" setsid bash -c ' - echo "$$" > "$0" - cd "$1" - exec "$2" run litellm --config "$3" --host 127.0.0.1 --port "$4" -' "${PROXY_PID_FILE}" "${WORKTREE}" "${WORKTREE_UV}" "${PROXY_CONFIG}" "${PROXY_PORT}" \ - >"${WORKDIR}/proxy.log" 2>&1 & -disown - -HEALTH_URL="http://127.0.0.1:${PROXY_PORT}/health/liveliness" -for _ in $(seq 1 45); do - if curl -fsS "${HEALTH_URL}" >/dev/null 2>&1; then - break - fi - sleep 2 -done -curl -fsS "${HEALTH_URL}" >/dev/null \ - || { tail -50 "${WORKDIR}/proxy.log" >&2; die "proxy did not become healthy"; } - -# --------------------------------------------------------------------------- -# 4. Run pytest -# --------------------------------------------------------------------------- - -RESULTS_JSON="${WORKDIR}/compat-results.json" -PYTEST_ARGS=( - tests/e2e/claude_code/ - --ignore=tests/e2e/claude_code/_driver_unit_tests - --ignore=tests/e2e/claude_code/_builder_unit_tests - --ignore=tests/e2e/claude_code/_publisher_unit_tests - --ignore=tests/e2e/claude_code/_pr_gate_unit_tests -) -if [[ -n "${PYTEST_K}" ]]; then - log "PYTEST_K set; narrowing to: ${PYTEST_K}" - PYTEST_ARGS+=(-k "${PYTEST_K}") -fi - -log "running pytest" -set +e -# Pytest only needs to talk to the loopback proxy at 127.0.0.1:${PROXY_PORT} -# — it has no legitimate reason to see ANTHROPIC_API_KEY / -# AWS_BEARER_TOKEN_BEDROCK / VERTEXAI_* / AZURE_FOUNDRY_* / -# AGENT_SHIN_GITHUB_TOKEN / GITHUB_TOKEN in its own env. The systemd -# unit's EnvironmentFile injects all of those into this script for the -# proxy to consume, and pytest inherits them by default. Wrap the -# invocation in `env -i` so: -# -# 1. test code under tests/e2e/claude_code/ (or anything it imports) -# cannot read provider/agent-shin creds out of `os.environ` and -# exfiltrate them via an outbound call from inside a conftest hook -# or a fixture (a sibling vector to the model-controlled Bash/Read -# concern handled by `cli_driver.py`'s own env scrub); -# 2. a model-directed `Read` tool call during a PDF/vision cell -# cannot reach /proc//environ and pull the creds out -# of the parent process the way it can today; -# 3. this matches the PR-gate pytest step in `.circleci/config.yml`, -# which already runs under `env -i` with the same minimal -# allowlist. -# -# `cli_driver.py` re-allowlists its own subset (PATH/USER/LOGNAME/etc.) -# when spawning the `claude` binary, so the CLI still finds Node + the -# claude shim on PATH and gets a fresh isolated HOME per invocation. -( - cd "${WORKTREE}" \ - && env -i \ - PATH="${PATH}" \ - HOME="${HOME}" \ - USER="${USER:-mateo}" \ - TERM="${TERM:-dumb}" \ - LANG="${LANG:-C.UTF-8}" \ - LC_ALL="${LC_ALL:-}" \ - TMPDIR="${TMPDIR:-/tmp}" \ - LITELLM_PROXY_BASE_URL="http://127.0.0.1:${PROXY_PORT}" \ - LITELLM_PROXY_API_KEY="${PROXY_API_KEY}" \ - COMPAT_RESULTS_PATH="${RESULTS_JSON}" \ - "${WORKTREE_UV}" run pytest "${PYTEST_ARGS[@]}" -) -PYTEST_EXIT=$? -set -e -log "pytest exit code: ${PYTEST_EXIT} (failures become 'fail' cells, not script errors)" -[[ -f "${RESULTS_JSON}" ]] || die "pytest did not produce ${RESULTS_JSON}" - -# --------------------------------------------------------------------------- -# 5. Build the matrix JSON -# --------------------------------------------------------------------------- - -MATRIX_JSON="${WORKDIR}/compatibility-matrix.json" -log "building ${MATRIX_JSON}" -( - cd "${WORKTREE}" \ - && "${WORKTREE_UV}" run python "${POPULATOR_DIR}/build_matrix.py" \ - --manifest "${WORKTREE}/tests/e2e/claude_code/manifest.yaml" \ - --results "${RESULTS_JSON}" \ - --output "${MATRIX_JSON}" \ - --litellm-version "${LITELLM_VERSION}" \ - --claude-code-version "${CLAUDE_CODE_VERSION}" -) - -# --------------------------------------------------------------------------- -# 6. Open a docs-repo PR -# --------------------------------------------------------------------------- - -if [[ "${SKIP_PUBLISH}" == "1" ]]; then - cp "${MATRIX_JSON}" "${LITELLM_REPO}/compatibility-matrix.json" - log "SKIP_PUBLISH=1; matrix written to ${LITELLM_REPO}/compatibility-matrix.json" - exit 0 -fi - -DATE_UTC="$(date -u +%Y-%m-%d)" -BRANCH_NAME="compat-matrix/${LITELLM_VERSION}-${CLAUDE_CODE_VERSION}-${DATE_UTC}" -DOCS_CLONE="${WORKDIR}/litellm-docs" -FORK_OWNER="${FORK_OWNER:-agent-shin}" -FORK_REPO="${FORK_REPO:-${FORK_OWNER}/litellm-docs}" - -log "cloning ${DOCS_REPO}@${DOCS_BRANCH}" -# Use the agent-shin token inline rather than the host gh-cli config. -# `BerriAI/litellm-docs` is a public repo so unauthenticated clone -# would also work, but passing the token explicitly means the systemd -# unit can hide `~/.config/gh` (`InaccessiblePaths=`) without breaking -# this clone — closing the model-directed `Read("/home/mateo/.config/gh/...")` -# exfiltration path on the cron VM. -GH_TOKEN="${AGENT_SHIN_GITHUB_TOKEN}" \ - gh repo clone "${DOCS_REPO}" "${DOCS_CLONE}" -- --depth 1 --branch "${DOCS_BRANCH}" - -cd "${DOCS_CLONE}" -git config user.email "litellm-bot@berri.ai" -git config user.name "litellm-compat-matrix-bot" -git checkout -b "${BRANCH_NAME}" - -mkdir -p "$(dirname "${DOCS_TARGET_PATH}")" -cp "${MATRIX_JSON}" "${DOCS_TARGET_PATH}" -git add "${DOCS_TARGET_PATH}" - -if git diff --cached --quiet; then - log "matrix JSON unchanged from ${DOCS_BRANCH}; skipping PR" - exit 0 -fi - -GENERATED_AT="$(jq -r '.generated_at' "${MATRIX_JSON}")" -COMMIT_MSG="$(cat </dev/null || true -git remote add fork "${FORK_PUSH_URL}" -git push --force --set-upstream fork "${BRANCH_NAME}" -git remote remove fork -unset FORK_PUSH_URL - -# Per-feature status table for the PR body. Reviewers triage from this. -PR_FEATURE_TABLE="$(jq -r ' - .features[] as $f - | "- **\($f.name)**: " + - ([ .providers[] as $p - | "\($p)=\($f.providers[$p].status // "not_tested")" - ] | join(", ")) -' "${MATRIX_JSON}")" - -PR_TITLE="chore(compat-matrix): refresh for ${LITELLM_VERSION} + claude-code ${CLAUDE_CODE_VERSION}" -PR_BODY="$(cat < ${DOCS_REPO}:${DOCS_BRANCH}" -# GH_TOKEN here is scoped to this single subshell so we don't bleed the -# fork token into the rest of the script (release-listing earlier uses -# ${GITHUB_TOKEN}, which may be a different identity). gh's --head accepts -# `OWNER:BRANCH` for cross-repo PRs from a fork. -# -# Reviewer assignment is done in a *separate* call below: as the PR -# author from a fork, agent-shin has no write/triage access on -# ${DOCS_REPO} and the `RequestReviewsByLogin` GraphQL mutation -# (which backs `gh pr create --reviewer` and `gh pr edit --add-reviewer`) -# rejects with "does not have the correct permissions". We use the -# collaborator-scoped ${GITHUB_TOKEN} for that instead. Don't fold -# --reviewer into `gh pr create` here -- it would fail the whole -# create on the very first cron run. -set +e -PR_OUT="$( - GH_TOKEN="${AGENT_SHIN_GITHUB_TOKEN}" gh pr create \ - --repo "${DOCS_REPO}" \ - --base "${DOCS_BRANCH}" \ - --head "${FORK_OWNER}:${BRANCH_NAME}" \ - --title "${PR_TITLE}" \ - --body "${PR_BODY}" 2>&1 -)" -PR_EXIT=$? -set -e -echo "${PR_OUT}" - -if [[ ${PR_EXIT} -ne 0 ]]; then - if grep -q "a pull request for branch.*already exists" <<<"${PR_OUT}"; then - log "PR already exists for ${FORK_OWNER}:${BRANCH_NAME}; updated branch in place" - else - die "gh pr create failed (exit ${PR_EXIT})" - fi -fi - -# Request reviews from PR_REVIEWERS using the collaborator-scoped -# ${GITHUB_TOKEN} (mateo-berri's token, already provisioned for release -# listing). This is idempotent: `gh pr edit --add-reviewer` is a no-op -# on a user who's already in reviewRequests, and silently re-adds -# anyone whose prior review was dismissed -- so same-day reruns stay -# clean. Reviewer-add failures are non-fatal: the matrix JSON has -# already landed on the PR; the worst case is a manual ping. -if [[ -n "${PR_REVIEWERS}" ]]; then - if [[ -z "${GITHUB_TOKEN:-}" ]]; then - log "WARN: PR_REVIEWERS set but GITHUB_TOKEN missing -- cannot request reviews; skipping" - else - log "requesting reviews from: ${PR_REVIEWERS}" - set +e - GH_TOKEN="${GITHUB_TOKEN}" gh pr edit \ - "${FORK_OWNER}:${BRANCH_NAME}" \ - --repo "${DOCS_REPO}" \ - --add-reviewer "${PR_REVIEWERS}" 2>&1 | sed 's/^/ /' - REVIEWER_EXIT=${PIPESTATUS[0]} - set -e - if [[ ${REVIEWER_EXIT} -ne 0 ]]; then - log "WARN: gh pr edit --add-reviewer exited ${REVIEWER_EXIT} (non-fatal)" - fi - fi -fi - -log "done" diff --git a/tests/e2e/claude_code/long_context_1m/test_anthropic.py b/tests/e2e/claude_code/long_context_1m/test_anthropic.py index fb74d5fd40d..b9bbd1c2fe7 100644 --- a/tests/e2e/claude_code/long_context_1m/test_anthropic.py +++ b/tests/e2e/claude_code/long_context_1m/test_anthropic.py @@ -54,25 +54,23 @@ $10/day on this row. from __future__ import annotations -import os from typing import 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" # Haiku 4.5 is excluded -- only Sonnet 4.6 and Opus 4.7 support the # 1M-context beta. See module docstring for the per-cell-aggregator # rationale. ANTHROPIC_MODELS: Sequence[str] = ( - "claude-sonnet-4-6", + "claude-sonnet-4-5", "claude-opus-4-7", ) @@ -155,26 +153,12 @@ def _build_long_prompt(target_tokens: int = TARGET_INPUT_TOKENS) -> str: return preamble + "".join(pad_lines) + closing +@pytest.mark.covers("llm.messages.anthropic.long_context_1m.nonstream.works") def test_long_context_1m_anthropic(compat_result): """Drive the `claude` CLI with a ~210k-token prompt and the `context-1m-2025-08-07` beta header; assert no 400 / 413 and a non-empty reply for Sonnet + Opus.""" - 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) long_prompt = _build_long_prompt() diff --git a/tests/e2e/claude_code/long_context_1m/test_azure.py b/tests/e2e/claude_code/long_context_1m/test_azure.py index 5800fdadbfc..d62214d2758 100644 --- a/tests/e2e/claude_code/long_context_1m/test_azure.py +++ b/tests/e2e/claude_code/long_context_1m/test_azure.py @@ -54,25 +54,23 @@ $10/day on this row. from __future__ import annotations -import os from typing import 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" # Haiku 4.5 is excluded -- only Sonnet 4.6 and Opus 4.7 support the # 1M-context beta. See module docstring for the per-cell-aggregator # rationale. AZURE_MODELS: Sequence[str] = ( - "claude-sonnet-4-6-azure", + "claude-sonnet-4-5-azure", "claude-opus-4-7-azure", ) @@ -155,26 +153,12 @@ def _build_long_prompt(target_tokens: int = TARGET_INPUT_TOKENS) -> str: return preamble + "".join(pad_lines) + closing +@pytest.mark.covers("llm.messages.azure_foundry.long_context_1m.nonstream.works") def test_long_context_1m_azure(compat_result): """Drive the `claude` CLI (Azure (Microsoft Foundry)) with a ~210k-token prompt and the `context-1m-2025-08-07` beta header; assert no 400 / 413 and a non-empty reply for Sonnet + Opus.""" - 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) long_prompt = _build_long_prompt() diff --git a/tests/e2e/claude_code/long_context_1m/test_bedrock_converse.py b/tests/e2e/claude_code/long_context_1m/test_bedrock_converse.py index 18587f7c2d6..3c2fd4f02cc 100644 --- a/tests/e2e/claude_code/long_context_1m/test_bedrock_converse.py +++ b/tests/e2e/claude_code/long_context_1m/test_bedrock_converse.py @@ -54,25 +54,23 @@ $10/day on this row. from __future__ import annotations -import os from typing import 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" # Haiku 4.5 is excluded -- only Sonnet 4.6 and Opus 4.7 support the # 1M-context beta. See module docstring for the per-cell-aggregator # rationale. BEDROCK_CONVERSE_MODELS: Sequence[str] = ( - "claude-sonnet-4-6-bedrock-converse", + "claude-sonnet-4-5-bedrock-converse", "claude-opus-4-7-bedrock-converse", ) @@ -155,26 +153,12 @@ def _build_long_prompt(target_tokens: int = TARGET_INPUT_TOKENS) -> str: return preamble + "".join(pad_lines) + closing +@pytest.mark.covers("llm.messages.bedrock_converse.long_context_1m.nonstream.works") def test_long_context_1m_bedrock_converse(compat_result): """Drive the `claude` CLI (Bedrock (Converse)) with a ~210k-token prompt and the `context-1m-2025-08-07` beta header; assert no 400 / 413 and a non-empty reply for Sonnet + Opus.""" - 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) long_prompt = _build_long_prompt() diff --git a/tests/e2e/claude_code/long_context_1m/test_bedrock_invoke.py b/tests/e2e/claude_code/long_context_1m/test_bedrock_invoke.py index 0270197ce2a..4801d405760 100644 --- a/tests/e2e/claude_code/long_context_1m/test_bedrock_invoke.py +++ b/tests/e2e/claude_code/long_context_1m/test_bedrock_invoke.py @@ -54,25 +54,23 @@ $10/day on this row. from __future__ import annotations -import os from typing import 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" # Haiku 4.5 is excluded -- only Sonnet 4.6 and Opus 4.7 support the # 1M-context beta. See module docstring for the per-cell-aggregator # rationale. BEDROCK_INVOKE_MODELS: Sequence[str] = ( - "claude-sonnet-4-6-bedrock-invoke", + "claude-sonnet-4-5-bedrock-invoke", "claude-opus-4-7-bedrock-invoke", ) @@ -155,26 +153,12 @@ def _build_long_prompt(target_tokens: int = TARGET_INPUT_TOKENS) -> str: return preamble + "".join(pad_lines) + closing +@pytest.mark.covers("llm.messages.bedrock_invoke.long_context_1m.nonstream.works") def test_long_context_1m_bedrock_invoke(compat_result): """Drive the `claude` CLI (Bedrock (Invoke)) with a ~210k-token prompt and the `context-1m-2025-08-07` beta header; assert no 400 / 413 and a non-empty reply for Sonnet + Opus.""" - 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) long_prompt = _build_long_prompt() diff --git a/tests/e2e/claude_code/long_context_1m/test_vertex_ai.py b/tests/e2e/claude_code/long_context_1m/test_vertex_ai.py index d2db4a1b4ee..efa96bf076d 100644 --- a/tests/e2e/claude_code/long_context_1m/test_vertex_ai.py +++ b/tests/e2e/claude_code/long_context_1m/test_vertex_ai.py @@ -54,25 +54,23 @@ $10/day on this row. from __future__ import annotations -import os from typing import 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" # Haiku 4.5 is excluded -- only Sonnet 4.6 and Opus 4.7 support the # 1M-context beta. See module docstring for the per-cell-aggregator # rationale. VERTEX_AI_MODELS: Sequence[str] = ( - "claude-sonnet-4-6-vertex", + "claude-sonnet-4-5-vertex", "claude-opus-4-7-vertex", ) @@ -155,26 +153,12 @@ def _build_long_prompt(target_tokens: int = TARGET_INPUT_TOKENS) -> str: return preamble + "".join(pad_lines) + closing +@pytest.mark.covers("llm.messages.vertex.long_context_1m.nonstream.works") def test_long_context_1m_vertex_ai(compat_result): """Drive the `claude` CLI (Vertex AI) with a ~210k-token prompt and the `context-1m-2025-08-07` beta header; assert no 400 / 413 and a non-empty reply for Sonnet + Opus.""" - 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) long_prompt = _build_long_prompt() diff --git a/tests/e2e/claude_code/matrix_builder.py b/tests/e2e/claude_code/matrix_builder.py index 5641e488da2..d9a13d17ea4 100644 --- a/tests/e2e/claude_code/matrix_builder.py +++ b/tests/e2e/claude_code/matrix_builder.py @@ -183,7 +183,7 @@ def build_from_paths( generated_at: str, output_path: Optional[Path] = None, ) -> Dict[str, Any]: - """I/O wrapper around build_matrix used by the publisher script.""" + """I/O wrapper around ``build_matrix``: reads the manifest and per-test results from disk, calls ``build_matrix``, and (optionally) writes the compat-matrix JSON to ``output_path``. Whatever orchestrator publishes the matrix (currently the ECR image) invokes this.""" manifest = load_manifest(manifest_path) results = load_results(results_path) matrix = build_matrix( diff --git a/tests/e2e/claude_code/passthrough/test_anthropic.py b/tests/e2e/claude_code/passthrough/test_anthropic.py index aa0443e0625..8382342ae12 100644 --- a/tests/e2e/claude_code/passthrough/test_anthropic.py +++ b/tests/e2e/claude_code/passthrough/test_anthropic.py @@ -29,7 +29,7 @@ from claude_code._passthrough import ( ANTHROPIC_MODELS = [ "claude-haiku-4-5", - "claude-sonnet-4-6", + "claude-sonnet-4-5", "claude-opus-4-7", ] diff --git a/tests/e2e/claude_code/passthrough/test_azure.py b/tests/e2e/claude_code/passthrough/test_azure.py index 09b0824047a..21100a49c16 100644 --- a/tests/e2e/claude_code/passthrough/test_azure.py +++ b/tests/e2e/claude_code/passthrough/test_azure.py @@ -44,7 +44,7 @@ from claude_code._passthrough import foundry_extra_env, run_passthrough_cell AZURE_MODELS = [ "claude-haiku-4-5", - "claude-sonnet-4-6", + "claude-sonnet-4-5", "claude-opus-4-7", ] diff --git a/tests/e2e/claude_code/passthrough/test_bedrock_invoke.py b/tests/e2e/claude_code/passthrough/test_bedrock_invoke.py index 6e84dea6779..f1f28ab5b4c 100644 --- a/tests/e2e/claude_code/passthrough/test_bedrock_invoke.py +++ b/tests/e2e/claude_code/passthrough/test_bedrock_invoke.py @@ -27,7 +27,7 @@ from claude_code._passthrough import bedrock_extra_env, run_passthrough_cell BEDROCK_INVOKE_MODELS = [ "claude-haiku-4-5-bedrock-invoke", - "claude-sonnet-4-6-bedrock-invoke", + "claude-sonnet-4-5-bedrock-invoke", "claude-opus-4-7-bedrock-invoke", ] diff --git a/tests/e2e/claude_code/passthrough/test_vertex_ai.py b/tests/e2e/claude_code/passthrough/test_vertex_ai.py index 5e3c6bce419..790f8b60c8f 100644 --- a/tests/e2e/claude_code/passthrough/test_vertex_ai.py +++ b/tests/e2e/claude_code/passthrough/test_vertex_ai.py @@ -30,7 +30,7 @@ from claude_code._passthrough import run_passthrough_cell, vertex_extra_env VERTEX_MODELS = [ "claude-haiku-4-5-vertex", - "claude-sonnet-4-6-vertex", + "claude-sonnet-4-5-vertex", "claude-opus-4-7-vertex", ] diff --git a/tests/e2e/claude_code/pdf_input/test_anthropic.py b/tests/e2e/claude_code/pdf_input/test_anthropic.py index 36fb69a1db6..21c8028ef1c 100644 --- a/tests/e2e/claude_code/pdf_input/test_anthropic.py +++ b/tests/e2e/claude_code/pdf_input/test_anthropic.py @@ -22,22 +22,19 @@ The (feature, provider) for this cell is inferred from the file path by from __future__ import annotations -import os - 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_MODELS = [ "claude-haiku-4-5", - "claude-sonnet-4-6", + "claude-sonnet-4-5", "claude-opus-4-7", ] @@ -108,24 +105,11 @@ def _build_minimal_pdf(marker: str) -> bytes: return bytes(out) +@pytest.mark.covers("llm.messages.anthropic.pdf_input.nonstream.works") def test_pdf_input_anthropic(compat_result, tmp_path): """Drive the `claude` CLI against the LiteLLM proxy with a PDF attached via the Read tool and assert the reply references it.""" - 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) pdf_path = tmp_path / "marker.pdf" pdf_path.write_bytes(_build_minimal_pdf(PDF_MARKER)) diff --git a/tests/e2e/claude_code/pdf_input/test_azure.py b/tests/e2e/claude_code/pdf_input/test_azure.py index 810c857e407..34ae3732b99 100644 --- a/tests/e2e/claude_code/pdf_input/test_azure.py +++ b/tests/e2e/claude_code/pdf_input/test_azure.py @@ -15,22 +15,19 @@ The (feature, provider) for this cell is inferred from the file path by from __future__ import annotations -import os - 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" AZURE_MODELS = [ "claude-haiku-4-5-azure", - "claude-sonnet-4-6-azure", + "claude-sonnet-4-5-azure", "claude-opus-4-7-azure", ] @@ -85,22 +82,9 @@ def _build_minimal_pdf(marker: str) -> bytes: return bytes(out) +@pytest.mark.covers("llm.messages.azure_foundry.pdf_input.nonstream.works") def test_pdf_input_azure(compat_result, tmp_path): - 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) pdf_path = tmp_path / "marker.pdf" pdf_path.write_bytes(_build_minimal_pdf(PDF_MARKER)) diff --git a/tests/e2e/claude_code/pdf_input/test_bedrock_converse.py b/tests/e2e/claude_code/pdf_input/test_bedrock_converse.py index 191a27c6d46..76aa84f0f47 100644 --- a/tests/e2e/claude_code/pdf_input/test_bedrock_converse.py +++ b/tests/e2e/claude_code/pdf_input/test_bedrock_converse.py @@ -21,22 +21,19 @@ The (feature, provider) for this cell is inferred from the file path by from __future__ import annotations -import os - 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" BEDROCK_CONVERSE_MODELS = [ "claude-haiku-4-5-bedrock-converse", - "claude-sonnet-4-6-bedrock-converse", + "claude-sonnet-4-5-bedrock-converse", "claude-opus-4-7-bedrock-converse", ] @@ -91,22 +88,9 @@ def _build_minimal_pdf(marker: str) -> bytes: return bytes(out) +@pytest.mark.covers("llm.messages.bedrock_converse.pdf_input.nonstream.works") def test_pdf_input_bedrock_converse(compat_result, tmp_path): - 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) pdf_path = tmp_path / "marker.pdf" pdf_path.write_bytes(_build_minimal_pdf(PDF_MARKER)) diff --git a/tests/e2e/claude_code/pdf_input/test_bedrock_invoke.py b/tests/e2e/claude_code/pdf_input/test_bedrock_invoke.py index 163cabb45a0..4450266bb6b 100644 --- a/tests/e2e/claude_code/pdf_input/test_bedrock_invoke.py +++ b/tests/e2e/claude_code/pdf_input/test_bedrock_invoke.py @@ -20,22 +20,19 @@ The (feature, provider) for this cell is inferred from the file path by from __future__ import annotations -import os - 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" BEDROCK_INVOKE_MODELS = [ "claude-haiku-4-5-bedrock-invoke", - "claude-sonnet-4-6-bedrock-invoke", + "claude-sonnet-4-5-bedrock-invoke", "claude-opus-4-7-bedrock-invoke", ] @@ -90,22 +87,9 @@ def _build_minimal_pdf(marker: str) -> bytes: return bytes(out) +@pytest.mark.covers("llm.messages.bedrock_invoke.pdf_input.nonstream.works") def test_pdf_input_bedrock_invoke(compat_result, tmp_path): - 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) pdf_path = tmp_path / "marker.pdf" pdf_path.write_bytes(_build_minimal_pdf(PDF_MARKER)) diff --git a/tests/e2e/claude_code/pdf_input/test_vertex_ai.py b/tests/e2e/claude_code/pdf_input/test_vertex_ai.py index 0d0573d05b3..b78f58cfda1 100644 --- a/tests/e2e/claude_code/pdf_input/test_vertex_ai.py +++ b/tests/e2e/claude_code/pdf_input/test_vertex_ai.py @@ -15,22 +15,19 @@ The (feature, provider) for this cell is inferred from the file path by from __future__ import annotations -import os - 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" VERTEX_AI_MODELS = [ "claude-haiku-4-5-vertex", - "claude-sonnet-4-6-vertex", + "claude-sonnet-4-5-vertex", "claude-opus-4-7-vertex", ] @@ -85,22 +82,9 @@ def _build_minimal_pdf(marker: str) -> bytes: return bytes(out) +@pytest.mark.covers("llm.messages.vertex.pdf_input.nonstream.works") def test_pdf_input_vertex_ai(compat_result, tmp_path): - 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) pdf_path = tmp_path / "marker.pdf" pdf_path.write_bytes(_build_minimal_pdf(PDF_MARKER)) diff --git a/tests/e2e/claude_code/prompt_caching_1h/test_anthropic.py b/tests/e2e/claude_code/prompt_caching_1h/test_anthropic.py index d81887231d8..637be1c551d 100644 --- a/tests/e2e/claude_code/prompt_caching_1h/test_anthropic.py +++ b/tests/e2e/claude_code/prompt_caching_1h/test_anthropic.py @@ -24,23 +24,21 @@ The (feature, provider) for this cell is inferred from the file path by from __future__ import annotations -import os from typing import Any, Mapping, Optional 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_MODELS = [ "claude-haiku-4-5", - "claude-sonnet-4-6", + "claude-sonnet-4-5", "claude-opus-4-7", ] @@ -65,25 +63,12 @@ def _cache_tokens(usage: Optional[Mapping[str, Any]]) -> int: return 0 +@pytest.mark.covers("llm.messages.anthropic.prompt_cache_1h.nonstream.works") def test_prompt_caching_1h_anthropic(compat_result): """Drive the `claude` CLI against the LiteLLM proxy with the 1h TTL opt-in env var set, and assert the upstream usage block surfaces a non-zero cache token count.""" - 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) outcomes = run_claude_models_parallel( models=ANTHROPIC_MODELS, diff --git a/tests/e2e/claude_code/prompt_caching_1h/test_azure.py b/tests/e2e/claude_code/prompt_caching_1h/test_azure.py index 416757f8691..f34557b3c5f 100644 --- a/tests/e2e/claude_code/prompt_caching_1h/test_azure.py +++ b/tests/e2e/claude_code/prompt_caching_1h/test_azure.py @@ -15,23 +15,21 @@ The (feature, provider) for this cell is inferred from the file path by from __future__ import annotations -import os from typing import Any, Mapping, Optional 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" AZURE_MODELS = [ "claude-haiku-4-5-azure", - "claude-sonnet-4-6-azure", + "claude-sonnet-4-5-azure", "claude-opus-4-7-azure", ] @@ -49,22 +47,9 @@ def _cache_tokens(usage: Optional[Mapping[str, Any]]) -> int: return 0 +@pytest.mark.covers("llm.messages.azure_foundry.prompt_cache_1h.nonstream.works") def test_prompt_caching_1h_azure(compat_result): - 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) outcomes = run_claude_models_parallel( models=AZURE_MODELS, diff --git a/tests/e2e/claude_code/prompt_caching_1h/test_bedrock_converse.py b/tests/e2e/claude_code/prompt_caching_1h/test_bedrock_converse.py index 5bc632c6f1b..bf62a49444c 100644 --- a/tests/e2e/claude_code/prompt_caching_1h/test_bedrock_converse.py +++ b/tests/e2e/claude_code/prompt_caching_1h/test_bedrock_converse.py @@ -20,23 +20,21 @@ The (feature, provider) for this cell is inferred from the file path by from __future__ import annotations -import os from typing import Any, Mapping, Optional 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" BEDROCK_CONVERSE_MODELS = [ "claude-haiku-4-5-bedrock-converse", - "claude-sonnet-4-6-bedrock-converse", + "claude-sonnet-4-5-bedrock-converse", "claude-opus-4-7-bedrock-converse", ] @@ -57,22 +55,9 @@ def _cache_tokens(usage: Optional[Mapping[str, Any]]) -> int: return 0 +@pytest.mark.covers("llm.messages.bedrock_converse.prompt_cache_1h.nonstream.works") def test_prompt_caching_1h_bedrock_converse(compat_result): - 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) outcomes = run_claude_models_parallel( models=BEDROCK_CONVERSE_MODELS, diff --git a/tests/e2e/claude_code/prompt_caching_1h/test_bedrock_invoke.py b/tests/e2e/claude_code/prompt_caching_1h/test_bedrock_invoke.py index 4501834956b..dc3468702d4 100644 --- a/tests/e2e/claude_code/prompt_caching_1h/test_bedrock_invoke.py +++ b/tests/e2e/claude_code/prompt_caching_1h/test_bedrock_invoke.py @@ -22,23 +22,21 @@ The (feature, provider) for this cell is inferred from the file path by from __future__ import annotations -import os from typing import Any, Mapping, Optional 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" BEDROCK_INVOKE_MODELS = [ "claude-haiku-4-5-bedrock-invoke", - "claude-sonnet-4-6-bedrock-invoke", + "claude-sonnet-4-5-bedrock-invoke", "claude-opus-4-7-bedrock-invoke", ] @@ -61,22 +59,9 @@ def _cache_tokens(usage: Optional[Mapping[str, Any]]) -> int: return 0 +@pytest.mark.covers("llm.messages.bedrock_invoke.prompt_cache_1h.nonstream.works") def test_prompt_caching_1h_bedrock_invoke(compat_result): - 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) outcomes = run_claude_models_parallel( models=BEDROCK_INVOKE_MODELS, diff --git a/tests/e2e/claude_code/prompt_caching_1h/test_vertex_ai.py b/tests/e2e/claude_code/prompt_caching_1h/test_vertex_ai.py index 09ded634b45..66cf961fcfc 100644 --- a/tests/e2e/claude_code/prompt_caching_1h/test_vertex_ai.py +++ b/tests/e2e/claude_code/prompt_caching_1h/test_vertex_ai.py @@ -15,23 +15,21 @@ The (feature, provider) for this cell is inferred from the file path by from __future__ import annotations -import os from typing import Any, Mapping, Optional 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" VERTEX_AI_MODELS = [ "claude-haiku-4-5-vertex", - "claude-sonnet-4-6-vertex", + "claude-sonnet-4-5-vertex", "claude-opus-4-7-vertex", ] @@ -49,22 +47,9 @@ def _cache_tokens(usage: Optional[Mapping[str, Any]]) -> int: return 0 +@pytest.mark.covers("llm.messages.vertex.prompt_cache_1h.nonstream.works") def test_prompt_caching_1h_vertex_ai(compat_result): - 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) outcomes = run_claude_models_parallel( models=VERTEX_AI_MODELS, diff --git a/tests/e2e/claude_code/prompt_caching_5m/test_anthropic.py b/tests/e2e/claude_code/prompt_caching_5m/test_anthropic.py index 4b20a65f31b..ef551beb45c 100644 --- a/tests/e2e/claude_code/prompt_caching_5m/test_anthropic.py +++ b/tests/e2e/claude_code/prompt_caching_5m/test_anthropic.py @@ -22,23 +22,21 @@ The (feature, provider) for this cell is inferred from the file path by from __future__ import annotations -import os from typing import Any, Mapping, Optional 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_MODELS = [ "claude-haiku-4-5", - "claude-sonnet-4-6", + "claude-sonnet-4-5", "claude-opus-4-7", ] @@ -56,24 +54,11 @@ def _cache_tokens(usage: Optional[Mapping[str, Any]]) -> int: return 0 +@pytest.mark.covers("llm.messages.anthropic.prompt_cache_5m.nonstream.works") def test_prompt_caching_5m_anthropic(compat_result): """Drive the `claude` CLI against the LiteLLM proxy and assert the upstream usage block surfaces a non-zero cache token count.""" - 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) outcomes = run_claude_models_parallel( models=ANTHROPIC_MODELS, diff --git a/tests/e2e/claude_code/prompt_caching_5m/test_azure.py b/tests/e2e/claude_code/prompt_caching_5m/test_azure.py index 22bd5aa7048..9d4137e0726 100644 --- a/tests/e2e/claude_code/prompt_caching_5m/test_azure.py +++ b/tests/e2e/claude_code/prompt_caching_5m/test_azure.py @@ -22,23 +22,21 @@ The (feature, provider) for this cell is inferred from the file path by from __future__ import annotations -import os from typing import Any, Mapping, Optional 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" AZURE_MODELS = [ "claude-haiku-4-5-azure", - "claude-sonnet-4-6-azure", + "claude-sonnet-4-5-azure", "claude-opus-4-7-azure", ] @@ -54,24 +52,11 @@ def _cache_tokens(usage: Optional[Mapping[str, Any]]) -> int: return 0 +@pytest.mark.covers("llm.messages.azure_foundry.prompt_cache_5m.nonstream.works") def test_prompt_caching_5m_azure(compat_result): """Drive the `claude` CLI against the LiteLLM proxy and assert the upstream usage block surfaces a non-zero cache token count.""" - 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) outcomes = run_claude_models_parallel( models=AZURE_MODELS, diff --git a/tests/e2e/claude_code/prompt_caching_5m/test_bedrock_converse.py b/tests/e2e/claude_code/prompt_caching_5m/test_bedrock_converse.py index 681a6ecce10..c9b34c010b0 100644 --- a/tests/e2e/claude_code/prompt_caching_5m/test_bedrock_converse.py +++ b/tests/e2e/claude_code/prompt_caching_5m/test_bedrock_converse.py @@ -15,23 +15,21 @@ The (feature, provider) for this cell is inferred from the file path by from __future__ import annotations -import os from typing import Any, Mapping, Optional 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" BEDROCK_CONVERSE_MODELS = [ "claude-haiku-4-5-bedrock-converse", - "claude-sonnet-4-6-bedrock-converse", + "claude-sonnet-4-5-bedrock-converse", "claude-opus-4-7-bedrock-converse", ] @@ -47,24 +45,11 @@ def _cache_tokens(usage: Optional[Mapping[str, Any]]) -> int: return 0 +@pytest.mark.covers("llm.messages.bedrock_converse.prompt_cache_5m.nonstream.works") def test_prompt_caching_5m_bedrock_converse(compat_result): """Drive the `claude` CLI against the LiteLLM proxy and assert the upstream usage block surfaces a non-zero cache token count.""" - 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) outcomes = run_claude_models_parallel( models=BEDROCK_CONVERSE_MODELS, diff --git a/tests/e2e/claude_code/prompt_caching_5m/test_bedrock_invoke.py b/tests/e2e/claude_code/prompt_caching_5m/test_bedrock_invoke.py index f1a3109b3a1..b95c509ba3c 100644 --- a/tests/e2e/claude_code/prompt_caching_5m/test_bedrock_invoke.py +++ b/tests/e2e/claude_code/prompt_caching_5m/test_bedrock_invoke.py @@ -15,23 +15,21 @@ The (feature, provider) for this cell is inferred from the file path by from __future__ import annotations -import os from typing import Any, Mapping, Optional 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" BEDROCK_INVOKE_MODELS = [ "claude-haiku-4-5-bedrock-invoke", - "claude-sonnet-4-6-bedrock-invoke", + "claude-sonnet-4-5-bedrock-invoke", "claude-opus-4-7-bedrock-invoke", ] @@ -47,24 +45,11 @@ def _cache_tokens(usage: Optional[Mapping[str, Any]]) -> int: return 0 +@pytest.mark.covers("llm.messages.bedrock_invoke.prompt_cache_5m.nonstream.works") def test_prompt_caching_5m_bedrock_invoke(compat_result): """Drive the `claude` CLI against the LiteLLM proxy and assert the upstream usage block surfaces a non-zero cache token count.""" - 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) outcomes = run_claude_models_parallel( models=BEDROCK_INVOKE_MODELS, diff --git a/tests/e2e/claude_code/prompt_caching_5m/test_vertex_ai.py b/tests/e2e/claude_code/prompt_caching_5m/test_vertex_ai.py index cc5d337dfbe..f79377b7372 100644 --- a/tests/e2e/claude_code/prompt_caching_5m/test_vertex_ai.py +++ b/tests/e2e/claude_code/prompt_caching_5m/test_vertex_ai.py @@ -15,23 +15,21 @@ The (feature, provider) for this cell is inferred from the file path by from __future__ import annotations -import os from typing import Any, Mapping, Optional 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" VERTEX_AI_MODELS = [ "claude-haiku-4-5-vertex", - "claude-sonnet-4-6-vertex", + "claude-sonnet-4-5-vertex", "claude-opus-4-7-vertex", ] @@ -47,24 +45,11 @@ def _cache_tokens(usage: Optional[Mapping[str, Any]]) -> int: return 0 +@pytest.mark.covers("llm.messages.vertex.prompt_cache_5m.nonstream.works") def test_prompt_caching_5m_vertex_ai(compat_result): """Drive the `claude` CLI against the LiteLLM proxy and assert the upstream usage block surfaces a non-zero cache token count.""" - 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) outcomes = run_claude_models_parallel( models=VERTEX_AI_MODELS, diff --git a/tests/e2e/claude_code/run_compat.sh b/tests/e2e/claude_code/run_compat.sh index 4d8d0b6d7b2..e383792f45e 100755 --- a/tests/e2e/claude_code/run_compat.sh +++ b/tests/e2e/claude_code/run_compat.sh @@ -11,9 +11,9 @@ # 4. If a provider has `rate_limited > 0`, halve its rate; else, double it. # 5. Repeat until the highest no-429 rate is found. # -# Required env (proxy connection): -# LITELLM_PROXY_BASE_URL e.g. http://localhost:4000 -# LITELLM_PROXY_API_KEY e.g. sk-1234 +# Required env (proxy connection), same names as the rest of tests/e2e: +# LITELLM_PROXY_URL e.g. http://localhost:4000 +# LITELLM_MASTER_KEY e.g. sk-1234 # # Optional env (rate limits, all default to 5 req/s; 0 disables a column): # LITELLM_COMPAT_RATE_ANTHROPIC @@ -32,8 +32,8 @@ set -euo pipefail -if [[ -z "${LITELLM_PROXY_BASE_URL:-}" || -z "${LITELLM_PROXY_API_KEY:-}" ]]; then - echo "error: LITELLM_PROXY_BASE_URL and LITELLM_PROXY_API_KEY must be set" >&2 +if [[ -z "${LITELLM_PROXY_URL:-}" || -z "${LITELLM_MASTER_KEY:-}" ]]; then + echo "error: LITELLM_PROXY_URL and LITELLM_MASTER_KEY must be set" >&2 exit 64 fi diff --git a/tests/e2e/claude_code/structured_outputs/test_anthropic.py b/tests/e2e/claude_code/structured_outputs/test_anthropic.py index 610d8433b72..3dc4c7ab8f2 100644 --- a/tests/e2e/claude_code/structured_outputs/test_anthropic.py +++ b/tests/e2e/claude_code/structured_outputs/test_anthropic.py @@ -48,24 +48,22 @@ tier so the matrix's "all three must pass" rule applies. from __future__ import annotations import json -import os import re from typing import Any, Mapping, Optional, Sequence, Tuple 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_MODELS = [ "claude-haiku-4-5", - "claude-sonnet-4-6", + "claude-sonnet-4-5", "claude-opus-4-7", ] @@ -152,26 +150,12 @@ def _validate_against_schema( return None +@pytest.mark.covers("llm.messages.anthropic.structured_output.nonstream.works") def test_structured_outputs_anthropic(compat_result): """Drive `claude --json-schema ...` against the LiteLLM proxy and assert the trailing `result` event contains a schema-conforming `structured_output`.""" - 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) outcomes = run_claude_models_parallel( models=ANTHROPIC_MODELS, diff --git a/tests/e2e/claude_code/structured_outputs/test_azure.py b/tests/e2e/claude_code/structured_outputs/test_azure.py index 290f9156910..7a776ed55ad 100644 --- a/tests/e2e/claude_code/structured_outputs/test_azure.py +++ b/tests/e2e/claude_code/structured_outputs/test_azure.py @@ -48,24 +48,22 @@ tier so the matrix's "all three must pass" rule applies. from __future__ import annotations import json -import os import re from typing import Any, Mapping, Optional, Sequence, Tuple 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" AZURE_MODELS = [ "claude-haiku-4-5-azure", - "claude-sonnet-4-6-azure", + "claude-sonnet-4-5-azure", "claude-opus-4-7-azure", ] @@ -152,26 +150,12 @@ def _validate_against_schema( return None +@pytest.mark.covers("llm.messages.azure_foundry.structured_output.nonstream.works") def test_structured_outputs_azure(compat_result): """Drive `claude --json-schema ...` against the LiteLLM proxy and assert the trailing `result` event contains a schema-conforming `structured_output`.""" - 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) outcomes = run_claude_models_parallel( models=AZURE_MODELS, diff --git a/tests/e2e/claude_code/structured_outputs/test_bedrock_converse.py b/tests/e2e/claude_code/structured_outputs/test_bedrock_converse.py index 5179014773c..345d7c327cf 100644 --- a/tests/e2e/claude_code/structured_outputs/test_bedrock_converse.py +++ b/tests/e2e/claude_code/structured_outputs/test_bedrock_converse.py @@ -48,24 +48,22 @@ tier so the matrix's "all three must pass" rule applies. from __future__ import annotations import json -import os import re from typing import Any, Mapping, Optional, Sequence, Tuple 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" BEDROCK_CONVERSE_MODELS = [ "claude-haiku-4-5-bedrock-converse", - "claude-sonnet-4-6-bedrock-converse", + "claude-sonnet-4-5-bedrock-converse", "claude-opus-4-7-bedrock-converse", ] @@ -152,26 +150,12 @@ def _validate_against_schema( return None +@pytest.mark.covers("llm.messages.bedrock_converse.structured_output.nonstream.works") def test_structured_outputs_bedrock_converse(compat_result): """Drive `claude --json-schema ...` against the LiteLLM proxy and assert the trailing `result` event contains a schema-conforming `structured_output`.""" - 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) outcomes = run_claude_models_parallel( models=BEDROCK_CONVERSE_MODELS, diff --git a/tests/e2e/claude_code/structured_outputs/test_bedrock_invoke.py b/tests/e2e/claude_code/structured_outputs/test_bedrock_invoke.py index 313a714be34..0cf48c72d4f 100644 --- a/tests/e2e/claude_code/structured_outputs/test_bedrock_invoke.py +++ b/tests/e2e/claude_code/structured_outputs/test_bedrock_invoke.py @@ -48,24 +48,22 @@ tier so the matrix's "all three must pass" rule applies. from __future__ import annotations import json -import os import re from typing import Any, Mapping, Optional, Sequence, Tuple 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" BEDROCK_INVOKE_MODELS = [ "claude-haiku-4-5-bedrock-invoke", - "claude-sonnet-4-6-bedrock-invoke", + "claude-sonnet-4-5-bedrock-invoke", "claude-opus-4-7-bedrock-invoke", ] @@ -152,26 +150,12 @@ def _validate_against_schema( return None +@pytest.mark.covers("llm.messages.bedrock_invoke.structured_output.nonstream.works") def test_structured_outputs_bedrock_invoke(compat_result): """Drive `claude --json-schema ...` against the LiteLLM proxy and assert the trailing `result` event contains a schema-conforming `structured_output`.""" - 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) outcomes = run_claude_models_parallel( models=BEDROCK_INVOKE_MODELS, diff --git a/tests/e2e/claude_code/structured_outputs/test_vertex_ai.py b/tests/e2e/claude_code/structured_outputs/test_vertex_ai.py index ec04c724193..24f5a0c35d4 100644 --- a/tests/e2e/claude_code/structured_outputs/test_vertex_ai.py +++ b/tests/e2e/claude_code/structured_outputs/test_vertex_ai.py @@ -48,24 +48,22 @@ tier so the matrix's "all three must pass" rule applies. from __future__ import annotations import json -import os import re from typing import Any, Mapping, Optional, Sequence, Tuple 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" VERTEX_AI_MODELS = [ "claude-haiku-4-5-vertex", - "claude-sonnet-4-6-vertex", + "claude-sonnet-4-5-vertex", "claude-opus-4-7-vertex", ] @@ -152,26 +150,12 @@ def _validate_against_schema( return None +@pytest.mark.covers("llm.messages.vertex.structured_output.nonstream.works") def test_structured_outputs_vertex_ai(compat_result): """Drive `claude --json-schema ...` against the LiteLLM proxy and assert the trailing `result` event contains a schema-conforming `structured_output`.""" - 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) outcomes = run_claude_models_parallel( models=VERTEX_AI_MODELS, diff --git a/tests/e2e/claude_code/test_config.yaml b/tests/e2e/claude_code/test_config.yaml index e9253da2b3c..ba57ca4ccb7 100644 --- a/tests/e2e/claude_code/test_config.yaml +++ b/tests/e2e/claude_code/test_config.yaml @@ -21,42 +21,54 @@ model_list: litellm_params: model: anthropic/claude-haiku-4-5 api_key: os.environ/ANTHROPIC_API_KEY - - model_name: claude-sonnet-4-6 + - model_name: claude-sonnet-4-5 litellm_params: - model: anthropic/claude-sonnet-4-6 + model: anthropic/claude-sonnet-4-5 api_key: os.environ/ANTHROPIC_API_KEY + extra_headers: + anthropic-beta: "context-1m-2025-08-07" - model_name: claude-opus-4-7 litellm_params: model: anthropic/claude-opus-4-7 api_key: os.environ/ANTHROPIC_API_KEY + extra_headers: + anthropic-beta: "context-1m-2025-08-07" # ---- Bedrock (InvokeModel) ---- - model_name: claude-haiku-4-5-bedrock-invoke litellm_params: model: bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0 aws_region_name: us-east-1 - - model_name: claude-sonnet-4-6-bedrock-invoke + - model_name: claude-sonnet-4-5-bedrock-invoke litellm_params: - model: bedrock/us.anthropic.claude-sonnet-4-6 + model: bedrock/us.anthropic.claude-sonnet-4-5-20250929-v1:0 aws_region_name: us-east-1 + extra_headers: + anthropic-beta: "context-1m-2025-08-07" - model_name: claude-opus-4-7-bedrock-invoke litellm_params: model: bedrock/us.anthropic.claude-opus-4-7 aws_region_name: us-east-1 + extra_headers: + anthropic-beta: "context-1m-2025-08-07" # ---- Bedrock (Converse) ---- - model_name: claude-haiku-4-5-bedrock-converse litellm_params: model: bedrock/converse/us.anthropic.claude-haiku-4-5-20251001-v1:0 aws_region_name: us-east-1 - - model_name: claude-sonnet-4-6-bedrock-converse + - model_name: claude-sonnet-4-5-bedrock-converse litellm_params: - model: bedrock/converse/us.anthropic.claude-sonnet-4-6 + model: bedrock/converse/us.anthropic.claude-sonnet-4-5-20250929-v1:0 aws_region_name: us-east-1 + extra_headers: + anthropic-beta: "context-1m-2025-08-07" - model_name: claude-opus-4-7-bedrock-converse litellm_params: model: bedrock/converse/us.anthropic.claude-opus-4-7 aws_region_name: us-east-1 + extra_headers: + anthropic-beta: "context-1m-2025-08-07" # ---- Vertex AI ---- # `use_in_pass_through: true` registers each deployment's @@ -70,37 +82,45 @@ model_list: litellm_params: model: vertex_ai/claude-haiku-4-5 vertex_project: os.environ/VERTEXAI_PROJECT - vertex_location: os.environ/VERTEXAI_LOCATION + vertex_location: global use_in_pass_through: true - - model_name: claude-sonnet-4-6-vertex + - model_name: claude-sonnet-4-5-vertex litellm_params: - model: vertex_ai/claude-sonnet-4-6 + model: vertex_ai/claude-sonnet-4-5 vertex_project: os.environ/VERTEXAI_PROJECT - vertex_location: os.environ/VERTEXAI_LOCATION + vertex_location: global use_in_pass_through: true + extra_headers: + anthropic-beta: "context-1m-2025-08-07" - model_name: claude-opus-4-7-vertex litellm_params: model: vertex_ai/claude-opus-4-7 vertex_project: os.environ/VERTEXAI_PROJECT - vertex_location: os.environ/VERTEXAI_LOCATION + vertex_location: global use_in_pass_through: true + extra_headers: + anthropic-beta: "context-1m-2025-08-07" # ---- Microsoft Foundry (Anthropic deployments on Azure) ---- - model_name: claude-haiku-4-5-azure litellm_params: model: azure_ai/claude-haiku-4-5 - api_base: os.environ/AZURE_FOUNDRY_API_BASE - api_key: os.environ/AZURE_FOUNDRY_API_KEY - - model_name: claude-sonnet-4-6-azure + api_base: os.environ/AZURE_AI_API_BASE + api_key: os.environ/AZURE_AI_API_KEY + - model_name: claude-sonnet-4-5-azure litellm_params: - model: azure_ai/claude-sonnet-4-6 - api_base: os.environ/AZURE_FOUNDRY_API_BASE - api_key: os.environ/AZURE_FOUNDRY_API_KEY + model: azure_ai/claude-sonnet-4-5 + api_base: os.environ/AZURE_AI_API_BASE + api_key: os.environ/AZURE_AI_API_KEY + extra_headers: + anthropic-beta: "context-1m-2025-08-07" - model_name: claude-opus-4-7-azure litellm_params: model: azure_ai/claude-opus-4-7 - api_base: os.environ/AZURE_FOUNDRY_API_BASE - api_key: os.environ/AZURE_FOUNDRY_API_KEY + api_base: os.environ/AZURE_AI_API_BASE + api_key: os.environ/AZURE_AI_API_KEY + extra_headers: + anthropic-beta: "context-1m-2025-08-07" general_settings: # Claude Code sends provider-specific headers (e.g. anthropic-beta) we diff --git a/tests/e2e/claude_code/thinking/test_anthropic.py b/tests/e2e/claude_code/thinking/test_anthropic.py index 1090d1b384e..ebb2445fb6d 100644 --- a/tests/e2e/claude_code/thinking/test_anthropic.py +++ b/tests/e2e/claude_code/thinking/test_anthropic.py @@ -20,23 +20,21 @@ still sees three rows for this (feature, provider). from __future__ import annotations -import os from typing import Any, Mapping, 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_MODELS = [ "claude-haiku-4-5", - "claude-sonnet-4-6", + "claude-sonnet-4-5", "claude-opus-4-7", ] @@ -76,24 +74,11 @@ def _has_thinking_block(events: Sequence[Mapping[str, Any]]) -> bool: return False +@pytest.mark.covers("llm.messages.anthropic.thinking.nonstream.works") def test_thinking_anthropic(compat_result): """Drive the `claude` CLI against the LiteLLM proxy with thinking enabled and assert a `thinking` content block was emitted.""" - 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) outcomes = run_claude_models_parallel( models=ANTHROPIC_MODELS, diff --git a/tests/e2e/claude_code/thinking/test_azure.py b/tests/e2e/claude_code/thinking/test_azure.py index 1fd5138d574..ffd5ca92df0 100644 --- a/tests/e2e/claude_code/thinking/test_azure.py +++ b/tests/e2e/claude_code/thinking/test_azure.py @@ -23,23 +23,21 @@ The (feature, provider) for this cell is inferred from the file path by from __future__ import annotations -import os from typing import Any, Mapping, 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" AZURE_MODELS = [ "claude-haiku-4-5-azure", - "claude-sonnet-4-6-azure", + "claude-sonnet-4-5-azure", "claude-opus-4-7-azure", ] @@ -64,24 +62,11 @@ def _has_thinking_block(events: Sequence[Mapping[str, Any]]) -> bool: return False +@pytest.mark.covers("llm.messages.azure_foundry.thinking.nonstream.works") def test_thinking_azure(compat_result): """Drive the `claude` CLI against the LiteLLM proxy with thinking enabled and assert a `thinking` content block was emitted.""" - 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) outcomes = run_claude_models_parallel( models=AZURE_MODELS, diff --git a/tests/e2e/claude_code/thinking/test_bedrock_converse.py b/tests/e2e/claude_code/thinking/test_bedrock_converse.py index 793ce8542da..0b409f18ea7 100644 --- a/tests/e2e/claude_code/thinking/test_bedrock_converse.py +++ b/tests/e2e/claude_code/thinking/test_bedrock_converse.py @@ -15,23 +15,21 @@ The (feature, provider) for this cell is inferred from the file path by from __future__ import annotations -import os from typing import Any, Mapping, 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" BEDROCK_CONVERSE_MODELS = [ "claude-haiku-4-5-bedrock-converse", - "claude-sonnet-4-6-bedrock-converse", + "claude-sonnet-4-5-bedrock-converse", "claude-opus-4-7-bedrock-converse", ] @@ -56,24 +54,11 @@ def _has_thinking_block(events: Sequence[Mapping[str, Any]]) -> bool: return False +@pytest.mark.covers("llm.messages.bedrock_converse.thinking.nonstream.works") def test_thinking_bedrock_converse(compat_result): """Drive the `claude` CLI against the LiteLLM proxy with thinking enabled and assert a `thinking` content block was emitted.""" - 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) outcomes = run_claude_models_parallel( models=BEDROCK_CONVERSE_MODELS, diff --git a/tests/e2e/claude_code/thinking/test_bedrock_invoke.py b/tests/e2e/claude_code/thinking/test_bedrock_invoke.py index e31b60eb004..a2c97eae321 100644 --- a/tests/e2e/claude_code/thinking/test_bedrock_invoke.py +++ b/tests/e2e/claude_code/thinking/test_bedrock_invoke.py @@ -15,23 +15,21 @@ The (feature, provider) for this cell is inferred from the file path by from __future__ import annotations -import os from typing import Any, Mapping, 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" BEDROCK_INVOKE_MODELS = [ "claude-haiku-4-5-bedrock-invoke", - "claude-sonnet-4-6-bedrock-invoke", + "claude-sonnet-4-5-bedrock-invoke", "claude-opus-4-7-bedrock-invoke", ] @@ -56,24 +54,11 @@ def _has_thinking_block(events: Sequence[Mapping[str, Any]]) -> bool: return False +@pytest.mark.covers("llm.messages.bedrock_invoke.thinking.nonstream.works") def test_thinking_bedrock_invoke(compat_result): """Drive the `claude` CLI against the LiteLLM proxy with thinking enabled and assert a `thinking` content block was emitted.""" - 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) outcomes = run_claude_models_parallel( models=BEDROCK_INVOKE_MODELS, diff --git a/tests/e2e/claude_code/thinking/test_vertex_ai.py b/tests/e2e/claude_code/thinking/test_vertex_ai.py index c5c7df1f9b8..f1a1c5b6cee 100644 --- a/tests/e2e/claude_code/thinking/test_vertex_ai.py +++ b/tests/e2e/claude_code/thinking/test_vertex_ai.py @@ -15,23 +15,21 @@ The (feature, provider) for this cell is inferred from the file path by from __future__ import annotations -import os from typing import Any, Mapping, 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" VERTEX_AI_MODELS = [ "claude-haiku-4-5-vertex", - "claude-sonnet-4-6-vertex", + "claude-sonnet-4-5-vertex", "claude-opus-4-7-vertex", ] @@ -56,24 +54,11 @@ def _has_thinking_block(events: Sequence[Mapping[str, Any]]) -> bool: return False +@pytest.mark.covers("llm.messages.vertex.thinking.nonstream.works") def test_thinking_vertex_ai(compat_result): """Drive the `claude` CLI against the LiteLLM proxy with thinking enabled and assert a `thinking` content block was emitted.""" - 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) outcomes = run_claude_models_parallel( models=VERTEX_AI_MODELS, diff --git a/tests/e2e/claude_code/thinking_with_tool_use/test_anthropic.py b/tests/e2e/claude_code/thinking_with_tool_use/test_anthropic.py index 2c573ea039e..7e39ea26d42 100644 --- a/tests/e2e/claude_code/thinking_with_tool_use/test_anthropic.py +++ b/tests/e2e/claude_code/thinking_with_tool_use/test_anthropic.py @@ -24,23 +24,21 @@ The (feature, provider) for this cell is inferred from the file path by from __future__ import annotations -import os from typing import Any, Mapping, 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_MODELS = [ "claude-haiku-4-5", - "claude-sonnet-4-6", + "claude-sonnet-4-5", "claude-opus-4-7", ] @@ -90,25 +88,12 @@ def _has_block_type( return False +@pytest.mark.covers("llm.messages.anthropic.thinking_with_tool_use.nonstream.works") def test_thinking_with_tool_use_anthropic(compat_result): """Drive the `claude` CLI against the LiteLLM proxy with thinking enabled and tool use, and assert both `thinking` and `tool_use` content blocks landed in the same turn.""" - 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) outcomes = run_claude_models_parallel( models=ANTHROPIC_MODELS, diff --git a/tests/e2e/claude_code/thinking_with_tool_use/test_azure.py b/tests/e2e/claude_code/thinking_with_tool_use/test_azure.py index 3d65e82cdec..0371a10f8a6 100644 --- a/tests/e2e/claude_code/thinking_with_tool_use/test_azure.py +++ b/tests/e2e/claude_code/thinking_with_tool_use/test_azure.py @@ -18,23 +18,21 @@ The (feature, provider) for this cell is inferred from the file path by from __future__ import annotations -import os from typing import Any, Mapping, 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" AZURE_MODELS = [ "claude-haiku-4-5-azure", - "claude-sonnet-4-6-azure", + "claude-sonnet-4-5-azure", "claude-opus-4-7-azure", ] @@ -71,22 +69,9 @@ def _has_block_type( return False +@pytest.mark.covers("llm.messages.azure_foundry.thinking_with_tool_use.nonstream.works") def test_thinking_with_tool_use_azure(compat_result): - 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) outcomes = run_claude_models_parallel( models=AZURE_MODELS, diff --git a/tests/e2e/claude_code/thinking_with_tool_use/test_bedrock_converse.py b/tests/e2e/claude_code/thinking_with_tool_use/test_bedrock_converse.py index eb916323546..026d2a3707f 100644 --- a/tests/e2e/claude_code/thinking_with_tool_use/test_bedrock_converse.py +++ b/tests/e2e/claude_code/thinking_with_tool_use/test_bedrock_converse.py @@ -23,23 +23,21 @@ The (feature, provider) for this cell is inferred from the file path by from __future__ import annotations -import os from typing import Any, Mapping, 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" BEDROCK_CONVERSE_MODELS = [ "claude-haiku-4-5-bedrock-converse", - "claude-sonnet-4-6-bedrock-converse", + "claude-sonnet-4-5-bedrock-converse", "claude-opus-4-7-bedrock-converse", ] @@ -76,22 +74,9 @@ def _has_block_type( return False +@pytest.mark.covers("llm.messages.bedrock_converse.thinking_with_tool_use.nonstream.works") def test_thinking_with_tool_use_bedrock_converse(compat_result): - 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) outcomes = run_claude_models_parallel( models=BEDROCK_CONVERSE_MODELS, diff --git a/tests/e2e/claude_code/thinking_with_tool_use/test_bedrock_invoke.py b/tests/e2e/claude_code/thinking_with_tool_use/test_bedrock_invoke.py index d1a61a59772..1dd4cf0a73c 100644 --- a/tests/e2e/claude_code/thinking_with_tool_use/test_bedrock_invoke.py +++ b/tests/e2e/claude_code/thinking_with_tool_use/test_bedrock_invoke.py @@ -25,23 +25,21 @@ The (feature, provider) for this cell is inferred from the file path by from __future__ import annotations -import os from typing import Any, Mapping, 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" BEDROCK_INVOKE_MODELS = [ "claude-haiku-4-5-bedrock-invoke", - "claude-sonnet-4-6-bedrock-invoke", + "claude-sonnet-4-5-bedrock-invoke", "claude-opus-4-7-bedrock-invoke", ] @@ -78,22 +76,9 @@ def _has_block_type( return False +@pytest.mark.covers("llm.messages.bedrock_invoke.thinking_with_tool_use.nonstream.works") def test_thinking_with_tool_use_bedrock_invoke(compat_result): - 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) outcomes = run_claude_models_parallel( models=BEDROCK_INVOKE_MODELS, diff --git a/tests/e2e/claude_code/thinking_with_tool_use/test_vertex_ai.py b/tests/e2e/claude_code/thinking_with_tool_use/test_vertex_ai.py index 285419c67f7..b25228edb55 100644 --- a/tests/e2e/claude_code/thinking_with_tool_use/test_vertex_ai.py +++ b/tests/e2e/claude_code/thinking_with_tool_use/test_vertex_ai.py @@ -23,23 +23,21 @@ The (feature, provider) for this cell is inferred from the file path by from __future__ import annotations -import os from typing import Any, Mapping, 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" VERTEX_AI_MODELS = [ "claude-haiku-4-5-vertex", - "claude-sonnet-4-6-vertex", + "claude-sonnet-4-5-vertex", "claude-opus-4-7-vertex", ] @@ -76,22 +74,9 @@ def _has_block_type( return False +@pytest.mark.covers("llm.messages.vertex.thinking_with_tool_use.nonstream.works") def test_thinking_with_tool_use_vertex_ai(compat_result): - 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) outcomes = run_claude_models_parallel( models=VERTEX_AI_MODELS, diff --git a/tests/e2e/claude_code/tool_search/test_anthropic.py b/tests/e2e/claude_code/tool_search/test_anthropic.py index 3495c882e06..7b8ea07aa07 100644 --- a/tests/e2e/claude_code/tool_search/test_anthropic.py +++ b/tests/e2e/claude_code/tool_search/test_anthropic.py @@ -43,45 +43,28 @@ per-cell aggregator. from __future__ import annotations -import os - import pytest +from claude_code._env import require_proxy from claude_code.http_probe import ( assert_tool_search_shape, probe_tool_search, ) -PROXY_BASE_URL_ENV = "LITELLM_PROXY_BASE_URL" -PROXY_API_KEY_ENV = "LITELLM_PROXY_API_KEY" ANTHROPIC_MODELS = [ "claude-haiku-4-5", - "claude-sonnet-4-6", + "claude-sonnet-4-5", "claude-opus-4-7", ] +@pytest.mark.covers("llm.messages.anthropic.tool_search.nonstream.works") def test_tool_search_anthropic(compat_result): """Probe `/v1/messages` with a `tool_search_tool_regex_20251119` tool and assert the proxy + upstream accept it for every Anthropic tier.""" - 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) failures = [] for model in ANTHROPIC_MODELS: diff --git a/tests/e2e/claude_code/tool_search/test_azure.py b/tests/e2e/claude_code/tool_search/test_azure.py index 1d9cb5673c5..4eee13e4ecc 100644 --- a/tests/e2e/claude_code/tool_search/test_azure.py +++ b/tests/e2e/claude_code/tool_search/test_azure.py @@ -43,45 +43,28 @@ per-cell aggregator. from __future__ import annotations -import os - import pytest +from claude_code._env import require_proxy from claude_code.http_probe import ( assert_tool_search_shape, probe_tool_search, ) -PROXY_BASE_URL_ENV = "LITELLM_PROXY_BASE_URL" -PROXY_API_KEY_ENV = "LITELLM_PROXY_API_KEY" AZURE_MODELS = [ "claude-haiku-4-5-azure", - "claude-sonnet-4-6-azure", + "claude-sonnet-4-5-azure", "claude-opus-4-7-azure", ] +@pytest.mark.covers("llm.messages.azure_foundry.tool_search.nonstream.works") def test_tool_search_azure(compat_result): """Probe `/v1/messages` with a `tool_search_tool_regex_20251119` tool and assert the proxy + upstream accept it for every Azure (Microsoft Foundry) tier.""" - 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) failures = [] for model in AZURE_MODELS: diff --git a/tests/e2e/claude_code/tool_search/test_bedrock_converse.py b/tests/e2e/claude_code/tool_search/test_bedrock_converse.py index 5ca0792529a..7951f8ecdb4 100644 --- a/tests/e2e/claude_code/tool_search/test_bedrock_converse.py +++ b/tests/e2e/claude_code/tool_search/test_bedrock_converse.py @@ -43,45 +43,28 @@ per-cell aggregator. from __future__ import annotations -import os - import pytest +from claude_code._env import require_proxy from claude_code.http_probe import ( assert_tool_search_shape, probe_tool_search, ) -PROXY_BASE_URL_ENV = "LITELLM_PROXY_BASE_URL" -PROXY_API_KEY_ENV = "LITELLM_PROXY_API_KEY" BEDROCK_CONVERSE_MODELS = [ "claude-haiku-4-5-bedrock-converse", - "claude-sonnet-4-6-bedrock-converse", + "claude-sonnet-4-5-bedrock-converse", "claude-opus-4-7-bedrock-converse", ] +@pytest.mark.covers("llm.messages.bedrock_converse.tool_search.nonstream.works") def test_tool_search_bedrock_converse(compat_result): """Probe `/v1/messages` with a `tool_search_tool_regex_20251119` tool and assert the proxy + upstream accept it for every Bedrock (Converse) tier.""" - 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) failures = [] for model in BEDROCK_CONVERSE_MODELS: diff --git a/tests/e2e/claude_code/tool_search/test_bedrock_invoke.py b/tests/e2e/claude_code/tool_search/test_bedrock_invoke.py index 21bb33e34bd..654c2aa18d1 100644 --- a/tests/e2e/claude_code/tool_search/test_bedrock_invoke.py +++ b/tests/e2e/claude_code/tool_search/test_bedrock_invoke.py @@ -43,45 +43,28 @@ per-cell aggregator. from __future__ import annotations -import os - import pytest +from claude_code._env import require_proxy from claude_code.http_probe import ( assert_tool_search_shape, probe_tool_search, ) -PROXY_BASE_URL_ENV = "LITELLM_PROXY_BASE_URL" -PROXY_API_KEY_ENV = "LITELLM_PROXY_API_KEY" BEDROCK_INVOKE_MODELS = [ "claude-haiku-4-5-bedrock-invoke", - "claude-sonnet-4-6-bedrock-invoke", + "claude-sonnet-4-5-bedrock-invoke", "claude-opus-4-7-bedrock-invoke", ] +@pytest.mark.covers("llm.messages.bedrock_invoke.tool_search.nonstream.works") def test_tool_search_bedrock_invoke(compat_result): """Probe `/v1/messages` with a `tool_search_tool_regex_20251119` tool and assert the proxy + upstream accept it for every Bedrock (Invoke) tier.""" - 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) failures = [] for model in BEDROCK_INVOKE_MODELS: diff --git a/tests/e2e/claude_code/tool_search/test_vertex_ai.py b/tests/e2e/claude_code/tool_search/test_vertex_ai.py index f91400b1817..f6ff855fa78 100644 --- a/tests/e2e/claude_code/tool_search/test_vertex_ai.py +++ b/tests/e2e/claude_code/tool_search/test_vertex_ai.py @@ -43,45 +43,28 @@ per-cell aggregator. from __future__ import annotations -import os - import pytest +from claude_code._env import require_proxy from claude_code.http_probe import ( assert_tool_search_shape, probe_tool_search, ) -PROXY_BASE_URL_ENV = "LITELLM_PROXY_BASE_URL" -PROXY_API_KEY_ENV = "LITELLM_PROXY_API_KEY" VERTEX_AI_MODELS = [ "claude-haiku-4-5-vertex", - "claude-sonnet-4-6-vertex", + "claude-sonnet-4-5-vertex", "claude-opus-4-7-vertex", ] +@pytest.mark.covers("llm.messages.vertex.tool_search.nonstream.works") def test_tool_search_vertex_ai(compat_result): """Probe `/v1/messages` with a `tool_search_tool_regex_20251119` tool and assert the proxy + upstream accept it for every Vertex AI tier.""" - 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) failures = [] for model in VERTEX_AI_MODELS: diff --git a/tests/e2e/claude_code/tool_use/test_anthropic.py b/tests/e2e/claude_code/tool_use/test_anthropic.py index 7d2aa4be683..9ff4c58907f 100644 --- a/tests/e2e/claude_code/tool_use/test_anthropic.py +++ b/tests/e2e/claude_code/tool_use/test_anthropic.py @@ -15,23 +15,21 @@ The (feature, provider) for this cell is inferred from the file path by from __future__ import annotations -import os from typing import Any, Mapping, 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_MODELS = [ "claude-haiku-4-5", - "claude-sonnet-4-6", + "claude-sonnet-4-5", "claude-opus-4-7", ] @@ -73,24 +71,11 @@ def _has_tool_use_event(events: Sequence[Mapping[str, Any]]) -> bool: return False +@pytest.mark.covers("llm.messages.anthropic.tool_use.nonstream.works") def test_tool_use_anthropic(compat_result): """Drive the `claude` CLI against the LiteLLM proxy and assert a tool call was emitted on the wire.""" - 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) outcomes = run_claude_models_parallel( models=ANTHROPIC_MODELS, diff --git a/tests/e2e/claude_code/tool_use/test_azure.py b/tests/e2e/claude_code/tool_use/test_azure.py index 484f50a5508..9e7398267c4 100644 --- a/tests/e2e/claude_code/tool_use/test_azure.py +++ b/tests/e2e/claude_code/tool_use/test_azure.py @@ -19,23 +19,21 @@ The (feature, provider) for this cell is inferred from the file path by from __future__ import annotations -import os from typing import Any, Mapping, 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" AZURE_MODELS = [ "claude-haiku-4-5-azure", - "claude-sonnet-4-6-azure", + "claude-sonnet-4-5-azure", "claude-opus-4-7-azure", ] @@ -67,24 +65,11 @@ def _has_tool_use_event(events: Sequence[Mapping[str, Any]]) -> bool: return False +@pytest.mark.covers("llm.messages.azure_foundry.tool_use.nonstream.works") def test_tool_use_azure(compat_result): """Drive the `claude` CLI against the LiteLLM proxy and assert a tool call was emitted on the wire.""" - 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) outcomes = run_claude_models_parallel( models=AZURE_MODELS, diff --git a/tests/e2e/claude_code/tool_use/test_bedrock_converse.py b/tests/e2e/claude_code/tool_use/test_bedrock_converse.py index 7d1b58fce90..33d4d3820d2 100644 --- a/tests/e2e/claude_code/tool_use/test_bedrock_converse.py +++ b/tests/e2e/claude_code/tool_use/test_bedrock_converse.py @@ -15,23 +15,21 @@ The (feature, provider) for this cell is inferred from the file path by from __future__ import annotations -import os from typing import Any, Mapping, 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" BEDROCK_CONVERSE_MODELS = [ "claude-haiku-4-5-bedrock-converse", - "claude-sonnet-4-6-bedrock-converse", + "claude-sonnet-4-5-bedrock-converse", "claude-opus-4-7-bedrock-converse", ] @@ -63,24 +61,11 @@ def _has_tool_use_event(events: Sequence[Mapping[str, Any]]) -> bool: return False +@pytest.mark.covers("llm.messages.bedrock_converse.tool_use.nonstream.works") def test_tool_use_bedrock_converse(compat_result): """Drive the `claude` CLI against the LiteLLM proxy and assert a tool call was emitted on the wire.""" - 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) outcomes = run_claude_models_parallel( models=BEDROCK_CONVERSE_MODELS, diff --git a/tests/e2e/claude_code/tool_use/test_bedrock_invoke.py b/tests/e2e/claude_code/tool_use/test_bedrock_invoke.py index 7d2b72b951d..47ae3aef1da 100644 --- a/tests/e2e/claude_code/tool_use/test_bedrock_invoke.py +++ b/tests/e2e/claude_code/tool_use/test_bedrock_invoke.py @@ -15,23 +15,21 @@ The (feature, provider) for this cell is inferred from the file path by from __future__ import annotations -import os from typing import Any, Mapping, 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" BEDROCK_INVOKE_MODELS = [ "claude-haiku-4-5-bedrock-invoke", - "claude-sonnet-4-6-bedrock-invoke", + "claude-sonnet-4-5-bedrock-invoke", "claude-opus-4-7-bedrock-invoke", ] @@ -63,24 +61,11 @@ def _has_tool_use_event(events: Sequence[Mapping[str, Any]]) -> bool: return False +@pytest.mark.covers("llm.messages.bedrock_invoke.tool_use.nonstream.works") def test_tool_use_bedrock_invoke(compat_result): """Drive the `claude` CLI against the LiteLLM proxy and assert a tool call was emitted on the wire.""" - 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) outcomes = run_claude_models_parallel( models=BEDROCK_INVOKE_MODELS, diff --git a/tests/e2e/claude_code/tool_use/test_vertex_ai.py b/tests/e2e/claude_code/tool_use/test_vertex_ai.py index 0a8ecc9f7a7..79a3016345c 100644 --- a/tests/e2e/claude_code/tool_use/test_vertex_ai.py +++ b/tests/e2e/claude_code/tool_use/test_vertex_ai.py @@ -15,23 +15,21 @@ The (feature, provider) for this cell is inferred from the file path by from __future__ import annotations -import os from typing import Any, Mapping, 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" VERTEX_AI_MODELS = [ "claude-haiku-4-5-vertex", - "claude-sonnet-4-6-vertex", + "claude-sonnet-4-5-vertex", "claude-opus-4-7-vertex", ] @@ -63,24 +61,11 @@ def _has_tool_use_event(events: Sequence[Mapping[str, Any]]) -> bool: return False +@pytest.mark.covers("llm.messages.vertex.tool_use.nonstream.works") def test_tool_use_vertex_ai(compat_result): """Drive the `claude` CLI against the LiteLLM proxy and assert a tool call was emitted on the wire.""" - 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) outcomes = run_claude_models_parallel( models=VERTEX_AI_MODELS, diff --git a/tests/e2e/claude_code/tool_use_streaming/test_anthropic.py b/tests/e2e/claude_code/tool_use_streaming/test_anthropic.py index 9aa94c89241..152652dcf3c 100644 --- a/tests/e2e/claude_code/tool_use_streaming/test_anthropic.py +++ b/tests/e2e/claude_code/tool_use_streaming/test_anthropic.py @@ -25,23 +25,21 @@ The (feature, provider) for this cell is inferred from the file path by from __future__ import annotations -import os from typing import Any, Mapping, 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_MODELS = [ "claude-haiku-4-5", - "claude-sonnet-4-6", + "claude-sonnet-4-5", "claude-opus-4-7", ] @@ -99,24 +97,11 @@ def _count_input_json_deltas(events: Sequence[Mapping[str, Any]]) -> int: ) +@pytest.mark.covers("llm.messages.anthropic.tool_use.stream.works") def test_tool_use_streaming_anthropic(compat_result): """Drive the `claude` CLI against the LiteLLM proxy and assert the proxy preserves fine-grained tool streaming end-to-end.""" - 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) outcomes = run_claude_models_parallel( models=ANTHROPIC_MODELS, diff --git a/tests/e2e/claude_code/tool_use_streaming/test_azure.py b/tests/e2e/claude_code/tool_use_streaming/test_azure.py index c73062b72cd..8a1cc1852dd 100644 --- a/tests/e2e/claude_code/tool_use_streaming/test_azure.py +++ b/tests/e2e/claude_code/tool_use_streaming/test_azure.py @@ -17,23 +17,21 @@ The (feature, provider) for this cell is inferred from the file path by from __future__ import annotations -import os from typing import Any, Mapping, 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" AZURE_MODELS = [ "claude-haiku-4-5-azure", - "claude-sonnet-4-6-azure", + "claude-sonnet-4-5-azure", "claude-opus-4-7-azure", ] @@ -84,22 +82,9 @@ def _count_input_json_deltas(events: Sequence[Mapping[str, Any]]) -> int: ) +@pytest.mark.covers("llm.messages.azure_foundry.tool_use.stream.works") def test_tool_use_streaming_azure(compat_result): - 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) outcomes = run_claude_models_parallel( models=AZURE_MODELS, diff --git a/tests/e2e/claude_code/tool_use_streaming/test_bedrock_converse.py b/tests/e2e/claude_code/tool_use_streaming/test_bedrock_converse.py index 3642551c7c3..3b04ed5962f 100644 --- a/tests/e2e/claude_code/tool_use_streaming/test_bedrock_converse.py +++ b/tests/e2e/claude_code/tool_use_streaming/test_bedrock_converse.py @@ -23,23 +23,21 @@ The (feature, provider) for this cell is inferred from the file path by from __future__ import annotations -import os from typing import Any, Mapping, 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" BEDROCK_CONVERSE_MODELS = [ "claude-haiku-4-5-bedrock-converse", - "claude-sonnet-4-6-bedrock-converse", + "claude-sonnet-4-5-bedrock-converse", "claude-opus-4-7-bedrock-converse", ] @@ -90,22 +88,9 @@ def _count_input_json_deltas(events: Sequence[Mapping[str, Any]]) -> int: ) +@pytest.mark.covers("llm.messages.bedrock_converse.tool_use.stream.works") def test_tool_use_streaming_bedrock_converse(compat_result): - 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) outcomes = run_claude_models_parallel( models=BEDROCK_CONVERSE_MODELS, diff --git a/tests/e2e/claude_code/tool_use_streaming/test_bedrock_invoke.py b/tests/e2e/claude_code/tool_use_streaming/test_bedrock_invoke.py index af4689b2847..c7b61129782 100644 --- a/tests/e2e/claude_code/tool_use_streaming/test_bedrock_invoke.py +++ b/tests/e2e/claude_code/tool_use_streaming/test_bedrock_invoke.py @@ -21,23 +21,21 @@ The (feature, provider) for this cell is inferred from the file path by from __future__ import annotations -import os from typing import Any, Mapping, 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" BEDROCK_INVOKE_MODELS = [ "claude-haiku-4-5-bedrock-invoke", - "claude-sonnet-4-6-bedrock-invoke", + "claude-sonnet-4-5-bedrock-invoke", "claude-opus-4-7-bedrock-invoke", ] @@ -88,22 +86,9 @@ def _count_input_json_deltas(events: Sequence[Mapping[str, Any]]) -> int: ) +@pytest.mark.covers("llm.messages.bedrock_invoke.tool_use.stream.works") def test_tool_use_streaming_bedrock_invoke(compat_result): - 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) outcomes = run_claude_models_parallel( models=BEDROCK_INVOKE_MODELS, diff --git a/tests/e2e/claude_code/tool_use_streaming/test_vertex_ai.py b/tests/e2e/claude_code/tool_use_streaming/test_vertex_ai.py index 19ef9a4e90e..2912e3aae3d 100644 --- a/tests/e2e/claude_code/tool_use_streaming/test_vertex_ai.py +++ b/tests/e2e/claude_code/tool_use_streaming/test_vertex_ai.py @@ -20,23 +20,21 @@ The (feature, provider) for this cell is inferred from the file path by from __future__ import annotations -import os from typing import Any, Mapping, 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" VERTEX_AI_MODELS = [ "claude-haiku-4-5-vertex", - "claude-sonnet-4-6-vertex", + "claude-sonnet-4-5-vertex", "claude-opus-4-7-vertex", ] @@ -87,22 +85,9 @@ def _count_input_json_deltas(events: Sequence[Mapping[str, Any]]) -> int: ) +@pytest.mark.covers("llm.messages.vertex.tool_use.stream.works") def test_tool_use_streaming_vertex_ai(compat_result): - 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) outcomes = run_claude_models_parallel( models=VERTEX_AI_MODELS, diff --git a/tests/e2e/claude_code/vision/test_anthropic.py b/tests/e2e/claude_code/vision/test_anthropic.py index 650940248ea..f681b2be5ae 100644 --- a/tests/e2e/claude_code/vision/test_anthropic.py +++ b/tests/e2e/claude_code/vision/test_anthropic.py @@ -25,22 +25,19 @@ the proxy must preserve. from __future__ import annotations import json -import os - 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_MODELS = [ "claude-haiku-4-5", - "claude-sonnet-4-6", + "claude-sonnet-4-5", "claude-opus-4-7", ] @@ -85,24 +82,11 @@ def _build_stdin_input() -> str: return json.dumps(user_event) + "\n" +@pytest.mark.covers("llm.messages.anthropic.vision.nonstream.works") def test_vision_anthropic(compat_result): """Drive the `claude` CLI against the LiteLLM proxy with an image attached via stream-json input and assert a non-empty reply.""" - 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) outcomes = run_claude_models_parallel( models=ANTHROPIC_MODELS, diff --git a/tests/e2e/claude_code/vision/test_azure.py b/tests/e2e/claude_code/vision/test_azure.py index 3b03c0f2b35..f0eaaad84a2 100644 --- a/tests/e2e/claude_code/vision/test_azure.py +++ b/tests/e2e/claude_code/vision/test_azure.py @@ -25,22 +25,19 @@ the proxy must preserve. from __future__ import annotations import json -import os - 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" AZURE_MODELS = [ "claude-haiku-4-5-azure", - "claude-sonnet-4-6-azure", + "claude-sonnet-4-5-azure", "claude-opus-4-7-azure", ] @@ -85,24 +82,11 @@ def _build_stdin_input() -> str: return json.dumps(user_event) + "\n" +@pytest.mark.covers("llm.messages.azure_foundry.vision.nonstream.works") def test_vision_azure(compat_result): """Drive the `claude` CLI against the LiteLLM proxy with an image attached via stream-json input and assert a non-empty reply.""" - 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) outcomes = run_claude_models_parallel( models=AZURE_MODELS, diff --git a/tests/e2e/claude_code/vision/test_bedrock_converse.py b/tests/e2e/claude_code/vision/test_bedrock_converse.py index 4201f9e64fc..2a5aba5a393 100644 --- a/tests/e2e/claude_code/vision/test_bedrock_converse.py +++ b/tests/e2e/claude_code/vision/test_bedrock_converse.py @@ -25,22 +25,19 @@ the proxy must preserve. from __future__ import annotations import json -import os - 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" BEDROCK_CONVERSE_MODELS = [ "claude-haiku-4-5-bedrock-converse", - "claude-sonnet-4-6-bedrock-converse", + "claude-sonnet-4-5-bedrock-converse", "claude-opus-4-7-bedrock-converse", ] @@ -85,24 +82,11 @@ def _build_stdin_input() -> str: return json.dumps(user_event) + "\n" +@pytest.mark.covers("llm.messages.bedrock_converse.vision.nonstream.works") def test_vision_bedrock_converse(compat_result): """Drive the `claude` CLI against the LiteLLM proxy with an image attached via stream-json input and assert a non-empty reply.""" - 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) outcomes = run_claude_models_parallel( models=BEDROCK_CONVERSE_MODELS, diff --git a/tests/e2e/claude_code/vision/test_bedrock_invoke.py b/tests/e2e/claude_code/vision/test_bedrock_invoke.py index d2e641f1462..5c995cd479e 100644 --- a/tests/e2e/claude_code/vision/test_bedrock_invoke.py +++ b/tests/e2e/claude_code/vision/test_bedrock_invoke.py @@ -25,22 +25,19 @@ the proxy must preserve. from __future__ import annotations import json -import os - 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" BEDROCK_INVOKE_MODELS = [ "claude-haiku-4-5-bedrock-invoke", - "claude-sonnet-4-6-bedrock-invoke", + "claude-sonnet-4-5-bedrock-invoke", "claude-opus-4-7-bedrock-invoke", ] @@ -85,24 +82,11 @@ def _build_stdin_input() -> str: return json.dumps(user_event) + "\n" +@pytest.mark.covers("llm.messages.bedrock_invoke.vision.nonstream.works") def test_vision_bedrock_invoke(compat_result): """Drive the `claude` CLI against the LiteLLM proxy with an image attached via stream-json input and assert a non-empty reply.""" - 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) outcomes = run_claude_models_parallel( models=BEDROCK_INVOKE_MODELS, diff --git a/tests/e2e/claude_code/vision/test_vertex_ai.py b/tests/e2e/claude_code/vision/test_vertex_ai.py index a39ef1a34b7..8d385e295d0 100644 --- a/tests/e2e/claude_code/vision/test_vertex_ai.py +++ b/tests/e2e/claude_code/vision/test_vertex_ai.py @@ -25,22 +25,19 @@ the proxy must preserve. from __future__ import annotations import json -import os - 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" VERTEX_AI_MODELS = [ "claude-haiku-4-5-vertex", - "claude-sonnet-4-6-vertex", + "claude-sonnet-4-5-vertex", "claude-opus-4-7-vertex", ] @@ -85,24 +82,11 @@ def _build_stdin_input() -> str: return json.dumps(user_event) + "\n" +@pytest.mark.covers("llm.messages.vertex.vision.nonstream.works") def test_vision_vertex_ai(compat_result): """Drive the `claude` CLI against the LiteLLM proxy with an image attached via stream-json input and assert a non-empty reply.""" - 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) outcomes = run_claude_models_parallel( models=VERTEX_AI_MODELS, diff --git a/tests/e2e/claude_code/web_search/test_anthropic.py b/tests/e2e/claude_code/web_search/test_anthropic.py index b8fa806f923..a20a2133dc9 100644 --- a/tests/e2e/claude_code/web_search/test_anthropic.py +++ b/tests/e2e/claude_code/web_search/test_anthropic.py @@ -27,23 +27,21 @@ The (feature, provider) for this cell is inferred from the file path by from __future__ import annotations -import os from typing import Any, Mapping, 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_MODELS = [ "claude-haiku-4-5", - "claude-sonnet-4-6", + "claude-sonnet-4-5", "claude-opus-4-7", ] @@ -86,26 +84,13 @@ def _has_web_search_tool_use(events: Sequence[Mapping[str, Any]]) -> bool: return False +@pytest.mark.covers("llm.messages.anthropic.web_search.nonstream.works") def test_web_search_anthropic(compat_result): """Drive the `claude` CLI against the LiteLLM proxy and assert the upstream emitted a `tool_use` block calling `WebSearch`, proving the proxy preserved both the request-side tool definition and the response-side tool_use block.""" - 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) outcomes = run_claude_models_parallel( models=ANTHROPIC_MODELS, diff --git a/tests/e2e/claude_code/web_search/test_azure.py b/tests/e2e/claude_code/web_search/test_azure.py index e70dc848dcf..8f9f638fbee 100644 --- a/tests/e2e/claude_code/web_search/test_azure.py +++ b/tests/e2e/claude_code/web_search/test_azure.py @@ -27,23 +27,21 @@ The (feature, provider) for this cell is inferred from the file path by from __future__ import annotations -import os from typing import Any, Mapping, 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" AZURE_MODELS = [ "claude-haiku-4-5-azure", - "claude-sonnet-4-6-azure", + "claude-sonnet-4-5-azure", "claude-opus-4-7-azure", ] @@ -86,26 +84,13 @@ def _has_web_search_tool_use(events: Sequence[Mapping[str, Any]]) -> bool: return False +@pytest.mark.covers("llm.messages.azure_foundry.web_search.nonstream.works") def test_web_search_azure(compat_result): """Drive the `claude` CLI against the LiteLLM proxy and assert the upstream emitted a `tool_use` block calling `WebSearch`, proving the proxy preserved both the request-side tool definition and the response-side tool_use block.""" - 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) outcomes = run_claude_models_parallel( models=AZURE_MODELS, diff --git a/tests/e2e/claude_code/web_search/test_bedrock_converse.py b/tests/e2e/claude_code/web_search/test_bedrock_converse.py index cbeea03df40..32f37b2be79 100644 --- a/tests/e2e/claude_code/web_search/test_bedrock_converse.py +++ b/tests/e2e/claude_code/web_search/test_bedrock_converse.py @@ -27,23 +27,21 @@ The (feature, provider) for this cell is inferred from the file path by from __future__ import annotations -import os from typing import Any, Mapping, 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" BEDROCK_CONVERSE_MODELS = [ "claude-haiku-4-5-bedrock-converse", - "claude-sonnet-4-6-bedrock-converse", + "claude-sonnet-4-5-bedrock-converse", "claude-opus-4-7-bedrock-converse", ] @@ -86,26 +84,13 @@ def _has_web_search_tool_use(events: Sequence[Mapping[str, Any]]) -> bool: return False +@pytest.mark.covers("llm.messages.bedrock_converse.web_search.nonstream.works") def test_web_search_bedrock_converse(compat_result): """Drive the `claude` CLI against the LiteLLM proxy and assert the upstream emitted a `tool_use` block calling `WebSearch`, proving the proxy preserved both the request-side tool definition and the response-side tool_use block.""" - 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) outcomes = run_claude_models_parallel( models=BEDROCK_CONVERSE_MODELS, diff --git a/tests/e2e/claude_code/web_search/test_bedrock_invoke.py b/tests/e2e/claude_code/web_search/test_bedrock_invoke.py index 86068e1e22b..68d1b30e83f 100644 --- a/tests/e2e/claude_code/web_search/test_bedrock_invoke.py +++ b/tests/e2e/claude_code/web_search/test_bedrock_invoke.py @@ -27,23 +27,21 @@ The (feature, provider) for this cell is inferred from the file path by from __future__ import annotations -import os from typing import Any, Mapping, 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" BEDROCK_INVOKE_MODELS = [ "claude-haiku-4-5-bedrock-invoke", - "claude-sonnet-4-6-bedrock-invoke", + "claude-sonnet-4-5-bedrock-invoke", "claude-opus-4-7-bedrock-invoke", ] @@ -86,26 +84,13 @@ def _has_web_search_tool_use(events: Sequence[Mapping[str, Any]]) -> bool: return False +@pytest.mark.covers("llm.messages.bedrock_invoke.web_search.nonstream.works") def test_web_search_bedrock_invoke(compat_result): """Drive the `claude` CLI against the LiteLLM proxy and assert the upstream emitted a `tool_use` block calling `WebSearch`, proving the proxy preserved both the request-side tool definition and the response-side tool_use block.""" - 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) outcomes = run_claude_models_parallel( models=BEDROCK_INVOKE_MODELS, diff --git a/tests/e2e/claude_code/web_search/test_vertex_ai.py b/tests/e2e/claude_code/web_search/test_vertex_ai.py index a33515771f3..540a8396c98 100644 --- a/tests/e2e/claude_code/web_search/test_vertex_ai.py +++ b/tests/e2e/claude_code/web_search/test_vertex_ai.py @@ -27,23 +27,21 @@ The (feature, provider) for this cell is inferred from the file path by from __future__ import annotations -import os from typing import Any, Mapping, 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" VERTEX_AI_MODELS = [ "claude-haiku-4-5-vertex", - "claude-sonnet-4-6-vertex", + "claude-sonnet-4-5-vertex", "claude-opus-4-7-vertex", ] @@ -86,26 +84,13 @@ def _has_web_search_tool_use(events: Sequence[Mapping[str, Any]]) -> bool: return False +@pytest.mark.covers("llm.messages.vertex.web_search.nonstream.works") def test_web_search_vertex_ai(compat_result): """Drive the `claude` CLI against the LiteLLM proxy and assert the upstream emitted a `tool_use` block calling `WebSearch`, proving the proxy preserved both the request-side tool definition and the response-side tool_use block.""" - 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) outcomes = run_claude_models_parallel( models=VERTEX_AI_MODELS, diff --git a/tests/e2e/conftest.py b/tests/e2e/conftest.py index 08df334d4b8..3aec104c861 100644 --- a/tests/e2e/conftest.py +++ b/tests/e2e/conftest.py @@ -15,13 +15,14 @@ shared fixtures build on it. import functools import sys +from collections.abc import Generator, Iterator from pathlib import Path -from typing import Iterator import pytest import requests from e2e_config import CONTROL_PLANE_BASE_URL, PROXY_BASE_URL +from e2e_result_reporter import covers_from_item, format_e2e_result_line, result_from_pytest from lifecycle import GatewayProvider, ResourceManager @@ -85,6 +86,30 @@ def pytest_runtest_call(item: pytest.Item) -> None: item.session.stash[_E2E_TEST_RAN] = True +@pytest.hookimpl(wrapper=True, tryfirst=True) +def pytest_runtest_makereport( + item: pytest.Item, call: pytest.CallInfo[object] +) -> Generator[None, pytest.TestReport, pytest.TestReport]: + """Emit one structured E2E_RESULT line per finished test for Loki/Grafana. + + Status-history panels should aggregate by package (and optional covers), not + scrape pytest progress basenames. See e2e_result_reporter.py. + """ + report = yield + result = result_from_pytest( + nodeid=str(report.nodeid), + when=str(report.when), + failed=bool(report.failed), + skipped=bool(report.skipped), + passed=bool(report.passed), + duration_seconds=float(report.duration), + covers=covers_from_item(item), + ) + if result is not None: + print(format_e2e_result_line(result), flush=True) + return report + + def pytest_sessionfinish(session: pytest.Session, exitstatus: int) -> None: """Once the whole e2e session is done (all suites), truncate the spend logs so the DB doesn't accumulate test rows. Sessions where no e2e test body ran leave diff --git a/tests/e2e/coverage_registry/README.md b/tests/e2e/coverage_registry/README.md index aef4c16c89a..ae08d61cacc 100644 --- a/tests/e2e/coverage_registry/README.md +++ b/tests/e2e/coverage_registry/README.md @@ -53,6 +53,11 @@ in `MODULE_ORDER`, in that order. Loki uses log-safe `module=` labels from `LOKI_MODULE_LABELS` (`core_llms`, `management_ui`, etc.) so existing JSON and Prometheus consumers keep their human-readable module names unchanged. +Live pass/fail is separate: each finished pytest node prints an `E2E_RESULT` +logfmt line (see `tests/e2e/e2e_result_reporter.py` and +`tests/e2e/grafana/status_history_panels.md`). Coverage answers "is there a +test for this cell?"; `E2E_RESULT` answers "did that run pass?" + The headline is overall coverage. The collector also lists markers that point at ids not in the registry, so a typo or an unenumerated behavior surfaces instead of being silently dropped. diff --git a/tests/e2e/coverage_registry/llm_claude_code_compat.yaml b/tests/e2e/coverage_registry/llm_claude_code_compat.yaml new file mode 100644 index 00000000000..6edf890f7ec --- /dev/null +++ b/tests/e2e/coverage_registry/llm_claude_code_compat.yaml @@ -0,0 +1,110 @@ +# Claude Code compatibility matrix: /v1/messages coverage across the five provider surfaces +# claude-code drives (anthropic direct, azure ai foundry, bedrock invoke, bedrock converse, +# vertex ai). Each row is one (feature x provider) cell in the matrix. The seven anthropic-direct +# rows already declared in llm_conversational.yaml are NOT duplicated here; the four other +# provider surfaces plus every feature not already listed for anthropic direct are declared below. +# +# Grammar: llm.messages....works +# route : anthropic | azure_foundry | bedrock_converse | bedrock_invoke | vertex +# capability : basic | tool_use | vision | thinking | prompt_cache_5m | prompt_cache_1h +# | structured_output | pdf_input | long_context_1m +# | thinking_with_tool_use | tool_search | count_tokens | web_search +# streaming : stream | nonstream + +# ---- basic / non-streaming ---- +- {id: llm.messages.azure_foundry.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: azure_foundry, capability: basic, streaming: nonstream, assertions: [works], source: "claude_code compat matrix", rationale: "Basic messaging over Azure AI Foundry Anthropic deployments"} +- {id: llm.messages.bedrock_converse.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: bedrock_converse, capability: basic, streaming: nonstream, assertions: [works], source: "claude_code compat matrix", rationale: "Basic messaging over Bedrock Converse Anthropic"} +- {id: llm.messages.bedrock_invoke.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: bedrock_invoke, capability: basic, streaming: nonstream, assertions: [works], source: "claude_code compat matrix", rationale: "Basic messaging over Bedrock Invoke Anthropic"} +- {id: llm.messages.vertex.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: vertex, capability: basic, streaming: nonstream, assertions: [works], source: "claude_code compat matrix", rationale: "Basic messaging over Vertex AI Anthropic"} + +# ---- basic / streaming ---- +- {id: llm.messages.azure_foundry.basic.stream.works, module: llm, tier: P1, subject_endpoint: messages, route: azure_foundry, capability: basic, streaming: stream, assertions: [works], source: "claude_code compat matrix", rationale: "Basic streaming over Azure AI Foundry"} +- {id: llm.messages.bedrock_converse.basic.stream.works, module: llm, tier: P1, subject_endpoint: messages, route: bedrock_converse, capability: basic, streaming: stream, assertions: [works], source: "claude_code compat matrix", rationale: "Basic streaming over Bedrock Converse"} +- {id: llm.messages.bedrock_invoke.basic.stream.works, module: llm, tier: P1, subject_endpoint: messages, route: bedrock_invoke, capability: basic, streaming: stream, assertions: [works], source: "claude_code compat matrix", rationale: "Basic streaming over Bedrock Invoke"} +- {id: llm.messages.vertex.basic.stream.works, module: llm, tier: P1, subject_endpoint: messages, route: vertex, capability: basic, streaming: stream, assertions: [works], source: "claude_code compat matrix", rationale: "Basic streaming over Vertex AI"} + +# ---- tool_use / non-streaming ---- +- {id: llm.messages.azure_foundry.tool_use.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: azure_foundry, capability: tool_use, streaming: nonstream, assertions: [works], source: "claude_code compat matrix", rationale: "Tool use over Azure AI Foundry"} +- {id: llm.messages.bedrock_converse.tool_use.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: bedrock_converse, capability: tool_use, streaming: nonstream, assertions: [works], source: "claude_code compat matrix", rationale: "Tool use over Bedrock Converse"} +- {id: llm.messages.bedrock_invoke.tool_use.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: bedrock_invoke, capability: tool_use, streaming: nonstream, assertions: [works], source: "claude_code compat matrix", rationale: "Tool use over Bedrock Invoke"} +- {id: llm.messages.vertex.tool_use.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: vertex, capability: tool_use, streaming: nonstream, assertions: [works], source: "claude_code compat matrix", rationale: "Tool use over Vertex AI"} + +# ---- tool_use / streaming ---- +- {id: llm.messages.azure_foundry.tool_use.stream.works, module: llm, tier: P1, subject_endpoint: messages, route: azure_foundry, capability: tool_use, streaming: stream, assertions: [works], source: "claude_code compat matrix", rationale: "Streaming tool use over Azure AI Foundry"} +- {id: llm.messages.bedrock_converse.tool_use.stream.works, module: llm, tier: P1, subject_endpoint: messages, route: bedrock_converse, capability: tool_use, streaming: stream, assertions: [works], source: "claude_code compat matrix", rationale: "Streaming tool use over Bedrock Converse"} +- {id: llm.messages.bedrock_invoke.tool_use.stream.works, module: llm, tier: P1, subject_endpoint: messages, route: bedrock_invoke, capability: tool_use, streaming: stream, assertions: [works], source: "claude_code compat matrix", rationale: "Streaming tool use over Bedrock Invoke"} +- {id: llm.messages.vertex.tool_use.stream.works, module: llm, tier: P1, subject_endpoint: messages, route: vertex, capability: tool_use, streaming: stream, assertions: [works], source: "claude_code compat matrix", rationale: "Streaming tool use over Vertex AI"} + +# ---- vision ---- +- {id: llm.messages.azure_foundry.vision.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: azure_foundry, capability: vision, streaming: nonstream, assertions: [works], source: "claude_code compat matrix", rationale: "Vision over Azure AI Foundry"} +- {id: llm.messages.bedrock_converse.vision.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: bedrock_converse, capability: vision, streaming: nonstream, assertions: [works], source: "claude_code compat matrix", rationale: "Vision over Bedrock Converse"} +- {id: llm.messages.bedrock_invoke.vision.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: bedrock_invoke, capability: vision, streaming: nonstream, assertions: [works], source: "claude_code compat matrix", rationale: "Vision over Bedrock Invoke"} +- {id: llm.messages.vertex.vision.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: vertex, capability: vision, streaming: nonstream, assertions: [works], source: "claude_code compat matrix", rationale: "Vision over Vertex AI"} + +# ---- thinking ---- +- {id: llm.messages.azure_foundry.thinking.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: azure_foundry, capability: thinking, streaming: nonstream, assertions: [works], source: "claude_code compat matrix", rationale: "Extended thinking over Azure AI Foundry"} +- {id: llm.messages.bedrock_converse.thinking.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: bedrock_converse, capability: thinking, streaming: nonstream, assertions: [works], source: "claude_code compat matrix", rationale: "Extended thinking over Bedrock Converse"} +- {id: llm.messages.bedrock_invoke.thinking.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: bedrock_invoke, capability: thinking, streaming: nonstream, assertions: [works], source: "claude_code compat matrix", rationale: "Extended thinking over Bedrock Invoke"} +- {id: llm.messages.vertex.thinking.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: vertex, capability: thinking, streaming: nonstream, assertions: [works], source: "claude_code compat matrix", rationale: "Extended thinking over Vertex AI"} + +# ---- prompt_cache_5m ---- +- {id: llm.messages.azure_foundry.prompt_cache_5m.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: azure_foundry, capability: prompt_cache_5m, streaming: nonstream, assertions: [works], source: "claude_code compat matrix", rationale: "5m prompt cache over Azure AI Foundry"} +- {id: llm.messages.bedrock_converse.prompt_cache_5m.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: bedrock_converse, capability: prompt_cache_5m, streaming: nonstream, assertions: [works], source: "claude_code compat matrix", rationale: "5m prompt cache over Bedrock Converse"} +- {id: llm.messages.bedrock_invoke.prompt_cache_5m.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: bedrock_invoke, capability: prompt_cache_5m, streaming: nonstream, assertions: [works], source: "claude_code compat matrix", rationale: "5m prompt cache over Bedrock Invoke"} +- {id: llm.messages.vertex.prompt_cache_5m.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: vertex, capability: prompt_cache_5m, streaming: nonstream, assertions: [works], source: "claude_code compat matrix", rationale: "5m prompt cache over Vertex AI"} + +# ---- prompt_cache_1h ---- +- {id: llm.messages.anthropic.prompt_cache_1h.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: anthropic, capability: prompt_cache_1h, streaming: nonstream, assertions: [works], source: "claude_code compat matrix", rationale: "1h prompt cache over Anthropic direct"} +- {id: llm.messages.azure_foundry.prompt_cache_1h.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: azure_foundry, capability: prompt_cache_1h, streaming: nonstream, assertions: [works], source: "claude_code compat matrix", rationale: "1h prompt cache over Azure AI Foundry"} +- {id: llm.messages.bedrock_converse.prompt_cache_1h.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: bedrock_converse, capability: prompt_cache_1h, streaming: nonstream, assertions: [works], source: "claude_code compat matrix", rationale: "1h prompt cache over Bedrock Converse"} +- {id: llm.messages.bedrock_invoke.prompt_cache_1h.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: bedrock_invoke, capability: prompt_cache_1h, streaming: nonstream, assertions: [works], source: "claude_code compat matrix", rationale: "1h prompt cache over Bedrock Invoke"} +- {id: llm.messages.vertex.prompt_cache_1h.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: vertex, capability: prompt_cache_1h, streaming: nonstream, assertions: [works], source: "claude_code compat matrix", rationale: "1h prompt cache over Vertex AI"} + +# ---- structured_output ---- +- {id: llm.messages.anthropic.structured_output.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: anthropic, capability: structured_output, streaming: nonstream, assertions: [works], source: "claude_code compat matrix", rationale: "Structured outputs (--json-schema) over Anthropic direct"} +- {id: llm.messages.azure_foundry.structured_output.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: azure_foundry, capability: structured_output, streaming: nonstream, assertions: [works], source: "claude_code compat matrix", rationale: "Structured outputs over Azure AI Foundry"} +- {id: llm.messages.bedrock_converse.structured_output.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: bedrock_converse, capability: structured_output, streaming: nonstream, assertions: [works], source: "claude_code compat matrix", rationale: "Structured outputs over Bedrock Converse"} +- {id: llm.messages.bedrock_invoke.structured_output.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: bedrock_invoke, capability: structured_output, streaming: nonstream, assertions: [works], source: "claude_code compat matrix", rationale: "Structured outputs over Bedrock Invoke"} +- {id: llm.messages.vertex.structured_output.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: vertex, capability: structured_output, streaming: nonstream, assertions: [works], source: "claude_code compat matrix", rationale: "Structured outputs over Vertex AI"} + +# ---- pdf_input ---- +- {id: llm.messages.anthropic.pdf_input.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: anthropic, capability: pdf_input, streaming: nonstream, assertions: [works], source: "claude_code compat matrix", rationale: "PDF document input over Anthropic direct"} +- {id: llm.messages.azure_foundry.pdf_input.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: azure_foundry, capability: pdf_input, streaming: nonstream, assertions: [works], source: "claude_code compat matrix", rationale: "PDF document input over Azure AI Foundry"} +- {id: llm.messages.bedrock_converse.pdf_input.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: bedrock_converse, capability: pdf_input, streaming: nonstream, assertions: [works], source: "claude_code compat matrix", rationale: "PDF document input over Bedrock Converse"} +- {id: llm.messages.bedrock_invoke.pdf_input.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: bedrock_invoke, capability: pdf_input, streaming: nonstream, assertions: [works], source: "claude_code compat matrix", rationale: "PDF document input over Bedrock Invoke"} +- {id: llm.messages.vertex.pdf_input.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: vertex, capability: pdf_input, streaming: nonstream, assertions: [works], source: "claude_code compat matrix", rationale: "PDF document input over Vertex AI"} + +# ---- long_context_1m ---- +- {id: llm.messages.anthropic.long_context_1m.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: anthropic, capability: long_context_1m, streaming: nonstream, assertions: [works], source: "claude_code compat matrix", rationale: "1M context beta over Anthropic direct"} +- {id: llm.messages.azure_foundry.long_context_1m.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: azure_foundry, capability: long_context_1m, streaming: nonstream, assertions: [works], source: "claude_code compat matrix", rationale: "1M context beta over Azure AI Foundry"} +- {id: llm.messages.bedrock_converse.long_context_1m.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: bedrock_converse, capability: long_context_1m, streaming: nonstream, assertions: [works], source: "claude_code compat matrix", rationale: "1M context beta over Bedrock Converse"} +- {id: llm.messages.bedrock_invoke.long_context_1m.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: bedrock_invoke, capability: long_context_1m, streaming: nonstream, assertions: [works], source: "claude_code compat matrix", rationale: "1M context beta over Bedrock Invoke"} +- {id: llm.messages.vertex.long_context_1m.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: vertex, capability: long_context_1m, streaming: nonstream, assertions: [works], source: "claude_code compat matrix", rationale: "1M context beta over Vertex AI"} + +# ---- thinking_with_tool_use ---- +- {id: llm.messages.anthropic.thinking_with_tool_use.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: anthropic, capability: thinking_with_tool_use, streaming: nonstream, assertions: [works], source: "claude_code compat matrix", rationale: "Thinking + tool_use interleaved over Anthropic direct"} +- {id: llm.messages.azure_foundry.thinking_with_tool_use.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: azure_foundry, capability: thinking_with_tool_use, streaming: nonstream, assertions: [works], source: "claude_code compat matrix", rationale: "Thinking + tool_use interleaved over Azure AI Foundry"} +- {id: llm.messages.bedrock_converse.thinking_with_tool_use.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: bedrock_converse, capability: thinking_with_tool_use, streaming: nonstream, assertions: [works], source: "claude_code compat matrix", rationale: "Thinking + tool_use interleaved over Bedrock Converse"} +- {id: llm.messages.bedrock_invoke.thinking_with_tool_use.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: bedrock_invoke, capability: thinking_with_tool_use, streaming: nonstream, assertions: [works], source: "claude_code compat matrix", rationale: "Thinking + tool_use interleaved over Bedrock Invoke"} +- {id: llm.messages.vertex.thinking_with_tool_use.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: vertex, capability: thinking_with_tool_use, streaming: nonstream, assertions: [works], source: "claude_code compat matrix", rationale: "Thinking + tool_use interleaved over Vertex AI"} + +# ---- tool_search ---- +- {id: llm.messages.anthropic.tool_search.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: anthropic, capability: tool_search, streaming: nonstream, assertions: [works], source: "claude_code compat matrix", rationale: "tool_search_tool_regex_20251119 discovery tool over Anthropic direct"} +- {id: llm.messages.azure_foundry.tool_search.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: azure_foundry, capability: tool_search, streaming: nonstream, assertions: [works], source: "claude_code compat matrix", rationale: "tool_search discovery tool over Azure AI Foundry"} +- {id: llm.messages.bedrock_converse.tool_search.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: bedrock_converse, capability: tool_search, streaming: nonstream, assertions: [works], source: "claude_code compat matrix", rationale: "tool_search discovery tool over Bedrock Converse"} +- {id: llm.messages.bedrock_invoke.tool_search.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: bedrock_invoke, capability: tool_search, streaming: nonstream, assertions: [works], source: "claude_code compat matrix", rationale: "tool_search discovery tool over Bedrock Invoke"} +- {id: llm.messages.vertex.tool_search.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: vertex, capability: tool_search, streaming: nonstream, assertions: [works], source: "claude_code compat matrix", rationale: "tool_search discovery tool over Vertex AI"} + +# ---- count_tokens ---- +- {id: llm.messages.anthropic.count_tokens.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: anthropic, capability: count_tokens, streaming: nonstream, assertions: [works], source: "claude_code compat matrix", rationale: "/v1/messages/count_tokens over Anthropic direct"} +- {id: llm.messages.azure_foundry.count_tokens.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: azure_foundry, capability: count_tokens, streaming: nonstream, assertions: [works], source: "claude_code compat matrix", rationale: "/v1/messages/count_tokens over Azure AI Foundry"} +- {id: llm.messages.bedrock_converse.count_tokens.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: bedrock_converse, capability: count_tokens, streaming: nonstream, assertions: [works], source: "claude_code compat matrix", rationale: "/v1/messages/count_tokens over Bedrock Converse"} +- {id: llm.messages.bedrock_invoke.count_tokens.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: bedrock_invoke, capability: count_tokens, streaming: nonstream, assertions: [works], source: "claude_code compat matrix", rationale: "/v1/messages/count_tokens over Bedrock Invoke"} +- {id: llm.messages.vertex.count_tokens.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: vertex, capability: count_tokens, streaming: nonstream, assertions: [works], source: "claude_code compat matrix", rationale: "/v1/messages/count_tokens over Vertex AI"} + +# ---- web_search ---- +- {id: llm.messages.anthropic.web_search.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: anthropic, capability: web_search, streaming: nonstream, assertions: [works], source: "claude_code compat matrix", rationale: "Web search server tool over Anthropic direct"} +- {id: llm.messages.azure_foundry.web_search.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: azure_foundry, capability: web_search, streaming: nonstream, assertions: [works], source: "claude_code compat matrix", rationale: "Web search server tool over Azure AI Foundry"} +- {id: llm.messages.bedrock_converse.web_search.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: bedrock_converse, capability: web_search, streaming: nonstream, assertions: [works], source: "claude_code compat matrix", rationale: "Web search server tool over Bedrock Converse"} +- {id: llm.messages.bedrock_invoke.web_search.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: bedrock_invoke, capability: web_search, streaming: nonstream, assertions: [works], source: "claude_code compat matrix", rationale: "Web search server tool over Bedrock Invoke"} +- {id: llm.messages.vertex.web_search.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: vertex, capability: web_search, streaming: nonstream, assertions: [works], source: "claude_code compat matrix", rationale: "Web search server tool over Vertex AI"} diff --git a/tests/e2e/coverage_registry/logging.yaml b/tests/e2e/coverage_registry/logging.yaml index afb6dbc964e..01a44539c88 100644 --- a/tests/e2e/coverage_registry/logging.yaml +++ b/tests/e2e/coverage_registry/logging.yaml @@ -1,11 +1,8 @@ # Logging integration delivery (behavior features). Grounded in litellm/integrations/. -- {id: logging.langfuse.success.logs_spend, module: logging, tier: P0, event: success, assertions: [logs_spend], exercised_on: [chat_completions, messages, embeddings], source: "integrations/langfuse/langfuse.py", rationale: "Primary tracing backend; cost accuracy"} -- {id: logging.langfuse.failure.logs_spend, module: logging, tier: P0, event: failure, assertions: [logs_spend], exercised_on: [chat_completions, messages], source: "integrations/langfuse/langfuse.py", rationale: "Failure path must still track spend"} -- {id: logging.langfuse.stream.logs_spend, module: logging, tier: P0, event: stream, assertions: [logs_spend], exercised_on: [chat_completions, messages], source: "integrations/langfuse/langfuse.py", rationale: "Streaming token counts aggregate"} - {id: logging.s3.success.writes_object, module: logging, tier: P0, event: success, assertions: [writes_object], exercised_on: [chat_completions, messages, embeddings], source: "integrations/s3_v2.py", rationale: "Primary audit trail; batch flush no-drop"} - {id: logging.s3.failure.writes_object, module: logging, tier: P0, event: failure, assertions: [writes_object], exercised_on: [chat_completions, messages], source: "integrations/s3_v2.py", rationale: "Failed calls persisted for compliance"} - {id: logging.gcs_bucket.success.writes_object, module: logging, tier: P0, event: success, assertions: [writes_object], exercised_on: [chat_completions, messages, embeddings], source: "integrations/gcs_bucket/gcs_bucket.py", rationale: "GCS parallel to S3"} -- {id: logging.datadog.success.exports_metric, module: logging, tier: P0, event: success, assertions: [exports_metric], exercised_on: [chat_completions, messages, embeddings], source: "integrations/datadog/datadog.py", rationale: "Powers dashboards/alerts; cardinality regressions common"} +- {id: logging.datadog.success.exports_metric, module: logging, tier: P0, event: success, assertions: [exports_metric], exercised_on: [chat_completions, messages, responses, embeddings], source: "integrations/datadog/datadog.py", rationale: "Powers dashboards/alerts; cardinality regressions common"} - {id: logging.datadog.failure.exports_metric, module: logging, tier: P0, event: failure, assertions: [exports_metric], exercised_on: [chat_completions], source: "integrations/datadog/datadog.py", rationale: "Failure metrics for alerting/SLO"} - {id: logging.prometheus.success.exports_metric, module: logging, tier: P0, event: success, assertions: [exports_metric], exercised_on: [chat_completions, messages, embeddings], source: "integrations/prometheus.py", rationale: "Standard OSS metrics; per-key cardinality (existing e2e)"} - {id: logging.otel.success.exports_metric, module: logging, tier: P0, event: success, assertions: [exports_metric], exercised_on: [chat_completions, messages, responses, embeddings], source: "integrations/otel/logger.py", rationale: "OTEL spans on every call path"} diff --git a/tests/e2e/coverage_registry/schema.py b/tests/e2e/coverage_registry/schema.py index 7482088d93f..a76774c3bde 100644 --- a/tests/e2e/coverage_registry/schema.py +++ b/tests/e2e/coverage_registry/schema.py @@ -54,13 +54,20 @@ LlmRoute = Literal[ LlmCapability = Literal[ "basic", + "count_tokens", + "long_context_1m", "mid_conversation_system", + "pdf_input", + "prompt_cache_1h", "prompt_cache_5m", "service_tier", "structured_output", "thinking", + "thinking_with_tool_use", + "tool_search", "tool_use", "vision", + "web_search", ] diff --git a/tests/e2e/docker-compose.yml b/tests/e2e/docker-compose.yml index b64f3d8dbfd..5f75409f025 100644 --- a/tests/e2e/docker-compose.yml +++ b/tests/e2e/docker-compose.yml @@ -1,5 +1,41 @@ # local setup to run e2e tests configs: + dd_sink_script: + content: | + # Minimal DataDog logs-intake sink for the logging suite: records every + # POST (gunzipping the compressed batches the integration sends) and + # replays them as JSON on GET /requests so tests can assert delivery. + import gzip, json + from http.server import BaseHTTPRequestHandler, HTTPServer + + REQUESTS = [] + + class Handler(BaseHTTPRequestHandler): + def do_POST(self): + body = self.rfile.read(int(self.headers.get("Content-Length", 0))) + if self.headers.get("Content-Encoding") == "gzip": + body = gzip.decompress(body) + REQUESTS.append({"path": self.path, "body": body.decode("utf-8", "replace")}) + self.send_response(202) + self.end_headers() + self.wfile.write(b"{}") + + def do_GET(self): + self.send_response(200) + if self.path == "/health": + self.send_header("Content-Type", "text/plain") + self.end_headers() + self.wfile.write(b"ok") + return + self.send_header("Content-Type", "application/json") + self.end_headers() + self.wfile.write(json.dumps({"requests": REQUESTS}).encode()) + + def log_message(self, *args): + pass + + HTTPServer(("0.0.0.0", 8080), Handler).serve_forever() + litellm_config: content: | general_settings: @@ -23,7 +59,7 @@ configs: # (PHOENIX_COLLECTOR_HTTP_ENDPOINT below points it at the jaeger service), # so gen-AI spans export through a preset-owned provider - the code path # where trace splits actually happen - with no cloud credentials needed. - callbacks: ["arize_phoenix"] + callbacks: ["arize_phoenix", "datadog"] router_settings: routing_strategy: simple-shuffle @@ -93,10 +129,15 @@ services: condition: service_healthy jaeger: condition: service_healthy + dd-sink: + condition: service_healthy env_file: .env environment: LITELLM_MASTER_KEY: sk-1234 STORE_MODEL_IN_DB: "True" + DD_API_KEY: local-sink-noauth + DD_SITE: datadoghq.com + DD_BASE_URL: http://dd-sink:8080 LITELLM_OTEL_V2: "true" PHOENIX_COLLECTOR_HTTP_ENDPOINT: http://jaeger:4318/v1/traces PHOENIX_API_KEY: local-jaeger-noauth @@ -116,6 +157,8 @@ services: MISTRAL_API_KEY: ${MISTRAL_API_KEY:-} AZURE_API_BASE: ${AZURE_API_BASE:-} AZURE_API_KEY: ${AZURE_API_KEY:-} + AZURE_AI_API_BASE: ${AZURE_AI_API_BASE:-} + AZURE_AI_API_KEY: ${AZURE_AI_API_KEY:-} ports: - "4000:4000" configs: @@ -155,3 +198,19 @@ services: interval: 3s timeout: 3s retries: 20 + +# throwaway DataDog logs-intake sink (records POSTs, replays on GET /requests; +# see E2E_DD_SINK_URL) + dd-sink: + image: python:3.12-alpine + command: ["python", "/sink.py"] + configs: + - source: dd_sink_script + target: /sink.py + ports: + - "9915:8080" + healthcheck: + test: ["CMD", "wget", "-qO-", "http://127.0.0.1:8080/health"] + interval: 3s + timeout: 3s + retries: 20 diff --git a/tests/e2e/e2e_config.py b/tests/e2e/e2e_config.py index 6e6c30709de..e84438430fd 100644 --- a/tests/e2e/e2e_config.py +++ b/tests/e2e/e2e_config.py @@ -32,6 +32,10 @@ CHEAP_OPENAI_MODEL = os.environ.get("E2E_CHEAP_OPENAI_MODEL", "gpt-5.5") # read exported spans back through it. OTEL_QUERY_URL = os.environ.get("E2E_OTEL_QUERY_URL", "http://localhost:16686").rstrip("/") +# Query URL of the compose stack's DataDog logs-intake sink (the `dd-sink` +# service records every intake POST and replays them on GET /requests). +DD_SINK_URL = os.environ.get("E2E_DD_SINK_URL", "http://localhost:9915").rstrip("/") + # Writes on the proxy are eventually consistent (e.g. spend rows flush on # proxy_batch_write_at, ~60s). Read-backs poll to this deadline, never sleep-once. POLL_TIMEOUT = float(os.environ.get("E2E_POLL_TIMEOUT", "120")) diff --git a/tests/e2e/e2e_gateway.py b/tests/e2e/e2e_gateway.py index d40b96d60fa..ad8b2e833a8 100644 --- a/tests/e2e/e2e_gateway.py +++ b/tests/e2e/e2e_gateway.py @@ -319,21 +319,31 @@ class Gateway: return self.transport.probe(path, params=params) -def build_gateway() -> Gateway: +def build_gateway( + *, + base_url: str = PROXY_BASE_URL, + master_key: str = MASTER_KEY, + control_plane_base_url: str = CONTROL_PLANE_BASE_URL, +) -> Gateway: """The Gateway every suite's client is built from: a SplitTransport that routes LLM calls to the data plane (PROXY_BASE_URL) and management/admin calls to the control plane (CONTROL_PLANE_BASE_URL), with the shared poll budget. The two - base URLs are the same for a monolithic proxy, so routing is then a no-op.""" + base URLs are the same for a monolithic proxy, so routing is then a no-op. + + The endpoints are injectable for callers that resolve the proxy some other + way than ``e2e_config``'s env names (see ``claude_code/_env.py``); they must + pass all three together, since a caller that overrides only the data plane + would leave management calls pointed at the env default.""" return Gateway( transport=SplitTransport( data=HttpTransport( - base_url=PROXY_BASE_URL, - master_key=MASTER_KEY, + base_url=base_url, + master_key=master_key, request_timeout=REQUEST_TIMEOUT, ), control=HttpTransport( - base_url=CONTROL_PLANE_BASE_URL, - master_key=MASTER_KEY, + base_url=control_plane_base_url, + master_key=master_key, request_timeout=REQUEST_TIMEOUT, ), ), diff --git a/tests/e2e/e2e_result_reporter.py b/tests/e2e/e2e_result_reporter.py new file mode 100644 index 00000000000..22f7581818f --- /dev/null +++ b/tests/e2e/e2e_result_reporter.py @@ -0,0 +1,144 @@ +"""Structured e2e result lines for Loki / Grafana status history. + +Pytest progress lines are a bad dashboard source: they only expose file basenames, +break under quiet modes, and force status-history rows to explode with suite growth. + +Each finished test emits one logfmt line: + + E2E_RESULT package=logging file=test_langfuse_e2e.py outcome=failed + duration_ms=1234 node_id=logging/test_langfuse_e2e.py::TestX::test_y + covers=logging.langfuse.team.success + +Grafana package status-history queries max(fail) by package over E2E_RESULT lines. +Drill-down uses node_id / covers in Explore, not status-history cardinality. +""" + +from __future__ import annotations + +from collections.abc import Iterable, Sequence +from dataclasses import dataclass +from pathlib import Path +from typing import Literal, Protocol, runtime_checkable + +Outcome = Literal["passed", "failed", "error", "skipped"] + + +@dataclass(frozen=True, slots=True) +class E2EResult: + package: str + file: str + outcome: Outcome + duration_ms: int + node_id: str + covers: tuple[str, ...] + + +@runtime_checkable +class _MarkerArgs(Protocol): + args: Sequence[object] + + +@runtime_checkable +class _ItemWithCovers(Protocol): + def iter_markers(self, name: str) -> Iterable[object]: ... + + +def package_from_nodeid(nodeid: str) -> str: + """Top-level suite package under tests/e2e/, or 'root' for top-level files. + + Pytest nodeids are relative to the invocation cwd. Repo-root runs look like + `tests/e2e/logging/...`; suite-cwd runs look like `logging/...`. Strip the + `tests/e2e` prefix so package is the suite dir either way. + """ + path_part = nodeid.split("::", 1)[0].replace("\\", "/") + parts = tuple(p for p in path_part.split("/") if p and p != ".") + if len(parts) >= 3 and parts[0] == "tests" and parts[1] == "e2e": + parts = parts[2:] + if len(parts) <= 1: + return "root" + return parts[0] + + +def file_from_nodeid(nodeid: str) -> str: + path_part = nodeid.split("::", 1)[0].replace("\\", "/") + return Path(path_part).name + + +def covers_from_item(item: object) -> tuple[str, ...]: + """Read @pytest.mark.covers cell ids from a pytest Item.""" + if not isinstance(item, _ItemWithCovers): + return () + return tuple( + dict.fromkeys( + arg + for marker in item.iter_markers(name="covers") + if isinstance(marker, _MarkerArgs) + for arg in marker.args + if isinstance(arg, str) and arg + ) + ) + + +def outcome_from_report(when: str, failed: bool, skipped: bool, passed: bool) -> Outcome | None: + """Map pytest TestReport fields to a terminal outcome. None if not final.""" + if when == "setup" and skipped: + return "skipped" + if when == "setup" and failed: + return "error" + if when != "call": + return None + if skipped: + return "skipped" + if failed: + return "failed" + if passed: + return "passed" + return "failed" + + +def _logfmt_escape(value: str) -> str: + if value == "": + return '""' + needs_quote = any(ch.isspace() or ch in "\"=\\" for ch in value) + if not needs_quote: + return value + escaped = value.replace("\\", "\\\\").replace('"', '\\"') + return f'"{escaped}"' + + +def format_e2e_result_line(result: E2EResult) -> str: + covers = ",".join(result.covers) + fields = ( + ("package", result.package), + ("file", result.file), + ("outcome", result.outcome), + ("duration_ms", str(result.duration_ms)), + ("node_id", result.node_id), + ("covers", covers), + ) + body = " ".join(f"{key}={_logfmt_escape(value)}" for key, value in fields) + return f"E2E_RESULT {body}" + + +def result_from_pytest( + *, + nodeid: str, + when: str, + failed: bool, + skipped: bool, + passed: bool, + duration_seconds: float, + covers: tuple[str, ...] = (), +) -> E2EResult | None: + outcome = outcome_from_report(when=when, failed=failed, skipped=skipped, passed=passed) + if outcome is None: + return None + duration_ms = max(0, int(round(duration_seconds * 1000))) + return E2EResult( + package=package_from_nodeid(nodeid), + file=file_from_nodeid(nodeid), + outcome=outcome, + duration_ms=duration_ms, + node_id=nodeid, + covers=covers, + ) diff --git a/tests/e2e/grafana/status_history_panels.md b/tests/e2e/grafana/status_history_panels.md new file mode 100644 index 00000000000..f8cda509c63 --- /dev/null +++ b/tests/e2e/grafana/status_history_panels.md @@ -0,0 +1,66 @@ +# Grafana: package status history for e2e + +Dashboard: [LiteLLM E2E](https://berriai.grafana.net/d/mup2cfn/litellm-e2e) (`mup2cfn`). + +The old **test suite status history** panel scraped pytest progress lines and +grouped by **file basename** (`test_foo.py`). That does not scale: multi-class +files collapse to one bit, and full `node_id` cardinality melts status-history. + +## Emitter + +After each test finishes, `tests/e2e/conftest.py` prints one logfmt line: + +``` +E2E_RESULT package=logging file=test_langfuse_e2e.py outcome=failed duration_ms=1500 node_id="logging/..." covers=cell.id +``` + +## Panel: package status history (replace panel 11) + +**Type:** Status history +**Interval:** 15m (or 1h for multi-day ranges) +**Description:** Per top-level package under `tests/e2e/`: red if any test failed or errored in the bucket. + +```logql +max by (package) ( + max_over_time( + {service_name="litellm-e2e"} + |= "E2E_RESULT" + | logfmt + | outcome != "" + | label_format result=`{{ if or (eq .outcome "failed") (eq .outcome "error") }}1{{ else }}0{{ end }}` + | unwrap result + [$__interval] + ) +) +``` + +Value mappings: `0` → Pass (green), `1` → Fail (red). + +If `service_name` is missing on older scrapes, use: + +```logql +{cluster="berrie-litellm-stage", pod=~"litellm-e2e-.+"} +``` + +instead of `{service_name="litellm-e2e"}`. + +## Panel: failed tests (logs drill-down) + +```logql +{service_name="litellm-e2e"} |= "E2E_RESULT" | logfmt | outcome=~"failed|error" +``` + +Show fields: `package`, `file`, `node_id`, `covers`, `duration_ms`. + +## Panel (optional): filter by package variable + +Dashboard variable `package` (custom or from label_values on E2E_RESULT): + +```logql +{service_name="litellm-e2e"} |= "E2E_RESULT" | logfmt | package=`$package` | outcome=~"failed|error" +``` + +## Do not + +- Put full `node_id` as the status-history series key (cardinality). +- Rely on `::S+ PASSED` progress regex as the primary signal once E2E_RESULT is live. diff --git a/tests/e2e/logging/conftest.py b/tests/e2e/logging/conftest.py index 43d279602ef..5ae791917fd 100644 --- a/tests/e2e/logging/conftest.py +++ b/tests/e2e/logging/conftest.py @@ -11,6 +11,7 @@ import os import pytest from logging_client import LangfuseCreds, LoggingClient, build_logging_client, load_langfuse_creds +from datadog_sink import DdSinkReader, build_dd_sink_reader from otel_client import OtelReader, build_otel_reader @@ -35,6 +36,12 @@ def otel_reader() -> OtelReader: return build_otel_reader() +@pytest.fixture(scope="session") +def dd_sink() -> DdSinkReader: + """Read-back client for the compose stack's DataDog logs-intake sink.""" + return build_dd_sink_reader() + + @pytest.fixture def datadog_creds() -> None: """Require Datadog shipping credentials. Hard-fail when absent; never skip.""" diff --git a/tests/e2e/logging/datadog_sink.py b/tests/e2e/logging/datadog_sink.py new file mode 100644 index 00000000000..5b5059d1428 --- /dev/null +++ b/tests/e2e/logging/datadog_sink.py @@ -0,0 +1,103 @@ +"""Read-back for the DataDog logging tests: typed models over the compose +stack's dd-sink service, which records every logs-intake POST the datadog +callback sends (gunzipped) and replays them as JSON. + +Delivery is judged on what the sink actually received, mirroring how the OTEL +tests read Jaeger; a failed sink query is a hard failure, never an empty +result. External reads go through ``e2e_http``. +""" + +from __future__ import annotations + +import time +from dataclasses import dataclass + +import pytest +from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError + +from e2e_config import DD_SINK_URL, POLL_INTERVAL, POLL_TIMEOUT +from e2e_http import URL, NoBody, Success, get + + +class DdSinkRequest(BaseModel): + model_config = ConfigDict(extra="ignore") + + path: str + body: str + + +class DdSinkRequests(BaseModel): + model_config = ConfigDict(extra="ignore") + + requests: list[DdSinkRequest] = [] + + +class DdLogEvent(BaseModel): + model_config = ConfigDict(extra="ignore") + + message: str + ddsource: str | None = None + service: str | None = None + status: str | None = None + + +_EVENT_BATCH: TypeAdapter[list[DdLogEvent]] = TypeAdapter(list[DdLogEvent]) + + +def _parse_batch(request: DdSinkRequest) -> list[DdLogEvent]: + """The intake accepts an array of events or a single event object.""" + try: + return _EVENT_BATCH.validate_json(request.body) + except ValidationError: + try: + return [DdLogEvent.model_validate_json(request.body)] + except ValidationError: + pytest.fail(f"dd-sink recorded a non-log body on {request.path}: {request.body[:200]}") + + +@dataclass(frozen=True, slots=True) +class DdSinkReader: + sink_url: str + + def _recorded_requests(self) -> list[DdSinkRequest]: + result = get( + URL(f"{self.sink_url}/requests"), + headers=NoBody(), + params=NoBody(), + response_type=DdSinkRequests, + timeout=30.0, + ) + match result: + case Success(data=page): + return page.requests + case failure: + pytest.fail(f"dd-sink query at {self.sink_url} failed: {failure}") + + def events_for_marker(self, marker: str) -> list[DdLogEvent]: + """Every log event across every recorded intake batch whose message + carries the marker. More than one hit for one call IS the + duplicate-delivery bug, so this never collapses to a single event.""" + events: list[DdLogEvent] = [] + for request in self._recorded_requests(): + if "/api/v2/logs" not in request.path: + continue + events.extend(event for event in _parse_batch(request) if marker in event.message) + return events + + def poll_events_for_marker(self, marker: str) -> list[DdLogEvent]: + """Poll until at least one matching event lands (the callback flushes + in periodic batches), then re-read after one more interval so a late + duplicate cannot hide from the exactly-one assertion. At the deadline + the last result is returned as-is.""" + deadline = time.monotonic() + POLL_TIMEOUT + while time.monotonic() < deadline: + events = self.events_for_marker(marker) + if events: + time.sleep(POLL_INTERVAL) + return self.events_for_marker(marker) + time.sleep(POLL_INTERVAL) + return self.events_for_marker(marker) + + +def build_dd_sink_reader() -> DdSinkReader: + return DdSinkReader(sink_url=DD_SINK_URL) diff --git a/tests/e2e/logging/logging_client.py b/tests/e2e/logging/logging_client.py index e37d6175705..8be573d72a9 100644 --- a/tests/e2e/logging/logging_client.py +++ b/tests/e2e/logging/logging_client.py @@ -18,7 +18,7 @@ import json import os import time from dataclasses import dataclass -from typing import Literal +from typing import Callable, Literal import pytest from pydantic import BaseModel, ConfigDict, Field, JsonValue, TypeAdapter, ValidationError @@ -28,6 +28,7 @@ from e2e_gateway import Gateway, build_gateway from e2e_http import ( URL, AuthHeaders, + require_successful_call, NoBody, StreamingResponse, Success, @@ -617,5 +618,20 @@ class LoggingClient: return self.list_langfuse_observations(creds, trace_id=gen.trace_id) or [gen] +def first_ok(client: LoggingClient, send: Callable[[], StreamingResponse]) -> StreamingResponse: + """First successful call on a fresh key. A fresh key may briefly 401 until + the data plane's auth cache picks it up, so retry on 401 to a deadline; a + 401 is rejected before the LLM call, so it cannot contaminate delivery or + trace assertions. Any other failure is behavior under test and fails hard.""" + deadline = time.monotonic() + client.gateway.poll_timeout + while True: + outcome = send() + if outcome.ok: + return outcome + if outcome.status_code != 401 or time.monotonic() >= deadline: + require_successful_call(outcome) + time.sleep(client.gateway.poll_interval) + + def build_logging_client() -> LoggingClient: return LoggingClient(gateway=build_gateway()) diff --git a/tests/e2e/logging/test_datadog_log_e2e.py b/tests/e2e/logging/test_datadog_log_e2e.py new file mode 100644 index 00000000000..4651bbb28ba --- /dev/null +++ b/tests/e2e/logging/test_datadog_log_e2e.py @@ -0,0 +1,162 @@ +"""Live e2e: DataDog log delivery for successful non-streaming calls. + +Covers logging.datadog.success.exports_metric: one successful call on each +route must reach the DataDog logs intake as EXACTLY ONE log event whose +message (the StandardLoggingPayload) carries the model, the token counts, and +the response cost. Delivery is judged on what the intake actually received: +the compose stack's dd-sink service records every batch the datadog callback +ships (DD_BASE_URL override) and the tests read it back, so a dropped event, a +duplicated event, or a payload missing the cost all fail here. + +Both halves of the contract are asserted: the recorded state (the proxy +reports the DataDogLogger callback active via /health/readiness/details) and +the enforced behavior (the event at the intake, with the cost cross-checked +exactly against the x-litellm-response-cost header of the very response the +caller received). +""" + +from __future__ import annotations + +import pytest +from pydantic import BaseModel, ConfigDict + +from datadog_sink import DdLogEvent, DdSinkReader +from e2e_config import CHEAP_ANTHROPIC_MODEL, CHEAP_OPENAI_MODEL, unique_marker +from e2e_http import NoBody, StreamingResponse +from lifecycle import ResourceManager +from logging_client import LoggingClient, first_ok + +pytestmark = pytest.mark.e2e + +#: The active DataDog callback's name in /health/readiness/details success_callbacks. +DD_LOGGER_NAME = "DataDogLogger" + + +class _DdMessagePayload(BaseModel): + """The fields of the StandardLoggingPayload the scenario pins.""" + + model_config = ConfigDict(extra="ignore") + + model_group: str + total_tokens: int + response_cost: float + status: str + call_type: str + + +def _assert_datadog_configured(client: LoggingClient) -> None: + """Recorded state: the proxy reports the DataDog callback among its active + callbacks, so a missing destination config fails here, before any + delivery-based assertion can time out confusingly.""" + result = client.gateway.probe("/health/readiness/details", params=NoBody()) + assert result.status_code == 200, ( + f"/health/readiness/details must answer 200, got {result.status_code}: {result.body[:300]}" + ) + assert DD_LOGGER_NAME in result.body, ( + f"the proxy must report the {DD_LOGGER_NAME} callback active " + f"(callbacks + DD_* env in the compose config); got: {result.body[:400]}" + ) + + +def _assert_exactly_one_event( + events: list[DdLogEvent], *, model_group: str, call_type: str, outcome: StreamingResponse +) -> None: + """The enforced behavior: the intake holds exactly one event for the call, + sourced from litellm, whose payload names the model group and call type, + counts real tokens, and carries the same cost the response header reported.""" + assert events, "no DataDog log event for this call reached the intake within the deadline" + assert len(events) == 1, ( + f"expected exactly ONE DataDog log event for the call, got {len(events)} - " + "more than one event for one call is the duplicate-delivery bug (see LIT-4447 " + "for the currently known /v1/messages instance)" + ) + event = events[0] + assert event.ddsource == "litellm", f"event ddsource must be litellm, got {event.ddsource!r}" + assert event.status == "info", f"success events ship at status info, got {event.status!r}" + + payload = _DdMessagePayload.model_validate_json(event.message) + assert payload.status == "success", f"payload status must be success, got {payload.status!r}" + assert payload.model_group == model_group, ( + f"payload model_group must be {model_group!r}, got {payload.model_group!r}" + ) + assert payload.call_type == call_type, ( + f"payload call_type must be {call_type!r}, got {payload.call_type!r}" + ) + assert payload.total_tokens > 0, f"payload must count real tokens, got {payload.total_tokens}" + assert outcome.response_cost is not None and outcome.response_cost > 0, ( + f"the response must report x-litellm-response-cost, got {outcome.response_cost!r}" + ) + assert abs(payload.response_cost - outcome.response_cost) < 1e-12, ( + f"payload response_cost {payload.response_cost} must equal the response header " + f"cost {outcome.response_cost}" + ) + + +class TestDataDogLogDelivery: + @pytest.mark.covers("logging.datadog.success.exports_metric", exercised_on=["chat_completions"]) + def test_chat_completions_emits_one_log_event( + self, client: LoggingClient, dd_sink: DdSinkReader, resources: ResourceManager + ) -> None: + """One successful non-streaming /chat/completions call must reach the + DataDog logs intake as exactly one log event whose payload carries the + model, the token counts, and the response cost.""" + _assert_datadog_configured(client) + + key = client.key_with_alias(f"dd-chat-{unique_marker()}", models=[CHEAP_ANTHROPIC_MODEL]) + resources.defer(lambda: client.delete_key(key)) + + marker = unique_marker() + outcome = first_ok( + client, + lambda: client.chat_raw(key, CHEAP_ANTHROPIC_MODEL, f"reply with one word {marker}", max_tokens=16), + ) + events = dd_sink.poll_events_for_marker(marker) + _assert_exactly_one_event( + events, model_group=CHEAP_ANTHROPIC_MODEL, call_type="acompletion", outcome=outcome + ) + + @pytest.mark.covers("logging.datadog.success.exports_metric", exercised_on=["messages"]) + def test_messages_emits_one_log_event( + self, client: LoggingClient, dd_sink: DdSinkReader, resources: ResourceManager + ) -> None: + """One successful non-streaming /v1/messages call must reach the + DataDog logs intake as exactly one log event whose payload carries the + model, the token counts, and the response cost. + + This currently fails on the known /v1/messages double-log (LIT-4447); it goes green when the fix lands.""" + _assert_datadog_configured(client) + + key = client.key_with_alias(f"dd-messages-{unique_marker()}", models=[CHEAP_ANTHROPIC_MODEL]) + resources.defer(lambda: client.delete_key(key)) + + marker = unique_marker() + outcome = first_ok( + client, + lambda: client.messages_raw(key, CHEAP_ANTHROPIC_MODEL, f"reply with one word {marker}", max_tokens=16), + ) + events = dd_sink.poll_events_for_marker(marker) + _assert_exactly_one_event( + events, model_group=CHEAP_ANTHROPIC_MODEL, call_type="anthropic_messages", outcome=outcome + ) + + @pytest.mark.covers("logging.datadog.success.exports_metric", exercised_on=["responses"]) + def test_responses_emits_one_log_event( + self, client: LoggingClient, dd_sink: DdSinkReader, resources: ResourceManager + ) -> None: + """One successful non-streaming /v1/responses call must reach the + DataDog logs intake as exactly one log event whose payload carries the + model, the token counts, and the response cost.""" + _assert_datadog_configured(client) + + key = client.key_with_alias(f"dd-responses-{unique_marker()}", models=[CHEAP_OPENAI_MODEL]) + resources.defer(lambda: client.delete_key(key)) + + marker = unique_marker() + outcome = first_ok( + client, + lambda: client.responses_raw(key, CHEAP_OPENAI_MODEL, f"reply with one word {marker}"), + ) + events = dd_sink.poll_events_for_marker(marker) + _assert_exactly_one_event( + events, model_group=CHEAP_OPENAI_MODEL, call_type="aresponses", outcome=outcome + ) diff --git a/tests/e2e/logging/test_langfuse_e2e.py b/tests/e2e/logging/test_langfuse_e2e.py deleted file mode 100644 index d014b5d8291..00000000000 --- a/tests/e2e/logging/test_langfuse_e2e.py +++ /dev/null @@ -1,534 +0,0 @@ -"""Live e2e: Langfuse OTEL logs_spend for registry cells in logging.yaml P0. - -Registry cells: -- logging.langfuse.success.logs_spend (exercised_on chat_completions, messages, embeddings) -- logging.langfuse.failure.logs_spend (exercised_on chat_completions, messages) -- logging.langfuse.stream.logs_spend (exercised_on chat_completions, messages) - -Integration under test is ``langfuse_otel`` (OTLP to Langfuse), not the classic -``langfuse`` SDK callback. StandardLoggingPayload.response_cost is the spend -source of truth. Generations are named ``litellm_request``; correlate by unique -prompt marker and user_api_key_alias in metadata. - -Dynamic credentials by product surface: -- team: POST /team/{id}/callback with callback_name=langfuse_otel -- user/key: key metadata.logging with callback_name=langfuse_otel -- org: organization + team under it + team callback (no org-level callback API) - -Extra success paths assert tool calls and applied guardrails land on the trace. -""" - -from __future__ import annotations - -import json - -import pytest - -from e2e_config import unique_marker -from e2e_http import StreamingResponse, require_successful_call -from lifecycle import ResourceManager -from logging_client import ( - INVALID_UPSTREAM_API_KEY, - WEATHER_TOOL, - LangfuseCreds, - LoggingClient, - completion_response_id, - costs_agree, - observation_has_guardrail, - observation_mentions_tool, - observation_spend, -) -from models import LiteLLMParamsBody - -pytestmark = pytest.mark.e2e - -DRIVER_MODEL = "gemini-2.5-flash" -FAIL_BACKEND = "openai/gpt-4o-mini" - - -def _json_blob(value: object) -> str: - return json.dumps(value, default=str) - - -def _assert_logs_spend( - client: LoggingClient, - *, - key: str, - outcome: StreamingResponse, - obs_cost: float | None, - scope: str, - require_positive: bool = True, -) -> None: - """logs_spend: Langfuse cost matches StandardLogging response_cost and proxy spend. - - Non-stream responses expose response_cost on x-litellm-response-cost. Streaming - sends headers before final cost is known, so stream paths rely on /spend/logs. - """ - if not require_positive: - assert obs_cost is not None, ( - f"{scope}: failure path must still track spend (0 is fine); cost={obs_cost!r}" - ) - return - - assert obs_cost is not None and obs_cost > 0, ( - f"{scope}: Langfuse must log positive spend; calculatedTotalCost={obs_cost!r}" - ) - # Stream responses send headers before final cost is known, so the cost header - # is often absent; non-stream must always expose x-litellm-response-cost. - if not outcome.is_streaming: - assert outcome.response_cost is not None and outcome.response_cost > 0, ( - f"{scope}: proxy must return positive x-litellm-response-cost; " - f"got {outcome.response_cost!r}" - ) - assert costs_agree(outcome.response_cost, obs_cost), ( - f"{scope}: Langfuse cost {obs_cost!r} disagrees with " - f"x-litellm-response-cost {outcome.response_cost!r}" - ) - elif outcome.response_cost is not None and outcome.response_cost > 0: - assert costs_agree(outcome.response_cost, obs_cost), ( - f"{scope}: Langfuse cost {obs_cost!r} disagrees with " - f"x-litellm-response-cost {outcome.response_cost!r}" - ) - spend_row = client.poll_proxy_spend_for_key( - key, - response_id=completion_response_id(outcome.body), - require_positive_spend=True, - ) - assert spend_row is not None and spend_row.spend is not None and spend_row.spend > 0, ( - f"{scope}: proxy /spend/logs never produced a positive spend row for key" - ) - assert costs_agree(spend_row.spend, obs_cost), ( - f"{scope}: Langfuse cost {obs_cost!r} disagrees with proxy spend " - f"{spend_row.spend!r} (request_id={spend_row.request_id!r})" - ) - - -class TestLangfuseTeamLogging: - """Team-scoped callback via POST /team/{id}/callback.""" - - def _team_key( - self, - client: LoggingClient, - resources: ResourceManager, - creds: LangfuseCreds, - *, - models: list[str], - organization_id: str | None = None, - ) -> tuple[str, str, str]: - marker = unique_marker() - key_alias = f"e2e-lf-team-key-{marker}" - team_id = client.create_team( - f"e2e-lf-team-{marker}", - models=models, - organization_id=organization_id, - ) - resources.defer(lambda: client.delete_team(team_id)) - client.add_team_langfuse_callback(team_id, creds) - key = client.key_with_alias(key_alias, models=models, team_id=team_id) - resources.defer(lambda: client.delete_key(key)) - return team_id, key, key_alias - - @pytest.mark.covers("logging.langfuse.success.logs_spend", exercised_on=["chat_completions"]) - def test_success_logs_spend( - self, - client: LoggingClient, - resources: ResourceManager, - langfuse_creds: LangfuseCreds, - ) -> None: - _, key, key_alias = self._team_key( - client, resources, langfuse_creds, models=[DRIVER_MODEL] - ) - prompt_marker = unique_marker() - outcome = client.chat_raw( - key, DRIVER_MODEL, f"reply with one word only {prompt_marker}" - ) - require_successful_call(outcome) - - obs = client.poll_langfuse_observation( - langfuse_creds, - key_alias=key_alias, - prompt_marker=prompt_marker, - require_positive_cost=True, - ) - assert obs is not None, ( - f"team scope: Langfuse never received generation for key_alias={key_alias!r}" - ) - _assert_logs_spend( - client, - key=key, - outcome=outcome, - obs_cost=observation_spend(obs), - scope="team-success", - ) - - @pytest.mark.covers("logging.langfuse.failure.logs_spend", exercised_on=["chat_completions"]) - def test_failure_logs_spend( - self, - client: LoggingClient, - resources: ResourceManager, - langfuse_creds: LangfuseCreds, - ) -> None: - """Provider-auth failure still ships a Langfuse observation with spend tracked. - - Uses a throwaway deployment whose upstream OpenAI key is - INVALID_UPSTREAM_API_KEY (not a LiteLLM virtual key). - """ - prompt_marker = unique_marker() - model_name = f"e2e-lf-fail-{prompt_marker}" - model_id = client.create_model( - model_name, - LiteLLMParamsBody(model=FAIL_BACKEND, api_key=INVALID_UPSTREAM_API_KEY), - ) - resources.defer(lambda: client.delete_model(model_id)) - - _, key, key_alias = self._team_key( - client, resources, langfuse_creds, models=[model_name] - ) - outcome = client.chat_raw(key, model_name, f"this must fail {prompt_marker}") - assert not outcome.ok, ( - f"expected upstream provider failure for {INVALID_UPSTREAM_API_KEY!r}, " - f"got {outcome.status_code}: {outcome.body[:200]}" - ) - - obs = client.poll_langfuse_observation( - langfuse_creds, - key_alias=key_alias, - prompt_marker=prompt_marker, - require_positive_cost=False, - ) - assert obs is not None, ( - f"team failure path: Langfuse never received generation for key_alias={key_alias!r}" - ) - _assert_logs_spend( - client, - key=key, - outcome=outcome, - obs_cost=observation_spend(obs), - scope="team-failure", - require_positive=False, - ) - - @pytest.mark.covers("logging.langfuse.stream.logs_spend", exercised_on=["chat_completions"]) - def test_stream_logs_spend( - self, - client: LoggingClient, - resources: ResourceManager, - langfuse_creds: LangfuseCreds, - ) -> None: - _, key, key_alias = self._team_key( - client, resources, langfuse_creds, models=[DRIVER_MODEL] - ) - prompt_marker = unique_marker() - outcome = client.chat_raw( - key, DRIVER_MODEL, f"reply with one word only {prompt_marker}", stream=True - ) - require_successful_call(outcome) - assert outcome.is_streaming - assert outcome.chunks > 0 - - obs = client.poll_langfuse_observation( - langfuse_creds, - key_alias=key_alias, - prompt_marker=prompt_marker, - require_positive_cost=True, - ) - assert obs is not None - # Streamed body is elided; correlate cost via header + key spend row. - _assert_logs_spend( - client, - key=key, - outcome=outcome, - obs_cost=observation_spend(obs), - scope="team-stream", - ) - - @pytest.mark.covers("logging.langfuse.success.logs_spend", exercised_on=["chat_completions"]) - def test_tool_calls_logged_with_cost( - self, - client: LoggingClient, - resources: ResourceManager, - langfuse_creds: LangfuseCreds, - ) -> None: - _, key, key_alias = self._team_key( - client, resources, langfuse_creds, models=[DRIVER_MODEL] - ) - prompt_marker = unique_marker() - outcome = client.chat_raw( - key, - DRIVER_MODEL, - f"Use get_weather for Paris. marker={prompt_marker}", - tools=[WEATHER_TOOL], - tool_choice="required", - max_tokens=128, - ) - require_successful_call(outcome) - assert "get_weather" in outcome.body or "tool_calls" in outcome.body, ( - f"gateway response must include a tool call; body={outcome.body[:300]}" - ) - - obs = client.poll_langfuse_observation( - langfuse_creds, - key_alias=key_alias, - prompt_marker=prompt_marker, - require_positive_cost=True, - ) - assert obs is not None - assert observation_mentions_tool(obs, "get_weather"), ( - f"Langfuse generation must record the tool; name={obs.name!r} " - f"input={str(obs.input)[:200]} output={str(obs.output)[:200]}" - ) - _assert_logs_spend( - client, - key=key, - outcome=outcome, - obs_cost=observation_spend(obs), - scope="team-tools", - ) - - @pytest.mark.covers("logging.langfuse.success.logs_spend", exercised_on=["chat_completions"]) - def test_tool_permission_guardrail_logged( - self, - client: LoggingClient, - resources: ResourceManager, - langfuse_creds: LangfuseCreds, - ) -> None: - """tool_permission post_call guardrail must appear on the Langfuse trace - (StandardLogging guardrail_information -> Langfuse guardrail span).""" - marker = unique_marker() - guardrail_name = f"e2e-lf-tool-perm-{marker}" - guardrail_id = client.create_tool_permission_guardrail( - guardrail_name, allowed_tool="get_weather" - ) - resources.defer(lambda: client.delete_guardrail(guardrail_id)) - - _, key, key_alias = self._team_key( - client, resources, langfuse_creds, models=[DRIVER_MODEL] - ) - prompt_marker = unique_marker() - outcome = client.chat_raw( - key, - DRIVER_MODEL, - f"Use get_weather for Berlin. marker={prompt_marker}", - tools=[WEATHER_TOOL], - tool_choice="required", - guardrails=[guardrail_name], - max_tokens=128, - ) - require_successful_call(outcome) - - observations = client.poll_langfuse_trace_observations( - langfuse_creds, key_alias=key_alias, prompt_marker=prompt_marker - ) - assert observations, ( - f"team+guardrail: no Langfuse observations for key_alias={key_alias!r}" - ) - gen = next( - ( - o - for o in observations - if prompt_marker in _json_blob(o.input) - or key_alias in _json_blob(o.metadata) - or o.name in (f"litellm:{key_alias}", "litellm_request") - ), - observations[0], - ) - _assert_logs_spend( - client, - key=key, - outcome=outcome, - obs_cost=observation_spend(gen), - scope="team-guardrail", - ) - assert any( - observation_has_guardrail(o, guardrail_name=guardrail_name) - or (o.name is not None and "guardrail" in o.name.lower()) - for o in observations - ), ( - f"Langfuse trace must include applied guardrail {guardrail_name!r}; " - f"observation names={[o.name for o in observations]}" - ) - - -class TestLangfuseUserKeyLogging: - """User-owned key with metadata.logging (key-level dynamic Langfuse credentials). - - Product surface: key metadata.logging on /key/generate, not a separate - /user/.../callback route. The key is bound to a real /user/new user_id. - """ - - def _user_key( - self, - client: LoggingClient, - resources: ResourceManager, - creds: LangfuseCreds, - *, - models: list[str], - ) -> tuple[str, str, str]: - marker = unique_marker() - key_alias = f"e2e-lf-user-key-{marker}" - user_id = client.create_user( - user_email=f"e2e-lf-user-{marker}@example.com", - user_id=f"e2e-lf-user-{marker}", - ) - resources.defer(lambda: client.delete_user(user_id)) - key = client.key_with_alias( - key_alias, - models=models, - user_id=user_id, - metadata=creds.key_logging_metadata(), - ) - resources.defer(lambda: client.delete_key(key)) - return user_id, key, key_alias - - @pytest.mark.covers("logging.langfuse.success.logs_spend", exercised_on=["chat_completions"]) - def test_success_logs_spend( - self, - client: LoggingClient, - resources: ResourceManager, - langfuse_creds: LangfuseCreds, - ) -> None: - user_id, key, key_alias = self._user_key( - client, resources, langfuse_creds, models=[DRIVER_MODEL] - ) - prompt_marker = unique_marker() - outcome = client.chat_raw( - key, DRIVER_MODEL, f"reply with one word only {prompt_marker}" - ) - require_successful_call(outcome) - - obs = client.poll_langfuse_observation( - langfuse_creds, - key_alias=key_alias, - prompt_marker=prompt_marker, - require_positive_cost=True, - ) - assert obs is not None, ( - f"user/key scope: Langfuse never received generation for key_alias={key_alias!r}" - ) - meta_blob = _json_blob(obs.metadata) - assert user_id in meta_blob or key_alias in (obs.name or ""), ( - f"user/key scope should attribute the user or key; metadata={meta_blob[:300]}" - ) - _assert_logs_spend( - client, - key=key, - outcome=outcome, - obs_cost=observation_spend(obs), - scope="user-key", - ) - - @pytest.mark.covers("logging.langfuse.success.logs_spend", exercised_on=["chat_completions"]) - def test_tool_calls_logged_with_cost( - self, - client: LoggingClient, - resources: ResourceManager, - langfuse_creds: LangfuseCreds, - ) -> None: - _, key, key_alias = self._user_key( - client, resources, langfuse_creds, models=[DRIVER_MODEL] - ) - prompt_marker = unique_marker() - outcome = client.chat_raw( - key, - DRIVER_MODEL, - f"Use get_weather for Tokyo. marker={prompt_marker}", - tools=[WEATHER_TOOL], - tool_choice="required", - max_tokens=128, - ) - require_successful_call(outcome) - - obs = client.poll_langfuse_observation( - langfuse_creds, - key_alias=key_alias, - prompt_marker=prompt_marker, - require_positive_cost=True, - ) - assert obs is not None - assert observation_mentions_tool(obs, "get_weather"), ( - f"user/key tool path: tool missing from Langfuse; output={str(obs.output)[:200]}" - ) - _assert_logs_spend( - client, - key=key, - outcome=outcome, - obs_cost=observation_spend(obs), - scope="user-key-tools", - ) - - -class TestLangfuseOrgScopedLogging: - """Org-scoped run: organization + team under it + team Langfuse callback. - - There is no /organization/.../callback today; logging attaches at the team - (or key) under the org. This class proves org-linked team keys still deliver - accurate Langfuse spend and team attribution (StandardLogging metadata - user_api_key_team_id / user_api_key_org_id). - """ - - def _org_team_key( - self, - client: LoggingClient, - resources: ResourceManager, - creds: LangfuseCreds, - *, - models: list[str], - ) -> tuple[str, str, str, str]: - marker = unique_marker() - key_alias = f"e2e-lf-org-key-{marker}" - org_id = client.create_org(f"e2e-lf-org-{marker}", models=models) - resources.defer(lambda: client.delete_org(org_id)) - team_id = client.create_team( - f"e2e-lf-org-team-{marker}", - models=models, - organization_id=org_id, - ) - resources.defer(lambda: client.delete_team(team_id)) - client.add_team_langfuse_callback(team_id, creds) - key = client.key_with_alias( - key_alias, - models=models, - team_id=team_id, - organization_id=org_id, - ) - resources.defer(lambda: client.delete_key(key)) - return org_id, team_id, key, key_alias - - @pytest.mark.covers("logging.langfuse.success.logs_spend", exercised_on=["chat_completions"]) - def test_success_logs_spend_with_team_attribution( - self, - client: LoggingClient, - resources: ResourceManager, - langfuse_creds: LangfuseCreds, - ) -> None: - org_id, team_id, key, key_alias = self._org_team_key( - client, resources, langfuse_creds, models=[DRIVER_MODEL] - ) - prompt_marker = unique_marker() - outcome = client.chat_raw( - key, DRIVER_MODEL, f"reply with one word only {prompt_marker}" - ) - require_successful_call(outcome) - - obs = client.poll_langfuse_observation( - langfuse_creds, - key_alias=key_alias, - prompt_marker=prompt_marker, - require_positive_cost=True, - ) - assert obs is not None, ( - f"org scope: Langfuse never received generation for key_alias={key_alias!r}" - ) - meta_blob = _json_blob(obs.metadata) - assert team_id in meta_blob, ( - f"org-scoped team key must stamp team_id on Langfuse metadata; " - f"team_id={team_id!r} metadata={meta_blob[:400]}" - ) - _ = org_id - _assert_logs_spend( - client, - key=key, - outcome=outcome, - obs_cost=observation_spend(obs), - scope="org-team", - ) diff --git a/tests/e2e/logging/test_otel_trace_e2e.py b/tests/e2e/logging/test_otel_trace_e2e.py index b00fd91be3c..e49dd311b32 100644 --- a/tests/e2e/logging/test_otel_trace_e2e.py +++ b/tests/e2e/logging/test_otel_trace_e2e.py @@ -18,15 +18,14 @@ destination's own query API - never proxy-side "export succeeded" logs). from __future__ import annotations import time -from collections.abc import Callable import pytest from pydantic import BaseModel, ConfigDict, ValidationError from e2e_config import CHEAP_ANTHROPIC_MODEL, CHEAP_OPENAI_MODEL, unique_marker -from e2e_http import NoBody, StreamingResponse, require_successful_call +from e2e_http import NoBody from lifecycle import ResourceManager -from logging_client import INVALID_UPSTREAM_API_KEY, LoggingClient +from logging_client import INVALID_UPSTREAM_API_KEY, LoggingClient, first_ok from models import LiteLLMParamsBody from otel_client import JaegerSpan, JaegerTrace, OtelReader @@ -60,22 +59,6 @@ def _assert_otel_destination_configured(client: LoggingClient) -> None: ) -def _first_ok(client: LoggingClient, send: Callable[[], StreamingResponse]) -> StreamingResponse: - """First successful call on a fresh key. A fresh key may briefly 401 until - the data plane's auth cache picks it up, so retry on 401 to a deadline; a - 401 is rejected before the LLM call so it exports no gen-AI span and cannot - contaminate the trace assertions. Any other failure is behavior under test - and fails hard.""" - deadline = time.monotonic() + client.gateway.poll_timeout - while True: - outcome = send() - if outcome.ok: - return outcome - if outcome.status_code != 401 or time.monotonic() >= deadline: - require_successful_call(outcome) - time.sleep(client.gateway.poll_interval) - - def _parent_ids(span_id: str, trace: JaegerTrace) -> list[str]: span = next(s for s in trace.spans if s.span_id == span_id) return [ref.span_id for ref in span.references if ref.ref_type == "CHILD_OF"] @@ -253,7 +236,7 @@ class TestOtelTraceCompleteness: resources.defer(lambda: client.delete_key(key)) marker = unique_marker() - outcome = _first_ok( + outcome = first_ok( client, lambda: client.chat_raw(key, MODEL, f"reply with one word {marker}", max_tokens=16) ) assert outcome.call_id is not None, "success response must carry x-litellm-call-id" @@ -286,7 +269,7 @@ class TestOtelTraceCompleteness: resources.defer(lambda: client.delete_key(key)) marker = unique_marker() - outcome = _first_ok( + outcome = first_ok( client, lambda: client.messages_raw(key, MODEL, f"reply with one word {marker}", max_tokens=16) ) assert outcome.call_id is not None, "success response must carry x-litellm-call-id" @@ -319,7 +302,7 @@ class TestOtelTraceCompleteness: resources.defer(lambda: client.delete_key(key)) marker = unique_marker() - outcome = _first_ok( + outcome = first_ok( client, lambda: client.responses_raw(key, CHEAP_OPENAI_MODEL, f"reply with one word {marker}"), ) @@ -360,7 +343,7 @@ class TestOtelTraceCompleteness: resources.defer(lambda: client.delete_key(key)) marker = unique_marker() - outcome = _first_ok( + outcome = first_ok( client, lambda: client.chat_raw(key, MODEL, f"reply with one word {marker}", stream=True, max_tokens=16), ) @@ -416,7 +399,7 @@ class TestOtelTraceCompleteness: resources.defer(lambda: client.delete_key(key)) marker = unique_marker() - outcome = _first_ok( + outcome = first_ok( client, lambda: client.messages_raw(key, MODEL, f"reply with one word {marker}", max_tokens=16, stream=True), ) @@ -474,7 +457,7 @@ class TestOtelTraceCompleteness: resources.defer(lambda: client.delete_key(key)) marker = unique_marker() - outcome = _first_ok( + outcome = first_ok( client, lambda: client.responses_raw(key, CHEAP_OPENAI_MODEL, f"reply with one word {marker}", stream=True), ) diff --git a/tests/e2e/management/conftest.py b/tests/e2e/management/conftest.py index 47c1d34baae..264108f6089 100644 --- a/tests/e2e/management/conftest.py +++ b/tests/e2e/management/conftest.py @@ -50,8 +50,10 @@ def ui_page(browser: "Browser") -> "Iterator[Page]": page.goto(f"{PROXY_BASE_URL}/ui/") page.fill("#username", UI_USERNAME) page.fill("#password", UI_PASSWORD) - page.click('input[type="submit"]') - page.wait_for_url("**/ui/**") + page.click('button[type="submit"]') + page.wait_for_function( + "() => document.cookie.includes('token=') || !document.querySelector('#username')" + ) yield page finally: context.close() diff --git a/tests/e2e/models.py b/tests/e2e/models.py index 4140967f3e0..c4ecd0cfa63 100644 --- a/tests/e2e/models.py +++ b/tests/e2e/models.py @@ -440,6 +440,8 @@ class LiteLLMParamsBody(BaseModel): aws_batch_role_arn: str | None = None input_cost_per_token: float | None = None output_cost_per_token: float | None = None + extra_headers: dict[str, str] | None = None + use_in_pass_through: bool | None = None ModelMode = Literal["batch", "realtime", "image_generation"] diff --git a/tests/e2e/test_e2e_gateway.py b/tests/e2e/test_e2e_gateway.py deleted file mode 100644 index 70e846ad8d5..00000000000 --- a/tests/e2e/test_e2e_gateway.py +++ /dev/null @@ -1,273 +0,0 @@ -"""Unit coverage for the Gateway model-management surface (create_model / -delete_model) and the bounded spend read-back (spend_logs_window). - -The batches conftest and several llm_translation tests register deployments at -runtime through gateway.create_model; when that method went missing, every batch -test errored at fixture setup (AttributeError) before a single request reached -the proxy. This pins the surface with a typed fake Transport so a rename or -signature drift fails here instead of in a live stage run. - -spend_logs_window exists because the unpaginated /spend/logs whole-table read -grew past the e2e runner's memory limit on stage and OOMKilled every run; these -tests pin its /spend/logs/v2 pagination and that SpendLogsParams can no longer -express the unfiltered read. -""" - -from dataclasses import dataclass, field -from datetime import datetime, timezone - -import pytest -from pydantic import BaseModel, ValidationError - -from batches.batch_client import BatchClient -from e2e_gateway import Gateway -from e2e_http import ( - AuthHeaders, - FileUploadForm, - ProbeResult, - Result, - StreamingResponse, - Success, - UnknownApiError, -) -from models import ( - LiteLLMParamsBody, - ModelDeleteBody, - ModelNewBody, - ModelNewResponse, - ModelsListResponse, - SpendLogsPage, - SpendLogsPageParams, - SpendLogsParams, -) - - -@dataclass -class _RecordingTransport: - """Typed fake fulfilling the Transport protocol; records every post and - answers with a canned success so the test asserts on what was sent. - - `get("/v1/models")` reports a created model as servable only after - `servable_after_gets` polls, so a test can drive the data-plane wait in - create_model.""" - - posts: list[tuple[str, BaseModel]] = field(default_factory=list) - servable_after_gets: int = 0 - models_error: UnknownApiError | None = None - model_gets: int = 0 - spend_total: int = 0 - spend_gets: list[SpendLogsPageParams] = field(default_factory=list) - _created: list[str] = field(default_factory=list) - - def post[R: BaseModel]( - self, path: str, *, headers: BaseModel, json: BaseModel, response_type: type[R] - ) -> Result[R]: - self.posts.append((path, json)) - if path == "/model/new" and isinstance(json, ModelNewBody): - self._created.append(json.model_name) - payload = ( - {"model_id": "registered-id"} if response_type is ModelNewResponse else {} - ) - return Success(data=response_type.model_validate(payload)) - - def stream( - self, path: str, *, headers: BaseModel, json: BaseModel - ) -> StreamingResponse: - raise AssertionError("stream is not part of model management") - - def send( - self, - path: str, - *, - headers: BaseModel, - json: BaseModel, - params: BaseModel | None = None, - stream: bool = False, - ) -> StreamingResponse: - raise AssertionError("send is not part of model management") - - def get[R: BaseModel]( - self, - path: str, - *, - headers: BaseModel, - params: BaseModel, - response_type: type[R], - ) -> Result[R]: - if path == "/v1/models" and response_type is ModelsListResponse: - self.model_gets += 1 - if self.models_error is not None: - return self.models_error - visible = self._created if self.model_gets > self.servable_after_gets else [] - return Success( - data=response_type.model_validate({"data": [{"id": name} for name in visible]}) - ) - if path == "/spend/logs/v2" and response_type is SpendLogsPage: - assert isinstance(params, SpendLogsPageParams) - self.spend_gets.append(params) - offset = (params.page - 1) * params.page_size - count = min(params.page_size, max(self.spend_total - offset, 0)) - return Success( - data=response_type.model_validate( - { - "data": [{"request_id": f"req-{offset + i}"} for i in range(count)], - "total": self.spend_total, - "page": params.page, - "page_size": params.page_size, - "total_pages": (self.spend_total + params.page_size - 1) // params.page_size, - } - ) - ) - raise AssertionError(f"unexpected get: {path}") - - def delete[R: BaseModel]( - self, path: str, *, headers: BaseModel, json: BaseModel, response_type: type[R] - ) -> Result[R]: - raise AssertionError("delete is not part of model management") - - def probe(self, path: str, *, params: BaseModel) -> ProbeResult: - raise AssertionError("probe is not part of model management") - - def upload[R: BaseModel]( - self, - path: str, - *, - headers: BaseModel, - form: FileUploadForm, - filename: str, - content: bytes, - params: BaseModel | None = None, - response_type: type[R], - ) -> Result[R]: - raise AssertionError("upload is not part of model management") - - def download(self, path: str, *, headers: BaseModel) -> StreamingResponse: - raise AssertionError("download is not part of model management") - - def bearer(self, key: str) -> AuthHeaders: - return AuthHeaders(authorization=f"Bearer {key}") - - @property - def master(self) -> AuthHeaders: - return self.bearer("sk-test-master") - - -def test_gateway_create_model_registers_deployment_and_returns_model_id() -> None: - transport = _RecordingTransport() - gateway = Gateway(transport=transport, poll_interval=0.0) - - model_id = gateway.create_model( - "e2e-test-model", LiteLLMParamsBody(model="openai/gpt-4o-mini") - ) - - assert model_id == "registered-id" - path, body = transport.posts[0] - assert path == "/model/new" - assert isinstance(body, ModelNewBody) - assert body.model_name == "e2e-test-model" - # No pinned model_id: the proxy assigns a unique one, so a fixed-name model - # re-registered after a failed teardown can't collide on the id constraint. - assert body.model_info.id is None - assert body.model_info.mode is None - # It confirmed data-plane visibility before returning. - assert transport.model_gets >= 1 - - -def test_gateway_create_model_waits_until_servable_on_the_data_plane() -> None: - # The model shows up on /v1/models only on the third poll (simulating the - # gateway's delayed DB reload in a split deployment); create_model must keep - # polling instead of returning after /model/new. - transport = _RecordingTransport(servable_after_gets=2) - gateway = Gateway(transport=transport, poll_interval=0.0) - - gateway.create_model("e2e-late-model", LiteLLMParamsBody(model="openai/gpt-4o-mini")) - - assert transport.model_gets == 3 - - -def test_gateway_create_model_fails_loudly_when_never_servable() -> None: - transport = _RecordingTransport(servable_after_gets=10**9) - gateway = Gateway(transport=transport, poll_timeout=0.05, poll_interval=0.0) - - with pytest.raises(AssertionError, match="never became servable"): - gateway.create_model("e2e-ghost-model", LiteLLMParamsBody(model="openai/gpt-4o-mini")) - - -def test_gateway_create_model_surfaces_the_last_data_plane_error() -> None: - transport = _RecordingTransport( - models_error=UnknownApiError(status_code=503, body="data plane down") - ) - gateway = Gateway(transport=transport, poll_timeout=0.05, poll_interval=0.0) - - with pytest.raises(AssertionError, match="data plane down") as excinfo: - gateway.create_model("e2e-flaky-model", LiteLLMParamsBody(model="openai/gpt-4o-mini")) - assert "503" in str(excinfo.value) - - -def test_batch_client_create_model_registers_a_batch_mode_deployment() -> None: - transport = _RecordingTransport() - client = BatchClient(gateway=Gateway(transport=transport, poll_interval=0.0)) - - model_id = client.create_model( - "e2e-batch-model", LiteLLMParamsBody(model="openai/gpt-4o-mini") - ) - - assert model_id == "registered-id" - path, body = transport.posts[0] - assert path == "/model/new" - assert isinstance(body, ModelNewBody) - assert body.model_info.mode == "batch" - - -def test_gateway_delete_model_posts_the_model_id() -> None: - transport = _RecordingTransport() - gateway = Gateway(transport=transport) - - gateway.delete_model("registered-id") - - path, body = transport.posts[0] - assert path == "/model/delete" - assert isinstance(body, ModelDeleteBody) - assert body.id == "registered-id" - - -WINDOW_START = datetime(2026, 7, 14, 12, 0, 0, tzinfo=timezone.utc) -WINDOW_END = datetime(2026, 7, 14, 14, 0, 0, tzinfo=timezone.utc) - - -def test_gateway_spend_logs_window_pages_through_every_row_in_the_window() -> None: - transport = _RecordingTransport(spend_total=250) - gateway = Gateway(transport=transport) - - rows = gateway.spend_logs_window(start=WINDOW_START, end=WINDOW_END) - - assert len(rows) == 250 - assert len({row.request_id for row in rows}) == 250 - assert [params.page for params in transport.spend_gets] == [1, 2, 3] - assert all(params.start_date == "2026-07-14 12:00:00" for params in transport.spend_gets) - assert all(params.end_date == "2026-07-14 14:00:00" for params in transport.spend_gets) - - -def test_gateway_spend_logs_window_stops_at_an_exact_page_boundary() -> None: - transport = _RecordingTransport(spend_total=200) - gateway = Gateway(transport=transport) - - rows = gateway.spend_logs_window(start=WINDOW_START, end=WINDOW_END) - - assert len(rows) == 200 - assert [params.page for params in transport.spend_gets] == [1, 2] - - -def test_gateway_spend_logs_window_returns_empty_for_an_empty_window() -> None: - transport = _RecordingTransport(spend_total=0) - gateway = Gateway(transport=transport) - - rows = gateway.spend_logs_window(start=WINDOW_START, end=WINDOW_END) - - assert rows == [] - assert [params.page for params in transport.spend_gets] == [1] - - -def test_spend_logs_params_rejects_the_unfiltered_whole_table_read() -> None: - with pytest.raises(ValidationError, match="spend_logs_window"): - SpendLogsParams() diff --git a/tests/e2e/test_lifecycle.py b/tests/e2e/test_lifecycle.py deleted file mode 100644 index d3c559dd2ed..00000000000 --- a/tests/e2e/test_lifecycle.py +++ /dev/null @@ -1,46 +0,0 @@ -"""Unit coverage for the lifecycle harness (lifecycle.run_case). - -Cases register cleanups progressively during init() (create team, then user, then -key), so a failure partway through init() must still release whatever was already -created on the long-lived shared proxy. This guards that contract. -""" - -from dataclasses import dataclass, field -from typing import Callable, List - -import pytest - -from lifecycle import run_case - - -@dataclass -class _PartialInitCase: - """init() registers a cleanup, then raises before finishing - mirroring a real - case that creates a resource, registers its delete, then fails on the next - step.""" - - released: List[str] = field(default_factory=list) - _undo: List[Callable[[], None]] = field(default_factory=list) - - def init(self) -> None: - self._undo.append(lambda: self.released.append("first")) - raise RuntimeError("init failed after registering the first resource") - - def run(self) -> None: - raise AssertionError("run() must not execute when init() failed") - - def teardown(self) -> None: - for undo in reversed(self._undo): - undo() - - -def test_run_case_releases_resources_when_init_fails_partway() -> None: - case = _PartialInitCase() - - with pytest.raises(RuntimeError, match="init failed"): - run_case(case) - - assert case.released == ["first"], ( - "a resource registered before init() failed must still be released, or it " - "leaks on the long-lived shared proxy" - ) diff --git a/tests/e2e/test_transport.py b/tests/e2e/test_transport.py deleted file mode 100644 index c7ce61b90c1..00000000000 --- a/tests/e2e/test_transport.py +++ /dev/null @@ -1,52 +0,0 @@ -"""Unit coverage for SplitTransport path routing (is_control_plane_path). - -Model-management calls (/model/new, /model/delete, /model/info) must go to the -control plane: the data-plane gateway does not serve management routes, so a -misrouted /model/new 404s and takes down every suite that registers deployments -at runtime (llm_translation, batches, access_control). /models must stay on the -data plane; it is the OpenAI-compatible list-models route, not a management -route. -""" - -import pytest - -from transport import is_control_plane_path - - -@pytest.mark.parametrize( - "path", - [ - "/model/new", - "/model/delete", - "/model/update", - "/model/info", - "/key/generate", - "/budget/new", - "/spend/logs", - "/end_user/daily/activity", - "/user/daily/activity", - "/team/daily/activity", - "/tag/daily/activity", - ], -) -def test_management_routes_go_to_the_control_plane(path: str) -> None: - assert is_control_plane_path(path), ( - f"{path} is a management route; sending it to the data plane 404s" - ) - - -@pytest.mark.parametrize( - "path", - [ - "/models", - "/v1/models", - "/chat/completions", - "/v1/messages", - "/embeddings", - "/anthropic/v1/messages", - ], -) -def test_llm_routes_stay_on_the_data_plane(path: str) -> None: - assert not is_control_plane_path(path), ( - f"{path} is an LLM route; it must go to the data plane" - ) diff --git a/tests/llm_translation/reasoning_effort_grid/grid_spec.py b/tests/llm_translation/reasoning_effort_grid/grid_spec.py index 762606bb3a0..c47fddb1d8d 100644 --- a/tests/llm_translation/reasoning_effort_grid/grid_spec.py +++ b/tests/llm_translation/reasoning_effort_grid/grid_spec.py @@ -212,12 +212,6 @@ AZURE_AI_MODELS: Tuple[ModelEntry, ...] = ( mode="adaptive", required_env=_AZURE_FOUNDRY_REQ, caps=_CAPS_XHIGH_MAX, - fail_reason=( - "claude-fable-5 has no deployment on the CI Microsoft Foundry " - "resource yet; Foundry returns DeploymentNotFound until someone " - "creates the fable-5 deployment, so this cell stays loud in CI. " - "Remove this fail_reason once the deployment exists." - ), ), ModelEntry( alias="azure-claude-opus-4-8", @@ -225,12 +219,6 @@ AZURE_AI_MODELS: Tuple[ModelEntry, ...] = ( mode="adaptive", required_env=_AZURE_FOUNDRY_REQ, caps=_CAPS_XHIGH_MAX, - fail_reason=( - "claude-opus-4-8 has no deployment on the CI Microsoft Foundry " - "resource yet; Foundry returns DeploymentNotFound until someone " - "creates the opus-4-8 deployment, so this cell stays loud in CI. " - "Remove this fail_reason once the deployment exists." - ), ), ModelEntry( alias="azure-claude-opus-4-7", diff --git a/tests/ocr_tests/test_ocr_azure_ai.py b/tests/ocr_tests/test_ocr_azure_ai.py index 172682175c5..acb44958fd9 100644 --- a/tests/ocr_tests/test_ocr_azure_ai.py +++ b/tests/ocr_tests/test_ocr_azure_ai.py @@ -23,7 +23,7 @@ class TestAzureAIOCR(BaseOCRTest): Return the base OCR call args for Azure AI. """ return { - "model": "azure_ai/mistral-document-ai-2505", + "model": "azure_ai/mistral-document-ai-2512", "api_key": os.getenv("AZURE_API_KEY"), "api_base": os.getenv("AZURE_API_BASE"), } diff --git a/tests/test_litellm/caching/test_dual_cache.py b/tests/test_litellm/caching/test_dual_cache.py index f4f88def78d..47be139eb5e 100644 --- a/tests/test_litellm/caching/test_dual_cache.py +++ b/tests/test_litellm/caching/test_dual_cache.py @@ -88,6 +88,60 @@ async def test_dual_cache_async_set_cache_injects_default_in_memory_ttl(): assert expiry <= after + 60 +@pytest.mark.asyncio +async def test_dual_cache_redis_backfill_injects_default_in_memory_ttl(): + """ + A Redis-hit backfill into the in-memory tier must honor + default_in_memory_ttl the same way the write paths do. Without it, the + backfilled entry falls to InMemoryCache's own default_ttl (600s), so a + replica that primed a management object (e.g. a virtual key's auth blob) + from Redis keeps serving it for 10 minutes after the object was updated + and invalidated, instead of re-reading within the configured TTL. + """ + in_memory_cache = InMemoryCache(default_ttl=600) + redis_cache = MagicMock() + redis_cache.async_get_cache = AsyncMock(return_value="redis_value") + dual_cache = DualCache( + in_memory_cache=in_memory_cache, + redis_cache=redis_cache, + default_in_memory_ttl=60, + ) + + before = time.time() + result = await dual_cache.async_get_cache(key="backfill_key") + after = time.time() + + assert result == "redis_value" + expiry = in_memory_cache.ttl_dict["backfill_key"] + assert expiry >= before + 60 + assert expiry <= after + 60 + + +@pytest.mark.asyncio +async def test_dual_cache_batch_redis_backfill_injects_default_in_memory_ttl(): + """async_batch_get_cache's Redis-to-memory backfill must honor + default_in_memory_ttl, same as the single-key path.""" + in_memory_cache = InMemoryCache(default_ttl=600) + mock_redis = MagicMock(spec=RedisCache) + mock_redis.async_batch_get_cache = AsyncMock( + return_value={"batch_backfill_key": "redis_value"} + ) + dual_cache = DualCache( + in_memory_cache=in_memory_cache, + redis_cache=mock_redis, + default_in_memory_ttl=60, + ) + + before = time.time() + result = await dual_cache.async_batch_get_cache(keys=["batch_backfill_key"]) + after = time.time() + + assert result == ["redis_value"] + expiry = in_memory_cache.ttl_dict["batch_backfill_key"] + assert expiry >= before + 60 + assert expiry <= after + 60 + + @pytest.mark.asyncio async def test_dual_cache_async_set_cache_respects_explicit_ttl(): """ diff --git a/tests/test_litellm/integrations/test_langfuse_otel.py b/tests/test_litellm/integrations/test_langfuse_otel.py index 3aade7514e4..2f2675ca790 100644 --- a/tests/test_litellm/integrations/test_langfuse_otel.py +++ b/tests/test_litellm/integrations/test_langfuse_otel.py @@ -412,6 +412,171 @@ class TestLangfuseOtelIntegration: # Endpoint assertion removed as side effect is gone +class TestLangfuseOtelKeyDynamicConfig: + """Key/team-scoped Langfuse credentials must define the full export target + (OTLP endpoint + auth), not just auth headers on the init-time exporter.""" + + CLEAN_ENV_VARS = [ + "LANGFUSE_PUBLIC_KEY", + "LANGFUSE_SECRET_KEY", + "LANGFUSE_HOST", + "LANGFUSE_OTEL_HOST", + "OTEL_EXPORTER", + "OTEL_EXPORTER_OTLP_PROTOCOL", + "OTEL_ENDPOINT", + "OTEL_EXPORTER_OTLP_ENDPOINT", + "OTEL_HEADERS", + "OTEL_EXPORTER_OTLP_HEADERS", + ] + + def _clean_env(self): + cleaned = {k: v for k, v in os.environ.items() if k not in self.CLEAN_ENV_VARS} + return patch.dict(os.environ, cleaned, clear=True) + + def _dynamic_params(self, **overrides): + from litellm.types.utils import StandardCallbackDynamicParams + + params = { + "langfuse_public_key": "key_public", + "langfuse_secret_key": "key_secret", + "langfuse_host": "https://langfuse.example.com", + } + params.update(overrides) + return StandardCallbackDynamicParams(**{k: v for k, v in params.items() if v is not None}) + + def test_construct_dynamic_otel_config_with_key_credentials(self): + with self._clean_env(): + logger = LangfuseOtelLogger() + config = logger.construct_dynamic_otel_config(self._dynamic_params()) + + assert config is not None + assert config.exporter == "otlp_http" + assert config.endpoint == "https://langfuse.example.com/api/public/otel" + + import base64 + + expected_auth = base64.b64encode(b"key_public:key_secret").decode() + assert config.headers == f"Authorization=Basic {expected_auth}" + + def test_construct_dynamic_otel_config_host_without_protocol(self): + with self._clean_env(): + logger = LangfuseOtelLogger() + config = logger.construct_dynamic_otel_config(self._dynamic_params(langfuse_host="langfuse.example.com")) + + assert config is not None + assert config.endpoint == "https://langfuse.example.com/api/public/otel" + + def test_construct_dynamic_otel_config_defaults_to_us_cloud(self): + with self._clean_env(): + logger = LangfuseOtelLogger() + config = logger.construct_dynamic_otel_config(self._dynamic_params(langfuse_host=None)) + + assert config is not None + assert config.endpoint == "https://us.cloud.langfuse.com/api/public/otel" + + def test_construct_dynamic_otel_config_falls_back_to_env_host(self): + with self._clean_env(): + with patch.dict(os.environ, {"LANGFUSE_HOST": "https://env-host.example.com"}): + logger = LangfuseOtelLogger() + config = logger.construct_dynamic_otel_config(self._dynamic_params(langfuse_host=None)) + + assert config is not None + assert config.endpoint == "https://env-host.example.com/api/public/otel" + + def test_construct_dynamic_otel_config_requires_both_keys(self): + with self._clean_env(): + logger = LangfuseOtelLogger() + + assert logger.construct_dynamic_otel_config(self._dynamic_params(langfuse_secret_key=None)) is None + assert logger.construct_dynamic_otel_config(self._dynamic_params(langfuse_public_key=None)) is None + + def test_key_dynamic_params_create_otlp_exporter_without_global_env(self): + """Without global LANGFUSE_* env vars, a request carrying key-scoped Langfuse + credentials must get a tracer exporting via OTLP HTTP to that key's host.""" + from opentelemetry.exporter.otlp.proto.http.trace_exporter import ( + OTLPSpanExporter, + ) + from opentelemetry.sdk.trace.export import BatchSpanProcessor + + with self._clean_env(): + logger = LangfuseOtelLogger() + assert logger.OTEL_EXPORTER == "console" + + tracer = logger.get_tracer_to_use_for_request( + {"standard_callback_dynamic_params": self._dynamic_params()} + ) + + assert tracer is not logger.tracer + assert len(logger._tracer_provider_cache) == 1 + + provider = next(iter(logger._tracer_provider_cache.values())) + span_processors = provider._active_span_processor._span_processors + assert len(span_processors) == 1 + assert isinstance(span_processors[0], BatchSpanProcessor) + + exporter = span_processors[0].span_exporter + assert isinstance(exporter, OTLPSpanExporter) + assert exporter._endpoint == "https://langfuse.example.com/api/public/otel/v1/traces" + + import base64 + + expected_auth = base64.b64encode(b"key_public:key_secret").decode() + assert exporter._headers == {"Authorization": f"Basic {expected_auth}"} + + def test_key_dynamic_params_reuse_cached_provider(self): + with self._clean_env(): + logger = LangfuseOtelLogger() + kwargs = {"standard_callback_dynamic_params": self._dynamic_params()} + logger.get_tracer_to_use_for_request(kwargs) + logger.get_tracer_to_use_for_request(kwargs) + + assert len(logger._tracer_provider_cache) == 1 + + def test_no_dynamic_params_keeps_default_tracer(self): + with self._clean_env(): + logger = LangfuseOtelLogger() + tracer = logger.get_tracer_to_use_for_request({}) + + assert tracer is logger.tracer + assert logger._tracer_provider_cache == {} + + def test_key_credentials_never_passed_to_debug_logger(self): + """The span-processor debug logs must receive a redacted header value, so the + key-scoped Langfuse secret never enters a log record regardless of downstream + handler configuration, while the exporter still gets the real header.""" + import base64 + + from opentelemetry.exporter.otlp.proto.http.trace_exporter import ( + OTLPSpanExporter, + ) + + from litellm.integrations import opentelemetry as otel_module + + secret = base64.b64encode(b"key_public:key_secret").decode() + + recorded_arguments = [] + + def _spy(message, *args, **kwargs): + recorded_arguments.append(" ".join(str(part) for part in (message, *args))) + + with self._clean_env(): + logger = LangfuseOtelLogger() + with patch.object(otel_module.verbose_logger, "debug", side_effect=_spy): + logger.get_tracer_to_use_for_request( + {"standard_callback_dynamic_params": self._dynamic_params()} + ) + + logged = "\n".join(recorded_arguments) + assert "initializing span processor" in logged + assert secret not in logged + assert f"Basic {secret}" not in logged + + provider = next(iter(logger._tracer_provider_cache.values())) + exporter = provider._active_span_processor._span_processors[0].span_exporter + assert isinstance(exporter, OTLPSpanExporter) + assert exporter._headers == {"Authorization": f"Basic {secret}"} + + class TestLangfuseOtelResponsesAPI: """Test suite for Langfuse OTEL integration with ResponsesAPI""" diff --git a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py index f99cfb953ba..6875894c1bf 100644 --- a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py +++ b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py @@ -2247,6 +2247,53 @@ def test_get_error_information_prefers_message_attribute_over_str(): assert result["error_class"] == "ProxyExceptionLike" +def test_get_error_information_budget_exceeded_structured_fields(): + """ + Regression for LIT-4458: a budget-rejected request's failure + StandardLoggingPayload must identify WHICH budget blocked the call + as structured fields, not only inside the free-text error_str + ("ExceededBudget: User=... over budget. Spend=..., Budget=..."). + + Asserts get_error_information copies entity_type / entity_id / + max_budget / current_cost off BudgetExceededError into + error_budget_entity_type / error_budget_entity_id / + error_budget_limit / error_budget_spend, and leaves all four None + for non-budget exceptions. + """ + from litellm.exceptions import BudgetExceededError + from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup + + exc = BudgetExceededError( + current_cost=3.4e-05, + max_budget=1e-06, + message="ExceededBudget: User=repro-user over budget. Spend=3.4e-05, Budget=1e-06", + entity_type="user", + entity_id="repro-user", + ) + + result = StandardLoggingPayloadSetup.get_error_information(exc) + assert result["error_budget_entity_type"] == "user" + assert result["error_budget_entity_id"] == "repro-user" + assert result["error_budget_limit"] == 1e-06 + assert result["error_budget_spend"] == 3.4e-05 + assert result["error_code"] == "429" + assert result["error_class"] == "BudgetExceededError" + assert result["error_rate_limit_type"] == "budget" + + legacy_exc = BudgetExceededError(current_cost=2.0, max_budget=1.0) + legacy_result = StandardLoggingPayloadSetup.get_error_information(legacy_exc) + assert legacy_result["error_budget_entity_type"] is None + assert legacy_result["error_budget_entity_id"] is None + assert legacy_result["error_budget_limit"] == 1.0 + assert legacy_result["error_budget_spend"] == 2.0 + + non_budget_result = StandardLoggingPayloadSetup.get_error_information(ValueError("boom")) + assert non_budget_result["error_budget_entity_type"] is None + assert non_budget_result["error_budget_entity_id"] is None + assert non_budget_result["error_budget_limit"] is None + assert non_budget_result["error_budget_spend"] is None + + def test_get_error_information_preserves_explicit_empty_message(): """ An exception that deliberately sets `.message = ""` must surface diff --git a/tests/test_litellm/litellm_core_utils/test_streaming_chunk_builder_utils.py b/tests/test_litellm/litellm_core_utils/test_streaming_chunk_builder_utils.py index ba95c90798e..be8c5a05601 100644 --- a/tests/test_litellm/litellm_core_utils/test_streaming_chunk_builder_utils.py +++ b/tests/test_litellm/litellm_core_utils/test_streaming_chunk_builder_utils.py @@ -956,3 +956,39 @@ def test_stream_chunk_builder_propagates_vertex_ai_metadata_from_dict_chunks(): assert response.model_dump()["vertex_ai_grounding_metadata"] == [ {"webSearchQueries": ["test query"]} ] + + +def test_cost_field_in_usage_chunks(): + chunk1_usage = Usage(completion_tokens=1, prompt_tokens=10, total_tokens=11) + chunk1 = ModelResponseStream( + id="chatcmpl-1", + created=1745513206, + model="openrouter/claude", + choices=[ + StreamingChoices(finish_reason=None, index=0, delta=Delta(content="Hi")) + ], + usage=chunk1_usage, + ) + + chunk2_usage = Usage( + completion_tokens=5, prompt_tokens=10, total_tokens=15, cost=0.00025 + ) + chunk2 = ModelResponseStream( + id="chatcmpl-1", + created=1745513207, + model="openrouter/claude", + choices=[ + StreamingChoices(finish_reason="stop", index=0, delta=Delta(content="")) + ], + usage=chunk2_usage, + ) + + processor = ChunkProcessor(chunks=[chunk1, chunk2]) + usage = processor.calculate_usage( + chunks=[chunk1, chunk2], model="openrouter/claude", completion_output="Hi" + ) + + assert hasattr(usage, "cost") + assert usage.cost == 0.00025 + assert usage.prompt_tokens == 10 + assert usage.completion_tokens == 5 diff --git a/tests/test_litellm/litellm_core_utils/test_streaming_handler.py b/tests/test_litellm/litellm_core_utils/test_streaming_handler.py index e430ce3b084..514714136fd 100644 --- a/tests/test_litellm/litellm_core_utils/test_streaming_handler.py +++ b/tests/test_litellm/litellm_core_utils/test_streaming_handler.py @@ -1393,6 +1393,190 @@ def test_has_any_special_delta_attributes( assert result is False +def test_calculate_total_usage_with_cost(): + from litellm.litellm_core_utils.streaming_handler import calculate_total_usage + + chunk1_usage = Usage(completion_tokens=1, prompt_tokens=10, total_tokens=11) + chunk1 = ModelResponseStream( + id="test-1", + created=1745513206, + model="openrouter/test", + choices=[ + StreamingChoices(finish_reason=None, index=0, delta=Delta(content="Hi")) + ], + usage=chunk1_usage, + ) + + chunk2_usage = Usage( + completion_tokens=5, prompt_tokens=10, total_tokens=15, cost=0.00025 + ) + chunk2 = ModelResponseStream( + id="test-1", + created=1745513207, + model="openrouter/test", + choices=[ + StreamingChoices(finish_reason="stop", index=0, delta=Delta(content="")) + ], + usage=chunk2_usage, + ) + + usage = calculate_total_usage([chunk1, chunk2]) + + assert hasattr(usage, "cost") + assert usage.cost == 0.00025 + assert usage.prompt_tokens == 10 + assert usage.completion_tokens == 5 + + +def test_calculate_total_usage_with_dict_usage_cost(): + """Regression: dict-shaped `usage` with a `cost` key must still surface + provider cost even though `hasattr` on a dict does not consult its keys.""" + from litellm.litellm_core_utils.streaming_handler import calculate_total_usage + + chunk = { + "usage": { + "prompt_tokens": 10, + "completion_tokens": 5, + "total_tokens": 15, + "cost": 0.00025, + } + } + + usage = calculate_total_usage([chunk]) + + assert usage.prompt_tokens == 10 + assert usage.completion_tokens == 5 + assert getattr(usage, "cost", None) == 0.00025 + + +@pytest.mark.asyncio +async def test_openrouter_streaming_cost_after_finish_reason(logging_obj: Logging): + from litellm.utils import ModelResponseListIterator + + chunk1 = ModelResponseStream( + id="chatcmpl-or", + created=1742056047, + model="openrouter/claude", + choices=[ + StreamingChoices( + finish_reason=None, index=0, delta=Delta(content="Hi", role="assistant") + ) + ], + usage=None, + ) + chunk2 = ModelResponseStream( + id="chatcmpl-or", + created=1742056048, + model="openrouter/claude", + choices=[ + StreamingChoices(finish_reason="stop", index=0, delta=Delta(content="")) + ], + usage=None, + ) + chunk3_usage = Usage( + completion_tokens=5, prompt_tokens=10, total_tokens=15, cost=0.00025 + ) + chunk3 = ModelResponseStream( + id="chatcmpl-or", + created=1742056049, + model="openrouter/claude", + choices=[ + StreamingChoices(finish_reason=None, index=0, delta=Delta(content="")) + ], + usage=chunk3_usage, + ) + + completion_stream = ModelResponseListIterator( + model_responses=[chunk1, chunk2, chunk3] + ) + response = CustomStreamWrapper( + completion_stream=completion_stream, + model="openrouter/claude", + custom_llm_provider="openrouter", + logging_obj=logging_obj, + stream_options={"include_usage": True}, + ) + + collected_chunks = [] + async for chunk in response: + collected_chunks.append(chunk) + + usage_chunks = [c for c in collected_chunks if hasattr(c, "usage") and c.usage] + assert len(usage_chunks) > 0 + assert hasattr(usage_chunks[-1].usage, "cost") + assert usage_chunks[-1].usage.cost == 0.00025 + + +def test_openrouter_streaming_cost_propagates_to_hidden_params(): + """ + Verify that provider-reported cost from usage.cost flows into + _hidden_params["additional_headers"]["llm_provider-x-litellm-response-cost"] + on the complete streaming response, so litellm's cost calculator uses it. + """ + import litellm + + chunk1 = ModelResponseStream( + id="chatcmpl-or", + created=1742056047, + model="openrouter/claude", + choices=[ + StreamingChoices( + finish_reason=None, index=0, delta=Delta(content="Hi", role="assistant") + ) + ], + usage=None, + ) + chunk2 = ModelResponseStream( + id="chatcmpl-or", + created=1742056048, + model="openrouter/claude", + choices=[ + StreamingChoices(finish_reason="stop", index=0, delta=Delta(content="")) + ], + usage=None, + ) + chunk3 = ModelResponseStream( + id="chatcmpl-or", + created=1742056049, + model="openrouter/claude", + choices=[ + StreamingChoices(finish_reason=None, index=0, delta=Delta(content="")) + ], + usage=Usage( + completion_tokens=5, prompt_tokens=10, total_tokens=15, cost=0.00025 + ), + ) + + # Build the complete response as stream_chunk_builder does + complete_response = litellm.stream_chunk_builder( + chunks=[chunk1, chunk2, chunk3], + messages=[{"role": "user", "content": "test"}], + ) + + assert complete_response is not None + assert hasattr(complete_response.usage, "cost") + assert complete_response.usage.cost == 0.00025 + + # Use the real propagation method from CustomStreamWrapper + CustomStreamWrapper._propagate_usage_cost_to_hidden_params(complete_response) + + assert "additional_headers" in complete_response._hidden_params + assert ( + complete_response._hidden_params["additional_headers"][ + "llm_provider-x-litellm-response-cost" + ] + == 0.00025 + ) + + # Verify the cost calculator would pick this up + from litellm.cost_calculator import get_response_cost_from_hidden_params + + provider_cost = get_response_cost_from_hidden_params( + complete_response._hidden_params + ) + assert provider_cost == 0.00025 + + def test_handle_special_delta_attributes( initialized_custom_stream_wrapper: CustomStreamWrapper, ): @@ -3059,3 +3243,115 @@ async def test_stream_chunk_builder_raise_and_usage_recovery_failure_does_not_cr chunks = [c async for c in response] assert len(chunks) > 0 + + +class TransportErrorAfterChunksIterator: + """Yields the given chunks, then raises the given exception once, then StopAsyncIteration.""" + + def __init__(self, model_responses, exception): + self.model_responses = model_responses + self.exception = exception + self.index = 0 + self.raised = False + + def __aiter__(self): + return self + + async def __anext__(self): + if self.index < len(self.model_responses): + chunk = self.model_responses[self.index] + self.index += 1 + return chunk + if not self.raised: + self.raised = True + raise self.exception + raise StopAsyncIteration + + +def _reset_test_chunk(content: Optional[str] = None, finish_reason: Optional[str] = None) -> ModelResponseStream: + return ModelResponseStream( + id="chatcmpl-reset-test", + created=1783458104, + model="stub-model", + object="chat.completion.chunk", + choices=[ + StreamingChoices( + index=0, + delta=Delta(content=content), + finish_reason=finish_reason, + ) + ], + ) + + +@pytest.mark.asyncio +async def test_transport_read_error_after_finish_reason_ends_stream_gracefully( + logging_obj: Logging, +): + """A trailing connection reset after the provider's finish chunk must not fail the stream.""" + import httpx + + completion_stream = TransportErrorAfterChunksIterator( + model_responses=[ + _reset_test_chunk(content="Hello"), + _reset_test_chunk(finish_reason="stop"), + ], + exception=httpx.ReadError("Response payload is not completed"), + ) + response = CustomStreamWrapper( + completion_stream=completion_stream, + model="hosted_vllm/stub-model", + custom_llm_provider="hosted_vllm", + logging_obj=logging_obj, + ) + + chunks = [chunk async for chunk in response] + + finish_reasons = [ + chunk.choices[0].finish_reason + for chunk in chunks + if chunk.choices and chunk.choices[0].finish_reason + ] + contents = [ + chunk.choices[0].delta.content + for chunk in chunks + if chunk.choices and chunk.choices[0].delta and chunk.choices[0].delta.content + ] + assert finish_reasons == ["stop"] + assert contents == ["Hello"] + + +@pytest.mark.asyncio +async def test_transport_read_error_before_finish_reason_raises(logging_obj: Logging): + """A connection reset before any finish chunk must surface, never end as a clean stop. + + Regression test for silent empty/truncated HTTP 200 streams: the aiohttp + transport used to swallow mid-stream connection resets, so the wrapper saw a + clean end-of-stream and fabricated finish_reason "stop". + """ + import httpx + + from litellm.exceptions import MidStreamFallbackError + + completion_stream = TransportErrorAfterChunksIterator( + model_responses=[_reset_test_chunk(content="Hel")], + exception=httpx.ReadError("Response payload is not completed"), + ) + response = CustomStreamWrapper( + completion_stream=completion_stream, + model="hosted_vllm/stub-model", + custom_llm_provider="hosted_vllm", + logging_obj=logging_obj, + ) + + received = [] + with pytest.raises(MidStreamFallbackError): + async for chunk in response: + received.append(chunk) + + fabricated_finish_reasons = [ + chunk.choices[0].finish_reason + for chunk in received + if chunk.choices and chunk.choices[0].finish_reason + ] + assert fabricated_finish_reasons == [] diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py index f9c55db72b5..dfe7e0c3a51 100644 --- a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py +++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py @@ -256,7 +256,14 @@ def test_translate_streaming_openai_chunk_to_anthropic_thinking_signature_block( } -def test_translate_streaming_openai_chunk_to_anthropic_raises_when_thinking_and_signature_content_block(): +def test_translate_streaming_openai_chunk_to_anthropic_content_block_thinking_and_signature(): + """The content-block classifier must treat a chunk carrying both ``thinking`` + and ``signature`` as a ``thinking`` block instead of raising. + + Such a chunk is the terminal signature event of an already-open thinking block, + so classifying it as ``thinking`` keeps the stream on the same block rather than + 500'ing. Before the fix this raised ``ValueError``. + """ choices = [ StreamingChoices( finish_reason=None, @@ -289,10 +296,14 @@ def test_translate_streaming_openai_chunk_to_anthropic_raises_when_thinking_and_ ) ] - with pytest.raises(ValueError): - LiteLLMAnthropicMessagesAdapter()._translate_streaming_openai_chunk_to_anthropic_content_block( - choices=choices - ) + ( + block_type, + content_block_start, + ) = LiteLLMAnthropicMessagesAdapter()._translate_streaming_openai_chunk_to_anthropic_content_block( + choices=choices + ) + + assert block_type == "thinking" def test_translate_anthropic_messages_to_openai_thinking_blocks(): @@ -738,7 +749,17 @@ def test_translate_streaming_openai_chunk_to_anthropic_with_thinking(): assert content_block_delta["signature"] == "sigsig" -def test_translate_streaming_openai_chunk_to_anthropic_raises_when_thinking_and_signature(): +def test_translate_streaming_openai_chunk_to_anthropic_emits_signature_when_thinking_and_signature(): + """A single streaming chunk carrying both ``thinking`` and ``signature`` must + translate to a ``signature_delta``, not crash. + + litellm's Anthropic streaming handler emits the ``signature_delta`` event as an + OpenAI chunk whose ``thinking_blocks`` entry re-states the full accumulated + thinking text alongside the signature (see anthropic/chat/handler.py). That text + was already streamed as ``thinking_delta`` chunks, so the signature must win and + the duplicate thinking must not be re-emitted. Before the fix this raised + ``ValueError`` and 500'd the whole stream, breaking Claude Code through the proxy. + """ choices = [ StreamingChoices( finish_reason=None, @@ -771,10 +792,25 @@ def test_translate_streaming_openai_chunk_to_anthropic_raises_when_thinking_and_ ) ] - with pytest.raises(ValueError): - LiteLLMAnthropicMessagesAdapter()._translate_streaming_openai_chunk_to_anthropic( - choices=choices - ) + adapter = LiteLLMAnthropicMessagesAdapter() + + ( + type_of_content, + content_block_delta, + ) = adapter._translate_streaming_openai_chunk_to_anthropic(choices=choices) + + assert type_of_content == "signature_delta" + assert content_block_delta["type"] == "signature_delta" + assert content_block_delta["signature"] == "sigsig" + + ( + block_type, + content_block_start, + ) = adapter._translate_streaming_openai_chunk_to_anthropic_content_block( + choices=choices + ) + + assert block_type == "thinking" def test_translate_anthropic_messages_to_openai_user_message_with_base64_image(): diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_combined_chunk.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_combined_chunk.py index d67de0dcaf8..6973340101e 100644 --- a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_combined_chunk.py +++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_combined_chunk.py @@ -12,6 +12,7 @@ content survives. import asyncio import json from types import SimpleNamespace +from typing import AsyncIterator from litellm.llms.anthropic.experimental_pass_through.adapters.streaming_iterator import ( AnthropicStreamWrapper, @@ -202,3 +203,89 @@ def test_split_clears_reasoning_and_thinking_on_finish_chunk(): assert content_chunk.choices[0].delta.thinking_blocks == [{"type": "thinking"}] assert finish_chunk.choices[0].delta.reasoning_content is None assert finish_chunk.choices[0].delta.thinking_blocks is None + + +def _thinking_delta_chunk(thinking: str) -> ModelResponseStream: + return ModelResponseStream( + choices=[ + StreamingChoices( + index=0, + delta=Delta( + reasoning_content=thinking, + thinking_blocks=[{"type": "thinking", "thinking": thinking, "signature": None}], + provider_specific_fields={ + "thinking_blocks": [{"type": "thinking", "thinking": thinking, "signature": None}] + }, + ), + finish_reason=None, + ) + ], + ) + + +def _signature_chunk(recap: str, signature: str) -> ModelResponseStream: + return ModelResponseStream( + choices=[ + StreamingChoices( + index=0, + delta=Delta( + reasoning_content=recap, + thinking_blocks=[{"type": "thinking", "thinking": recap, "signature": signature}], + provider_specific_fields={ + "thinking_blocks": [{"type": "thinking", "thinking": recap, "signature": signature}] + }, + ), + finish_reason=None, + ) + ], + ) + + +def test_thinking_then_signature_chunk_does_not_crash_stream(): + """Regression for the /v1/messages streaming crash reported on autoroute. + + Anthropic streams extended thinking as incremental ``thinking_delta`` chunks, then a + closing chunk that recaps the full accumulated thinking AND carries the signature. The + adapter used to raise ``ValueError`` on that closing chunk, killing the whole stream. It + must instead emit a single ``signature_delta`` for the recap chunk and never re-emit the + recap thinking, so the incremental thinking text is not duplicated. + """ + chunks = [ + _thinking_delta_chunk("First, "), + _thinking_delta_chunk("reason."), + _signature_chunk("First, reason.", "sig-abc"), + ModelResponseStream( + choices=[StreamingChoices(index=0, delta=Delta(content="Done"), finish_reason=None)], + ), + ModelResponseStream( + choices=[StreamingChoices(index=0, delta=Delta(), finish_reason="stop")], + usage=Usage(prompt_tokens=5, completion_tokens=3, total_tokens=8), + ), + ] + + async def _aiter() -> "AsyncIterator[ModelResponseStream]": + for chunk in chunks: + yield chunk + + wrapper = AnthropicStreamWrapper(completion_stream=_aiter(), model="claude-haiku-4-5") + sse = _collect_async(wrapper) + + signature_deltas = [ + json.loads(line[len("data: ") :]) + for block in sse.split("\n\n") + for line in block.splitlines() + if line.startswith("data: ") and '"signature_delta"' in line + ] + assert len(signature_deltas) == 1 + assert signature_deltas[0]["delta"]["signature"] == "sig-abc" + + thinking_text = "".join( + json.loads(line[len("data: ") :])["delta"]["thinking"] + for block in sse.split("\n\n") + for line in block.splitlines() + if line.startswith("data: ") and '"thinking_delta"' in line + ) + assert thinking_text == "First, reason." + + assert "message_stop" in sse + assert "Done" in sse diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_first_delta.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_first_delta.py index 6ed79763753..19ec1a04b45 100644 --- a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_first_delta.py +++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_first_delta.py @@ -438,6 +438,58 @@ async def test_empty_reasoning_delta_mid_thinking_block_is_suppressed_async(): _assert_empty_reasoning_delta_suppressed(await _drain_async(wrapper)) +def _full_snapshot_signature_chunks() -> List[MagicMock]: + """Mirror litellm's real Anthropic streaming: incremental ``thinking_delta`` + chunks (empty signature), then a terminal chunk whose ``thinking_blocks`` entry + re-states the *full accumulated thinking text* together with the signature + (anthropic/chat/handler.py builds the signature_delta event this way), then the + answer text. + """ + return [ + _thinking_chunk("Let me "), + _thinking_chunk("think about it."), + _thinking_chunk("Let me think about it.", signature="sig-abc"), + _make_chunk(Delta(content="42")), + _make_chunk(Delta(content=None), finish_reason="stop"), + ] + + +def _assert_full_snapshot_signature_handled(events: List[dict]) -> None: + _assert_deltas_match_their_block_type(events) + # The full-text snapshot on the signature chunk must NOT be re-emitted as an + # extra thinking_delta (it was already streamed incrementally) - otherwise the + # client renders the reasoning twice. + assert _thinking_deltas(events) == ["Let me ", "think about it."] + assert "".join(_thinking_deltas(events)) == "Let me think about it." + assert _signature_deltas(events) == ["sig-abc"] + assert _text_deltas(events) == ["42"] + + +def test_full_thinking_snapshot_with_signature_emits_signature_only_sync(): + """Regression: a terminal thinking chunk carrying both the full thinking text + and the signature used to raise ``ValueError`` (500) mid-stream, breaking every + Claude Code request routed through the proxy with an extended-thinking model. It + must instead emit a single ``signature_delta`` without duplicating the thinking. + """ + wrapper = AnthropicStreamWrapper( + completion_stream=iter(_full_snapshot_signature_chunks()), + model="claude-x", + ) + _assert_full_snapshot_signature_handled(_drain_sync(wrapper)) + + +@pytest.mark.asyncio +async def test_full_thinking_snapshot_with_signature_emits_signature_only_async(): + """Async twin - the proxy serves the async iterator, so the crash must be gone + on that path too. + """ + wrapper = AnthropicStreamWrapper( + completion_stream=_AsyncStream(_full_snapshot_signature_chunks()), + model="claude-x", + ) + _assert_full_snapshot_signature_handled(await _drain_async(wrapper)) + + def test_empty_content_chunk_mid_text_block_is_suppressed_sync(): """An empty-content chunk arriving mid-text-block (no transition) used to emit a pointless ``text_delta {"text": ""}``; it must be dropped without diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_anthropic_messages_effort.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_anthropic_messages_effort.py index 5254808e315..e9d4d625421 100644 --- a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_anthropic_messages_effort.py +++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_anthropic_messages_effort.py @@ -32,6 +32,37 @@ def _transform(model, params, litellm_params=None): ) +def test_adaptive_thinking_only_translated_to_legacy_for_haiku_4_5(): + """The minimal autoroute repro: Claude Code sends bare ``thinking={type: adaptive}`` + (no ``output_config``) and the complexity router picks Haiku 4.5, which does not + support adaptive thinking. Anthropic 400s with "adaptive thinking is not supported on + this model" unless the flag is dropped, so it must be translated to the legacy extended + thinking the model does support rather than forwarded raw.""" + result = _transform("claude-haiku-4-5", {"max_tokens": 8192, "thinking": {"type": "adaptive"}}) + + assert result["thinking"] == { + "type": "enabled", + "budget_tokens": DEFAULT_REASONING_EFFORT_MEDIUM_THINKING_BUDGET, + } + assert "output_config" not in result + + +def test_adaptive_thinking_only_dropped_for_non_reasoning_model(): + """Bare adaptive thinking on a model with no reasoning support at all is silently + dropped so the request still succeeds instead of being rejected.""" + result = _transform("claude-3-5-haiku-latest", {"max_tokens": 8192, "thinking": {"type": "adaptive"}}) + + assert "thinking" not in result + + +def test_adaptive_thinking_only_preserved_for_4_6(): + """A 4.6+ model natively supports adaptive thinking, so a bare adaptive flag must not + be rewritten even without output_config.""" + result = _transform("claude-sonnet-4-6", {"max_tokens": 8192, "thinking": {"type": "adaptive"}}) + + assert result["thinking"] == {"type": "adaptive"} + + def test_effort_translated_to_legacy_thinking_for_haiku_4_5(): """Core regression: Claude Code sends adaptive thinking + effort to Haiku 4.5 (thinking-capable, pre-4.6). Effort must be translated to legacy extended diff --git a/tests/test_litellm/llms/custom_httpx/test_aiohttp_transport.py b/tests/test_litellm/llms/custom_httpx/test_aiohttp_transport.py index 0c5e386c438..d1bc356662f 100644 --- a/tests/test_litellm/llms/custom_httpx/test_aiohttp_transport.py +++ b/tests/test_litellm/llms/custom_httpx/test_aiohttp_transport.py @@ -81,7 +81,7 @@ class MockContent: def __init__(self, chunks=None, exception_to_raise=None, exception_at_chunk=None): self.chunks = chunks or [b"chunk1", b"chunk2", b"chunk3"] self.exception_to_raise = exception_to_raise - self.exception_at_chunk = exception_at_chunk or (len(self.chunks) - 1) + self.exception_at_chunk = exception_at_chunk if exception_at_chunk is not None else (len(self.chunks) - 1) self.chunk_index = 0 async def iter_chunked(self, chunk_size): @@ -107,15 +107,11 @@ async def test_aiohttp_response_stream_normal_flow(): @pytest.mark.asyncio -async def test_transfer_encoding_error_no_httpx_read_error(): - """Test that TransferEncodingError doesn't get converted to httpx.ReadError""" - - # Create a TransferEncodingError wrapped in ClientPayloadError (like in real scenarios) +async def test_client_payload_error_mid_stream_raises_read_error(): + """A connection reset mid-body must surface as httpx.ReadError, not truncate silently""" transfer_error = aiohttp.http_exceptions.TransferEncodingError( message="400, message: Not enough data for satisfy transfer length header." ) - - # Wrap it in ClientPayloadError as aiohttp does client_payload_error = aiohttp.ClientPayloadError( "Response payload is not completed" ) @@ -124,47 +120,100 @@ async def test_transfer_encoding_error_no_httpx_read_error(): mock_response = MockAiohttpResponse( content_chunks=[b"chunk1", b"chunk2", b"chunk3"], exception_to_raise=client_payload_error, - exception_at_chunk=1, # Error occurs at chunk 1 + exception_at_chunk=1, ) stream = AiohttpResponseStream(mock_response) # type: ignore received_chunks = [] - # This should NOT raise httpx.ReadError or any other exception - # It should handle the error gracefully and just return what was received - async for chunk in stream: - received_chunks.append(chunk) - print(f"received_chunks: {received_chunks}") + with pytest.raises(httpx.ReadError): + async for chunk in stream: + received_chunks.append(chunk) - # Should have received the first chunk before the error assert received_chunks == [b"chunk1"] - assert len(received_chunks) == 1 + assert mock_response.closed is True @pytest.mark.asyncio -async def test_client_payload_error_graceful_handling(): - """Test that ClientPayloadError is handled gracefully without stacktrace""" - # Create a ClientPayloadError directly +async def test_client_payload_error_before_first_chunk_raises_read_error(): + """A connection reset before any body byte must surface, not yield an empty 200 body""" client_error = aiohttp.client_exceptions.ClientPayloadError( "Response payload is not completed" ) mock_response = MockAiohttpResponse( - content_chunks=[b"data1", b"data2", b"data3"], + content_chunks=[b"data1", b"data2"], exception_to_raise=client_error, - exception_at_chunk=2, # Error occurs at chunk 2 + exception_at_chunk=0, ) stream = AiohttpResponseStream(mock_response) # type: ignore received_chunks = [] - # This should handle the error gracefully without raising - async for chunk in stream: - received_chunks.append(chunk) + with pytest.raises(httpx.ReadError): + async for chunk in stream: + received_chunks.append(chunk) - # Should have received chunks before the error - assert received_chunks == [b"data1", b"data2"] - assert len(received_chunks) == 2 + assert received_chunks == [] + assert mock_response.closed is True + + +@pytest.mark.asyncio +async def test_connection_closed_runtime_error_raises_read_error(): + """aiohttp's bare RuntimeError('Connection closed.') must surface as httpx.ReadError""" + mock_response = MockAiohttpResponse( + content_chunks=[b"data1", b"data2"], + exception_to_raise=RuntimeError("Connection closed."), + exception_at_chunk=1, + ) + + stream = AiohttpResponseStream(mock_response) # type: ignore + received_chunks = [] + + with pytest.raises(httpx.ReadError): + async for chunk in stream: + received_chunks.append(chunk) + + assert received_chunks == [b"data1"] + assert mock_response.closed is True + + +@pytest.mark.asyncio +async def test_unrelated_runtime_error_propagates_unmapped(): + """RuntimeErrors other than 'Connection closed' must propagate untouched""" + mock_response = MockAiohttpResponse( + content_chunks=[b"data1"], + exception_to_raise=RuntimeError("something else broke"), + exception_at_chunk=0, + ) + + stream = AiohttpResponseStream(mock_response) # type: ignore + + with pytest.raises(RuntimeError, match="something else broke"): + async for _ in stream: + pass + + +@pytest.mark.asyncio +async def test_transfer_encoding_error_raises_read_error(): + """A raw TransferEncodingError mid-body must surface as httpx.ReadError""" + mock_response = MockAiohttpResponse( + content_chunks=[b"data1", b"data2"], + exception_to_raise=aiohttp.http_exceptions.TransferEncodingError( + message="Not enough data to satisfy transfer length header." + ), + exception_at_chunk=1, + ) + + stream = AiohttpResponseStream(mock_response) # type: ignore + received_chunks = [] + + with pytest.raises(httpx.ReadError): + async for chunk in stream: + received_chunks.append(chunk) + + assert received_chunks == [b"data1"] + assert mock_response.closed is True @pytest.mark.asyncio diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_partial_update.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_partial_update.py index 41d61c8f508..968fafc0e0e 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_partial_update.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_partial_update.py @@ -203,6 +203,7 @@ async def test_auth_type_switch_clears_stale_flow_scoped_fields(): data_dict = await _run_update_with_existing(data, existing_auth_type="oauth2") for stale_field in ( + "issuer", "authorization_url", "token_url", "registration_url", @@ -217,6 +218,165 @@ async def test_auth_type_switch_clears_stale_flow_scoped_fields(): assert _credentials_cleared(data_dict["credentials"]) +@pytest.mark.asyncio +async def test_url_change_clears_stale_discovered_oauth_fields(): + """Re-pointing the server url at a potentially different upstream must clear the discovered or + trust-on-first-use OAuth issuer and endpoints, so the new upstream re-discovers instead of + anchoring on the previous upstream's issuer (RFC 8414 §3.3 against a stale anchor).""" + mock_prisma = _mock_prisma() + existing = MagicMock() + existing.auth_type = "oauth2" + existing.url = "https://old.example.com/mcp" + existing.credentials = None + mock_prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=existing) + + data = UpdateMCPServerRequest(server_id="my-test-server", url="https://new.example.com/mcp") + await update_mcp_server(mock_prisma, data, "test-user") + data_dict = mock_prisma.db.litellm_mcpservertable.update.call_args[1]["data"] + + assert data_dict["url"] == "https://new.example.com/mcp" + for stale_field in ("issuer", "authorization_url", "token_url", "registration_url"): + assert data_dict[stale_field] is None, f"{stale_field} must be cleared on url change" + + +@pytest.mark.asyncio +async def test_url_change_clears_stale_oauth_fields_even_when_resubmitted_unchanged(): + """The edit form re-sends every field, so a URL change arrives WITH the previous upstream's issuer + and endpoints in the payload. Those resubmitted-unchanged values are stale and must still clear + (otherwise they survive the url change and win in the resolution merge). A genuinely new value the + caller changed in the same submit is kept.""" + mock_prisma = _mock_prisma() + existing = MagicMock() + existing.auth_type = "oauth2" + existing.url = "https://old.example.com/mcp" + existing.credentials = None + existing.issuer = "https://old-idp.example.com" + existing.token_url = "https://old-idp.example.com/token" + existing.authorization_url = "https://old-idp.example.com/authorize" + mock_prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=existing) + + data = UpdateMCPServerRequest( + server_id="my-test-server", + url="https://new.example.com/mcp", + issuer="https://old-idp.example.com", # resubmitted unchanged -> stale, must clear + token_url="https://old-idp.example.com/token", # resubmitted unchanged -> stale, must clear + authorization_url="https://new-idp.example.com/authorize", # genuinely changed -> kept + ) + await update_mcp_server(mock_prisma, data, "test-user") + data_dict = mock_prisma.db.litellm_mcpservertable.update.call_args[1]["data"] + + assert data_dict["issuer"] is None + assert data_dict["token_url"] is None + assert data_dict["authorization_url"] == "https://new-idp.example.com/authorize" + + +@pytest.mark.asyncio +async def test_clearing_pinned_issuer_clears_stale_oauth_endpoints(): + """Clearing a previously pinned issuer must not revive the endpoints resolved under it. Under an + issuer anchor the endpoints come solely from the issuer document and are not persisted, but a row + that was resource-rooted before the pin can still hold stale authorization_url/token_url; clearing + the anchor without clearing those would let them win the resolution merge and be posted to without + fresh discovery (RFC 8414 §3.3 provenance).""" + mock_prisma = _mock_prisma() + existing = MagicMock() + existing.auth_type = "oauth2" + existing.url = "https://same.example.com/mcp" + existing.credentials = None + existing.issuer = "https://pinned-idp.example.com" + existing.token_url = "https://pinned-idp.example.com/token" + existing.authorization_url = "https://pinned-idp.example.com/authorize" + mock_prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=existing) + + data = UpdateMCPServerRequest( + server_id="my-test-server", + issuer="", # admin clears the anchor; url and auth_type unchanged + token_url="https://pinned-idp.example.com/token", + authorization_url="https://pinned-idp.example.com/authorize", + ) + await update_mcp_server(mock_prisma, data, "test-user") + data_dict = mock_prisma.db.litellm_mcpservertable.update.call_args[1]["data"] + + assert data_dict["token_url"] is None + assert data_dict["authorization_url"] is None + + +@pytest.mark.asyncio +async def test_repointing_pinned_issuer_clears_stale_endpoints_keeps_new_issuer(): + """Re-pointing the issuer to a different authorization server invalidates the old issuer's + endpoints while keeping the new issuer the admin submitted.""" + mock_prisma = _mock_prisma() + existing = MagicMock() + existing.auth_type = "oauth2" + existing.url = "https://same.example.com/mcp" + existing.credentials = None + existing.issuer = "https://old-idp.example.com" + existing.token_url = "https://old-idp.example.com/token" + existing.authorization_url = "https://old-idp.example.com/authorize" + mock_prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=existing) + + data = UpdateMCPServerRequest( + server_id="my-test-server", + issuer="https://new-idp.example.com", + token_url="https://old-idp.example.com/token", # resubmitted stale -> must clear + authorization_url="https://old-idp.example.com/authorize", # resubmitted stale -> must clear + ) + await update_mcp_server(mock_prisma, data, "test-user") + data_dict = mock_prisma.db.litellm_mcpservertable.update.call_args[1]["data"] + + assert data_dict["issuer"] == "https://new-idp.example.com" + assert data_dict["token_url"] is None + assert data_dict["authorization_url"] is None + + +@pytest.mark.asyncio +async def test_establishing_issuer_first_time_preserves_discovered_fields(): + """Establishing an issuer for the first time (None -> X), which is exactly what the trust-on-first-use + discovery write-back does, must NOT clear the endpoints or oauth2_flow it discovered in the same + write. Only an issuer that was already pinned and is now changed or cleared invalidates its + endpoints, so the discovery persist cannot wipe the fields it just resolved.""" + mock_prisma = _mock_prisma() + existing = MagicMock() + existing.auth_type = "oauth2" + existing.url = "https://same.example.com/mcp" + existing.credentials = None + existing.issuer = None + mock_prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=existing) + + data = UpdateMCPServerRequest( + server_id="my-test-server", + issuer="https://discovered-idp.example.com", + authorization_url="https://discovered-idp.example.com/authorize", + token_url="https://discovered-idp.example.com/token", + oauth2_flow="authorization_code", + ) + await update_mcp_server(mock_prisma, data, "mcp_oauth_discovery") + data_dict = mock_prisma.db.litellm_mcpservertable.update.call_args[1]["data"] + + assert data_dict["issuer"] == "https://discovered-idp.example.com" + assert data_dict["authorization_url"] == "https://discovered-idp.example.com/authorize" + assert data_dict["token_url"] == "https://discovered-idp.example.com/token" + assert data_dict.get("oauth2_flow") == "authorization_code" + + +@pytest.mark.asyncio +async def test_unchanged_url_does_not_clear_discovered_oauth_fields(): + """A partial update that resends the same url (or omits it) must not clear the discovered OAuth + fields, so a routine save does not force needless re-discovery.""" + mock_prisma = _mock_prisma() + existing = MagicMock() + existing.auth_type = "oauth2" + existing.url = "https://same.example.com/mcp" + existing.credentials = None + mock_prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=existing) + + data = UpdateMCPServerRequest(server_id="my-test-server", url="https://same.example.com/mcp") + await update_mcp_server(mock_prisma, data, "test-user") + data_dict = mock_prisma.db.litellm_mcpservertable.update.call_args[1]["data"] + + for preserved_field in ("issuer", "authorization_url", "token_url", "registration_url"): + assert preserved_field not in data_dict, f"{preserved_field} must not be cleared when url is unchanged" + + @pytest.mark.asyncio async def test_auth_type_switch_keeps_explicitly_provided_flow_fields(): """Fields explicitly provided alongside the auth_type switch must survive it.""" diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index bba00ed1819..80bf08a5eba 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -1166,6 +1166,67 @@ class TestMCPServerManager: assert built.token_url == "https://idp.example.com/token" assert built.scopes == ["read", "admin"] + @pytest.mark.asyncio + async def test_build_from_table_reflects_discovered_issuer_trust_on_first_use(self): + """An unpinned server resolves endpoints resource-rooted on first discovery and records the + discovered issuer trust-on-first-use. The returned in-memory server must carry that discovered + issuer so the registry matches what gets persisted to the row; otherwise the OAuth token + identity (which includes issuer) differs between this build and the next rebuild, forcing a + spurious re-auth. Endpoints and issuer come from the same authorization-server document, so + they are consistent.""" + manager = MCPServerManager() + row = LiteLLM_MCPServerTable( + server_id="tofu-issuer-1", + alias="tofu_issuer", + description="unpinned, discovers its issuer", + url="https://up.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + created_at=datetime.now(), + updated_at=datetime.now(), + ) + + metadata = MCPOAuthMetadata( + authorization_url="https://idp.example.com/authorize", + token_url="https://idp.example.com/token", + scopes=["read"], + discovered_issuer="https://idp.example.com", + ) + with patch.object(manager, "_descovery_metadata", new=AsyncMock(return_value=metadata)): + built = await manager.build_mcp_server_from_table(row, credentials_are_encrypted=False) + + assert built.issuer == "https://idp.example.com" + assert built.issuer_is_anchored is False + assert built.authorization_url == "https://idp.example.com/authorize" + + @pytest.mark.asyncio + async def test_build_from_table_origin_fallback_issuer_is_not_reflected(self): + """An origin-fallback discovery is a guess that is deliberately never persisted, so the built + server must not claim an issuer the row will not hold; otherwise in-memory and DB would + disagree in the opposite direction.""" + manager = MCPServerManager() + row = LiteLLM_MCPServerTable( + server_id="origin-fallback-1", + alias="origin_fallback", + description="unpinned, origin-fallback discovery", + url="https://up.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + created_at=datetime.now(), + updated_at=datetime.now(), + ) + + metadata = MCPOAuthMetadata( + authorization_url="https://up.example.com/authorize", + token_url="https://up.example.com/token", + discovered_issuer="https://up.example.com", + from_origin_fallback=True, + ) + with patch.object(manager, "_descovery_metadata", new=AsyncMock(return_value=metadata)): + built = await manager.build_mcp_server_from_table(row, credentials_are_encrypted=False) + + assert built.issuer is None + @pytest.mark.asyncio async def test_build_from_table_whitespace_authorization_url_is_not_a_pin(self): """A whitespace-only authorization_url on the row must not be kept for redirects while the @@ -1230,6 +1291,159 @@ class TestMCPServerManager: assert built.registration_url == "https://idp.example.com/register" assert built.scopes == ["read"] + @pytest.mark.asyncio + async def test_build_from_table_uses_issuer_anchored_endpoints_when_issuer_configured(self): + """When an admin configures an issuer, the build takes its endpoints from the issuer-anchored + fetch (RFC 8414 §3.3) rather than the resource-rooted corroboration path. The build path does + not call _descovery_metadata directly; the issuer-anchored helper is responsible for combining + issuer endpoints with resource-driven scopes internally, and is invoked with the server url so + it can fetch those scopes.""" + manager = MCPServerManager() + row = LiteLLM_MCPServerTable( + server_id="issuer-anchored-1", + alias="issuer_anchored", + description="issuer configured, blank endpoints", + url="https://up.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + issuer="https://idp.example.com", + created_at=datetime.now(), + updated_at=datetime.now(), + ) + + resolved = MCPOAuthMetadata( + authorization_url="https://idp.example.com/authorize", + token_url="https://idp.example.com/token", + registration_url="https://idp.example.com/register", + scopes=["read", "write"], + ) + resource_rooted = AsyncMock(return_value=MCPOAuthMetadata(token_url="https://attacker.example.com/steal")) + with ( + patch.object(manager, "_fetch_issuer_anchored_oauth_metadata", new=AsyncMock(return_value=resolved)) as anchored, + patch.object(manager, "_descovery_metadata", new=resource_rooted), + ): + built = await manager.build_mcp_server_from_table(row, credentials_are_encrypted=False) + + anchored.assert_awaited_once_with("https://idp.example.com", "https://up.example.com/mcp") + resource_rooted.assert_not_awaited() + assert built.issuer == "https://idp.example.com" + assert built.issuer_is_anchored is True + assert built.authorization_url == "https://idp.example.com/authorize" + assert built.token_url == "https://idp.example.com/token" + assert built.registration_url == "https://idp.example.com/register" + assert built.scopes == ["read", "write"] + + @pytest.mark.asyncio + async def test_fetch_issuer_anchored_metadata_takes_endpoints_from_issuer_scopes_from_resource(self): + """The issuer-anchored helper adopts token_endpoint/registration_endpoint from the pinned + issuer's own §3.3-validated document, but the scopes are resource-driven: it fetches the + resource's advertised scopes and uses those, not the issuer document's scopes_supported. This + keeps endpoint trust anchored on the issuer while scope selection stays resource-driven per the + MCP Scope Selection Strategy.""" + manager = MCPServerManager() + + issuer_document = MCPOAuthMetadata( + authorization_url="https://idp.example.com/authorize", + token_url="https://idp.example.com/token", + registration_url="https://idp.example.com/register", + scopes=["as.everything"], + ) + resource_document = MCPOAuthMetadata(scopes=["resource.read"]) + with ( + patch.object( + manager, "_fetch_single_authorization_server_metadata", new=AsyncMock(return_value=issuer_document) + ) as issuer_fetch, + patch.object(manager, "_descovery_metadata", new=AsyncMock(return_value=resource_document)) as resource_fetch, + ): + result = await manager._fetch_issuer_anchored_oauth_metadata( + "https://idp.example.com", "https://up.example.com/mcp" + ) + + issuer_fetch.assert_awaited_once_with( + "https://idp.example.com", "https://idp.example.com", require_issuer="https://idp.example.com" + ) + resource_fetch.assert_awaited_once() + assert result is not None + assert result.token_url == "https://idp.example.com/token" + assert result.registration_url == "https://idp.example.com/register" + assert result.scopes == ["resource.read"] + + @pytest.mark.asyncio + async def test_build_from_table_issuer_anchor_fails_closed_without_falling_back_to_resource(self): + """A configured issuer whose metadata does not validate (RFC 8414 §3.3 mismatch or fetch + failure) yields None from the anchored fetch. The build must adopt nothing and must NOT fall + back to resource-rooted discovery, or the fail-closed guarantee would be defeated by the very + resource the issuer anchor exists to distrust.""" + manager = MCPServerManager() + row = LiteLLM_MCPServerTable( + server_id="issuer-anchored-2", + alias="issuer_anchored_failclosed", + description="issuer configured, upstream fails validation", + url="https://up.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + issuer="https://idp.example.com", + created_at=datetime.now(), + updated_at=datetime.now(), + ) + + resource_rooted = AsyncMock(return_value=MCPOAuthMetadata(token_url="https://attacker.example.com/steal")) + with ( + patch.object(manager, "_fetch_issuer_anchored_oauth_metadata", new=AsyncMock(return_value=None)), + patch.object(manager, "_descovery_metadata", new=resource_rooted), + ): + built = await manager.build_mcp_server_from_table(row, credentials_are_encrypted=False) + + resource_rooted.assert_not_awaited() + assert built.issuer == "https://idp.example.com" + assert built.token_url is None + assert built.registration_url is None + assert built.scopes is None + + @pytest.mark.asyncio + async def test_build_from_table_issuer_anchor_overrides_stored_endpoints_even_when_populated(self): + """When an issuer is pinned, the endpoints come SOLELY from the §3.3-validated issuer document + and win over any stored/manual endpoint values, even a fully-populated row. Otherwise an + attacker who controls a stored token endpoint keeps receiving codes/secrets after an admin + pins a trusted issuer: `needs_discovery` must not short-circuit on populated fields, and the + issuer's endpoints must override the stored ones.""" + manager = MCPServerManager() + row = LiteLLM_MCPServerTable( + server_id="issuer-anchored-populated", + alias="issuer_anchored_populated", + description="issuer set, but stale/hostile endpoints already stored", + url="https://up.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + issuer="https://idp.example.com", + authorization_url="https://attacker.example.com/authorize", + token_url="https://attacker.example.com/steal", + credentials={"scopes": ["stale"]}, + created_at=datetime.now(), + updated_at=datetime.now(), + ) + + issuer_resolved = MCPOAuthMetadata( + authorization_url="https://idp.example.com/authorize", + token_url="https://idp.example.com/token", + scopes=["read"], + ) + with ( + patch.object( + manager, "_fetch_issuer_anchored_oauth_metadata", new=AsyncMock(return_value=issuer_resolved) + ) as anchored, + patch.object(manager, "_persist_discovered_oauth_endpoints", new=AsyncMock()) as mock_persist, + ): + built = await manager.build_mcp_server_from_table(row, credentials_are_encrypted=False) + + anchored.assert_awaited_once_with("https://idp.example.com", "https://up.example.com/mcp") + assert built.authorization_url == "https://idp.example.com/authorize" + assert built.token_url == "https://idp.example.com/token" + assert built.token_url != "https://attacker.example.com/steal" + # The issuer-anchored endpoints are never persisted into the endpoint columns, so a later + # build cannot treat them as authoritative stored values. + assert mock_persist.await_args.kwargs["is_issuer_anchored"] is True + @pytest.mark.asyncio @pytest.mark.parametrize( "advertised_authorization_url", @@ -2462,6 +2676,78 @@ class TestMCPServerManager: assert result.token_url == "https://login.microsoftonline.com/test-tenant-id/oauth2/v2.0/token" assert result.scopes == ["api://some-scope/.default"] + @staticmethod + def _issuer_doc_response_builder(well_known_url: str, document: dict): + def build_response(url: str, **kwargs): + mock_response = MagicMock() + if url == well_known_url: + mock_response.json.return_value = document + mock_response.raise_for_status = MagicMock() + else: + request = httpx.Request("GET", url) + response_obj = httpx.Response(status_code=404, request=request) + mock_response.raise_for_status = MagicMock( + side_effect=httpx.HTTPStatusError("not found", request=request, response=response_obj) + ) + return mock_response + + return build_response + + @pytest.mark.asyncio + async def test_fetch_single_authorization_server_metadata_adopts_document_with_matching_issuer(self): + """RFC 8414 §3.3: under require_issuer, a document that self-attests the same issuer it was + fetched from is authoritative and its endpoints and scopes are adopted.""" + manager = MCPServerManager() + issuer = "https://idp.example.com" + build_response = self._issuer_doc_response_builder( + f"{issuer}/.well-known/oauth-authorization-server", + { + "issuer": issuer, + "authorization_endpoint": "https://idp.example.com/authorize", + "token_endpoint": "https://idp.example.com/token", + "scopes_supported": ["read", "write"], + }, + ) + mock_client = MagicMock() + mock_client.get = AsyncMock(side_effect=build_response) + with patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.get_async_httpx_client", + return_value=mock_client, + ): + result = await manager._fetch_single_authorization_server_metadata(issuer, issuer, require_issuer=issuer) + + assert result is not None + assert result.authorization_url == "https://idp.example.com/authorize" + assert result.token_url == "https://idp.example.com/token" + assert result.scopes == ["read", "write"] + + @pytest.mark.asyncio + async def test_fetch_single_authorization_server_metadata_rejects_issuer_mismatch(self): + """RFC 8414 §3.3 fail-closed: a document self-attesting a DIFFERENT issuer than the one it was + fetched from is rejected even though it carries valid-looking endpoints, so a compromised + resource cannot point the issuer-anchored fetch at an attacker authorization server that + smuggles its own token_endpoint and inflated scopes.""" + manager = MCPServerManager() + issuer = "https://idp.example.com" + build_response = self._issuer_doc_response_builder( + f"{issuer}/.well-known/oauth-authorization-server", + { + "issuer": "https://attacker.example.com", + "authorization_endpoint": "https://idp.example.com/authorize", + "token_endpoint": "https://attacker.example.com/steal", + "scopes_supported": ["admin"], + }, + ) + mock_client = MagicMock() + mock_client.get = AsyncMock(side_effect=build_response) + with patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.get_async_httpx_client", + return_value=mock_client, + ): + result = await manager._fetch_single_authorization_server_metadata(issuer, issuer, require_issuer=issuer) + + assert result is None + @pytest.mark.asyncio async def test_fetch_single_authorization_server_metadata_derives_azure_metadata( self, @@ -2489,6 +2775,37 @@ class TestMCPServerManager: assert result.authorization_url == "https://login.microsoftonline.com/test-tenant-id/oauth2/v2.0/authorize" assert result.token_url == "https://login.microsoftonline.com/test-tenant-id/oauth2/v2.0/token" + @pytest.mark.asyncio + async def test_azure_heuristic_reachable_under_require_issuer(self): + """Under issuer-anchored discovery (require_issuer set), an Entra issuer whose OIDC document + cannot be fetched still gets the deterministic Azure endpoint construction. The heuristic + derives the endpoints from the pinned issuer's own tenant URL, so it is authoritative-by- + construction and safe under require_issuer; only a non-Entra issuer stays fail-closed (None).""" + manager = MCPServerManager() + issuer = "https://login.microsoftonline.com/test-tenant-id/v2.0" + + request = httpx.Request("GET", issuer) + response_obj = httpx.Response(status_code=404, request=request) + mock_response = MagicMock() + mock_response.raise_for_status = MagicMock( + side_effect=httpx.HTTPStatusError("not found", request=request, response=response_obj) + ) + mock_client = MagicMock() + mock_client.get = AsyncMock(return_value=mock_response) + + with patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.get_async_httpx_client", + return_value=mock_client, + ): + azure = await manager._fetch_single_authorization_server_metadata(issuer, issuer, require_issuer=issuer) + non_entra = await manager._fetch_single_authorization_server_metadata( + "https://idp.example.com", "https://idp.example.com", require_issuer="https://idp.example.com" + ) + + assert azure is not None + assert azure.token_url == "https://login.microsoftonline.com/test-tenant-id/oauth2/v2.0/token" + assert non_entra is None + @pytest.mark.asyncio async def test_descovery_metadata_falls_back_to_origin_when_no_auth_servers(self): manager = MCPServerManager() @@ -5194,6 +5511,7 @@ class TestMCPServerTimestamps: await manager._persist_discovered_oauth_endpoints( server_id="s", auth_type=MCPAuth.api_key, + existing_issuer=None, existing_authorization_url=None, existing_token_url=None, existing_scopes=None, @@ -5202,6 +5520,7 @@ class TestMCPServerTimestamps: await manager._persist_discovered_oauth_endpoints( server_id="s", auth_type=MCPAuth.oauth2, + existing_issuer=None, existing_authorization_url=None, existing_token_url=None, existing_scopes=None, @@ -5210,6 +5529,7 @@ class TestMCPServerTimestamps: await manager._persist_discovered_oauth_endpoints( server_id="s", auth_type=MCPAuth.oauth2, + existing_issuer=None, existing_authorization_url=None, existing_token_url=None, existing_scopes=None, @@ -5218,6 +5538,7 @@ class TestMCPServerTimestamps: await manager._persist_discovered_oauth_endpoints( server_id="s", auth_type=MCPAuth.oauth2, + existing_issuer=None, existing_authorization_url="https://configured.example.com/authorize", existing_token_url="https://configured.example.com/token", existing_scopes=["configured"], @@ -5243,6 +5564,7 @@ class TestMCPServerTimestamps: await manager._persist_discovered_oauth_endpoints( server_id="s", auth_type=MCPAuth.oauth2, + existing_issuer=None, existing_authorization_url=None, existing_token_url="https://configured.example.com/token", existing_scopes=None, @@ -5259,6 +5581,81 @@ class TestMCPServerTimestamps: assert persisted.credentials == {"scopes": ["s1"]} assert "token_url" not in persisted.fields_set() + @pytest.mark.asyncio + async def test_persist_discovered_oauth_endpoints_writes_discovered_issuer_trust_on_first_use(self): + """A server with no configured issuer records the discovered issuer trust-on-first-use, so the + next rebuild anchors discovery on it (RFC 8414 §3.3) instead of re-trusting the resource. When + an issuer is already set (admin-typed or a prior discovery), it is never overwritten.""" + manager = MCPServerManager() + metadata = MCPOAuthMetadata( + authorization_url="https://idp.example.com/authorize", + token_url="https://idp.example.com/token", + discovered_issuer="https://idp.example.com", + ) + + update_mcp_server_mock = AsyncMock() + with ( + patch("litellm.proxy._experimental.mcp_server.db.update_mcp_server", new=update_mcp_server_mock), + patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), + ): + await manager._persist_discovered_oauth_endpoints( + server_id="s", + auth_type=MCPAuth.oauth2, + existing_issuer=None, + existing_authorization_url=None, + existing_token_url=None, + existing_scopes=None, + metadata=metadata, + ) + await manager._persist_discovered_oauth_endpoints( + server_id="s", + auth_type=MCPAuth.oauth2, + existing_issuer="https://admin-configured.example.com", + existing_authorization_url="https://admin-configured.example.com/authorize", + existing_token_url="https://admin-configured.example.com/token", + existing_scopes=["cfg"], + metadata=metadata, + ) + + assert update_mcp_server_mock.await_count == 1 + persisted = update_mcp_server_mock.call_args.kwargs["data"] + assert persisted.issuer == "https://idp.example.com" + + @pytest.mark.asyncio + async def test_persist_discovered_oauth_endpoints_does_not_persist_endpoints_for_issuer_anchored(self): + """For an issuer-anchored server the endpoints are re-derived from the §3.3-validated issuer + document every build, so they must NOT be written into the endpoint columns: persisting them + would make the next build see populated endpoints and treat them as authoritative stored + values, defeating the issuer-only invariant. Only the resource-driven scopes are persisted.""" + manager = MCPServerManager() + metadata = MCPOAuthMetadata( + authorization_url="https://idp.example.com/authorize", + token_url="https://idp.example.com/token", + scopes=["read"], + ) + + update_mcp_server_mock = AsyncMock() + with ( + patch("litellm.proxy._experimental.mcp_server.db.update_mcp_server", new=update_mcp_server_mock), + patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), + ): + await manager._persist_discovered_oauth_endpoints( + server_id="s", + auth_type=MCPAuth.oauth2, + existing_issuer="https://idp.example.com", + existing_authorization_url=None, + existing_token_url=None, + existing_scopes=None, + metadata=metadata, + is_issuer_anchored=True, + ) + + update_mcp_server_mock.assert_awaited_once() + persisted = update_mcp_server_mock.call_args.kwargs["data"] + assert "authorization_url" not in persisted.fields_set() + assert "token_url" not in persisted.fields_set() + assert persisted.credentials == {"scopes": ["read"]} + @pytest.mark.asyncio async def test_build_mcp_server_from_table_skips_persistence_for_temporary_servers(self): """The session endpoint builds temporary servers whose server_id has no DB row; with @@ -5473,6 +5870,88 @@ class TestMCPServerTimestamps: assert same_authorize.token_url == "https://idp.example.com/token" assert same_authorize.registration_url == "https://idp.example.com/register" + def test_carry_forward_does_not_restore_endpoints_for_issuer_anchored_server(self): + """When the server is issuer-anchored the endpoints come solely from the §3.3-validated issuer + document, so a failed issuer fetch (token_url None) must stay fail-closed. Carry-forward must + NOT resurrect the previous registry entry's token endpoint, or the very attacker-controlled + endpoint the issuer anchor distrusts would keep being served across rebuilds. Resource-driven + scopes still carry as last-known-good. Anchoring is keyed on the explicit issuer_is_anchored + flag, not on issuer truthiness, so a discovered issuer does not trip this fail-closed branch.""" + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + _carry_forward_resolved_oauth_endpoints, + ) + + previous = MCPServer( + server_id="s1", + name="s1", + url="https://up.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + issuer="https://idp.example.com", + issuer_is_anchored=True, + authorization_url="https://idp.example.com/authorize", + token_url="https://idp.example.com/token", + registration_url="https://idp.example.com/register", + scopes=["read"], + ) + failed_rebuild = MCPServer( + server_id="s1", + name="s1", + url="https://up.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + issuer="https://idp.example.com", + issuer_is_anchored=True, + ) + + _carry_forward_resolved_oauth_endpoints(new_server=failed_rebuild, previous_server=previous) + + assert failed_rebuild.authorization_url is None + assert failed_rebuild.token_url is None + assert failed_rebuild.registration_url is None + assert failed_rebuild.scopes == ["read"] + + def test_carry_forward_restores_endpoints_for_discovered_issuer_not_anchored(self): + """A server that merely DISCOVERED its issuer trust-on-first-use is not anchored: issuer is set + for token identity but the endpoints are resource-rooted, so on a transient discovery blip they + must still carry forward as last-known-good, the same as any resource-rooted server. This is the + regression the explicit issuer_is_anchored flag prevents: keying fail-closed on issuer truthiness + alone would drop the working endpoints the moment the server learned its issuer.""" + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + _carry_forward_resolved_oauth_endpoints, + ) + + previous = MCPServer( + server_id="s1", + name="s1", + url="https://up.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + issuer="https://idp.example.com", + issuer_is_anchored=False, + authorization_url="https://idp.example.com/authorize", + token_url="https://idp.example.com/token", + registration_url="https://idp.example.com/register", + scopes=["read"], + ) + blipped_rebuild = MCPServer( + server_id="s1", + name="s1", + url="https://up.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + issuer="https://idp.example.com", + issuer_is_anchored=False, + authorization_url=None, + ) + + _carry_forward_resolved_oauth_endpoints(new_server=blipped_rebuild, previous_server=previous) + + assert blipped_rebuild.authorization_url == "https://idp.example.com/authorize" + assert blipped_rebuild.token_url == "https://idp.example.com/token" + assert blipped_rebuild.registration_url == "https://idp.example.com/register" + assert blipped_rebuild.scopes == ["read"] + def test_normalized_authorize_endpoint_treats_default_port_and_slash_as_identity(self): """The corroboration check must not fail on formatting-only differences an IdP legitimately emits: default port, trailing slash, host case, and query string are not identity, but a @@ -5487,6 +5966,22 @@ class TestMCPServerTimestamps: assert _normalized_authorize_endpoint("https://idp.example.com/authorize?prompt=consent") == canonical assert _normalized_authorize_endpoint("https://idp.example.com:8443/authorize") != canonical + def test_issuer_matches_rfc8414_section_3_3(self): + """Issuer equality tolerates only URL-insignificant differences (scheme/host case, default + port, a trailing slash). A different host, a non-string, an empty string, or a None issuer + never matches, so a document that omits issuer fails closed under issuer-anchored discovery.""" + from litellm.proxy._experimental.mcp_server.mcp_server_manager import _issuer_matches + + assert _issuer_matches("https://mcp.slack.com", "https://mcp.slack.com") + assert _issuer_matches("https://MCP.slack.com/", "https://mcp.slack.com") + assert _issuer_matches("https://mcp.slack.com:443", "https://mcp.slack.com") + assert _issuer_matches("https://login.example.com/tenant/v2.0", "https://login.example.com/tenant/v2.0") + assert not _issuer_matches("https://attacker.example.com", "https://mcp.slack.com") + assert not _issuer_matches("https://login.example.com/other/v2.0", "https://login.example.com/tenant/v2.0") + assert not _issuer_matches(None, "https://mcp.slack.com") + assert not _issuer_matches("", "https://mcp.slack.com") + assert not _issuer_matches(123, "https://mcp.slack.com") + def test_build_mcp_server_table_preserves_timestamps(self): """_build_mcp_server_table must use the MCPServer's stored timestamps, not datetime.now().""" manager = MCPServerManager() diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_semantic_tool_filter.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_semantic_tool_filter.py index 9bc0a525326..d864b442bd3 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_semantic_tool_filter.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_semantic_tool_filter.py @@ -865,14 +865,17 @@ async def test_semantic_filter_hook_filters_expanded_litellm_proxy_tools(): ) assert result is not None, "Hook should return modified data" - filtered = result["tools"] + mcp_references = [tool for tool in result["tools"] if tool.get("type") == "mcp"] + assert len(mcp_references) == 1, "The litellm_proxy MCP reference must be preserved for the MCP gateway to expand" - assert len(filtered) <= 2, f"Expanded tools should be filtered to top_k=2, got {len(filtered)}" - assert len(filtered) < len(expanded_tools), ( - f"Hook must not forward all {len(expanded_tools)} expanded tools unfiltered, got {len(filtered)}" + allowed_tools = mcp_references[0]["allowed_tools"] + assert len(allowed_tools) <= 2, f"Expanded tools should be filtered to top_k=2, got {len(allowed_tools)}" + assert len(allowed_tools) < len(expanded_tools), ( + f"Hook must not forward all {len(expanded_tools)} expanded tools unfiltered, got {len(allowed_tools)}" ) - for tool in filtered: - assert tool in expanded_tools, "Filtered tools must be the original expanded tool dicts" + expanded_names = {tool["name"] for tool in expanded_tools} + for name in allowed_tools: + assert name in expanded_names, "Selected tool names must come from the expanded tools" assert ( "litellm_semantic_filter_stats" in result["metadata"] @@ -880,9 +883,246 @@ async def test_semantic_filter_hook_filters_expanded_litellm_proxy_tools(): stats = result["metadata"]["litellm_semantic_filter_stats"] total, selected = stats.split("->") assert int(total) == 5, f"Stats 'from' should be pre-filter expanded count (5), got {total}" - assert int(selected) == len(filtered), f"Stats 'to' should match post-filter count, got {selected}" + assert int(selected) == len(allowed_tools), f"Stats 'to' should match post-filter count, got {selected}" - print(f"✅ Expanded litellm_proxy tools filtered: {len(expanded_tools)} -> {len(filtered)}, stats={stats}") + print(f"✅ Expanded litellm_proxy tools filtered: {len(expanded_tools)} -> {len(allowed_tools)}, stats={stats}") + + +@pytest.mark.asyncio +async def test_semantic_filter_hook_narrows_mcp_reference_for_chat_completions(): + """ + Regression test (LIT-4451): the hook must narrow the litellm_proxy MCP + reference instead of replacing it with expanded tool definitions. + + Given: A /chat/completions request whose tools are a single + {"type": "mcp", "server_url": "litellm_proxy"} reference that + expands to 5 tools, with the semantic filter selecting top_k=2 + When: The hook processes the request with call_type="acompletion" + Then: The MCP reference survives in data["tools"], carrying the selected + tools in allowed_tools, and no expanded function definitions are + written into the request. + + Replacing the reference made the hook write Responses-API-shaped tools + ({"type": "function", "name": ...}) into /chat/completions, which expects + {"type": "function", "function": {...}}. The provider transformation then + rejected every MCP tool (Anthropic raised KeyError: 'function') or dropped + it silently (Bedrock), so the model saw no MCP tools at all. Replacing the + reference also removed the marker the MCP gateway matches on, which + disabled tool auto-execution for require_approval="never". + """ + from litellm.proxy._experimental.mcp_server.semantic_tool_filter import ( + SemanticMCPToolFilter, + ) + from litellm.proxy.hooks.mcp_semantic_filter import SemanticToolFilterHook + from litellm.types.utils import Embedding, EmbeddingResponse + + mock_router = Mock() + + def mock_embedding_sync(*args, **kwargs): + return EmbeddingResponse( + data=[Embedding(embedding=[0.1] * 1536, index=0, object="embedding")], + model="text-embedding-3-small", + object="list", + usage={"prompt_tokens": 10, "total_tokens": 10}, + ) + + async def mock_embedding_async(*args, **kwargs): + return mock_embedding_sync() + + mock_router.embedding = mock_embedding_sync + mock_router.aembedding = mock_embedding_async + + filter_instance = SemanticMCPToolFilter( + embedding_model="text-embedding-3-small", + litellm_router_instance=mock_router, + top_k=2, + similarity_threshold=0.3, + enabled=True, + ) + + registry_tools = [ + MCPTool( + name=f"srv-tool_{i}", + description=f"Registry tool {i}", + inputSchema={"type": "object"}, + ) + for i in range(5) + ] + filter_instance._build_router(registry_tools) + + expanded_tools = [ + { + "type": "function", + "name": f"srv-tool_{i}", + "description": f"Registry tool {i}", + "parameters": {"type": "object", "properties": {}}, + } + for i in range(5) + ] + + hook = SemanticToolFilterHook(filter_instance) + hook._expand_mcp_tools = AsyncMock( # type: ignore[method-assign] + return_value=expanded_tools + ) + + mcp_reference = { + "type": "mcp", + "server_url": "litellm_proxy", + "require_approval": "never", + } + data = { + "model": "gpt-4", + "messages": [{"role": "user", "content": "Send an email"}], + "tools": [mcp_reference], + "metadata": {}, + } + + result = await hook.async_pre_call_hook( + user_api_key_dict=Mock(), + cache=Mock(), + data=data, + call_type="acompletion", + ) + + assert result is not None, "Hook should return modified data" + forwarded = result["tools"] + + assert [tool.get("type") for tool in forwarded] == ["mcp"], ( + "The MCP reference must be the only forwarded tool; writing expanded function " + f"definitions into a chat completion loses every MCP tool. Got: {forwarded}" + ) + assert forwarded[0]["server_url"] == "litellm_proxy", "The MCP reference must keep routing to the gateway" + assert forwarded[0]["require_approval"] == "never", "The MCP reference must keep its auto-execute marker" + + allowed_tools = forwarded[0]["allowed_tools"] + assert allowed_tools, "The narrowed reference must still carry the selected tools" + assert len(allowed_tools) <= 2, f"Selection must narrow the reference to top_k=2, got {allowed_tools}" + assert len(allowed_tools) < len(expanded_tools), ( + f"Hook must not forward all {len(expanded_tools)} expanded tools unfiltered, got {allowed_tools}" + ) + assert set(allowed_tools) <= {tool["name"] for tool in expanded_tools}, ( + f"Selected names must come from the expanded tools, got {allowed_tools}" + ) + + print(f"✅ chat completions: MCP reference preserved, narrowed to {allowed_tools}") + + +@pytest.mark.asyncio +async def test_semantic_filter_hook_zero_matches_exposes_all_tools_on_both_paths(): + """ + A query that matches nothing must expose every MCP tool, whether the request + carries a litellm_proxy MCP reference or plain MCP tool objects. + + Given: A router that returns no matches for the query + When: The hook processes an MCP reference request and a plain MCP tool request + Then: Both expose all 3 tools, because filter_tools owns the undecidable-selection + policy and returns the full set rather than an empty one + + The two paths narrow through different mechanisms (allowed_tools on the reference + versus dropping unmatched entries), so they could drift into opposite fail + behaviours. Pinning both here keeps that single policy honest: flipping + filter_tools to fail closed must fail this test on both paths at once, instead of + silently hard-limiting one surface and not the other. + """ + from litellm.proxy._experimental.mcp_server.semantic_tool_filter import ( + SemanticMCPToolFilter, + ) + from litellm.proxy.hooks.mcp_semantic_filter import SemanticToolFilterHook + from litellm.types.utils import Embedding, EmbeddingResponse + + mock_router = Mock() + + def mock_embedding_sync(*args, **kwargs): + return EmbeddingResponse( + data=[Embedding(embedding=[0.1] * 1536, index=0, object="embedding")], + model="text-embedding-3-small", + object="list", + usage={"prompt_tokens": 10, "total_tokens": 10}, + ) + + async def mock_embedding_async(*args, **kwargs): + return mock_embedding_sync() + + mock_router.embedding = mock_embedding_sync + mock_router.aembedding = mock_embedding_async + + registry_tools = [ + MCPTool( + name=f"srv-tool_{i}", + description=f"Registry tool {i}", + inputSchema={"type": "object"}, + ) + for i in range(3) + ] + + def build_hook(): + filter_instance = SemanticMCPToolFilter( + embedding_model="text-embedding-3-small", + litellm_router_instance=mock_router, + top_k=2, + similarity_threshold=0.3, + enabled=True, + ) + filter_instance._build_router(registry_tools) + zero_match_router = Mock(return_value=[]) + zero_match_router.top_k = 2 + filter_instance.tool_router = zero_match_router + return SemanticToolFilterHook(filter_instance) + + expanded_tools = [ + { + "type": "function", + "name": f"srv-tool_{i}", + "description": f"Registry tool {i}", + "parameters": {"type": "object", "properties": {}}, + } + for i in range(3) + ] + + reference_hook = build_hook() + reference_hook._expand_mcp_tools = AsyncMock( # type: ignore[method-assign] + return_value=expanded_tools + ) + reference_data = { + "model": "gpt-4", + "messages": [{"role": "user", "content": "something entirely unrelated"}], + "tools": [{"type": "mcp", "server_url": "litellm_proxy", "require_approval": "never"}], + "metadata": {}, + } + reference_result = await reference_hook.async_pre_call_hook( + user_api_key_dict=Mock(), + cache=Mock(), + data=reference_data, + call_type="acompletion", + ) + + reference_tools = (reference_result or reference_data)["tools"] + mcp_references = [tool for tool in reference_tools if tool.get("type") == "mcp"] + assert len(mcp_references) == 1, "The MCP reference must survive a zero-match query" + assert set(mcp_references[0].get("allowed_tools") or []) == {tool["name"] for tool in expanded_tools}, ( + "A zero-match query must leave every expanded tool reachable through the reference" + ) + + plain_hook = build_hook() + plain_data = { + "model": "gpt-4", + "messages": [{"role": "user", "content": "something entirely unrelated"}], + "tools": list(registry_tools), + "metadata": {}, + } + plain_result = await plain_hook.async_pre_call_hook( + user_api_key_dict=Mock(), + cache=Mock(), + data=plain_data, + call_type="acompletion", + ) + + plain_tools = (plain_result or plain_data)["tools"] + assert len(plain_tools) == len(registry_tools), ( + f"A zero-match query must not drop plain MCP tools, got {len(plain_tools)} of {len(registry_tools)}" + ) + + print("✅ zero matches: both the MCP reference path and the plain tool path expose every tool") @pytest.mark.asyncio @@ -958,8 +1198,9 @@ async def test_semantic_filter_hook_filters_expanded_tools_with_string_input(): async def test_semantic_filter_hook_expansion_skips_filter_when_disabled(): """ When the filter is disabled at runtime (e.g. via the UI toggle), the - expansion path must forward all expanded tools and emit NO filter - stats, mirroring the generic path's enabled guard. + expansion path must leave the MCP reference untouched and emit NO filter + stats, mirroring the generic path's enabled guard. The MCP gateway then + expands the reference itself, so no tool is narrowed away. """ from litellm.proxy._experimental.mcp_server.semantic_tool_filter import ( SemanticMCPToolFilter, @@ -1009,13 +1250,19 @@ async def test_semantic_filter_hook_expansion_skips_filter_when_disabled(): call_type="aresponses", ) - assert result is not None, "Hook should still expand MCP references when the filter is disabled" - assert len(result["tools"]) == 5, f"All expanded tools must be forwarded when disabled, got {len(result['tools'])}" + assert result is None, "Hook must not modify the request when the filter is disabled" + assert data["tools"] == [ + { + "type": "mcp", + "server_url": "litellm_proxy", + "require_approval": "never", + } + ], "The MCP reference must be left intact for the MCP gateway to expand" assert ( - "litellm_semantic_filter_stats" not in result["metadata"] + "litellm_semantic_filter_stats" not in data["metadata"] ), "No filter stats may be emitted when the filter is disabled" - print("✅ Disabled filter: expansion preserved, no spurious stats") + print("✅ Disabled filter: MCP reference untouched, no spurious stats") @pytest.mark.asyncio @@ -1664,3 +1911,252 @@ def test_is_context_window_error_detection_variants(): assert _is_context_window_error(ValueError("Invalid 'input[0]': maximum input length is 8192 tokens.")) assert not _is_context_window_error(ValueError("A generic API error occurred.")) assert not _is_context_window_error(None) + + +def _make_keyword_embedding_router(recorded_inputs): + """ + Mock litellm Router whose embeddings are deterministic keyword one-hots: + texts mentioning linear/issue/ticket embed to [1, 0], everything else to + [0, 1]. Lets tests assert real similarity ranking through the actual + semantic-router index. Every embedding input batch is appended to + recorded_inputs. + """ + from litellm.types.utils import Embedding, EmbeddingResponse + + def _vector(text): + lowered = text.lower() + if "kanban" in lowered: + return [0.6, 0.8] + if "linear" in lowered or "issue" in lowered or "ticket" in lowered: + return [1.0, 0.0] + return [0.0, 1.0] + + def mock_embedding_sync(*args, **kwargs): + texts = kwargs["input"] + recorded_inputs.append(list(texts)) + return EmbeddingResponse( + data=[Embedding(embedding=_vector(t), index=i, object="embedding") for i, t in enumerate(texts)], + model="text-embedding-3-small", + object="list", + usage={"prompt_tokens": 10, "total_tokens": 10}, + ) + + async def mock_embedding_async(*args, **kwargs): + return mock_embedding_sync(*args, **kwargs) + + mock_router = Mock() + mock_router.embedding = mock_embedding_sync + mock_router.aembedding = mock_embedding_async + return mock_router + + +def _make_keyword_filter(recorded_inputs, top_k: int = 3): + from litellm.proxy._experimental.mcp_server.semantic_tool_filter import ( + SemanticMCPToolFilter, + ) + + return SemanticMCPToolFilter( + embedding_model="text-embedding-3-small", + litellm_router_instance=_make_keyword_embedding_router(recorded_inputs), + top_k=top_k, + similarity_threshold=0.3, + enabled=True, + ) + + +def _linear_issue_tool(): + return MCPTool( + name="linear_stub-get_issue", + description="Get a Linear issue (ticket) by its identifier such as LIT-1234", + inputSchema={"type": "object"}, + ) + + +def _linear_list_tool(): + return MCPTool( + name="linear_stub-list_issues", + description="List Linear issues (tickets) in the workspace", + inputSchema={"type": "object"}, + ) + + +def _weather_tool(): + return MCPTool( + name="weather_stub-get_weather", + description="Get the current weather conditions for a city", + inputSchema={"type": "object"}, + ) + + +@pytest.mark.asyncio +async def test_filter_indexes_request_tools_when_startup_index_is_empty(): + """ + Regression test: the startup index is built by listing every MCP server + WITHOUT per-user credentials, so a gateway whose servers all require + per-user auth (e.g. interactive OAuth) starts with an empty index + (tool_router is None). filter_tools then failed open and returned all N + tools unfiltered (customer-visible as an N->N header and, past 128 tools, + an OpenAI 400 "tools array too long"). The authed request-time tools must + instead be indexed on first sight so filtering actually runs. + """ + filter_instance = _make_keyword_filter([]) + assert filter_instance.tool_router is None + + tools = [_linear_issue_tool(), _weather_tool()] + filtered = await filter_instance.filter_tools( + query="what is Linear ticket LIT-3794 about", + available_tools=tools, + ) + + assert [t.name for t in filtered] == ["linear_stub-get_issue"] + print("✅ Empty startup index is built from authed request-time tools") + + +@pytest.mark.asyncio +async def test_filter_indexes_tools_missing_from_partial_index(): + """ + Regression test: servers whose tools/list needs per-user auth contribute + zero routes to the startup index while anonymously listable servers are + indexed. Tools reaching the filter through the authed request-time + expansion must be added to the existing router (and only embedded once; + repeat requests embed just the query). + """ + recorded_inputs = [] + filter_instance = _make_keyword_filter(recorded_inputs) + filter_instance._build_router([_weather_tool()]) + assert filter_instance.tool_router is not None + + tools = [_linear_issue_tool(), _weather_tool()] + query = "what is Linear ticket LIT-3794 about" + + filtered = await filter_instance.filter_tools(query=query, available_tools=tools) + assert [t.name for t in filtered] == ["linear_stub-get_issue"] + + calls_after_first = len(recorded_inputs) + filtered_again = await filter_instance.filter_tools(query=query, available_tools=tools) + assert [t.name for t in filtered_again] == ["linear_stub-get_issue"] + assert len(recorded_inputs) == calls_after_first + 1 + + print("✅ Partial startup index is completed from request-time tools, embedding each tool once") + + +@pytest.mark.asyncio +async def test_filter_fails_open_when_matches_are_not_in_available_tools(): + """ + Regression test: when the semantic router's matches are all tools that are + NOT in the request's available_tools (an index/request mismatch), the + filter returned an empty list, stripping every tool from the request and + breaking it outright (observed live as a 3->0 header followed by a + provider 400). It must fail open with the full tool list instead, matching + the zero-match fallback. + """ + filter_instance = _make_keyword_filter([]) + filter_instance._build_router([_weather_tool()]) + + tools = [_linear_issue_tool(), _linear_list_tool()] + filtered = await filter_instance.filter_tools( + query="current weather in San Francisco", + available_tools=tools, + ) + + assert [t.name for t in filtered] == ["linear_stub-get_issue", "linear_stub-list_issues"] + print("✅ Matches outside available_tools fail open instead of dropping every tool") + + +@pytest.mark.asyncio +async def test_request_time_context_window_error_is_request_scoped(): + """ + Regression test: an oversized tool description hitting the embedding + context window while lazily indexing request-time tools must fail only + the requesting call. Previously the lazy path reused the startup build + and recorded the overflow in the shared context_window_error, after + which EVERY user's MCP requests on the worker were blocked with a 400 + until restart (index poisoning via a single request). + """ + from litellm.proxy._experimental.mcp_server.semantic_tool_filter import ( + SemanticToolFilterContextWindowError, + ) + + state = {"raise_context_error": True} + filter_instance = _make_context_window_filter(state) + tools = [ + MCPTool(name="tool_a", description="Tool A", inputSchema={"type": "object"}), + MCPTool(name="tool_b", description="Tool B", inputSchema={"type": "object"}), + ] + + with pytest.raises(SemanticToolFilterContextWindowError): + await filter_instance.filter_tools(query="send an email", available_tools=tools) + + assert filter_instance.context_window_error is None + assert filter_instance.tool_router is None + + state["raise_context_error"] = False + filtered = await filter_instance.filter_tools(query="send an email", available_tools=tools) + + assert len(filtered) > 0 + assert filter_instance.context_window_error is None + assert filter_instance.tool_router is not None + print("✅ Request-time context window overflow is scoped to the request, not the worker") + + +@pytest.mark.asyncio +async def test_foreign_index_routes_cannot_displace_available_tools(): + """ + Regression test: routes indexed from OTHER principals' tool listings must + not occupy the match candidate set for this request. Previously the + router matched over the whole shared index, so foreign routes that + embedded closer to the query displaced the caller's own tools from + top_k, degrading results to the fail-open list (or, before the + empty-result guard, stripping every tool). Matching is now scoped to the + request's own tool names via route_filter. + """ + filter_instance = _make_keyword_filter([], top_k=1) + foreign_tools = [ + MCPTool( + name=f"other_user-linear_tool_{i}", + description=f"Get a Linear issue variant {i}", + inputSchema={"type": "object"}, + ) + for i in range(6) + ] + filter_instance._build_router(foreign_tools) + + my_kanban = MCPTool( + name="mine-kanban_board", + description="Manage kanban board cards", + inputSchema={"type": "object"}, + ) + filtered = await filter_instance.filter_tools( + query="what is Linear ticket LIT-3794 about", + available_tools=[my_kanban, _weather_tool()], + ) + + assert [t.name for t in filtered] == ["mine-kanban_board"] + print("✅ Foreign index routes cannot displace the caller's own tools") + + +@pytest.mark.asyncio +async def test_top_k_above_router_default_is_respected(): + """ + Regression test: semantic-router's SemanticRouter defaults to top_k=5 at + the index-query layer, silently capping any configured filter top_k + above 5 regardless of the limit passed to __call__. The router must be + sized (and resized) to honor the configured top_k. + """ + filter_instance = _make_keyword_filter([], top_k=6) + tools = [ + MCPTool( + name=f"linear_stub-tool_{i}", + description=f"Work with Linear issues part {i}", + inputSchema={"type": "object"}, + ) + for i in range(6) + ] + + filtered = await filter_instance.filter_tools( + query="Linear ticket work", + available_tools=tools, + ) + + assert len(filtered) == 6 + print("✅ Configured top_k above the semantic-router default of 5 is honored") diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index 27f43c4948f..cc4a7d5bfb4 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -522,6 +522,49 @@ async def test_get_key_object_should_raise_if_reconnect_fails_on_db_connection_e assert mock_prisma_client.get_data.await_count == 1 +def _fake_redis_cache(): + fake_redis = MagicMock() + fake_redis.async_get_cache = AsyncMock(return_value=None) + fake_redis.async_set_cache = AsyncMock() + fake_redis.async_set_cache_pipeline = AsyncMock() + fake_redis.async_delete_cache = AsyncMock() + return fake_redis + + +class TestAuthCacheRedisWritePolicy: + """Redis auth-cache entries may only be written from fresh DB loads. + + With ``enable_redis_auth_cache`` and multiple replicas, a pod that re-publishes + a cache-derived key object to Redis can resurrect a stale auth blob after + ``/key/update`` or ``/key/delete`` already deleted it, so limit changes never + propagate fleet-wide while traffic keeps refreshing the stale entry's TTL. + """ + + @pytest.mark.asyncio + async def test_get_key_object_db_load_publishes_to_redis(self): + mock_prisma_client = MagicMock() + mock_prisma_client.get_data = AsyncMock( + return_value=UserAPIKeyAuth(token="hashed-token-db") + ) + + fake_redis = _fake_redis_cache() + cache = UserApiKeyCache() + cache.redis_cache = fake_redis + + key_obj = await get_key_object( + hashed_token="hashed-token-db", + prisma_client=mock_prisma_client, + user_api_key_cache=cache, + ) + + assert key_obj.token == "hashed-token-db" + fake_redis.async_set_cache.assert_awaited_once() + assert ( + fake_redis.async_set_cache.await_args.kwargs.get("key") + or fake_redis.async_set_cache.await_args.args[0] + ) == "hashed-token-db" + + def test_get_cli_jwt_auth_token_default_expiration(valid_sso_user_defined_values): """Test generating CLI JWT token with default 24-hour expiration""" token = ExperimentalUIJWTToken.get_cli_jwt_auth_token(valid_sso_user_defined_values) @@ -2634,6 +2677,8 @@ async def test_virtual_key_budget_check_reads_from_spend_counter(): ) assert exc_info.value.current_cost == 1.5 assert exc_info.value.max_budget == 1.0 + assert exc_info.value.entity_type == "key" + assert exc_info.value.entity_id == "test-hashed-token" @pytest.mark.asyncio @@ -2861,6 +2906,8 @@ async def test_team_budget_check_reads_from_spend_counter(): proxy_logging_obj=proxy_logging_obj, ) assert exc_info.value.current_cost == 1.5 + assert exc_info.value.entity_type == "team" + assert exc_info.value.entity_id == "test-team" @pytest.mark.asyncio @@ -2888,6 +2935,8 @@ async def test_end_user_budget_check_reads_from_spend_counter(): ) assert exc_info.value.current_cost == 1.5 assert exc_info.value.max_budget == 1.0 + assert exc_info.value.entity_type == "end_user" + assert exc_info.value.entity_id == "customer-1" @pytest.mark.asyncio @@ -2926,6 +2975,8 @@ async def test_tag_budget_check_reads_from_spend_counter(): ) assert exc_info.value.current_cost == 1.5 assert exc_info.value.max_budget == 1.0 + assert exc_info.value.entity_type == "tag" + assert exc_info.value.entity_id == "paid-tag" @pytest.mark.asyncio @@ -2976,6 +3027,8 @@ async def test_team_member_budget_check_reads_from_spend_counter(): proxy_logging_obj=proxy_logging_obj, ) assert exc_info.value.current_cost == 1.5 + assert exc_info.value.entity_type == "team_member" + assert exc_info.value.entity_id == "test-user:test-team" class TestGuardrailModificationCheck: diff --git a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py index a0248963cf1..9ac22086d92 100644 --- a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py +++ b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py @@ -1,3 +1,4 @@ +import asyncio import json import os import sys @@ -4161,6 +4162,99 @@ async def test_auth_path_caches_team_object_under_canonical_team_id_key(): assert cache.get_cache(key=None) is None +@pytest.mark.asyncio +async def test_auth_does_not_rewrite_cached_key_object_back_into_cache(): + """A cache-hit auth must not write the token back into the cache. + + Re-writing on every auth let a replica holding a stale in-memory token + republish it to shared Redis with a fresh TTL on each request, so + /key/update and /key/delete never propagated across replicas or regional + Redis while the key kept calling (stale auth re-cache feedback loop). + Only the DB-load paths (IdentityStore._resolve_key / get_key_object) may + populate the cache. + """ + from fastapi import Request + from starlette.datastructures import URL + + import litellm.proxy.proxy_server as _proxy_server_mod + from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache + from litellm.proxy.proxy_server import hash_token + + api_key = "sk-lit-cached-key-no-rewrite" + hashed_key = hash_token(api_key) + + key_cache = UserApiKeyCache() + stale_token = UserAPIKeyAuth( + api_key=api_key, + token=hashed_key, + metadata={"model_rpm_limit": {"gpt-5.4-mini": 3}}, + last_refreshed_at=1000.0, + ) + await key_cache.async_set_cache( + key=hashed_key, value=stale_token, model_type=UserAPIKeyAuth + ) + + fetch_from_db = AsyncMock( + side_effect=AssertionError("cache-hit auth must not touch the DB") + ) + + proxy_logging_obj = MagicMock() + proxy_logging_obj.internal_usage_cache = MagicMock() + proxy_logging_obj.internal_usage_cache.dual_cache = AsyncMock() + proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None) + + attrs = { + "prisma_client": MagicMock(), + "user_api_key_cache": key_cache, + "proxy_logging_obj": proxy_logging_obj, + "master_key": "sk-test-master", + "general_settings": {}, + "llm_model_list": [], + "llm_router": None, + "open_telemetry_logger": None, + "model_max_budget_limiter": MagicMock(), + "user_custom_auth": None, + "jwt_handler": None, + "litellm_proxy_admin_name": "admin", + } + originals = {a: getattr(_proxy_server_mod, a, None) for a in attrs} + try: + for k, v in attrs.items(): + setattr(_proxy_server_mod, k, v) + request = Request(scope={"type": "http"}) + request._url = URL(url="/chat/completions") + with patch( + "litellm.proxy.auth.resolvers.store._fetch_key_object_from_db_with_reconnect", + fetch_from_db, + ): + result = await _user_api_key_auth_builder( + request=request, + api_key=f"Bearer {api_key}", + azure_api_key_header="", + anthropic_api_key_header=None, + google_ai_studio_api_key_header=None, + azure_apim_header=None, + request_data={}, + ) + pending = [t for t in asyncio.all_tasks() if t is not asyncio.current_task()] + if pending: + await asyncio.wait(pending, timeout=5) + + assert result.token == hashed_key + fetch_from_db.assert_not_called() + + cached_after = await key_cache.async_get_cache( + key=hashed_key, model_type=UserAPIKeyAuth + ) + assert cached_after is not None + assert cached_after.last_refreshed_at == 1000.0 + assert cached_after.metadata == {"model_rpm_limit": {"gpt-5.4-mini": 3}} + finally: + for k, v in originals.items(): + setattr(_proxy_server_mod, k, v) + + class TestCheckKeyModelBudgetWithFallback: """`_check_key_model_budget_with_fallback` must reroute a request to the first configured `budget_fallbacks` entry still within its own budget, diff --git a/tests/test_litellm/proxy/client/cli/autoroute/test_commands.py b/tests/test_litellm/proxy/client/cli/autoroute/test_commands.py index e0b5d3a71af..9efde03e04c 100644 --- a/tests/test_litellm/proxy/client/cli/autoroute/test_commands.py +++ b/tests/test_litellm/proxy/client/cli/autoroute/test_commands.py @@ -69,6 +69,25 @@ class TestUpCommand: assert result.exception is None or isinstance(result.exception, SystemExit) assert "lite autoroute configure" in result.output + def test_refuses_with_actionable_error_when_proxy_runtime_missing(self, monkeypatch, tmp_path): + """`up` launches a real litellm proxy, which the thin `litellm[cli]` install cannot run. + It must fail fast with an actionable message pointing at the proxy install, before it ever + tries to launch the doomed subprocess (which would otherwise die with a bare ImportError).""" + config_path, _log_path, _settings_path, _backup_path, _pid_record_path = _patch_paths(monkeypatch, tmp_path) + config_path.write_text(yaml.safe_dump({"model_list": []})) + monkeypatch.setattr(commands_module, "missing_proxy_runtime_modules", lambda: ("fastapi", "websockets")) + + def _fail_if_launched(*args, **kwargs): + raise AssertionError("launch_proxy must not run when the proxy runtime is missing") + + monkeypatch.setattr(commands_module, "launch_proxy", _fail_if_launched) + + result = self.runner.invoke(up) + + assert result.exit_code != 0 + assert "fastapi, websockets" in result.output + assert "litellm[proxy]" in result.output + def test_refuses_when_pid_record_exists_and_process_still_running(self, monkeypatch, tmp_path): config_path, _log_path, _settings_path, _backup_path, pid_record_path = _patch_paths(monkeypatch, tmp_path) config_path.write_text(yaml.safe_dump({"model_list": []})) diff --git a/tests/test_litellm/proxy/client/cli/autoroute/test_config.py b/tests/test_litellm/proxy/client/cli/autoroute/test_config.py index 9fa01524ef3..f8d82476ef0 100644 --- a/tests/test_litellm/proxy/client/cli/autoroute/test_config.py +++ b/tests/test_litellm/proxy/client/cli/autoroute/test_config.py @@ -3,10 +3,12 @@ from typing import Any, Dict, Tuple import pytest from litellm.proxy.client.cli.commands.autoroute.config import ( + DEFAULT_KEYWORD_TIER_RULES, AutorouteConfig, ConfigGenerationError, DiscoveredModel, HeuristicClassifier, + KeywordTierRule, LLMClassifier, NoSemanticMatching, SemanticMatching, @@ -129,6 +131,31 @@ class TestBuildGeneratedModelList: assert router_config["keyword_tier_rules"] assert "classifier_type" not in router_config + def test_semantic_matching_defaults_emit_builtin_keyword_rules(self): + config = _base_config(semantic_matching=SemanticMatching(embedding_model="text-embedding-3-small")) + autorouter = next(m for m in build_generated_model_list(config) if m["model_name"] == "autorouter") + router_config = autorouter["litellm_params"]["complexity_router_config"] + assert router_config["keyword_tier_rules"] == [ + {"keywords": list(rule.keywords), "tier": rule.tier} for rule in DEFAULT_KEYWORD_TIER_RULES + ] + + def test_semantic_matching_serializes_custom_keyword_rules(self): + config = _base_config( + semantic_matching=SemanticMatching( + embedding_model="text-embedding-3-small", + keyword_tier_rules=( + KeywordTierRule(keywords=("yo", "sup"), tier="SIMPLE"), + KeywordTierRule(keywords=("architect", "design a system"), tier="COMPLEX"), + ), + ) + ) + autorouter = next(m for m in build_generated_model_list(config) if m["model_name"] == "autorouter") + router_config = autorouter["litellm_params"]["complexity_router_config"] + assert router_config["keyword_tier_rules"] == [ + {"keywords": ["yo", "sup"], "tier": "SIMPLE"}, + {"keywords": ["architect", "design a system"], "tier": "COMPLEX"}, + ] + def test_complexity_router_config_reflects_adaptive(self): config = _base_config(adaptive=True) autorouter = next(m for m in build_generated_model_list(config) if m["model_name"] == "autorouter") diff --git a/tests/test_litellm/proxy/client/cli/autoroute/test_process.py b/tests/test_litellm/proxy/client/cli/autoroute/test_process.py index 8e7f355adc1..a4f85ea44ff 100644 --- a/tests/test_litellm/proxy/client/cli/autoroute/test_process.py +++ b/tests/test_litellm/proxy/client/cli/autoroute/test_process.py @@ -14,6 +14,7 @@ from litellm.proxy.client.cli.commands.autoroute.process import ( clear_pid_record, is_running, launch_proxy, + missing_proxy_runtime_modules, poll_liveliness, read_pid_record, write_pid_record, @@ -137,3 +138,21 @@ class TestPollLiveliness: assert "exited early" in str(exc_info.value) assert "crash log line" in str(exc_info.value) + + +class TestMissingProxyRuntimeModules: + def test_flags_absent_modules_only(self, monkeypatch): + """A thin litellm[cli] install lacks the proxy runtime; the missing ones must be reported + (by name, for an actionable error) while modules that are importable are not.""" + monkeypatch.setattr( + process_module, + "_PROXY_RUNTIME_MODULES", + ("os", "litellm_autoroute_definitely_absent_pkg", "socket"), + ) + + assert missing_proxy_runtime_modules() == ("litellm_autoroute_definitely_absent_pkg",) + + def test_empty_when_all_present(self, monkeypatch): + monkeypatch.setattr(process_module, "_PROXY_RUNTIME_MODULES", ("os", "socket")) + + assert missing_proxy_runtime_modules() == () diff --git a/tests/test_litellm/proxy/client/cli/autoroute/test_wizard.py b/tests/test_litellm/proxy/client/cli/autoroute/test_wizard.py index 75431a2c588..2b9240aafc7 100644 --- a/tests/test_litellm/proxy/client/cli/autoroute/test_wizard.py +++ b/tests/test_litellm/proxy/client/cli/autoroute/test_wizard.py @@ -156,7 +156,7 @@ class TestRunConfigureWizardSemanticMatching: tmp_path, CHAT_AND_EMBEDDING_GROUPS, _SIMPLE_TIER_PICKS, - input_str="n\ny\nn\n", + input_str="n\ny\n\n\n\n\nn\n", embedding_pick="text-embedding-3-small", ) @@ -165,6 +165,42 @@ class TestRunConfigureWizardSemanticMatching: assert router_config["semantic_keyword_matching"] is True assert router_config["embedding_model"] == "text-embedding-3-small" + def test_blank_keyword_answers_keep_the_builtin_defaults(self, tmp_path): + result, config_path = _run( + tmp_path, + CHAT_AND_EMBEDDING_GROUPS, + _SIMPLE_TIER_PICKS, + input_str="n\ny\n\n\n\n\nn\n", + embedding_pick="text-embedding-3-small", + ) + + assert result.exit_code == 0, result.output + router_config = _router_config(config_path) + assert router_config["keyword_tier_rules"] == [ + {"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"}, + ] + + def test_custom_keyword_answers_are_recorded_per_tier(self, tmp_path): + result, config_path = _run( + tmp_path, + CHAT_AND_EMBEDDING_GROUPS, + _SIMPLE_TIER_PICKS, + input_str="n\ny\nyo, sup\n\nbuild a service, migrate\nderive, prove rigorously\nn\n", + embedding_pick="text-embedding-3-small", + ) + + assert result.exit_code == 0, result.output + router_config = _router_config(config_path) + assert router_config["keyword_tier_rules"] == [ + {"keywords": ["yo", "sup"], "tier": "SIMPLE"}, + {"keywords": ["explain", "how does"], "tier": "MEDIUM"}, + {"keywords": ["build a service", "migrate"], "tier": "COMPLEX"}, + {"keywords": ["derive", "prove rigorously"], "tier": "REASONING"}, + ] + class TestRunConfigureWizardAdaptive: def test_accepting_adaptive_sets_adaptive_flag(self, tmp_path): diff --git a/tests/test_litellm/proxy/client/cli/test_auth_commands.py b/tests/test_litellm/proxy/client/cli/test_auth_commands.py index e101515d4b1..2fbc9c5c82f 100644 --- a/tests/test_litellm/proxy/client/cli/test_auth_commands.py +++ b/tests/test_litellm/proxy/client/cli/test_auth_commands.py @@ -78,7 +78,7 @@ class TestPollingErrorSurfacing: result = CliRunner().invoke(login, obj=mock_context.obj) assert result.exit_code == 0 - assert "❌ Authentication failed:" in result.output + assert "Authentication failed:" in result.output assert "CLI login session not found or expired." in result.output assert "Authentication timed out" not in result.output @@ -414,7 +414,7 @@ class TestLoginCommand: result = self.runner.invoke(login, obj=mock_context.obj) assert result.exit_code == 0 - assert "✅ Login successful!" in result.output + assert "Login successful!" in result.output assert "Automatically assigned to team: team-1" in result.output # Verify browser was opened with correct URL @@ -456,7 +456,7 @@ class TestLoginCommand: result = self.runner.invoke(login, obj=mock_context.obj) assert result.exit_code == 0 - assert "❌ Authentication timed out" in result.output + assert "Authentication timed out" in result.output def test_login_http_error(self): """Test login with HTTP error""" @@ -476,7 +476,7 @@ class TestLoginCommand: result = self.runner.invoke(login, obj=mock_context.obj) assert result.exit_code == 0 - assert "❌ Authentication timed out" in result.output + assert "Authentication timed out" in result.output def test_login_request_exception(self): """Test login with request exception""" @@ -497,7 +497,7 @@ class TestLoginCommand: result = self.runner.invoke(login, obj=mock_context.obj) assert result.exit_code == 0 - assert "❌ Authentication timed out" in result.output + assert "Authentication timed out" in result.output def test_login_keyboard_interrupt(self): """Test login cancelled by user""" @@ -512,7 +512,7 @@ class TestLoginCommand: result = self.runner.invoke(login, obj=mock_context.obj) assert result.exit_code == 0 - assert "❌ Authentication cancelled by user" in result.output + assert "Authentication cancelled by user" in result.output def test_login_no_api_key_in_response(self): """Test login when response doesn't contain API key""" @@ -536,7 +536,7 @@ class TestLoginCommand: result = self.runner.invoke(login, obj=mock_context.obj) assert result.exit_code == 0 - assert "❌ Authentication timed out" in result.output + assert "Authentication timed out" in result.output def test_login_general_exception(self): """Test login with general exception (not requests exception)""" @@ -551,7 +551,7 @@ class TestLoginCommand: result = self.runner.invoke(login, obj=mock_context.obj) assert result.exit_code == 0 - assert "❌ Authentication failed: Invalid value" in result.output + assert "Authentication failed: Invalid value" in result.output class TestLogoutCommand: @@ -567,7 +567,7 @@ class TestLogoutCommand: result = self.runner.invoke(logout) assert result.exit_code == 0 - assert "✅ Logged out successfully" in result.output + assert "Logged out successfully" in result.output mock_clear.assert_called_once() @@ -591,7 +591,7 @@ class TestWhoamiCommand: result = self.runner.invoke(whoami) assert result.exit_code == 0 - assert "✅ Authenticated" in result.output + assert "Authenticated" in result.output assert "test@example.com" in result.output assert "test-user-123" in result.output assert "admin" in result.output @@ -603,7 +603,7 @@ class TestWhoamiCommand: result = self.runner.invoke(whoami) assert result.exit_code == 0 - assert "❌ Not authenticated" in result.output + assert "Not authenticated" in result.output assert "Run 'lite login'" in result.output def test_whoami_old_token(self): @@ -619,8 +619,8 @@ class TestWhoamiCommand: result = self.runner.invoke(whoami) assert result.exit_code == 0 - assert "✅ Authenticated" in result.output - assert "⚠️ Warning: Token is more than 24 hours old" in result.output + assert "Authenticated" in result.output + assert "Warning: Token is more than 24 hours old" in result.output def test_whoami_missing_fields(self): """Test whoami with token missing some fields""" @@ -633,7 +633,7 @@ class TestWhoamiCommand: result = self.runner.invoke(whoami) assert result.exit_code == 0 - assert "✅ Authenticated" in result.output + assert "Authenticated" in result.output assert "Unknown" in result.output # Should show "Unknown" for missing fields def test_whoami_no_timestamp(self): @@ -655,7 +655,7 @@ class TestWhoamiCommand: result = self.runner.invoke(whoami) assert result.exit_code == 0 - assert "✅ Authenticated" in result.output + assert "Authenticated" in result.output # Should calculate age based on timestamp=0 assert "Token age:" in result.output @@ -714,7 +714,7 @@ class TestCLIKeyRegenerationFlow: result = self.runner.invoke(login, obj=mock_context.obj) assert result.exit_code == 0 - assert "✅ Login successful!" in result.output + assert "Login successful!" in result.output assert "team-beta" in result.output # Ensure we surface the human-readable team alias to the user assert "Beta Team" in result.output @@ -774,7 +774,7 @@ class TestCLIKeyRegenerationFlow: result = self.runner.invoke(login, obj=mock_context.obj) assert result.exit_code == 0 - assert "✅ Login successful!" in result.output + assert "Login successful!" in result.output # Verify browser was opened mock_browser.assert_called_once() diff --git a/tests/test_litellm/proxy/client/cli/test_global_options.py b/tests/test_litellm/proxy/client/cli/test_global_options.py index 53b7e4dbc29..8df763d35c2 100644 --- a/tests/test_litellm/proxy/client/cli/test_global_options.py +++ b/tests/test_litellm/proxy/client/cli/test_global_options.py @@ -1,6 +1,7 @@ # stdlib imports import os import sys +from pathlib import Path from unittest.mock import Mock, patch import pytest @@ -11,6 +12,7 @@ sys.path.insert( ) # Adds the parent directory to the system path +import litellm.proxy.client.cli from litellm._version import version as litellm_version from litellm.proxy.client.cli import cli @@ -36,6 +38,19 @@ def test_cli_version_flag(cli_runner): assert "LiteLLM Proxy Server Version: 1.2.3" in result.output +def test_cli_source_is_ascii_only(): + """Non-ASCII output (emoji, box-drawing chars) raises UnicodeEncodeError on legacy Windows + consoles (cp1252), so the whole CLI package must stay ASCII-only.""" + cli_root = Path(litellm.proxy.client.cli.__file__).parent + offenders = [ + f"{path.relative_to(cli_root)}:{line_number}: {line.strip()}" + for path in sorted(cli_root.rglob("*.py")) + for line_number, line in enumerate(path.read_text(encoding="utf-8").splitlines(), start=1) + if not line.isascii() + ] + assert offenders == [] + + def test_base_url_trailing_slash_normalized(cli_runner): """A trailing slash on --base-url must not produce a double slash (e.g. '//sso/cli/start').""" with ( diff --git a/tests/test_litellm/proxy/client/cli/test_keys_commands.py b/tests/test_litellm/proxy/client/cli/test_keys_commands.py index 2c134f9defb..977aec9f5b7 100644 --- a/tests/test_litellm/proxy/client/cli/test_keys_commands.py +++ b/tests/test_litellm/proxy/client/cli/test_keys_commands.py @@ -262,7 +262,7 @@ def test_keys_import_actual_import_success(mock_keys_client, cli_runner): assert result.exit_code == 0 assert "Found 1 keys in source instance" in result.output - assert "✓ Imported key: import-key-1" in result.output + assert "Imported key: import-key-1" in result.output assert "Successfully imported: 1" in result.output assert "Failed to import: 0" in result.output @@ -481,8 +481,8 @@ def test_keys_import_partial_failure(mock_keys_client, cli_runner): ) assert result.exit_code == 0 # Command completes even with partial failures - assert "✓ Imported key: success-key" in result.output - assert "✗ Failed to import key fail-key" in result.output + assert "Imported key: success-key" in result.output + assert "Failed to import key fail-key" in result.output assert "Successfully imported: 1" in result.output assert "Failed to import: 1" in result.output assert "Total keys processed: 2" in result.output diff --git a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py index c29cdaf4171..4c17c5d3482 100644 --- a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py +++ b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py @@ -13,7 +13,10 @@ from datetime import datetime, timezone from unittest.mock import AsyncMock, MagicMock, call, patch import pytest +from redis.exceptions import DataError +import litellm +from litellm.proxy._types import Litellm_EntityType from litellm.proxy.db.db_spend_update_writer import DBSpendUpdateWriter @@ -1418,7 +1421,6 @@ async def test_batch_database_updates_isolation_on_failure(): org_id="org1", end_user_id="eu1", prisma_client=MagicMock(), - user_api_key_cache=MagicMock(), litellm_proxy_budget_name="budget", payload={"key": "value"}, ) @@ -1818,3 +1820,85 @@ async def test_update_database_does_not_deepcopy_on_request_path(): fake_payload["nested"]["a"] = 999 assert batch_payload["model"] == "gpt-4" assert batch_payload["nested"]["a"] == 1 + + +@pytest.mark.asyncio +async def test_spend_update_path_never_queries_user_cache_with_none_user_id(): + """ + When user_id is None, the spend-update path must not perform a user-cache + lookup at all. With a Redis-backed auth cache (enable_redis_auth_cache), + a lookup with key=None raises redis.exceptions.DataError, which aborted + _update_user_db before any spend updates were enqueued. + + This test fails on the old code twice over: the cache mock records the + forbidden lookup, and the DataError it raises kills the end-user spend + update that must survive. + """ + db_writer = DBSpendUpdateWriter() + + strict_redis_backed_cache = MagicMock() + strict_redis_backed_cache.async_get_cache = AsyncMock( + side_effect=DataError("Invalid input of type: 'NoneType'") + ) + + with ( + patch.object(litellm, "max_budget", 0), + patch("litellm.proxy.proxy_server.disable_spend_logs", True), + patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), + patch("litellm.proxy.proxy_server.user_api_key_cache", strict_redis_backed_cache), + patch("litellm.proxy.proxy_server.litellm_proxy_budget_name", "litellm-proxy-budget"), + patch( + "litellm.proxy.spend_tracking.spend_tracking_utils.get_logging_payload", + return_value={ + "startTime": "2024-01-01T00:00:00", + "endTime": "2024-01-01T00:01:00", + "model": "gpt-4", + "custom_llm_provider": "openai", + "spend": 0.0, + }, + ), + ): + await db_writer.update_database( + token=None, + user_id=None, + end_user_id="end-user-1", + team_id=None, + org_id=None, + kwargs={"model": "gpt-4", "custom_llm_provider": "openai"}, + completion_response=MagicMock(), + start_time=datetime.now(), + end_time=datetime.now(), + response_cost=0.1, + ) + await asyncio.sleep(0) + + strict_redis_backed_cache.async_get_cache.assert_not_called() + + queued = await db_writer.spend_update_queue.flush_all_updates_from_in_memory_queue() + end_user_updates = [u for u in queued if u["entity_type"] == Litellm_EntityType.END_USER] + assert len(end_user_updates) == 1 + assert end_user_updates[0]["entity_id"] == "end-user-1" + assert all(u["entity_type"] != Litellm_EntityType.USER for u in queued) + + +@pytest.mark.asyncio +async def test_update_user_db_enqueues_user_spend_without_cache_dependency(): + """ + _update_user_db needs no cache handle: it enqueues the user spend update + (and the end-user one) purely from the ids it is given. + """ + db_writer = DBSpendUpdateWriter() + + with patch.object(litellm, "max_budget", 0): + await db_writer._update_user_db( + response_cost=0.25, + user_id="user-123", + prisma_client=MagicMock(), + litellm_proxy_budget_name="litellm-proxy-budget", + end_user_id="end-user-9", + ) + + queued = await db_writer.spend_update_queue.flush_all_updates_from_in_memory_queue() + by_type = {u["entity_type"]: u["entity_id"] for u in queued} + assert by_type[Litellm_EntityType.USER] == "user-123" + assert by_type[Litellm_EntityType.END_USER] == "end-user-9" diff --git a/tests/test_litellm/proxy/db/test_query_engine_reaper.py b/tests/test_litellm/proxy/db/test_query_engine_reaper.py new file mode 100644 index 00000000000..efcecb4bc08 --- /dev/null +++ b/tests/test_litellm/proxy/db/test_query_engine_reaper.py @@ -0,0 +1,244 @@ +import os +import signal +import subprocess +import sys +import time +from unittest.mock import MagicMock, patch + +import pytest + +from litellm.proxy.db.query_engine_reaper import ( + REAPER_THREAD_NAME, + _read_comm_and_ppid, + _reaper_loop, + _send_signal, + _try_reap, + list_orphaned_engine_pids, + reap_orphaned_engines, + set_child_subreaper, + start_query_engine_reaper, + terminate_and_reap, + terminate_and_reap_all, +) + + +def _write_stat(proc_root, pid, comm, ppid): + pid_dir = proc_root / str(pid) + pid_dir.mkdir() + (pid_dir / "stat").write_text(f"{pid} ({comm}) S {ppid} {pid} {pid} 0 -1 4194304 100 0 0 0") + + +class TestReadCommAndPpid: + def test_parses_comm_and_ppid(self, tmp_path): + _write_stat(tmp_path, 137, "query-engine-de", 81) + assert _read_comm_and_ppid(137, str(tmp_path)) == ("query-engine-de", 81) + + def test_comm_containing_parens_and_spaces(self, tmp_path): + _write_stat(tmp_path, 42, "weird) (name", 1) + assert _read_comm_and_ppid(42, str(tmp_path)) == ("weird) (name", 1) + + def test_missing_pid_returns_none(self, tmp_path): + assert _read_comm_and_ppid(999, str(tmp_path)) is None + + def test_malformed_stat_returns_none(self, tmp_path): + pid_dir = tmp_path / "55" + pid_dir.mkdir() + (pid_dir / "stat").write_text("garbage with no parens") + assert _read_comm_and_ppid(55, str(tmp_path)) is None + + def test_truncated_fields_after_comm_returns_none(self, tmp_path): + pid_dir = tmp_path / "56" + pid_dir.mkdir() + (pid_dir / "stat").write_text("56 (proc) S") + assert _read_comm_and_ppid(56, str(tmp_path)) is None + + def test_non_numeric_ppid_returns_none(self, tmp_path): + pid_dir = tmp_path / "57" + pid_dir.mkdir() + (pid_dir / "stat").write_text("57 (proc) S notanint 57") + assert _read_comm_and_ppid(57, str(tmp_path)) is None + + +class TestListOrphanedEnginePids: + def test_finds_only_engine_children_of_parent(self, tmp_path): + _write_stat(tmp_path, 137, "query-engine-de", 1) + _write_stat(tmp_path, 138, "query-engine-de", 1) + _write_stat(tmp_path, 260, "python", 1) + _write_stat(tmp_path, 285, "query-engine-de", 260) + (tmp_path / "not-a-pid").mkdir() + + assert sorted(list_orphaned_engine_pids(1, proc_root=str(tmp_path))) == [137, 138] + + def test_no_matches_returns_empty(self, tmp_path): + _write_stat(tmp_path, 260, "python", 1) + assert list_orphaned_engine_pids(1, proc_root=str(tmp_path)) == () + + def test_missing_proc_root_returns_empty(self, tmp_path): + assert list_orphaned_engine_pids(1, proc_root=str(tmp_path / "absent")) == () + + +class TestSetChildSubreaper: + def test_matches_platform_capability(self): + result = set_child_subreaper() + if sys.platform.startswith("linux"): + assert result is True + else: + assert result is False + + +@pytest.mark.skipif(sys.platform == "win32", reason="POSIX signals and waitpid") +class TestSignalHelpers: + def test_try_reap_true_for_non_child_pid(self): + assert _try_reap(1) is True + + def test_send_signal_swallows_missing_pid(self): + child = subprocess.Popen([sys.executable, "-c", "pass"]) + child.wait() + _send_signal(child.pid, signal.SIGTERM) + + +class TestReaperLoop: + def test_survives_scan_failure_and_continues(self): + calls = [] + + def flaky_scan(parent_pid, proc_root="/proc"): + calls.append(parent_pid) + if len(calls) == 1: + raise RuntimeError("scan blew up") + raise KeyboardInterrupt + + with ( + patch( + "litellm.proxy.db.query_engine_reaper.reap_orphaned_engines", + side_effect=flaky_scan, + ), + patch("litellm.proxy.db.query_engine_reaper.time.sleep"), + pytest.raises(KeyboardInterrupt), + ): + _reaper_loop(1234) + + assert calls == [1234, 1234] + + +@pytest.mark.skipif(sys.platform == "win32", reason="POSIX signals and waitpid") +class TestTerminateAndReap: + def test_sigterm_terminates_and_reaps_child(self): + child = subprocess.Popen([sys.executable, "-c", "import time; time.sleep(300)"]) + terminate_and_reap(child.pid, grace_seconds=10.0) + + with pytest.raises((ChildProcessError, OSError)): + os.waitpid(child.pid, os.WNOHANG) + child.returncode = -signal.SIGTERM + + def test_escalates_to_sigkill_when_sigterm_ignored(self): + child = subprocess.Popen( + [ + sys.executable, + "-c", + "import signal, time; signal.signal(signal.SIGTERM, signal.SIG_IGN); time.sleep(300)", + ] + ) + deadline = time.monotonic() + 5 + while time.monotonic() < deadline: + probe = subprocess.run( + [sys.executable, "-c", f"import os, signal; os.kill({child.pid}, 0)"], + capture_output=True, + ) + if probe.returncode == 0: + break + time.sleep(0.05) + time.sleep(0.3) + + terminate_and_reap(child.pid, grace_seconds=0.5) + + with pytest.raises((ChildProcessError, OSError)): + os.waitpid(child.pid, os.WNOHANG) + child.returncode = -signal.SIGKILL + + +class TestReapOrphanedEngines: + def test_terminates_each_orphan(self, tmp_path): + _write_stat(tmp_path, 137, "query-engine-de", 1) + _write_stat(tmp_path, 138, "query-engine-de", 1) + _write_stat(tmp_path, 285, "query-engine-de", 260) + + with patch("litellm.proxy.db.query_engine_reaper.terminate_and_reap_all") as mock_terminate: + acted_on = reap_orphaned_engines(1, proc_root=str(tmp_path)) + + assert sorted(acted_on) == [137, 138] + assert sorted(mock_terminate.call_args.args[0]) == [137, 138] + + def test_no_orphans_no_kills(self, tmp_path): + _write_stat(tmp_path, 285, "query-engine-de", 260) + + with patch("litellm.proxy.db.query_engine_reaper.terminate_and_reap_all") as mock_terminate: + assert reap_orphaned_engines(1, proc_root=str(tmp_path)) == () + + mock_terminate.assert_not_called() + + +@pytest.mark.skipif(sys.platform == "win32", reason="POSIX signals and waitpid") +class TestTerminateAndReapAll: + def test_batch_shares_one_grace_period(self): + children = [ + subprocess.Popen( + [ + sys.executable, + "-c", + "import signal, time; signal.signal(signal.SIGTERM, signal.SIG_IGN); time.sleep(300)", + ] + ) + for _ in range(3) + ] + time.sleep(0.5) + + start = time.monotonic() + terminate_and_reap_all(tuple(child.pid for child in children), grace_seconds=1.0) + elapsed = time.monotonic() - start + + assert elapsed < 3.0 + for child in children: + with pytest.raises((ChildProcessError, OSError)): + os.waitpid(child.pid, os.WNOHANG) + child.returncode = -signal.SIGKILL + + +class TestStartQueryEngineReaper: + def test_noop_on_non_linux(self): + with patch("litellm.proxy.db.query_engine_reaper.sys.platform", "darwin"): + assert start_query_engine_reaper() is None + + def test_starts_daemon_thread_on_linux(self): + with ( + patch("litellm.proxy.db.query_engine_reaper.sys.platform", "linux"), + patch( + "litellm.proxy.db.query_engine_reaper.threading.enumerate", + return_value=[], + ), + patch("litellm.proxy.db.query_engine_reaper.set_child_subreaper") as mock_subreaper, + patch("litellm.proxy.db.query_engine_reaper.threading.Thread") as mock_thread_cls, + ): + thread = start_query_engine_reaper() + + mock_subreaper.assert_called_once() + mock_thread_cls.assert_called_once() + assert mock_thread_cls.call_args.kwargs["daemon"] is True + assert mock_thread_cls.call_args.kwargs["args"] == (os.getpid(),) + mock_thread_cls.return_value.start.assert_called_once() + assert thread is mock_thread_cls.return_value + + def test_second_call_returns_existing_thread(self): + existing = MagicMock() + existing.name = REAPER_THREAD_NAME + with ( + patch("litellm.proxy.db.query_engine_reaper.sys.platform", "linux"), + patch( + "litellm.proxy.db.query_engine_reaper.threading.enumerate", + return_value=[existing], + ), + patch("litellm.proxy.db.query_engine_reaper.threading.Thread") as mock_thread_cls, + ): + thread = start_query_engine_reaper() + + assert thread is existing + mock_thread_cls.assert_not_called() diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_model_armor.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_model_armor.py index 19c200bdaf0..07c40aa763d 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_model_armor.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_model_armor.py @@ -2205,21 +2205,20 @@ async def test_pre_call_file_id_reference_skipped_when_fail_open(): @pytest.mark.asyncio -async def test_pre_call_blocks_when_attachment_count_exceeds_cap(): - """More attachments than the per-request cap fail closed by default to bound scan fan-out.""" - from litellm.proxy.guardrails.guardrail_hooks.model_armor.file_scanning import ( - MAX_FILE_ATTACHMENTS_PER_REQUEST, - ) - - guardrail = _make_guardrail() - pdf_b64 = base64.b64encode(PDF_BYTES).decode("utf-8") - block = { - "type": "file", - "file": {"file_data": f"data:application/pdf;base64,{pdf_b64}"}, - } +async def test_pre_call_file_id_reference_passthrough_when_skip_unscannable_enabled(): + """skip_unscannable_attachments lets a file_id reference through even with fail_on_error=True.""" + guardrail = _make_guardrail(skip_unscannable_attachments=True) request_data = { "model": "gpt-4", - "messages": [{"role": "user", "content": [block] * (MAX_FILE_ATTACHMENTS_PER_REQUEST + 1)}], + "messages": [ + { + "role": "user", + "content": [ + {"type": "text", "text": "summarize this"}, + {"type": "file", "file": {"file_id": "file-abc123"}}, + ], + } + ], "metadata": {"guardrails": ["model-armor-test"]}, } @@ -2227,8 +2226,109 @@ async def test_pre_call_blocks_when_attachment_count_exceeds_cap(): guardrail.async_handler, "post", AsyncMock(return_value=_armor_response(blocked=False)), + ) as mock_post: + await guardrail.async_pre_call_hook( + user_api_key_dict=UserAPIKeyAuth(), + cache=MagicMock(spec=DualCache), + data=request_data, + call_type="completion", + ) + + assert _byte_items_sent(mock_post) == [] + assert _text_payloads_sent(mock_post) == ["summarize this"] + + +@pytest.mark.asyncio +async def test_pre_call_gs_uri_reference_passthrough_when_skip_unscannable_enabled(): + """A gs:// document reference passes through when skip_unscannable_attachments is enabled.""" + guardrail = _make_guardrail(skip_unscannable_attachments=True) + request_data = { + "model": "gpt-4", + "messages": [ + { + "role": "user", + "content": [ + { + "type": "file", + "file": {"file_data": "gs://my-bucket/report.pdf", "filename": "report.pdf"}, + } + ], + } + ], + "metadata": {"guardrails": ["model-armor-test"]}, + } + + with patch.object( + guardrail.async_handler, + "post", + AsyncMock(return_value=_armor_response(blocked=False)), + ) as mock_post: + await guardrail.async_pre_call_hook( + user_api_key_dict=UserAPIKeyAuth(), + cache=MagicMock(spec=DualCache), + data=request_data, + call_type="completion", + ) + + assert _byte_items_sent(mock_post) == [] + + +def test_initialize_guardrail_forwards_skip_unscannable_attachments(): + """skip_unscannable_attachments configured in litellm_params reaches the guardrail instance.""" + from litellm.proxy.guardrails.guardrail_hooks.model_armor import initialize_guardrail + from litellm.types.guardrails import Guardrail, LitellmParams + + litellm_params = LitellmParams( + guardrail="model_armor", + mode="pre_call", + template_id="demo-template", + project_id="demo-project", + skip_unscannable_attachments=True, + ) + guardrail = initialize_guardrail( + litellm_params=litellm_params, + guardrail=Guardrail(guardrail_name="model-armor-config-test"), + ) + + assert guardrail.optional_params.get("skip_unscannable_attachments") is True + + +def test_initialize_guardrail_skip_unscannable_defaults_false(): + """A config that omits skip_unscannable_attachments keeps the secure default (block).""" + from litellm.proxy.guardrails.guardrail_hooks.model_armor import initialize_guardrail + from litellm.types.guardrails import Guardrail, LitellmParams + + litellm_params = LitellmParams( + guardrail="model_armor", + mode="pre_call", + template_id="demo-template", + project_id="demo-project", + ) + guardrail = initialize_guardrail( + litellm_params=litellm_params, + guardrail=Guardrail(guardrail_name="model-armor-config-default"), + ) + + assert guardrail.optional_params.get("skip_unscannable_attachments") is False + + +@pytest.mark.asyncio +async def test_skip_unscannable_still_fails_closed_on_api_error(): + """skip_unscannable_attachments only affects references; a real API error still fails closed.""" + guardrail = _make_guardrail(skip_unscannable_attachments=True, fail_on_error=True) + pdf_b64 = base64.b64encode(PDF_BYTES).decode("utf-8") + request_data = { + "model": "gpt-4", + "messages": [_file_message(pdf_b64)], + "metadata": {"guardrails": ["model-armor-test"]}, + } + + with patch.object( + guardrail.async_handler, + "post", + AsyncMock(side_effect=Exception("model armor upstream 500")), ): - with pytest.raises(HTTPException) as exc_info: + with pytest.raises(Exception) as exc_info: await guardrail.async_pre_call_hook( user_api_key_dict=UserAPIKeyAuth(), cache=MagicMock(spec=DualCache), @@ -2236,8 +2336,35 @@ async def test_pre_call_blocks_when_attachment_count_exceeds_cap(): call_type="completion", ) - assert exc_info.value.status_code == 400 - assert "per-request scan limit" in str(exc_info.value.detail) + assert "model armor upstream 500" in str(exc_info.value) + + +@pytest.mark.asyncio +async def test_pre_call_scans_every_attachment_without_a_count_cap(): + """There is no per-request attachment cap: every scannable attachment is submitted to Model Armor.""" + guardrail = _make_guardrail() + pdf_b64 = base64.b64encode(PDF_BYTES).decode("utf-8") + block = { + "type": "file", + "file": {"file_data": f"data:application/pdf;base64,{pdf_b64}"}, + } + count = 25 + request_data = { + "model": "gpt-4", + "messages": [{"role": "user", "content": [block] * count}], + "metadata": {"guardrails": ["model-armor-test"]}, + } + + mock_post = AsyncMock(return_value=_armor_response(blocked=False)) + with patch.object(guardrail.async_handler, "post", mock_post): + await guardrail.async_pre_call_hook( + user_api_key_dict=UserAPIKeyAuth(), + cache=MagicMock(spec=DualCache), + data=request_data, + call_type="completion", + ) + + assert len(_byte_items_sent(mock_post)) == count @pytest.mark.asyncio diff --git a/tests/test_litellm/proxy/guardrails/test_guardrail_registry.py b/tests/test_litellm/proxy/guardrails/test_guardrail_registry.py index 0ef9ad857f9..26feddadf79 100644 --- a/tests/test_litellm/proxy/guardrails/test_guardrail_registry.py +++ b/tests/test_litellm/proxy/guardrails/test_guardrail_registry.py @@ -44,9 +44,7 @@ def test_update_in_memory_guardrail(): "123", Guardrail( guardrail_name="test-guardrail", - litellm_params=LitellmParams( - guardrail="test-guardrail", mode="pre_call", default_on=True - ), + litellm_params=LitellmParams(guardrail="test-guardrail", mode="pre_call", default_on=True), ), ) @@ -56,10 +54,7 @@ def test_update_in_memory_guardrail(): ) is True ) - assert ( - handler.guardrail_id_to_custom_guardrail["123"].event_hook - is GuardrailEventHooks.pre_call - ) + assert handler.guardrail_id_to_custom_guardrail["123"].event_hook is GuardrailEventHooks.pre_call def _make_guardrail(guardrail_id: str, name: str = "g") -> Guardrail: @@ -135,6 +130,34 @@ def test_delete_in_memory_guardrail_clears_source_marker(): assert handler.get_source("a") is None +def test_list_config_guardrails_excludes_db_sourced(): + """LIT-2529: read surfaces union DB rows with config guardrails; db-sourced + in-memory entries would double-count (or resurrect stale ones), so exclude them.""" + handler = InMemoryGuardrailHandler() + handler.IN_MEMORY_GUARDRAILS["cfg"] = _make_guardrail("cfg", name="config-one") + handler._sources["cfg"] = "config" + handler.IN_MEMORY_GUARDRAILS["db"] = _make_guardrail("db", name="db-one") + handler._sources["db"] = "db" + + config_guardrails = handler.list_config_guardrails() + + assert [g["guardrail_id"] for g in config_guardrails] == ["cfg"] + + +def test_get_config_guardrail_by_id_returns_config_only(): + """LIT-2529: the detail/logs fallback must return config-owned guardrails and + treat a db-sourced (stale) or missing id as a miss.""" + handler = InMemoryGuardrailHandler() + handler.IN_MEMORY_GUARDRAILS["cfg"] = _make_guardrail("cfg", name="config-one") + handler._sources["cfg"] = "config" + handler.IN_MEMORY_GUARDRAILS["db"] = _make_guardrail("db", name="db-one") + handler._sources["db"] = "db" + + assert handler.get_config_guardrail_by_id("cfg")["guardrail_name"] == "config-one" + assert handler.get_config_guardrail_by_id("db") is None + assert handler.get_config_guardrail_by_id("missing") is None + + def test_initialize_guardrail_early_return_updates_source_marker(): """ When initialize_guardrail is called for a guardrail that already exists @@ -152,9 +175,7 @@ def test_initialize_guardrail_early_return_updates_source_marker(): g = Guardrail( guardrail_id="collide", guardrail_name="bedrock", - litellm_params=LitellmParams( - guardrail="bedrock", mode="pre_call", default_on=False - ), + litellm_params=LitellmParams(guardrail="bedrock", mode="pre_call", default_on=False), ) handler.initialize_guardrail(guardrail=g, source="config") @@ -331,10 +352,7 @@ def test_repeated_db_sync_does_not_accumulate_runner_instances(): def distinct_runner_instances() -> int: seen = set() for callback in litellm.logging_callback_manager._get_all_callbacks(): - if ( - isinstance(callback, CustomGuardrail) - and getattr(callback, "guardrail_name", None) == name - ): + if isinstance(callback, CustomGuardrail) and getattr(callback, "guardrail_name", None) == name: seen.add(id(callback)) return len(seen) diff --git a/tests/test_litellm/proxy/guardrails/test_init_guardrails.py b/tests/test_litellm/proxy/guardrails/test_init_guardrails.py index a511229942a..83593c20110 100644 --- a/tests/test_litellm/proxy/guardrails/test_init_guardrails.py +++ b/tests/test_litellm/proxy/guardrails/test_init_guardrails.py @@ -5,9 +5,7 @@ from unittest.mock import MagicMock, patch import pytest -sys.path.insert( - 0, os.path.abspath("../../..") -) # Adds the parent directory to the system path +sys.path.insert(0, os.path.abspath("../../..")) # Adds the parent directory to the system path from litellm.proxy.guardrails.guardrail_registry import InMemoryGuardrailHandler from litellm.types.guardrails import SupportedGuardrailIntegrations @@ -36,8 +34,31 @@ def test_initialize_presidio_guardrail(): ) assert result["guardrail_name"] == "test_presidio_guardrail" - assert ( - result["litellm_params"].guardrail - == SupportedGuardrailIntegrations.PRESIDIO.value - ) + assert result["litellm_params"].guardrail == SupportedGuardrailIntegrations.PRESIDIO.value assert result["litellm_params"].mode == "pre_call" + + +def test_initialize_guardrail_preserves_guardrail_info(): + """ + Regression (LIT-2529): initialize_guardrail must carry guardrail_info into the + stored in-memory Guardrail. Dropping it left the Guardrail Monitor's usage + endpoints unable to render type/description for YAML-defined guardrails. + """ + test_guardrail = { + "guardrail_name": "test_presidio_with_info", + "litellm_params": { + "guardrail": SupportedGuardrailIntegrations.PRESIDIO.value, + "mode": "pre_call", + "presidio_analyzer_api_base": "https://fakelink.com/v1/presidio/analyze", + "presidio_anonymizer_api_base": "https://fakelink.com/v1/presidio/anonymize", + }, + "guardrail_info": {"type": "PII", "description": "masks PII"}, + } + + guardrail_handler = InMemoryGuardrailHandler() + result = guardrail_handler.initialize_guardrail(guardrail=test_guardrail) + + assert result is not None + assert result["guardrail_info"] == {"type": "PII", "description": "masks PII"} + stored = guardrail_handler.IN_MEMORY_GUARDRAILS[result["guardrail_id"]] + assert stored["guardrail_info"] == {"type": "PII", "description": "masks PII"} diff --git a/tests/test_litellm/proxy/guardrails/test_usage_endpoints.py b/tests/test_litellm/proxy/guardrails/test_usage_endpoints.py new file mode 100644 index 00000000000..bf7b1b3b238 --- /dev/null +++ b/tests/test_litellm/proxy/guardrails/test_usage_endpoints.py @@ -0,0 +1,239 @@ +""" +Tests for the /guardrails/usage/* endpoints backing the dashboard Guardrail Monitor. + +Regression (LIT-2529): guardrails defined in config.yaml live only in +IN_MEMORY_GUARDRAIL_HANDLER, so the monitor's overview/detail/logs endpoints — +which read the litellm_guardrailstable Prisma table — could not see them: +detail 404'd, overview omitted them (or rendered them as Custom/Guardrail +orphans), and logs missed their logical-name alias. +""" + +import os +import sys +from datetime import datetime +from typing import Any, Optional +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +sys.path.insert(0, os.path.abspath("../../..")) + +from fastapi import HTTPException + +from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth +from litellm.proxy.guardrails.guardrail_registry import InMemoryGuardrailHandler +from litellm.proxy.guardrails.usage_endpoints import ( + guardrails_usage_detail, + guardrails_usage_logs, + guardrails_usage_overview, +) +from litellm.types.guardrails import Guardrail, LitellmParams + +ADMIN = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN) +# Query() defaults don't resolve to None when the handler is called directly. +START, END = "2026-04-20", "2026-04-27" + + +def _config_handler(*guardrails: Guardrail) -> InMemoryGuardrailHandler: + """A real handler seeded with config-sourced YAML guardrails (no callbacks).""" + handler = InMemoryGuardrailHandler() + for g in guardrails: + gid = g["guardrail_id"] + handler.IN_MEMORY_GUARDRAILS[gid] = g + handler._sources[gid] = "config" + return handler + + +def _yaml_guardrail( + guardrail_id: str = "yaml-1", + name: str = "yaml-pii", + provider: str = "presidio", + info: Optional[dict] = None, +) -> Guardrail: + return Guardrail( + guardrail_id=guardrail_id, + guardrail_name=name, + litellm_params=LitellmParams(guardrail=provider, mode="pre_call"), + guardrail_info=info if info is not None else {"type": "PII", "description": "yaml-defined"}, + ) + + +def _db_row(guardrail_id: str = "db-1", name: str = "db-guard", provider: str = "aim") -> Any: + """A Prisma-style row: attribute access, litellm_params/guardrail_info as plain dicts.""" + row = MagicMock(spec=["guardrail_id", "guardrail_name", "litellm_params", "guardrail_info"]) + row.guardrail_id = guardrail_id + row.guardrail_name = name + row.litellm_params = {"guardrail": provider, "mode": "pre_call"} + row.guardrail_info = {"type": "ContentSafety", "description": "db-defined"} + return row + + +def _metric(guardrail_id: str, date: str = "2026-04-25", requests: int = 10, passed: int = 8, blocked: int = 2) -> Any: + m = MagicMock() + m.guardrail_id = guardrail_id + m.date = date + m.requests_evaluated = requests + m.passed_count = passed + m.blocked_count = blocked + m.flagged_count = 0 + return m + + +def _prisma( + *, + find_many=None, + find_unique=None, + metrics=None, + index_find_many=None, +) -> MagicMock: + client = MagicMock() + db = client.db + db.litellm_guardrailstable.find_many = AsyncMock(return_value=find_many or []) + db.litellm_guardrailstable.find_unique = AsyncMock(return_value=find_unique) + db.litellm_dailyguardrailmetrics.find_many = AsyncMock(return_value=metrics or []) + db.litellm_spendlogguardrailindex.find_many = AsyncMock(return_value=index_find_many or []) + db.litellm_spendlogguardrailindex.count = AsyncMock(return_value=0) + db.litellm_spendlogs.find_many = AsyncMock(return_value=[]) + return client + + +def _patches(prisma: MagicMock, handler: InMemoryGuardrailHandler): + return ( + patch("litellm.proxy.proxy_server.prisma_client", prisma), + patch("litellm.proxy.guardrails.guardrail_registry.IN_MEMORY_GUARDRAIL_HANDLER", handler), + ) + + +# ---- detail ----------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_detail_returns_yaml_guardrail_when_db_misses(): + prisma = _prisma(find_unique=None) + handler = _config_handler(_yaml_guardrail()) + p1, p2 = _patches(prisma, handler) + with p1, p2: + resp = await guardrails_usage_detail( + guardrail_id="yaml-1", start_date=START, end_date=END, user_api_key_dict=ADMIN + ) + assert resp.guardrail_id == "yaml-1" + assert resp.guardrail_name == "yaml-pii" + assert resp.provider == "presidio" # coerced from the LitellmParams pydantic model + assert resp.type == "PII" # from guardrail_info + assert resp.description == "yaml-defined" + + +@pytest.mark.asyncio +async def test_detail_404_when_neither_db_nor_config(): + prisma = _prisma(find_unique=None) + handler = _config_handler() # empty + p1, p2 = _patches(prisma, handler) + with p1, p2, pytest.raises(HTTPException) as exc: + await guardrails_usage_detail(guardrail_id="ghost", start_date=START, end_date=END, user_api_key_dict=ADMIN) + assert exc.value.status_code == 404 + + +@pytest.mark.asyncio +async def test_detail_does_not_surface_db_sourced_in_memory_entry(): + """A stale in-memory entry (source=db, gone from DB) must 404, not resurface.""" + prisma = _prisma(find_unique=None) + handler = InMemoryGuardrailHandler() + stale = _yaml_guardrail(guardrail_id="stale-1", name="stale") + handler.IN_MEMORY_GUARDRAILS["stale-1"] = stale + handler._sources["stale-1"] = "db" + p1, p2 = _patches(prisma, handler) + with p1, p2, pytest.raises(HTTPException) as exc: + await guardrails_usage_detail(guardrail_id="stale-1", start_date=START, end_date=END, user_api_key_dict=ADMIN) + assert exc.value.status_code == 404 + + +@pytest.mark.asyncio +async def test_detail_db_row_still_resolves(): + prisma = _prisma(find_unique=_db_row(guardrail_id="db-1", provider="aim")) + handler = _config_handler() + p1, p2 = _patches(prisma, handler) + with p1, p2: + resp = await guardrails_usage_detail( + guardrail_id="db-1", start_date=START, end_date=END, user_api_key_dict=ADMIN + ) + assert resp.provider == "aim" + assert resp.type == "ContentSafety" + + +# ---- overview --------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_overview_includes_yaml_guardrail_with_no_metrics(): + """The core bug: a YAML guardrail with zero metrics must still appear as a row.""" + prisma = _prisma(find_many=[]) # no DB guardrails + handler = _config_handler(_yaml_guardrail()) + p1, p2 = _patches(prisma, handler) + with p1, p2: + resp = await guardrails_usage_overview(start_date=START, end_date=END, user_api_key_dict=ADMIN) + rows = [r for r in resp.rows if r.id == "yaml-1"] + assert len(rows) == 1 + assert rows[0].name == "yaml-pii" + assert rows[0].provider == "presidio" + assert rows[0].type == "PII" + assert rows[0].requestsEvaluated == 0 + + +@pytest.mark.asyncio +async def test_overview_yaml_metrics_matched_by_logical_name(): + """Daily metrics are keyed by logical name; the YAML row must pick them up.""" + prisma = _prisma( + find_many=[], + metrics=[_metric("yaml-pii", requests=10, blocked=2)], # keyed by name, not uuid + ) + handler = _config_handler(_yaml_guardrail(guardrail_id="yaml-uuid", name="yaml-pii")) + p1, p2 = _patches(prisma, handler) + with p1, p2: + resp = await guardrails_usage_overview(start_date=START, end_date=END, user_api_key_dict=ADMIN) + rows = [r for r in resp.rows if r.id == "yaml-uuid"] + assert len(rows) == 1 + assert rows[0].requestsEvaluated == 10 + assert rows[0].failRate == 20.0 + # must not also emit an orphan row keyed by the logical name + assert [r for r in resp.rows if r.id == "yaml-pii"] == [] + + +@pytest.mark.asyncio +async def test_overview_excludes_db_sourced_in_memory_entry(): + """union must not resurrect a stale db-sourced in-memory guardrail.""" + prisma = _prisma(find_many=[]) + handler = InMemoryGuardrailHandler() + handler.IN_MEMORY_GUARDRAILS["cfg"] = _yaml_guardrail(guardrail_id="cfg", name="cfg-guard") + handler._sources["cfg"] = "config" + handler.IN_MEMORY_GUARDRAILS["stale"] = _yaml_guardrail(guardrail_id="stale", name="stale-guard") + handler._sources["stale"] = "db" + p1, p2 = _patches(prisma, handler) + with p1, p2: + resp = await guardrails_usage_overview(start_date=START, end_date=END, user_api_key_dict=ADMIN) + ids = {r.id for r in resp.rows} + assert "cfg" in ids + assert "stale" not in ids + + +# ---- logs ------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_logs_resolves_config_guardrail_logical_name(): + """The index query must include the YAML guardrail's logical name alias.""" + prisma = _prisma(find_unique=None) + handler = _config_handler(_yaml_guardrail(guardrail_id="yaml-uuid", name="yaml-pii")) + p1, p2 = _patches(prisma, handler) + with p1, p2: + await guardrails_usage_logs( + guardrail_id="yaml-uuid", + policy_id=None, + page=1, + page_size=50, + action=None, + start_date=START, + end_date=END, + user_api_key_dict=ADMIN, + ) + where = prisma.db.litellm_spendlogguardrailindex.find_many.call_args.kwargs["where"] + assert where["guardrail_id"] == {"in": ["yaml-uuid", "yaml-pii"]} diff --git a/tests/test_litellm/proxy/management_endpoints/scim/test_scim_patch_user.py b/tests/test_litellm/proxy/management_endpoints/scim/test_scim_patch_user.py index 2a2bed13bf4..f8995a6f4da 100644 --- a/tests/test_litellm/proxy/management_endpoints/scim/test_scim_patch_user.py +++ b/tests/test_litellm/proxy/management_endpoints/scim/test_scim_patch_user.py @@ -1,9 +1,10 @@ from unittest.mock import AsyncMock, MagicMock, patch import pytest +from fastapi import HTTPException from litellm.proxy._types import LiteLLM_UserTable -from litellm.proxy.management_endpoints.scim.scim_v2 import patch_user +from litellm.proxy.management_endpoints.scim.scim_v2 import _apply_patch_ops, patch_user from litellm.types.proxy.management_endpoints.scim_v2 import ( SCIMPatchOp, SCIMPatchOperation, @@ -329,3 +330,138 @@ async def test_patch_user_multiple_fields_without_path(): assert update_data["user_alias"] == "New Display Name" assert "" not in metadata # Ensure no empty string key assert result.active is False + + +def _user_with_metadata(metadata): + return LiteLLM_UserTable( + user_id="user-mva", + user_email="mva@example.com", + user_alias=None, + teams=[], + metadata=metadata, + ) + + +def test_apply_patch_ops_replace_entitlements_writes_canonical_key(): + """A PATCH on path=entitlements must persist under scim_entitlements, not + fall through to the generic handler's raw path key""" + patch_ops = SCIMPatchOp( + Operations=[ + SCIMPatchOperation( + op="replace", + path="entitlements", + value=[{"value": "jira-software", "display": "Jira Software"}], + ) + ] + ) + + update_data, _ = _apply_patch_ops( + existing_user=_user_with_metadata({}), patch_ops=patch_ops + ) + + metadata = update_data["metadata"] + assert metadata["scim_entitlements"] == [ + {"value": "jira-software", "display": "Jira Software"} + ] + assert "entitlements" not in metadata + + +def test_apply_patch_ops_add_roles_appends_to_existing(): + patch_ops = SCIMPatchOp( + Operations=[ + SCIMPatchOperation(op="add", path="roles", value=[{"value": "admin"}]) + ] + ) + + update_data, _ = _apply_patch_ops( + existing_user=_user_with_metadata({"scim_roles": [{"value": "viewer"}]}), + patch_ops=patch_ops, + ) + + assert update_data["metadata"]["scim_roles"] == [ + {"value": "viewer"}, + {"value": "admin"}, + ] + + +def test_apply_patch_ops_remove_entitlements_clears_canonical_key(): + patch_ops = SCIMPatchOp( + Operations=[SCIMPatchOperation(op="remove", path="entitlements")] + ) + + update_data, _ = _apply_patch_ops( + existing_user=_user_with_metadata( + {"scim_entitlements": [{"value": "jira-software"}]} + ), + patch_ops=patch_ops, + ) + + assert "scim_entitlements" not in update_data["metadata"] + + +def test_apply_patch_ops_pathless_value_dict_handles_roles(): + patch_ops = SCIMPatchOp( + Operations=[ + SCIMPatchOperation( + op="replace", + value={"roles": [{"value": "engineering-admin", "primary": True}]}, + ) + ] + ) + + update_data, _ = _apply_patch_ops( + existing_user=_user_with_metadata({}), patch_ops=patch_ops + ) + + assert update_data["metadata"]["scim_roles"] == [ + {"value": "engineering-admin", "primary": True} + ] + + +def test_apply_patch_ops_invalid_entitlements_value_raises_400(): + patch_ops = SCIMPatchOp( + Operations=[ + SCIMPatchOperation( + op="replace", path="entitlements", value=[{"display": "no value"}] + ) + ] + ) + + with pytest.raises(HTTPException) as exc_info: + _apply_patch_ops(existing_user=_user_with_metadata({}), patch_ops=patch_ops) + + assert exc_info.value.status_code == 400 + + +def test_apply_patch_ops_add_without_value_raises_400_naming_value_member(): + patch_ops = SCIMPatchOp( + Operations=[SCIMPatchOperation(op="add", path="entitlements")] + ) + + with pytest.raises(HTTPException) as exc_info: + _apply_patch_ops(existing_user=_user_with_metadata({}), patch_ops=patch_ops) + + assert exc_info.value.status_code == 400 + assert "value" in str(exc_info.value.detail) + + +def test_apply_patch_ops_filtered_path_raises_400_instead_of_junk_metadata(): + """A filtered path must fail loudly rather than fall through to the generic + handler, which would write a junk metadata key while reporting success""" + patch_ops = SCIMPatchOp( + Operations=[ + SCIMPatchOperation( + op="remove", path='roles[value eq "engineering-admin"]' + ) + ] + ) + + with pytest.raises(HTTPException) as exc_info: + _apply_patch_ops( + existing_user=_user_with_metadata( + {"scim_roles": [{"value": "engineering-admin"}]} + ), + patch_ops=patch_ops, + ) + + assert exc_info.value.status_code == 400 diff --git a/tests/test_litellm/proxy/management_endpoints/scim/test_scim_transformations.py b/tests/test_litellm/proxy/management_endpoints/scim/test_scim_transformations.py index ad0e7010325..458c7c42eb6 100644 --- a/tests/test_litellm/proxy/management_endpoints/scim/test_scim_transformations.py +++ b/tests/test_litellm/proxy/management_endpoints/scim/test_scim_transformations.py @@ -15,6 +15,7 @@ from litellm.proxy.management_endpoints.scim.scim_transformations import ( from litellm.types.proxy.management_endpoints.scim_v2 import ( SCIM_ENTERPRISE_USER_SCHEMA, SCIMEnterpriseUser, + SCIMMultiValuedAttribute, SCIMPatchOperation, SCIMUser, ) @@ -179,6 +180,74 @@ class TestScimTransformations: assert scim_user.enterprise_user.department == "Platform" assert SCIM_ENTERPRISE_USER_SCHEMA in scim_user.schemas + @pytest.mark.asyncio + async def test_transform_user_with_entitlements_and_roles_metadata( + self, mock_prisma_client + ): + mock_client, mock_find_unique = mock_prisma_client + mock_find_unique.return_value = None + + user = LiteLLM_UserTable( + user_id="user-entitled", + user_email="entitled@example.com", + user_alias=None, + teams=[], + created_at=None, + updated_at=None, + metadata={ + "scim_entitlements": [ + {"value": "jira-software", "display": "Jira Software"} + ], + "scim_roles": [{"value": "engineering-admin", "primary": True}], + }, + ) + + with patch("litellm.proxy.proxy_server.prisma_client", mock_client): + scim_user = await ScimTransformations.transform_litellm_user_to_scim_user( + user + ) + + assert scim_user.entitlements is not None + assert scim_user.entitlements[0].value == "jira-software" + assert scim_user.entitlements[0].display == "Jira Software" + assert scim_user.roles is not None + assert scim_user.roles[0].value == "engineering-admin" + assert scim_user.roles[0].primary is True + + @pytest.mark.asyncio + async def test_transform_user_with_malformed_directory_metadata_fails_soft( + self, mock_prisma_client + ): + """Metadata is writable outside the SCIM surface; a corrupted value on one + user must omit the attribute, not fail the whole directory response""" + mock_client, mock_find_unique = mock_prisma_client + mock_find_unique.return_value = None + + user = LiteLLM_UserTable( + user_id="user-corrupt", + user_email="corrupt@example.com", + user_alias=None, + teams=[], + created_at=None, + updated_at=None, + metadata={ + "scim_entitlements": [{"display": 123}], + "scim_roles": {"value": "not-a-list"}, + "scim_enterprise": {"manager": 42}, + }, + ) + + with patch("litellm.proxy.proxy_server.prisma_client", mock_client): + scim_user = await ScimTransformations.transform_litellm_user_to_scim_user( + user + ) + + assert scim_user.id == "user-corrupt" + assert scim_user.entitlements is None + assert scim_user.roles is None + assert scim_user.enterprise_user is None + assert SCIM_ENTERPRISE_USER_SCHEMA not in scim_user.schemas + @pytest.mark.asyncio async def test_transform_user_without_enterprise_metadata_omits_schema( self, mock_user, mock_prisma_client @@ -223,6 +292,27 @@ class TestScimTransformations: dumped_ent = with_enterprise.model_dump(by_alias=True) assert dumped_ent[SCIM_ENTERPRISE_USER_SCHEMA]["costCenter"] == "CC-42" + def test_scim_user_serialization_omits_absent_entitlements_and_roles(self): + without_attrs = SCIMUser( + schemas=["urn:ietf:params:scim:schemas:core:2.0:User"], + id="user-1", + userName="user@example.com", + ) + dumped = without_attrs.model_dump(by_alias=True) + assert "entitlements" not in dumped + assert "roles" not in dumped + + with_attrs = SCIMUser( + schemas=["urn:ietf:params:scim:schemas:core:2.0:User"], + id="user-2", + userName="entitled@example.com", + entitlements=[SCIMMultiValuedAttribute(value="jira-software")], + roles=[SCIMMultiValuedAttribute(value="engineering-admin")], + ) + dumped_attrs = with_attrs.model_dump(by_alias=True) + assert dumped_attrs["entitlements"][0]["value"] == "jira-software" + assert dumped_attrs["roles"][0]["value"] == "engineering-admin" + @pytest.mark.asyncio async def test_transform_litellm_team_to_scim_group( self, mock_team, mock_prisma_client diff --git a/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_endpoints.py b/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_endpoints.py index f39ff93cee7..f27f1197090 100644 --- a/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_endpoints.py @@ -172,6 +172,70 @@ async def test_create_user_ingests_enterprise_extension(mocker, monkeypatch): } +@pytest.mark.asyncio +async def test_create_user_ingests_entitlements_and_roles(mocker, monkeypatch): + """A SCIM create payload carrying entitlements and roles should land in the + created user's metadata under scim_entitlements and scim_roles""" + + scim_user = SCIMUser.model_validate( + { + "schemas": ["urn:ietf:params:scim:schemas:core:2.0:User"], + "userName": "entitled-user", + "name": {"familyName": "User", "givenName": "Entitled"}, + "emails": [{"value": "entitled@example.com"}], + "entitlements": [ + { + "value": "jira-software", + "display": "Jira Software", + "type": "app", + "primary": True, + }, + "bare-entitlement", + ], + "roles": [{"value": "engineering-admin", "type": "role"}], + } + ) + + mock_prisma_client = mocker.MagicMock() + mock_prisma_client.db = mocker.MagicMock() + mock_prisma_client.db.litellm_usertable = mocker.MagicMock() + mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=None) + mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=None) + + monkeypatch.setattr("litellm.default_internal_user_params", None, raising=False) + + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception", + AsyncMock(return_value=mock_prisma_client), + ) + + new_user_mock = mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2.new_user", + AsyncMock(return_value=NewUserRequest(user_id="entitled-user")), + ) + + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2.ScimTransformations.transform_litellm_user_to_scim_user", + AsyncMock(return_value=scim_user), + ) + + await create_user(user=scim_user) + + created_metadata = new_user_mock.call_args.kwargs["data"].metadata + assert created_metadata["scim_entitlements"] == [ + { + "value": "jira-software", + "display": "Jira Software", + "type": "app", + "primary": True, + }, + {"value": "bare-entitlement"}, + ] + assert created_metadata["scim_roles"] == [ + {"value": "engineering-admin", "type": "role"} + ] + + @pytest.mark.asyncio async def test_create_user_uses_default_internal_user_params_role(mocker, monkeypatch): """If role is set in default_internal_user_params, new user should use that role""" diff --git a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py index f7b2df45a85..4936191c344 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py @@ -555,6 +555,85 @@ async def test_new_team_with_mcp_tool_permissions(mock_db_client, mock_admin_aut assert created_permission_data["mcp_servers"] == ["server_a", "server_b"] +@pytest.mark.parametrize( + "user_role,user_id,flag_value,expected", + [ + (LitellmUserRoles.PROXY_ADMIN, "admin-1", True, False), + (LitellmUserRoles.PROXY_ADMIN, "admin-1", False, True), + (LitellmUserRoles.PROXY_ADMIN, "admin-1", None, True), + (LitellmUserRoles.INTERNAL_USER, "user-1", True, True), + (LitellmUserRoles.ORG_ADMIN, "org-admin-1", True, True), + (LitellmUserRoles.PROXY_ADMIN, None, False, False), + ], +) +def test_should_auto_add_team_creator(user_role, user_id, flag_value, expected): + from litellm.proxy.management_endpoints.team_endpoints import ( + _should_auto_add_team_creator, + ) + + general_settings = ( + {} if flag_value is None else {"disable_auto_add_proxy_admin_to_teams": flag_value} + ) + auth = UserAPIKeyAuth(user_role=user_role, user_id=user_id) + assert _should_auto_add_team_creator(auth, general_settings) is expected + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "disable_flag,expect_creator_added", [(True, False), (False, True)] +) +async def test_new_team_disable_auto_add_proxy_admin_flag( + mock_db_client, disable_flag, expect_creator_added +): + """ + When general_settings.disable_auto_add_proxy_admin_to_teams is True, a proxy + admin calling /team/new must NOT be auto-added to the team's members. When + the flag is off, the creator is auto-added as a team admin (default + behavior, regression guard for LIT-3739). + """ + mock_db_client.jsonify_team_object = lambda db_data: db_data + mock_db_client.get_data = AsyncMock(return_value=None) + mock_db_client.update_data = AsyncMock(return_value=MagicMock()) + mock_db_client.db = MagicMock() + + team_create_result = MagicMock(team_id="team-789") + team_create_result.model_dump.return_value = {"team_id": "team-789"} + mock_db_client.db.litellm_teamtable = MagicMock() + mock_db_client.db.litellm_teamtable.create = AsyncMock( + return_value=team_create_result + ) + mock_db_client.db.litellm_teamtable.count = AsyncMock(return_value=0) + mock_db_client.db.litellm_usertable = MagicMock() + mock_db_client.db.litellm_usertable.update = AsyncMock(return_value=MagicMock()) + + from fastapi import Request + + from litellm.proxy._types import NewTeamRequest + from litellm.proxy.management_endpoints.team_endpoints import new_team + + admin_auth = UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin-user-1" + ) + + with patch( + "litellm.proxy.proxy_server.general_settings", + {"disable_auto_add_proxy_admin_to_teams": disable_flag}, + ), patch( + "litellm.proxy.management_endpoints.team_endpoints._add_team_members_to_team", + new_callable=AsyncMock, + ) as mock_add_members: + await new_team( + data=NewTeamRequest(team_alias="flag-test-team"), + http_request=MagicMock(spec=Request), + user_api_key_dict=admin_auth, + ) + + mock_add_members.assert_called_once() + member_add_request = mock_add_members.call_args.kwargs["data"] + member_user_ids = [m.user_id for m in member_add_request.member] + assert ("admin-user-1" in member_user_ids) is expect_creator_added + + @pytest.mark.asyncio async def test_team_update_object_permissions_existing_permission(monkeypatch): """ diff --git a/tests/test_litellm/proxy/test_budget_reservation.py b/tests/test_litellm/proxy/test_budget_reservation.py index 540f017ee88..0b304f2fec7 100644 --- a/tests/test_litellm/proxy/test_budget_reservation.py +++ b/tests/test_litellm/proxy/test_budget_reservation.py @@ -150,8 +150,41 @@ async def test_reservation_blocks_over_budget_non_throttled_key( await _reserve(valid_token, 0.6, key_cache, proxy_logging_obj) await _reserve(valid_token, 0.6, key_cache, proxy_logging_obj) # counter -> 1.0 - with pytest.raises(litellm.BudgetExceededError): + with pytest.raises(litellm.BudgetExceededError) as exc_info: await _reserve(valid_token, 0.6, key_cache, proxy_logging_obj) + assert exc_info.value.entity_type == "key" + assert exc_info.value.entity_id == "key-no-optin-over" + + +@pytest.mark.asyncio +async def test_over_budget_window_counter_tags_clean_entity_id(): + from litellm.proxy.spend_tracking.budget_reservation import ( + _apply_over_budget_reservation_policy, + _BudgetCounter, + ) + + counter = _BudgetCounter( + counter_key="spend:key:test-token:window:1d", + max_budget=1.0, + fallback_spend=0.0, + entity_type="Key", + entity_id="test-token:1d", + spend_log_entity_id="test-token", + ) + + with pytest.raises(litellm.BudgetExceededError) as exc_info: + await _apply_over_budget_reservation_policy( + counter=counter, + valid_token=None, + entry={"counter_key": counter.counter_key}, + applied_entries=[], + reservation_cost=0.5, + current_spend=2.0, + ) + assert exc_info.value.entity_type == "key" + assert exc_info.value.entity_id == "test-token" + assert exc_info.value.max_budget == 1.0 + assert exc_info.value.current_cost == 2.0 def test_should_not_serialize_budget_reservation_on_user_api_key_auth(): diff --git a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py index 913ac116866..d2b8b7ec23d 100644 --- a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py +++ b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py @@ -2688,6 +2688,70 @@ def test_get_sanitized_user_information_from_key_includes_guardrails_metadata(): assert result["user_api_key_auth_metadata"]["other_field"] == "value" +def test_user_and_team_spend_and_budget_flow_to_standard_logging_metadata(): + """ + Full flow: UserAPIKeyAuth -> get_sanitized_user_information_from_key -> + get_standard_logging_metadata. User-level and team-level spend + max budget + must reach the StandardLoggingPayload metadata that custom loggers receive, + alongside the key-level values + """ + from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup + + user_api_key_dict = UserAPIKeyAuth( + api_key="test-key-hash", + spend=1.5, + max_budget=10.0, + user_id="test-user", + user_spend=25.5, + user_max_budget=100.0, + team_id="test-team", + team_spend=250.75, + team_max_budget=1000.0, + ) + + sanitized = LiteLLMProxyRequestSetup.get_sanitized_user_information_from_key( + user_api_key_dict=user_api_key_dict + ) + + assert sanitized["user_api_key_spend"] == 1.5 + assert sanitized["user_api_key_max_budget"] == 10.0 + assert sanitized["user_api_key_user_spend"] == 25.5 + assert sanitized["user_api_key_user_max_budget"] == 100.0 + assert sanitized["user_api_key_team_spend"] == 250.75 + assert sanitized["user_api_key_team_max_budget"] == 1000.0 + + logging_metadata = StandardLoggingPayloadSetup.get_standard_logging_metadata( + dict(sanitized) + ) + + assert logging_metadata["user_api_key_user_spend"] == 25.5 + assert logging_metadata["user_api_key_user_max_budget"] == 100.0 + assert logging_metadata["user_api_key_team_spend"] == 250.75 + assert logging_metadata["user_api_key_team_max_budget"] == 1000.0 + + +def test_user_and_team_spend_and_budget_default_to_none_in_standard_logging_metadata(): + """ + Keys with no user or team level budgets report None for the new fields in the + StandardLoggingPayload metadata instead of raising + """ + from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup + + user_api_key_dict = UserAPIKeyAuth(api_key="test-key-hash") + + sanitized = LiteLLMProxyRequestSetup.get_sanitized_user_information_from_key( + user_api_key_dict=user_api_key_dict + ) + logging_metadata = StandardLoggingPayloadSetup.get_standard_logging_metadata( + dict(sanitized) + ) + + assert logging_metadata["user_api_key_user_spend"] is None + assert logging_metadata["user_api_key_user_max_budget"] is None + assert logging_metadata["user_api_key_team_spend"] is None + assert logging_metadata["user_api_key_team_max_budget"] is None + + @pytest.mark.asyncio async def test_team_guardrails_append_to_key_guardrails(): """ diff --git a/tests/test_litellm/proxy/test_proxy_cli.py b/tests/test_litellm/proxy/test_proxy_cli.py index 88dbec4020f..6b0c0dba40f 100644 --- a/tests/test_litellm/proxy/test_proxy_cli.py +++ b/tests/test_litellm/proxy/test_proxy_cli.py @@ -1,3 +1,4 @@ +import inspect import os import sys from pathlib import Path @@ -15,6 +16,8 @@ sys.path.insert( import builtins import types +import uvicorn + from litellm.proxy.proxy_cli import ProxyInitializationHelpers @@ -135,6 +138,16 @@ class TestProxyInitializationHelpers: ) assert args["timeout_worker_healthcheck"] == 15 + def test_installed_uvicorn_supports_worker_flags(self): + params = inspect.signature(uvicorn.Config.__init__).parameters + assert "timeout_worker_healthcheck" in params + assert "limit_max_requests_jitter" in params + + args = ProxyInitializationHelpers._get_default_unvicorn_init_args( + "localhost", 8000, timeout_worker_healthcheck=30 + ) + assert args["timeout_worker_healthcheck"] == 30 + def test_get_reload_options_no_config_still_watches_env(self): opts = ProxyInitializationHelpers._get_reload_options(None) assert opts["reload"] is True @@ -1557,6 +1570,85 @@ class TestProxyInitializationHelpers: mock_uvicorn_run.assert_called_once() +class TestQueryEngineReaperWiring: + def _invoke_run_server(self, args): + from click.testing import CliRunner + + from litellm.proxy.proxy_cli import run_server + + runner = CliRunner() + clean_env = { + k: v + for k, v in os.environ.items() + if k not in ("DATABASE_URL", "DIRECT_URL") + } + with ( + patch.dict(os.environ, clean_env, clear=True), + patch.dict( + "sys.modules", + { + "proxy_server": MagicMock( + app=MagicMock(), + ProxyConfig=MagicMock(), + KeyManagementSettings=MagicMock(), + save_worker_config=MagicMock(), + ) + }, + ), + patch("uvicorn.run") as mock_uvicorn_run, + patch( + "litellm.proxy.proxy_cli.start_query_engine_reaper" + ) as mock_start_reaper, + patch( + "litellm.proxy.proxy_cli.ProxyInitializationHelpers._get_default_unvicorn_init_args" + ) as mock_get_args, + ): + mock_get_args.return_value = { + "app": "litellm.proxy.proxy_server:app", + "host": "localhost", + "port": 8000, + } + result = runner.invoke(run_server, args) + return result, mock_uvicorn_run, mock_start_reaper + + def test_multi_worker_uvicorn_starts_reaper(self): + result, mock_uvicorn_run, mock_start_reaper = self._invoke_run_server( + ["--local", "--num_workers", "2"] + ) + assert result.exit_code == 0, f"exit_code={result.exit_code}, output={result.output}" + mock_uvicorn_run.assert_called_once() + mock_start_reaper.assert_called_once() + + def test_single_worker_uvicorn_does_not_start_reaper(self): + result, mock_uvicorn_run, mock_start_reaper = self._invoke_run_server( + ["--local", "--num_workers", "1"] + ) + assert result.exit_code == 0, f"exit_code={result.exit_code}, output={result.output}" + mock_uvicorn_run.assert_called_once() + mock_start_reaper.assert_not_called() + + @pytest.mark.skipif(os.name == "nt", reason="gunicorn server path skips Windows") + def test_gunicorn_arbiter_starts_reaper(self): + pytest.importorskip("gunicorn") + + with ( + patch("gunicorn.app.base.BaseApplication.run"), + patch( + "litellm.proxy.proxy_cli.start_query_engine_reaper" + ) as mock_start_reaper, + ): + ProxyInitializationHelpers._run_gunicorn_server( + host="127.0.0.1", + port=4010, + app=MagicMock(), + num_workers=1, + ssl_certfile_path=None, + ssl_keyfile_path=None, + ) + + mock_start_reaper.assert_called_once() + + class TestRunServerDbSetup: """Tests for run_server's prisma setup_database behavior.""" diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 2f0924e9192..54db0c0fd4f 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -2588,6 +2588,84 @@ async def test_load_config_max_budget_env_var_coerced_to_float(tmp_path, monkeyp litellm.max_budget = original_max_budget +@pytest.mark.asyncio +async def test_load_config_default_internal_user_params_max_budget_scientific_notation(tmp_path): + """ + Helm's toYaml renders large floats in scientific notation without a + decimal mantissa (e.g. 1e+09), which PyYAML parses as a string. + load_config must coerce default_internal_user_params.max_budget to + float, otherwise every consumer of the raw dict (/user/new, SSO, + SCIM user creation) passes the string to Prisma, which rejects it + since max_budget must be Float or Null. Keys outside the coercion + (including ones not on DefaultInternalUserParams, like + auto_create_key) must pass through unchanged. + """ + from litellm.proxy.proxy_server import ProxyConfig + + config_file = tmp_path / "config.yaml" + config_file.write_text( + "model_list: []\n" + "litellm_settings:\n" + " default_internal_user_params:\n" + " user_role: internal_user\n" + " max_budget: 1e+09\n" + " budget_duration: 30d\n" + " auto_create_key: false\n" + ) + + original_params = litellm.default_internal_user_params + try: + await ProxyConfig().load_config(router=MagicMock(), config_file_path=str(config_file)) + assert litellm.default_internal_user_params == { + "user_role": "internal_user", + "max_budget": 1000000000.0, + "budget_duration": "30d", + "auto_create_key": False, + } + assert isinstance(litellm.default_internal_user_params["max_budget"], float) + finally: + litellm.default_internal_user_params = original_params + + +@pytest.mark.asyncio +async def test_load_config_default_internal_user_params_without_max_budget(tmp_path): + """ + default_internal_user_params without max_budget (or with an explicit + null) must be stored as-is and not gain a max_budget key. + """ + from litellm.proxy.proxy_server import ProxyConfig + + absent_config_file = tmp_path / "absent_config.yaml" + absent_config_file.write_text( + "model_list: []\n" + "litellm_settings:\n" + " default_internal_user_params:\n" + " user_role: internal_user\n" + ) + + null_config_file = tmp_path / "null_config.yaml" + null_config_file.write_text( + "model_list: []\n" + "litellm_settings:\n" + " default_internal_user_params:\n" + " user_role: internal_user\n" + " max_budget: null\n" + ) + + original_params = litellm.default_internal_user_params + try: + await ProxyConfig().load_config(router=MagicMock(), config_file_path=str(absent_config_file)) + assert litellm.default_internal_user_params == {"user_role": "internal_user"} + + await ProxyConfig().load_config(router=MagicMock(), config_file_path=str(null_config_file)) + assert litellm.default_internal_user_params == { + "user_role": "internal_user", + "max_budget": None, + } + finally: + litellm.default_internal_user_params = original_params + + @pytest.mark.asyncio async def test_load_config_user_url_validation_handles_null_and_string_false(tmp_path, monkeypatch): from litellm.proxy.proxy_server import ProxyConfig @@ -4490,6 +4568,128 @@ async def test_update_cache_pipeline_honors_user_api_key_cache_ttl(): setattr(litellm.proxy.proxy_server, "user_api_key_cache", original_cache) +@pytest.mark.asyncio +async def test_spend_tracking_never_writes_the_auth_object_back(): + """Spend tracking must never write the auth object back into the cache. + + Writing the mutated auth object back after every priced request let a + stale copy be re-published with a fresh TTL: to shared Redis it defeated + /key/update and /key/delete across replicas, and even a local-only write + could race an invalidation and resurrect a revoked key on this worker. + Spend is tracked through the spend:key:* counters, so the auth object is + only ever written by the DB-load paths. + """ + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache + + original_cache = litellm.proxy.proxy_server.user_api_key_cache + cache = UserApiKeyCache() + setattr(litellm.proxy.proxy_server, "user_api_key_cache", cache) + try: + hashed_token = "spend-tracking-no-writeback-token" + await cache.async_set_cache( + key=hashed_token, + value=UserAPIKeyAuth(token=hashed_token, spend=1.0), + model_type=UserAPIKeyAuth, + ) + with ( + patch.object( + cache, "async_set_cache_pipeline", new=AsyncMock() + ) as mock_pipeline, + patch.object(cache, "async_set_cache", new=AsyncMock()) as mock_set, + ): + await litellm.proxy.proxy_server.update_cache( + token=hashed_token, + user_id=None, + end_user_id=None, + team_id=None, + response_cost=5.0, + parent_otel_span=None, + ) + pending = [ + t for t in asyncio.all_tasks() if t is not asyncio.current_task() + ] + if pending: + await asyncio.wait(pending, timeout=5) + + key_pipeline_writes = [ + call + for call in mock_pipeline.call_args_list + if any(k == hashed_token for k, _ in call.kwargs["cache_list"]) + ] + assert key_pipeline_writes == [] + mock_set.assert_not_called() + finally: + setattr(litellm.proxy.proxy_server, "user_api_key_cache", original_cache) + + +@pytest.mark.asyncio +async def test_update_cache_global_proxy_spend_scalar_stays_shared(): + """ + The proxy-wide spend estimate must keep flowing to Redis when the spend + writeback goes per-pod: the global max_budget check reads the + ``{litellm_proxy_admin_name}:spend`` cache entry between authoritative DB + reloads, so keeping it pod-local would let traffic spread across replicas + exceed the proxy budget by roughly a factor of the replica count within a + cache TTL. Sharing this scalar is safe because it carries no limits or + permissions, so it cannot resurrect an invalidated auth blob. + """ + from litellm.caching.caching import DualCache + + admin_name = litellm.proxy.proxy_server.litellm_proxy_admin_name + global_key = "{}:spend".format(admin_name) + + async def fake_get(key, **kwargs): + if key == "user-lit": + return {"user_id": "user-lit", "spend": 1.0} + if key == global_key: + return 10.0 + return None + + original_cache = litellm.proxy.proxy_server.user_api_key_cache + cache = DualCache(default_in_memory_ttl=300) + setattr(litellm.proxy.proxy_server, "user_api_key_cache", cache) + try: + with patch.object( + cache, "async_get_cache", new=AsyncMock(side_effect=fake_get) + ): + with patch.object( + cache, "async_set_cache_pipeline", new=AsyncMock() + ) as mock_set_cache: + await litellm.proxy.proxy_server.update_cache( + token=None, + user_id="user-lit", + end_user_id=None, + team_id=None, + response_cost=5.0, + parent_otel_span=None, + ) + + pending = [ + t for t in asyncio.all_tasks() if t is not asyncio.current_task() + ] + if pending: + await asyncio.wait(pending, timeout=5) + + calls = mock_set_cache.await_args_list + local_keys = [ + k + for c in calls + if c.kwargs.get("local_only") is True + for k, _ in c.kwargs["cache_list"] + ] + shared_keys = [ + k + for c in calls + if c.kwargs.get("local_only") is not True + for k, _ in c.kwargs["cache_list"] + ] + assert "user-lit" in local_keys + assert global_key not in local_keys + assert shared_keys == [global_key] + finally: + setattr(litellm.proxy.proxy_server, "user_api_key_cache", original_cache) + + @pytest.mark.asyncio async def test_init_sso_settings_in_db(): """ @@ -6187,6 +6387,32 @@ async def test_update_general_settings_store_model_in_db_false(): assert ps.general_settings["store_model_in_db"] is False +@pytest.mark.asyncio +@pytest.mark.parametrize( + "db_value,expected", + [(True, True), (False, False), ("true", True), ("false", False), (None, None)], +) +async def test_update_general_settings_disable_auto_add_proxy_admin_to_teams(db_value, expected): + """ + Verify _update_general_settings propagates disable_auto_add_proxy_admin_to_teams + from the DB config into the live general_settings dict, so a UI toggle via + /config/field/update takes effect on the next config poll instead of + requiring a proxy restart. + """ + from litellm.proxy.proxy_server import ProxyConfig + + proxy_config = ProxyConfig() + + with patch("litellm.proxy.proxy_server.general_settings", {}): + await proxy_config._update_general_settings( + db_general_settings={"disable_auto_add_proxy_admin_to_teams": db_value} + ) + + import litellm.proxy.proxy_server as ps + + assert ps.general_settings["disable_auto_add_proxy_admin_to_teams"] is expected + + @pytest.mark.asyncio async def test_update_general_settings_store_model_in_db_string_normalization(): """ diff --git a/tests/test_litellm/proxy/test_proxy_types.py b/tests/test_litellm/proxy/test_proxy_types.py index bc77a9ba3c0..c8e0b3a730a 100644 --- a/tests/test_litellm/proxy/test_proxy_types.py +++ b/tests/test_litellm/proxy/test_proxy_types.py @@ -124,3 +124,18 @@ def test_proxy_exception_str_returns_message(): "param": "key", "code": "401", } + + +def test_key_request_router_settings_keeps_enable_tag_filtering(): + """``router_settings`` on key requests validates through + ``UpdateRouterConfig``; a field missing from that model is silently + dropped at parse time, so a key's "Enable Tag Filtering" toggle would + never reach the DB even though the team path (plain dict) kept it.""" + from litellm.proxy._types import GenerateKeyRequest + + req = GenerateKeyRequest(router_settings={"enable_tag_filtering": True, "num_retries": 2}) + + assert req.router_settings is not None + dumped = req.router_settings.model_dump(exclude_none=True) + assert dumped["enable_tag_filtering"] is True + assert dumped["num_retries"] == 2 diff --git a/tests/test_litellm/proxy/test_route_llm_request.py b/tests/test_litellm/proxy/test_route_llm_request.py index 303871e3981..f506b9665a6 100644 --- a/tests/test_litellm/proxy/test_route_llm_request.py +++ b/tests/test_litellm/proxy/test_route_llm_request.py @@ -382,8 +382,8 @@ async def test_route_request_with_router_settings_override(): "num_retries": 5, "timeout": 30, "model_group_retry_policy": {"gpt-3.5-turbo": {"RateLimitErrorRetries": 3}}, - # These settings should be ignored (not in per_request_settings list) "routing_strategy": "least-busy", + # This setting should be ignored (not in per_request_settings list) "model_group_alias": {"alias": "real_model"}, }, } @@ -400,8 +400,8 @@ async def test_route_request_with_router_settings_override(): assert call_kwargs["num_retries"] == 5 assert call_kwargs["timeout"] == 30 assert call_kwargs["model_group_retry_policy"] == {"gpt-3.5-turbo": {"RateLimitErrorRetries": 3}} + assert call_kwargs["routing_strategy"] == "least-busy" # Verify unsupported settings were NOT merged - assert "routing_strategy" not in call_kwargs assert "model_group_alias" not in call_kwargs # Verify router_settings_override was removed from data assert "router_settings_override" not in call_kwargs @@ -819,3 +819,71 @@ async def test_route_request_realtime_transcription_session_resolves_credentials ) assert mock_handler.call_args.kwargs["api_key"] == "transcription-key" + + +@pytest.mark.asyncio +async def test_route_request_merges_enable_tag_filtering_from_override(): + """Key/team router_settings carry enable_tag_filtering; the override + whitelist must forward it to the router call or the team's tag-routing + toggle saved in the UI is silently ignored at request time.""" + data = { + "model": "gpt-3.5-turbo", + "messages": [{"role": "user", "content": "Hello"}], + "router_settings_override": { + "enable_tag_filtering": True, + }, + } + + llm_router = MagicMock() + llm_router.acompletion.return_value = "success" + + response = await route_request(data, llm_router, None, "acompletion") + + assert response == "success" + call_kwargs = llm_router.acompletion.call_args[1] + assert call_kwargs["enable_tag_filtering"] is True + + +@pytest.mark.asyncio +async def test_route_request_strips_client_supplied_enable_tag_filtering(): + """enable_tag_filtering influences deployment selection and is only + trusted when it comes from key/team router_settings via + router_settings_override. A caller putting it in the request body must + not reach the router with it.""" + data = { + "model": "gpt-3.5-turbo", + "messages": [{"role": "user", "content": "Hello"}], + "enable_tag_filtering": True, + } + + llm_router = MagicMock() + llm_router.acompletion.return_value = "ok" + + await route_request(data, llm_router, None, "acompletion") + + call_kwargs = llm_router.acompletion.call_args[1] + assert "enable_tag_filtering" not in call_kwargs + assert "enable_tag_filtering" not in data + + +@pytest.mark.asyncio +async def test_route_request_override_enable_tag_filtering_beats_body_value(): + """A client-sent enable_tag_filtering must not shadow the key/team + setting: the body copy is stripped first, so the override value is the + one the router sees.""" + data = { + "model": "gpt-3.5-turbo", + "messages": [{"role": "user", "content": "Hello"}], + "enable_tag_filtering": False, + "router_settings_override": { + "enable_tag_filtering": True, + }, + } + + llm_router = MagicMock() + llm_router.acompletion.return_value = "ok" + + await route_request(data, llm_router, None, "acompletion") + + call_kwargs = llm_router.acompletion.call_args[1] + assert call_kwargs["enable_tag_filtering"] is True diff --git a/tests/test_litellm/responses/test_streaming_iterator.py b/tests/test_litellm/responses/test_streaming_iterator.py index 3e4b07fce18..5b0f40fdf27 100644 --- a/tests/test_litellm/responses/test_streaming_iterator.py +++ b/tests/test_litellm/responses/test_streaming_iterator.py @@ -5,13 +5,18 @@ completion_start_time = end_time.""" import json from datetime import datetime +from typing import Optional from unittest.mock import Mock +import httpx import pytest from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig -from litellm.responses.streaming_iterator import ResponsesAPIStreamingIterator +from litellm.responses.streaming_iterator import ( + ResponsesAPIStreamingIterator, + SyncResponsesAPIStreamingIterator, +) from litellm.types.llms.openai import ( ResponseCompletedEvent, ResponsesAPIResponse, @@ -23,19 +28,7 @@ def _sse_event(payload: dict) -> bytes: return f"data: {json.dumps(payload)}\n\n".encode("utf-8") -def _make_iterator( - *, - sse_events: list[bytes], - logging_obj: LiteLLMLoggingObj, -) -> ResponsesAPIStreamingIterator: - async def aiter_bytes(): - for evt in sse_events: - yield evt - - mock_response = Mock() - mock_response.headers = {} - mock_response.aiter_bytes = aiter_bytes - +def _mock_config() -> Mock: mock_config = Mock(spec=BaseResponsesAPIConfig) mock_responses_api_response = Mock(spec=ResponsesAPIResponse) mock_responses_api_response.id = "resp_ttft" @@ -52,17 +45,68 @@ def _make_iterator( return stub mock_config.transform_streaming_response.side_effect = _transform + return mock_config + + +def _make_iterator( + *, + sse_events: list[bytes], + logging_obj: LiteLLMLoggingObj, + trailing_error: Optional[Exception] = None, +) -> ResponsesAPIStreamingIterator: + async def aiter_bytes(): + for evt in sse_events: + yield evt + if trailing_error is not None: + raise trailing_error + + mock_response = Mock() + mock_response.headers = {} + mock_response.aiter_bytes = aiter_bytes return ResponsesAPIStreamingIterator( response=mock_response, model="gpt-4o-mini", - responses_api_provider_config=mock_config, + responses_api_provider_config=_mock_config(), logging_obj=logging_obj, litellm_metadata={}, custom_llm_provider="openai", ) +def _make_sync_iterator( + *, + sse_events: list[bytes], + logging_obj: LiteLLMLoggingObj, + trailing_error: Optional[Exception] = None, +) -> SyncResponsesAPIStreamingIterator: + def iter_bytes(): + for evt in sse_events: + yield evt + if trailing_error is not None: + raise trailing_error + + mock_response = Mock() + mock_response.headers = {} + mock_response.iter_bytes = iter_bytes + + return SyncResponsesAPIStreamingIterator( + response=mock_response, + model="gpt-4o-mini", + responses_api_provider_config=_mock_config(), + logging_obj=logging_obj, + litellm_metadata={}, + custom_llm_provider="openai", + ) + + +def _logging_obj_stub() -> Mock: + logging_obj = Mock(spec=LiteLLMLoggingObj) + logging_obj.completion_start_time = None + logging_obj.model_call_details = {"litellm_params": {}} + return logging_obj + + @pytest.mark.asyncio async def test_responses_streaming_stamps_completion_start_time_on_first_chunk(): """Without the fix, `logging_obj.completion_start_time` stays None across the @@ -122,3 +166,72 @@ async def test_responses_streaming_does_not_reset_prior_completion_start_time(): logging_obj._update_completion_start_time.assert_not_called() assert logging_obj.completion_start_time == prior + + +_COMPLETE_STREAM_EVENTS = [ + _sse_event({"type": "response.created"}), + _sse_event({"type": "response.output_text.delta", "delta": "hi"}), + _sse_event({"type": "response.completed"}), +] + +_TRAILING_ERRORS = [ + httpx.ReadError("Response payload is not completed"), + httpx.RemoteProtocolError("peer closed connection without sending complete message body"), +] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("trailing_error", _TRAILING_ERRORS, ids=type) +async def test_transport_error_after_completed_event_ends_stream_cleanly(trailing_error): + """A sloppy connection close after `response.completed` must not turn a + complete stream into an error (regression guard for the transport no longer + swallowing ClientPayloadError/TransferEncodingError).""" + iterator = _make_iterator( + sse_events=_COMPLETE_STREAM_EVENTS, + logging_obj=_logging_obj_stub(), + trailing_error=trailing_error, + ) + + seen = [event.type async for event in iterator] + + assert ResponsesAPIStreamEvents.RESPONSE_COMPLETED in seen + + +@pytest.mark.asyncio +async def test_transport_error_before_completed_event_raises(): + """A connection lost before any terminal event is a real failure and must + surface, not end the stream as if it completed.""" + iterator = _make_iterator( + sse_events=_COMPLETE_STREAM_EVENTS[:-1], + logging_obj=_logging_obj_stub(), + trailing_error=httpx.ReadError("Response payload is not completed"), + ) + + with pytest.raises(httpx.ReadError): + async for _ in iterator: + pass + + +@pytest.mark.parametrize("trailing_error", _TRAILING_ERRORS, ids=type) +def test_sync_transport_error_after_completed_event_ends_stream_cleanly(trailing_error): + iterator = _make_sync_iterator( + sse_events=_COMPLETE_STREAM_EVENTS, + logging_obj=_logging_obj_stub(), + trailing_error=trailing_error, + ) + + seen = [event.type for event in iterator] + + assert ResponsesAPIStreamEvents.RESPONSE_COMPLETED in seen + + +def test_sync_transport_error_before_completed_event_raises(): + iterator = _make_sync_iterator( + sse_events=_COMPLETE_STREAM_EVENTS[:-1], + logging_obj=_logging_obj_stub(), + trailing_error=httpx.ReadError("Response payload is not completed"), + ) + + with pytest.raises(httpx.ReadError): + for _ in iterator: + pass diff --git a/tests/test_litellm/router_strategy/test_router_routing_groups.py b/tests/test_litellm/router_strategy/test_router_routing_groups.py index 7368c72836d..b8dcdacd8a3 100644 --- a/tests/test_litellm/router_strategy/test_router_routing_groups.py +++ b/tests/test_litellm/router_strategy/test_router_routing_groups.py @@ -629,3 +629,100 @@ async def test_async_dispatch_falls_back_to_sync_for_usage_based_routing_v1(): ) assert v1_spy.called, "async dispatch must route v1 strategy through sync method" + + +def test_request_routing_strategy_override_beats_top_level(): + router = _build_router(routing_strategy="least-busy") + strategy, selector = router._get_routing_context( + "other-model", {"routing_strategy": "simple-shuffle"} + ) + assert strategy == "simple-shuffle" + assert selector is None + + +def test_request_routing_strategy_override_beats_explicit_group(): + router = _build_router( + routing_strategy="simple-shuffle", + routing_groups=[ + { + "group_name": "fast", + "models": ["filtered-model"], + "routing_strategy": "latency-based-routing", + } + ], + ) + strategy, _ = router._get_routing_context( + "filtered-model", {"routing_strategy": "least-busy"} + ) + assert strategy == "least-busy" + + +def test_request_routing_strategy_override_builds_and_caches_selector(): + router = _build_router(routing_strategy="simple-shuffle") + strategy, selector = router._get_routing_context( + "other-model", {"routing_strategy": "latency-based-routing"} + ) + assert strategy == "latency-based-routing" + assert selector is not None + _, selector_again = router._get_routing_context( + "other-model", {"routing_strategy": "latency-based-routing"} + ) + assert selector_again is selector + + +def test_request_routing_strategy_override_matching_global_reuses_default_selector(): + router = _build_router(routing_strategy="least-busy") + _, selector = router._get_routing_context( + "other-model", {"routing_strategy": "least-busy"} + ) + assert selector is router.leastbusy_logger + assert router._override_selectors == {} + + +def test_invalid_request_routing_strategy_override_falls_back(): + router = _build_router(routing_strategy="least-busy") + strategy, selector = router._get_routing_context( + "other-model", {"routing_strategy": "not-a-real-strategy"} + ) + assert strategy == "least-busy" + assert selector is router.leastbusy_logger + + +def test_no_override_key_keeps_existing_behavior(): + router = _build_router(routing_strategy="least-busy") + strategy, _ = router._get_routing_context("other-model", {"messages": []}) + assert strategy == "least-busy" + strategy_none_kwargs, _ = router._get_routing_context("other-model", None) + assert strategy_none_kwargs == "least-busy" + + +def test_request_routing_strategy_override_helper_validates_directly(): + router = _build_router(routing_strategy="least-busy") + assert router._get_request_routing_strategy_override({"routing_strategy": "simple-shuffle"}) == "simple-shuffle" + assert router._get_request_routing_strategy_override({"routing_strategy": RoutingStrategy.LEAST_BUSY}) == "least-busy" + assert router._get_request_routing_strategy_override({"routing_strategy": "lar1"}) is None + assert router._get_request_routing_strategy_override({"routing_strategy": {"bad": "type"}}) is None + assert router._get_request_routing_strategy_override({}) is None + assert router._get_request_routing_strategy_override(None) is None + + +def test_override_strategy_selector_helper_builds_per_strategy(): + router = _build_router(routing_strategy="least-busy") + latency_selector = router._get_override_strategy_selector("latency-based-routing") + assert latency_selector is not None + assert router._get_override_strategy_selector("latency-based-routing") is latency_selector + assert router._get_override_strategy_selector("least-busy") is router.leastbusy_logger + assert router._get_override_strategy_selector("simple-shuffle") is None + + +def test_strategy_reinit_unregisters_override_selectors(): + router = _build_router(routing_strategy="least-busy") + override_selector = router._get_override_strategy_selector("latency-based-routing") + assert override_selector is not None + assert any(id(cb) == id(override_selector) for cb in litellm.callbacks) + + router.update_settings(routing_strategy="latency-based-routing") + + assert router._override_selectors == {} + assert not any(id(cb) == id(override_selector) for cb in litellm.callbacks) + assert router._get_override_strategy_selector("latency-based-routing") is router.lowestlatency_logger diff --git a/tests/test_litellm/router_strategy/test_router_tag_routing.py b/tests/test_litellm/router_strategy/test_router_tag_routing.py index eb289095c51..98506aad594 100644 --- a/tests/test_litellm/router_strategy/test_router_tag_routing.py +++ b/tests/test_litellm/router_strategy/test_router_tag_routing.py @@ -1019,3 +1019,99 @@ async def test_negation_removes_tag_regex_deployment_falls_to_ban_only(): mock_response="hi", ) assert response._hidden_params["model_id"] == "openai-deployment" + + +@pytest.mark.asyncio() +async def test_request_level_enable_tag_filtering_applies_when_global_off(): + """ + A request carrying enable_tag_filtering=True (set by the proxy from key/team + router_settings) must activate tag filtering even when the router-level flag + is off. Without this, a team's "Enable Tag Filtering" toggle saved in the UI + is silently ignored at request time. + """ + router = litellm.Router( + model_list=[ + { + "model_name": "gpt-4", + "litellm_params": { + "model": "gpt-4o", + "api_base": "https://exampleopenaiendpoint-production.up.railway.app/", + "tags": ["teamA"], + }, + "model_info": {"id": "team-a-deployment"}, + }, + { + "model_name": "gpt-4", + "litellm_params": { + "model": "gpt-4o-mini", + "api_base": "https://exampleopenaiendpoint-production.up.railway.app/", + "tags": ["teamB"], + }, + "model_info": {"id": "team-b-deployment"}, + }, + ], + enable_tag_filtering=False, + ) + + for _ in range(5): + response = await router.acompletion( + model="gpt-4", + messages=[{"role": "user", "content": "hi"}], + metadata={"tags": ["teamA"]}, + enable_tag_filtering=True, + mock_response="hi", + ) + assert response._hidden_params["model_id"] == "team-a-deployment" + + for _ in range(5): + response = await router.acompletion( + model="gpt-4", + messages=[{"role": "user", "content": "hi"}], + metadata={"tags": ["teamB"]}, + enable_tag_filtering=True, + mock_response="hi", + ) + assert response._hidden_params["model_id"] == "team-b-deployment" + + +@pytest.mark.asyncio() +async def test_request_level_enable_tag_filtering_false_cannot_disable_global(): + """ + A request-level enable_tag_filtering=False must not bypass a router-level + True: tag filtering can be an operator-level restriction on which + deployments a caller may reach, so per-request settings may only scope + down, never escape the global policy. + """ + router = litellm.Router( + model_list=[ + { + "model_name": "gpt-4", + "litellm_params": { + "model": "gpt-4o", + "api_base": "https://exampleopenaiendpoint-production.up.railway.app/", + "tags": ["teamA"], + }, + "model_info": {"id": "team-a-deployment"}, + }, + { + "model_name": "gpt-4", + "litellm_params": { + "model": "gpt-4o-mini", + "api_base": "https://exampleopenaiendpoint-production.up.railway.app/", + "tags": ["teamB"], + }, + "model_info": {"id": "team-b-deployment"}, + }, + ], + enable_tag_filtering=True, + ) + + for _ in range(5): + response = await router.acompletion( + model="gpt-4", + messages=[{"role": "user", "content": "hi"}], + metadata={"tags": ["teamA"]}, + enable_tag_filtering=False, + mock_response="hi", + ) + assert response._hidden_params["model_id"] == "team-a-deployment" diff --git a/tests/test_litellm/test_redis.py b/tests/test_litellm/test_redis.py index 4b9f13c340b..0818237655d 100644 --- a/tests/test_litellm/test_redis.py +++ b/tests/test_litellm/test_redis.py @@ -724,3 +724,16 @@ def test_connection_pool_without_ssl_kwarg_uses_plain_connection(monkeypatch): call_kwargs = mock_pool.call_args.kwargs assert call_kwargs.get("connection_class") is not async_redis.SSLConnection assert "ssl" not in call_kwargs + + +def test_connection_pool_env_redis_ssl_false_uses_plain_connection(monkeypatch): + """REDIS_SSL=false from the environment must not select SSLConnection.""" + monkeypatch.delenv("REDIS_URL", raising=False) + monkeypatch.delenv("REDIS_CLUSTER_NODES", raising=False) + monkeypatch.setenv("REDIS_SSL", "false") + + pool = get_redis_connection_pool(host="plain-host", port=6379) + + assert pool is not None + assert pool.connection_class is async_redis.Connection + assert "ssl" not in pool.connection_kwargs diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index 9c4d83ff7ea..c2c98c8869c 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -1455,6 +1455,132 @@ def test_model_group_info_cost_none_when_db_model_info_has_no_cost(): assert result.output_cost_per_token is None +@pytest.mark.parametrize( + "value,expected", + [ + ("1e-05", 1e-05), + ("0.00001", 1e-05), + (1e-05, 1e-05), + (5, 5.0), + (None, None), + ("not-a-number", None), + ], +) +def test_cost_value_as_float(value, expected): + from litellm.router import _cost_value_as_float + + assert _cost_value_as_float(value) == expected + + +def test_model_group_info_with_stringified_cost_values(): + """ + YAML 1.2 parsers emit '1e-05' (integer mantissa) as a string, so cost + values in deployment model_info can arrive as str. Aggregating the model + group must not raise TypeError('>' between str and float) and must return + float costs. + """ + router = litellm.Router( + model_list=[ + { + "model_name": "my-custom-model", + "litellm_params": { + "model": "openai/my-custom-backend-1", + "api_key": "fake", + }, + "model_info": { + "input_cost_per_token": "1e-05", + "output_cost_per_token": "1e-05", + }, + }, + { + "model_name": "my-custom-model", + "litellm_params": { + "model": "openai/my-custom-backend-2", + "api_key": "fake", + }, + "model_info": { + "input_cost_per_token": "2e-05", + "output_cost_per_token": "2e-05", + }, + }, + ] + ) + + def _model_info_with_str_costs(model_id: str, model_name: str): + for model in router.model_list: + if model["model_info"]["id"] == model_id: + return { + "key": model_name, + "input_cost_per_token": model["model_info"]["input_cost_per_token"], + "output_cost_per_token": model["model_info"]["output_cost_per_token"], + "litellm_provider": "openai", + "mode": "chat", + } + return None + + with patch.object( + router, "get_deployment_model_info", side_effect=_model_info_with_str_costs + ): + result = router._set_model_group_info( + model_group="my-custom-model", + user_facing_model_group_name="my-custom-model", + ) + + assert result is not None + assert result.input_cost_per_token == 2e-05 + assert result.output_cost_per_token == 2e-05 + assert isinstance(result.input_cost_per_token, float) + assert isinstance(result.output_cost_per_token, float) + + +def test_model_group_info_db_fallback_with_stringified_cost_values(): + """ + Fallback path: when get_deployment_model_info returns nothing, costs are + read straight from the deployment's model_info dict, which can hold + stringified floats parsed from YAML. They must be coerced to float. + """ + router = litellm.Router( + model_list=[ + { + "model_name": "my-custom-model", + "litellm_params": { + "model": "openai/my-custom-backend-1", + "api_key": "fake", + }, + "model_info": { + "input_cost_per_token": "1e-05", + "output_cost_per_token": "3e-05", + }, + }, + { + "model_name": "my-custom-model", + "litellm_params": { + "model": "openai/my-custom-backend-2", + "api_key": "fake", + }, + "model_info": { + "input_cost_per_token": "2e-05", + "output_cost_per_token": "2e-05", + }, + }, + ] + ) + + with patch.object( + router, "get_deployment_model_info", side_effect=Exception("not found") + ): + result = router._set_model_group_info( + model_group="my-custom-model", + user_facing_model_group_name="my-custom-model", + ) + + assert result is not None + assert result.input_cost_per_token == 2e-05 + assert result.output_cost_per_token == 3e-05 + assert isinstance(result.input_cost_per_token, float) + assert isinstance(result.output_cost_per_token, float) + + def test_get_model_access_groups_caching(): """ Test that get_model_access_groups caches the no-args result diff --git a/ui/litellm-dashboard/eslint-suppressions.json b/ui/litellm-dashboard/eslint-suppressions.json index 9b55bf6ca0d..cad2874c1e6 100644 --- a/ui/litellm-dashboard/eslint-suppressions.json +++ b/ui/litellm-dashboard/eslint-suppressions.json @@ -950,19 +950,6 @@ "count": 1 } }, - "src/app/(dashboard)/policies/_components/policy_table.test.tsx": { - "react/display-name": { - "count": 1 - } - }, - "src/app/(dashboard)/policies/_components/policy_table.tsx": { - "no-nested-ternary": { - "count": 2 - }, - "no-restricted-imports": { - "count": 1 - } - }, "src/app/(dashboard)/policies/_components/policy_test_panel.tsx": { "no-restricted-imports": { "count": 1 @@ -997,11 +984,6 @@ "count": 2 } }, - "src/app/(dashboard)/projects/_components/ProjectsPage.tsx": { - "react-hooks/set-state-in-effect": { - "count": 1 - } - }, "src/app/(dashboard)/prompts/_components/add_prompt_form.tsx": { "no-restricted-imports": { "count": 1 @@ -1098,14 +1080,6 @@ "count": 2 } }, - "src/app/(dashboard)/prompts/_components/prompt_table.tsx": { - "no-nested-ternary": { - "count": 1 - }, - "no-restricted-imports": { - "count": 1 - } - }, "src/app/(dashboard)/router-settings/_components/general_settings.tsx": { "no-nested-ternary": { "count": 3 @@ -1153,14 +1127,6 @@ "count": 1 } }, - "src/app/(dashboard)/skills/_components/plugin_table.tsx": { - "no-nested-ternary": { - "count": 1 - }, - "no-restricted-imports": { - "count": 1 - } - }, "src/app/(dashboard)/tag-management/_components/components/CreateTagModal.tsx": { "no-restricted-imports": { "count": 1 @@ -1320,11 +1286,6 @@ "count": 1 } }, - "src/app/(dashboard)/vector-stores/_components/VectorStoreTable.tsx": { - "no-restricted-imports": { - "count": 1 - } - }, "src/app/(dashboard)/vector-stores/_components/index.tsx": { "no-restricted-imports": { "count": 1 @@ -1448,22 +1409,6 @@ "count": 1 } }, - "src/components/DeletedKeysPage/DeletedKeysTable/DeletedKeysTable.tsx": { - "no-nested-ternary": { - "count": 1 - }, - "no-restricted-imports": { - "count": 1 - } - }, - "src/components/DeletedTeamsPage/DeletedTeamsTable/DeletedTeamsTable.tsx": { - "no-nested-ternary": { - "count": 1 - }, - "no-restricted-imports": { - "count": 1 - } - }, "src/components/EntityUsageExport/ExportSummary.tsx": { "no-restricted-imports": { "count": 1 @@ -1578,14 +1523,6 @@ "count": 1 } }, - "src/components/Settings/LoggingAndAlerts/LoggingCallbacks/LoggingCallbacksTable.tsx": { - "no-nested-ternary": { - "count": 1 - }, - "no-restricted-imports": { - "count": 1 - } - }, "src/components/Settings/RouterSettings/Fallbacks/AddFallbacks.tsx": { "no-restricted-imports": { "count": 1 @@ -2145,11 +2082,6 @@ "count": 1 } }, - "src/components/pass_through_settings.tsx": { - "no-restricted-imports": { - "count": 1 - } - }, "src/components/per_user_usage.tsx": { "no-restricted-imports": { "count": 1 diff --git a/ui/litellm-dashboard/src/app/(dashboard)/budgets/_components/BudgetTable.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/budgets/_components/BudgetTable.test.tsx new file mode 100644 index 00000000000..9a7a7bd2eb9 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/budgets/_components/BudgetTable.test.tsx @@ -0,0 +1,87 @@ +import { screen, within } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import { beforeEach, describe, expect, it, vi } from "vitest"; +import { renderWithProviders } from "@/../tests/test-utils"; +import BudgetTable from "./BudgetTable"; +import { budgetItem } from "@/app/(dashboard)/hooks/budgets/useBudgets"; + +const makeBudget = (overrides: Partial = {}): budgetItem => ({ + budget_id: "budget-1", + max_budget: 100, + tpm_limit: 1000, + rpm_limit: 10, + updated_at: "2024-01-01T00:00:00Z", + ...overrides, +}); + +const defaultProps = { + budgets: [makeBudget()], + isLoading: false, + canModify: true, + onEditClick: vi.fn(), + onDeleteClick: vi.fn(), +}; + +describe("BudgetTable", () => { + beforeEach(() => { + vi.clearAllMocks(); + }); + + it("should display budget information", () => { + renderWithProviders(); + expect(screen.getByText("budget-1")).toBeInTheDocument(); + expect(screen.getByText("$100.00")).toBeInTheDocument(); + expect(screen.getByText("1000")).toBeInTheDocument(); + expect(screen.getByText("10")).toBeInTheDocument(); + }); + + it("should show n/a for missing rate limits and Unlimited for a missing max budget", () => { + renderWithProviders( + , + ); + expect(screen.getAllByText("n/a")).toHaveLength(2); + expect(screen.getByText("Unlimited")).toBeInTheDocument(); + }); + + it("should sort budgets by updated_at descending", () => { + const budgets = [ + makeBudget({ budget_id: "budget-old", updated_at: "2024-01-01T00:00:00Z" }), + makeBudget({ budget_id: "budget-new", updated_at: "2024-06-01T00:00:00Z" }), + ]; + renderWithProviders(); + const rows = screen.getAllByRole("row").slice(1); + expect(within(rows[0]).getByText("budget-new")).toBeInTheDocument(); + expect(within(rows[1]).getByText("budget-old")).toBeInTheDocument(); + }); + + it("should call onEditClick from the actions menu", async () => { + const user = userEvent.setup(); + renderWithProviders(); + await user.click(screen.getByTestId("budget-actions-budget-1")); + await user.click(await screen.findByTestId("budget-action-edit")); + expect(defaultProps.onEditClick).toHaveBeenCalledWith(defaultProps.budgets[0]); + }); + + it("should call onDeleteClick from the actions menu", async () => { + const user = userEvent.setup(); + renderWithProviders(); + await user.click(screen.getByTestId("budget-actions-budget-1")); + await user.click(await screen.findByTestId("budget-action-delete")); + expect(defaultProps.onDeleteClick).toHaveBeenCalledWith(defaultProps.budgets[0]); + }); + + it("should not render the actions menu when the user cannot modify budgets", () => { + renderWithProviders(); + expect(screen.queryByTestId("budget-actions-budget-1")).not.toBeInTheDocument(); + }); + + it("should show skeleton rows when loading", () => { + renderWithProviders(); + expect(screen.getAllByTestId("skeleton-row").length).toBeGreaterThan(0); + }); + + it("should show the empty state when there are no budgets", () => { + renderWithProviders(); + expect(screen.getByText("No budgets yet")).toBeInTheDocument(); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/budgets/_components/BudgetTable.tsx b/ui/litellm-dashboard/src/app/(dashboard)/budgets/_components/BudgetTable.tsx new file mode 100644 index 00000000000..4bc06425f80 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/budgets/_components/BudgetTable.tsx @@ -0,0 +1,57 @@ +"use client"; + +import { Inbox } from "lucide-react"; +import React, { useMemo } from "react"; + +import { DataTable } from "@/components/shared/DataTable"; +import { budgetItem } from "@/app/(dashboard)/hooks/budgets/useBudgets"; + +import { getBudgetTableColumns } from "./BudgetTableColumns"; + +interface BudgetTableProps { + budgets: budgetItem[]; + isLoading: boolean; + canModify: boolean; + onEditClick: (budget: budgetItem) => void; + onDeleteClick: (budget: budgetItem) => void; +} + +function EmptyState() { + return ( +
+
+ +
+
No budgets yet
+
+ Create a budget to set spend, TPM and RPM limits for customers. +
+
+ ); +} + +const BudgetTable: React.FC = ({ budgets, isLoading, canModify, onEditClick, onDeleteClick }) => { + const rows = useMemo( + () => [...budgets].sort((a, b) => new Date(b.updated_at).getTime() - new Date(a.updated_at).getTime()), + [budgets], + ); + + const columns = useMemo( + () => getBudgetTableColumns({ canModify, onEditClick, onDeleteClick }), + [canModify, onEditClick, onDeleteClick], + ); + + return ( + budget.budget_id || String(index)} + isLoading={isLoading} + loadingMessage="Loading budgets…" + noDataMessage={} + size="compact" + /> + ); +}; + +export default BudgetTable; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/budgets/_components/BudgetTableColumns.tsx b/ui/litellm-dashboard/src/app/(dashboard)/budgets/_components/BudgetTableColumns.tsx new file mode 100644 index 00000000000..456ab9d6b68 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/budgets/_components/BudgetTableColumns.tsx @@ -0,0 +1,124 @@ +"use client"; + +import { ColumnDef } from "@tanstack/react-table"; +import { MoreHorizontal, Pencil, Trash2 } from "lucide-react"; + +import { IdCell, MoneyCell } from "@/components/shared/table_cells"; +import { budgetItem } from "@/app/(dashboard)/hooks/budgets/useBudgets"; +import { buttonVariants } from "@/components/ui/button"; +import { + DropdownMenu, + DropdownMenuContent, + DropdownMenuItem, + DropdownMenuSeparator, + DropdownMenuTrigger, +} from "@/components/ui/dropdown-menu"; +import { cn } from "@/lib/cva.config"; + +function RateLimitCell({ value }: { value: number | null }) { + if (value == null) { + return n/a; + } + return {value}; +} + +interface BudgetRowActionsProps { + budget: budgetItem; + onEditClick: (budget: budgetItem) => void; + onDeleteClick: (budget: budgetItem) => void; +} + +function BudgetRowActions({ budget, onEditClick, onDeleteClick }: BudgetRowActionsProps) { + return ( + + + + + + onEditClick(budget)}> + + Edit budget + + + onDeleteClick(budget)} + > + + Delete budget + + + + ); +} + +interface BudgetTableColumnsDeps { + canModify: boolean; + onEditClick: (budget: budgetItem) => void; + onDeleteClick: (budget: budgetItem) => void; +} + +export const getBudgetTableColumns = ({ + canModify, + onEditClick, + onDeleteClick, +}: BudgetTableColumnsDeps): ColumnDef[] => [ + { + id: "budget_id", + accessorKey: "budget_id", + meta: { title: "Budget ID" }, + header: "Budget ID", + size: 220, + enableSorting: false, + cell: ({ row }) => , + }, + { + id: "max_budget", + accessorKey: "max_budget", + meta: { title: "Max Budget", numeric: true }, + header: "Max Budget", + size: 120, + enableSorting: false, + cell: ({ row }) => , + }, + { + id: "tpm_limit", + accessorKey: "tpm_limit", + meta: { title: "TPM", numeric: true }, + header: "TPM", + size: 100, + enableSorting: false, + cell: ({ row }) => , + }, + { + id: "rpm_limit", + accessorKey: "rpm_limit", + meta: { title: "RPM", numeric: true }, + header: "RPM", + size: 100, + enableSorting: false, + cell: ({ row }) => , + }, + ...(canModify + ? [ + { + id: "actions", + meta: { className: "text-right", headerClassName: "text-right" }, + header: () => Actions, + size: 64, + enableSorting: false, + enableHiding: false, + cell: ({ row }) => ( +
+ +
+ ), + } satisfies ColumnDef, + ] + : []), +]; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/budgets/_components/budget_panel.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/budgets/_components/budget_panel.test.tsx index f4d70a5e8f8..392616f1935 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/budgets/_components/budget_panel.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/budgets/_components/budget_panel.test.tsx @@ -1,5 +1,6 @@ import { fireEvent, render, waitFor, screen } from "@testing-library/react"; import { act } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; import { afterEach, describe, expect, it, vi } from "vitest"; import BudgetPanel from "./budget_panel"; @@ -57,7 +58,8 @@ describe("Budget Panel", () => { }); }); - it("should open delete modal when clicking delete icon", async () => { + it("should open delete modal from the actions menu", async () => { + const user = userEvent.setup(); vi.mocked(useBudgets).mockReturnValue({ data: [ { @@ -77,11 +79,8 @@ describe("Budget Panel", () => { expect(screen.getByText("budget-to-delete")).toBeInTheDocument(); }); - const deleteButton = screen.getByTestId("delete-budget-button"); - - act(() => { - fireEvent.click(deleteButton); - }); + await user.click(screen.getByTestId("budget-actions-budget-to-delete")); + await user.click(await screen.findByTestId("budget-action-delete")); await waitFor(() => { expect(screen.getByText("Delete Budget?")).toBeInTheDocument(); @@ -89,6 +88,7 @@ describe("Budget Panel", () => { }); it("should successfully delete a budget", async () => { + const user = userEvent.setup(); const deleteMutateAsync = vi.fn().mockResolvedValue(undefined); vi.mocked(useBudgets).mockReturnValue({ data: [ @@ -113,17 +113,13 @@ describe("Budget Panel", () => { expect(screen.getByText("budget-to-delete")).toBeInTheDocument(); }); - // Open delete modal - const deleteButton = screen.getByTestId("delete-budget-button"); - act(() => { - fireEvent.click(deleteButton); - }); + await user.click(screen.getByTestId("budget-actions-budget-to-delete")); + await user.click(await screen.findByTestId("budget-action-delete")); await waitFor(() => { expect(screen.getByText("Delete Budget?")).toBeInTheDocument(); }); - // Confirm delete const confirmButton = screen.getByRole("button", { name: /delete/i }); act(() => { fireEvent.click(confirmButton); @@ -148,6 +144,7 @@ describe("Budget Panel", () => { }); it("should handle delete error", async () => { + const user = userEvent.setup(); const deleteMutateAsync = vi.fn().mockRejectedValue(new Error("Delete failed")); vi.mocked(useBudgets).mockReturnValue({ data: [ @@ -172,17 +169,13 @@ describe("Budget Panel", () => { expect(screen.getByText("budget-to-delete")).toBeInTheDocument(); }); - // Open delete modal - const deleteButton = screen.getByTestId("delete-budget-button"); - act(() => { - fireEvent.click(deleteButton); - }); + await user.click(screen.getByTestId("budget-actions-budget-to-delete")); + await user.click(await screen.findByTestId("budget-action-delete")); await waitFor(() => { expect(screen.getByText("Delete Budget?")).toBeInTheDocument(); }); - // Confirm delete const confirmButton = screen.getByRole("button", { name: /delete/i }); act(() => { fireEvent.click(confirmButton); @@ -193,7 +186,8 @@ describe("Budget Panel", () => { }); }); - it("should open edit modal when clicking edit icon", async () => { + it("should open edit modal from the actions menu", async () => { + const user = userEvent.setup(); vi.mocked(useBudgets).mockReturnValue({ data: [ { @@ -213,11 +207,8 @@ describe("Budget Panel", () => { expect(screen.getByText("budget-to-edit")).toBeInTheDocument(); }); - const editButton = screen.getByTestId("edit-budget-button"); - - act(() => { - fireEvent.click(editButton); - }); + await user.click(screen.getByTestId("budget-actions-budget-to-edit")); + await user.click(await screen.findByTestId("budget-action-edit")); await waitFor(() => { expect(screen.getByText("Edit Budget")).toBeInTheDocument(); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/budgets/_components/budget_panel.tsx b/ui/litellm-dashboard/src/app/(dashboard)/budgets/_components/budget_panel.tsx index af15a99f0b4..6d5c0c7be08 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/budgets/_components/budget_panel.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/budgets/_components/budget_panel.tsx @@ -3,30 +3,14 @@ * */ -import { - Button, - Card, - Tab, - TabGroup, - Table, - TableBody, - TableCell, - TableHead, - TableHeaderCell, - TableRow, - TabList, - TabPanel, - TabPanels, - Text, -} from "@tremor/react"; +import { Button, Tab, TabGroup, TabList, TabPanel, TabPanels, Text } from "@tremor/react"; import React, { useState } from "react"; import { Prism as SyntaxHighlighter } from "react-syntax-highlighter"; import DeleteResourceModal from "@/components/common_components/DeleteResourceModal"; -import TableIconActionButton from "@/components/common_components/IconActionButton/TableIconActionButtons/TableIconActionButton"; import NotificationsManager from "@/components/molecules/notifications_manager"; import { useBudgets, useDeleteBudget, budgetItem } from "@/app/(dashboard)/hooks/budgets/useBudgets"; -import { MoneyCell } from "@/components/shared/table_cells"; import BudgetModal from "./budget_modal"; +import BudgetTable from "./BudgetTable"; import EditBudgetModal from "./edit_budget_modal"; import { CREATE_END_USER_CURL_COMMAND, CHAT_COMPLETIONS_CURL_COMMAND, OPENAI_SDK_PYTHON_CODE } from "./constants"; import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; @@ -46,7 +30,7 @@ const BudgetPanel: React.FC = ({ accessToken }) => { // Admin Viewer follows the read-parity rule: see budgets, no writes. const canModify = isProxyAdminRole(userRole ?? ""); - const { data: budgetList = [] } = useBudgets(); + const { data: budgetList = [], isLoading } = useBudgets(); const deleteBudget = useDeleteBudget(); const handleEditCall = async (budget: budgetItem) => { @@ -109,51 +93,14 @@ const BudgetPanel: React.FC = ({ accessToken }) => { existingBudget={selectedBudget} /> )} - - Create a budget to assign to customers. - - - - Budget ID - Max Budget - TPM - RPM - - - - - {budgetList - .slice() - .sort((a, b) => new Date(b.updated_at).getTime() - new Date(a.updated_at).getTime()) - .map((value: budgetItem) => ( - - {value.budget_id} - - - - {value.tpm_limit ? value.tpm_limit : "n/a"} - {value.rpm_limit ? value.rpm_limit : "n/a"} - {canModify && ( - <> - handleEditCall(value)} - dataTestId="edit-budget-button" - /> - handleDeleteClick(value)} - dataTestId="delete-budget-button" - /> - - )} - - ))} - -
-
+ Create a budget to assign to customers. + => { +): UseQueryResult => { const { accessToken } = useAuthorized(); - return useQuery({ + return useQuery({ queryKey: deletedKeyKeys.list({ page, limit: pageSize, ...options }), queryFn: async () => await keyListCall(accessToken!, page, pageSize, { ...options, status: "deleted" }), enabled: Boolean(accessToken), diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/useAuthorized.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/useAuthorized.test.ts index 94f9d9173f0..7076e69edc2 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/useAuthorized.test.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/useAuthorized.test.ts @@ -1,7 +1,7 @@ /* @vitest-environment jsdom */ import React from "react"; import { renderHook, waitFor } from "@testing-library/react"; -import { afterEach, describe, expect, it, vi } from "vitest"; +import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; import useAuthorized from "./useAuthorized"; @@ -26,12 +26,6 @@ const { buildLoginUrlWithReturnMock: vi.fn((baseUrl: string) => baseUrl), })); -vi.mock("next/navigation", () => ({ - useRouter: () => ({ - replace: replaceMock, - }), -})); - vi.mock("@/components/networking", async (importOriginal) => { const actual = await importOriginal(); return { @@ -91,7 +85,28 @@ const clearCookie = () => { }; describe("useAuthorized", () => { + const originalLocation = window.location; + + beforeEach(() => { + Object.defineProperty(window, "location", { + value: { + href: "http://proxy.example/ui/?page=api-keys", + origin: "http://proxy.example", + hostname: "proxy.example", + pathname: "/ui/", + search: "?page=api-keys", + protocol: "http:", + replace: replaceMock, + }, + writable: true, + }); + }); + afterEach(() => { + Object.defineProperty(window, "location", { + value: originalLocation, + writable: true, + }); replaceMock.mockReset(); clearTokenCookiesMock.mockReset(); getProxyBaseUrlMock.mockClear(); @@ -164,7 +179,7 @@ describe("useAuthorized", () => { expect(clearTokenCookiesMock).toHaveBeenCalled(); }); - expect(replaceMock).toHaveBeenCalledWith("http://proxy.example/ui/login"); + expect(replaceMock).toHaveBeenCalledWith("http://proxy.example/ui/login/"); expect(result.current.accessToken).toBeNull(); expect(result.current.userRole).toBe("Undefined Role"); }); @@ -197,7 +212,7 @@ describe("useAuthorized", () => { const { result } = renderHook(() => useAuthorized(), { wrapper }); await waitFor(() => { - expect(replaceMock).toHaveBeenCalledWith("http://proxy.example/ui/login"); + expect(replaceMock).toHaveBeenCalledWith("http://proxy.example/ui/login/"); }); expect(result.current.accessToken).toBe("api-key-123"); @@ -221,7 +236,7 @@ describe("useAuthorized", () => { const { result } = renderHook(() => useAuthorized(), { wrapper }); await waitFor(() => { - expect(replaceMock).toHaveBeenCalledWith("http://proxy.example/ui/login"); + expect(replaceMock).toHaveBeenCalledWith("http://proxy.example/ui/login/"); }); expect(clearTokenCookiesMock).not.toHaveBeenCalled(); @@ -256,7 +271,7 @@ describe("useAuthorized", () => { expect(clearTokenCookiesMock).toHaveBeenCalled(); }); - expect(replaceMock).toHaveBeenCalledWith("http://proxy.example/ui/login"); + expect(replaceMock).toHaveBeenCalledWith("http://proxy.example/ui/login/"); expect(checkTokenValidityMock).toHaveBeenCalledWith(token); }); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/useAuthorized.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/useAuthorized.ts index 8f8c403a4e9..bb22ebf5edc 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/useAuthorized.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/useAuthorized.ts @@ -3,14 +3,12 @@ import { getProxyBaseUrl } from "@/components/networking"; import { clearTokenCookies, getCookie } from "@/utils/cookieUtils"; import { checkTokenValidity, decodeToken } from "@/utils/jwtUtils"; -import { buildLoginUrlWithReturn, storeReturnUrl } from "@/utils/returnUrlUtils"; -import { useRouter } from "next/navigation"; +import { buildLoginUrlWithReturn, getLoginUrl, storeReturnUrl } from "@/utils/returnUrlUtils"; import { useCallback, useEffect, useMemo } from "react"; import { formatUserRole } from "@/utils/roles"; import { useUIConfig } from "./uiConfig/useUIConfig"; const useAuthorized = () => { - const router = useRouter(); const { data: uiConfig, isLoading: isUIConfigLoading } = useUIConfig(); const token = typeof document !== "undefined" ? getCookie("token") : null; @@ -23,10 +21,10 @@ const useAuthorized = () => { // Helper function to redirect to login while preserving the current URL const redirectToLogin = useCallback(() => { storeReturnUrl(); - const baseLoginUrl = `${getProxyBaseUrl()}/ui/login`; + const baseLoginUrl = getLoginUrl(getProxyBaseUrl()); const loginUrlWithReturn = buildLoginUrlWithReturn(baseLoginUrl); - router.replace(loginUrlWithReturn); - }, [router]); + window.location.replace(loginUrlWithReturn); + }, []); // Single useEffect for all redirect logic useEffect(() => { diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/OAuthFormFields.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/OAuthFormFields.tsx index fc72981edd0..76e4342f52e 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/OAuthFormFields.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/OAuthFormFields.tsx @@ -167,6 +167,17 @@ const OAuthFormFields: React.FC = ({ > + @@ -1588,6 +1607,7 @@ const MCPServerEdit: React.FC = ({ oauthFlowTypeValue ?? oauth2FlowToFormValue(mcpServer.oauth2_flow) ?? OAUTH_FLOW.INTERACTIVE, static_headers: currentStaticHeaders ?? mcpServer.static_headers, credentials: currentCredentials, + issuer: currentIssuer ?? mcpServer.issuer, authorization_url: currentAuthorizationUrl ?? mcpServer.authorization_url, token_url: currentTokenUrl ?? mcpServer.token_url, registration_url: currentRegistrationUrl ?? mcpServer.registration_url, diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/ModelsAndEndpointsView.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/ModelsAndEndpointsView.tsx index 4bf4af0c7f6..aac5405ce6b 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/ModelsAndEndpointsView.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/ModelsAndEndpointsView.tsx @@ -26,7 +26,7 @@ import HealthCheckComponent from "../../../components/model_dashboard/HealthChec import ModelGroupAliasSettings from "../../../components/model_group_alias_settings"; import ModelInfoView from "../../../components/model_info_view"; import NotificationsManager from "../../../components/molecules/notifications_manager"; -import PassThroughSettings from "../../../components/pass_through_settings"; +import PassThroughSettings from "../../../components/PassThroughSettings/PassThroughSettings"; import TeamInfoView from "../../../components/team/TeamInfo"; import useAuthorized from "../hooks/useAuthorized"; @@ -396,7 +396,6 @@ const ModelsAndEndpointsView: React.FC = ({ premiumUser, te accessToken={accessToken} userRole={userRole} userID={userID} - modelData={processedModelData} premiumUser={premiumUser} /> diff --git a/ui/litellm-dashboard/src/app/(dashboard)/page.tsx b/ui/litellm-dashboard/src/app/(dashboard)/page.tsx index 6c0d780183a..02b2ccf5357 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/page.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/page.tsx @@ -7,6 +7,7 @@ import { useAuth } from "@/contexts/AuthContext"; import { buildLoginUrlWithReturn, consumeReturnUrl, + getLoginUrl, isValidReturnUrl, normalizeUrlForCompare, storeReturnUrl, @@ -33,7 +34,7 @@ function CreateKeyPageContent() { // Store the current URL so we can redirect back after login storeReturnUrl(); // Build login URL with return URL parameter - const baseLoginUrl = (proxyBaseUrl || "") + "/ui/login"; + const baseLoginUrl = getLoginUrl(proxyBaseUrl || ""); const dest = buildLoginUrlWithReturn(baseLoginUrl); // Replace instead of assigning to avoid back-button loops window.location.replace(dest); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/policies/_components/policy_table.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/policies/_components/PolicyTable.test.tsx similarity index 58% rename from ui/litellm-dashboard/src/app/(dashboard)/policies/_components/policy_table.test.tsx rename to ui/litellm-dashboard/src/app/(dashboard)/policies/_components/PolicyTable.test.tsx index 934114f5405..06c939aa151 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/policies/_components/policy_table.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/policies/_components/PolicyTable.test.tsx @@ -1,47 +1,11 @@ import React from "react"; -import { screen } from "@testing-library/react"; +import { screen, within } from "@testing-library/react"; import userEvent from "@testing-library/user-event"; import { renderWithProviders } from "@/../tests/test-utils"; import { beforeEach, describe, expect, it, vi } from "vitest"; -import PolicyTable from "./policy_table"; +import PolicyTable from "./PolicyTable"; import { Policy } from "@/components/policies/types"; -vi.mock("@heroicons/react/outline", () => ({ - TrashIcon: function TrashIcon() { - return null; - }, - PencilIcon: function PencilIcon() { - return null; - }, - SwitchVerticalIcon: function SwitchVerticalIcon() { - return null; - }, - ChevronUpIcon: function ChevronUpIcon() { - return null; - }, - ChevronDownIcon: function ChevronDownIcon() { - return null; - }, -})); - -vi.mock("@tremor/react", async (importOriginal) => { - const actual = await importOriginal(); - return { - ...actual, - Button: React.forwardRef(({ children, ...props }, ref) => - React.createElement("button", { ...props, ref }, children), - ), - Icon: ({ icon: IconComp, onClick, className }: any) => - React.createElement( - "button", - { type: "button", onClick, className }, - IconComp?.displayName ?? IconComp?.name ?? "icon", - ), - Tooltip: ({ children }: { children?: React.ReactNode }) => React.createElement(React.Fragment, null, children), - Badge: ({ children }: { children?: React.ReactNode }) => React.createElement("span", null, children), - }; -}); - const makePolicy = (overrides: Partial = {}): Policy => ({ policy_id: "policy-id-1", policy_name: "test-policy", @@ -71,20 +35,21 @@ describe("PolicyTable", () => { renderWithProviders(); expect(screen.getByText("Name")).toBeInTheDocument(); expect(screen.getByText("Description")).toBeInTheDocument(); - expect(screen.getByText("Actions")).toBeInTheDocument(); + expect(screen.getByText("Guardrails (Add)")).toBeInTheDocument(); + expect(screen.getByText("Created At")).toBeInTheDocument(); }); - it("should show a loading message when isLoading is true", () => { + it("should show skeleton rows when isLoading is true", () => { renderWithProviders(); - expect(screen.getByText(/loading/i)).toBeInTheDocument(); + expect(screen.getAllByTestId("skeleton-row").length).toBeGreaterThan(0); }); - it("should show 'No policies found' when there are no policies", () => { + it("should show the empty state when there are no policies", () => { renderWithProviders(); - expect(screen.getByText(/no policies found/i)).toBeInTheDocument(); + expect(screen.getByText("No policies found")).toBeInTheDocument(); }); - it("should render a button with the policy name for each grouped policy", () => { + it("should render a clickable name cell for each grouped policy", () => { const policies = [ makePolicy({ policy_name: "alpha-policy", policy_id: "id-1" }), makePolicy({ policy_name: "beta-policy", policy_id: "id-2" }), @@ -94,7 +59,18 @@ describe("PolicyTable", () => { expect(screen.getByRole("button", { name: "beta-policy" })).toBeInTheDocument(); }); - it("should call onViewClick with the policy_id when the policy name button is clicked", async () => { + it("should sort rows by policy name ascending by default", () => { + const policies = [ + makePolicy({ policy_name: "zeta-policy", policy_id: "id-z" }), + makePolicy({ policy_name: "alpha-policy", policy_id: "id-a" }), + ]; + renderWithProviders(); + const rows = screen.getAllByRole("row").slice(1); + expect(within(rows[0]).getByText("alpha-policy")).toBeInTheDocument(); + expect(within(rows[1]).getByText("zeta-policy")).toBeInTheDocument(); + }); + + it("should call onViewClick with the policy_id when the policy name is clicked", async () => { const user = userEvent.setup(); const policy = makePolicy({ policy_name: "my-policy", policy_id: "view-id-1" }); renderWithProviders(); @@ -102,36 +78,46 @@ describe("PolicyTable", () => { expect(defaultProps.onViewClick).toHaveBeenCalledWith("view-id-1"); }); - it("should call onDeleteClick with policy_id and policy_name when the delete icon is clicked", async () => { + it("should call onDeleteClick with policy_id and policy_name from the actions menu", async () => { const user = userEvent.setup(); const policy = makePolicy({ policy_name: "del-policy", policy_id: "del-id-1" }); renderWithProviders(); - await user.click(screen.getByRole("button", { name: /TrashIcon/i })); + await user.click(screen.getByTestId("policy-actions-del-id-1")); + await user.click(await screen.findByTestId("policy-action-delete")); expect(defaultProps.onDeleteClick).toHaveBeenCalledWith("del-id-1", "del-policy"); }); - it("should call onEditClick with the policy when the edit icon is clicked", async () => { + it("should call onEditClick with the policy from the actions menu", async () => { const user = userEvent.setup(); const policy = makePolicy({ policy_name: "edit-policy", policy_id: "edit-id-1" }); renderWithProviders(); - await user.click(screen.getByRole("button", { name: /PencilIcon/i })); + await user.click(screen.getByTestId("policy-actions-edit-id-1")); + await user.click(await screen.findByTestId("policy-action-edit")); expect(defaultProps.onEditClick).toHaveBeenCalledWith(policy); }); - it("should not show admin action icons for non-admins", () => { + it("should not show the actions menu for non-admins", () => { const policy = makePolicy(); renderWithProviders(); - expect(screen.queryByRole("button", { name: /TrashIcon/i })).not.toBeInTheDocument(); - expect(screen.queryByRole("button", { name: /PencilIcon/i })).not.toBeInTheDocument(); + expect(screen.queryByTestId(`policy-actions-${policy.policy_id}`)).not.toBeInTheDocument(); }); it("should show a version badge when multiple versions of the same policy name exist", () => { - const policies = [ - makePolicy({ policy_name: "versioned", policy_id: "v1", version_status: "published", version_number: 1 }), - makePolicy({ policy_name: "versioned", policy_id: "v2", version_status: "production", version_number: 2 }), - ]; + const publishedVersion: Partial = { + policy_name: "versioned", + policy_id: "v1", + version_status: "published", + version_number: 1, + }; + const productionVersion: Partial = { + policy_name: "versioned", + policy_id: "v2", + version_status: "production", + version_number: 2, + }; + const policies = [makePolicy(publishedVersion), makePolicy(productionVersion)]; renderWithProviders(); - expect(screen.getByText(/2 version/i)).toBeInTheDocument(); + expect(screen.getByText("2 versions")).toBeInTheDocument(); }); it("should group policies with the same name into a single row", () => { @@ -140,10 +126,10 @@ describe("PolicyTable", () => { makePolicy({ policy_name: "shared", policy_id: "s2", version_status: "production" }), ]; renderWithProviders(); - expect(screen.getAllByRole("button", { name: "shared" })).toHaveLength(1); + expect(screen.getAllByText("shared")).toHaveLength(1); }); - it("should show an overflow tag when more than 2 guardrails_add exist", () => { + it("should show an overflow badge when more than 2 guardrails_add exist", () => { const policy = makePolicy({ guardrails_add: ["g1", "g2", "g3", "g4"] }); renderWithProviders(); expect(screen.getByText("+2")).toBeInTheDocument(); @@ -156,7 +142,7 @@ describe("PolicyTable", () => { makePolicy({ policy_name: "grouped", policy_id: "prod-id", version_status: "production" }), ]; renderWithProviders(); - await user.click(screen.getByRole("button", { name: "grouped" })); + await user.click(screen.getByRole("button", { name: /grouped/ })); expect(defaultProps.onViewClick).toHaveBeenCalledWith("prod-id"); }); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/policies/_components/PolicyTable.tsx b/ui/litellm-dashboard/src/app/(dashboard)/policies/_components/PolicyTable.tsx new file mode 100644 index 00000000000..d6e841c2119 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/policies/_components/PolicyTable.tsx @@ -0,0 +1,82 @@ +"use client"; + +import { SortingState } from "@tanstack/react-table"; +import { Inbox } from "lucide-react"; +import React, { useMemo, useState } from "react"; + +import { DataTable } from "@/components/shared/DataTable"; +import { Policy } from "@/components/policies/types"; + +import { getPolicyTableColumns, PolicyRow } from "./PolicyTableColumns"; + +/** One row per policy name; primaryPolicy is used for display and for Edit (FlowBuilder loads all versions) */ +function groupPoliciesByName(policies: Policy[]): PolicyRow[] { + const names = Array.from(new Set(policies.map((policy) => policy.policy_name || "(unnamed)"))); + return names.map((policyName) => { + const versions = policies.filter((policy) => (policy.policy_name || "(unnamed)") === policyName); + const primary = + versions.find((version) => version.version_status === "production") ?? + [...versions].sort((a, b) => (b.version_number ?? 0) - (a.version_number ?? 0))[0]; + return { policy_name: policyName, primaryPolicy: primary, versionCount: versions.length }; + }); +} + +interface PolicyTableProps { + policies: Policy[]; + isLoading: boolean; + onDeleteClick: (policyId: string, policyName: string) => void; + onEditClick: (policy: Policy) => void; + onViewClick: (policyId: string) => void; + isAdmin?: boolean; +} + +const DEFAULT_SORTING: SortingState = [{ id: "policy_name", desc: false }]; + +function EmptyState() { + return ( +
+
+ +
+
No policies found
+
+ Create a policy to bundle guardrails and apply them across teams. +
+
+ ); +} + +const PolicyTable: React.FC = ({ + policies, + isLoading, + onDeleteClick, + onEditClick, + onViewClick, + isAdmin = false, +}) => { + const [sorting, setSorting] = useState(DEFAULT_SORTING); + + const rows = useMemo(() => groupPoliciesByName(policies), [policies]); + + const columns = useMemo(() => { + const deps = { isAdmin, onViewClick, onEditClick, onDeleteClick }; + return getPolicyTableColumns(deps); + }, [isAdmin, onViewClick, onEditClick, onDeleteClick]); + + return ( + row.policy_name} + sortingMode="client" + sorting={sorting} + onSortingChange={setSorting} + isLoading={isLoading} + loadingMessage="Loading policies…" + noDataMessage={} + size="compact" + /> + ); +}; + +export default PolicyTable; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/policies/_components/PolicyTableColumns.tsx b/ui/litellm-dashboard/src/app/(dashboard)/policies/_components/PolicyTableColumns.tsx new file mode 100644 index 00000000000..dd036d83283 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/policies/_components/PolicyTableColumns.tsx @@ -0,0 +1,207 @@ +"use client"; + +import { ColumnDef } from "@tanstack/react-table"; +import { MoreHorizontal, Pencil, Trash2 } from "lucide-react"; + +import { DataTableSortHeader } from "@/components/shared/DataTable"; +import { DateCell, IdentityCell, StatusBadge } from "@/components/shared/table_cells"; +import { Policy } from "@/components/policies/types"; +import { buttonVariants } from "@/components/ui/button"; +import { + DropdownMenu, + DropdownMenuContent, + DropdownMenuItem, + DropdownMenuSeparator, + DropdownMenuTrigger, +} from "@/components/ui/dropdown-menu"; +import { cn } from "@/lib/cva.config"; + +export interface PolicyRow { + policy_name: string; + primaryPolicy: Policy; + versionCount: number; +} + +function GuardrailChips({ guardrails, tone }: { guardrails: string[]; tone: "success" | "error" }) { + if (guardrails.length === 0) { + return -; + } + return ( +
+ {guardrails.slice(0, 2).map((guardrail) => ( + + ))} + {guardrails.length > 2 && ( + + )} +
+ ); +} + +interface PolicyRowActionsProps { + policy: Policy; + onEditClick: (policy: Policy) => void; + onDeleteClick: (policyId: string, policyName: string) => void; +} + +function PolicyRowActions({ policy, onEditClick, onDeleteClick }: PolicyRowActionsProps) { + return ( + + + + + + onEditClick(policy)}> + + Edit policy + + + onDeleteClick(policy.policy_id, policy.policy_name || "Unnamed Policy")} + > + + Delete policy + + + + ); +} + +interface PolicyTableColumnsDeps { + isAdmin: boolean; + onViewClick: (policyId: string) => void; + onEditClick: (policy: Policy) => void; + onDeleteClick: (policyId: string, policyName: string) => void; +} + +export const getPolicyTableColumns = ({ + isAdmin, + onViewClick, + onEditClick, + onDeleteClick, +}: PolicyTableColumnsDeps): ColumnDef[] => [ + { + id: "policy_name", + accessorKey: "policy_name", + meta: { title: "Name", skeleton: "twoLine" }, + header: ({ column }) => , + size: 220, + enableSorting: true, + cell: ({ row }) => ( + 1 ? ( + + ) : undefined + } + onClick={() => onViewClick(row.original.primaryPolicy.policy_id)} + /> + ), + }, + { + id: "description", + accessorFn: (row) => row.primaryPolicy.description ?? "", + meta: { title: "Description" }, + header: "Description", + size: 220, + enableSorting: false, + cell: ({ row }) => { + const description = row.original.primaryPolicy.description; + if (!description) { + return -; + } + return ( + + {description} + + ); + }, + }, + { + id: "inherit", + accessorFn: (row) => row.primaryPolicy.inherit ?? "", + meta: { title: "Inherits From", skeleton: "badge" }, + header: "Inherits From", + size: 150, + enableSorting: false, + cell: ({ row }) => { + const inherit = row.original.primaryPolicy.inherit; + if (!inherit) { + return -; + } + return ; + }, + }, + { + id: "guardrails_add", + meta: { title: "Guardrails (Add)", skeleton: "chips" }, + header: "Guardrails (Add)", + size: 180, + enableSorting: false, + cell: ({ row }) => , + }, + { + id: "guardrails_remove", + meta: { title: "Guardrails (Remove)", skeleton: "chips" }, + header: "Guardrails (Remove)", + size: 180, + enableSorting: false, + cell: ({ row }) => , + }, + { + id: "model_condition", + meta: { title: "Model Condition" }, + header: "Model Condition", + size: 160, + enableSorting: false, + cell: ({ row }) => { + const model = row.original.primaryPolicy.condition?.model; + if (!model) { + return -; + } + return ( + + {model} + + ); + }, + }, + { + id: "created_at", + accessorFn: (row) => row.primaryPolicy.created_at ?? "", + meta: { title: "Created At" }, + header: ({ column }) => , + size: 150, + enableSorting: true, + cell: ({ row }) => , + }, + ...(isAdmin + ? [ + { + id: "actions", + meta: { className: "text-right", headerClassName: "text-right" }, + header: () => Actions, + size: 64, + enableSorting: false, + enableHiding: false, + cell: ({ row }) => ( +
+ +
+ ), + } satisfies ColumnDef, + ] + : []), +]; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/policies/_components/index.tsx b/ui/litellm-dashboard/src/app/(dashboard)/policies/_components/index.tsx index 0545e2f86dc..df55d2c386b 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/policies/_components/index.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/policies/_components/index.tsx @@ -5,7 +5,7 @@ import { Alert } from "antd"; import MessageManager from "@/components/molecules/message_manager"; import { InfoCircleOutlined } from "@ant-design/icons"; import { isAdminRole } from "@/utils/roles"; -import PolicyTable from "./policy_table"; +import PolicyTable from "./PolicyTable"; import PolicyInfoView from "./policy_info"; import AddPolicyForm from "./add_policy_form"; import { FlowBuilderPage } from "./pipeline_flow_builder"; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/policies/_components/policy_table.tsx b/ui/litellm-dashboard/src/app/(dashboard)/policies/_components/policy_table.tsx deleted file mode 100644 index 278773069e6..00000000000 --- a/ui/litellm-dashboard/src/app/(dashboard)/policies/_components/policy_table.tsx +++ /dev/null @@ -1,326 +0,0 @@ -import React, { useMemo, useState } from "react"; -import { Table, TableBody, TableCell, TableHead, TableHeaderCell, TableRow, Icon, Button, Badge } from "@tremor/react"; -import { TrashIcon, PencilIcon, SwitchVerticalIcon, ChevronUpIcon, ChevronDownIcon } from "@heroicons/react/outline"; -import { Tooltip, Tag } from "antd"; -import { - ColumnDef, - flexRender, - getCoreRowModel, - getSortedRowModel, - SortingState, - useReactTable, -} from "@tanstack/react-table"; -import { DateCell } from "@/components/shared/table_cells"; -import { Policy } from "@/components/policies/types"; - -/** One row per policy name; primaryPolicy is used for display and for Edit (FlowBuilder loads all versions) */ -interface PolicyRow { - policy_name: string; - primaryPolicy: Policy; - versionCount: number; -} - -function groupPoliciesByName(policies: Policy[]): PolicyRow[] { - const byName = new Map(); - for (const p of policies) { - const name = p.policy_name || "(unnamed)"; - if (!byName.has(name)) byName.set(name, []); - byName.get(name)!.push(p); - } - const rows: PolicyRow[] = []; - for (const [policyName, versions] of byName) { - // Prefer production, then highest version_number - const primary = - versions.find((v) => v.version_status === "production") ?? - [...versions].sort((a, b) => (b.version_number ?? 0) - (a.version_number ?? 0))[0] ?? - versions[0]; - rows.push({ policy_name: policyName, primaryPolicy: primary, versionCount: versions.length }); - } - return rows.sort((a, b) => a.policy_name.localeCompare(b.policy_name)); -} - -interface PolicyTableProps { - policies: Policy[]; - isLoading: boolean; - onDeleteClick: (policyId: string, policyName: string) => void; - onEditClick: (policy: Policy) => void; - onViewClick: (policyId: string) => void; - isAdmin?: boolean; -} - -const PolicyTable: React.FC = ({ - policies, - isLoading, - onDeleteClick, - onEditClick, - onViewClick, - isAdmin = false, -}) => { - const [sorting, setSorting] = useState([{ id: "policy_name", desc: false }]); - - const rows = useMemo(() => groupPoliciesByName(policies), [policies]); - - const columns: ColumnDef[] = [ - { - header: "Name", - accessorKey: "policy_name", - cell: ({ row }) => { - const { primaryPolicy, versionCount } = row.original; - return ( -
- 1 ? ` (${versionCount} versions)` : ""}`} - > - - - {versionCount > 1 && ( - - {versionCount} version{versionCount !== 1 ? "s" : ""} - - )} -
- ); - }, - }, - { - header: "Description", - accessorFn: (row) => row.primaryPolicy.description ?? "", - cell: ({ row }) => { - const policy = row.original.primaryPolicy; - return ( - - {policy.description || "-"} - - ); - }, - }, - { - header: "Inherits From", - accessorFn: (row) => row.primaryPolicy.inherit ?? "", - cell: ({ row }) => { - const policy = row.original.primaryPolicy; - return policy.inherit ? ( - - {policy.inherit} - - ) : ( - - - ); - }, - }, - { - header: "Guardrails (Add)", - accessorFn: (row) => (row.primaryPolicy.guardrails_add ?? []).join(", "), - cell: ({ row }) => { - const policy = row.original.primaryPolicy; - const guardrails = policy.guardrails_add || []; - if (guardrails.length === 0) { - return -; - } - return ( -
- {guardrails.slice(0, 2).map((g, i) => ( - - {g} - - ))} - {guardrails.length > 2 && ( - - +{guardrails.length - 2} - - )} -
- ); - }, - }, - { - header: "Guardrails (Remove)", - accessorFn: (row) => (row.primaryPolicy.guardrails_remove ?? []).join(", "), - cell: ({ row }) => { - const policy = row.original.primaryPolicy; - const guardrails = policy.guardrails_remove || []; - if (guardrails.length === 0) { - return -; - } - return ( -
- {guardrails.slice(0, 2).map((g, i) => ( - - {g} - - ))} - {guardrails.length > 2 && ( - - +{guardrails.length - 2} - - )} -
- ); - }, - }, - { - header: "Model Condition", - accessorFn: (row) => { - const m = row.primaryPolicy.condition?.model; - return typeof m === "string" ? m : JSON.stringify(m ?? ""); - }, - cell: ({ row }) => { - const policy = row.original.primaryPolicy; - const modelCondition = policy.condition?.model; - if (!modelCondition) { - return -; - } - return ( - - - {typeof modelCondition === "string" - ? modelCondition.length > 20 - ? modelCondition.slice(0, 20) + "..." - : modelCondition - : "Multiple"} - - - ); - }, - }, - { - header: "Created At", - id: "created_at", - accessorFn: (row) => row.primaryPolicy.created_at ?? "", - cell: ({ row }) => , - }, - { - id: "actions", - header: "Actions", - cell: ({ row }) => { - const { primaryPolicy } = row.original; - const policy = primaryPolicy; - return ( -
- {isAdmin && ( - <> - - onEditClick(policy)} - className="cursor-pointer hover:text-blue-500" - /> - - - - policy.policy_id && onDeleteClick(policy.policy_id, policy.policy_name || "Unnamed Policy") - } - className="cursor-pointer hover:text-red-500" - /> - - - )} -
- ); - }, - }, - ]; - - const table = useReactTable({ - data: rows, - columns, - state: { - sorting, - }, - onSortingChange: setSorting, - getCoreRowModel: getCoreRowModel(), - getSortedRowModel: getSortedRowModel(), - enableSorting: true, - }); - - return ( -
-
- - - {table.getHeaderGroups().map((headerGroup) => ( - - {headerGroup.headers.map((header) => ( - -
-
- {header.isPlaceholder ? null : flexRender(header.column.columnDef.header, header.getContext())} -
- {header.id !== "actions" && ( -
- {header.column.getIsSorted() ? ( - { - asc: , - desc: , - }[header.column.getIsSorted() as string] - ) : ( - - )} -
- )} -
-
- ))} -
- ))} -
- - {isLoading ? ( - - -
-

Loading...

-
-
-
- ) : rows.length > 0 ? ( - table.getRowModel().rows.map((row) => ( - - {row.getVisibleCells().map((cell) => ( - - {flexRender(cell.column.columnDef.cell, cell.getContext())} - - ))} - - )) - ) : ( - - -
-

No policies found

-
-
-
- )} -
-
-
-
- ); -}; - -export default PolicyTable; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/projects/_components/ProjectDetailsPage.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/projects/_components/ProjectDetailsPage.test.tsx index 6f8676360c4..61d42f2aac6 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/projects/_components/ProjectDetailsPage.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/projects/_components/ProjectDetailsPage.test.tsx @@ -22,6 +22,12 @@ vi.mock("@/components/common_components/DefaultProxyAdminTag", () => ({ default: ({ userId }: { userId: string }) => {userId}, })); +vi.mock("./ProjectKeysSection", () => ({ + ProjectKeysSection: ({ projectId }: { projectId: string }) => ( +
{projectId}
+ ), +})); + const mockProject: ProjectResponse = { project_id: "proj-1", project_alias: "My Project", @@ -98,6 +104,11 @@ describe("ProjectDetail", () => { expect(screen.getByRole("heading", { name: "My Project" })).toBeInTheDocument(); }); + it("should render the project keys section for the project", () => { + renderWithProviders(); + expect(screen.getByTestId("project-keys-section")).toHaveTextContent("proj-1"); + }); + it("should display 'Active' for a non-blocked project", () => { renderWithProviders(); expect(screen.getByText("Active")).toBeInTheDocument(); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/projects/_components/ProjectDetailsPage.tsx b/ui/litellm-dashboard/src/app/(dashboard)/projects/_components/ProjectDetailsPage.tsx index 2bc8bc31129..fd043031b26 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/projects/_components/ProjectDetailsPage.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/projects/_components/ProjectDetailsPage.tsx @@ -17,10 +17,11 @@ import { } from "antd"; import { LoadingOutlined } from "@ant-design/icons"; import { BarChart } from "@/components/shared/charts"; -import { ArrowLeftIcon, DollarSignIcon, EditIcon, KeyIcon, UsersIcon } from "lucide-react"; +import { ArrowLeftIcon, DollarSignIcon, EditIcon, UsersIcon } from "lucide-react"; import { useMemo, useState } from "react"; import DefaultProxyAdminTag from "@/components/common_components/DefaultProxyAdminTag"; import { EditProjectModal } from "./ProjectModals/EditProjectModal"; +import { ProjectKeysSection } from "./ProjectKeysSection"; const { Title, Text } = Typography; const { Content } = Layout; @@ -203,17 +204,7 @@ export function ProjectDetail({ projectId, onBack }: ProjectDetailProps) { {/* Keys & Team */} - - - Keys - - } - style={{ height: "100%" }} - > - - + { isLoading: false, }); renderWithProviders(); - expect(screen.getByText("42 keys")).toBeInTheDocument(); + expect(screen.getByTestId("pagination-range")).toHaveTextContent("of 42"); }); it("should show 'No keys found' when the project has no keys", () => { diff --git a/ui/litellm-dashboard/src/app/(dashboard)/projects/_components/ProjectKeysSection.tsx b/ui/litellm-dashboard/src/app/(dashboard)/projects/_components/ProjectKeysSection.tsx index 9f60db596e8..4266b238e21 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/projects/_components/ProjectKeysSection.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/projects/_components/ProjectKeysSection.tsx @@ -1,6 +1,6 @@ import { useKeys } from "@/app/(dashboard)/hooks/keys/useKeys"; -import { LoadingOutlined } from "@ant-design/icons"; -import { Card, Flex, Input, Pagination, Spin } from "antd"; +import { PaginationState } from "@tanstack/react-table"; +import { Card, Flex, Input } from "antd"; import { KeyIcon, SearchIcon } from "lucide-react"; import { useEffect, useState } from "react"; import { ProjectKeysTable } from "./ProjectKeysTable"; @@ -12,17 +12,16 @@ interface ProjectKeysSectionProps { const PAGE_SIZE = 5; export function ProjectKeysSection({ projectId }: ProjectKeysSectionProps) { - const [page, setPage] = useState(1); + const [pagination, setPagination] = useState({ pageIndex: 0, pageSize: PAGE_SIZE }); const [keyAlias, setKeyAlias] = useState(""); - const { data, isLoading } = useKeys(page, PAGE_SIZE, { + const { data, isLoading } = useKeys(pagination.pageIndex + 1, pagination.pageSize, { projectID: projectId, selectedKeyAlias: keyAlias || null, }); - // Reset to page 1 when filter changes useEffect(() => { - setPage(1); + setPagination((current) => ({ ...current, pageIndex: 0 })); }, [keyAlias]); const keys = data?.keys ?? []; @@ -38,7 +37,7 @@ export function ProjectKeysSection({ projectId }: ProjectKeysSectionProps) { } style={{ height: "100%" }} > - + } placeholder="Filter by key name..." @@ -48,19 +47,13 @@ export function ProjectKeysSection({ projectId }: ProjectKeysSectionProps) { allowClear size="small" /> - `${total} keys`} - /> } /> } : false} + totalCount={totalCount} + isLoading={isLoading} + pagination={pagination} + onPaginationChange={setPagination} /> ); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/projects/_components/ProjectKeysTable.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/projects/_components/ProjectKeysTable.test.tsx index c0f12f77cfb..3685d30b414 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/projects/_components/ProjectKeysTable.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/projects/_components/ProjectKeysTable.test.tsx @@ -1,4 +1,5 @@ import { describe, it, expect, vi } from "vitest"; +import userEvent from "@testing-library/user-event"; import { renderWithProviders, screen } from "../../../../../tests/test-utils"; import { ProjectKeysTable } from "./ProjectKeysTable"; import { KeyResponse } from "@/components/key_team_helpers/key_list"; @@ -7,6 +8,13 @@ vi.mock("@/components/common_components/DefaultProxyAdminTag", () => ({ default: ({ userId }: { userId: string }) => {userId}, })); +const defaultProps = { + totalCount: 0, + isLoading: false, + pagination: { pageIndex: 0, pageSize: 5 }, + onPaginationChange: vi.fn(), +}; + function makeKey(overrides: Partial = {}): KeyResponse { return { token: "tok-abc123", @@ -70,52 +78,78 @@ function makeKey(overrides: Partial = {}): KeyResponse { describe("ProjectKeysTable", () => { it("should render", () => { - renderWithProviders(); + renderWithProviders(); expect(screen.getByRole("table")).toBeInTheDocument(); }); it("should display 'No keys found' when the keys list is empty", () => { - renderWithProviders(); + renderWithProviders(); expect(screen.getByText("No keys found")).toBeInTheDocument(); }); it("should display the key alias when provided", () => { - renderWithProviders(); + renderWithProviders(); expect(screen.getByText("My API Key")).toBeInTheDocument(); }); it("should display '—' when the key alias is null", () => { // Provide a user_id so only the alias column shows "—" (not the owner column too) - renderWithProviders(); + renderWithProviders( + , + ); expect(screen.getByText("—")).toBeInTheDocument(); }); it("should display the owner using user.user_email when available", () => { - const key = makeKey({ user: { user_id: "u1", user_email: "alice@example.com" } }); - renderWithProviders(); + const key = makeKey({ user: { user_id: "u1", user_email: "alice@example.com", user_alias: null } }); + renderWithProviders(); expect(screen.getByTestId("owner-tag")).toHaveTextContent("alice@example.com"); }); it("should fall back to user_id when user.user_email is absent", () => { const key = makeKey({ user_id: "user-99" }); - renderWithProviders(); + renderWithProviders(); expect(screen.getByTestId("owner-tag")).toHaveTextContent("user-99"); }); it("should display 'Never' in the Last Active column when last_active is null", () => { - renderWithProviders(); + renderWithProviders(); expect(screen.getByText("Never")).toBeInTheDocument(); }); it("should display a formatted date in the Last Active column when last_active is provided", () => { - renderWithProviders(); + renderWithProviders( + , + ); expect(screen.queryByText("Never")).not.toBeInTheDocument(); }); it("should render multiple keys as separate rows", () => { const keys = [makeKey({ token: "tok-1", key_alias: "Key One" }), makeKey({ token: "tok-2", key_alias: "Key Two" })]; - renderWithProviders(); + renderWithProviders(); expect(screen.getByText("Key One")).toBeInTheDocument(); expect(screen.getByText("Key Two")).toBeInTheDocument(); }); + + it("should show skeleton rows while loading", () => { + renderWithProviders(); + expect(screen.getAllByTestId("skeleton-row").length).toBeGreaterThan(0); + expect(screen.queryByText("No keys found")).not.toBeInTheDocument(); + }); + + it("should show the server-side total in the pagination footer", () => { + renderWithProviders(); + expect(screen.getByTestId("pagination-range")).toHaveTextContent("Showing 1-5 of 42"); + expect(screen.getByTestId("pagination-page")).toHaveTextContent("Page 1 of 9"); + }); + + it("should request the next page through the pagination footer", async () => { + const user = userEvent.setup(); + const onPaginationChange = vi.fn(); + renderWithProviders( + , + ); + await user.click(screen.getByTestId("pagination-next")); + expect(onPaginationChange).toHaveBeenCalled(); + }); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/projects/_components/ProjectKeysTable.tsx b/ui/litellm-dashboard/src/app/(dashboard)/projects/_components/ProjectKeysTable.tsx index 8269c843b98..080aad7b26d 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/projects/_components/ProjectKeysTable.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/projects/_components/ProjectKeysTable.tsx @@ -1,59 +1,59 @@ +"use client"; + +import { OnChangeFn, PaginationState } from "@tanstack/react-table"; +import { KeyRound } from "lucide-react"; +import { useMemo } from "react"; + import { KeyResponse } from "@/components/key_team_helpers/key_list"; -import { Empty, Table, Tooltip } from "antd"; -import type { ColumnsType } from "antd/es/table"; -import type { SpinProps } from "antd"; -import DefaultProxyAdminTag from "@/components/common_components/DefaultProxyAdminTag"; -import { DateCell } from "@/components/shared/table_cells"; +import { DataTable } from "@/components/shared/DataTable"; + +import { getProjectKeysTableColumns } from "./ProjectKeysTableColumns"; interface ProjectKeysTableProps { keys: KeyResponse[]; - loading?: boolean | SpinProps; + totalCount: number; + isLoading: boolean; + pagination: PaginationState; + onPaginationChange: OnChangeFn; } -const columns: ColumnsType = [ - { - title: "Key Name", - dataIndex: "key_alias", - key: "key_alias", - render: (alias: string | null) => alias || "—", - }, - { - title: "Owner", - key: "owner", - render: (_: unknown, record: KeyResponse) => { - const email = record.user?.user_email ?? record.user_id ?? null; - if (!email) return "—"; - return ( - - - - ); - }, - }, - { - title: "Created", - dataIndex: "created_at", - key: "created_at", - render: (date: string) => , - }, - { - title: "Last Active", - dataIndex: "last_active", - key: "last_active", - render: (date: string | null) => , - }, -]; +const PAGE_SIZE_OPTIONS = [5, 10, 25]; -export function ProjectKeysTable({ keys, loading }: ProjectKeysTableProps) { +function EmptyState() { return ( - +
+ +
+
No keys found
+
Keys created in this project will show up here.
+ + ); +} + +export function ProjectKeysTable({ + keys, + totalCount, + isLoading, + pagination, + onPaginationChange, +}: ProjectKeysTableProps) { + const columns = useMemo(() => getProjectKeysTableColumns(), []); + + return ( + }} + getRowId={(key, index) => key.token || String(index)} + paginationMode="server" + pagination={pagination} + onPaginationChange={onPaginationChange} + rowCount={totalCount} + pageSizeOptions={PAGE_SIZE_OPTIONS} + isLoading={isLoading} + loadingMessage="Loading keys…" + noDataMessage={} + size="compact" /> ); } diff --git a/ui/litellm-dashboard/src/app/(dashboard)/projects/_components/ProjectKeysTableColumns.tsx b/ui/litellm-dashboard/src/app/(dashboard)/projects/_components/ProjectKeysTableColumns.tsx new file mode 100644 index 00000000000..b04a32844ea --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/projects/_components/ProjectKeysTableColumns.tsx @@ -0,0 +1,62 @@ +"use client"; + +import { ColumnDef } from "@tanstack/react-table"; + +import DefaultProxyAdminTag from "@/components/common_components/DefaultProxyAdminTag"; +import { KeyResponse } from "@/components/key_team_helpers/key_list"; +import { CellTooltip, DateCell } from "@/components/shared/table_cells"; + +function OwnerCell({ record }: { record: KeyResponse }) { + const email = record.user?.user_email ?? record.user_id ?? null; + if (!email) return —; + return ( + + + + } + /> + ); +} + +export const getProjectKeysTableColumns = (): ColumnDef[] => [ + { + id: "key_alias", + accessorKey: "key_alias", + meta: { title: "Key Name" }, + header: "Key Name", + enableSorting: false, + cell: ({ row }) => ( + + {row.original.key_alias || "—"} + + ), + }, + { + id: "owner", + meta: { title: "Owner" }, + header: "Owner", + enableSorting: false, + cell: ({ row }) => , + }, + { + id: "created_at", + accessorKey: "created_at", + meta: { title: "Created" }, + header: "Created", + size: 130, + enableSorting: false, + cell: ({ row }) => , + }, + { + id: "last_active", + accessorKey: "last_active", + meta: { title: "Last Active" }, + header: "Last Active", + size: 130, + enableSorting: false, + cell: ({ row }) => , + }, +]; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/projects/_components/ProjectsPage.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/projects/_components/ProjectsPage.test.tsx index 73232baa03a..1c06b61698b 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/projects/_components/ProjectsPage.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/projects/_components/ProjectsPage.test.tsx @@ -1,6 +1,6 @@ import { describe, it, expect, vi, beforeEach } from "vitest"; import userEvent from "@testing-library/user-event"; -import { renderWithProviders, screen, waitFor } from "../../../../../tests/test-utils"; +import { renderWithProviders, screen, waitFor, within } from "../../../../../tests/test-utils"; import { ProjectsPage } from "./ProjectsPage"; import { ProjectResponse } from "@/app/(dashboard)/hooks/projects/useProjects"; @@ -141,7 +141,63 @@ describe("ProjectsPage", () => { it("should show the total project count in the pagination", () => { mockUseProjects.mockReturnValue({ data: mockProjects, isLoading: false }); renderWithProviders(); - expect(screen.getByText("2 projects")).toBeInTheDocument(); + expect(screen.getByTestId("pagination-range")).toHaveTextContent("Showing 1-2 of 2"); + }); + + it("should show skeleton rows while projects are loading", () => { + mockUseProjects.mockReturnValue({ data: undefined, isLoading: true }); + renderWithProviders(); + expect(screen.getAllByTestId("skeleton-row").length).toBeGreaterThan(0); + }); + + it("should show the empty state when there are no projects", () => { + mockUseProjects.mockReturnValue({ data: [], isLoading: false }); + renderWithProviders(); + expect(screen.getByText("No projects yet")).toBeInTheDocument(); + }); + + it("should show the filtered empty state when a search matches nothing", async () => { + const user = userEvent.setup(); + mockUseProjects.mockReturnValue({ data: mockProjects, isLoading: false }); + renderWithProviders(); + await user.type(screen.getByPlaceholderText(/search projects/i), "zzz-no-match"); + await waitFor(() => { + expect(screen.getByText("No matching projects")).toBeInTheDocument(); + }); + }); + + it("should sort by name when the Name header is clicked", async () => { + const user = userEvent.setup(); + mockUseProjects.mockReturnValue({ data: mockProjects, isLoading: false }); + renderWithProviders(); + + await user.click(screen.getByRole("button", { name: /^name$/i })); + let rows = screen.getAllByRole("row").slice(1); + expect(within(rows[0]).getByText("Alpha Project")).toBeInTheDocument(); + + await user.click(screen.getByRole("button", { name: /^name$/i })); + rows = screen.getAllByRole("row").slice(1); + expect(within(rows[0]).getByText("Beta Project")).toBeInTheDocument(); + }); + + it("should reset to the first page when the search text changes", async () => { + const user = userEvent.setup(); + const manyProjects = Array.from({ length: 12 }, (_, i) => ({ + ...mockProjects[0], + project_id: `proj-${i + 1}`, + project_alias: `Project ${String(i + 1).padStart(2, "0")}`, + })); + mockUseProjects.mockReturnValue({ data: manyProjects, isLoading: false }); + renderWithProviders(); + + await user.click(screen.getByTestId("pagination-next")); + expect(screen.getByTestId("pagination-page")).toHaveTextContent("Page 2 of 2"); + + await user.type(screen.getByPlaceholderText(/search projects/i), "Project 01"); + await waitFor(() => { + expect(screen.getByText("Project 01")).toBeInTheDocument(); + expect(screen.getByTestId("pagination-page")).toHaveTextContent("Page 1 of 1"); + }); }); it("should resolve team alias from the teams list in the Team column", () => { diff --git a/ui/litellm-dashboard/src/app/(dashboard)/projects/_components/ProjectsPage.tsx b/ui/litellm-dashboard/src/app/(dashboard)/projects/_components/ProjectsPage.tsx index be989229022..58c2c4c3ad8 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/projects/_components/ProjectsPage.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/projects/_components/ProjectsPage.tsx @@ -1,27 +1,12 @@ -import { useProjects, ProjectResponse } from "@/app/(dashboard)/hooks/projects/useProjects"; +import { useProjects } from "@/app/(dashboard)/hooks/projects/useProjects"; import { useTeams } from "@/app/(dashboard)/hooks/teams/useTeams"; -import { DateCell, IdCell } from "@/components/shared/table_cells"; -import { LoadingOutlined, PlusOutlined } from "@ant-design/icons"; -import { - Button, - Card, - Flex, - Input, - Layout, - Pagination, - Space, - Spin, - Table, - Tag, - theme, - Tooltip, - Typography, -} from "antd"; -import type { ColumnsType } from "antd/es/table"; -import { LayersIcon, SearchIcon } from "lucide-react"; -import { useEffect, useMemo, useState } from "react"; +import { PlusOutlined } from "@ant-design/icons"; +import { Button, Flex, Input, Layout, Space, theme, Typography } from "antd"; +import { SearchIcon } from "lucide-react"; +import { useMemo, useState } from "react"; import { CreateProjectModal } from "./ProjectModals/CreateProjectModal"; import { ProjectDetail } from "./ProjectDetailsPage"; +import { ProjectsTable } from "./ProjectsTable"; const { Title, Text } = Typography; const { Content } = Layout; @@ -34,14 +19,7 @@ export function ProjectsPage() { const [selectedProjectId, setSelectedProjectId] = useState(null); const [isCreateModalVisible, setIsCreateModalVisible] = useState(false); const [searchText, setSearchText] = useState(""); - const [currentPage, setCurrentPage] = useState(1); - const pageSize = 10; - useEffect(() => { - setCurrentPage(1); - }, [searchText]); - - // Build a team_id → team_alias lookup from the teams list const teamAliasMap = useMemo(() => { const map = new Map(); for (const team of teams ?? []) { @@ -50,7 +28,6 @@ export function ProjectsPage() { return map; }, [teams]); - // ---------- filtered data ---------- const filteredProjects = useMemo(() => { const list = projects ?? []; if (!searchText) return list; @@ -66,78 +43,6 @@ export function ProjectsPage() { }); }, [projects, searchText, teamAliasMap]); - // ---------- Ant Design columns ---------- - const columns: ColumnsType = [ - { - title: "ID", - dataIndex: "project_id", - key: "project_id", - width: 170, - render: (id: string) => , - }, - { - title: "Name", - dataIndex: "project_alias", - key: "project_alias", - sorter: (a, b) => (a.project_alias ?? "").localeCompare(b.project_alias ?? ""), - render: (alias: string | null) => alias ?? "—", - }, - { - title: "Team", - key: "team", - sorter: (a, b) => { - const aAlias = teamAliasMap.get(a.team_id ?? "") ?? ""; - const bAlias = teamAliasMap.get(b.team_id ?? "") ?? ""; - return aAlias.localeCompare(bAlias); - }, - render: (_: unknown, record: ProjectResponse) => { - if (!record.team_id) return "—"; - const alias = teamAliasMap.get(record.team_id); - if (alias) return alias; - if (isTeamsLoading) return } size="small" />; - return record.team_id; - }, - }, - { - title: "Models", - key: "models", - render: (_: unknown, record: ProjectResponse) => { - const models = record.models ?? []; - return ( - 0 ? models.join(", ") : "No models"}> - - - - {models.length} - - - - ); - }, - }, - { - title: "Status", - dataIndex: "blocked", - key: "status", - render: (blocked: boolean) => {blocked ? "Blocked" : "Active"}, - }, - { - title: "Created", - dataIndex: "created_at", - key: "created_at", - sorter: (a, b) => new Date(a.created_at).getTime() - new Date(b.created_at).getTime(), - responsive: ["lg"], - render: (date: string) => , - }, - { - title: "Updated", - dataIndex: "updated_at", - key: "updated_at", - responsive: ["xl"], - render: (date: string) => , - }, - ]; - if (selectedProjectId) { return setSelectedProjectId(null)} />; } @@ -156,34 +61,25 @@ export function ProjectsPage() { - - - } - placeholder="Search projects by name, ID, description, or team..." - style={{ maxWidth: 400 }} - value={searchText} - onChange={(e) => setSearchText(e.target.value)} - allowClear - /> - setCurrentPage(page)} - size="small" - showTotal={(total) => `${total} projects`} - showSizeChanger={false} - /> - -
+ } + placeholder="Search projects by name, ID, description, or team..." + style={{ maxWidth: 400 }} + value={searchText} + onChange={(e) => setSearchText(e.target.value)} + allowClear /> - + + + 0} + onProjectClick={setSelectedProjectId} + teamAliasMap={teamAliasMap} + isTeamsLoading={isTeamsLoading} + /> setIsCreateModalVisible(false)} /> diff --git a/ui/litellm-dashboard/src/app/(dashboard)/projects/_components/ProjectsTable.tsx b/ui/litellm-dashboard/src/app/(dashboard)/projects/_components/ProjectsTable.tsx new file mode 100644 index 00000000000..ad85019ca6d --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/projects/_components/ProjectsTable.tsx @@ -0,0 +1,70 @@ +"use client"; + +import { SortingState } from "@tanstack/react-table"; +import { FolderKanban } from "lucide-react"; +import { useMemo, useState } from "react"; + +import { ProjectResponse } from "@/app/(dashboard)/hooks/projects/useProjects"; +import { DataTable } from "@/components/shared/DataTable"; + +import { getProjectsTableColumns } from "./ProjectsTableColumns"; + +interface ProjectsTableProps { + projects: ProjectResponse[]; + isLoading: boolean; + isFiltered: boolean; + onProjectClick: (projectId: string) => void; + teamAliasMap: Map; + isTeamsLoading: boolean; +} + +const PAGE_SIZE_OPTIONS = [10, 25, 50]; + +function EmptyState({ isFiltered }: { isFiltered: boolean }) { + return ( +
+
+ +
+
+ {isFiltered ? "No matching projects" : "No projects yet"} +
+
+ {isFiltered ? "Try a different search term." : "Create a project to organize keys within your teams."} +
+
+ ); +} + +export function ProjectsTable({ + projects, + isLoading, + isFiltered, + onProjectClick, + teamAliasMap, + isTeamsLoading, +}: ProjectsTableProps) { + const [sorting, setSorting] = useState([]); + + const columns = useMemo(() => { + const deps = { onProjectClick, teamAliasMap, isTeamsLoading }; + return getProjectsTableColumns(deps); + }, [onProjectClick, teamAliasMap, isTeamsLoading]); + + return ( + project.project_id || String(index)} + sortingMode="client" + sorting={sorting} + onSortingChange={setSorting} + paginationMode="client" + pageSizeOptions={PAGE_SIZE_OPTIONS} + isLoading={isLoading} + loadingMessage="Loading projects…" + noDataMessage={} + size="compact" + /> + ); +} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/projects/_components/ProjectsTableColumns.tsx b/ui/litellm-dashboard/src/app/(dashboard)/projects/_components/ProjectsTableColumns.tsx new file mode 100644 index 00000000000..46fe259aed2 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/projects/_components/ProjectsTableColumns.tsx @@ -0,0 +1,144 @@ +"use client"; + +import { ColumnDef } from "@tanstack/react-table"; +import { LayersIcon } from "lucide-react"; + +import { ProjectResponse } from "@/app/(dashboard)/hooks/projects/useProjects"; +import { DataTableSortHeader } from "@/components/shared/DataTable"; +import { CellTooltip, DateCell, IdentityCell, StatusBadge } from "@/components/shared/table_cells"; +import { Badge } from "@/components/ui/badge"; +import { Skeleton } from "@/components/ui/skeleton"; + +function ProjectTeamCell({ + project, + teamAliasMap, + isTeamsLoading, +}: { + project: ProjectResponse; + teamAliasMap: Map; + isTeamsLoading: boolean; +}) { + if (!project.team_id) return —; + const alias = teamAliasMap.get(project.team_id); + if (alias) { + return ( + + {alias} + + ); + } + if (isTeamsLoading) return ; + return ( + + {project.team_id} + + ); +} + +function ProjectModelsCell({ project }: { project: ProjectResponse }) { + const models = project.models ?? []; + return ( + 0 ? models.join(", ") : "No models"} + trigger={ + + + {models.length} + + } + /> + ); +} + +interface ProjectsTableColumnsDeps { + onProjectClick: (projectId: string) => void; + teamAliasMap: Map; + isTeamsLoading: boolean; +} + +export const getProjectsTableColumns = ({ + onProjectClick, + teamAliasMap, + isTeamsLoading, +}: ProjectsTableColumnsDeps): ColumnDef[] => [ + { + id: "project_id", + accessorKey: "project_id", + meta: { title: "ID" }, + header: "ID", + size: 190, + enableSorting: false, + cell: ({ row }) => ( + onProjectClick(row.original.project_id)} + /> + ), + }, + { + id: "project_alias", + accessorFn: (row) => row.project_alias ?? "", + meta: { title: "Name" }, + header: ({ column }) => , + size: 200, + enableSorting: true, + cell: ({ row }) => ( + + {row.original.project_alias ?? "—"} + + ), + }, + { + id: "team", + accessorFn: (row) => teamAliasMap.get(row.team_id ?? "") ?? "", + meta: { title: "Team" }, + header: ({ column }) => , + size: 180, + enableSorting: true, + cell: ({ row }) => ( + + ), + }, + { + id: "models", + meta: { title: "Models", skeleton: "badge" }, + header: "Models", + size: 110, + enableSorting: false, + cell: ({ row }) => , + }, + { + id: "status", + accessorKey: "blocked", + meta: { title: "Status", skeleton: "badge" }, + header: "Status", + size: 110, + enableSorting: false, + cell: ({ row }) => ( + + ), + }, + { + id: "created_at", + accessorKey: "created_at", + sortingFn: "datetime", + meta: { title: "Created" }, + header: ({ column }) => , + size: 140, + enableSorting: true, + cell: ({ row }) => , + }, + { + id: "updated_at", + accessorKey: "updated_at", + meta: { title: "Updated" }, + header: "Updated", + size: 140, + enableSorting: false, + cell: ({ row }) => , + }, +]; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/prompts/_components/PromptTable.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/prompts/_components/PromptTable.test.tsx new file mode 100644 index 00000000000..52efd6407f8 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/prompts/_components/PromptTable.test.tsx @@ -0,0 +1,104 @@ +import { render, screen, within } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import { beforeEach, describe, expect, it, vi } from "vitest"; + +import { PromptSpec } from "@/components/networking"; + +import PromptTable from "./PromptTable"; + +vi.mock("@/components/networking", () => ({ + modelHubCall: vi.fn().mockResolvedValue({ data: [] }), +})); + +const mockPrompts: PromptSpec[] = [ + { + prompt_id: "prompt-newer", + litellm_params: { prompt_id: "prompt-newer" }, + prompt_info: { prompt_type: "dotprompt" }, + created_at: "2025-01-15T10:30:00Z", + updated_at: "2025-01-15T11:00:00Z", + environment: "production", + created_by: "user-1", + }, + { + prompt_id: "prompt-older", + litellm_params: { prompt_id: "prompt-older" }, + prompt_info: { prompt_type: "dotprompt" }, + created_at: "2024-01-10T09:15:00Z", + updated_at: "2024-01-12T14:20:00Z", + }, +]; + +const mockOnPromptClick = vi.fn(); +const mockOnDeleteClick = vi.fn(); + +const defaultProps = { + promptsList: mockPrompts, + isLoading: false, + onPromptClick: mockOnPromptClick, + onDeleteClick: mockOnDeleteClick, + accessToken: null, + isAdmin: true, +}; + +describe("PromptTable", () => { + beforeEach(() => { + vi.clearAllMocks(); + }); + + it("should render every column header", () => { + render(); + for (const header of ["Prompt ID", "Model", "Created At", "Updated At", "Environment", "Created By", "Type"]) { + expect(screen.getByText(header)).toBeInTheDocument(); + } + }); + + it("should display the empty state when data is empty", () => { + render(); + expect(screen.getByText("No prompts yet")).toBeInTheDocument(); + }); + + it("should sort by created date descending by default", () => { + render(); + const rows = screen.getAllByRole("row").slice(1); + expect(within(rows[0]).getByText("prompt-newer")).toBeInTheDocument(); + expect(within(rows[1]).getByText("prompt-older")).toBeInTheDocument(); + }); + + it("should call onPromptClick when the prompt ID is clicked", async () => { + const user = userEvent.setup(); + render(); + await user.click(screen.getByRole("button", { name: "prompt-newer" })); + expect(mockOnPromptClick).toHaveBeenCalledWith("prompt-newer"); + }); + + it("should label the environment and default missing environments to development", () => { + render(); + expect(screen.getByText("production")).toBeInTheDocument(); + expect(screen.getByText("development")).toBeInTheDocument(); + }); + + it("should delete a prompt through the actions menu when admin", async () => { + const user = userEvent.setup(); + render(); + await user.click(screen.getByTestId("prompt-actions-prompt-newer")); + await user.click(await screen.findByTestId("prompt-action-delete")); + expect(mockOnDeleteClick).toHaveBeenCalledWith("prompt-newer", "prompt-newer"); + }); + + it("should copy the prompt ID through the actions menu", async () => { + const user = userEvent.setup(); + render(); + await user.click(screen.getByTestId("prompt-actions-prompt-newer")); + await user.click(await screen.findByTestId("prompt-action-copy")); + expect(await window.navigator.clipboard.readText()).toBe("prompt-newer"); + }); + + it("should hide the delete action for non-admins but keep copy available", async () => { + const user = userEvent.setup(); + render(); + await user.click(screen.getByTestId("prompt-actions-prompt-newer")); + expect(await screen.findByTestId("prompt-action-copy")).toBeInTheDocument(); + expect(screen.queryByTestId("prompt-action-delete")).not.toBeInTheDocument(); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/prompts/_components/PromptTable.tsx b/ui/litellm-dashboard/src/app/(dashboard)/prompts/_components/PromptTable.tsx new file mode 100644 index 00000000000..47d4f64f254 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/prompts/_components/PromptTable.tsx @@ -0,0 +1,89 @@ +"use client"; + +import { SortingState } from "@tanstack/react-table"; +import { Inbox } from "lucide-react"; +import React, { useEffect, useMemo, useState } from "react"; + +import { DataTable } from "@/components/shared/DataTable"; +import { modelHubCall, PromptSpec } from "@/components/networking"; + +import { getPromptTableColumns } from "./PromptTableColumns"; +import { ModelGroupInfo } from "./prompt_utils"; + +interface PromptTableProps { + promptsList: PromptSpec[]; + isLoading: boolean; + onPromptClick?: (id: string) => void; + onDeleteClick?: (id: string, name: string) => void; + accessToken: string | null; + isAdmin: boolean; +} + +const DEFAULT_SORTING: SortingState = [{ id: "created_at", desc: true }]; + +function EmptyState() { + return ( +
+
+ +
+
No prompts yet
+
Add a prompt to start managing reusable templates.
+
+ ); +} + +const PromptTable: React.FC = ({ + promptsList, + isLoading, + onPromptClick, + onDeleteClick, + accessToken, + isAdmin, +}) => { + const [sorting, setSorting] = useState(DEFAULT_SORTING); + const [modelHubData, setModelHubData] = useState>(new Map()); + + useEffect(() => { + const fetchModelHubData = async () => { + if (!accessToken) return; + + try { + const response = await modelHubCall(accessToken); + if (response?.data) { + const modelMap = new Map(); + response.data.forEach((model: ModelGroupInfo) => { + modelMap.set(model.model_group, model); + }); + setModelHubData(modelMap); + } + } catch (error) { + console.error("Error fetching model hub data:", error); + } + }; + + fetchModelHubData(); + }, [accessToken]); + + const columns = useMemo( + () => getPromptTableColumns({ modelHubData, isAdmin, onPromptClick, onDeleteClick }), + [modelHubData, isAdmin, onPromptClick, onDeleteClick], + ); + + return ( + prompt.prompt_id || String(index)} + sortingMode="client" + sorting={sorting} + onSortingChange={setSorting} + isLoading={isLoading} + loadingMessage="Loading prompts…" + noDataMessage={} + size="compact" + /> + ); +}; + +export default PromptTable; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/prompts/_components/PromptTableColumns.tsx b/ui/litellm-dashboard/src/app/(dashboard)/prompts/_components/PromptTableColumns.tsx new file mode 100644 index 00000000000..ae584ef6df6 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/prompts/_components/PromptTableColumns.tsx @@ -0,0 +1,220 @@ +"use client"; + +import { ColumnDef } from "@tanstack/react-table"; +import { Copy, MoreHorizontal, Trash2 } from "lucide-react"; + +import { DataTableSortHeader } from "@/components/shared/DataTable"; +import { CellTooltip, DateCell, IdentityCell, StatusBadge, StatusTone } from "@/components/shared/table_cells"; +import { PromptSpec } from "@/components/networking"; +import { getProviderLogoAndName } from "@/components/provider_info_helpers"; +import { buttonVariants } from "@/components/ui/button"; +import { + DropdownMenu, + DropdownMenuContent, + DropdownMenuItem, + DropdownMenuSeparator, + DropdownMenuTrigger, +} from "@/components/ui/dropdown-menu"; +import { cn } from "@/lib/cva.config"; +import { copyToClipboard } from "@/utils/dataUtils"; + +import { extractModel, getProviderFromModelHub, ModelGroupInfo } from "./prompt_utils"; + +const ENVIRONMENT_TONE: Record = { + production: "error", + staging: "warning", + development: "success", +}; + +function PromptModelCell({ prompt, modelHubData }: { prompt: PromptSpec; modelHubData: Map }) { + const model = extractModel(prompt); + if (!model) { + return -; + } + + const provider = getProviderFromModelHub(model, modelHubData); + const { logo } = provider ? getProviderLogoAndName(provider) : { logo: "" }; + + return ( + + {logo ? ( + { + (event.currentTarget as HTMLImageElement).style.display = "none"; + }} + /> + ) : ( + + {provider?.charAt(0) || "-"} + + )} + {model} + + } + /> + ); +} + +interface PromptRowActionsProps { + prompt: PromptSpec; + isAdmin: boolean; + onDeleteClick?: (id: string, name: string) => void; +} + +function PromptRowActions({ prompt, isAdmin, onDeleteClick }: PromptRowActionsProps) { + return ( + + + + + + void copyToClipboard(prompt.prompt_id, "Prompt ID copied")} + > + + Copy prompt ID + + {isAdmin && ( + <> + + onDeleteClick?.(prompt.prompt_id, prompt.prompt_id || "Unknown Prompt")} + > + + Delete + + + )} + + + ); +} + +interface PromptTableColumnsDeps { + modelHubData: Map; + isAdmin: boolean; + onPromptClick?: (id: string) => void; + onDeleteClick?: (id: string, name: string) => void; +} + +export const getPromptTableColumns = ({ + modelHubData, + isAdmin, + onPromptClick, + onDeleteClick, +}: PromptTableColumnsDeps): ColumnDef[] => [ + { + id: "prompt_id", + accessorKey: "prompt_id", + meta: { title: "Prompt ID" }, + header: ({ column }) => , + size: 220, + enableSorting: true, + cell: ({ row }) => ( + onPromptClick(row.original.prompt_id) : undefined} + /> + ), + }, + { + id: "model", + meta: { title: "Model" }, + header: "Model", + size: 200, + enableSorting: false, + cell: ({ row }) => , + }, + { + id: "created_at", + accessorKey: "created_at", + sortingFn: "datetime", + meta: { title: "Created At" }, + header: ({ column }) => , + size: 160, + enableSorting: true, + cell: ({ row }) => , + }, + { + id: "updated_at", + accessorKey: "updated_at", + sortingFn: "datetime", + meta: { title: "Updated At" }, + header: ({ column }) => , + size: 160, + enableSorting: true, + cell: ({ row }) => , + }, + { + id: "environment", + accessorKey: "environment", + meta: { title: "Environment", skeleton: "badge" }, + header: "Environment", + size: 130, + enableSorting: false, + cell: ({ row }) => { + const environment = row.original.environment || "development"; + return ; + }, + }, + { + id: "created_by", + accessorKey: "created_by", + meta: { title: "Created By" }, + header: "Created By", + size: 160, + enableSorting: false, + cell: ({ row }) => { + const createdBy = row.original.created_by; + return ( + + {createdBy || "-"} + + ); + }, + }, + { + id: "prompt_type", + accessorKey: "prompt_info.prompt_type", + meta: { title: "Type" }, + header: "Type", + size: 140, + enableSorting: false, + cell: ({ row }) => { + const promptType = row.original.prompt_info.prompt_type; + return ( + + {promptType} + + ); + }, + }, + { + id: "actions", + meta: { className: "text-right", headerClassName: "text-right" }, + header: () => Actions, + size: 64, + enableSorting: false, + enableHiding: false, + cell: ({ row }) => ( +
+ +
+ ), + }, +]; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/prompts/_components/index.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/prompts/_components/index.test.tsx new file mode 100644 index 00000000000..99c58e2b98f --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/prompts/_components/index.test.tsx @@ -0,0 +1,51 @@ +import { render, screen } from "@testing-library/react"; +import { beforeEach, describe, expect, it, vi } from "vitest"; + +import { getPromptsList } from "@/components/networking"; + +import PromptsPanel from "./index"; + +vi.mock("@/components/networking", () => ({ + getPromptsList: vi.fn(), + deletePromptCall: vi.fn(), +})); + +vi.mock("./PromptTable", () => ({ + __esModule: true, + default: ({ isLoading }: { isLoading: boolean }) => ( +
{isLoading ? "table-loading" : "table-loaded"}
+ ), +})); + +vi.mock("./prompt_info", () => ({ __esModule: true, default: () => null })); +vi.mock("./add_prompt_form", () => ({ __esModule: true, default: () => null })); +vi.mock("./prompt_editor_view", () => ({ __esModule: true, default: () => null })); + +const mockGetPromptsList = vi.mocked(getPromptsList); + +describe("PromptsPanel loading state", () => { + beforeEach(() => { + vi.clearAllMocks(); + }); + + it("should resolve the loading state when accessToken is null instead of showing the skeleton forever", async () => { + render(); + expect(await screen.findByText("table-loaded")).toBeInTheDocument(); + expect(mockGetPromptsList).not.toHaveBeenCalled(); + }); + + it("should show the loading state until the prompt fetch settles", async () => { + let resolveFetch: (value: { prompts: never[] }) => void = () => {}; + mockGetPromptsList.mockReturnValue( + new Promise((resolve) => { + resolveFetch = resolve; + }), + ); + render(); + expect(screen.getByText("table-loading")).toBeInTheDocument(); + + resolveFetch({ prompts: [] }); + expect(await screen.findByText("table-loaded")).toBeInTheDocument(); + expect(mockGetPromptsList).toHaveBeenCalledWith("sk-test", undefined); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/prompts/_components/index.tsx b/ui/litellm-dashboard/src/app/(dashboard)/prompts/_components/index.tsx index 6e0e13d5181..de461ebd86d 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/prompts/_components/index.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/prompts/_components/index.tsx @@ -3,7 +3,7 @@ import React, { useState, useEffect } from "react"; import { Button } from "@tremor/react"; import { Modal, Select } from "antd"; import { getPromptsList, PromptSpec, ListPromptsResponse, deletePromptCall } from "@/components/networking"; -import PromptTable from "./prompt_table"; +import PromptTable from "./PromptTable"; import PromptInfoView from "./prompt_info"; import AddPromptForm from "./add_prompt_form"; import PromptEditorView from "./prompt_editor_view"; @@ -17,7 +17,7 @@ interface PromptsProps { const PromptsPanel: React.FC = ({ accessToken, userRole }) => { const [promptsList, setPromptsList] = useState([]); - const [isLoading, setIsLoading] = useState(false); + const [isLoading, setIsLoading] = useState(true); const [selectedEnvironment, setSelectedEnvironment] = useState(undefined); const [selectedPromptId, setSelectedPromptId] = useState(null); const [isAddModalVisible, setIsAddModalVisible] = useState(false); @@ -32,6 +32,7 @@ const PromptsPanel: React.FC = ({ accessToken, userRole }) => { const fetchPrompts = async () => { if (!accessToken) { + setIsLoading(false); return; } diff --git a/ui/litellm-dashboard/src/app/(dashboard)/prompts/_components/prompt_table.tsx b/ui/litellm-dashboard/src/app/(dashboard)/prompts/_components/prompt_table.tsx deleted file mode 100644 index 51a03d19e83..00000000000 --- a/ui/litellm-dashboard/src/app/(dashboard)/prompts/_components/prompt_table.tsx +++ /dev/null @@ -1,284 +0,0 @@ -import React, { useState, useEffect } from "react"; -import { Table, TableBody, TableCell, TableHead, TableHeaderCell, TableRow, Button } from "@tremor/react"; -import { SwitchVerticalIcon, ChevronUpIcon, ChevronDownIcon, TrashIcon } from "@heroicons/react/outline"; -import { Tooltip } from "antd"; -import { PromptSpec, modelHubCall } from "@/components/networking"; -import { DateCell, IdCell } from "@/components/shared/table_cells"; -import { - ColumnDef, - flexRender, - getCoreRowModel, - getSortedRowModel, - SortingState, - useReactTable, -} from "@tanstack/react-table"; -import { getProviderLogoAndName } from "@/components/provider_info_helpers"; -import { extractModel, getProviderFromModelHub } from "./prompt_utils"; - -interface PromptTableProps { - promptsList: PromptSpec[]; - isLoading: boolean; - onPromptClick?: (id: string) => void; - onDeleteClick?: (id: string, name: string) => void; - accessToken: string | null; - isAdmin: boolean; -} - -interface ModelGroupInfo { - model_group: string; - providers: string[]; - [key: string]: any; -} - -const PromptTable: React.FC = ({ - promptsList, - isLoading, - onPromptClick, - onDeleteClick, - accessToken, - isAdmin, -}) => { - const [sorting, setSorting] = useState([{ id: "created_at", desc: true }]); - const [modelHubData, setModelHubData] = useState>(new Map()); - - useEffect(() => { - const fetchModelHubData = async () => { - if (!accessToken) return; - - try { - const response = await modelHubCall(accessToken); - if (response?.data) { - const modelMap = new Map(); - response.data.forEach((model: ModelGroupInfo) => { - modelMap.set(model.model_group, model); - }); - setModelHubData(modelMap); - } - } catch (error) { - console.error("Error fetching model hub data:", error); - } - }; - - fetchModelHubData(); - }, [accessToken]); - - const columns: ColumnDef[] = [ - { - header: "Prompt ID", - accessorKey: "prompt_id", - cell: (info: any) => , - }, - { - header: "Model", - accessorKey: "model", - cell: ({ row }) => { - const prompt = row.original; - const model = extractModel(prompt); - - if (!model) { - return -; - } - - const provider = getProviderFromModelHub(model, modelHubData); - const { logo } = getProviderLogoAndName(provider || ""); - - return ( - -
- {/* Provider Icon */} -
- {provider && logo ? ( - {`${provider} { - const target = e.currentTarget as HTMLImageElement; - const parent = target.parentElement; - if (!parent || !parent.contains(target)) { - return; - } - - try { - const fallbackDiv = document.createElement("div"); - fallbackDiv.className = - "w-4 h-4 rounded-full bg-gray-200 flex items-center justify-center text-xs"; - fallbackDiv.textContent = provider?.charAt(0) || "-"; - parent.replaceChild(fallbackDiv, target); - } catch (error) { - console.error("Failed to replace provider logo fallback:", error); - } - }} - /> - ) : ( -
-
- )} -
- - {/* Model Name */} - {model} -
-
- ); - }, - }, - { - header: "Created At", - accessorKey: "created_at", - cell: ({ row }) => , - }, - { - header: "Updated At", - accessorKey: "updated_at", - cell: ({ row }) => , - }, - { - header: "Environment", - accessorKey: "environment", - cell: ({ row }) => { - const prompt = row.original; - const env = prompt.environment || "development"; - const colorMap: Record = { - production: "text-red-600 bg-red-50", - staging: "text-yellow-600 bg-yellow-50", - development: "text-green-600 bg-green-50", - }; - return ( - {env} - ); - }, - }, - { - header: "Created By", - accessorKey: "created_by", - cell: ({ row }) => { - const prompt = row.original; - return {prompt.created_by || "-"}; - }, - }, - { - header: "Type", - accessorKey: "prompt_info.prompt_type", - cell: ({ row }) => { - const prompt = row.original; - return ( - - {prompt.prompt_info.prompt_type} - - ); - }, - }, - ...(isAdmin - ? [ - { - header: "Actions", - id: "actions", - enableSorting: false, - cell: ({ row }: any) => { - const prompt = row.original; - const promptName = prompt.prompt_id || "Unknown Prompt"; - - return ( -
- -
- ); - }, - }, - ] - : []), - ]; - - const table = useReactTable({ - data: promptsList, - columns, - state: { - sorting, - }, - onSortingChange: setSorting, - getCoreRowModel: getCoreRowModel(), - getSortedRowModel: getSortedRowModel(), - enableSorting: true, - }); - - return ( -
-
-
- - {table.getHeaderGroups().map((headerGroup) => ( - - {headerGroup.headers.map((header) => ( - -
-
- {header.isPlaceholder ? null : flexRender(header.column.columnDef.header, header.getContext())} -
-
- {header.column.getIsSorted() ? ( - { - asc: , - desc: , - }[header.column.getIsSorted() as string] - ) : ( - - )} -
-
-
- ))} -
- ))} -
- - {isLoading ? ( - - -
-

Loading...

-
-
-
- ) : promptsList.length > 0 ? ( - table.getRowModel().rows.map((row) => ( - - {row.getVisibleCells().map((cell) => ( - - {flexRender(cell.column.columnDef.cell, cell.getContext())} - - ))} - - )) - ) : ( - - -
-

No prompts found

-
-
-
- )} -
-
- - - ); -}; - -export default PromptTable; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/prompts/_components/prompt_utils.tsx b/ui/litellm-dashboard/src/app/(dashboard)/prompts/_components/prompt_utils.tsx index 6e0bd096af0..5c21a7d9ac6 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/prompts/_components/prompt_utils.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/prompts/_components/prompt_utils.tsx @@ -1,7 +1,7 @@ import { PromptSpec } from "@/components/networking"; import { getVersionNumber } from "./prompt_editor_view/utils"; -interface ModelGroupInfo { +export interface ModelGroupInfo { model_group: string; providers: string[]; [key: string]: any; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/search-tools/_components/SearchToolColumn.tsx b/ui/litellm-dashboard/src/app/(dashboard)/search-tools/_components/SearchToolColumn.tsx deleted file mode 100644 index 3d9fdb2866e..00000000000 --- a/ui/litellm-dashboard/src/app/(dashboard)/search-tools/_components/SearchToolColumn.tsx +++ /dev/null @@ -1,104 +0,0 @@ -import { Tag } from "antd"; -import { ColumnsType } from "antd/es/table"; -import TableIconActionButton from "@/components/common_components/IconActionButton/TableIconActionButtons/TableIconActionButton"; -import { DateCell, IdCell } from "@/components/shared/table_cells"; -import { SearchTool } from "./types"; - -export const searchToolColumns = ( - onView: (searchToolId: string) => void, - onEdit: (searchToolId: string) => void, - onDelete: (searchToolId: string) => void, - availableProviders: Array<{ provider_name: string; ui_friendly_name: string }>, -): ColumnsType => [ - { - title: "Search Tool ID", - dataIndex: "search_tool_id", - key: "search_tool_id", - render: (_, tool) => { - const isFromConfig = tool.is_from_config; - - if (isFromConfig) { - return -; - } - - return ; - }, - }, - { - title: "Name", - dataIndex: "search_tool_name", - key: "search_tool_name", - render: (name: string) => {name}, - }, - { - title: "Provider", - key: "provider", - render: (_, tool) => { - const provider = tool.litellm_params.search_provider; - const providerInfo = availableProviders.find((p) => p.provider_name === provider); - const displayName = providerInfo?.ui_friendly_name || provider; - - return {displayName}; - }, - }, - { - title: "Created At", - dataIndex: "created_at", - key: "created_at", - render: (_, tool) => { - return ; - }, - }, - { - title: "Updated At", - dataIndex: "updated_at", - key: "updated_at", - render: (_, tool) => { - return ; - }, - }, - { - title: "Source", - key: "source", - render: (_, tool) => { - const isFromConfig = tool.is_from_config ?? false; - - return {isFromConfig ? "Config" : "DB"}; - }, - }, - { - title: "Actions", - key: "actions", - render: (_, tool) => { - const toolId = tool.search_tool_id; - const isFromConfig = tool.is_from_config ?? false; - - return ( -
- { - if (toolId && !isFromConfig) { - onEdit(toolId); - } - }} - /> - { - if (toolId && !isFromConfig) { - onDelete(toolId); - } - }} - /> -
- ); - }, - }, -]; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/search-tools/_components/SearchToolTable.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/search-tools/_components/SearchToolTable.test.tsx new file mode 100644 index 00000000000..c4c05f97dce --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/search-tools/_components/SearchToolTable.test.tsx @@ -0,0 +1,117 @@ +import { screen, within } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import { beforeEach, describe, expect, it, vi } from "vitest"; +import { renderWithProviders } from "@/../tests/test-utils"; +import SearchToolTable from "./SearchToolTable"; +import { AvailableSearchProvider, SearchTool } from "./types"; + +const makeSearchTool = (overrides: Partial = {}): SearchTool => ({ + search_tool_id: "tool-1", + search_tool_name: "Perplexity Search", + litellm_params: { + search_provider: "perplexity", + }, + created_at: "2024-01-15T10:30:00Z", + updated_at: "2024-01-16T10:30:00Z", + ...overrides, +}); + +const availableProviders: AvailableSearchProvider[] = [ + { provider_name: "perplexity", ui_friendly_name: "Perplexity AI" }, +]; + +const defaultProps = { + searchTools: [makeSearchTool()], + isLoading: false, + availableProviders, + onView: vi.fn(), + onEdit: vi.fn(), + onDelete: vi.fn(), +}; + +describe("SearchToolTable", () => { + beforeEach(() => { + vi.clearAllMocks(); + }); + + it("should display search tool information with the friendly provider name", () => { + renderWithProviders(); + expect(screen.getByText("Perplexity Search")).toBeInTheDocument(); + expect(screen.getByText("tool-1")).toBeInTheDocument(); + expect(screen.getByText("Perplexity AI")).toBeInTheDocument(); + expect(screen.getByText("DB")).toBeInTheDocument(); + }); + + it("should call onView when the search tool ID is clicked", async () => { + const user = userEvent.setup(); + renderWithProviders(); + await user.click(screen.getByRole("button", { name: /tool-1/ })); + expect(defaultProps.onView).toHaveBeenCalledWith("tool-1"); + }); + + it("should call onEdit and onDelete from the actions menu for a DB tool", async () => { + const user = userEvent.setup(); + renderWithProviders(); + + await user.click(screen.getByTestId("search-tool-actions-tool-1")); + await user.click(await screen.findByTestId("search-tool-action-edit")); + expect(defaultProps.onEdit).toHaveBeenCalledWith("tool-1"); + + await user.click(screen.getByTestId("search-tool-actions-tool-1")); + await user.click(await screen.findByTestId("search-tool-action-delete")); + expect(defaultProps.onDelete).toHaveBeenCalledWith("tool-1"); + }); + + it("should show a dash instead of a clickable ID for config tools", () => { + const configTool = makeSearchTool({ search_tool_id: "config-tool", is_from_config: true }); + renderWithProviders(); + expect(screen.queryByRole("button", { name: /config-tool/ })).not.toBeInTheDocument(); + }); + + it("should disable Edit and Delete for config tools and suppress their callbacks", async () => { + const user = userEvent.setup(); + const configTool = makeSearchTool({ search_tool_id: "config-tool", is_from_config: true }); + renderWithProviders(); + + await user.click(screen.getByTestId("search-tool-actions-config-tool")); + + const editItem = await screen.findByTestId("search-tool-action-edit"); + const deleteItem = await screen.findByTestId("search-tool-action-delete"); + expect(editItem).toHaveAttribute("data-disabled"); + expect(deleteItem).toHaveAttribute("data-disabled"); + + await user.click(editItem); + await user.click(deleteItem); + expect(defaultProps.onEdit).not.toHaveBeenCalled(); + expect(defaultProps.onDelete).not.toHaveBeenCalled(); + }); + + it("should sort tools by created_at descending by default", () => { + const tools = [ + makeSearchTool({ + search_tool_id: "tool-old", + search_tool_name: "older-tool", + created_at: "2024-01-01T00:00:00Z", + }), + makeSearchTool({ + search_tool_id: "tool-new", + search_tool_name: "newer-tool", + created_at: "2024-06-01T00:00:00Z", + }), + ]; + renderWithProviders(); + const rows = screen.getAllByRole("row").slice(1); + expect(within(rows[0]).getByText("newer-tool")).toBeInTheDocument(); + expect(within(rows[1]).getByText("older-tool")).toBeInTheDocument(); + }); + + it("should show skeleton rows when loading", () => { + renderWithProviders(); + expect(screen.getAllByTestId("skeleton-row").length).toBeGreaterThan(0); + }); + + it("should show the empty state when there are no search tools", () => { + renderWithProviders(); + expect(screen.getByText("No search tools configured")).toBeInTheDocument(); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/search-tools/_components/SearchToolTable.tsx b/ui/litellm-dashboard/src/app/(dashboard)/search-tools/_components/SearchToolTable.tsx new file mode 100644 index 00000000000..70fc6a376df --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/search-tools/_components/SearchToolTable.tsx @@ -0,0 +1,66 @@ +"use client"; + +import { SortingState } from "@tanstack/react-table"; +import { Inbox } from "lucide-react"; +import React, { useMemo, useState } from "react"; + +import { DataTable } from "@/components/shared/DataTable"; + +import { getSearchToolTableColumns, searchToolKey } from "./SearchToolTableColumns"; +import { AvailableSearchProvider, SearchTool } from "./types"; + +interface SearchToolTableProps { + searchTools: SearchTool[]; + isLoading: boolean; + availableProviders: AvailableSearchProvider[]; + onView: (searchToolId: string) => void; + onEdit: (searchToolId: string) => void; + onDelete: (searchToolId: string) => void; +} + +const DEFAULT_SORTING: SortingState = [{ id: "created_at", desc: true }]; + +function EmptyState() { + return ( +
+
+ +
+
No search tools configured
+
Add a search tool to enable web search for your models.
+
+ ); +} + +const SearchToolTable: React.FC = ({ + searchTools, + isLoading, + availableProviders, + onView, + onEdit, + onDelete, +}) => { + const [sorting, setSorting] = useState(DEFAULT_SORTING); + + const columns = useMemo(() => { + const deps = { availableProviders, onView, onEdit, onDelete }; + return getSearchToolTableColumns(deps); + }, [availableProviders, onView, onEdit, onDelete]); + + return ( + searchToolKey(tool) || String(index)} + sortingMode="client" + sorting={sorting} + onSortingChange={setSorting} + isLoading={isLoading} + loadingMessage="Loading search tools…" + noDataMessage={} + size="compact" + /> + ); +}; + +export default SearchToolTable; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/search-tools/_components/SearchToolTableColumns.tsx b/ui/litellm-dashboard/src/app/(dashboard)/search-tools/_components/SearchToolTableColumns.tsx new file mode 100644 index 00000000000..550bf3dc7bd --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/search-tools/_components/SearchToolTableColumns.tsx @@ -0,0 +1,168 @@ +"use client"; + +import { ColumnDef } from "@tanstack/react-table"; +import { MoreHorizontal, Pencil, Trash2 } from "lucide-react"; + +import { DataTableSortHeader } from "@/components/shared/DataTable"; +import { DateCell, IdentityCell, StatusBadge } from "@/components/shared/table_cells"; +import { buttonVariants } from "@/components/ui/button"; +import { + DropdownMenu, + DropdownMenuContent, + DropdownMenuItem, + DropdownMenuSeparator, + DropdownMenuTrigger, +} from "@/components/ui/dropdown-menu"; +import { cn } from "@/lib/cva.config"; + +import { AvailableSearchProvider, SearchTool } from "./types"; + +const CONFIG_EDIT_HINT = "Config search tools cannot be edited on the dashboard. Please edit the config file."; +const CONFIG_DELETE_HINT = "Config search tools cannot be deleted on the dashboard. Please edit the config file."; + +export const searchToolKey = (tool: SearchTool): string => tool.search_tool_id || tool.search_tool_name; + +interface SearchToolRowActionsProps { + tool: SearchTool; + onEdit: (searchToolId: string) => void; + onDelete: (searchToolId: string) => void; +} + +function SearchToolRowActions({ tool, onEdit, onDelete }: SearchToolRowActionsProps) { + const isFromConfig = tool.is_from_config ?? false; + const toolId = tool.search_tool_id; + + return ( + + + + + + toolId && onEdit(toolId)} + > + + Edit search tool + + + toolId && onDelete(toolId)} + > + + Delete search tool + + + + ); +} + +interface SearchToolTableColumnsDeps { + availableProviders: AvailableSearchProvider[]; + onView: (searchToolId: string) => void; + onEdit: (searchToolId: string) => void; + onDelete: (searchToolId: string) => void; +} + +export const getSearchToolTableColumns = ({ + availableProviders, + onView, + onEdit, + onDelete, +}: SearchToolTableColumnsDeps): ColumnDef[] => [ + { + id: "search_tool_id", + accessorKey: "search_tool_id", + meta: { title: "Search Tool ID" }, + header: ({ column }) => , + size: 200, + enableSorting: true, + cell: ({ row }) => { + const tool = row.original; + const toolId = tool.search_tool_id; + if (tool.is_from_config || !toolId) { + return -; + } + return ( + onView(toolId)} /> + ); + }, + }, + { + id: "search_tool_name", + accessorKey: "search_tool_name", + meta: { title: "Name" }, + header: ({ column }) => , + size: 200, + enableSorting: true, + cell: ({ row }) => ( + + {row.original.search_tool_name || "-"} + + ), + }, + { + id: "provider", + meta: { title: "Provider" }, + header: "Provider", + size: 160, + enableSorting: false, + cell: ({ row }) => { + const provider = row.original.litellm_params.search_provider; + const providerInfo = availableProviders.find((candidate) => candidate.provider_name === provider); + return {providerInfo?.ui_friendly_name || provider}; + }, + }, + { + id: "created_at", + accessorKey: "created_at", + meta: { title: "Created At" }, + header: ({ column }) => , + size: 130, + enableSorting: true, + cell: ({ row }) => , + }, + { + id: "updated_at", + accessorKey: "updated_at", + meta: { title: "Updated At" }, + header: ({ column }) => , + size: 130, + enableSorting: true, + cell: ({ row }) => , + }, + { + id: "source", + meta: { title: "Source", skeleton: "badge" }, + header: "Source", + size: 100, + enableSorting: false, + cell: ({ row }) => { + const isFromConfig = row.original.is_from_config ?? false; + return ; + }, + }, + { + id: "actions", + meta: { className: "text-right", headerClassName: "text-right" }, + header: () => Actions, + size: 64, + enableSorting: false, + enableHiding: false, + cell: ({ row }) => ( +
+ +
+ ), + }, +]; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/search-tools/_components/SearchTools.tsx b/ui/litellm-dashboard/src/app/(dashboard)/search-tools/_components/SearchTools.tsx index 65aee211271..c9e2a46a861 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/search-tools/_components/SearchTools.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/search-tools/_components/SearchTools.tsx @@ -1,8 +1,7 @@ import { isAdminRole } from "@/utils/roles"; -import { LoadingOutlined } from "@ant-design/icons"; import { useQuery } from "@tanstack/react-query"; import { Button, Text, Title } from "@tremor/react"; -import { Form, Input, Modal, Select, Spin, Table } from "antd"; +import { Form, Input, Modal, Select } from "antd"; import React, { useState } from "react"; import DeleteResourceModal from "@/components/common_components/DeleteResourceModal"; import NotificationsManager from "@/components/molecules/notifications_manager"; @@ -13,7 +12,7 @@ import { updateSearchTool, } from "@/components/networking"; import CreateSearchTool from "./CreateSearchTools"; -import { searchToolColumns } from "./SearchToolColumn"; +import SearchToolTable from "./SearchToolTable"; import { SearchToolView } from "./SearchToolView"; import { AvailableSearchProvider, SearchTool } from "./types"; @@ -58,34 +57,29 @@ const SearchTools: React.FC = ({ accessToken, userRole, userID const [isEditModalVisible, setEditModalVisible] = useState(false); const [form] = Form.useForm(); - const columns = React.useMemo( - () => - searchToolColumns( - (toolId: string) => { - setSelectedToolId(toolId); - setEditTool(false); - }, - (toolId: string) => { - const tool = searchTools?.find((t) => t.search_tool_id === toolId); - if (tool) { - form.setFieldsValue({ - search_tool_name: tool.search_tool_name, - search_provider: tool.litellm_params.search_provider, - api_key: tool.litellm_params.api_key, - api_base: tool.litellm_params.api_base, - timeout: tool.litellm_params.timeout, - max_retries: tool.litellm_params.max_retries, - description: tool.search_tool_info?.description, - }); - setSelectedToolId(toolId); - setEditModalVisible(true); - } - }, - handleDelete, - availableProviders, - ), - [availableProviders, searchTools, form], - ); + const handleView = (toolId: string) => { + setSelectedToolId(toolId); + setEditTool(false); + }; + + const handleEditOpen = (toolId: string) => { + const tool = searchTools?.find((t) => t.search_tool_id === toolId); + if (!tool) { + return; + } + const editFormValues = { + search_tool_name: tool.search_tool_name, + search_provider: tool.litellm_params.search_provider, + api_key: tool.litellm_params.api_key, + api_base: tool.litellm_params.api_base, + timeout: tool.litellm_params.timeout, + max_retries: tool.litellm_params.max_retries, + description: tool.search_tool_info?.description, + }; + form.setFieldsValue(editFormValues); + setSelectedToolId(toolId); + setEditModalVisible(true); + }; function handleDelete(toolId: string) { setToolToDelete(toolId); @@ -220,19 +214,14 @@ const SearchTools: React.FC = ({ accessToken, userRole, userID /> ) : (
- } size="large"> - record.search_tool_id || record.search_tool_name} - pagination={false} - locale={{ - emptyText: "No search tools configured", - }} - size="small" - /> - + ); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/skills/_components/ClaudeCodePluginsPanel.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/skills/_components/ClaudeCodePluginsPanel.test.tsx new file mode 100644 index 00000000000..52f3dc21b7a --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/skills/_components/ClaudeCodePluginsPanel.test.tsx @@ -0,0 +1,50 @@ +import { render, screen } from "@testing-library/react"; +import { beforeEach, describe, expect, it, vi } from "vitest"; + +import { getClaudeCodePluginsList } from "@/components/networking"; + +import ClaudeCodePluginsPanel from "./ClaudeCodePluginsPanel"; + +vi.mock("@/components/networking", () => ({ + getClaudeCodePluginsList: vi.fn(), + deleteClaudeCodePlugin: vi.fn(), +})); + +vi.mock("./PluginTable", () => ({ + __esModule: true, + default: ({ isLoading }: { isLoading: boolean }) => ( +
{isLoading ? "table-loading" : "table-loaded"}
+ ), +})); + +vi.mock("./add_plugin_form", () => ({ __esModule: true, default: () => null })); +vi.mock("@/components/claude_code_plugins/skill_detail", () => ({ __esModule: true, default: () => null })); + +const mockGetClaudeCodePluginsList = vi.mocked(getClaudeCodePluginsList); + +describe("ClaudeCodePluginsPanel loading state", () => { + beforeEach(() => { + vi.clearAllMocks(); + }); + + it("should resolve the loading state when accessToken is null instead of showing the skeleton forever", async () => { + render(); + expect(await screen.findByText("table-loaded")).toBeInTheDocument(); + expect(mockGetClaudeCodePluginsList).not.toHaveBeenCalled(); + }); + + it("should show the loading state until the skills fetch settles", async () => { + let resolveFetch: (value: { plugins: never[]; count: number }) => void = () => {}; + mockGetClaudeCodePluginsList.mockReturnValue( + new Promise((resolve) => { + resolveFetch = resolve; + }), + ); + render(); + expect(screen.getByText("table-loading")).toBeInTheDocument(); + + resolveFetch({ plugins: [], count: 0 }); + expect(await screen.findByText("table-loaded")).toBeInTheDocument(); + expect(mockGetClaudeCodePluginsList).toHaveBeenCalledWith("sk-test", false); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/skills/_components/ClaudeCodePluginsPanel.tsx b/ui/litellm-dashboard/src/app/(dashboard)/skills/_components/ClaudeCodePluginsPanel.tsx index 178cd36857f..5a638f9ae79 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/skills/_components/ClaudeCodePluginsPanel.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/skills/_components/ClaudeCodePluginsPanel.tsx @@ -3,7 +3,7 @@ import { Button } from "@tremor/react"; import { Modal } from "antd"; import { getClaudeCodePluginsList, deleteClaudeCodePlugin } from "@/components/networking"; import AddPluginForm from "./add_plugin_form"; -import PluginTable from "./plugin_table"; +import PluginTable from "./PluginTable"; import SkillDetail from "@/components/claude_code_plugins/skill_detail"; import { isAdminRole } from "@/utils/roles"; import NotificationsManager from "@/components/molecules/notifications_manager"; @@ -17,7 +17,7 @@ interface ClaudeCodePluginsPanelProps { const ClaudeCodePluginsPanel: React.FC = ({ accessToken, userRole }) => { const [pluginsList, setPluginsList] = useState([]); const [isAddModalVisible, setIsAddModalVisible] = useState(false); - const [isLoading, setIsLoading] = useState(false); + const [isLoading, setIsLoading] = useState(true); const [isDeleting, setIsDeleting] = useState(false); const [pluginToDelete, setPluginToDelete] = useState<{ name: string; @@ -28,7 +28,10 @@ const ClaudeCodePluginsPanel: React.FC = ({ accessT const isAdmin = userRole ? isAdminRole(userRole) : false; const fetchPlugins = async () => { - if (!accessToken) return; + if (!accessToken) { + setIsLoading(false); + return; + } setIsLoading(true); try { @@ -95,7 +98,6 @@ const ClaudeCodePluginsPanel: React.FC = ({ accessT pluginsList={pluginsList} isLoading={isLoading} onDeleteClick={handleDeleteClick} - accessToken={accessToken} isAdmin={isAdmin} onPluginClick={(id) => { const skill = pluginsList.find((p) => p.id === id); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/skills/_components/PluginTable.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/skills/_components/PluginTable.test.tsx new file mode 100644 index 00000000000..66a4d7e2524 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/skills/_components/PluginTable.test.tsx @@ -0,0 +1,113 @@ +import { render, screen, within } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import { beforeEach, describe, expect, it, vi } from "vitest"; + +import { Plugin } from "@/components/claude_code_plugins/types"; + +import PluginTable from "./PluginTable"; + +const mockPlugins: Plugin[] = [ + { + id: "plugin-id-newer", + name: "newer-skill", + version: "1.2.0", + description: "A skill for testing", + source: { source: "github", repo: "org/newer-skill" }, + category: "development", + enabled: true, + created_at: "2025-01-15T10:30:00Z", + }, + { + id: "plugin-id-older", + name: "older-skill", + source: { source: "github", repo: "org/older-skill" }, + enabled: false, + created_at: "2024-01-10T09:15:00Z", + }, +]; + +const mockOnDeleteClick = vi.fn(); +const mockOnPluginClick = vi.fn(); + +const defaultProps = { + pluginsList: mockPlugins, + isLoading: false, + onDeleteClick: mockOnDeleteClick, + isAdmin: true, + onPluginClick: mockOnPluginClick, +}; + +describe("PluginTable", () => { + beforeEach(() => { + vi.clearAllMocks(); + }); + + it("should render every column header", () => { + render(); + for (const header of ["Skill Name", "Version", "Description", "Category", "Public", "Created At"]) { + expect(screen.getByText(header)).toBeInTheDocument(); + } + }); + + it("should display the empty state when data is empty", () => { + render(); + expect(screen.getByText("No skills found")).toBeInTheDocument(); + }); + + it("should sort by created date descending by default", () => { + render(); + const rows = screen.getAllByRole("row").slice(1); + expect(within(rows[0]).getByText("newer-skill")).toBeInTheDocument(); + expect(within(rows[1]).getByText("older-skill")).toBeInTheDocument(); + }); + + it("should call onPluginClick with the plugin ID when the skill name is clicked", async () => { + const user = userEvent.setup(); + render(); + await user.click(screen.getByRole("button", { name: "newer-skill" })); + expect(mockOnPluginClick).toHaveBeenCalledWith("plugin-id-newer"); + }); + + it("should not navigate when clicking elsewhere in the row", async () => { + const user = userEvent.setup(); + render(); + await user.click(screen.getByText("A skill for testing")); + expect(mockOnPluginClick).not.toHaveBeenCalled(); + }); + + it("should badge the category and fall back to Uncategorized", () => { + render(); + expect(screen.getByText("development")).toBeInTheDocument(); + expect(screen.getByText("Uncategorized")).toBeInTheDocument(); + }); + + it("should show whether the skill is public", () => { + render(); + expect(screen.getByText("Yes")).toBeInTheDocument(); + expect(screen.getByText("No")).toBeInTheDocument(); + }); + + it("should delete a skill through the actions menu when admin", async () => { + const user = userEvent.setup(); + render(); + await user.click(screen.getByTestId("plugin-actions-newer-skill")); + await user.click(await screen.findByTestId("plugin-action-delete")); + expect(mockOnDeleteClick).toHaveBeenCalledWith("newer-skill", "newer-skill"); + }); + + it("should copy the skill ID through the actions menu", async () => { + const user = userEvent.setup(); + render(); + await user.click(screen.getByTestId("plugin-actions-newer-skill")); + await user.click(await screen.findByTestId("plugin-action-copy")); + expect(await window.navigator.clipboard.readText()).toBe("plugin-id-newer"); + }); + + it("should hide the delete action for non-admins but keep copy available", async () => { + const user = userEvent.setup(); + render(); + await user.click(screen.getByTestId("plugin-actions-newer-skill")); + expect(await screen.findByTestId("plugin-action-copy")).toBeInTheDocument(); + expect(screen.queryByTestId("plugin-action-delete")).not.toBeInTheDocument(); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/skills/_components/PluginTable.tsx b/ui/litellm-dashboard/src/app/(dashboard)/skills/_components/PluginTable.tsx new file mode 100644 index 00000000000..c581b0dfdeb --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/skills/_components/PluginTable.tsx @@ -0,0 +1,58 @@ +"use client"; + +import { SortingState } from "@tanstack/react-table"; +import { Inbox } from "lucide-react"; +import React, { useMemo, useState } from "react"; + +import { DataTable } from "@/components/shared/DataTable"; +import { Plugin } from "@/components/claude_code_plugins/types"; + +import { getPluginTableColumns } from "./PluginTableColumns"; + +interface PluginTableProps { + pluginsList: Plugin[]; + isLoading: boolean; + onDeleteClick: (pluginName: string, displayName: string) => void; + isAdmin: boolean; + onPluginClick: (pluginId: string) => void; +} + +const DEFAULT_SORTING: SortingState = [{ id: "created_at", desc: true }]; + +function EmptyState() { + return ( +
+
+ +
+
No skills found
+
Add one to get started.
+
+ ); +} + +const PluginTable: React.FC = ({ pluginsList, isLoading, onDeleteClick, isAdmin, onPluginClick }) => { + const [sorting, setSorting] = useState(DEFAULT_SORTING); + + const columns = useMemo( + () => getPluginTableColumns({ isAdmin, onPluginClick, onDeleteClick }), + [isAdmin, onPluginClick, onDeleteClick], + ); + + return ( + plugin.id || String(index)} + sortingMode="client" + sorting={sorting} + onSortingChange={setSorting} + isLoading={isLoading} + loadingMessage="Loading skills…" + noDataMessage={} + size="compact" + /> + ); +}; + +export default PluginTable; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/skills/_components/PluginTableColumns.tsx b/ui/litellm-dashboard/src/app/(dashboard)/skills/_components/PluginTableColumns.tsx new file mode 100644 index 00000000000..95c9924b375 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/skills/_components/PluginTableColumns.tsx @@ -0,0 +1,180 @@ +"use client"; + +import { ColumnDef } from "@tanstack/react-table"; +import { Copy, MoreHorizontal, Trash2 } from "lucide-react"; + +import { DataTableSortHeader } from "@/components/shared/DataTable"; +import { DateCell, IdentityCell, StatusBadge } from "@/components/shared/table_cells"; +import { getCategoryBadgeColor } from "@/components/claude_code_plugins/helpers"; +import { Plugin } from "@/components/claude_code_plugins/types"; +import { Badge } from "@/components/ui/badge"; +import { buttonVariants } from "@/components/ui/button"; +import { + DropdownMenu, + DropdownMenuContent, + DropdownMenuItem, + DropdownMenuSeparator, + DropdownMenuTrigger, +} from "@/components/ui/dropdown-menu"; +import { cn } from "@/lib/cva.config"; +import { copyToClipboard } from "@/utils/dataUtils"; + +const CATEGORY_BADGE_CLASS: Record, string> = { + blue: "border-blue-200 bg-blue-50 text-blue-600", + green: "border-green-200 bg-green-50 text-green-600", + purple: "border-purple-200 bg-purple-50 text-purple-600", + red: "border-red-200 bg-red-50 text-red-600", + orange: "border-orange-200 bg-orange-50 text-orange-600", + yellow: "border-yellow-200 bg-yellow-50 text-yellow-600", + gray: "border-gray-200 bg-gray-50 text-gray-600", +}; + +function PluginCategoryBadge({ category }: { category?: string }) { + return ( + + {category || "Uncategorized"} + + ); +} + +interface PluginRowActionsProps { + plugin: Plugin; + isAdmin: boolean; + onDeleteClick: (pluginName: string, displayName: string) => void; +} + +function PluginRowActions({ plugin, isAdmin, onDeleteClick }: PluginRowActionsProps) { + return ( + + + + + + void copyToClipboard(plugin.id, "Skill ID copied")} + > + + Copy skill ID + + {isAdmin && ( + <> + + onDeleteClick(plugin.name, plugin.name)} + > + + Delete + + + )} + + + ); +} + +interface PluginTableColumnsDeps { + isAdmin: boolean; + onPluginClick: (pluginId: string) => void; + onDeleteClick: (pluginName: string, displayName: string) => void; +} + +export const getPluginTableColumns = ({ + isAdmin, + onPluginClick, + onDeleteClick, +}: PluginTableColumnsDeps): ColumnDef[] => [ + { + id: "name", + accessorKey: "name", + meta: { title: "Skill Name" }, + header: ({ column }) => , + size: 220, + enableSorting: true, + cell: ({ row }) => ( + onPluginClick(row.original.id)} + /> + ), + }, + { + id: "version", + accessorKey: "version", + meta: { title: "Version" }, + header: "Version", + size: 100, + enableSorting: false, + cell: ({ row }) => {row.original.version || "N/A"}, + }, + { + id: "description", + accessorKey: "description", + meta: { title: "Description" }, + header: "Description", + size: 300, + enableSorting: false, + cell: ({ row }) => { + const description = row.original.description; + return ( + + {description || "No description"} + + ); + }, + }, + { + id: "category", + accessorKey: "category", + meta: { title: "Category", skeleton: "badge" }, + header: "Category", + size: 150, + enableSorting: false, + cell: ({ row }) => , + }, + { + id: "enabled", + accessorKey: "enabled", + meta: { title: "Public", skeleton: "badge" }, + header: "Public", + size: 100, + enableSorting: false, + cell: ({ row }) => ( + + ), + }, + { + id: "created_at", + accessorKey: "created_at", + sortingFn: "datetime", + meta: { title: "Created At" }, + header: ({ column }) => , + size: 160, + enableSorting: true, + cell: ({ row }) => , + }, + { + id: "actions", + meta: { className: "text-right", headerClassName: "text-right" }, + header: () => Actions, + size: 64, + enableSorting: false, + enableHiding: false, + cell: ({ row }) => ( +
+ +
+ ), + }, +]; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/skills/_components/plugin_table.tsx b/ui/litellm-dashboard/src/app/(dashboard)/skills/_components/plugin_table.tsx deleted file mode 100644 index eb1c495374a..00000000000 --- a/ui/litellm-dashboard/src/app/(dashboard)/skills/_components/plugin_table.tsx +++ /dev/null @@ -1,245 +0,0 @@ -import { CopyOutlined } from "@ant-design/icons"; -import { ChevronDownIcon, ChevronUpIcon, SwitchVerticalIcon, TrashIcon } from "@heroicons/react/outline"; -import { - ColumnDef, - flexRender, - getCoreRowModel, - getSortedRowModel, - SortingState, - useReactTable, -} from "@tanstack/react-table"; -import { Badge, Button, Table, TableBody, TableCell, TableHead, TableHeaderCell, TableRow } from "@tremor/react"; -import { Tooltip } from "antd"; -import React, { useState } from "react"; -import { DateCell, IdCell, StatusBadge } from "@/components/shared/table_cells"; -import NotificationsManager from "@/components/molecules/notifications_manager"; -import { getCategoryBadgeColor } from "@/components/claude_code_plugins/helpers"; -import { Plugin } from "@/components/claude_code_plugins/types"; - -interface PluginTableProps { - pluginsList: Plugin[]; - isLoading: boolean; - onDeleteClick: (pluginName: string, displayName: string) => void; - accessToken: string | null; - isAdmin: boolean; - onPluginClick: (pluginId: string) => void; -} - -const PluginTable: React.FC = ({ - pluginsList, - isLoading, - onDeleteClick, - accessToken, - isAdmin, - onPluginClick, -}) => { - const [sorting, setSorting] = useState([{ id: "created_at", desc: true }]); - - const copyToClipboard = (text: string) => { - navigator.clipboard.writeText(text); - NotificationsManager.success("Copied to clipboard!"); - }; - - const columns: ColumnDef[] = [ - { - header: "Skill Name", - accessorKey: "name", - cell: ({ row }) => { - const plugin = row.original; - return ( -
- onPluginClick(plugin.id)} /> - - { - e.stopPropagation(); - copyToClipboard(plugin.id); - }} - className="cursor-pointer text-gray-500 hover:text-blue-500 text-xs" - /> - -
- ); - }, - }, - { - header: "Version", - accessorKey: "version", - cell: ({ row }) => { - const version = row.original.version || "N/A"; - return {version}; - }, - }, - { - header: "Description", - accessorKey: "description", - cell: ({ row }) => { - const description = row.original.description || "No description"; - return ( - - {description} - - ); - }, - }, - { - header: "Category", - accessorKey: "category", - cell: ({ row }) => { - const category = row.original.category; - if (!category) { - return ( - - Uncategorized - - ); - } - const badgeColor = getCategoryBadgeColor(category); - return ( - - {category} - - ); - }, - }, - { - header: "Public", - accessorKey: "enabled", - cell: ({ row }) => { - const plugin = row.original; - return ; - }, - }, - { - header: "Created At", - accessorKey: "created_at", - cell: ({ row }) => , - }, - ...(isAdmin - ? [ - { - header: "Actions", - id: "actions", - enableSorting: false, - cell: ({ row }: any) => { - const plugin = row.original; - - return ( -
- -
- ); - }, - }, - ] - : []), - ]; - - const table = useReactTable({ - data: pluginsList, - columns, - state: { - sorting, - }, - onSortingChange: setSorting, - getCoreRowModel: getCoreRowModel(), - getSortedRowModel: getSortedRowModel(), - enableSorting: true, - }); - - return ( -
-
-
- - {table.getHeaderGroups().map((headerGroup) => ( - - {headerGroup.headers.map((header) => ( - -
-
- {header.isPlaceholder ? null : flexRender(header.column.columnDef.header, header.getContext())} -
- {header.column.getCanSort() && ( -
- {header.column.getIsSorted() ? ( - { - asc: , - desc: , - }[header.column.getIsSorted() as string] - ) : ( - - )} -
- )} -
-
- ))} -
- ))} -
- - {isLoading ? ( - - -
-

Loading...

-
-
-
- ) : pluginsList && pluginsList.length > 0 ? ( - table.getRowModel().rows.map((row) => ( - onPluginClick(row.original.id)} - > - {row.getVisibleCells().map((cell) => ( - - {flexRender(cell.column.columnDef.cell, cell.getContext())} - - ))} - - )) - ) : ( - - -
-

No skills found. Add one to get started.

-
-
-
- )} -
-
-
- - ); -}; - -export default PluginTable; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/vector-stores/_components/DocumentsTable.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/vector-stores/_components/DocumentsTable.test.tsx index c2425c8d651..bbeef5c2216 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/vector-stores/_components/DocumentsTable.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/vector-stores/_components/DocumentsTable.test.tsx @@ -1,18 +1,10 @@ -import { render, screen, fireEvent, act } from "@testing-library/react"; -import { describe, it, expect, vi } from "vitest"; -import DocumentsTable from "./DocumentsTable"; +import { render, screen } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import { beforeEach, describe, expect, it, vi } from "vitest"; + import { DocumentUpload } from "@/components/vector_store_management/types"; -// Mock antd message -vi.mock("antd", async () => { - const actual = await vi.importActual("antd"); - return { - ...actual, - message: { - success: vi.fn(), - }, - }; -}); +import DocumentsTable from "./DocumentsTable"; describe("DocumentsTable", () => { const mockDocuments: DocumentUpload[] = [ @@ -39,9 +31,12 @@ describe("DocumentsTable", () => { }, ]; - it("should render the table successfully", () => { - const onRemove = vi.fn(); - render(); + beforeEach(() => { + vi.clearAllMocks(); + }); + + it("should render every document row", () => { + render(); expect(screen.getByText("test1.pdf")).toBeInTheDocument(); expect(screen.getByText("test2.txt")).toBeInTheDocument(); @@ -49,8 +44,7 @@ describe("DocumentsTable", () => { }); it("should display correct status badges", () => { - const onRemove = vi.fn(); - render(); + render(); expect(screen.getByText("Ready")).toBeInTheDocument(); expect(screen.getByText("Uploading")).toBeInTheDocument(); @@ -58,45 +52,46 @@ describe("DocumentsTable", () => { }); it("should display file sizes", () => { - const onRemove = vi.fn(); - render(); + render(); expect(screen.getByText(/1000.00 KB/)).toBeInTheDocument(); expect(screen.getByText(/1.95 MB/)).toBeInTheDocument(); expect(screen.getByText(/500.00 KB/)).toBeInTheDocument(); }); - it("should call onRemove when delete button is clicked", () => { + it("should call onRemove through the actions menu", async () => { + const user = userEvent.setup(); const onRemove = vi.fn(); render(); - const deleteButtons = screen.getAllByLabelText(/delete/i); - - act(() => { - fireEvent.click(deleteButtons[0]); - }); + await user.click(screen.getByTestId("document-actions-1")); + await user.click(await screen.findByTestId("document-action-remove")); expect(onRemove).toHaveBeenCalledWith("1"); }); - it("should show empty state when no documents", () => { - const onRemove = vi.fn(); - render(); + it("should copy the document ID through the actions menu", async () => { + const user = userEvent.setup(); + render(); - expect(screen.getByText(/No documents uploaded yet/)).toBeInTheDocument(); + await user.click(screen.getByTestId("document-actions-2")); + await user.click(await screen.findByTestId("document-action-copy")); + + expect(await window.navigator.clipboard.readText()).toBe("2"); }); - it("should have action buttons for each document", () => { - const onRemove = vi.fn(); - render(); + it("should show the empty state when no documents", () => { + render(); - // Each document should have 3 action buttons (view, copy, delete) - const viewButtons = screen.getAllByLabelText(/eye/i); - const copyButtons = screen.getAllByLabelText(/copy/i); - const deleteButtons = screen.getAllByLabelText(/delete/i); + expect(screen.getByText("No documents uploaded yet")).toBeInTheDocument(); + expect(screen.getByText("Upload documents above to get started.")).toBeInTheDocument(); + }); - expect(viewButtons).toHaveLength(3); - expect(copyButtons).toHaveLength(3); - expect(deleteButtons).toHaveLength(3); + it("should render one actions menu per document", () => { + render(); + + for (const doc of mockDocuments) { + expect(screen.getByTestId(`document-actions-${doc.uid}`)).toBeInTheDocument(); + } }); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/vector-stores/_components/DocumentsTable.tsx b/ui/litellm-dashboard/src/app/(dashboard)/vector-stores/_components/DocumentsTable.tsx index d4c45a6b751..416fcb37a10 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/vector-stores/_components/DocumentsTable.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/vector-stores/_components/DocumentsTable.tsx @@ -1,95 +1,40 @@ -import React from "react"; -import { Table, Tooltip } from "antd"; -import MessageManager from "@/components/molecules/message_manager"; -import { EyeOutlined, CopyOutlined, DeleteOutlined } from "@ant-design/icons"; -import { StatusBadge, type StatusTone } from "@/components/shared/table_cells"; +"use client"; + +import { Inbox } from "lucide-react"; +import React, { useMemo } from "react"; + +import { DataTable } from "@/components/shared/DataTable"; import { DocumentUpload } from "@/components/vector_store_management/types"; +import { getDocumentsTableColumns } from "./DocumentsTableColumns"; + interface DocumentsTableProps { documents: DocumentUpload[]; onRemove: (uid: string) => void; } +function EmptyState() { + return ( +
+
+ +
+
No documents uploaded yet
+
Upload documents above to get started.
+
+ ); +} + const DocumentsTable: React.FC = ({ documents, onRemove }) => { - const handleCopyId = (uid: string) => { - navigator.clipboard.writeText(uid); - MessageManager.success("Document ID copied to clipboard"); - }; - - const getStatusBadge = (status: DocumentUpload["status"]) => { - const statusConfig: Record = { - uploading: { tone: "info", label: "Uploading" }, - done: { tone: "success", label: "Ready" }, - error: { tone: "error", label: "Error" }, - removed: { tone: "neutral", label: "Removed" }, - }; - - const config: { tone: StatusTone; label: string } = statusConfig[status] ?? { tone: "neutral", label: status }; - return ; - }; - - const formatFileSize = (bytes?: number) => { - if (!bytes) return "-"; - const kb = bytes / 1024; - if (kb < 1024) return `${kb.toFixed(2)} KB`; - return `${(kb / 1024).toFixed(2)} MB`; - }; - - const columns = [ - { - title: "Name", - dataIndex: "name", - key: "name", - render: (name: string, record: DocumentUpload) => ( -
- {name} - {record.size && ({formatFileSize(record.size)})} -
- ), - }, - { - title: "Status", - dataIndex: "status", - key: "status", - width: 150, - render: (status: DocumentUpload["status"]) => getStatusBadge(status), - }, - { - title: "Actions", - key: "actions", - width: 120, - render: (_: any, record: DocumentUpload) => ( -
- - {}} /> - - - handleCopyId(record.uid)} - /> - - - onRemove(record.uid)} - /> - -
- ), - }, - ]; + const columns = useMemo(() => getDocumentsTableColumns({ onRemove }), [onRemove]); return ( - document.uid || String(index)} + noDataMessage={} + size="compact" /> ); }; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/vector-stores/_components/DocumentsTableColumns.tsx b/ui/litellm-dashboard/src/app/(dashboard)/vector-stores/_components/DocumentsTableColumns.tsx new file mode 100644 index 00000000000..9be9f806bf9 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/vector-stores/_components/DocumentsTableColumns.tsx @@ -0,0 +1,110 @@ +"use client"; + +import { ColumnDef } from "@tanstack/react-table"; +import { Copy, MoreHorizontal, Trash2 } from "lucide-react"; + +import { StatusBadge, type StatusTone } from "@/components/shared/table_cells"; +import { buttonVariants } from "@/components/ui/button"; +import { + DropdownMenu, + DropdownMenuContent, + DropdownMenuItem, + DropdownMenuTrigger, +} from "@/components/ui/dropdown-menu"; +import { DocumentUpload } from "@/components/vector_store_management/types"; +import { cn } from "@/lib/cva.config"; +import { copyToClipboard } from "@/utils/dataUtils"; + +const STATUS_CONFIG: Record = { + uploading: { tone: "info", label: "Uploading" }, + done: { tone: "success", label: "Ready" }, + error: { tone: "error", label: "Error" }, + removed: { tone: "neutral", label: "Removed" }, +}; + +function formatFileSize(bytes?: number): string { + if (!bytes) return "-"; + const kb = bytes / 1024; + if (kb < 1024) return `${kb.toFixed(2)} KB`; + return `${(kb / 1024).toFixed(2)} MB`; +} + +function DocumentRowActions({ document, onRemove }: { document: DocumentUpload; onRemove: (uid: string) => void }) { + return ( + + + + + + void copyToClipboard(document.uid, "Document ID copied to clipboard")} + > + + Copy document ID + + onRemove(document.uid)} + > + + Remove + + + + ); +} + +interface DocumentsTableColumnsDeps { + onRemove: (uid: string) => void; +} + +export const getDocumentsTableColumns = ({ onRemove }: DocumentsTableColumnsDeps): ColumnDef[] => [ + { + id: "name", + accessorKey: "name", + meta: { title: "Name" }, + header: "Name", + enableSorting: false, + cell: ({ row }) => ( +
+ + {row.original.name} + + {row.original.size ? ( + ({formatFileSize(row.original.size)}) + ) : null} +
+ ), + }, + { + id: "status", + accessorKey: "status", + meta: { title: "Status", skeleton: "badge" }, + header: "Status", + size: 150, + enableSorting: false, + cell: ({ row }) => { + const config = STATUS_CONFIG[row.original.status] ?? { tone: "neutral", label: row.original.status }; + return ; + }, + }, + { + id: "actions", + meta: { className: "text-right", headerClassName: "text-right" }, + header: () => Actions, + size: 64, + enableSorting: false, + enableHiding: false, + cell: ({ row }) => ( +
+ +
+ ), + }, +]; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/vector-stores/_components/VectorStoreTable.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/vector-stores/_components/VectorStoreTable.test.tsx index 45ca40152fa..7446d0efe3c 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/vector-stores/_components/VectorStoreTable.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/vector-stores/_components/VectorStoreTable.test.tsx @@ -1,90 +1,46 @@ -import { render, screen } from "@testing-library/react"; +import { render, screen, within } from "@testing-library/react"; import userEvent from "@testing-library/user-event"; import { beforeEach, describe, expect, it, vi } from "vitest"; -import VectorStoreTable from "./VectorStoreTable"; + import { VectorStore } from "@/components/vector_store_management/types"; -// Mock dependencies -const mockGetProviderLogoAndName = vi.fn(); -const mockTableIconActionButton = vi.fn(); +import VectorStoreTable from "./VectorStoreTable"; vi.mock("@/components/provider_info_helpers", () => ({ - getProviderLogoAndName: (...args: any[]) => mockGetProviderLogoAndName(...args), -})); - -vi.mock("@/components/common_components/IconActionButton/TableIconActionButtons/TableIconActionButton", () => ({ - default: (props: any) => { - mockTableIconActionButton(props); - return ( - - ); + getProviderLogoAndName: (provider: string) => { + const providerMap: Record = { + openai: { displayName: "OpenAI", logo: "/openai-logo.png" }, + azure: { displayName: "Azure", logo: "/azure-logo.png" }, + }; + return providerMap[provider] || { displayName: provider, logo: "" }; }, })); -// Mock Tremor components to avoid complex styling issues -vi.mock("@tremor/react", () => ({ - Table: ({ children, ...props }: any) =>
{children}
, - TableHead: ({ children, ...props }: any) => {children}, - TableBody: ({ children, ...props }: any) => {children}, - TableRow: ({ children, ...props }: any) => {children}, - TableHeaderCell: ({ children, ...props }: any) => {children}, - TableCell: ({ children, ...props }: any) => {children}, -})); - -// Mock antd Tooltip -vi.mock("antd", () => ({ - Tooltip: ({ children, title }: any) => ( -
- {children} -
- ), -})); - -// Mock Heroicons -vi.mock("@heroicons/react/outline", () => ({ - ChevronDownIcon: (props: any) =>
, - ChevronUpIcon: (props: any) =>
, - SwitchVerticalIcon: (props: any) =>
, -})); - -// Test data const mockVectorStores: VectorStore[] = [ { - vector_store_id: "short-id", + vector_store_id: "vs-newer", custom_llm_provider: "openai", vector_store_name: "My OpenAI Store", vector_store_description: "A store for OpenAI vectors", + vector_store_metadata: { + ingested_files: [ + { filename: "a.pdf", ingested_at: "2024-01-15T10:00:00Z" }, + { filename: "b.pdf", ingested_at: "2024-01-15T10:00:00Z" }, + ], + }, created_at: "2024-01-15T10:30:00Z", updated_at: "2024-01-15T11:00:00Z", - created_by: "user-1", - updated_by: "user-1", }, { - vector_store_id: "very-long-vector-store-id-that-should-be-truncated", + vector_store_id: "vs-older", custom_llm_provider: "azure", - vector_store_name: undefined, // Test missing name - vector_store_description: "A store for Azure vectors with a very long description that should show a tooltip", + vector_store_name: undefined, + vector_store_description: undefined, created_at: "2024-01-10T09:15:00Z", updated_at: "2024-01-12T14:20:00Z", }, - { - vector_store_id: "store-3", - custom_llm_provider: "pg_vector", - vector_store_name: "PostgreSQL Store", - vector_store_description: undefined, // Test missing description - created_at: "2024-01-05T08:00:00Z", - updated_at: "2024-01-08T16:45:00Z", - }, ]; -// Mock functions const mockOnView = vi.fn(); const mockOnEdit = vi.fn(); const mockOnDelete = vi.fn(); @@ -96,319 +52,71 @@ const defaultProps = { onDelete: mockOnDelete, }; -// Helper function to render component -const renderComponent = (props = {}) => { - return render(); -}; - describe("VectorStoreTable", () => { beforeEach(() => { vi.clearAllMocks(); - - // Setup default mock returns for getProviderLogoAndName - mockGetProviderLogoAndName.mockImplementation((provider: string) => { - const providerMap: Record = { - openai: { displayName: "OpenAI", logo: "/openai-logo.png" }, - azure: { displayName: "Azure", logo: "/azure-logo.png" }, - pg_vector: { displayName: "PostgreSQL Vector", logo: "/pg-logo.png" }, - }; - return providerMap[provider] || { displayName: provider, logo: "" }; - }); }); - describe("Rendering", () => { - it("should render the table with data", () => { - renderComponent(); - expect(screen.getByRole("table")).toBeInTheDocument(); - }); - - it("should render table headers", () => { - renderComponent(); - expect(screen.getByText("Vector Store ID")).toBeInTheDocument(); - expect(screen.getByText("Name")).toBeInTheDocument(); - expect(screen.getByText("Description")).toBeInTheDocument(); - expect(screen.getByText("Provider")).toBeInTheDocument(); - expect(screen.getByText("Created At")).toBeInTheDocument(); - expect(screen.getByText("Updated At")).toBeInTheDocument(); - // Check that we have the expected number of header cells (7 data + 1 actions) - const headers = screen.getAllByRole("columnheader"); - expect(headers).toHaveLength(8); - }); - - it("should render all vector store rows", () => { - renderComponent(); - expect(screen.getAllByRole("row")).toHaveLength(mockVectorStores.length + 1); // +1 for header row - }); - - it("should render empty state when no data", () => { - renderComponent({ data: [] }); - expect(screen.getByText("No vector stores found")).toBeInTheDocument(); - }); + it("should render every column header", () => { + render(); + for (const header of ["Vector Store ID", "Name", "Description", "Files", "Provider", "Created At", "Updated At"]) { + expect(screen.getByText(header)).toBeInTheDocument(); + } }); - describe("Vector Store ID Column", () => { - it("should render short vector store IDs fully", () => { - renderComponent(); - expect(screen.getByText("short-id")).toBeInTheDocument(); - }); - - it("should truncate long vector store IDs", () => { - renderComponent(); - const idButton = screen.getByText("very-long-vector-store-id-that-should-be-truncated"); - expect(idButton).toHaveClass("truncate", "max-w-[15ch]"); - }); - - it("should make vector store ID clickable", async () => { - const user = userEvent.setup(); - renderComponent(); - const idButton = screen.getByText("short-id"); - await user.click(idButton); - expect(mockOnView).toHaveBeenCalledWith("short-id"); - }); - - it("should have correct styling for vector store ID button", () => { - renderComponent(); - const idButton = screen.getByText("short-id").closest("button"); - expect(idButton).toHaveClass("font-mono", "text-blue-500", "bg-blue-50", "hover:bg-blue-100"); - }); + it("should display the empty state when data is empty", () => { + render(); + expect(screen.getByText("No vector stores")).toBeInTheDocument(); }); - describe("Name Column", () => { - it("should render vector store name", () => { - renderComponent(); - expect(screen.getByText("My OpenAI Store")).toBeInTheDocument(); - }); - - it("should render fallback for missing name", () => { - renderComponent(); - const fallbackElements = screen.getAllByText("-"); - expect(fallbackElements.length).toBe(5); // One for missing name, one for missing description, three for missing files (one per store) - }); - - it("should wrap name in tooltip", () => { - renderComponent(); - const tooltips = screen.getAllByTestId("tooltip"); - const nameTooltip = tooltips.find((t) => t.getAttribute("data-title") === "My OpenAI Store"); - expect(nameTooltip).toBeInTheDocument(); - }); + it("should sort by created date descending by default", () => { + render(); + const rows = screen.getAllByRole("row").slice(1); + expect(within(rows[0]).getByText("vs-newer")).toBeInTheDocument(); + expect(within(rows[1]).getByText("vs-older")).toBeInTheDocument(); }); - describe("Description Column", () => { - it("should render vector store description", () => { - renderComponent(); - expect(screen.getByText("A store for OpenAI vectors")).toBeInTheDocument(); - }); - - it("should render fallback for missing description", () => { - renderComponent(); - const fallbackElements = screen.getAllByText("-"); - expect(fallbackElements.length).toBe(5); // One for missing name, one for missing description, three for missing files (one per store) - }); - - it("should wrap description in tooltip", () => { - renderComponent(); - const tooltips = screen.getAllByTestId("tooltip"); - const descTooltip = tooltips.find( - (t) => - t.getAttribute("data-title") === - "A store for Azure vectors with a very long description that should show a tooltip", - ); - expect(descTooltip).toBeInTheDocument(); - }); + it("should call onView when the vector store ID is clicked", async () => { + const user = userEvent.setup(); + render(); + await user.click(screen.getByRole("button", { name: "vs-newer" })); + expect(mockOnView).toHaveBeenCalledWith("vs-newer"); }); - describe("Provider Column", () => { - it("should render provider display name", () => { - renderComponent(); - expect(screen.getByText("OpenAI")).toBeInTheDocument(); - expect(screen.getByText("Azure")).toBeInTheDocument(); - expect(screen.getByText("PostgreSQL Vector")).toBeInTheDocument(); - }); - - it("should render provider logo when available", () => { - renderComponent(); - const logos = screen.getAllByRole("img"); - expect(logos).toHaveLength(3); // All providers have logos in our mock - expect(logos[0]).toHaveAttribute("src", "/openai-logo.png"); - expect(logos[0]).toHaveAttribute("alt", "OpenAI"); - }); - - it("should call getProviderLogoAndName for each provider", () => { - renderComponent(); - expect(mockGetProviderLogoAndName).toHaveBeenCalledWith("openai"); - expect(mockGetProviderLogoAndName).toHaveBeenCalledWith("azure"); - expect(mockGetProviderLogoAndName).toHaveBeenCalledWith("pg_vector"); - }); + it("should render provider display names", () => { + render(); + expect(screen.getByText("OpenAI")).toBeInTheDocument(); + expect(screen.getByText("Azure")).toBeInTheDocument(); }); - describe("Date Columns", () => { - it("should render created at dates", () => { - renderComponent(); - const dateElements = screen.getAllByText(/Jan \d+, 2024/); - expect(dateElements.length).toBe(6); // 3 created_at + 3 updated_at dates - }); - - it("should render updated at dates", () => { - renderComponent(); - const dateElements = screen.getAllByText(/Jan \d+, 2024/); - expect(dateElements.length).toBe(6); // 3 created_at + 3 updated_at dates - }); + it("should summarize ingested files and fall back to a dash without files", () => { + render(); + expect(screen.getByText("2 files")).toBeInTheDocument(); + const olderRow = screen.getAllByRole("row").slice(1)[1]; + expect(within(olderRow).getAllByText("-").length).toBeGreaterThan(0); }); - describe("Actions Column", () => { - it("should render edit and delete action buttons for each row", () => { - renderComponent(); - expect(screen.getAllByTestId("action-button-edit")).toHaveLength(mockVectorStores.length); - expect(screen.getAllByTestId("action-button-delete")).toHaveLength(mockVectorStores.length); - }); - - it("should call onEdit when edit button is clicked", async () => { - const user = userEvent.setup(); - renderComponent(); - const editButtons = screen.getAllByTestId("action-button-edit"); - await user.click(editButtons[0]); - expect(mockOnEdit).toHaveBeenCalledWith("short-id"); - }); - - it("should call onDelete when delete button is clicked", async () => { - const user = userEvent.setup(); - renderComponent(); - const deleteButtons = screen.getAllByTestId("action-button-delete"); - await user.click(deleteButtons[0]); - expect(mockOnDelete).toHaveBeenCalledWith("short-id"); - }); - - it("should pass correct props to TableIconActionButton", () => { - renderComponent(); - expect(mockTableIconActionButton).toHaveBeenCalledWith( - expect.objectContaining({ - variant: "Edit", - tooltipText: "Edit vector store", - onClick: expect.any(Function), - }), - ); - expect(mockTableIconActionButton).toHaveBeenCalledWith( - expect.objectContaining({ - variant: "Delete", - tooltipText: "Delete vector store", - onClick: expect.any(Function), - }), - ); - }); + it("should edit a vector store through the actions menu", async () => { + const user = userEvent.setup(); + render(); + await user.click(screen.getByTestId("vector-store-actions-vs-newer")); + await user.click(await screen.findByTestId("vector-store-action-edit")); + expect(mockOnEdit).toHaveBeenCalledWith("vs-newer"); }); - describe("Sorting", () => { - it("should initialize with created_at descending sort", () => { - renderComponent(); - // The table should initialize with sorting state - expect(screen.getByTestId("chevron-down")).toBeInTheDocument(); - }); - - it("should render sort icons for sortable columns", () => { - renderComponent(); - // Should have sort icons for Created At and Updated At columns - const sortIcons = screen.getAllByTestId(/^chevron-(up|down)$|^switch-vertical$/); - expect(sortIcons.length).toBeGreaterThan(0); - }); - - it("should make header cells clickable for sorting", () => { - renderComponent(); - const headerCells = screen.getAllByRole("columnheader"); - const sortableHeaders = headerCells.filter((cell) => cell.textContent !== ""); - expect(sortableHeaders.length).toBeGreaterThan(0); - }); - - it("should show ascending icon when sorted ascending", () => { - renderComponent(); - // Initially shows descending, but we can test the logic by checking the icons are present - expect(screen.getByTestId("chevron-down")).toBeInTheDocument(); - }); + it("should delete a vector store through the actions menu", async () => { + const user = userEvent.setup(); + render(); + await user.click(screen.getByTestId("vector-store-actions-vs-newer")); + await user.click(await screen.findByTestId("vector-store-action-delete")); + expect(mockOnDelete).toHaveBeenCalledWith("vs-newer"); }); - describe("Styling and Layout", () => { - it("should apply correct CSS classes to table container", () => { - renderComponent(); - const tableContainer = screen.getByRole("table").parentElement?.parentElement; - expect(tableContainer).toHaveClass("rounded-lg", "custom-border", "relative"); - }); - - it("should apply overflow styling to table wrapper", () => { - renderComponent(); - const tableWrapper = screen.getByRole("table").parentElement; - expect(tableWrapper).toHaveClass("overflow-x-auto"); - }); - - it("should apply sticky styling to actions column", () => { - renderComponent(); - const headerCells = screen.getAllByRole("columnheader"); - const actionsHeader = headerCells[headerCells.length - 1]; - expect(actionsHeader).toHaveClass("sticky", "right-0", "bg-white"); - }); - - it("should apply sticky styling to action cells", () => { - renderComponent(); - const rows = screen.getAllByRole("row").slice(1); // Skip header row - rows.forEach((row) => { - const cells = row.querySelectorAll("td"); - const lastCell = cells[cells.length - 1]; - expect(lastCell).toHaveClass("sticky", "right-0", "bg-white"); - }); - }); - }); - - describe("Table Row Styling", () => { - it("should apply correct height to table rows", () => { - renderComponent(); - const rows = screen.getAllByRole("row").slice(1); // Skip header row - rows.forEach((row) => { - expect(row).toHaveClass("h-8"); - }); - }); - - it("should apply correct cell padding and styling", () => { - renderComponent(); - const cells = screen.getAllByRole("cell"); - cells.forEach((cell) => { - expect(cell).toHaveClass("py-0.5", "max-h-8", "overflow-hidden", "text-ellipsis", "whitespace-nowrap"); - }); - }); - }); - - describe("Empty State", () => { - it("should render single row with centered message when no data", () => { - renderComponent({ data: [] }); - const rows = screen.getAllByRole("row"); - expect(rows).toHaveLength(2); // Header + empty state row - expect(screen.getByText("No vector stores found")).toBeInTheDocument(); - }); - - it("should span all columns in empty state", () => { - renderComponent({ data: [] }); - const emptyCell = screen.getByText("No vector stores found").closest("td"); - expect(emptyCell).toHaveAttribute("colSpan", "8"); // 7 data columns + 1 actions column - }); - }); - - describe("Data Edge Cases", () => { - it("should handle vector stores with minimal data", () => { - const minimalData: VectorStore[] = [ - { - vector_store_id: "minimal", - custom_llm_provider: "test", - created_at: "2024-01-01T00:00:00Z", - updated_at: "2024-01-01T00:00:00Z", - }, - ]; - - renderComponent({ data: minimalData }); - expect(screen.getByText("minimal")).toBeInTheDocument(); - expect(screen.getAllByText("-")).toHaveLength(3); // Name, description, and files fallbacks - }); - - it("should handle single vector store", () => { - const singleData = [mockVectorStores[0]]; - renderComponent({ data: singleData }); - expect(screen.getAllByRole("row")).toHaveLength(2); // Header + 1 data row - }); + it("should copy the vector store ID through the actions menu", async () => { + const user = userEvent.setup(); + render(); + await user.click(screen.getByTestId("vector-store-actions-vs-newer")); + await user.click(await screen.findByTestId("vector-store-action-copy")); + expect(await window.navigator.clipboard.readText()).toBe("vs-newer"); }); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/vector-stores/_components/VectorStoreTable.tsx b/ui/litellm-dashboard/src/app/(dashboard)/vector-stores/_components/VectorStoreTable.tsx index 074af873601..2f8508dc7c6 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/vector-stores/_components/VectorStoreTable.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/vector-stores/_components/VectorStoreTable.tsx @@ -1,215 +1,57 @@ -import { ChevronDownIcon, ChevronUpIcon, SwitchVerticalIcon } from "@heroicons/react/outline"; -import { - ColumnDef, - flexRender, - getCoreRowModel, - getSortedRowModel, - SortingState, - useReactTable, -} from "@tanstack/react-table"; -import { Table, TableBody, TableCell, TableHead, TableHeaderCell, TableRow } from "@tremor/react"; -import { Tooltip } from "antd"; -import React from "react"; -import { DateCell, IdCell } from "@/components/shared/table_cells"; -import TableIconActionButton from "@/components/common_components/IconActionButton/TableIconActionButtons/TableIconActionButton"; -import { getProviderLogoAndName } from "@/components/provider_info_helpers"; +"use client"; + +import { SortingState } from "@tanstack/react-table"; +import { Inbox } from "lucide-react"; +import React, { useMemo, useState } from "react"; + +import { DataTable } from "@/components/shared/DataTable"; import { VectorStore } from "@/components/vector_store_management/types"; +import { getVectorStoreTableColumns } from "./VectorStoreTableColumns"; + interface VectorStoreTableProps { data: VectorStore[]; onView: (vectorStoreId: string) => void; onEdit: (vectorStoreId: string) => void; onDelete: (vectorStoreId: string) => void; + isLoading?: boolean; } -const VectorStoreTable: React.FC = ({ data, onView, onEdit, onDelete }) => { - const [sorting, setSorting] = React.useState([{ id: "created_at", desc: true }]); - - const columns: ColumnDef[] = [ - { - header: "Vector Store ID", - accessorKey: "vector_store_id", - cell: ({ row }) => , - }, - { - header: "Name", - accessorKey: "vector_store_name", - cell: ({ row }) => { - const vectorStore = row.original; - return ( - - {vectorStore.vector_store_name || "-"} - - ); - }, - }, - { - header: "Description", - accessorKey: "vector_store_description", - cell: ({ row }) => { - const vectorStore = row.original; - return ( - - {vectorStore.vector_store_description || "-"} - - ); - }, - }, - { - header: "Files", - accessorKey: "vector_store_metadata", - cell: ({ row }) => { - const vectorStore = row.original; - const ingestedFiles = vectorStore.vector_store_metadata?.ingested_files || []; - - if (ingestedFiles.length === 0) { - return -; - } - - const filenames = ingestedFiles.map((file) => file.filename || file.file_url || "Unknown").join(", "); - - const displayText = - ingestedFiles.length === 1 - ? ingestedFiles[0].filename || ingestedFiles[0].file_url || "1 file" - : `${ingestedFiles.length} files`; - - return ( - - {displayText} - - ); - }, - }, - { - header: "Provider", - accessorKey: "custom_llm_provider", - cell: ({ row }) => { - const vectorStore = row.original; - const { displayName, logo } = getProviderLogoAndName(vectorStore.custom_llm_provider); - return ( -
- {logo && {displayName}} - {displayName} -
- ); - }, - }, - { - header: "Created At", - accessorKey: "created_at", - sortingFn: "datetime", - cell: ({ row }) => , - }, - { - header: "Updated At", - accessorKey: "updated_at", - sortingFn: "datetime", - cell: ({ row }) => , - }, - { - id: "actions", - header: "", - cell: ({ row }) => { - const vectorStore = row.original; - return ( -
- onEdit(vectorStore.vector_store_id)} - /> - onDelete(vectorStore.vector_store_id)} - /> -
- ); - }, - }, - ]; - - const table = useReactTable({ - data, - columns, - state: { - sorting, - }, - onSortingChange: setSorting, - getCoreRowModel: getCoreRowModel(), - getSortedRowModel: getSortedRowModel(), - enableSorting: true, - }); +const DEFAULT_SORTING: SortingState = [{ id: "created_at", desc: true }]; +function EmptyState() { return ( -
-
- - - {table.getHeaderGroups().map((headerGroup) => ( - - {headerGroup.headers.map((header) => ( - -
-
- {header.isPlaceholder ? null : flexRender(header.column.columnDef.header, header.getContext())} -
- {header.id !== "actions" && ( -
- {header.column.getIsSorted() ? ( - { - asc: , - desc: , - }[header.column.getIsSorted() as string] - ) : ( - - )} -
- )} -
-
- ))} -
- ))} -
- - {table.getRowModel().rows.length > 0 ? ( - table.getRowModel().rows.map((row) => ( - - {row.getVisibleCells().map((cell) => ( - - {flexRender(cell.column.columnDef.cell, cell.getContext())} - - ))} - - )) - ) : ( - - -
-

No vector stores found

-
-
-
- )} -
-
+
+
+ +
+
No vector stores
+
+ Connect a vector store to enable retrieval-augmented generation.
); +} + +const VectorStoreTable: React.FC = ({ data, onView, onEdit, onDelete, isLoading = false }) => { + const [sorting, setSorting] = useState(DEFAULT_SORTING); + + const columns = useMemo(() => getVectorStoreTableColumns({ onView, onEdit, onDelete }), [onView, onEdit, onDelete]); + + return ( + vectorStore.vector_store_id || String(index)} + sortingMode="client" + sorting={sorting} + onSortingChange={setSorting} + isLoading={isLoading} + loadingMessage="Loading vector stores…" + noDataMessage={} + size="compact" + /> + ); }; export default VectorStoreTable; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/vector-stores/_components/VectorStoreTableColumns.tsx b/ui/litellm-dashboard/src/app/(dashboard)/vector-stores/_components/VectorStoreTableColumns.tsx new file mode 100644 index 00000000000..cf162578177 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/vector-stores/_components/VectorStoreTableColumns.tsx @@ -0,0 +1,211 @@ +"use client"; + +import { ColumnDef } from "@tanstack/react-table"; +import { Copy, MoreHorizontal, Pencil, Trash2 } from "lucide-react"; + +import { DataTableSortHeader } from "@/components/shared/DataTable"; +import { CellTooltip, DateCell, IdentityCell } from "@/components/shared/table_cells"; +import { getProviderLogoAndName } from "@/components/provider_info_helpers"; +import { buttonVariants } from "@/components/ui/button"; +import { + DropdownMenu, + DropdownMenuContent, + DropdownMenuItem, + DropdownMenuSeparator, + DropdownMenuTrigger, +} from "@/components/ui/dropdown-menu"; +import { VectorStore } from "@/components/vector_store_management/types"; +import { cn } from "@/lib/cva.config"; +import { copyToClipboard } from "@/utils/dataUtils"; + +function VectorStoreProviderCell({ provider }: { provider: string }) { + const { displayName, logo } = getProviderLogoAndName(provider); + return ( +
+ {logo ? ( + { + (event.currentTarget as HTMLImageElement).style.display = "none"; + }} + /> + ) : null} + {displayName} +
+ ); +} + +function VectorStoreFilesCell({ vectorStore }: { vectorStore: VectorStore }) { + const ingestedFiles = vectorStore.vector_store_metadata?.ingested_files || []; + if (ingestedFiles.length === 0) { + return -; + } + + const filenames = ingestedFiles.map((file) => file.filename || file.file_url || "Unknown").join(", "); + const displayText = + ingestedFiles.length === 1 + ? ingestedFiles[0].filename || ingestedFiles[0].file_url || "1 file" + : `${ingestedFiles.length} files`; + + return ( + {displayText}} + /> + ); +} + +interface VectorStoreRowActionsProps { + vectorStore: VectorStore; + onEdit: (vectorStoreId: string) => void; + onDelete: (vectorStoreId: string) => void; +} + +function VectorStoreRowActions({ vectorStore, onEdit, onDelete }: VectorStoreRowActionsProps) { + return ( + + + + + + onEdit(vectorStore.vector_store_id)}> + + Edit + + void copyToClipboard(vectorStore.vector_store_id, "Vector store ID copied")} + > + + Copy vector store ID + + + onDelete(vectorStore.vector_store_id)} + > + + Delete + + + + ); +} + +interface VectorStoreTableColumnsDeps { + onView: (vectorStoreId: string) => void; + onEdit: (vectorStoreId: string) => void; + onDelete: (vectorStoreId: string) => void; +} + +export const getVectorStoreTableColumns = ({ + onView, + onEdit, + onDelete, +}: VectorStoreTableColumnsDeps): ColumnDef[] => [ + { + id: "vector_store_id", + accessorKey: "vector_store_id", + meta: { title: "Vector Store ID" }, + header: ({ column }) => , + size: 220, + enableSorting: true, + cell: ({ row }) => ( + onView(row.original.vector_store_id)} + /> + ), + }, + { + id: "vector_store_name", + accessorKey: "vector_store_name", + meta: { title: "Name" }, + header: ({ column }) => , + size: 200, + enableSorting: true, + cell: ({ row }) => { + const name = row.original.vector_store_name; + return ( + + {name || "-"} + + ); + }, + }, + { + id: "vector_store_description", + accessorKey: "vector_store_description", + meta: { title: "Description" }, + header: "Description", + size: 280, + enableSorting: false, + cell: ({ row }) => { + const description = row.original.vector_store_description; + return ( + + {description || "-"} + + ); + }, + }, + { + id: "files", + meta: { title: "Files" }, + header: "Files", + size: 160, + enableSorting: false, + cell: ({ row }) => , + }, + { + id: "provider", + accessorKey: "custom_llm_provider", + meta: { title: "Provider" }, + header: "Provider", + size: 160, + enableSorting: false, + cell: ({ row }) => , + }, + { + id: "created_at", + accessorKey: "created_at", + sortingFn: "datetime", + meta: { title: "Created At" }, + header: ({ column }) => , + size: 150, + enableSorting: true, + cell: ({ row }) => , + }, + { + id: "updated_at", + accessorKey: "updated_at", + sortingFn: "datetime", + meta: { title: "Updated At" }, + header: ({ column }) => , + size: 150, + enableSorting: true, + cell: ({ row }) => , + }, + { + id: "actions", + meta: { className: "text-right", headerClassName: "text-right" }, + header: () => Actions, + size: 64, + enableSorting: false, + enableHiding: false, + cell: ({ row }) => ( +
+ +
+ ), + }, +]; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/vector-stores/_components/index.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/vector-stores/_components/index.test.tsx new file mode 100644 index 00000000000..2931372f384 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/vector-stores/_components/index.test.tsx @@ -0,0 +1,53 @@ +import { render, screen } from "@testing-library/react"; +import { beforeEach, describe, expect, it, vi } from "vitest"; + +import { vectorStoreListCall } from "@/components/networking"; + +import VectorStoreManagement from "./index"; + +vi.mock("@/components/networking", () => ({ + vectorStoreListCall: vi.fn(), + vectorStoreDeleteCall: vi.fn(), + credentialListCall: vi.fn(), +})); + +vi.mock("./VectorStoreTable", () => ({ + __esModule: true, + default: ({ isLoading }: { isLoading?: boolean }) => ( +
{isLoading ? "table-loading" : "table-loaded"}
+ ), +})); + +vi.mock("./VectorStoreForm", () => ({ __esModule: true, default: () => null })); +vi.mock("./vector_store_info", () => ({ __esModule: true, default: () => null })); +vi.mock("./CreateVectorStore", () => ({ __esModule: true, default: () => null })); +vi.mock("./TestVectorStoreTab", () => ({ __esModule: true, default: () => null })); + +const mockVectorStoreListCall = vi.mocked(vectorStoreListCall); + +describe("VectorStoreManagement loading state", () => { + beforeEach(() => { + vi.clearAllMocks(); + }); + + it("should resolve the loading state when accessToken is null instead of showing the skeleton forever", async () => { + render(); + expect(await screen.findByText("table-loaded")).toBeInTheDocument(); + expect(mockVectorStoreListCall).not.toHaveBeenCalled(); + }); + + it("should show the loading state until the vector store fetch settles", async () => { + let resolveFetch: (value: { data: never[] }) => void = () => {}; + mockVectorStoreListCall.mockReturnValue( + new Promise((resolve) => { + resolveFetch = resolve; + }), + ); + render(); + expect(screen.getByText("table-loading")).toBeInTheDocument(); + + resolveFetch({ data: [] }); + expect(await screen.findByText("table-loaded")).toBeInTheDocument(); + expect(mockVectorStoreListCall).toHaveBeenCalledWith("sk-test"); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/vector-stores/_components/index.tsx b/ui/litellm-dashboard/src/app/(dashboard)/vector-stores/_components/index.tsx index 565e645f7b5..5d3f81f0275 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/vector-stores/_components/index.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/vector-stores/_components/index.tsx @@ -36,6 +36,7 @@ interface VectorStoreProps { const VectorStoreManagement: React.FC = ({ accessToken, userID, userRole }) => { const [vectorStores, setVectorStores] = useState([]); + const [isLoadingVectorStores, setIsLoadingVectorStores] = useState(true); const [isCreateModalVisible, setIsCreateModalVisible] = useState(false); const [isDeleteModalOpen, setIsDeleteModalOpen] = useState(false); const [vectorStoreToDelete, setVectorStoreToDelete] = useState(null); @@ -46,13 +47,18 @@ const VectorStoreManagement: React.FC = ({ accessToken, userID const [isDeleting, setIsDeleting] = useState(false); const fetchVectorStores = async () => { - if (!accessToken) return; + if (!accessToken) { + setIsLoadingVectorStores(false); + return; + } try { const response = await vectorStoreListCall(accessToken); setVectorStores(response.data || []); } catch (error) { console.error("Error fetching vector stores:", error); NotificationsManager.fromBackend("Error fetching vector stores: " + error); + } finally { + setIsLoadingVectorStores(false); } }; @@ -181,6 +187,7 @@ const VectorStoreManagement: React.FC = ({ accessToken, userID { expect(screen.getByText("The invitation link may be invalid or expired.")).toBeInTheDocument(); }); - it("should render a Back to Login link pointing to /ui/login", () => { + it("should render a Back to Login link pointing to /ui/login/", () => { render(); // antd Button with href renders as an
element const link = screen.getByRole("link", { name: "Back to Login" }); - expect(link).toHaveAttribute("href", "/ui/login"); + expect(link).toHaveAttribute("href", "/ui/login/"); }); }); diff --git a/ui/litellm-dashboard/src/app/onboarding/OnboardingErrorView.tsx b/ui/litellm-dashboard/src/app/onboarding/OnboardingErrorView.tsx index ca0f57c56ce..3de9a9ffaae 100644 --- a/ui/litellm-dashboard/src/app/onboarding/OnboardingErrorView.tsx +++ b/ui/litellm-dashboard/src/app/onboarding/OnboardingErrorView.tsx @@ -1,5 +1,6 @@ import React from "react"; import { Alert, Button } from "antd"; +import { getLoginUrl } from "@/utils/returnUrlUtils"; export function OnboardingErrorView() { return ( @@ -11,7 +12,7 @@ export function OnboardingErrorView() { showIcon />
- +
); diff --git a/ui/litellm-dashboard/src/components/AIHub/ModelHubTable.test.tsx b/ui/litellm-dashboard/src/components/AIHub/ModelHubTable.test.tsx index 3a22a55298e..dfe307c7edf 100644 --- a/ui/litellm-dashboard/src/components/AIHub/ModelHubTable.test.tsx +++ b/ui/litellm-dashboard/src/components/AIHub/ModelHubTable.test.tsx @@ -1,5 +1,5 @@ import * as networking from "@/components/networking"; -import { afterEach, describe, expect, it, vi } from "vitest"; +import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; import { renderWithProviders, screen, waitFor } from "../../../tests/test-utils"; import ModelHubTable from "./ModelHubTable"; @@ -7,6 +7,7 @@ const mockUseUISettings = vi.hoisted(() => vi.fn()); const mockGetCookie = vi.hoisted(() => vi.fn()); const mockCheckTokenValidity = vi.hoisted(() => vi.fn()); const mockRouterReplace = vi.hoisted(() => vi.fn()); +const mockLocationReplace = vi.hoisted(() => vi.fn()); vi.mock("@/components/networking", () => ({ getUiConfig: vi.fn(), @@ -43,7 +44,28 @@ vi.mock("@/utils/jwtUtils", () => ({ })); describe("ModelHubTable", () => { + const originalLocation = window.location; + + beforeEach(() => { + Object.defineProperty(window, "location", { + value: { + href: "http://localhost:4000/ui/model_hub_table", + origin: "http://localhost:4000", + hostname: "localhost", + pathname: "/ui/model_hub_table", + search: "", + protocol: "http:", + replace: mockLocationReplace, + }, + writable: true, + }); + }); + afterEach(() => { + Object.defineProperty(window, "location", { + value: originalLocation, + writable: true, + }); vi.clearAllMocks(); }); @@ -60,6 +82,7 @@ describe("ModelHubTable", () => { mockGetCookie.mockReturnValue(tokenValue); mockCheckTokenValidity.mockReturnValue(isTokenValid); mockRouterReplace.mockClear(); + mockLocationReplace.mockClear(); // Setup other required mocks vi.mocked(networking.getUiConfig).mockResolvedValue({ @@ -92,9 +115,10 @@ describe("ModelHubTable", () => { await waitFor(() => { if (shouldRedirect) { - expect(mockRouterReplace).toHaveBeenCalledWith("http://localhost:4000/ui/login"); - } else { + expect(mockLocationReplace).toHaveBeenCalledWith("http://localhost:4000/ui/login/"); expect(mockRouterReplace).not.toHaveBeenCalled(); + } else { + expect(mockLocationReplace).not.toHaveBeenCalled(); } }); }); diff --git a/ui/litellm-dashboard/src/components/AIHub/ModelHubTable.tsx b/ui/litellm-dashboard/src/components/AIHub/ModelHubTable.tsx index eb47f775d90..1f64d175052 100644 --- a/ui/litellm-dashboard/src/components/AIHub/ModelHubTable.tsx +++ b/ui/litellm-dashboard/src/components/AIHub/ModelHubTable.tsx @@ -33,6 +33,7 @@ import { Prism as SyntaxHighlighter } from "react-syntax-highlighter"; import { useUISettings } from "@/app/(dashboard)/hooks/uiSettings/useUISettings"; import { checkTokenValidity } from "@/utils/jwtUtils"; import { getCookie } from "@/utils/cookieUtils"; +import { getLoginUrl } from "@/utils/returnUrlUtils"; interface ModelHubTableProps { accessToken: string | null; @@ -108,12 +109,12 @@ const ModelHubTable: React.FC = ({ accessToken, publicPage, // If token is invalid, redirect to login if (!isTokenValid) { - router.replace(`${getProxyBaseUrl()}/ui/login`); + window.location.replace(getLoginUrl(getProxyBaseUrl())); return; } } // If require_auth_for_public_ai_hub is false, allow public access (no change) - }, [isUISettingsLoading, publicPage, uiSettings, router]); + }, [isUISettingsLoading, publicPage, uiSettings]); useEffect(() => { const fetchData = async (accessToken: string) => { diff --git a/ui/litellm-dashboard/src/components/DashboardHeader.tsx b/ui/litellm-dashboard/src/components/DashboardHeader.tsx index 8b928f8e6fe..babf007d475 100644 --- a/ui/litellm-dashboard/src/components/DashboardHeader.tsx +++ b/ui/litellm-dashboard/src/components/DashboardHeader.tsx @@ -18,7 +18,7 @@ import WorkerDropdown from "@/components/Navbar/WorkerDropdown/WorkerDropdown"; import { useWorker } from "@/hooks/useWorker"; import { useDisableShowPrompts } from "@/app/(dashboard)/hooks/useDisableShowPrompts"; import { clearTokenCookies } from "@/utils/cookieUtils"; -import { clearStoredReturnUrl } from "@/utils/returnUrlUtils"; +import { clearStoredReturnUrl, getLoginUrl } from "@/utils/returnUrlUtils"; interface DashboardHeaderProps { page: string; @@ -37,7 +37,7 @@ export function DashboardHeader({ page }: DashboardHeaderProps) { clearStoredReturnUrl(); localStorage.removeItem("litellm_selected_worker_id"); localStorage.removeItem("litellm_worker_url"); - window.location.href = `/ui/login?worker=${encodeURIComponent(workerId)}`; + window.location.href = `${getLoginUrl()}?worker=${encodeURIComponent(workerId)}`; }; return ( diff --git a/ui/litellm-dashboard/src/components/DeletedKeysPage/DeletedKeysPage.test.tsx b/ui/litellm-dashboard/src/components/DeletedKeysPage/DeletedKeysPage.test.tsx index 46df98a31de..05361f1c858 100644 --- a/ui/litellm-dashboard/src/components/DeletedKeysPage/DeletedKeysPage.test.tsx +++ b/ui/litellm-dashboard/src/components/DeletedKeysPage/DeletedKeysPage.test.tsx @@ -1,4 +1,5 @@ import { screen } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; import { vi, it, expect, beforeEach, MockedFunction } from "vitest"; import { renderWithProviders } from "../../../tests/test-utils"; import DeletedKeysPage from "./DeletedKeysPage"; @@ -13,6 +14,9 @@ const mockUseDeletedKeys = useDeletedKeys as MockedFunction { current_page: 1, total_pages: 1, }, - isPending: false, - isFetching: false, - } as any); + isLoading: false, + } as unknown as ReturnType); }); it("should render DeletedKeysPage component", () => { @@ -89,14 +92,32 @@ it("should render DeletedKeysPage component", () => { expect(screen.getByText("Test Key Alias")).toBeInTheDocument(); }); -it("should handle loading state", () => { +it("should show skeleton rows while the initial load is pending", () => { mockUseDeletedKeys.mockReturnValue({ data: undefined, - isPending: true, - isFetching: false, - } as any); + isLoading: true, + } as unknown as ReturnType); renderWithProviders(); - expect(screen.getByText("🚅 Loading keys...")).toBeInTheDocument(); + expect(screen.getAllByTestId("skeleton-row").length).toBeGreaterThan(0); +}); + +it("should request the next page from the hook when the pagination next button is clicked", async () => { + const user = userEvent.setup(); + mockUseDeletedKeys.mockReturnValue({ + data: { + keys: [mockDeletedKey], + total_count: 120, + current_page: 1, + total_pages: 3, + }, + isLoading: false, + } as unknown as ReturnType); + + renderWithProviders(); + + expect(mockUseDeletedKeys).toHaveBeenLastCalledWith(1, 50); + await user.click(screen.getByTestId("pagination-next")); + expect(mockUseDeletedKeys).toHaveBeenLastCalledWith(2, 50); }); diff --git a/ui/litellm-dashboard/src/components/DeletedKeysPage/DeletedKeysPage.tsx b/ui/litellm-dashboard/src/components/DeletedKeysPage/DeletedKeysPage.tsx index 8523710719e..78d500f5e64 100644 --- a/ui/litellm-dashboard/src/components/DeletedKeysPage/DeletedKeysPage.tsx +++ b/ui/litellm-dashboard/src/components/DeletedKeysPage/DeletedKeysPage.tsx @@ -1,5 +1,6 @@ "use client"; import { useState } from "react"; +import { PaginationState } from "@tanstack/react-table"; import { Alert } from "antd"; import { useDeletedKeys } from "@/app/(dashboard)/hooks/keys/useKeys"; import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; @@ -7,10 +8,9 @@ import { DeletedKeysTable } from "./DeletedKeysTable/DeletedKeysTable"; export default function DeletedKeysPage() { const { premiumUser } = useAuthorized(); - const [pageIndex, setPageIndex] = useState(0); - const [pageSize] = useState(50); + const [pagination, setPagination] = useState({ pageIndex: 0, pageSize: 50 }); - const { data: keysData, isPending: isLoading, isFetching } = useDeletedKeys(pageIndex + 1, pageSize); + const { data: keysData, isLoading } = useDeletedKeys(pagination.pageIndex + 1, pagination.pageSize); return (
@@ -27,10 +27,8 @@ export default function DeletedKeysPage() { keys={keysData?.keys || []} totalCount={keysData?.total_count || 0} isLoading={isLoading} - isFetching={isFetching} - pageIndex={pageIndex} - pageSize={pageSize} - onPageChange={setPageIndex} + pagination={pagination} + onPaginationChange={setPagination} />
); diff --git a/ui/litellm-dashboard/src/components/DeletedKeysPage/DeletedKeysTable/DeletedKeysTable.test.tsx b/ui/litellm-dashboard/src/components/DeletedKeysPage/DeletedKeysTable/DeletedKeysTable.test.tsx index 081ae0a80b5..7e30ef2c135 100644 --- a/ui/litellm-dashboard/src/components/DeletedKeysPage/DeletedKeysTable/DeletedKeysTable.test.tsx +++ b/ui/litellm-dashboard/src/components/DeletedKeysPage/DeletedKeysTable/DeletedKeysTable.test.tsx @@ -1,101 +1,88 @@ -import { screen } from "@testing-library/react"; +import { screen, within } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; import { vi, it, expect, beforeEach } from "vitest"; import { renderWithProviders } from "../../../../tests/test-utils"; import { DeletedKeysTable } from "./DeletedKeysTable"; import { DeletedKeyResponse } from "@/app/(dashboard)/hooks/keys/useKeys"; -const mockDeletedKey: DeletedKeyResponse = { - token: "sk-1234567890abcdef", - token_id: "key-1", - key_name: "test-key", - key_alias: "Test Key Alias", - spend: 5.5, - max_budget: 100, - expires: "2024-12-31T23:59:59Z", - models: ["gpt-3.5-turbo"], - aliases: {}, - config: {}, - user_id: "user-1", - team_id: "team-1", - max_parallel_requests: 10, - metadata: {}, - tpm_limit: 1000, - rpm_limit: 100, - duration: "30d", - budget_duration: "1m", - budget_reset_at: "2024-12-01T00:00:00Z", - allowed_cache_controls: [], - allowed_routes: [], - permissions: {}, - model_spend: {}, - model_max_budget: {}, - soft_budget_cooldown: false, - blocked: false, - litellm_budget_table: {}, - organization_id: "org-1", - created_at: "2024-11-01T10:00:00Z", - updated_at: "2024-11-15T10:00:00Z", - team_spend: 5.5, - team_alias: "Test Team", - team_tpm_limit: 5000, - team_rpm_limit: 500, - team_max_budget: 500, - team_models: ["gpt-3.5-turbo"], - team_blocked: false, - soft_budget: 50, - team_model_aliases: {}, - team_member_spend: 0, - team_metadata: {}, - end_user_id: "end-user-1", - end_user_tpm_limit: 100, - end_user_rpm_limit: 10, - end_user_max_budget: 10, - last_refreshed_at: Date.now(), - api_key: "sk-1234567890abcdef", - user_role: "user", - rpm_limit_per_model: {}, - tpm_limit_per_model: {}, - user_tpm_limit: 1000, - user_rpm_limit: 100, - user_email: "user@example.com", - deleted_at: "2024-11-15T10:00:00Z", - deleted_by: "user-1", +const makeDeletedKey = (overrides: Partial = {}): DeletedKeyResponse => + ({ + token: "sk-1234567890abcdef", + token_id: "key-1", + key_name: "test-key", + key_alias: "Test Key Alias", + spend: 5.5, + max_budget: 100, + models: ["gpt-3.5-turbo"], + user_id: "user-1", + team_id: "team-1", + organization_id: "org-1", + created_at: "2024-11-01T10:00:00Z", + updated_at: "2024-11-15T10:00:00Z", + created_by: "creator-1", + team_alias: "Test Team", + user_email: "user@example.com", + deleted_at: "2024-11-15T10:00:00Z", + deleted_by: "user-1", + ...overrides, + }) as DeletedKeyResponse; + +const defaultProps = { + keys: [makeDeletedKey()], + totalCount: 1, + isLoading: false, + pagination: { pageIndex: 0, pageSize: 50 }, + onPaginationChange: vi.fn(), }; beforeEach(() => { vi.clearAllMocks(); }); -it("should render DeletedKeysTable component", () => { - renderWithProviders( - , - ); - - expect(screen.getByText("Test Key Alias")).toBeInTheDocument(); -}); - -it("should display key information correctly", () => { - renderWithProviders( - , - ); +it("should display key information", () => { + renderWithProviders(); expect(screen.getByText("Test Key Alias")).toBeInTheDocument(); expect(screen.getByText("sk-1234567890abcdef")).toBeInTheDocument(); - expect(screen.getByText("Showing 1 - 1 of 1 results")).toBeInTheDocument(); + expect(screen.getByText("user@example.com")).toBeInTheDocument(); +}); + +it("should show the total count in the pagination footer", () => { + renderWithProviders(); + + expect(screen.getByTestId("pagination-range")).toHaveTextContent("Showing 1-50 of 120"); + expect(screen.getByTestId("pagination-page")).toHaveTextContent("Page 1 of 3"); +}); + +it("should propagate pagination changes when the next page button is clicked", async () => { + const user = userEvent.setup(); + renderWithProviders(); + + await user.click(screen.getByTestId("pagination-next")); + + expect(defaultProps.onPaginationChange).toHaveBeenCalled(); +}); + +it("should sort the current page by deleted_at descending by default", () => { + const keys = [ + makeDeletedKey({ token: "sk-older", key_alias: "older-key", deleted_at: "2024-01-01T10:00:00Z" }), + makeDeletedKey({ token: "sk-newer", key_alias: "newer-key", deleted_at: "2024-06-01T10:00:00Z" }), + ]; + renderWithProviders(); + + const rows = screen.getAllByRole("row").slice(1); + expect(within(rows[0]).getByText("newer-key")).toBeInTheDocument(); + expect(within(rows[1]).getByText("older-key")).toBeInTheDocument(); +}); + +it("should show skeleton rows when loading", () => { + renderWithProviders(); + + expect(screen.getAllByTestId("skeleton-row").length).toBeGreaterThan(0); +}); + +it("should show the empty state when there are no deleted keys", () => { + renderWithProviders(); + + expect(screen.getByText("No deleted keys found")).toBeInTheDocument(); }); diff --git a/ui/litellm-dashboard/src/components/DeletedKeysPage/DeletedKeysTable/DeletedKeysTable.tsx b/ui/litellm-dashboard/src/components/DeletedKeysPage/DeletedKeysTable/DeletedKeysTable.tsx index d4a120d0589..bc6941a7860 100644 --- a/ui/litellm-dashboard/src/components/DeletedKeysPage/DeletedKeysTable/DeletedKeysTable.tsx +++ b/ui/litellm-dashboard/src/components/DeletedKeysPage/DeletedKeysTable/DeletedKeysTable.tsx @@ -1,364 +1,63 @@ "use client"; -import { DateCell, IdCell, MoneyCell } from "@/components/shared/table_cells"; -import { ChevronDownIcon, ChevronUpIcon, SwitchVerticalIcon } from "@heroicons/react/outline"; -import { - ColumnDef, - flexRender, - getCoreRowModel, - getPaginationRowModel, - getSortedRowModel, - PaginationState, - SortingState, - useReactTable, -} from "@tanstack/react-table"; -import { Table, TableBody, TableCell, TableHead, TableHeaderCell, TableRow } from "@tremor/react"; -import { Tooltip } from "antd"; -import React, { useState } from "react"; -import { KeyResponse } from "../../key_team_helpers/key_list"; + +import { OnChangeFn, PaginationState, SortingState } from "@tanstack/react-table"; +import { Inbox } from "lucide-react"; +import { useMemo, useState } from "react"; + +import { DataTable } from "@/components/shared/DataTable"; +import { DeletedKeyResponse } from "@/app/(dashboard)/hooks/keys/useKeys"; + +import { getDeletedKeysTableColumns } from "./DeletedKeysTableColumns"; interface DeletedKeysTableProps { - keys: KeyResponse[]; + keys: DeletedKeyResponse[]; totalCount: number; isLoading: boolean; - isFetching: boolean; - pageIndex: number; - pageSize: number; - onPageChange: (pageIndex: number) => void; + pagination: PaginationState; + onPaginationChange: OnChangeFn; +} + +const DEFAULT_SORTING: SortingState = [{ id: "deleted_at", desc: true }]; + +function EmptyState() { + return ( +
+
+ +
+
No deleted keys found
+
Keys deleted from this proxy will show up here.
+
+ ); } export function DeletedKeysTable({ keys, totalCount, isLoading, - isFetching, - pageIndex, - pageSize, - onPageChange, + pagination, + onPaginationChange, }: DeletedKeysTableProps) { - const [sorting, setSorting] = useState([ - { - id: "deleted_at", - desc: true, - }, - ]); + const [sorting, setSorting] = useState(DEFAULT_SORTING); - const [tablePagination, setTablePagination] = useState({ - pageIndex, - pageSize, - }); - - // Sync pagination state when prop changes - React.useEffect(() => { - setTablePagination({ pageIndex, pageSize }); - }, [pageIndex, pageSize]); - - const columns: ColumnDef[] = [ - { - id: "token", - accessorKey: "token", - header: "Key ID", - size: 150, - maxSize: 250, - cell: (info) => , - }, - { - id: "key_alias", - accessorKey: "key_alias", - header: "Key Alias", - size: 150, - maxSize: 200, - cell: (info) => { - const value = info.getValue() as string; - return ( - - {value ?? "-"} - - ); - }, - }, - { - id: "team_alias", - accessorKey: "team_alias", - header: "Team Alias", - size: 120, - maxSize: 180, - cell: (info) => { - const value = info.getValue() as string; - return {value || "-"}; - }, - }, - { - id: "spend", - accessorKey: "spend", - header: "Spend (USD)", - size: 100, - maxSize: 140, - cell: (info) => , - }, - { - id: "max_budget", - accessorKey: "max_budget", - header: "Budget (USD)", - size: 110, - maxSize: 150, - cell: (info) => ( - - ), - }, - { - id: "user_email", - accessorKey: "user_email", - header: "User Email", - size: 160, - maxSize: 250, - cell: (info) => { - const value = info.getValue() as string; - return ( - - {value ?? "-"} - - ); - }, - }, - { - id: "user_id", - accessorKey: "user_id", - header: "User ID", - size: 120, - maxSize: 200, - cell: (info) => , - }, - { - id: "created_at", - accessorKey: "created_at", - header: "Created At", - size: 120, - maxSize: 140, - cell: (info) => , - }, - { - id: "created_by", - accessorKey: "created_by", - header: "Created By", - size: 120, - maxSize: 180, - cell: (info) => { - const value = (info.row.original as any).created_by as string | null | undefined; - return ( - - {value || "-"} - - ); - }, - }, - { - id: "deleted_at", - accessorKey: "deleted_at", - header: "Deleted At", - size: 120, - maxSize: 140, - cell: (info) => ( - - ), - }, - { - id: "deleted_by", - accessorKey: "deleted_by", - header: "Deleted By", - size: 120, - maxSize: 180, - cell: (info) => { - const value = (info.row.original as any).deleted_by as string | null | undefined; - return ( - - {value || "-"} - - ); - }, - }, - ]; - - const table = useReactTable({ - data: keys, - columns, - columnResizeMode: "onChange", - columnResizeDirection: "ltr", - state: { - sorting, - pagination: tablePagination, - }, - onSortingChange: setSorting, - onPaginationChange: (updater) => { - const newPagination = typeof updater === "function" ? updater(tablePagination) : updater; - setTablePagination(newPagination); - onPageChange(newPagination.pageIndex); - }, - getCoreRowModel: getCoreRowModel(), - getSortedRowModel: getSortedRowModel(), - getPaginationRowModel: getPaginationRowModel(), - enableSorting: true, - manualSorting: false, - manualPagination: true, - pageCount: Math.ceil(totalCount / pageSize), - }); - - const { pageIndex: currentPageIndex } = table.getState().pagination; - const start = currentPageIndex * pageSize + 1; - const end = Math.min((currentPageIndex + 1) * pageSize, totalCount); - const rangeLabel = `${start} - ${end}`; + const columns = useMemo(() => getDeletedKeysTableColumns(), []); return ( -
-
-
- {isLoading || isFetching ? ( - Loading... - ) : ( - - Showing {rangeLabel} of {totalCount} results - - )} - -
- {isLoading || isFetching ? ( - Loading... - ) : ( - - Page {currentPageIndex + 1} of {table.getPageCount()} - - )} - - - - -
-
-
-
-
- - - {table.getHeaderGroups().map((headerGroup) => ( - - {headerGroup.headers.map((header) => ( - { - const resizer = document.querySelector(`[data-header-id="${header.id}"] .resizer`); - if (resizer) { - (resizer as HTMLElement).style.opacity = "0.5"; - } - }} - onMouseLeave={() => { - const resizer = document.querySelector(`[data-header-id="${header.id}"] .resizer`); - if (resizer && !header.column.getIsResizing()) { - (resizer as HTMLElement).style.opacity = "0"; - } - }} - onClick={header.column.getToggleSortingHandler()} - > -
-
- {header.isPlaceholder - ? null - : flexRender(header.column.columnDef.header, header.getContext())} -
-
- {header.column.getIsSorted() ? ( - { - asc: , - desc: , - }[header.column.getIsSorted() as string] - ) : ( - - )} -
-
header.column.resetSize()} - onMouseDown={header.getResizeHandler()} - onTouchStart={header.getResizeHandler()} - className={`resizer ${table.options.columnResizeDirection} ${header.column.getIsResizing() ? "isResizing" : ""}`} - style={{ - position: "absolute", - right: 0, - top: 0, - height: "100%", - width: "5px", - background: header.column.getIsResizing() ? "#3b82f6" : "transparent", - cursor: "col-resize", - userSelect: "none", - touchAction: "none", - opacity: header.column.getIsResizing() ? 1 : 0, - }} - /> -
- - ))} - - ))} - - - {isLoading || isFetching ? ( - - -
-

🚅 Loading keys...

-
-
-
- ) : keys.length > 0 ? ( - table.getRowModel().rows.map((row) => ( - - {row.getVisibleCells().map((cell) => ( - - {flexRender(cell.column.columnDef.cell, cell.getContext())} - - ))} - - )) - ) : ( - - -
-

No deleted keys found

-
-
-
- )} -
-
-
-
-
-
-
+ key.token || String(index)} + sortingMode="client" + sorting={sorting} + onSortingChange={setSorting} + paginationMode="server" + pagination={pagination} + onPaginationChange={onPaginationChange} + rowCount={totalCount} + isLoading={isLoading} + loadingMessage="Loading deleted keys…" + noDataMessage={} + size="compact" + /> ); } diff --git a/ui/litellm-dashboard/src/components/DeletedKeysPage/DeletedKeysTable/DeletedKeysTableColumns.tsx b/ui/litellm-dashboard/src/components/DeletedKeysPage/DeletedKeysTable/DeletedKeysTableColumns.tsx new file mode 100644 index 00000000000..aa7d6380cc3 --- /dev/null +++ b/ui/litellm-dashboard/src/components/DeletedKeysPage/DeletedKeysTable/DeletedKeysTableColumns.tsx @@ -0,0 +1,130 @@ +"use client"; + +import { ColumnDef } from "@tanstack/react-table"; + +import { DataTableSortHeader } from "@/components/shared/DataTable"; +import { DateCell, IdCell, MoneyCell } from "@/components/shared/table_cells"; +import { DeletedKeyResponse } from "@/app/(dashboard)/hooks/keys/useKeys"; + +function TruncatedTextCell({ value }: { value: string | null | undefined }) { + if (!value) { + return -; + } + return ( + + {value} + + ); +} + +export const getDeletedKeysTableColumns = (): ColumnDef[] => [ + { + id: "token", + accessorKey: "token", + meta: { title: "Key ID" }, + header: "Key ID", + size: 150, + enableSorting: false, + cell: ({ row }) => , + }, + { + id: "key_alias", + accessorKey: "key_alias", + meta: { title: "Key Alias" }, + header: "Key Alias", + size: 150, + enableSorting: false, + cell: ({ row }) => { + const value = row.original.key_alias; + if (!value) { + return -; + } + return ( + + {value} + + ); + }, + }, + { + id: "team_alias", + accessorKey: "team_alias", + meta: { title: "Team Alias" }, + header: "Team Alias", + size: 120, + enableSorting: false, + cell: ({ row }) => , + }, + { + id: "spend", + accessorKey: "spend", + meta: { title: "Spend (USD)", numeric: true }, + header: ({ column }) => , + size: 100, + enableSorting: true, + cell: ({ row }) => , + }, + { + id: "max_budget", + accessorKey: "max_budget", + meta: { title: "Budget (USD)", numeric: true }, + header: "Budget (USD)", + size: 110, + enableSorting: false, + cell: ({ row }) => , + }, + { + id: "user_email", + accessorKey: "user_email", + meta: { title: "User Email" }, + header: "User Email", + size: 160, + enableSorting: false, + cell: ({ row }) => , + }, + { + id: "user_id", + accessorKey: "user_id", + meta: { title: "User ID" }, + header: "User ID", + size: 120, + enableSorting: false, + cell: ({ row }) => , + }, + { + id: "created_at", + accessorKey: "created_at", + meta: { title: "Created At" }, + header: ({ column }) => , + size: 120, + enableSorting: true, + cell: ({ row }) => , + }, + { + id: "created_by", + accessorKey: "created_by", + meta: { title: "Created By" }, + header: "Created By", + size: 120, + enableSorting: false, + cell: ({ row }) => , + }, + { + id: "deleted_at", + accessorKey: "deleted_at", + meta: { title: "Deleted At" }, + header: ({ column }) => , + size: 120, + enableSorting: true, + cell: ({ row }) => , + }, + { + id: "deleted_by", + accessorKey: "deleted_by", + meta: { title: "Deleted By" }, + header: "Deleted By", + size: 120, + enableSorting: false, + cell: ({ row }) => , + }, +]; diff --git a/ui/litellm-dashboard/src/components/DeletedTeamsPage/DeletedTeamsPage.test.tsx b/ui/litellm-dashboard/src/components/DeletedTeamsPage/DeletedTeamsPage.test.tsx index 77d2e94067e..a3eceb4b458 100644 --- a/ui/litellm-dashboard/src/components/DeletedTeamsPage/DeletedTeamsPage.test.tsx +++ b/ui/litellm-dashboard/src/components/DeletedTeamsPage/DeletedTeamsPage.test.tsx @@ -32,9 +32,8 @@ beforeEach(() => { mockUseDeletedTeams.mockReturnValue({ data: [mockDeletedTeam], - isPending: false, - isFetching: false, - } as any); + isLoading: false, + } as unknown as ReturnType); }); it("should render DeletedTeamsPage component", () => { @@ -43,14 +42,13 @@ it("should render DeletedTeamsPage component", () => { expect(screen.getByText("Test Team")).toBeInTheDocument(); }); -it("should handle loading state", () => { +it("should show skeleton rows while the initial load is pending", () => { mockUseDeletedTeams.mockReturnValue({ data: undefined, - isPending: true, - isFetching: false, - } as any); + isLoading: true, + } as unknown as ReturnType); renderWithProviders(); - expect(screen.getByText("🚅 Loading teams...")).toBeInTheDocument(); + expect(screen.getAllByTestId("skeleton-row").length).toBeGreaterThan(0); }); diff --git a/ui/litellm-dashboard/src/components/DeletedTeamsPage/DeletedTeamsPage.tsx b/ui/litellm-dashboard/src/components/DeletedTeamsPage/DeletedTeamsPage.tsx index 30803d0a004..0265a3b623e 100644 --- a/ui/litellm-dashboard/src/components/DeletedTeamsPage/DeletedTeamsPage.tsx +++ b/ui/litellm-dashboard/src/components/DeletedTeamsPage/DeletedTeamsPage.tsx @@ -6,7 +6,7 @@ import { DeletedTeamsTable } from "./DeletedTeamsTable/DeletedTeamsTable"; export default function DeletedTeamsPage() { const { premiumUser } = useAuthorized(); - const { data: teamsData, isPending: isLoading, isFetching } = useDeletedTeams(1, 100); + const { data: teamsData, isLoading } = useDeletedTeams(1, 100); return (
@@ -19,7 +19,7 @@ export default function DeletedTeamsPage() { description="Deleted team auditing is graduating from beta into our Enterprise audit & compliance suite." /> )} - +
); } diff --git a/ui/litellm-dashboard/src/components/DeletedTeamsPage/DeletedTeamsTable/DeletedTeamsTable.test.tsx b/ui/litellm-dashboard/src/components/DeletedTeamsPage/DeletedTeamsTable/DeletedTeamsTable.test.tsx index 358a9e90fa1..c0cc5a342a8 100644 --- a/ui/litellm-dashboard/src/components/DeletedTeamsPage/DeletedTeamsTable/DeletedTeamsTable.test.tsx +++ b/ui/litellm-dashboard/src/components/DeletedTeamsPage/DeletedTeamsTable/DeletedTeamsTable.test.tsx @@ -1,10 +1,10 @@ -import { screen } from "@testing-library/react"; +import { screen, within } from "@testing-library/react"; import { vi, it, expect, beforeEach } from "vitest"; import { renderWithProviders } from "../../../../tests/test-utils"; import { DeletedTeamsTable } from "./DeletedTeamsTable"; import { DeletedTeam } from "@/app/(dashboard)/hooks/teams/useTeams"; -const mockDeletedTeam: DeletedTeam = { +const makeDeletedTeam = (overrides: Partial = {}): DeletedTeam => ({ team_id: "team-1", team_alias: "Test Team", models: ["gpt-3.5-turbo", "gpt-4"], @@ -19,22 +19,41 @@ const mockDeletedTeam: DeletedTeam = { deleted_at: "2024-11-15T10:00:00Z", deleted_by: "user-1", spend: 100.5, -}; + ...overrides, +}); beforeEach(() => { vi.clearAllMocks(); }); -it("should render DeletedTeamsTable component", () => { - renderWithProviders(); - - expect(screen.getByText("Test Team")).toBeInTheDocument(); -}); - -it("should display team information correctly", () => { - renderWithProviders(); +it("should display team information", () => { + renderWithProviders(); expect(screen.getByText("Test Team")).toBeInTheDocument(); expect(screen.getByText("team-1")).toBeInTheDocument(); - expect(screen.getByText("Showing 1 team")).toBeInTheDocument(); + expect(screen.getByText("org-1")).toBeInTheDocument(); +}); + +it("should sort teams by deleted_at descending by default", () => { + const teams = [ + makeDeletedTeam({ team_id: "team-old", team_alias: "older-team", deleted_at: "2024-01-01T10:00:00Z" }), + makeDeletedTeam({ team_id: "team-new", team_alias: "newer-team", deleted_at: "2024-06-01T10:00:00Z" }), + ]; + renderWithProviders(); + + const rows = screen.getAllByRole("row").slice(1); + expect(within(rows[0]).getByText("newer-team")).toBeInTheDocument(); + expect(within(rows[1]).getByText("older-team")).toBeInTheDocument(); +}); + +it("should show skeleton rows when loading", () => { + renderWithProviders(); + + expect(screen.getAllByTestId("skeleton-row").length).toBeGreaterThan(0); +}); + +it("should show the empty state when there are no deleted teams", () => { + renderWithProviders(); + + expect(screen.getByText("No deleted teams found")).toBeInTheDocument(); }); diff --git a/ui/litellm-dashboard/src/components/DeletedTeamsPage/DeletedTeamsTable/DeletedTeamsTable.tsx b/ui/litellm-dashboard/src/components/DeletedTeamsPage/DeletedTeamsTable/DeletedTeamsTable.tsx index ddfd5cf73b6..9578a52453f 100644 --- a/ui/litellm-dashboard/src/components/DeletedTeamsPage/DeletedTeamsTable/DeletedTeamsTable.tsx +++ b/ui/litellm-dashboard/src/components/DeletedTeamsPage/DeletedTeamsTable/DeletedTeamsTable.tsx @@ -1,299 +1,50 @@ "use client"; -import { DateCell, IdCell, MoneyCell } from "@/components/shared/table_cells"; -import { ChevronDownIcon, ChevronUpIcon, SwitchVerticalIcon } from "@heroicons/react/outline"; -import { - ColumnDef, - flexRender, - getCoreRowModel, - getSortedRowModel, - SortingState, - useReactTable, -} from "@tanstack/react-table"; -import { Table, TableBody, TableCell, TableHead, TableHeaderCell, TableRow, Badge, Text } from "@tremor/react"; -import { Tooltip } from "antd"; -import React, { useState } from "react"; + +import { SortingState } from "@tanstack/react-table"; +import { Inbox } from "lucide-react"; +import { useMemo, useState } from "react"; + +import { DataTable } from "@/components/shared/DataTable"; import { DeletedTeam } from "@/app/(dashboard)/hooks/teams/useTeams"; -import { getModelDisplayName } from "@/components/key_team_helpers/fetch_available_models_team_key"; + +import { getDeletedTeamsTableColumns } from "./DeletedTeamsTableColumns"; interface DeletedTeamsTableProps { teams: DeletedTeam[]; isLoading: boolean; - isFetching: boolean; } -export function DeletedTeamsTable({ teams, isLoading, isFetching }: DeletedTeamsTableProps) { - const [sorting, setSorting] = useState([ - { - id: "deleted_at", - desc: true, - }, - ]); - - const columns: ColumnDef[] = [ - { - id: "team_alias", - accessorKey: "team_alias", - header: "Team Name", - size: 150, - maxSize: 200, - cell: (info) => { - const value = info.getValue() as string; - return ( - - {value || "-"} - - ); - }, - }, - { - id: "team_id", - accessorKey: "team_id", - header: "Team ID", - size: 150, - maxSize: 250, - cell: (info) => , - }, - { - id: "created_at", - accessorKey: "created_at", - header: "Created", - size: 120, - maxSize: 140, - cell: (info) => , - }, - { - id: "spend", - accessorKey: "spend", - header: "Spend (USD)", - size: 100, - maxSize: 140, - cell: (info) => , - }, - { - id: "max_budget", - accessorKey: "max_budget", - header: "Budget (USD)", - size: 110, - maxSize: 150, - cell: (info) => ( - - ), - }, - { - id: "models", - accessorKey: "models", - header: "Models", - size: 200, - maxSize: 300, - cell: (info) => { - const models = info.getValue() as string[]; - if (!Array.isArray(models) || models.length === 0) { - return ( - - All Proxy Models - - ); - } - return ( -
- {models.slice(0, 3).map((model: string, index: number) => - model === "all-proxy-models" ? ( - - All Proxy Models - - ) : ( - - - {model.length > 30 ? `${getModelDisplayName(model).slice(0, 30)}...` : getModelDisplayName(model)} - - - ), - )} - {models.length > 3 && ( - - - +{models.length - 3} {models.length - 3 === 1 ? "more model" : "more models"} - - - )} -
- ); - }, - }, - { - id: "organization_id", - accessorKey: "organization_id", - header: "Organization", - size: 150, - maxSize: 200, - cell: (info) => , - }, - { - id: "deleted_at", - accessorKey: "deleted_at", - header: "Deleted At", - size: 120, - maxSize: 140, - cell: (info) => , - }, - { - id: "deleted_by", - accessorKey: "deleted_by", - header: "Deleted By", - size: 120, - maxSize: 180, - cell: (info) => { - const value = (info.row.original as any).deleted_by as string | null | undefined; - return ( - - {value || "-"} - - ); - }, - }, - ]; - - const table = useReactTable({ - data: teams, - columns, - columnResizeMode: "onChange", - columnResizeDirection: "ltr", - state: { - sorting, - }, - onSortingChange: setSorting, - getCoreRowModel: getCoreRowModel(), - getSortedRowModel: getSortedRowModel(), - enableSorting: true, - manualSorting: false, - }); +const DEFAULT_SORTING: SortingState = [{ id: "deleted_at", desc: true }]; +function EmptyState() { return ( -
-
-
- {isLoading || isFetching ? ( - Loading... - ) : ( - - Showing {teams.length} {teams.length === 1 ? "team" : "teams"} - - )} -
-
-
-
- - - {table.getHeaderGroups().map((headerGroup) => ( - - {headerGroup.headers.map((header) => ( - { - const resizer = document.querySelector(`[data-header-id="${header.id}"] .resizer`); - if (resizer) { - (resizer as HTMLElement).style.opacity = "0.5"; - } - }} - onMouseLeave={() => { - const resizer = document.querySelector(`[data-header-id="${header.id}"] .resizer`); - if (resizer && !header.column.getIsResizing()) { - (resizer as HTMLElement).style.opacity = "0"; - } - }} - onClick={header.column.getToggleSortingHandler()} - > -
-
- {header.isPlaceholder - ? null - : flexRender(header.column.columnDef.header, header.getContext())} -
-
- {header.column.getIsSorted() ? ( - { - asc: , - desc: , - }[header.column.getIsSorted() as string] - ) : ( - - )} -
-
header.column.resetSize()} - onMouseDown={header.getResizeHandler()} - onTouchStart={header.getResizeHandler()} - className={`resizer ${table.options.columnResizeDirection} ${header.column.getIsResizing() ? "isResizing" : ""}`} - style={{ - position: "absolute", - right: 0, - top: 0, - height: "100%", - width: "5px", - background: header.column.getIsResizing() ? "#3b82f6" : "transparent", - cursor: "col-resize", - userSelect: "none", - touchAction: "none", - opacity: header.column.getIsResizing() ? 1 : 0, - }} - /> -
- - ))} - - ))} - - - {isLoading || isFetching ? ( - - -
-

🚅 Loading teams...

-
-
-
- ) : teams.length > 0 ? ( - table.getRowModel().rows.map((row) => ( - - {row.getVisibleCells().map((cell) => ( - - {flexRender(cell.column.columnDef.cell, cell.getContext())} - - ))} - - )) - ) : ( - - -
-

No deleted teams found

-
-
-
- )} -
-
-
-
-
+
+
+
+
No deleted teams found
+
Teams deleted from this proxy will show up here.
); } + +export function DeletedTeamsTable({ teams, isLoading }: DeletedTeamsTableProps) { + const [sorting, setSorting] = useState(DEFAULT_SORTING); + + const columns = useMemo(() => getDeletedTeamsTableColumns(), []); + + return ( + team.team_id || String(index)} + sortingMode="client" + sorting={sorting} + onSortingChange={setSorting} + isLoading={isLoading} + loadingMessage="Loading deleted teams…" + noDataMessage={} + size="compact" + /> + ); +} diff --git a/ui/litellm-dashboard/src/components/DeletedTeamsPage/DeletedTeamsTable/DeletedTeamsTableColumns.tsx b/ui/litellm-dashboard/src/components/DeletedTeamsPage/DeletedTeamsTable/DeletedTeamsTableColumns.tsx new file mode 100644 index 00000000000..e36077fd2c3 --- /dev/null +++ b/ui/litellm-dashboard/src/components/DeletedTeamsPage/DeletedTeamsTable/DeletedTeamsTableColumns.tsx @@ -0,0 +1,111 @@ +"use client"; + +import { ColumnDef } from "@tanstack/react-table"; + +import { DataTableSortHeader } from "@/components/shared/DataTable"; +import { DateCell, IdCell, ModelsCell, MoneyCell } from "@/components/shared/table_cells"; +import { DeletedTeam } from "@/app/(dashboard)/hooks/teams/useTeams"; + +export const getDeletedTeamsTableColumns = (): ColumnDef[] => [ + { + id: "team_alias", + accessorKey: "team_alias", + meta: { title: "Team Name" }, + header: "Team Name", + size: 150, + enableSorting: false, + cell: ({ row }) => { + const value = row.original.team_alias; + if (!value) { + return -; + } + return ( + + {value} + + ); + }, + }, + { + id: "team_id", + accessorKey: "team_id", + meta: { title: "Team ID" }, + header: "Team ID", + size: 150, + enableSorting: false, + cell: ({ row }) => , + }, + { + id: "created_at", + accessorKey: "created_at", + meta: { title: "Created" }, + header: ({ column }) => , + size: 120, + enableSorting: true, + cell: ({ row }) => , + }, + { + id: "spend", + accessorKey: "spend", + meta: { title: "Spend (USD)", numeric: true }, + header: ({ column }) => , + size: 100, + enableSorting: true, + cell: ({ row }) => , + }, + { + id: "max_budget", + accessorKey: "max_budget", + meta: { title: "Budget (USD)", numeric: true }, + header: "Budget (USD)", + size: 110, + enableSorting: false, + cell: ({ row }) => , + }, + { + id: "models", + accessorKey: "models", + meta: { title: "Models", skeleton: "chips" }, + header: "Models", + size: 200, + enableSorting: false, + cell: ({ row }) => , + }, + { + id: "organization_id", + accessorKey: "organization_id", + meta: { title: "Organization" }, + header: "Organization", + size: 150, + enableSorting: false, + cell: ({ row }) => , + }, + { + id: "deleted_at", + accessorKey: "deleted_at", + meta: { title: "Deleted At" }, + header: ({ column }) => , + size: 120, + enableSorting: true, + cell: ({ row }) => , + }, + { + id: "deleted_by", + accessorKey: "deleted_by", + meta: { title: "Deleted By" }, + header: "Deleted By", + size: 120, + enableSorting: false, + cell: ({ row }) => { + const value = row.original.deleted_by; + if (!value) { + return -; + } + return ( + + {value} + + ); + }, + }, +]; diff --git a/ui/litellm-dashboard/src/components/PassThroughSettings/PassThroughEndpointsTable.test.tsx b/ui/litellm-dashboard/src/components/PassThroughSettings/PassThroughEndpointsTable.test.tsx new file mode 100644 index 00000000000..040fa463f9c --- /dev/null +++ b/ui/litellm-dashboard/src/components/PassThroughSettings/PassThroughEndpointsTable.test.tsx @@ -0,0 +1,130 @@ +import { render, screen } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import { beforeEach, describe, expect, it, vi } from "vitest"; + +import { PassThroughEndpointsTable } from "./PassThroughEndpointsTable"; +import type { passThroughItem } from "./PassThroughSettings"; + +const endpoints: passThroughItem[] = [ + { + id: "ep-1", + path: "/v1/rerank", + target: "https://api.cohere.com/v1/rerank", + headers: { Authorization: "Bearer secret-value" }, + auth: true, + methods: ["POST"], + }, + { + id: "ep-2", + path: "/bria", + target: "https://engine.prod.bria-api.com", + headers: {}, + auth: false, + }, +]; + +const defaultProps = { + endpoints, + isLoading: false, + onEndpointClick: vi.fn(), + onDeleteClick: vi.fn(), +}; + +describe("PassThroughEndpointsTable", () => { + beforeEach(() => { + vi.clearAllMocks(); + }); + + it("should render a row per endpoint with path and target", () => { + render(); + expect(screen.getByText("/v1/rerank")).toBeInTheDocument(); + expect(screen.getByText("https://api.cohere.com/v1/rerank")).toBeInTheDocument(); + expect(screen.getByText("/bria")).toBeInTheDocument(); + }); + + it("should open the endpoint when its ID is clicked", async () => { + const user = userEvent.setup(); + const onEndpointClick = vi.fn(); + render(); + await user.click(screen.getByRole("button", { name: "ep-1" })); + expect(onEndpointClick).toHaveBeenCalledWith("ep-1"); + }); + + it("should show method chips, or ALL when no methods are set", () => { + render(); + expect(screen.getByText("POST")).toBeInTheDocument(); + expect(screen.getByText("ALL")).toBeInTheDocument(); + }); + + it("should show authentication as Yes or No", () => { + render(); + expect(screen.getByText("Yes")).toBeInTheDocument(); + expect(screen.getByText("No")).toBeInTheDocument(); + }); + + it("should mask headers until the visibility toggle is clicked", async () => { + const user = userEvent.setup(); + render(); + + expect(screen.queryByText(/secret-value/)).not.toBeInTheDocument(); + const toggles = screen.getAllByRole("button", { name: "Show headers" }); + await user.click(toggles[0]); + expect(screen.getByText(/secret-value/)).toBeInTheDocument(); + }); + + it("should edit and delete an endpoint through the actions menu", async () => { + const user = userEvent.setup(); + const onEndpointClick = vi.fn(); + const onDeleteClick = vi.fn(); + render( + , + ); + + await user.click(screen.getByTestId("endpoint-actions-ep-1")); + await user.click(await screen.findByTestId("endpoint-action-edit")); + expect(onEndpointClick).toHaveBeenCalledWith("ep-1"); + + await user.click(screen.getByTestId("endpoint-actions-ep-1")); + await user.click(await screen.findByTestId("endpoint-action-delete")); + expect(onDeleteClick).toHaveBeenCalledWith("ep-1"); + }); + + it("should disable edit and delete for endpoints without an id", async () => { + const user = userEvent.setup(); + const onEndpointClick = vi.fn(); + const onDeleteClick = vi.fn(); + const endpointWithoutId: passThroughItem = { path: "/legacy", target: "https://legacy.example.com", headers: {} }; + render( + , + ); + + await user.click(screen.getByTestId("endpoint-actions-/legacy")); + const editItem = await screen.findByTestId("endpoint-action-edit"); + const deleteItem = await screen.findByTestId("endpoint-action-delete"); + + expect(editItem).toHaveAttribute("data-disabled"); + expect(deleteItem).toHaveAttribute("data-disabled"); + + await user.click(editItem); + await user.click(deleteItem); + + expect(onEndpointClick).not.toHaveBeenCalled(); + expect(onDeleteClick).not.toHaveBeenCalled(); + }); + + it("should show the empty state when there are no endpoints", () => { + render(); + expect(screen.getByText("No pass-through endpoints configured")).toBeInTheDocument(); + }); + + it("should show skeleton rows while loading", () => { + render(); + expect(screen.getAllByTestId("skeleton-row").length).toBeGreaterThan(0); + expect(screen.queryByText("No pass-through endpoints configured")).not.toBeInTheDocument(); + }); +}); diff --git a/ui/litellm-dashboard/src/components/PassThroughSettings/PassThroughEndpointsTable.tsx b/ui/litellm-dashboard/src/components/PassThroughSettings/PassThroughEndpointsTable.tsx new file mode 100644 index 00000000000..754e7ff68dd --- /dev/null +++ b/ui/litellm-dashboard/src/components/PassThroughSettings/PassThroughEndpointsTable.tsx @@ -0,0 +1,52 @@ +"use client"; + +import { Waypoints } from "lucide-react"; +import { useMemo } from "react"; + +import { DataTable } from "@/components/shared/DataTable"; + +import { getPassThroughEndpointsTableColumns } from "./PassThroughEndpointsTableColumns"; +import type { passThroughItem } from "./PassThroughSettings"; + +interface PassThroughEndpointsTableProps { + endpoints: passThroughItem[]; + isLoading: boolean; + onEndpointClick: (endpointId: string) => void; + onDeleteClick: (endpointId: string) => void; +} + +function EmptyState() { + return ( +
+
+ +
+
No pass-through endpoints configured
+
Add a pass-through endpoint to route custom paths.
+
+ ); +} + +export function PassThroughEndpointsTable({ + endpoints, + isLoading, + onEndpointClick, + onDeleteClick, +}: PassThroughEndpointsTableProps) { + const columns = useMemo( + () => getPassThroughEndpointsTableColumns({ onEndpointClick, onDeleteClick }), + [onEndpointClick, onDeleteClick], + ); + + return ( + endpoint.id || endpoint.path || String(index)} + isLoading={isLoading} + loadingMessage="Loading pass-through endpoints…" + noDataMessage={} + size="compact" + /> + ); +} diff --git a/ui/litellm-dashboard/src/components/PassThroughSettings/PassThroughEndpointsTableColumns.tsx b/ui/litellm-dashboard/src/components/PassThroughSettings/PassThroughEndpointsTableColumns.tsx new file mode 100644 index 00000000000..d22b274861a --- /dev/null +++ b/ui/litellm-dashboard/src/components/PassThroughSettings/PassThroughEndpointsTableColumns.tsx @@ -0,0 +1,203 @@ +"use client"; + +import { ColumnDef } from "@tanstack/react-table"; +import { Eye, EyeOff, Info, MoreHorizontal, Pencil, Trash2 } from "lucide-react"; +import React, { useState } from "react"; + +import { CellTooltip, IdentityCell, StatusBadge } from "@/components/shared/table_cells"; +import { Badge } from "@/components/ui/badge"; +import { buttonVariants } from "@/components/ui/button"; +import { + DropdownMenu, + DropdownMenuContent, + DropdownMenuItem, + DropdownMenuSeparator, + DropdownMenuTrigger, +} from "@/components/ui/dropdown-menu"; +import { cn } from "@/lib/cva.config"; + +import type { passThroughItem } from "./PassThroughSettings"; + +function HeaderWithTooltip({ title, tooltip }: { title: string; tooltip: string }) { + return ( +
+ {title} + } /> +
+ ); +} + +function HeadersCell({ value }: { value: object }) { + const [showHeaders, setShowHeaders] = useState(false); + const headerString = JSON.stringify(value); + + return ( +
+ {showHeaders ? headerString : "••••••••"} + +
+ ); +} + +function MethodsCell({ methods }: { methods: string[] | undefined }) { + if (!methods || methods.length === 0) { + return ALL; + } + return ( +
+ {methods.map((method) => ( + + {method} + + ))} +
+ ); +} + +interface EndpointRowActionsProps { + endpoint: passThroughItem; + onEndpointClick: (endpointId: string) => void; + onDeleteClick: (endpointId: string) => void; +} + +function EndpointRowActions({ endpoint, onEndpointClick, onDeleteClick }: EndpointRowActionsProps) { + const endpointId = endpoint.id; + return ( + + + + + + endpointId && onEndpointClick(endpointId)} + > + + Edit + + + endpointId && onDeleteClick(endpointId)} + > + + Delete + + + + ); +} + +interface PassThroughEndpointsTableColumnsDeps { + onEndpointClick: (endpointId: string) => void; + onDeleteClick: (endpointId: string) => void; +} + +export const getPassThroughEndpointsTableColumns = ({ + onEndpointClick, + onDeleteClick, +}: PassThroughEndpointsTableColumnsDeps): ColumnDef[] => [ + { + id: "id", + accessorKey: "id", + meta: { title: "ID" }, + header: "ID", + size: 190, + enableSorting: false, + cell: ({ row }) => { + const endpointId = row.original.id; + if (!endpointId) return —; + return ( + onEndpointClick(endpointId)} + /> + ); + }, + }, + { + id: "path", + accessorKey: "path", + meta: { title: "Path" }, + header: "Path", + size: 200, + enableSorting: false, + cell: ({ row }) => ( + + {row.original.path} + + ), + }, + { + id: "target", + accessorKey: "target", + meta: { title: "Target" }, + header: "Target", + size: 240, + enableSorting: false, + cell: ({ row }) => ( + + {row.original.target} + + ), + }, + { + id: "methods", + meta: { title: "Methods", skeleton: "chips" }, + header: () => , + size: 150, + enableSorting: false, + cell: ({ row }) => , + }, + { + id: "auth", + accessorKey: "auth", + meta: { title: "Authentication", skeleton: "badge" }, + header: () => , + size: 140, + enableSorting: false, + cell: ({ row }) => ( + + ), + }, + { + id: "headers", + meta: { title: "Headers" }, + header: "Headers", + size: 180, + enableSorting: false, + cell: ({ row }) => , + }, + { + id: "actions", + meta: { className: "text-right", headerClassName: "text-right" }, + header: () => Actions, + size: 64, + enableSorting: false, + enableHiding: false, + cell: ({ row }) => ( +
+ +
+ ), + }, +]; diff --git a/ui/litellm-dashboard/src/components/PassThroughSettings/PassThroughSettings.test.tsx b/ui/litellm-dashboard/src/components/PassThroughSettings/PassThroughSettings.test.tsx new file mode 100644 index 00000000000..90269a916f7 --- /dev/null +++ b/ui/litellm-dashboard/src/components/PassThroughSettings/PassThroughSettings.test.tsx @@ -0,0 +1,120 @@ +import { act, render, screen, waitFor } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import { beforeEach, describe, expect, it, vi } from "vitest"; + +import { deletePassThroughEndpointsCall, getPassThroughEndpointsCall } from "../networking"; +import PassThroughSettings from "./PassThroughSettings"; +import type { PassThroughEndpointsTable } from "./PassThroughEndpointsTable"; + +vi.mock("../networking", () => ({ + getPassThroughEndpointsCall: vi.fn(), + deletePassThroughEndpointsCall: vi.fn(), +})); + +vi.mock("../add_pass_through", () => ({ + default: () =>
, +})); + +vi.mock("../pass_through_info", () => ({ + default: ({ endpointData }: { endpointData: { id?: string } }) => ( +
{endpointData.id}
+ ), +})); + +vi.mock("../molecules/notifications_manager", () => ({ + __esModule: true, + default: { success: vi.fn(), fromBackend: vi.fn() }, +})); + +vi.mock("./PassThroughEndpointsTable", () => ({ + PassThroughEndpointsTable: (props: React.ComponentProps) => ( +
+ {props.endpoints.map((endpoint) => ( +
+ + +
+ ))} +
+ ), +})); + +const mockGetEndpoints = vi.mocked(getPassThroughEndpointsCall); +const mockDeleteEndpoint = vi.mocked(deletePassThroughEndpointsCall); + +const defaultProps = { + accessToken: "token", + userRole: "Admin", + userID: "user-1", + premiumUser: false, +}; + +const endpoint = { id: "ep-1", path: "/v1/rerank", target: "https://example.com", headers: {} }; + +describe("PassThroughSettings", () => { + beforeEach(() => { + vi.clearAllMocks(); + mockGetEndpoints.mockResolvedValue({ endpoints: [endpoint] }); + }); + + it("should render nothing without an access token", () => { + const { container } = render(); + expect(container).toBeEmptyDOMElement(); + }); + + it("should hold the table in loading state until the fetch settles", async () => { + let resolveEndpoints: (value: { endpoints: (typeof endpoint)[] }) => void = () => {}; + mockGetEndpoints.mockReturnValue( + new Promise((resolve) => { + resolveEndpoints = resolve; + }), + ); + + render(); + expect(screen.getByTestId("endpoints-table")).toHaveAttribute("data-loading", "true"); + + await act(async () => { + resolveEndpoints({ endpoints: [endpoint] }); + }); + + await waitFor(() => { + expect(screen.getByTestId("endpoints-table")).toHaveAttribute("data-loading", "false"); + }); + }); + + it("should resolve loading without fetching when the user id is missing", async () => { + render(); + + await waitFor(() => { + expect(screen.getByTestId("endpoints-table")).toHaveAttribute("data-loading", "false"); + }); + expect(mockGetEndpoints).not.toHaveBeenCalled(); + }); + + it("should swap to the endpoint info view when an endpoint is opened", async () => { + const user = userEvent.setup(); + render(); + + await user.click(await screen.findByText("open-ep-1")); + expect(screen.getByTestId("endpoint-info")).toHaveTextContent("ep-1"); + }); + + it("should confirm before deleting an endpoint", async () => { + const user = userEvent.setup(); + mockDeleteEndpoint.mockResolvedValue(undefined); + render(); + + await user.click(await screen.findByText("delete-ep-1")); + expect(screen.getByText("Delete Pass-Through Endpoint")).toBeInTheDocument(); + expect(mockDeleteEndpoint).not.toHaveBeenCalled(); + + await user.click(screen.getByRole("button", { name: "Delete" })); + await waitFor(() => { + expect(mockDeleteEndpoint).toHaveBeenCalledWith("token", "ep-1"); + }); + }); +}); diff --git a/ui/litellm-dashboard/src/components/PassThroughSettings/PassThroughSettings.tsx b/ui/litellm-dashboard/src/components/PassThroughSettings/PassThroughSettings.tsx new file mode 100644 index 00000000000..975ba0b70e1 --- /dev/null +++ b/ui/litellm-dashboard/src/components/PassThroughSettings/PassThroughSettings.tsx @@ -0,0 +1,176 @@ +import React, { useState, useEffect } from "react"; +import { Button } from "@/components/ui/button"; +import { deletePassThroughEndpointsCall, getPassThroughEndpointsCall } from "../networking"; +import AddPassThroughEndpoint from "../add_pass_through"; +import PassThroughInfoView from "../pass_through_info"; +import NotificationsManager from "../molecules/notifications_manager"; +import { PassThroughEndpointsTable } from "./PassThroughEndpointsTable"; + +interface PassThroughSettingsProps { + accessToken: string | null; + userRole: string | null; + userID: string | null; + premiumUser?: boolean; +} + +export interface passThroughItem { + id?: string; + path: string; + target: string; + headers: object; + include_subpath?: boolean; + cost_per_request?: number; + timeout?: number; + auth?: boolean; + methods?: string[]; + guardrails?: Record; + default_query_params?: Record; +} + +const PassThroughSettings: React.FC = ({ accessToken, userRole, userID, premiumUser }) => { + const [generalSettings, setGeneralSettings] = useState([]); + const [isLoading, setIsLoading] = useState(true); + const [selectedEndpointId, setSelectedEndpointId] = useState(null); + const [isDeleteModalOpen, setIsDeleteModalOpen] = useState(false); + const [endpointToDelete, setEndpointToDelete] = useState(null); + + useEffect(() => { + const fetchEndpoints = async () => { + if (!accessToken || !userRole || !userID) { + setIsLoading(false); + return; + } + try { + const data = await getPassThroughEndpointsCall(accessToken); + setGeneralSettings(data["endpoints"]); + } finally { + setIsLoading(false); + } + }; + fetchEndpoints(); + }, [accessToken, userRole, userID]); + + const handleEndpointUpdated = () => { + if (accessToken) { + getPassThroughEndpointsCall(accessToken).then((data) => { + setGeneralSettings(data["endpoints"]); + }); + } + }; + + const handleDelete = (endpointId: string) => { + setEndpointToDelete(endpointId); + setIsDeleteModalOpen(true); + }; + + const confirmDelete = async () => { + if (endpointToDelete == null || !accessToken) { + return; + } + + try { + await deletePassThroughEndpointsCall(accessToken, endpointToDelete); + + const updatedSettings = generalSettings.filter((setting) => setting.id !== endpointToDelete); + setGeneralSettings(updatedSettings); + + NotificationsManager.success("Endpoint deleted successfully."); + } catch (error) { + console.error("Error deleting the endpoint:", error); + NotificationsManager.fromBackend("Error deleting the endpoint: " + error); + } + + setIsDeleteModalOpen(false); + setEndpointToDelete(null); + }; + + const cancelDelete = () => { + setIsDeleteModalOpen(false); + setEndpointToDelete(null); + }; + + if (!accessToken) { + return null; + } + + if (selectedEndpointId) { + const selectedEndpoint = generalSettings.find((endpoint) => endpoint.id === selectedEndpointId); + + if (!selectedEndpoint) { + return
Endpoint not found
; + } + + return ( + setSelectedEndpointId(null)} + accessToken={accessToken} + isAdmin={userRole === "Admin" || userRole === "admin"} + premiumUser={premiumUser} + onEndpointUpdated={handleEndpointUpdated} + /> + ); + } + + return ( +
+
+

Pass Through Endpoints

+

Configure and manage your pass-through endpoints

+
+ + + + + + {isDeleteModalOpen && ( +
+
+ + + + +
+
+
+
+

Delete Pass-Through Endpoint

+
+

+ Are you sure you want to delete this pass-through endpoint? This action cannot be undone. +

+
+
+
+
+
+ + +
+
+
+
+ )} +
+ ); +}; + +export default PassThroughSettings; diff --git a/ui/litellm-dashboard/src/components/Settings/LoggingAndAlerts/LoggingCallbacks/LoggingCallbacksTable.test.tsx b/ui/litellm-dashboard/src/components/Settings/LoggingAndAlerts/LoggingCallbacks/LoggingCallbacksTable.test.tsx index 0533b98b762..3ef7dee01b0 100644 --- a/ui/litellm-dashboard/src/components/Settings/LoggingAndAlerts/LoggingCallbacks/LoggingCallbacksTable.test.tsx +++ b/ui/litellm-dashboard/src/components/Settings/LoggingAndAlerts/LoggingCallbacks/LoggingCallbacksTable.test.tsx @@ -1,28 +1,43 @@ -import { render } from "@testing-library/react"; -import { describe, expect, it } from "vitest"; +import { render, screen } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import { beforeEach, describe, expect, it, vi } from "vitest"; + import { LoggingCallbacksTable } from "./LoggingCallbacksTable"; +const baseVars = { + SLACK_WEBHOOK_URL: null, + LANGFUSE_PUBLIC_KEY: null, + LANGFUSE_SECRET_KEY: null, + LANGFUSE_HOST: null, + OPENMETER_API_KEY: null, +}; + describe("LoggingCallbacksTable", () => { + beforeEach(() => { + vi.clearAllMocks(); + }); + it("should render", () => { - const { getByText } = render(); - expect(getByText("Active Logging Callbacks")).toBeInTheDocument(); + render(); + expect(screen.getByText("Active Logging Callbacks")).toBeInTheDocument(); + }); + + it("should show the empty state when there are no callbacks", () => { + render(); + expect(screen.getByText("No callbacks configured")).toBeInTheDocument(); + expect(screen.getByText("Add your first callback to start logging data to external services.")).toBeInTheDocument(); + }); + + it("should show skeleton rows while loading", () => { + render(); + expect(screen.getAllByTestId("skeleton-row").length).toBeGreaterThan(0); + expect(screen.queryByText("No callbacks configured")).not.toBeInTheDocument(); }); it('should map "otel" to "OpenTelemetry" on the table', () => { - const { getByText } = render( + render( { }} />, ); - expect(getByText("OpenTelemetry")).toBeInTheDocument(); + expect(screen.getByText("OpenTelemetry")).toBeInTheDocument(); }); it("should fallback to original callback name when not in availableCallbacks", () => { - const { getByText } = render( + render( , ); - expect(getByText("custom_callback_x")).toBeInTheDocument(); + expect(screen.getByText("custom_callback_x")).toBeInTheDocument(); + }); + + it("should call onAdd when the Add Callback button is clicked", async () => { + const user = userEvent.setup(); + const onAdd = vi.fn(); + render(); + await user.click(screen.getByRole("button", { name: /add callback/i })); + expect(onAdd).toHaveBeenCalled(); + }); + + it("should test, edit, and delete a callback through the actions menu", async () => { + const user = userEvent.setup(); + const onTest = vi.fn(); + const onEdit = vi.fn(); + const onDelete = vi.fn(); + const callback = { name: "langfuse", type: "success" as const, variables: baseVars }; + render( + , + ); + + await user.click(screen.getByTestId("callback-actions-langfuse-success")); + await user.click(await screen.findByTestId("callback-action-test")); + expect(onTest).toHaveBeenCalledWith(callback); + + await user.click(screen.getByTestId("callback-actions-langfuse-success")); + await user.click(await screen.findByTestId("callback-action-edit")); + expect(onEdit).toHaveBeenCalledWith(callback); + + await user.click(screen.getByTestId("callback-actions-langfuse-success")); + await user.click(await screen.findByTestId("callback-action-delete")); + expect(onDelete).toHaveBeenCalledWith(callback); }); // Regression: `/get_callbacks` returns the same `name` twice when a @@ -61,16 +102,9 @@ describe("LoggingCallbacksTable", () => { // → POST to spend-log on both 200 and 4xx/5xx). The UI used to ignore // the `type` field and render every row as "Success", masking the // failure registration. Reading `record.type` fixes the badge AND - // composing the rowKey with type avoids React's duplicate-key warning. + // composing the row id with type avoids React's duplicate-key warning. it("renders distinct Success and Failure badges for same-name dual registration", () => { - const baseVars = { - SLACK_WEBHOOK_URL: null, - LANGFUSE_PUBLIC_KEY: null, - LANGFUSE_SECRET_KEY: null, - LANGFUSE_HOST: null, - OPENMETER_API_KEY: null, - }; - const { getAllByText, getByText } = render( + render( { }} />, ); - // Both rows show the same display name, but distinct mode badges. - expect(getAllByText("Custom Callback API")).toHaveLength(2); - expect(getByText("Success")).toBeInTheDocument(); - expect(getByText("Failure")).toBeInTheDocument(); + expect(screen.getAllByText("Custom Callback API")).toHaveLength(2); + expect(screen.getByText("Success")).toBeInTheDocument(); + expect(screen.getByText("Failure")).toBeInTheDocument(); }); }); diff --git a/ui/litellm-dashboard/src/components/Settings/LoggingAndAlerts/LoggingCallbacks/LoggingCallbacksTable.tsx b/ui/litellm-dashboard/src/components/Settings/LoggingAndAlerts/LoggingCallbacks/LoggingCallbacksTable.tsx index 4f1889cbc99..d27b0cfd340 100644 --- a/ui/litellm-dashboard/src/components/Settings/LoggingAndAlerts/LoggingCallbacks/LoggingCallbacksTable.tsx +++ b/ui/litellm-dashboard/src/components/Settings/LoggingAndAlerts/LoggingCallbacks/LoggingCallbacksTable.tsx @@ -1,118 +1,75 @@ -import { Button } from "@tremor/react"; -import type { TableProps } from "antd"; -import { Table } from "antd"; -import Title from "antd/es/typography/Title"; -import React from "react"; -import { StatusBadge, type StatusTone } from "@/components/shared/table_cells"; -import TableIconActionButton from "../../../common_components/IconActionButton/TableIconActionButtons/TableIconActionButton"; +"use client"; + +import { Inbox, Plus } from "lucide-react"; +import React, { useMemo } from "react"; + +import { DataTable } from "@/components/shared/DataTable"; +import { Button } from "@/components/ui/button"; + +import { + AvailableCallbacks, + CallbackRow, + callbackRowMode, + getLoggingCallbacksTableColumns, +} from "./LoggingCallbacksTableColumns"; import { AlertingObject } from "./types"; type LoggingCallbacksProps = { callbacks: AlertingObject[]; - availableCallbacks?: Record< - string, - { - litellm_callback_name: string; - litellm_callback_params: string[]; - ui_callback_name: string; - } - >; + availableCallbacks?: AvailableCallbacks; + isLoading?: boolean; onTest?: (callback: AlertingObject) => void | Promise; onEdit?: (callback: AlertingObject) => void; onDelete?: (callback: AlertingObject) => void; onAdd?: () => void; }; -type CallbackRow = AlertingObject & { - id?: string; - mode?: "success" | "failure" | "info" | string; -}; - -const CALLBACK_MODES: { value: string; label: string }[] = [ - { value: "success", label: "Success" }, - { value: "failure", label: "Failure" }, - { value: "success_and_failure", label: "Success & Failure" }, -]; +function EmptyState() { + return ( +
+
+ +
+
No callbacks configured
+
+ Add your first callback to start logging data to external services. +
+
+ ); +} export const LoggingCallbacksTable: React.FC = ({ callbacks, availableCallbacks = {}, + isLoading = false, onTest = () => {}, onEdit = () => {}, onDelete = () => {}, onAdd = () => {}, }) => { - const columns: TableProps["columns"] = [ - { - title: Callback Name, - dataIndex: "name", - key: "name", - render: (_: string, record: CallbackRow) => { - const id = record.name; - const displayName = availableCallbacks[id]?.ui_callback_name || id; - return
{displayName}
; - }, - }, - { - title: Mode, - key: "mode", - render: (_: unknown, record: CallbackRow) => { - // Backend sends `type` (success | failure); legacy in-memory rows - // from add-callback flow set `mode`. Read both so newly-added rows - // and server-fetched rows both render correctly. - const mode = record.type || record.mode || "success"; - const label = CALLBACK_MODES.find((m) => m.value === mode)?.label || mode; - const tone: StatusTone = mode === "success" ? "success" : mode === "failure" ? "error" : "info"; - return ; - }, - width: 240, - }, - { - title: Actions, - key: "actions", - align: "right", - render: (_: unknown, record: CallbackRow) => ( -
- onTest(record)} /> - onEdit(record)} /> - onDelete(record)} /> -
- ), - width: 240, - }, - ]; + const columns = useMemo(() => { + const deps = { availableCallbacks, onTest, onEdit, onDelete }; + return getLoggingCallbacksTableColumns(deps); + }, [availableCallbacks, onTest, onEdit, onDelete]); + return ( - <> -
- -
- Active Logging Callbacks -
- {/* Empty state */} - {callbacks.length === 0 ? ( -
-
-

No callbacks configured

-

Add your first callback to start logging data to external services.

-
-
- ) : ( -
- `${record.name}-${record.type || record.mode || "success"}`} - pagination={false} - rowClassName={() => "hover:bg-gray-50"} - /> - - )} - + `${callback.name || index}-${callbackRowMode(callback)}`} + isLoading={isLoading} + loadingMessage="Loading callbacks…" + noDataMessage={} + size="compact" + /> + ); }; diff --git a/ui/litellm-dashboard/src/components/Settings/LoggingAndAlerts/LoggingCallbacks/LoggingCallbacksTableColumns.tsx b/ui/litellm-dashboard/src/components/Settings/LoggingAndAlerts/LoggingCallbacks/LoggingCallbacksTableColumns.tsx new file mode 100644 index 00000000000..2263fe03b3d --- /dev/null +++ b/ui/litellm-dashboard/src/components/Settings/LoggingAndAlerts/LoggingCallbacks/LoggingCallbacksTableColumns.tsx @@ -0,0 +1,134 @@ +"use client"; + +import { ColumnDef } from "@tanstack/react-table"; +import { MoreHorizontal, Pencil, Play, Trash2 } from "lucide-react"; + +import { StatusBadge, type StatusTone } from "@/components/shared/table_cells"; +import { buttonVariants } from "@/components/ui/button"; +import { + DropdownMenu, + DropdownMenuContent, + DropdownMenuItem, + DropdownMenuSeparator, + DropdownMenuTrigger, +} from "@/components/ui/dropdown-menu"; +import { cn } from "@/lib/cva.config"; + +import { AlertingObject } from "./types"; + +export type CallbackRow = AlertingObject & { + mode?: "success" | "failure" | "info" | string; +}; + +export interface AvailableCallbackMeta { + litellm_callback_name: string; + litellm_callback_params: string[]; + ui_callback_name: string; +} + +export type AvailableCallbacks = Record; + +export const callbackRowMode = (record: CallbackRow): string => record.type || record.mode || "success"; + +const CALLBACK_MODE_LABELS: Record = { + success: "Success", + failure: "Failure", + success_and_failure: "Success & Failure", +}; + +function callbackModeTone(mode: string): StatusTone { + if (mode === "success") return "success"; + if (mode === "failure") return "error"; + return "info"; +} + +interface CallbackRowActionsProps { + callback: CallbackRow; + onTest: (callback: AlertingObject) => void | Promise; + onEdit: (callback: AlertingObject) => void; + onDelete: (callback: AlertingObject) => void; +} + +function CallbackRowActions({ callback, onTest, onEdit, onDelete }: CallbackRowActionsProps) { + return ( + + + + + + void onTest(callback)}> + + Test + + onEdit(callback)}> + + Edit + + + onDelete(callback)}> + + Delete + + + + ); +} + +interface LoggingCallbacksTableColumnsDeps { + availableCallbacks: AvailableCallbacks; + onTest: (callback: AlertingObject) => void | Promise; + onEdit: (callback: AlertingObject) => void; + onDelete: (callback: AlertingObject) => void; +} + +export const getLoggingCallbacksTableColumns = ({ + availableCallbacks, + onTest, + onEdit, + onDelete, +}: LoggingCallbacksTableColumnsDeps): ColumnDef[] => [ + { + id: "name", + accessorKey: "name", + meta: { title: "Callback Name" }, + header: "Callback Name", + enableSorting: false, + cell: ({ row }) => { + const id = row.original.name; + const displayName = availableCallbacks[id]?.ui_callback_name || id; + return ( + + {displayName} + + ); + }, + }, + { + id: "mode", + meta: { title: "Mode", skeleton: "badge" }, + header: "Mode", + size: 240, + enableSorting: false, + cell: ({ row }) => { + const mode = callbackRowMode(row.original); + return ; + }, + }, + { + id: "actions", + meta: { className: "text-right", headerClassName: "text-right" }, + header: () => Actions, + size: 64, + enableSorting: false, + enableHiding: false, + cell: ({ row }) => ( +
+ +
+ ), + }, +]; diff --git a/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.tsx b/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.tsx index 8724c27b41a..5122d54db9b 100644 --- a/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.tsx +++ b/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.tsx @@ -21,6 +21,7 @@ import { getSemanticConfigError, } from "./build_complexity_router_config"; import { buildAutoRouterTestTargets, AutoRouterTestTarget } from "./build_auto_router_test_targets"; +import { getSemanticRouterError } from "./build_semantic_router_validation"; import AutoRouterConnectionTest from "./auto_router_connection_test"; import NotificationManager from "../molecules/notifications_manager"; @@ -164,23 +165,13 @@ const AddAutoRouterTab: React.FC = ({ form, handleOk, acc }; const submitSemanticRouter = (name: string) => { - if (!form.getFieldValue("auto_router_default_model")) { - NotificationManager.fromBackend("Please select a Default Model"); - return; - } - - if (!routerConfig || !routerConfig.routes || routerConfig.routes.length === 0) { - NotificationManager.fromBackend("Please configure at least one route for the auto router"); - return; - } - - const invalidRoutes = routerConfig.routes.filter( - (route: any) => !route.name || !route.description || route.utterances.length === 0, - ); - if (invalidRoutes.length > 0) { - NotificationManager.fromBackend( - "Please ensure all routes have a target model, description, and at least one utterance", - ); + const validationError = getSemanticRouterError({ + defaultModel: form.getFieldValue("auto_router_default_model"), + embeddingModel: form.getFieldValue("auto_router_embedding_model"), + routerConfig, + }); + if (validationError) { + NotificationManager.fromBackend(validationError); return; } @@ -358,18 +349,18 @@ const AddAutoRouterTab: React.FC = ({ form, handleOk, acc diff --git a/ui/litellm-dashboard/src/components/add_model/build_semantic_router_validation.test.ts b/ui/litellm-dashboard/src/components/add_model/build_semantic_router_validation.test.ts new file mode 100644 index 00000000000..a5556813cf4 --- /dev/null +++ b/ui/litellm-dashboard/src/components/add_model/build_semantic_router_validation.test.ts @@ -0,0 +1,67 @@ +import { getSemanticRouterError, SemanticRouterConfig } from "./build_semantic_router_validation"; + +const validRouterConfig: SemanticRouterConfig = { + routes: [{ name: "gpt-4o", description: "general chat", utterances: ["hello there"] }], +}; + +describe("getSemanticRouterError", () => { + it("requires an embedding model once the default model and routes are configured", () => { + expect( + getSemanticRouterError({ + defaultModel: "gpt-4o", + embeddingModel: undefined, + routerConfig: validRouterConfig, + }), + ).toBe("Please select an Embedding Model"); + }); + + it("treats an empty embedding model string as missing", () => { + expect( + getSemanticRouterError({ + defaultModel: "gpt-4o", + embeddingModel: "", + routerConfig: validRouterConfig, + }), + ).toBe("Please select an Embedding Model"); + }); + + it("passes when an embedding model is selected", () => { + expect( + getSemanticRouterError({ + defaultModel: "gpt-4o", + embeddingModel: "text-embedding-3-large", + routerConfig: validRouterConfig, + }), + ).toBeNull(); + }); + + it("flags a missing default model before checking the embedding model", () => { + expect( + getSemanticRouterError({ + defaultModel: undefined, + embeddingModel: undefined, + routerConfig: validRouterConfig, + }), + ).toBe("Please select a Default Model"); + }); + + it("flags missing routes before checking the embedding model", () => { + expect( + getSemanticRouterError({ + defaultModel: "gpt-4o", + embeddingModel: undefined, + routerConfig: { routes: [] }, + }), + ).toBe("Please configure at least one route for the auto router"); + }); + + it("validates route completeness after the embedding model is set", () => { + expect( + getSemanticRouterError({ + defaultModel: "gpt-4o", + embeddingModel: "text-embedding-3-large", + routerConfig: { routes: [{ name: "gpt-4o", description: "", utterances: [] }] }, + }), + ).toBe("Please ensure all routes have a target model, description, and at least one utterance"); + }); +}); diff --git a/ui/litellm-dashboard/src/components/add_model/build_semantic_router_validation.ts b/ui/litellm-dashboard/src/components/add_model/build_semantic_router_validation.ts new file mode 100644 index 00000000000..847ddee9ae1 --- /dev/null +++ b/ui/litellm-dashboard/src/components/add_model/build_semantic_router_validation.ts @@ -0,0 +1,29 @@ +export interface SemanticRouterRoute { + name?: string; + description?: string; + utterances?: unknown[]; +} + +export interface SemanticRouterConfig { + routes?: SemanticRouterRoute[]; +} + +export interface SemanticRouterValidationParams { + defaultModel: string | undefined; + embeddingModel: string | undefined; + routerConfig: SemanticRouterConfig | null | undefined; +} + +export const getSemanticRouterError = ({ + defaultModel, + embeddingModel, + routerConfig, +}: SemanticRouterValidationParams): string | null => { + if (!defaultModel) return "Please select a Default Model"; + if (!routerConfig?.routes || routerConfig.routes.length === 0) + return "Please configure at least one route for the auto router"; + if (!embeddingModel) return "Please select an Embedding Model"; + if (routerConfig.routes.some((route) => !route.name || !route.description || (route.utterances?.length ?? 0) === 0)) + return "Please ensure all routes have a target model, description, and at least one utterance"; + return null; +}; diff --git a/ui/litellm-dashboard/src/components/add_pass_through.tsx b/ui/litellm-dashboard/src/components/add_pass_through.tsx index af71d3f97fe..c0343e268a1 100644 --- a/ui/litellm-dashboard/src/components/add_pass_through.tsx +++ b/ui/litellm-dashboard/src/components/add_pass_through.tsx @@ -11,7 +11,7 @@ import NumericalInput from "./shared/numerical_input"; import { InfoCircleOutlined, ApiOutlined } from "@ant-design/icons"; import KeyValueInput from "./key_value_input"; import QueryParamInput from "./query_param_input"; -import { passThroughItem } from "./pass_through_settings"; +import { passThroughItem } from "./PassThroughSettings/PassThroughSettings"; import RoutePreview from "./route_preview"; import NotificationsManager from "./molecules/notifications_manager"; import PassThroughSecuritySection from "./common_components/PassThroughSecuritySection"; diff --git a/ui/litellm-dashboard/src/components/edit_auto_router/edit_auto_router_modal.tsx b/ui/litellm-dashboard/src/components/edit_auto_router/edit_auto_router_modal.tsx index ec54c9b7bad..f85cd16486a 100644 --- a/ui/litellm-dashboard/src/components/edit_auto_router/edit_auto_router_modal.tsx +++ b/ui/litellm-dashboard/src/components/edit_auto_router/edit_auto_router_modal.tsx @@ -367,15 +367,18 @@ const EditAutoRouterModal: React.FC = ({ {/* Embedding Model */} - + { setShowCustomEmbeddingModel(value === "custom"); }} options={[...modelOptions, { value: "custom", label: "Enter custom model name" }]} showSearch={true} - allowClear /> diff --git a/ui/litellm-dashboard/src/components/mcp_tools/types.tsx b/ui/litellm-dashboard/src/components/mcp_tools/types.tsx index 04766aad7b4..dd7ed2bfbe4 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/types.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/types.tsx @@ -81,6 +81,7 @@ export const getOAuthAuthorizationIdentity = (values: Record): client_id: credentials.client_id ?? null, client_secret: credentials.client_secret ?? null, scopes: credentials.scopes ?? null, + issuer: values.issuer ?? null, authorization_url: values.authorization_url ?? null, token_url: values.token_url ?? null, registration_url: values.registration_url ?? null, @@ -341,6 +342,7 @@ export interface MCPServer { transport?: string | null; auth_type?: string | null; oauth2_flow?: string | null; + issuer?: string | null; authorization_url?: string | null; token_url?: string | null; registration_url?: string | null; diff --git a/ui/litellm-dashboard/src/components/navbar.tsx b/ui/litellm-dashboard/src/components/navbar.tsx index 1a0b1c38520..40638e7b8ba 100644 --- a/ui/litellm-dashboard/src/components/navbar.tsx +++ b/ui/litellm-dashboard/src/components/navbar.tsx @@ -5,7 +5,7 @@ import { useWorker } from "@/hooks/useWorker"; import { getProxyBaseUrl } from "@/components/networking"; import { useTheme } from "@/contexts/ThemeContext"; import { clearTokenCookies } from "@/utils/cookieUtils"; -import { clearStoredReturnUrl } from "@/utils/returnUrlUtils"; +import { clearStoredReturnUrl, getLoginUrl } from "@/utils/returnUrlUtils"; import useProxySettings from "@/app/(dashboard)/hooks/proxySettings/useProxySettings"; import { DownOutlined, MenuFoldOutlined, MenuUnfoldOutlined } from "@ant-design/icons"; import { Tag } from "antd"; @@ -56,7 +56,7 @@ const Navbar: React.FC = ({ clearStoredReturnUrl(); localStorage.removeItem("litellm_selected_worker_id"); localStorage.removeItem("litellm_worker_url"); - window.location.href = `/ui/login?worker=${encodeURIComponent(workerId)}`; + window.location.href = `${getLoginUrl()}?worker=${encodeURIComponent(workerId)}`; }; return ( diff --git a/ui/litellm-dashboard/src/components/pass_through_settings.tsx b/ui/litellm-dashboard/src/components/pass_through_settings.tsx deleted file mode 100644 index 6c72b8a62ab..00000000000 --- a/ui/litellm-dashboard/src/components/pass_through_settings.tsx +++ /dev/null @@ -1,304 +0,0 @@ -import React, { useState, useEffect } from "react"; -import { Text, Button, Icon, Title } from "@tremor/react"; -import { deletePassThroughEndpointsCall, getPassThroughEndpointsCall } from "./networking"; -import { Badge, Tooltip } from "antd"; -import { PencilAltIcon, TrashIcon, InformationCircleIcon } from "@heroicons/react/outline"; -import AddPassThroughEndpoint from "./add_pass_through"; -import PassThroughInfoView from "./pass_through_info"; -import { DataTable } from "./view_logs/table"; -import { ColumnDef } from "@tanstack/react-table"; -import { IdCell, StatusBadge } from "@/components/shared/table_cells"; -import { Eye, EyeOff } from "lucide-react"; -import NotificationsManager from "./molecules/notifications_manager"; - -interface GeneralSettingsPageProps { - accessToken: string | null; - userRole: string | null; - userID: string | null; - modelData: any; - premiumUser?: boolean; -} - -interface routingStrategyArgs { - ttl?: number; - lowest_latency_buffer?: number; -} - -interface nestedFieldItem { - field_name: string; - field_type: string; - field_value: any; - field_description: string; - stored_in_db: boolean | null; -} - -export interface passThroughItem { - id?: string; - path: string; - target: string; - headers: object; - include_subpath?: boolean; - cost_per_request?: number; - timeout?: number; - auth?: boolean; - methods?: string[]; - guardrails?: Record; - default_query_params?: Record; -} - -// Password field component for headers -const PasswordField: React.FC<{ value: object }> = ({ value }) => { - const [showPassword, setShowPassword] = useState(false); - const headerString = JSON.stringify(value); - - return ( -
- {showPassword ? headerString : "••••••••"} - -
- ); -}; - -const PassThroughSettings: React.FC = ({ - accessToken, - userRole, - userID, - modelData, - premiumUser, -}) => { - const [generalSettings, setGeneralSettings] = useState([]); - const [selectedEndpointId, setSelectedEndpointId] = useState(null); - const [isDeleteModalOpen, setIsDeleteModalOpen] = useState(false); - const [endpointToDelete, setEndpointToDelete] = useState(null); - - useEffect(() => { - if (!accessToken || !userRole || !userID) { - return; - } - getPassThroughEndpointsCall(accessToken).then((data) => { - let general_settings = data["endpoints"]; - setGeneralSettings(general_settings); - }); - }, [accessToken, userRole, userID]); - - const handleEndpointUpdated = () => { - // Refresh the endpoints list when an endpoint is updated - if (accessToken) { - getPassThroughEndpointsCall(accessToken).then((data) => { - let general_settings = data["endpoints"]; - setGeneralSettings(general_settings); - }); - } - }; - - const handleDelete = async (endpointId: string) => { - // Set the endpoint to delete and open the confirmation modal - setEndpointToDelete(endpointId); - setIsDeleteModalOpen(true); - }; - - const confirmDelete = async () => { - if (endpointToDelete == null || !accessToken) { - return; - } - - try { - await deletePassThroughEndpointsCall(accessToken, endpointToDelete); - - const updatedSettings = generalSettings.filter((setting) => setting.id !== endpointToDelete); - setGeneralSettings(updatedSettings); - - NotificationsManager.success("Endpoint deleted successfully."); - } catch (error) { - console.error("Error deleting the endpoint:", error); - NotificationsManager.fromBackend("Error deleting the endpoint: " + error); - } - - // Close the confirmation modal and reset the endpointToDelete - setIsDeleteModalOpen(false); - setEndpointToDelete(null); - }; - - const cancelDelete = () => { - // Close the confirmation modal and reset the endpointToDelete - setIsDeleteModalOpen(false); - setEndpointToDelete(null); - }; - - const handleResetField = (endpointId: string, idx: number) => { - // Use handleDelete instead of direct deletion - handleDelete(endpointId); - }; - - // Define columns for the DataTable - const columns: ColumnDef[] = [ - { - header: "ID", - accessorKey: "id", - cell: (info: any) => , - }, - { - header: "Path", - accessorKey: "path", - }, - { - header: "Target", - accessorKey: "target", - cell: (info: any) => {info.getValue()}, - }, - { - header: () => ( -
- Methods - - - -
- ), - accessorKey: "methods", - cell: (info: any) => { - const methods = info.getValue(); - if (!methods || methods.length === 0) { - return ALL; - } - return ( -
- {methods.map((method: string) => ( - - {method} - - ))} -
- ); - }, - }, - { - header: () => ( -
- Authentication - - - -
- ), - accessorKey: "auth", - cell: (info: any) => ( - - ), - }, - { - header: "Headers", - accessorKey: "headers", - cell: (info: any) => , - }, - { - header: "Actions", - id: "actions", - cell: ({ row }) => ( -
- row.original.id && setSelectedEndpointId(row.original.id)} - title="Edit" - /> - handleResetField(row.original.id!, row.index)} - title="Delete" - /> -
- ), - }, - ]; - - if (!accessToken) { - return null; - } - - // If a specific endpoint is selected, show the info view - if (selectedEndpointId) { - // Find the endpoint by ID to get the endpoint data for the info view - const selectedEndpoint = generalSettings.find((endpoint) => endpoint.id === selectedEndpointId); - - if (!selectedEndpoint) { - return
Endpoint not found
; - } - - return ( - setSelectedEndpointId(null)} - accessToken={accessToken} - isAdmin={userRole === "Admin" || userRole === "admin"} - premiumUser={premiumUser} - onEndpointUpdated={handleEndpointUpdated} - /> - ); - } - - return ( -
-
- Pass Through Endpoints - Configure and manage your pass-through endpoints -
- - - - - - {isDeleteModalOpen && ( -
-
- - - {/* Modal Panel */} - - - {/* Confirmation Modal Content */} -
-
-
-
-

Delete Pass-Through Endpoint

-
-

- Are you sure you want to delete this pass-through endpoint? This action cannot be undone. -

-
-
-
-
-
- - -
-
-
-
- )} -
- ); -}; - -export default PassThroughSettings; diff --git a/ui/litellm-dashboard/src/components/settings.test.tsx b/ui/litellm-dashboard/src/components/settings.test.tsx index 7776d0c7082..ebc1a5a7a6c 100644 --- a/ui/litellm-dashboard/src/components/settings.test.tsx +++ b/ui/litellm-dashboard/src/components/settings.test.tsx @@ -1,4 +1,5 @@ -import { act, fireEvent, render, waitFor } from "@testing-library/react"; +import { act, render, screen, waitFor } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; import { beforeAll, beforeEach, describe, expect, it, vi } from "vitest"; import { alertingSettingsCall, getCallbackConfigsCall, getCallbacksCall } from "./networking"; import Settings from "./settings"; @@ -160,7 +161,8 @@ describe("Settings", () => { mockGetCallbackConfigsCall.mockResolvedValue([mockCallbackConfig]); - const { getByText, container } = render(); + const user = userEvent.setup(); + const { getByText } = render(); await waitFor(() => { expect(getByText("Active Logging Callbacks")).toBeInTheDocument(); @@ -170,18 +172,8 @@ describe("Settings", () => { expect(getByText("Langfuse")).toBeInTheDocument(); }); - const actionsCell = container.querySelector('[class*="flex justify-end gap-2"]'); - expect(actionsCell).toBeTruthy(); - - const icons = actionsCell?.querySelectorAll("svg"); - expect(icons?.length).toBeGreaterThanOrEqual(2); - - const editIconParent = icons?.[1]?.closest('[class*="cursor-pointer"]'); - expect(editIconParent).toBeTruthy(); - - act(() => { - fireEvent.click(editIconParent!); - }); + await user.click(screen.getByTestId("callback-actions-langfuse-success")); + await user.click(await screen.findByTestId("callback-action-edit")); await waitFor(() => { expect(getByText("Edit Callback Settings")).toBeInTheDocument(); @@ -194,6 +186,42 @@ describe("Settings", () => { }); }); + it("should hold the callbacks table in loading state until the fetch settles", async () => { + let resolveCallbacks: (value: { + callbacks: never[]; + available_callbacks: never[]; + alerts: never[]; + }) => void = () => {}; + mockGetCallbacksCall.mockReturnValue( + new Promise((resolve) => { + resolveCallbacks = resolve; + }), + ); + + render(); + + expect(screen.getAllByTestId("skeleton-row").length).toBeGreaterThan(0); + + await act(async () => { + resolveCallbacks({ callbacks: [], available_callbacks: [], alerts: [] }); + }); + + await waitFor(() => { + expect(screen.queryByTestId("skeleton-row")).not.toBeInTheDocument(); + }); + expect(screen.getByText("No callbacks configured")).toBeInTheDocument(); + }); + + it("should resolve loading without fetching when the user id is missing", async () => { + render(); + + await waitFor(() => { + expect(screen.queryByTestId("skeleton-row")).not.toBeInTheDocument(); + }); + expect(mockGetCallbacksCall).not.toHaveBeenCalled(); + expect(screen.getByText("No callbacks configured")).toBeInTheDocument(); + }); + it("should display CloudZero Cost Tracking tab", async () => { const { getByText } = render(); diff --git a/ui/litellm-dashboard/src/components/settings.tsx b/ui/litellm-dashboard/src/components/settings.tsx index 46811e0106b..b3f33133a80 100644 --- a/ui/litellm-dashboard/src/components/settings.tsx +++ b/ui/litellm-dashboard/src/components/settings.tsx @@ -217,6 +217,7 @@ const buildCallbackPayload = (formValues: Record, callbackName: str const Settings: React.FC = ({ accessToken, userRole, userID, premiumUser }) => { const [callbacks, setCallbacks] = useState([]); + const [isLoadingCallbacks, setIsLoadingCallbacks] = useState(true); const [alerts, setAlerts] = useState([]); const [isModalVisible, setIsModalVisible] = useState(false); const [addForm] = Form.useForm(); @@ -293,29 +294,35 @@ const Settings: React.FC = ({ accessToken, userRole, userID, }; useEffect(() => { - if (!accessToken || !userRole || !userID) { - return; - } - getCallbacksCall(accessToken, userID, userRole).then((data) => { - setCallbacks(data.callbacks); - setAllCallbacks(data.available_callbacks); - // setCallbacks(callbacks_data); - - let alerts_data = data.alerts; - if (alerts_data) { - if (alerts_data.length > 0) { - let _alert_info = alerts_data[0]; - let catch_all_webhook = _alert_info.variables.SLACK_WEBHOOK_URL; - - let active_alerts = _alert_info.active_alerts; - setActiveAlerts(active_alerts); - setCatchAllWebhookURL(catch_all_webhook); - setAlertToWebhooks(_alert_info.alerts_to_webhook); - } + const fetchCallbacks = async () => { + if (!accessToken || !userRole || !userID) { + setIsLoadingCallbacks(false); + return; } + try { + const data = await getCallbacksCall(accessToken, userID, userRole); + setCallbacks(data.callbacks); + setAllCallbacks(data.available_callbacks); - setAlerts(alerts_data); - }); + let alerts_data = data.alerts; + if (alerts_data) { + if (alerts_data.length > 0) { + let _alert_info = alerts_data[0]; + let catch_all_webhook = _alert_info.variables.SLACK_WEBHOOK_URL; + + let active_alerts = _alert_info.active_alerts; + setActiveAlerts(active_alerts); + setCatchAllWebhookURL(catch_all_webhook); + setAlertToWebhooks(_alert_info.alerts_to_webhook); + } + } + + setAlerts(alerts_data); + } finally { + setIsLoadingCallbacks(false); + } + }; + fetchCallbacks(); }, [accessToken, userRole, userID]); const isAlertOn = (alertName: string) => { @@ -581,6 +588,7 @@ const Settings: React.FC = ({ accessToken, userRole, userID, setShowAddCallbacksModal(true)} onEdit={(cb) => { setSelectedEditCallback(cb); diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 7a6088cbc59..e47ef4a9c59 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -22543,6 +22543,11 @@ export interface components { * @description connect to a postgres db - needed for generating temporary keys + tracking spend / key */ database_url?: string | null; + /** + * Disable Auto Add Proxy Admin To Teams + * @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. + */ + disable_auto_add_proxy_admin_to_teams?: boolean | null; /** * Disable Budget Reservation * @description If True, disables the optimistic per-request budget reservation introduced in v1.84.0. WARNING: This weakens hard budget enforcement. Without the reservation, a burst of concurrent requests from a single key can each pass the read-time spend check before any of them is charged, allowing a configured budget to be exceeded under high concurrency. Budgets are still evaluated on every request at read time, so an already-exhausted budget is still rejected. Enable only if your deployment is experiencing phantom BudgetExceededError responses caused by leaked reservations (see GitHub issue #27639). A proxy-level WARNING is logged on every request while this flag is active as a reminder that hard enforcement is relaxed. @@ -32308,6 +32313,8 @@ export interface components { }[] | null; /** Cooldown Time */ cooldown_time?: number | null; + /** Enable Tag Filtering */ + enable_tag_filtering?: boolean | null; /** Fallbacks */ fallbacks?: { [key: string]: unknown; diff --git a/ui/litellm-dashboard/src/utils/returnUrlUtils.test.ts b/ui/litellm-dashboard/src/utils/returnUrlUtils.test.ts index 56de299049c..0f3a8b4bf69 100644 --- a/ui/litellm-dashboard/src/utils/returnUrlUtils.test.ts +++ b/ui/litellm-dashboard/src/utils/returnUrlUtils.test.ts @@ -3,6 +3,7 @@ import { clearStoredReturnUrl, consumeReturnUrl, getCurrentUrl, + getLoginUrl, getReturnUrl, getReturnUrlFromParams, getStoredReturnUrl, @@ -98,6 +99,29 @@ describe("returnUrlUtils", () => { }); }); + describe("getLoginUrl", () => { + it("should build a relative login URL with a trailing slash", () => { + expect(getLoginUrl()).toBe("/ui/login/"); + }); + + it("should prepend the given base URL and keep the trailing slash", () => { + expect(getLoginUrl("http://proxy.example")).toBe("http://proxy.example/ui/login/"); + }); + + it("should keep the trailing slash before the query when composed with buildLoginUrlWithReturn", () => { + Object.defineProperty(window, "location", { + value: { + ...window.location, + href: "http://localhost:3000/ui?page=api-keys", + }, + writable: true, + }); + + const loginUrl = buildLoginUrlWithReturn(getLoginUrl()); + expect(loginUrl).toBe("/ui/login/?redirect_to=http%3A%2F%2Flocalhost%3A3000%2Fui%3Fpage%3Dapi-keys"); + }); + }); + describe("buildLoginUrlWithReturn", () => { it("should build login URL with return URL parameter", () => { Object.defineProperty(window, "location", { diff --git a/ui/litellm-dashboard/src/utils/returnUrlUtils.ts b/ui/litellm-dashboard/src/utils/returnUrlUtils.ts index 76562a0122a..b3cc5345bb1 100644 --- a/ui/litellm-dashboard/src/utils/returnUrlUtils.ts +++ b/ui/litellm-dashboard/src/utils/returnUrlUtils.ts @@ -13,6 +13,10 @@ const RETURN_URL_COOKIE_NAME = "litellm_return_url"; const RETURN_URL_PARAM = "redirect_to"; +export function getLoginUrl(baseUrl: string = ""): string { + return `${baseUrl}/ui/login/`; +} + /** * Gets the current URL with all query parameters. * Returns null if running on server-side. diff --git a/ui/litellm-dashboard/tests/CreateKeyPage.expiredToken.test.tsx b/ui/litellm-dashboard/tests/CreateKeyPage.expiredToken.test.tsx index 62346eff057..211a754dad4 100644 --- a/ui/litellm-dashboard/tests/CreateKeyPage.expiredToken.test.tsx +++ b/ui/litellm-dashboard/tests/CreateKeyPage.expiredToken.test.tsx @@ -232,7 +232,7 @@ describe("CreateKeyPage auth behavior", () => { // Assert: we eventually redirect to SSO login with return URL (single replace, not assign/href) await waitFor(() => { expect(window.location.replace).toHaveBeenCalledWith( - expect.stringContaining("https://example.com/ui/login?redirect_to="), + expect.stringContaining("https://example.com/ui/login/?redirect_to="), ); }); diff --git a/uv.lock b/uv.lock index 22e3a3feece..8dfc4bd5fcc 100644 --- a/uv.lock +++ b/uv.lock @@ -10,7 +10,7 @@ resolution-markers = [ ] [options] -exclude-newer = "2026-07-13T03:38:04.421387Z" +exclude-newer = "2026-07-13T21:03:01.672393Z" exclude-newer-span = "P3D" [manifest] @@ -222,9 +222,9 @@ name = "aiologic" version = "0.17.0" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "sniffio" }, + { name = "sniffio", marker = "python_full_version < '3.13'" }, { name = "typing-extensions", marker = "python_full_version < '3.13'" }, - { name = "wrapt" }, + { name = "wrapt", marker = "python_full_version < '3.13'" }, ] sdist = { url = "https://files.pythonhosted.org/packages/53/a7/809482759f40079f4c4328c7318bf569ae25d457f5017aad30a1b9aafedc/aiologic-0.17.0.tar.gz", hash = "sha256:65aa058e858c94cd208badb188e7f00b54dcabb3ba85b34f794db98074d108b9", size = 251625, upload-time = "2026-06-14T12:24:35.367Z" } wheels = [ @@ -516,14 +516,14 @@ name = "aurelio-sdk" version = "0.0.19" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "aiofiles" }, - { name = "aiohttp" }, - { name = "colorlog" }, - { name = "pydantic" }, - { name = "python-dotenv" }, - { name = "requests" }, - { name = "requests-toolbelt" }, - { name = "tornado" }, + { name = "aiofiles", marker = "python_full_version < '3.14'" }, + { name = "aiohttp", marker = "python_full_version < '3.14'" }, + { name = "colorlog", marker = "python_full_version < '3.14'" }, + { name = "pydantic", marker = "python_full_version < '3.14'" }, + { name = "python-dotenv", marker = "python_full_version < '3.14'" }, + { name = "requests", marker = "python_full_version < '3.14'" }, + { name = "requests-toolbelt", marker = "python_full_version < '3.14'" }, + { name = "tornado", marker = "python_full_version < '3.14'" }, ] sdist = { url = "https://files.pythonhosted.org/packages/27/0e/c2e369ad173fb3d76448e46d10beb3dcc53388318933ddf8169a3f21a810/aurelio_sdk-0.0.19.tar.gz", hash = "sha256:14107e7440ff2efd0b4a08c52fb595e7680bd4bc973a0ddfb3b64157c6666b91", size = 15258, upload-time = "2025-03-24T14:37:32.203Z" } wheels = [ @@ -1047,7 +1047,7 @@ name = "coloredlogs" version = "15.0.1" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "humanfriendly" }, + { name = "humanfriendly", marker = "python_full_version < '3.14'" }, ] sdist = { url = "https://files.pythonhosted.org/packages/cc/c7/eed8f27100517e8c0e6b923d5f0845d0cb99763da6fdee00478f91db7325/coloredlogs-15.0.1.tar.gz", hash = "sha256:7c991aa71a4577af2f82600d8f8f3a89f936baeaf9b50a9c197da014e5bf16b0", size = 278520, upload-time = "2021-06-11T10:22:45.202Z" } wheels = [ @@ -1059,7 +1059,7 @@ name = "colorlog" version = "6.10.1" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "colorama", marker = "sys_platform == 'win32'" }, + { name = "colorama", marker = "python_full_version < '3.14' and sys_platform == 'win32'" }, ] sdist = { url = "https://files.pythonhosted.org/packages/a2/61/f083b5ac52e505dfc1c624eafbf8c7589a0d7f32daa398d2e7590efa5fda/colorlog-6.10.1.tar.gz", hash = "sha256:eb4ae5cb65fe7fec7773c2306061a8e63e02efc2c72eba9d27b0fa23c94f1321", size = 17162, upload-time = "2025-10-16T16:14:11.978Z" } wheels = [ @@ -1420,7 +1420,7 @@ name = "culsans" version = "0.11.0" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "aiologic" }, + { name = "aiologic", marker = "python_full_version < '3.13'" }, { name = "typing-extensions", marker = "python_full_version < '3.13'" }, ] sdist = { url = "https://files.pythonhosted.org/packages/d9/e3/49afa1bc180e0d28008ec6bcdf82a4072d1c7a41032b5b759b60814ca4b0/culsans-0.11.0.tar.gz", hash = "sha256:0b43d0d05dce6106293d114c86e3fb4bfc63088cfe8ff08ed3fe36891447fe33", size = 107546, upload-time = "2025-12-31T23:15:38.196Z" } @@ -2966,7 +2966,7 @@ name = "humanfriendly" version = "10.0" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "pyreadline3", marker = "sys_platform == 'win32'" }, + { name = "pyreadline3", marker = "python_full_version < '3.14' and sys_platform == 'win32'" }, ] sdist = { url = "https://files.pythonhosted.org/packages/cc/3f/2c29224acb2e2df4d2046e4c73ee2662023c58ff5b113c4c1adac0886c43/humanfriendly-10.0.tar.gz", hash = "sha256:6b0b831ce8f15f7300721aa49829fc4e83921a9a301cc7f606be6686a2288ddc", size = 360702, upload-time = "2021-09-17T21:40:43.31Z" } wheels = [ @@ -3293,10 +3293,10 @@ name = "jsonschema-path" version = "0.3.4" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "pathable" }, - { name = "pyyaml" }, - { name = "referencing" }, - { name = "requests" }, + { name = "pathable", marker = "python_full_version < '3.14'" }, + { name = "pyyaml", marker = "python_full_version < '3.14'" }, + { name = "referencing", marker = "python_full_version < '3.14'" }, + { name = "requests", marker = "python_full_version < '3.14'" }, ] sdist = { url = "https://files.pythonhosted.org/packages/6e/45/41ebc679c2a4fced6a722f624c18d658dee42612b83ea24c1caf7c0eb3a8/jsonschema_path-0.3.4.tar.gz", hash = "sha256:8365356039f16cc65fddffafda5f58766e34bebab7d6d105616ab52bc4297001", size = 11159, upload-time = "2025-01-24T14:33:16.547Z" } wheels = [ @@ -4122,12 +4122,12 @@ proxy-dev = [ [[package]] name = "litellm-enterprise" -version = "0.1.50" +version = "0.1.51" source = { editable = "enterprise" } [[package]] name = "litellm-proxy-extras" -version = "0.4.77" +version = "0.4.78" source = { editable = "litellm-proxy-extras" } [[package]] @@ -4481,7 +4481,7 @@ version = "0.4.1" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "numpy", version = "1.26.4", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.12'" }, - { name = "numpy", version = "2.4.4", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.12'" }, + { name = "numpy", version = "2.4.4", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.12' and python_full_version < '3.14'" }, ] sdist = { url = "https://files.pythonhosted.org/packages/fd/15/76f86faa0902836cc133939732f7611ace68cf54148487a99c539c272dc8/ml_dtypes-0.4.1.tar.gz", hash = "sha256:fad5f2de464fd09127e49b7fd1252b9006fb43d2edc1ff112d390c324af5ca7a", size = 692594, upload-time = "2024-09-13T19:07:11.624Z" } wheels = [ @@ -4980,14 +4980,14 @@ name = "openapi-core" version = "0.22.0" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "isodate" }, - { name = "jsonschema" }, - { name = "jsonschema-path" }, - { name = "more-itertools" }, - { name = "openapi-schema-validator" }, - { name = "openapi-spec-validator" }, - { name = "typing-extensions" }, - { name = "werkzeug" }, + { name = "isodate", marker = "python_full_version < '3.14'" }, + { name = "jsonschema", marker = "python_full_version < '3.14'" }, + { name = "jsonschema-path", marker = "python_full_version < '3.14'" }, + { name = "more-itertools", marker = "python_full_version < '3.14'" }, + { name = "openapi-schema-validator", marker = "python_full_version < '3.14'" }, + { name = "openapi-spec-validator", marker = "python_full_version < '3.14'" }, + { name = "typing-extensions", marker = "python_full_version < '3.14'" }, + { name = "werkzeug", marker = "python_full_version < '3.14'" }, ] sdist = { url = "https://files.pythonhosted.org/packages/fd/65/ee75f25b9459a02df6f713f8ffde5dacb57b8b4e45145cde4cab28b5abba/openapi_core-0.22.0.tar.gz", hash = "sha256:b30490dfa74e3aac2276105525590135212352f5dd7e5acf8f62f6a89ed6f2d0", size = 109242, upload-time = "2025-12-22T19:19:49.608Z" } wheels = [ @@ -4999,9 +4999,9 @@ name = "openapi-schema-validator" version = "0.6.3" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "jsonschema" }, - { name = "jsonschema-specifications" }, - { name = "rfc3339-validator" }, + { name = "jsonschema", marker = "python_full_version < '3.14'" }, + { name = "jsonschema-specifications", marker = "python_full_version < '3.14'" }, + { name = "rfc3339-validator", marker = "python_full_version < '3.14'" }, ] sdist = { url = "https://files.pythonhosted.org/packages/8b/f3/5507ad3325169347cd8ced61c232ff3df70e2b250c49f0fe140edb4973c6/openapi_schema_validator-0.6.3.tar.gz", hash = "sha256:f37bace4fc2a5d96692f4f8b31dc0f8d7400fd04f3a937798eaf880d425de6ee", size = 11550, upload-time = "2025-01-10T18:08:22.268Z" } wheels = [ @@ -5013,10 +5013,10 @@ name = "openapi-spec-validator" version = "0.7.2" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "jsonschema" }, - { name = "jsonschema-path" }, - { name = "lazy-object-proxy" }, - { name = "openapi-schema-validator" }, + { name = "jsonschema", marker = "python_full_version < '3.14'" }, + { name = "jsonschema-path", marker = "python_full_version < '3.14'" }, + { name = "lazy-object-proxy", marker = "python_full_version < '3.14'" }, + { name = "openapi-schema-validator", marker = "python_full_version < '3.14'" }, ] sdist = { url = "https://files.pythonhosted.org/packages/82/af/fe2d7618d6eae6fb3a82766a44ed87cd8d6d82b4564ed1c7cfb0f6378e91/openapi_spec_validator-0.7.2.tar.gz", hash = "sha256:cc029309b5c5dbc7859df0372d55e9d1ff43e96d678b9ba087f7c56fc586f734", size = 36855, upload-time = "2025-06-07T14:48:56.299Z" } wheels = [ @@ -7163,16 +7163,16 @@ name = "redisvl" version = "0.4.1" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "coloredlogs" }, - { name = "ml-dtypes" }, + { name = "coloredlogs", marker = "python_full_version < '3.14'" }, + { name = "ml-dtypes", marker = "python_full_version < '3.14'" }, { name = "numpy", version = "1.26.4", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.12'" }, - { name = "numpy", version = "2.4.4", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.12'" }, - { name = "pydantic" }, - { name = "python-ulid" }, - { name = "pyyaml" }, - { name = "redis" }, - { name = "tabulate" }, - { name = "tenacity" }, + { name = "numpy", version = "2.4.4", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.12' and python_full_version < '3.14'" }, + { name = "pydantic", marker = "python_full_version < '3.14'" }, + { name = "python-ulid", marker = "python_full_version < '3.14'" }, + { name = "pyyaml", marker = "python_full_version < '3.14'" }, + { name = "redis", marker = "python_full_version < '3.14'" }, + { name = "tabulate", marker = "python_full_version < '3.14'" }, + { name = "tenacity", marker = "python_full_version < '3.14'" }, ] sdist = { url = "https://files.pythonhosted.org/packages/21/33/ab14865a0b2a31b1d003c29e7e8ea3a7a2f2c8ecb24e58e58d606e1f031b/redisvl-0.4.1.tar.gz", hash = "sha256:fd6a36426ba94792c0efca20915c31232d4ee3cc58eb23794a62c142696401e6", size = 77688, upload-time = "2025-02-21T22:51:41.389Z" } wheels = [ @@ -7406,7 +7406,7 @@ name = "rfc3339-validator" version = "0.1.4" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "six" }, + { name = "six", marker = "python_full_version < '3.14'" }, ] sdist = { url = "https://files.pythonhosted.org/packages/28/ea/a9387748e2d111c3c2b275ba970b735e04e15cdb1eb30693b6b5708c4dbd/rfc3339_validator-0.1.4.tar.gz", hash = "sha256:138a2abdf93304ad60530167e51d2dfb9549521a836871b88d7f4695d0022f6b", size = 5513, upload-time = "2021-05-12T16:37:54.178Z" } wheels = [ @@ -7855,20 +7855,20 @@ name = "semantic-router" version = "0.1.15" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "aiohttp" }, - { name = "aurelio-sdk" }, - { name = "colorama" }, - { name = "colorlog" }, - { name = "litellm" }, + { name = "aiohttp", marker = "python_full_version < '3.14'" }, + { name = "aurelio-sdk", marker = "python_full_version < '3.14'" }, + { name = "colorama", marker = "python_full_version < '3.14'" }, + { name = "colorlog", marker = "python_full_version < '3.14'" }, + { name = "litellm", marker = "python_full_version < '3.14'" }, { name = "numpy", version = "1.26.4", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.12'" }, - { name = "numpy", version = "2.4.4", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.12'" }, - { name = "openai" }, - { name = "pydantic" }, - { name = "pyyaml" }, - { name = "regex" }, - { name = "tiktoken" }, - { name = "tornado" }, - { name = "urllib3" }, + { name = "numpy", version = "2.4.4", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.12' and python_full_version < '3.14'" }, + { name = "openai", marker = "python_full_version < '3.14'" }, + { name = "pydantic", marker = "python_full_version < '3.14'" }, + { name = "pyyaml", marker = "python_full_version < '3.14'" }, + { name = "regex", marker = "python_full_version < '3.14'" }, + { name = "tiktoken", marker = "python_full_version < '3.14'" }, + { name = "tornado", marker = "python_full_version < '3.14'" }, + { name = "urllib3", marker = "python_full_version < '3.14'" }, ] sdist = { url = "https://files.pythonhosted.org/packages/dc/a9/1a689e916e8b280f1fd8fb335cc059be626a22fe4533baa045d32fcd6de5/semantic_router-0.1.15.tar.gz", hash = "sha256:328256ddc3c2b713101ec69561d6585aecbf1198ea3461e1486289d8c3a35288", size = 95605, upload-time = "2026-05-23T12:58:15.444Z" } wheels = [ @@ -8958,16 +8958,16 @@ wheels = [ [[package]] name = "uvicorn" -version = "0.33.0" +version = "0.51.0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "click" }, { name = "h11" }, { name = "typing-extensions", marker = "python_full_version < '3.11'" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/cb/81/a083ae41716b00df56d45d4b5f6ca8e90fc233a62e6c04ab3ad3c476b6c4/uvicorn-0.33.0.tar.gz", hash = "sha256:3577119f82b7091cf4d3d4177bfda0bae4723ed92ab1439e8d779de880c9cc59", size = 76590, upload-time = "2024-12-14T11:14:46.526Z" } +sdist = { url = "https://files.pythonhosted.org/packages/a2/65/b7c6c443ccc58678c91e1e973bbe2a878591538655d6e1d47f24ba1c51f3/uvicorn-0.51.0.tar.gz", hash = "sha256:f6f4b69b657c312f516dd2d268ab9ae6f254b11e4bac504f37b2ab58b24dd0b0", size = 94412, upload-time = "2026-07-08T10:59:05.962Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/98/79/2e2620337ef1e4ef7a058b351603b765f59ac28e6e3ac7c5e7cdee9ea1ab/uvicorn-0.33.0-py3-none-any.whl", hash = "sha256:2c30de4aeea83661a520abab179b24084a0019c0c1bbe137e5409f741cbde5f8", size = 62297, upload-time = "2024-12-14T11:14:43.408Z" }, + { url = "https://files.pythonhosted.org/packages/45/ec/dbb7e5a6b91f86bfb9eb7d2988a2730907b6a729875b949c7f022e8b88fa/uvicorn-0.51.0-py3-none-any.whl", hash = "sha256:5d38af6cd620f2ae3849fb44fd4879e0890aa1febe8d47eb355fb45d93fe6a5b", size = 73219, upload-time = "2026-07-08T10:59:04.44Z" }, ] [[package]]