diff --git a/.github/workflows/helm_unit_test.yml b/.github/workflows/helm_unit_test.yml index f95848945a0..f3d7bdf36cc 100644 --- a/.github/workflows/helm_unit_test.yml +++ b/.github/workflows/helm_unit_test.yml @@ -22,6 +22,12 @@ jobs: with: persist-credentials: false + - name: Check Lens Compose configuration + run: | + python3 -m unittest discover -s deploy/lens -p 'test_*.py' + python3 deploy/lens/configure.py --version 1.2.3 --env-file "$RUNNER_TEMP/lens.env" + docker compose --env-file "$RUNNER_TEMP/lens.env" -f deploy/lens/stack.yaml config --quiet + - name: Set up Helm 3.11.1 uses: azure/setup-helm@1a275c3b69536ee54be43f2070a358922e12c8d4 # v4.3.1 with: diff --git a/backend/routes/allowlist.py b/backend/routes/allowlist.py index ea56a1fd6df..44541276be0 100644 --- a/backend/routes/allowlist.py +++ b/backend/routes/allowlist.py @@ -23,6 +23,7 @@ BACKEND_PATH_PREFIXES: tuple[str, ...] = ( "/end_user/", "/sso/", "/liteadmin/slack/connect/", + "/moyai/connect/", "/login", "/v2/login", "/v3/login", diff --git a/deploy/lens/README.md b/deploy/lens/README.md index d5b9409e583..379913b3b7f 100644 --- a/deploy/lens/README.md +++ b/deploy/lens/README.md @@ -6,21 +6,20 @@ Agent exporters send traces directly to Lens. LiteLLM sends its optional request ## New local installation -Install Docker with Compose and Git, then build the gateway and Lens from one checkout: +Install Docker with Compose, Python 3.10 or later, and Git. Clone LiteLLM, select a published release that includes Lens, and start the existing Compose stack: ```bash git clone https://github.com/BerriAI/litellm.git cd litellm -export LITELLM_RELEASE_TAG="sha-$(git rev-parse HEAD)" -export LITELLM_MASTER_KEY="sk-$(openssl rand -hex 24)" -export LITELLM_LENS_SERVICE_TOKEN="$(openssl rand -hex 32)" -export OPENAI_API_KEY='' -docker compose -f docker/docker-compose.tracing.yml up -d --build +python3 deploy/lens/configure.py --version +docker compose --env-file deploy/lens/.env -f deploy/lens/stack.yaml up -d --wait ``` -Save the generated keys privately and reuse them when restarting or upgrading. This stack binds to localhost and uses development database passwords; use your normal secrets, TLS, backups, and ingress for a hosted deployment +The configuration command generates your keys and database passwords once, saves them in `deploy/lens/.env` with owner-only permissions, and preserves them on subsequent runs. Back up this file alongside your database volumes. Both images use the selected release; there is no local image build -Open `http://localhost:4002/ui/` and sign in as `admin` with `LITELLM_MASTER_KEY`. Under **Lens > Traces > Set up tracing**, generate a tracing key and copy the ingestion URL. Local exporters use `http://localhost:4318`. Model calls keep their existing LiteLLM URL and model key +Open `http://localhost:4000/ui/` and sign in as `admin` using `LITELLM_MASTER_KEY` from the saved file. Open **Lens**, select your framework, generate a tracing key, and copy the displayed configuration. The trace endpoint is already filled in. Keep your agent's existing model credentials; the tracing key only authorizes trace uploads + +PostgreSQL and ClickHouse use persistent Docker volumes and have no host ports. The dashboard and trace listener bind to localhost. Use your normal TLS and ingress for a hosted deployment. Stop the stack with `docker compose --env-file deploy/lens/.env -f deploy/lens/stack.yaml down`; omit `-v` to retain data Under **Lens > Investigations > Connect worker**, choose an analysis model and monthly budget. The deployed service connects automatically after you save these settings. There is no worker command or second token to copy @@ -68,14 +67,30 @@ Lens does not need provider credentials, PostgreSQL credentials, a GPU, or the L ### Kubernetes with Helm -Both `helm/litellm` and `helm/litellm-helm` support the Lens service. Keep your existing release, namespace, values, and database configuration. Create two Secrets through your normal secret manager: `litellm-lens-service` with key `service-token`, and `litellm-lens-clickhouse` with key `url` +Both `helm/litellm` and `helm/litellm-helm` support Lens. Keep your existing chart, release name, namespace, and values. Add: + +```yaml +lensWorker: + enabled: true +``` + +Then run your usual Helm deployment command using the matching published chart. The chart supplies the matching Lens image, generates the shared service secret, starts a single ClickHouse instance with a persistent volume, and connects the services. Your cluster needs a default storage class, or set `lensWorker.clickhouse.storageClassName`. Bundled storage defaults to 20 GiB; set `lensWorker.clickhouse.storage` before installation to choose another size + +When your chart manages an ingress with one hostname, the chart fills in the public tracing address and routes `/lens-ingest` directly to Lens. TLS is detected from `ingress.tls` or an ALB certificate annotation. With custom ingress, multiple hostnames, or TLS terminated elsewhere, set the address explicitly: + +```yaml +lensWorker: + enabled: true + publicUrl: https:///lens-ingest +``` + +For a dedicated trace hostname, configure `lensWorker.ingress.enabled`, `host`, `className`, and `tls`. Its hostname supplies the public address unless you override `publicUrl`. Internal Lens routes stay private + +To use an existing ClickHouse database and secrets managed by your platform, keep these overrides: ```yaml lensWorker: enabled: true - image: - repository: - digest: sha256: serviceTokenSecret: name: litellm-lens-service key: service-token @@ -84,21 +99,13 @@ lensWorker: key: url clickhouseDatabase: litellm retentionDays: 14 - publicUrl: https:///lens-ingest ``` -Set `clickhouseDatabase` and `retentionDays` to your existing database and retention before upgrading +Supplying `clickhouseSecret.name` uses that database and disables bundled storage. Keep your database name and retention policy. For GitOps tools that render Helm without cluster access, supply both existing secrets so rendering cannot regenerate credentials -When the chart's main ingress is enabled, it routes `/lens-ingest` directly to Lens. With a custom ingress, add that route yourself. For a dedicated hostname, use `lensWorker.ingress.enabled`, `host`, `className`, and `tls`, and set `publicUrl` to that hostname. The chart connects LiteLLM to Lens internally and gives both services the shared secret +Normal Helm upgrades reuse the generated credentials. Secrets are retained on uninstall, and the ClickHouse volume is retained by Kubernetes. Back them up together. Treat changing the database, storage class, or secret reference as an infrastructure change, not a routine version update -Update your existing component image overrides to matching builds, then use the chart from that checkout: - -```bash -helm upgrade --install litellm ./helm/litellm \ - --namespace litellm -f values.yaml --wait -``` - -Use `./helm/litellm-helm` if that is your existing chart. `lensWorker.replicaCount` scales ingestion and investigations. Each replica needs access to the same ClickHouse and gateway. Credentials refresh every 30 seconds; a newly created key may briefly receive a retryable 429. Revocations propagate on refresh, and a replica stops accepting traces when its credential snapshot reaches 90 seconds +After deployment, open **Lens**. If it was already open, click **Check setup**. The setup section moves to your framework and tracing key when Lens is reachable and storage is ready. Investigation setup asks for the analysis model and budget; the installed service connects automatically ## Upgrade @@ -106,13 +113,9 @@ Upgrade LiteLLM and Lens from the same source commit and release identity. For a Keep the same databases, encryption keys, shared service secret, and public ingestion URL. Pause scheduled investigations and finish or cancel active runs, update both images through your usual deployment process, then check ingestion and run an investigation before resuming schedules. Do not run `docker compose down -v` -When upgrading from the Python worker, replace it with the Rust Lens service, move the existing ClickHouse connection to Lens, and configure the service URLs and secret on LiteLLM. Existing trace data remains in the same ClickHouse database; findings and settings remain in PostgreSQL. Stop the old worker. Generate dedicated tracing keys and change agent exporters to the ingestion URL. A virtual model key no longer authorizes uploads; the old gateway upload endpoints return 410 with setup guidance - -If you retain an explicit `LENS_WORKER_TOKEN`, it remains an optional investigation credential. Normal setup uses the shared service connection and registers one managed worker identity. Configure the analysis model and billing key in the dashboard; provider keys stay on LiteLLM - ## Development -`make lens-dev` starts LiteLLM, the Rust Lens service, and the hot-reload dashboard. Set `LENS_DEV_PROXY_PORT` and `LENS_DEV_UI_PORT` to change the local ports. For containers, pass the same release identity to both builds. Unversioned or incompatible workers are refused before claiming work +`make lens-dev` starts LiteLLM, Lens, and the hot-reload dashboard. Set `LENS_DEV_PROXY_PORT` and `LENS_DEV_UI_PORT` to change the local ports. For containers, pass the same release identity to both builds. Unversioned or incompatible workers are refused before claiming work ## Configure a lens diff --git a/deploy/lens/config.yaml b/deploy/lens/config.yaml index 43bfe32ec26..f9eb15865a7 100644 --- a/deploy/lens/config.yaml +++ b/deploy/lens/config.yaml @@ -1,3 +1,5 @@ +model_list: [] + general_settings: master_key: os.environ/LITELLM_MASTER_KEY tracing: diff --git a/deploy/lens/configure.py b/deploy/lens/configure.py new file mode 100644 index 00000000000..866d986c570 --- /dev/null +++ b/deploy/lens/configure.py @@ -0,0 +1,73 @@ +from __future__ import annotations + +import argparse +import os +import re +import secrets +import shlex +import sys +from pathlib import Path +from typing import Final + +SECRET_NAMES: Final = ( + "LITELLM_MASTER_KEY", + "LITELLM_SALT_KEY", + "LITELLM_LENS_SERVICE_TOKEN", + "POSTGRES_PASSWORD", + "CLICKHOUSE_PASSWORD", +) + + +def environment_content(path: Path) -> str: + if not path.exists(): + return "".join( + f"{name}={'sk-' if name.endswith('KEY') else ''}{secrets.token_hex(32)}\n" for name in SECRET_NAMES + ) + saved: Final = path.read_text().splitlines() + values: Final = dict(line.split("=", 1) for line in saved if "=" in line) + if any(not values.get(name) for name in SECRET_NAMES): + raise ValueError(f"{path} is incomplete. Restore your saved credentials before continuing") + return "\n".join(line for line in saved if not line.startswith("LITELLM_VERSION=")) + "\n" + + +def configure(path: Path, version: str) -> None: + release: Final = version.removeprefix("v") + if not re.fullmatch(r"[0-9]+\.[0-9]+\.[0-9]+(?:[-.][a-zA-Z0-9.-]+)?", release): + raise ValueError("Use a published release version, such as 1.82.0 or 1.82.0-nightly") + if path.is_symlink(): + raise ValueError(f"Refusing to replace a symlink: {path}") + existing: Final = path.exists() + content: Final = environment_content(path) + temporary: Final = path.with_name(f".{path.name}.{secrets.token_hex(8)}") + descriptor: Final = os.open(temporary, os.O_WRONLY | os.O_CREAT | os.O_EXCL, 0o600) + try: + with os.fdopen(descriptor, "w") as output: + output.write(content + f"LITELLM_VERSION={release}\n") + if existing: + os.replace(temporary, path) + else: + os.link(temporary, path) + finally: + temporary.unlink(missing_ok=True) + + +def main() -> None: + parser: Final = argparse.ArgumentParser(description="Create or update the configuration for the Lens Compose stack") + parser.add_argument("--version", required=True, help="Published LiteLLM release; Lens uses the matching version") + parser.add_argument("--env-file", type=Path, default=Path(__file__).with_name(".env")) + arguments: Final = parser.parse_args() + try: + configure(arguments.env_file, arguments.version) + except (OSError, ValueError) as error: + parser.exit(1, f"Could not configure Lens: {error}\n") + sys.stdout.write( + f"Saved {arguments.env_file}. Existing keys and database passwords are preserved\n" + f"Start with: docker compose --env-file {shlex.quote(str(arguments.env_file))} " + "-f deploy/lens/stack.yaml up -d --wait\n" + "Open http://localhost:4000/ui/ and sign in as admin with LITELLM_MASTER_KEY from the saved file\n" + "Back up this file with your database volumes. Do not commit it\n" + ) + + +if __name__ == "__main__": + main() diff --git a/deploy/lens/test_configure.py b/deploy/lens/test_configure.py new file mode 100644 index 00000000000..dc61871ab74 --- /dev/null +++ b/deploy/lens/test_configure.py @@ -0,0 +1,52 @@ +import tempfile +import unittest +from pathlib import Path + +from configure import SECRET_NAMES, configure + + +class ComposeConfigurationTests(unittest.TestCase): + def test_restart_and_upgrade_preserve_private_credentials_and_custom_settings(self) -> None: + with tempfile.TemporaryDirectory() as directory: + path = Path(directory) / ".env" + configure(path, "v1.2.3") + original = dict(line.split("=", 1) for line in path.read_text().splitlines()) + self.assertEqual(path.stat().st_mode & 0o777, 0o600) + self.assertEqual(len({original[name] for name in SECRET_NAMES}), len(SECRET_NAMES)) + self.assertTrue(all(len(original[name]) >= 64 for name in SECRET_NAMES)) + with path.open("a") as output: + output.write("LITELLM_LENS_PUBLIC_URL=https://traces.example/prefix\n") + configure(path, "1.2.3") + configure(path, "v1.2.4-nightly") + updated = dict(line.split("=", 1) for line in path.read_text().splitlines()) + self.assertEqual({name: updated[name] for name in SECRET_NAMES}, {name: original[name] for name in SECRET_NAMES}) + self.assertEqual(updated["LITELLM_VERSION"], "1.2.4-nightly") + self.assertEqual(updated["LITELLM_LENS_PUBLIC_URL"], "https://traces.example/prefix") + self.assertEqual(path.stat().st_mode & 0o777, 0o600) + + def test_incomplete_configuration_is_never_replaced_with_new_database_passwords(self) -> None: + with tempfile.TemporaryDirectory() as directory: + path = Path(directory) / ".env" + original = "POSTGRES_PASSWORD=existing\n" + path.write_text(original) + with self.assertRaisesRegex(ValueError, "incomplete"): + configure(path, "1.2.3") + self.assertEqual(path.read_text(), original) + + def test_invalid_release_and_symlink_leave_existing_files_untouched(self) -> None: + with tempfile.TemporaryDirectory() as directory: + path = Path(directory) / ".env" + target = Path(directory) / "saved" + target.write_text("preserve") + path.symlink_to(target) + with self.assertRaisesRegex(ValueError, "symlink"): + configure(path, "1.2.3") + self.assertEqual(target.read_text(), "preserve") + path.unlink() + with self.assertRaisesRegex(ValueError, "published release"): + configure(path, "1.2.3\nPOSTGRES_PASSWORD=replaced") + self.assertFalse(path.exists()) + + +if __name__ == "__main__": + unittest.main() diff --git a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py index c07bb216602..28c4cf5e66f 100644 --- a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py +++ b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py @@ -255,7 +255,7 @@ def _storage_metadata_of(file_object: OpenAIFileObject | None) -> Mapping[str, s _MANAGED_FILES_TARGET: Final = "managed_files" -class PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): +class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): # Class variables or attributes def __init__(self, internal_usage_cache: InternalUsageCache, prisma_client: PrismaClient): self.internal_usage_cache = internal_usage_cache @@ -2110,4 +2110,4 @@ class PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): verbose_logger.debug( f"Converted file {file_id} from storage backend to base64 with format {content_type}" ) -_PROXY_LiteLLMManagedFiles = PROXY_LiteLLMManagedFiles +PROXY_LiteLLMManagedFiles = _PROXY_LiteLLMManagedFiles diff --git a/enterprise/litellm_enterprise/proxy/hooks/managed_vector_stores.py b/enterprise/litellm_enterprise/proxy/hooks/managed_vector_stores.py index c53f359a468..0e9657aba96 100644 --- a/enterprise/litellm_enterprise/proxy/hooks/managed_vector_stores.py +++ b/enterprise/litellm_enterprise/proxy/hooks/managed_vector_stores.py @@ -38,7 +38,7 @@ else: PrismaClient = Any -class PROXY_LiteLLMManagedVectorStores( +class _PROXY_LiteLLMManagedVectorStores( CustomLogger, BaseManagedResource[VectorStoreCreateResponse] ): """ @@ -462,4 +462,4 @@ class PROXY_LiteLLMManagedVectorStores( parent_otel_span=parent_otel_span, resource_id_key="vector_store_id", ) -_PROXY_LiteLLMManagedVectorStores = PROXY_LiteLLMManagedVectorStores +PROXY_LiteLLMManagedVectorStores = _PROXY_LiteLLMManagedVectorStores diff --git a/gateway/routes/allowlist.py b/gateway/routes/allowlist.py index 31c05f21b1b..3cfe9569b34 100644 --- a/gateway/routes/allowlist.py +++ b/gateway/routes/allowlist.py @@ -62,6 +62,8 @@ GATEWAY_PATH_PREFIXES: tuple[str, ...] = ( "/rerank", "/v1/decisions", "/decisions", + "/v1/systemone", + "/systemone", "/v1/ocr", "/ocr", "/v1/rag/", diff --git a/helm/litellm-helm/templates/_helpers.tpl b/helm/litellm-helm/templates/_helpers.tpl index 5e5b47b586a..6bf4023b353 100644 --- a/helm/litellm-helm/templates/_helpers.tpl +++ b/helm/litellm-helm/templates/_helpers.tpl @@ -371,3 +371,30 @@ shutdown drain window. {{- $_ := set $labels "app.kubernetes.io/name" (printf "%s-lens-worker" (include "litellm.name" . | trunc 51 | trimSuffix "-")) -}} {{- toYaml $labels -}} {{- end -}} + +{{- define "litellm.lensWorker.serviceTokenSecretName" -}} +{{- .Values.lensWorker.serviceTokenSecret.name | default (printf "%s-lens-service" (include "litellm.fullname" .)) -}} +{{- end -}} + +{{- define "litellm.lensWorker.bundledClickhouse" -}} +{{- if and .Values.lensWorker.enabled .Values.lensWorker.clickhouse.enabled (not .Values.lensWorker.clickhouseSecret.name) -}}true{{- end -}} +{{- end -}} + +{{- define "litellm.lensWorker.publicUrl" -}} +{{- if .Values.lensWorker.publicUrl -}} +{{- .Values.lensWorker.publicUrl -}} +{{- else if .Values.lensWorker.ingress.enabled -}} +{{- $tls := or (not (empty .Values.lensWorker.ingress.tls)) (hasKey .Values.lensWorker.ingress.annotations "alb.ingress.kubernetes.io/certificate-arn") -}} +{{- printf "%s://%s" (ternary "https" "http" $tls) (required "lensWorker.ingress.host is required" .Values.lensWorker.ingress.host) -}} +{{- else if and .Values.ingress.enabled (eq (len .Values.ingress.hosts) 1) -}} +{{- $host := required "ingress.hosts[0].host is required" (first .Values.ingress.hosts).host -}} +{{- $tls := or (not (empty .Values.ingress.tls)) (hasKey .Values.ingress.annotations "alb.ingress.kubernetes.io/certificate-arn") -}} +{{- printf "%s://%s/lens-ingest" (ternary "https" "http" $tls) $host -}} +{{- else -}} +{{- fail "lensWorker.publicUrl is required when there is no single ingress hostname" -}} +{{- end -}} +{{- end -}} + +{{- define "litellm.lensWorker.clickhouseName" -}} +{{- printf "%s-lens-clickhouse" (include "litellm.fullname" . | trunc 47 | trimSuffix "-") -}} +{{- end -}} diff --git a/helm/litellm-helm/templates/deployment.yaml b/helm/litellm-helm/templates/deployment.yaml index 4aac75fba8c..bfca52a927b 100644 --- a/helm/litellm-helm/templates/deployment.yaml +++ b/helm/litellm-helm/templates/deployment.yaml @@ -60,11 +60,11 @@ spec: - name: LITELLM_LENS_URL value: {{ printf "http://%s-lens-worker:%v" (include "litellm.fullname" .) .Values.lensWorker.service.port | quote }} - name: LITELLM_LENS_PUBLIC_URL - value: {{ required "lensWorker.publicUrl is required" .Values.lensWorker.publicUrl | quote }} + value: {{ include "litellm.lensWorker.publicUrl" . | quote }} - name: LITELLM_LENS_SERVICE_TOKEN valueFrom: secretKeyRef: - name: {{ required "lensWorker.serviceTokenSecret.name is required" .Values.lensWorker.serviceTokenSecret.name | quote }} + name: {{ include "litellm.lensWorker.serviceTokenSecretName" . | quote }} key: {{ .Values.lensWorker.serviceTokenSecret.key | quote }} {{- end }} {{- include "litellm.proxyEnv" . | nindent 12 }} diff --git a/helm/litellm-helm/templates/lens/clickhouse.yaml b/helm/litellm-helm/templates/lens/clickhouse.yaml new file mode 100644 index 00000000000..e046f86caf2 --- /dev/null +++ b/helm/litellm-helm/templates/lens/clickhouse.yaml @@ -0,0 +1,98 @@ +{{- if include "litellm.lensWorker.bundledClickhouse" . }} +{{- $name := include "litellm.lensWorker.clickhouseName" . }} +apiVersion: v1 +kind: Service +metadata: + name: {{ $name }} +spec: + clusterIP: None + selector: + app.kubernetes.io/instance: {{ .Release.Name }} + app.kubernetes.io/component: lens-clickhouse + ports: + - name: http + port: 8123 + targetPort: http +--- +apiVersion: apps/v1 +kind: StatefulSet +metadata: + name: {{ $name }} +spec: + serviceName: {{ $name }} + replicas: 1 + selector: + matchLabels: + app.kubernetes.io/instance: {{ .Release.Name }} + app.kubernetes.io/component: lens-clickhouse + template: + metadata: + labels: + app.kubernetes.io/instance: {{ .Release.Name }} + app.kubernetes.io/component: lens-clickhouse + spec: + automountServiceAccountToken: false + {{- with .Values.imagePullSecrets }} + imagePullSecrets: + {{- toYaml . | nindent 8 }} + {{- end }} + securityContext: + runAsNonRoot: true + runAsUser: 101 + runAsGroup: 101 + fsGroup: 101 + seccompProfile: + type: RuntimeDefault + containers: + - name: clickhouse + image: {{ .Values.lensWorker.clickhouse.image | quote }} + securityContext: + allowPrivilegeEscalation: false + capabilities: + drop: [ALL] + env: + - name: CLICKHOUSE_USER + value: default + - name: CLICKHOUSE_PASSWORD + valueFrom: + secretKeyRef: + name: {{ $name }} + key: password + - name: CLICKHOUSE_DEFAULT_ACCESS_MANAGEMENT + value: "1" + ports: + - name: http + containerPort: 8123 + startupProbe: + httpGet: + path: /ping + port: http + periodSeconds: 5 + timeoutSeconds: 3 + failureThreshold: 60 + readinessProbe: + httpGet: + path: /ping + port: http + livenessProbe: + httpGet: + path: /ping + port: http + timeoutSeconds: 3 + resources: + {{- toYaml .Values.lensWorker.clickhouse.resources | nindent 12 }} + volumeMounts: + - name: data + mountPath: /var/lib/clickhouse + volumeClaimTemplates: + - metadata: + name: data + spec: + accessModes: [ReadWriteOnce] + {{- if ne .Values.lensWorker.clickhouse.storageClassName nil }} + storageClassName: {{ .Values.lensWorker.clickhouse.storageClassName | quote }} + {{- end }} + resources: + requests: + storage: {{ .Values.lensWorker.clickhouse.storage | quote }} +{{- end }} diff --git a/helm/litellm-helm/templates/lens/deployment.yaml b/helm/litellm-helm/templates/lens/deployment.yaml index dee141acde8..b1693dfd9ff 100644 --- a/helm/litellm-helm/templates/lens/deployment.yaml +++ b/helm/litellm-helm/templates/lens/deployment.yaml @@ -45,13 +45,23 @@ spec: - name: LITELLM_LENS_SERVICE_TOKEN valueFrom: secretKeyRef: - name: {{ required "lensWorker.serviceTokenSecret.name is required" .Values.lensWorker.serviceTokenSecret.name | quote }} + name: {{ include "litellm.lensWorker.serviceTokenSecretName" . | quote }} key: {{ .Values.lensWorker.serviceTokenSecret.key | quote }} + {{- if include "litellm.lensWorker.bundledClickhouse" . }} + - name: CLICKHOUSE_HOST + value: {{ include "litellm.lensWorker.clickhouseName" . | quote }} + - name: CLICKHOUSE_PASSWORD + valueFrom: + secretKeyRef: + name: {{ include "litellm.lensWorker.clickhouseName" . | quote }} + key: password + {{- else }} - name: CLICKHOUSE_URL valueFrom: secretKeyRef: name: {{ required "lensWorker.clickhouseSecret.name is required" .Values.lensWorker.clickhouseSecret.name | quote }} key: {{ .Values.lensWorker.clickhouseSecret.key | quote }} + {{- end }} - name: CLICKHOUSE_DATABASE value: {{ .Values.lensWorker.clickhouseDatabase | quote }} - name: AGENT_TRACING_RETENTION_DAYS diff --git a/helm/litellm-helm/templates/lens/secrets.yaml b/helm/litellm-helm/templates/lens/secrets.yaml new file mode 100644 index 00000000000..3809167c1da --- /dev/null +++ b/helm/litellm-helm/templates/lens/secrets.yaml @@ -0,0 +1,27 @@ +{{- if and .Values.lensWorker.enabled (not .Values.lensWorker.serviceTokenSecret.name) }} +{{- $name := include "litellm.lensWorker.serviceTokenSecretName" . }} +{{- $existing := lookup "v1" "Secret" .Release.Namespace $name }} +apiVersion: v1 +kind: Secret +metadata: + name: {{ $name }} + annotations: + helm.sh/resource-policy: keep +type: Opaque +data: + {{ .Values.lensWorker.serviceTokenSecret.key }}: {{ if $existing }}{{ required "Saved Lens service secret is missing its key" (index $existing.data .Values.lensWorker.serviceTokenSecret.key) | quote }}{{ else }}{{ randAlphaNum 64 | b64enc | quote }}{{ end }} +{{- end }} +{{- if include "litellm.lensWorker.bundledClickhouse" . }} +{{- $name := include "litellm.lensWorker.clickhouseName" . }} +{{- $existing := lookup "v1" "Secret" .Release.Namespace $name }} +--- +apiVersion: v1 +kind: Secret +metadata: + name: {{ $name }} + annotations: + helm.sh/resource-policy: keep +type: Opaque +data: + password: {{ if $existing }}{{ required "Saved Lens ClickHouse secret is missing its password" (index $existing.data "password") | quote }}{{ else }}{{ randAlphaNum 64 | b64enc | quote }}{{ end }} +{{- end }} diff --git a/helm/litellm-helm/tests/lens_endpoint_defaults_tests.yaml b/helm/litellm-helm/tests/lens_endpoint_defaults_tests.yaml new file mode 100644 index 00000000000..22dcc7e9f9c --- /dev/null +++ b/helm/litellm-helm/tests/lens_endpoint_defaults_tests.yaml @@ -0,0 +1,84 @@ +suite: Lens endpoint defaults +templates: + - deployment.yaml + - configmap-litellm.yaml +set: + lensWorker.enabled: true + ingress.enabled: true + ingress.hosts: [{host: gateway.example, paths: [{path: /, pathType: Prefix}]}] + ingress.tls: + - hosts: [gateway.example] + secretName: tls +tests: + - it: derives the trace endpoint from the deployment hostname and TLS + template: deployment.yaml + asserts: + - contains: + path: spec.template.spec.containers[0].env + content: + name: LITELLM_LENS_PUBLIC_URL + value: https://gateway.example/lens-ingest + - it: preserves an explicitly configured public address + set: + lensWorker.publicUrl: https://custom.example/prefix + template: deployment.yaml + asserts: + - contains: + path: spec.template.spec.containers[0].env + content: + name: LITELLM_LENS_PUBLIC_URL + value: https://custom.example/prefix + - it: uses the dedicated Lens ingress when configured + set: + lensWorker.ingress.enabled: true + lensWorker.ingress.host: traces.example + lensWorker.ingress.tls: + - hosts: [traces.example] + secretName: traces-tls + template: deployment.yaml + asserts: + - contains: + path: spec.template.spec.containers[0].env + content: + name: LITELLM_LENS_PUBLIC_URL + value: https://traces.example + - it: uses HTTPS for a dedicated ALB ingress with certificate annotations + set: + lensWorker.ingress.enabled: true + lensWorker.ingress.host: traces.example + lensWorker.ingress.className: alb + lensWorker.ingress.tls: [] + lensWorker.ingress.annotations: + alb.ingress.kubernetes.io/certificate-arn: arn:aws:acm:us-west-2:123456789012:certificate/test + alb.ingress.kubernetes.io/listen-ports: '[{"HTTP":80},{"HTTPS":443}]' + alb.ingress.kubernetes.io/ssl-redirect: "443" + template: deployment.yaml + asserts: + - contains: + path: spec.template.spec.containers[0].env + content: + name: LITELLM_LENS_PUBLIC_URL + value: https://traces.example + - it: uses HTTP for a dedicated ingress without its own TLS + set: + lensWorker.ingress.enabled: true + lensWorker.ingress.host: traces.example + lensWorker.ingress.tls: [] + lensWorker.ingress.annotations: {} + template: deployment.yaml + asserts: + - contains: + path: spec.template.spec.containers[0].env + content: + name: LITELLM_LENS_PUBLIC_URL + value: http://traces.example + - it: uses HTTP when ingress has no TLS + set: + ingress.tls: [] + template: deployment.yaml + asserts: + - contains: + path: spec.template.spec.containers[0].env + content: + name: LITELLM_LENS_PUBLIC_URL + value: http://gateway.example/lens-ingest diff --git a/helm/litellm-helm/tests/lens_saved_secrets_tests.yaml b/helm/litellm-helm/tests/lens_saved_secrets_tests.yaml new file mode 100644 index 00000000000..1af9bc77594 --- /dev/null +++ b/helm/litellm-helm/tests/lens_saved_secrets_tests.yaml @@ -0,0 +1,45 @@ +suite: Lens preserves credentials across Helm upgrades +templates: + - lens/secrets.yaml +set: + fullnameOverride: lens-test + lensWorker.enabled: true +release: + namespace: lens + name: lens-test + upgrade: true +kubernetesProvider: + scheme: + v1/Secret: + gvr: + version: v1 + resource: secrets + namespaced: true + objects: + - apiVersion: v1 + kind: Secret + metadata: + name: lens-test-lens-service + namespace: lens + data: + service-token: c2F2ZWQtc2VydmljZS10b2tlbg== + - apiVersion: v1 + kind: Secret + metadata: + name: lens-test-lens-clickhouse + namespace: lens + data: + password: c2F2ZWQtZGF0YWJhc2UtcGFzc3dvcmQ= +tests: + - it: reuses the service credential instead of breaking running services + documentIndex: 0 + asserts: + - equal: + path: data.service-token + value: c2F2ZWQtc2VydmljZS10b2tlbg== + - it: reuses the database password instead of locking out stored traces + documentIndex: 1 + asserts: + - equal: + path: data.password + value: c2F2ZWQtZGF0YWJhc2UtcGFzc3dvcmQ= diff --git a/helm/litellm-helm/tests/lens_service_tests.yaml b/helm/litellm-helm/tests/lens_service_tests.yaml index 197f447f5f3..483db4d42c2 100644 --- a/helm/litellm-helm/tests/lens_service_tests.yaml +++ b/helm/litellm-helm/tests/lens_service_tests.yaml @@ -117,7 +117,7 @@ tests: lensWorker.clickhouseSecret.name: lens-storage asserts: - failedTemplate: - errorMessage: lensWorker.publicUrl is required + errorMessage: lensWorker.publicUrl is required when there is no single ingress hostname - it: omits Lens connection settings when disabled in deployment.yaml template: deployment.yaml asserts: diff --git a/helm/litellm-helm/tests/lens_setup_tests.yaml b/helm/litellm-helm/tests/lens_setup_tests.yaml new file mode 100644 index 00000000000..25b2da62743 --- /dev/null +++ b/helm/litellm-helm/tests/lens_setup_tests.yaml @@ -0,0 +1,160 @@ +suite: Lens managed setup +set: + fullnameOverride: lens-test + lensWorker.enabled: true + lensWorker.publicUrl: https://traces.example +release: + namespace: lens + name: lens-test +templates: + - lens/deployment.yaml + - lens/clickhouse.yaml + - lens/secrets.yaml +tests: + - it: generates both private credentials for a new installation + template: lens/secrets.yaml + asserts: + - hasDocuments: + count: 2 + - matchRegex: + path: data.service-token + pattern: '^[A-Za-z0-9+/]{86}==$' + documentIndex: 0 + - matchRegex: + path: data.password + pattern: '^[A-Za-z0-9+/]{86}==$' + documentIndex: 1 + - equal: + path: metadata.annotations["helm.sh/resource-policy"] + value: keep + - it: connects Lens to its private bundled storage + template: lens/deployment.yaml + asserts: + - contains: + path: spec.template.spec.containers[0].env + content: + name: CLICKHOUSE_HOST + value: lens-test-lens-clickhouse + - contains: + path: spec.template.spec.containers[0].env + content: + name: CLICKHOUSE_PASSWORD + valueFrom: + secretKeyRef: + name: lens-test-lens-clickhouse + key: password + - it: stores traces on a persistent volume with the chosen storage class + template: lens/clickhouse.yaml + documentIndex: 1 + set: + lensWorker.clickhouse.storage: 40Gi + lensWorker.clickhouse.storageClassName: fast + asserts: + - equal: + path: spec.volumeClaimTemplates[0].spec + value: + accessModes: [ReadWriteOnce] + storageClassName: fast + resources: + requests: + storage: 40Gi + - contains: + path: spec.template.spec.containers[0].env + content: + name: CLICKHOUSE_PASSWORD + valueFrom: + secretKeyRef: + name: lens-test-lens-clickhouse + key: password + - it: keeps an existing external database instead of creating a new one + template: lens/clickhouse.yaml + set: + lensWorker.clickhouseSecret.name: external-clickhouse + asserts: + - hasDocuments: + count: 0 + - it: leaves supplied secrets under their existing manager + template: lens/secrets.yaml + set: + lensWorker.clickhouseSecret.name: external-clickhouse + lensWorker.serviceTokenSecret.name: external-service + asserts: + - hasDocuments: + count: 0 + - it: creates no credentials or database when Lens is disabled + templates: + - lens/secrets.yaml + - lens/clickhouse.yaml + set: + lensWorker.enabled: false + asserts: + - hasDocuments: + count: 0 + - it: requires an external database when bundled storage is explicitly disabled + template: lens/deployment.yaml + set: + lensWorker.clickhouse.enabled: false + asserts: + - failedTemplate: + errorMessage: lensWorker.clickhouseSecret.name is required + - it: keeps storage names valid and uses the same name for the connection + set: + fullnameOverride: aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa + asserts: + - matchRegex: + path: metadata.name + pattern: '^[a-z]([-a-z0-9]{0,61}[a-z0-9])?$' + template: lens/clickhouse.yaml + documentIndex: 0 + - matchRegex: + path: metadata.name + pattern: '^[a-z]([-a-z0-9]{0,61}[a-z0-9])?$' + template: lens/clickhouse.yaml + documentIndex: 1 + - equal: + path: metadata.name + value: aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa-lens-clickhouse + template: lens/clickhouse.yaml + documentIndex: 0 + - equal: + path: spec.serviceName + value: aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa-lens-clickhouse + template: lens/clickhouse.yaml + documentIndex: 1 + - equal: + path: metadata.name + value: aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa-lens-clickhouse + template: lens/secrets.yaml + documentIndex: 1 + - contains: + path: spec.template.spec.containers[0].env + content: + name: CLICKHOUSE_HOST + value: aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa-lens-clickhouse + template: lens/deployment.yaml + - contains: + path: spec.template.spec.containers[0].env + content: + name: CLICKHOUSE_PASSWORD + valueFrom: + secretKeyRef: + name: aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa-lens-clickhouse + key: password + template: lens/deployment.yaml + + - it: allows cold storage startup and inherits registry credentials + template: lens/clickhouse.yaml + documentIndex: 1 + set: + imagePullSecrets: [{name: registry-auth}] + asserts: + - equal: + path: spec.template.spec.imagePullSecrets + value: [{name: registry-auth}] + - equal: + path: spec.template.spec.containers[0].startupProbe + value: + httpGet: {path: /ping, port: http} + periodSeconds: 5 + timeoutSeconds: 3 + failureThreshold: 60 diff --git a/helm/litellm-helm/values.yaml b/helm/litellm-helm/values.yaml index 821557fd116..57929738cd5 100644 --- a/helm/litellm-helm/values.yaml +++ b/helm/litellm-helm/values.yaml @@ -667,6 +667,17 @@ lensWorker: serviceTokenSecret: name: "" key: service-token + clickhouse: + enabled: true + image: clickhouse/clickhouse-server:26.9.6.6 + storage: 20Gi + storageClassName: null + resources: + requests: + cpu: 100m + memory: 512Mi + limits: + memory: 2Gi clickhouseDatabase: litellm retentionDays: 14 clickhouseSecret: diff --git a/helm/litellm/templates/_helpers.tpl b/helm/litellm/templates/_helpers.tpl index 9c76e2da748..2728d0271ec 100644 --- a/helm/litellm/templates/_helpers.tpl +++ b/helm/litellm/templates/_helpers.tpl @@ -520,11 +520,11 @@ shutdown drain window. - name: LITELLM_LENS_URL value: {{ printf "http://%s-lens-worker:%v" (include "litellm.fullname" .) .Values.lensWorker.service.port | quote }} - name: LITELLM_LENS_PUBLIC_URL - value: {{ required "lensWorker.publicUrl is required" .Values.lensWorker.publicUrl | quote }} + value: {{ include "litellm.lensWorker.publicUrl" . | quote }} - name: LITELLM_LENS_SERVICE_TOKEN valueFrom: secretKeyRef: - name: {{ required "lensWorker.serviceTokenSecret.name is required" .Values.lensWorker.serviceTokenSecret.name | quote }} + name: {{ include "litellm.lensWorker.serviceTokenSecretName" . | quote }} key: {{ .Values.lensWorker.serviceTokenSecret.key | quote }} {{- end }} {{- end -}} @@ -534,3 +534,29 @@ shutdown drain window. {{- $_ := set $labels "app.kubernetes.io/name" (printf "%s-lens-worker" (include "litellm.name" . | trunc 51 | trimSuffix "-")) -}} {{- toYaml $labels -}} {{- end -}} + +{{- define "litellm.lensWorker.serviceTokenSecretName" -}} +{{- .Values.lensWorker.serviceTokenSecret.name | default (printf "%s-lens-service" (include "litellm.fullname" .)) -}} +{{- end -}} + +{{- define "litellm.lensWorker.bundledClickhouse" -}} +{{- if and .Values.lensWorker.enabled .Values.lensWorker.clickhouse.enabled (not .Values.lensWorker.clickhouseSecret.name) -}}true{{- end -}} +{{- end -}} + +{{- define "litellm.lensWorker.publicUrl" -}} +{{- if .Values.lensWorker.publicUrl -}} +{{- .Values.lensWorker.publicUrl -}} +{{- else if .Values.lensWorker.ingress.enabled -}} +{{- $tls := or (not (empty .Values.lensWorker.ingress.tls)) (hasKey .Values.lensWorker.ingress.annotations "alb.ingress.kubernetes.io/certificate-arn") -}} +{{- printf "%s://%s" (ternary "https" "http" $tls) (required "lensWorker.ingress.host is required" .Values.lensWorker.ingress.host) -}} +{{- else if and .Values.ingress.enabled .Values.ingress.host -}} +{{- $tls := or (not (empty .Values.ingress.tls)) (hasKey .Values.ingress.annotations "alb.ingress.kubernetes.io/certificate-arn") -}} +{{- printf "%s://%s/lens-ingest" (ternary "https" "http" $tls) .Values.ingress.host -}} +{{- else -}} +{{- fail "lensWorker.publicUrl is required when there is no single ingress hostname" -}} +{{- end -}} +{{- end -}} + +{{- define "litellm.lensWorker.clickhouseName" -}} +{{- printf "%s-lens-clickhouse" (include "litellm.fullname" . | trunc 47 | trimSuffix "-") -}} +{{- end -}} diff --git a/helm/litellm/templates/lens/clickhouse.yaml b/helm/litellm/templates/lens/clickhouse.yaml new file mode 100644 index 00000000000..e046f86caf2 --- /dev/null +++ b/helm/litellm/templates/lens/clickhouse.yaml @@ -0,0 +1,98 @@ +{{- if include "litellm.lensWorker.bundledClickhouse" . }} +{{- $name := include "litellm.lensWorker.clickhouseName" . }} +apiVersion: v1 +kind: Service +metadata: + name: {{ $name }} +spec: + clusterIP: None + selector: + app.kubernetes.io/instance: {{ .Release.Name }} + app.kubernetes.io/component: lens-clickhouse + ports: + - name: http + port: 8123 + targetPort: http +--- +apiVersion: apps/v1 +kind: StatefulSet +metadata: + name: {{ $name }} +spec: + serviceName: {{ $name }} + replicas: 1 + selector: + matchLabels: + app.kubernetes.io/instance: {{ .Release.Name }} + app.kubernetes.io/component: lens-clickhouse + template: + metadata: + labels: + app.kubernetes.io/instance: {{ .Release.Name }} + app.kubernetes.io/component: lens-clickhouse + spec: + automountServiceAccountToken: false + {{- with .Values.imagePullSecrets }} + imagePullSecrets: + {{- toYaml . | nindent 8 }} + {{- end }} + securityContext: + runAsNonRoot: true + runAsUser: 101 + runAsGroup: 101 + fsGroup: 101 + seccompProfile: + type: RuntimeDefault + containers: + - name: clickhouse + image: {{ .Values.lensWorker.clickhouse.image | quote }} + securityContext: + allowPrivilegeEscalation: false + capabilities: + drop: [ALL] + env: + - name: CLICKHOUSE_USER + value: default + - name: CLICKHOUSE_PASSWORD + valueFrom: + secretKeyRef: + name: {{ $name }} + key: password + - name: CLICKHOUSE_DEFAULT_ACCESS_MANAGEMENT + value: "1" + ports: + - name: http + containerPort: 8123 + startupProbe: + httpGet: + path: /ping + port: http + periodSeconds: 5 + timeoutSeconds: 3 + failureThreshold: 60 + readinessProbe: + httpGet: + path: /ping + port: http + livenessProbe: + httpGet: + path: /ping + port: http + timeoutSeconds: 3 + resources: + {{- toYaml .Values.lensWorker.clickhouse.resources | nindent 12 }} + volumeMounts: + - name: data + mountPath: /var/lib/clickhouse + volumeClaimTemplates: + - metadata: + name: data + spec: + accessModes: [ReadWriteOnce] + {{- if ne .Values.lensWorker.clickhouse.storageClassName nil }} + storageClassName: {{ .Values.lensWorker.clickhouse.storageClassName | quote }} + {{- end }} + resources: + requests: + storage: {{ .Values.lensWorker.clickhouse.storage | quote }} +{{- end }} diff --git a/helm/litellm/templates/lens/deployment.yaml b/helm/litellm/templates/lens/deployment.yaml index 93772a7a650..2ba13609188 100644 --- a/helm/litellm/templates/lens/deployment.yaml +++ b/helm/litellm/templates/lens/deployment.yaml @@ -45,13 +45,23 @@ spec: - name: LITELLM_LENS_SERVICE_TOKEN valueFrom: secretKeyRef: - name: {{ required "lensWorker.serviceTokenSecret.name is required" .Values.lensWorker.serviceTokenSecret.name | quote }} + name: {{ include "litellm.lensWorker.serviceTokenSecretName" . | quote }} key: {{ .Values.lensWorker.serviceTokenSecret.key | quote }} + {{- if include "litellm.lensWorker.bundledClickhouse" . }} + - name: CLICKHOUSE_HOST + value: {{ include "litellm.lensWorker.clickhouseName" . | quote }} + - name: CLICKHOUSE_PASSWORD + valueFrom: + secretKeyRef: + name: {{ include "litellm.lensWorker.clickhouseName" . | quote }} + key: password + {{- else }} - name: CLICKHOUSE_URL valueFrom: secretKeyRef: name: {{ required "lensWorker.clickhouseSecret.name is required" .Values.lensWorker.clickhouseSecret.name | quote }} key: {{ .Values.lensWorker.clickhouseSecret.key | quote }} + {{- end }} - name: CLICKHOUSE_DATABASE value: {{ .Values.lensWorker.clickhouseDatabase | quote }} - name: AGENT_TRACING_RETENTION_DAYS diff --git a/helm/litellm/templates/lens/secrets.yaml b/helm/litellm/templates/lens/secrets.yaml new file mode 100644 index 00000000000..3809167c1da --- /dev/null +++ b/helm/litellm/templates/lens/secrets.yaml @@ -0,0 +1,27 @@ +{{- if and .Values.lensWorker.enabled (not .Values.lensWorker.serviceTokenSecret.name) }} +{{- $name := include "litellm.lensWorker.serviceTokenSecretName" . }} +{{- $existing := lookup "v1" "Secret" .Release.Namespace $name }} +apiVersion: v1 +kind: Secret +metadata: + name: {{ $name }} + annotations: + helm.sh/resource-policy: keep +type: Opaque +data: + {{ .Values.lensWorker.serviceTokenSecret.key }}: {{ if $existing }}{{ required "Saved Lens service secret is missing its key" (index $existing.data .Values.lensWorker.serviceTokenSecret.key) | quote }}{{ else }}{{ randAlphaNum 64 | b64enc | quote }}{{ end }} +{{- end }} +{{- if include "litellm.lensWorker.bundledClickhouse" . }} +{{- $name := include "litellm.lensWorker.clickhouseName" . }} +{{- $existing := lookup "v1" "Secret" .Release.Namespace $name }} +--- +apiVersion: v1 +kind: Secret +metadata: + name: {{ $name }} + annotations: + helm.sh/resource-policy: keep +type: Opaque +data: + password: {{ if $existing }}{{ required "Saved Lens ClickHouse secret is missing its password" (index $existing.data "password") | quote }}{{ else }}{{ randAlphaNum 64 | b64enc | quote }}{{ end }} +{{- end }} diff --git a/helm/litellm/tests/lens_endpoint_defaults_tests.yaml b/helm/litellm/tests/lens_endpoint_defaults_tests.yaml new file mode 100644 index 00000000000..2d98dcf471e --- /dev/null +++ b/helm/litellm/tests/lens_endpoint_defaults_tests.yaml @@ -0,0 +1,86 @@ +suite: Lens endpoint defaults +templates: + - gateway/deployment.yaml + - gateway/configmap.yaml +set: + lensWorker.enabled: true + ingress.enabled: true + ingress.host: gateway.example + ingress.tls: + - hosts: [gateway.example] + secretName: tls +tests: + - it: derives the trace endpoint from the deployment hostname and TLS + template: gateway/deployment.yaml + asserts: + - contains: + path: spec.template.spec.containers[0].env + content: + name: LITELLM_LENS_PUBLIC_URL + value: https://gateway.example/lens-ingest + - it: preserves an explicitly configured public address + set: + lensWorker.publicUrl: https://custom.example/prefix + template: gateway/deployment.yaml + asserts: + - contains: + path: spec.template.spec.containers[0].env + content: + name: LITELLM_LENS_PUBLIC_URL + value: https://custom.example/prefix + - it: uses the dedicated Lens ingress when configured + set: + lensWorker.ingress.enabled: true + lensWorker.ingress.host: traces.example + lensWorker.ingress.tls: + - hosts: [traces.example] + secretName: traces-tls + template: gateway/deployment.yaml + asserts: + - contains: + path: spec.template.spec.containers[0].env + content: + name: LITELLM_LENS_PUBLIC_URL + value: https://traces.example + - it: uses HTTPS for a dedicated ALB ingress with certificate annotations + set: + lensWorker.ingress.enabled: true + lensWorker.ingress.host: traces.example + lensWorker.ingress.className: alb + lensWorker.ingress.tls: [] + lensWorker.ingress.annotations: + alb.ingress.kubernetes.io/certificate-arn: arn:aws:acm:us-west-2:123456789012:certificate/test + alb.ingress.kubernetes.io/listen-ports: '[{"HTTP":80},{"HTTPS":443}]' + alb.ingress.kubernetes.io/ssl-redirect: "443" + template: gateway/deployment.yaml + asserts: + - contains: + path: spec.template.spec.containers[0].env + content: + name: LITELLM_LENS_PUBLIC_URL + value: https://traces.example + - it: uses HTTP for a dedicated ingress without its own TLS + set: + lensWorker.ingress.enabled: true + lensWorker.ingress.host: traces.example + lensWorker.ingress.tls: [] + lensWorker.ingress.annotations: {} + template: gateway/deployment.yaml + asserts: + - contains: + path: spec.template.spec.containers[0].env + content: + name: LITELLM_LENS_PUBLIC_URL + value: http://traces.example + - it: uses HTTP when ingress has no TLS + set: + ingress.tls: [] + template: gateway/deployment.yaml + asserts: + - contains: + path: spec.template.spec.containers[0].env + content: + name: LITELLM_LENS_PUBLIC_URL + value: http://gateway.example/lens-ingest +values: + - ./values/required.yaml diff --git a/helm/litellm/tests/lens_saved_secrets_tests.yaml b/helm/litellm/tests/lens_saved_secrets_tests.yaml new file mode 100644 index 00000000000..294d9adfe96 --- /dev/null +++ b/helm/litellm/tests/lens_saved_secrets_tests.yaml @@ -0,0 +1,47 @@ +suite: Lens preserves credentials across Helm upgrades +templates: + - lens/secrets.yaml +set: + fullnameOverride: lens-test + lensWorker.enabled: true +release: + namespace: lens + name: lens-test + upgrade: true +kubernetesProvider: + scheme: + v1/Secret: + gvr: + version: v1 + resource: secrets + namespaced: true + objects: + - apiVersion: v1 + kind: Secret + metadata: + name: lens-test-lens-service + namespace: lens + data: + service-token: c2F2ZWQtc2VydmljZS10b2tlbg== + - apiVersion: v1 + kind: Secret + metadata: + name: lens-test-lens-clickhouse + namespace: lens + data: + password: c2F2ZWQtZGF0YWJhc2UtcGFzc3dvcmQ= +tests: + - it: reuses the service credential instead of breaking running services + documentIndex: 0 + asserts: + - equal: + path: data.service-token + value: c2F2ZWQtc2VydmljZS10b2tlbg== + - it: reuses the database password instead of locking out stored traces + documentIndex: 1 + asserts: + - equal: + path: data.password + value: c2F2ZWQtZGF0YWJhc2UtcGFzc3dvcmQ= +values: + - ./values/required.yaml diff --git a/helm/litellm/tests/lens_service_tests.yaml b/helm/litellm/tests/lens_service_tests.yaml index c5504025572..81ad93bbeee 100644 --- a/helm/litellm/tests/lens_service_tests.yaml +++ b/helm/litellm/tests/lens_service_tests.yaml @@ -136,7 +136,7 @@ tests: lensWorker.clickhouseSecret.name: lens-storage asserts: - failedTemplate: - errorMessage: lensWorker.publicUrl is required + errorMessage: lensWorker.publicUrl is required when there is no single ingress hostname - it: omits Lens connection settings when disabled in gateway/deployment.yaml template: gateway/deployment.yaml asserts: diff --git a/helm/litellm/tests/lens_setup_tests.yaml b/helm/litellm/tests/lens_setup_tests.yaml new file mode 100644 index 00000000000..d052ae81c78 --- /dev/null +++ b/helm/litellm/tests/lens_setup_tests.yaml @@ -0,0 +1,162 @@ +suite: Lens managed setup +set: + fullnameOverride: lens-test + lensWorker.enabled: true + lensWorker.publicUrl: https://traces.example +release: + namespace: lens + name: lens-test +templates: + - lens/deployment.yaml + - lens/clickhouse.yaml + - lens/secrets.yaml +tests: + - it: generates both private credentials for a new installation + template: lens/secrets.yaml + asserts: + - hasDocuments: + count: 2 + - matchRegex: + path: data.service-token + pattern: '^[A-Za-z0-9+/]{86}==$' + documentIndex: 0 + - matchRegex: + path: data.password + pattern: '^[A-Za-z0-9+/]{86}==$' + documentIndex: 1 + - equal: + path: metadata.annotations["helm.sh/resource-policy"] + value: keep + - it: connects Lens to its private bundled storage + template: lens/deployment.yaml + asserts: + - contains: + path: spec.template.spec.containers[0].env + content: + name: CLICKHOUSE_HOST + value: lens-test-lens-clickhouse + - contains: + path: spec.template.spec.containers[0].env + content: + name: CLICKHOUSE_PASSWORD + valueFrom: + secretKeyRef: + name: lens-test-lens-clickhouse + key: password + - it: stores traces on a persistent volume with the chosen storage class + template: lens/clickhouse.yaml + documentIndex: 1 + set: + lensWorker.clickhouse.storage: 40Gi + lensWorker.clickhouse.storageClassName: fast + asserts: + - equal: + path: spec.volumeClaimTemplates[0].spec + value: + accessModes: [ReadWriteOnce] + storageClassName: fast + resources: + requests: + storage: 40Gi + - contains: + path: spec.template.spec.containers[0].env + content: + name: CLICKHOUSE_PASSWORD + valueFrom: + secretKeyRef: + name: lens-test-lens-clickhouse + key: password + - it: keeps an existing external database instead of creating a new one + template: lens/clickhouse.yaml + set: + lensWorker.clickhouseSecret.name: external-clickhouse + asserts: + - hasDocuments: + count: 0 + - it: leaves supplied secrets under their existing manager + template: lens/secrets.yaml + set: + lensWorker.clickhouseSecret.name: external-clickhouse + lensWorker.serviceTokenSecret.name: external-service + asserts: + - hasDocuments: + count: 0 + - it: creates no credentials or database when Lens is disabled + templates: + - lens/secrets.yaml + - lens/clickhouse.yaml + set: + lensWorker.enabled: false + asserts: + - hasDocuments: + count: 0 + - it: requires an external database when bundled storage is explicitly disabled + template: lens/deployment.yaml + set: + lensWorker.clickhouse.enabled: false + asserts: + - failedTemplate: + errorMessage: lensWorker.clickhouseSecret.name is required + - it: keeps storage names valid and uses the same name for the connection + set: + fullnameOverride: aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa + asserts: + - matchRegex: + path: metadata.name + pattern: '^[a-z]([-a-z0-9]{0,61}[a-z0-9])?$' + template: lens/clickhouse.yaml + documentIndex: 0 + - matchRegex: + path: metadata.name + pattern: '^[a-z]([-a-z0-9]{0,61}[a-z0-9])?$' + template: lens/clickhouse.yaml + documentIndex: 1 + - equal: + path: metadata.name + value: aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa-lens-clickhouse + template: lens/clickhouse.yaml + documentIndex: 0 + - equal: + path: spec.serviceName + value: aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa-lens-clickhouse + template: lens/clickhouse.yaml + documentIndex: 1 + - equal: + path: metadata.name + value: aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa-lens-clickhouse + template: lens/secrets.yaml + documentIndex: 1 + - contains: + path: spec.template.spec.containers[0].env + content: + name: CLICKHOUSE_HOST + value: aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa-lens-clickhouse + template: lens/deployment.yaml + - contains: + path: spec.template.spec.containers[0].env + content: + name: CLICKHOUSE_PASSWORD + valueFrom: + secretKeyRef: + name: aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa-lens-clickhouse + key: password + template: lens/deployment.yaml + - it: allows cold storage startup and inherits registry credentials + template: lens/clickhouse.yaml + documentIndex: 1 + set: + imagePullSecrets: [{name: registry-auth}] + asserts: + - equal: + path: spec.template.spec.imagePullSecrets + value: [{name: registry-auth}] + - equal: + path: spec.template.spec.containers[0].startupProbe + value: + httpGet: {path: /ping, port: http} + periodSeconds: 5 + timeoutSeconds: 3 + failureThreshold: 60 + +values: + - ./values/required.yaml diff --git a/helm/litellm/tests/lens_worker_tests.yaml b/helm/litellm/tests/lens_worker_tests.yaml index b93797dae18..b935734b4b3 100644 --- a/helm/litellm/tests/lens_worker_tests.yaml +++ b/helm/litellm/tests/lens_worker_tests.yaml @@ -77,13 +77,20 @@ tests: asserts: - hasDocuments: count: 0 - - it: requires a shared service secret when enabled + - it: uses the managed service secret when none is supplied template: lens/deployment.yaml set: lensWorker.enabled: true + lensWorker.publicUrl: https://traces.example asserts: - - failedTemplate: - errorMessage: lensWorker.serviceTokenSecret.name is required + - contains: + path: spec.template.spec.containers[0].env + content: + name: LITELLM_LENS_SERVICE_TOKEN + valueFrom: + secretKeyRef: + name: RELEASE-NAME-litellm-lens-service + key: service-token - it: uses the chart release and a secret without granting Kubernetes access template: lens/deployment.yaml chart: diff --git a/helm/litellm/values.yaml b/helm/litellm/values.yaml index 3fd3245d166..5851c5a6cb1 100644 --- a/helm/litellm/values.yaml +++ b/helm/litellm/values.yaml @@ -644,6 +644,17 @@ lensWorker: serviceTokenSecret: name: "" key: service-token + clickhouse: + enabled: true + image: clickhouse/clickhouse-server:26.9.6.6 + storage: 20Gi + storageClassName: null + resources: + requests: + cpu: 100m + memory: 512Mi + limits: + memory: 2Gi clickhouseDatabase: litellm retentionDays: 14 clickhouseSecret: diff --git a/litellm-rust/crates/lens/Cargo.toml b/litellm-rust/crates/lens/Cargo.toml index d513e979f97..620d37cef7c 100644 --- a/litellm-rust/crates/lens/Cargo.toml +++ b/litellm-rust/crates/lens/Cargo.toml @@ -42,5 +42,6 @@ prettyplease = "0.2" [dev-dependencies] rstest.workspace = true +tokio = { workspace = true, features = ["test-util"] } wiremock.workspace = true uuid.workspace = true diff --git a/litellm-rust/crates/lens/src/ingest.rs b/litellm-rust/crates/lens/src/ingest.rs index a490e9ee31e..1e890f87ec2 100644 --- a/litellm-rust/crates/lens/src/ingest.rs +++ b/litellm-rust/crates/lens/src/ingest.rs @@ -46,6 +46,15 @@ pub fn response(content_type: Option<&str>, outcome: Result<(), Error>) -> Respo .map(|_| StatusCode::OK) .unwrap_or_else(|error| error.status()); let message = status.canonical_reason().unwrap_or("Trace request failed"); + let rpc_code = match status { + StatusCode::OK => 0, + StatusCode::BAD_REQUEST => 3, + StatusCode::UNAUTHORIZED => 16, + StatusCode::PAYLOAD_TOO_LARGE | StatusCode::TOO_MANY_REQUESTS => 8, + StatusCode::CONFLICT => 10, + StatusCode::SERVICE_UNAVAILABLE => 14, + _ => 2, + }; let protobuf = content_type.is_some_and(|value| { value .split(';') @@ -58,7 +67,7 @@ pub fn response(content_type: Option<&str>, outcome: Result<(), Error>) -> Respo Vec::new() } else { OtlpError { - code: 0, + code: rpc_code, message: message.into(), } .encode_to_vec() @@ -70,7 +79,7 @@ pub fn response(content_type: Option<&str>, outcome: Result<(), Error>) -> Respo if outcome.is_ok() { b"{}".to_vec() } else { - serde_json::json!({"code": 0, "message": message}) + serde_json::json!({"code": rpc_code, "message": message}) .to_string() .into_bytes() }, @@ -169,3 +178,67 @@ async fn store( .await?; Ok(()) } + +#[cfg(test)] +mod tests { + use super::{OtlpError, response}; + use crate::Error; + use axum::body::to_bytes; + use prost::Message; + use rstest::rstest; + + #[rstest] + #[case::invalid_request(Error::InvalidRequest, 400, 3, false)] + #[case::unauthenticated(Error::Unauthorized, 401, 16, false)] + #[case::payload_too_large(Error::TooLarge, 413, 8, false)] + #[case::credentials_pending(Error::CredentialsPending, 429, 8, true)] + #[case::storage_unavailable(Error::Unavailable, 503, 14, true)] + #[case::conflict(Error::TraceChanged, 409, 10, false)] + #[tokio::test] + async fn rejected_batches_have_matching_http_and_rpc_errors( + #[case] error: Error, + #[case] http_status: u16, + #[case] rpc_code: i32, + #[case] retryable: bool, + #[values("application/json", "application/x-protobuf")] content_type: &str, + ) { + let reply = response(Some(content_type), Err(error)); + assert_eq!(reply.status().as_u16(), http_status); + assert_eq!(reply.headers()["content-type"], content_type); + assert_eq!( + reply + .headers() + .get("retry-after") + .map(|v| v.to_str().unwrap()), + retryable.then_some("5") + ); + let message = reply.status().canonical_reason().unwrap(); + let body = to_bytes(reply.into_body(), 1024).await.unwrap(); + if content_type == "application/x-protobuf" { + let status = OtlpError::decode(body).unwrap(); + assert_eq!(status.code, rpc_code); + assert_eq!(status.message, message); + } else { + let status: serde_json::Value = serde_json::from_slice(&body).unwrap(); + assert_eq!( + status, + serde_json::json!({"code": rpc_code, "message": message}) + ); + } + } + + #[rstest] + #[case::json("application/json", b"{}")] + #[case::protobuf("application/x-protobuf", b"")] + #[tokio::test] + async fn accepted_batches_keep_the_empty_export_response( + #[case] content_type: &str, + #[case] expected: &[u8], + ) { + let reply = response(Some(content_type), Ok(())); + assert_eq!(reply.status(), 200); + assert_eq!(reply.headers()["content-type"], content_type); + assert!(!reply.headers().contains_key("retry-after")); + assert_eq!(to_bytes(reply.into_body(), 1024).await.unwrap(), expected); + } +} diff --git a/litellm-rust/crates/lens/src/lib.rs b/litellm-rust/crates/lens/src/lib.rs index 25830fda896..a3038a0f975 100644 --- a/litellm-rust/crates/lens/src/lib.rs +++ b/litellm-rust/crates/lens/src/lib.rs @@ -26,6 +26,7 @@ use litellm_traces_clickhouse::InsertTable; use serde_json::Value; use std::{ collections::BTreeMap, + future::Future, sync::{ Arc, atomic::{AtomicBool, Ordering}, @@ -33,6 +34,7 @@ use std::{ time::Duration, }; pub use storage::Storage; +use tokio::sync::Semaphore; #[allow( dead_code, @@ -46,7 +48,8 @@ pub use storage::Storage; pub mod wire { include!(concat!(env!("OUT_DIR"), "/wire.rs")); } -use tokio::sync::Semaphore; + +const READ_QUEUE_WAIT: Duration = Duration::from_secs(10); pub struct State { pub credentials: Arc, @@ -80,6 +83,15 @@ impl State { } } +async fn wait_for_read_slot

( + acquire: impl Future>, +) -> Result { + tokio::time::timeout(READ_QUEUE_WAIT, acquire) + .await + .map_err(|_| Error::Unavailable)? + .map_err(|_| Error::Unavailable) +} + pub fn router(state: Arc) -> Router { let public = Router::new() .route("/health/live", get(|| async { StatusCode::OK })) @@ -104,6 +116,7 @@ pub fn router(state: Arc) -> Router { Router::new() .route("/internal/read", post(read)) .route("/internal/spend", post(spend)) + .route("/internal/feedback", post(feedback)) .route("/internal/credentials", post(credentials)) .route("/internal/status", get(status)), ) @@ -125,10 +138,7 @@ async fn receipt( ) -> Result, Error> { let tenant = state.credentials.tenant(&headers)?; state.require_storage()?; - let _permit = state - .read_slots - .try_acquire() - .map_err(|_| Error::Unavailable)?; + let _permit = wait_for_read_slot(state.read_slots.acquire()).await?; let body = tokio::time::timeout(Duration::from_secs(5), to_bytes(body, 64 * 1024)) .await .map_err(|_| Error::Unavailable)? @@ -206,11 +216,7 @@ async fn read( ) -> Result, Error> { auth::authorize_service(&headers, &state.service_token)?; state.require_storage()?; - let permit = state - .read_slots - .clone() - .try_acquire_owned() - .map_err(|_| Error::Unavailable)?; + let permit = wait_for_read_slot(state.read_slots.clone().acquire_owned()).await?; let body = tokio::time::timeout(Duration::from_secs(10), to_bytes(body, 1024 * 1024)) .await .map_err(|_| Error::Unavailable)? @@ -228,6 +234,23 @@ async fn spend( AppState(state): AppState>, headers: HeaderMap, body: Body, +) -> Result { + insert(state, headers, body, InsertTable::SpendLogs).await +} + +async fn feedback( + AppState(state): AppState>, + headers: HeaderMap, + body: Body, +) -> Result { + insert(state, headers, body, InsertTable::LensFeedback).await +} + +async fn insert( + state: Arc, + headers: HeaderMap, + body: Body, + table: InsertTable, ) -> Result { auth::authorize_service(&headers, &state.service_token)?; state.require_storage()?; @@ -251,7 +274,7 @@ async fn spend( &state.storage.client, state.storage.config.storage().writer(), state.storage.config.storage().database(), - InsertTable::SpendLogs, + table, rows, ) .await?; @@ -279,3 +302,39 @@ pub async fn provision(state: Arc) { tokio::time::sleep(Duration::from_secs(10)).await; } } + +#[cfg(test)] +mod tests { + use super::{Error, READ_QUEUE_WAIT, wait_for_read_slot}; + use std::sync::Arc; + use tokio::sync::Semaphore; + + #[tokio::test] + async fn ninth_read_waits_for_a_permit_and_succeeds() { + let slots = Arc::new(Semaphore::new(8)); + let permits = (0..8) + .map(|_| slots.clone().try_acquire_owned().expect("available permit")) + .collect::>(); + let waiting_slots = slots.clone(); + let waiting = + tokio::spawn(async move { wait_for_read_slot(waiting_slots.acquire_owned()).await }); + + tokio::task::yield_now().await; + assert!(!waiting.is_finished()); + drop(permits); + assert!(waiting.await.expect("joined read").is_ok()); + } + + #[tokio::test(start_paused = true)] + async fn read_queue_timeout_returns_unavailable() { + let slots = Arc::new(Semaphore::new(0)); + let waiting = tokio::spawn(wait_for_read_slot(slots.acquire_owned())); + + tokio::task::yield_now().await; + tokio::time::advance(READ_QUEUE_WAIT).await; + assert!(matches!( + waiting.await.expect("joined read"), + Err(Error::Unavailable) + )); + } +} diff --git a/litellm-rust/crates/lens/tests/receiver.rs b/litellm-rust/crates/lens/tests/receiver.rs index 1679fc2540d..bda1cecbbb3 100644 --- a/litellm-rust/crates/lens/tests/receiver.rs +++ b/litellm-rust/crates/lens/tests/receiver.rs @@ -116,6 +116,74 @@ async fn agent_picker_query_preserves_scope_through_the_internal_read_route() { assert_eq!(response.json::().await.unwrap(), result); } +#[rstest] +#[tokio::test] +async fn feedback_summary_query_preserves_scope_through_the_internal_read_route() { + let store = MockServer::start().await; + let result = json!({"data": [{ + "trace_id": "1234567890abcdef1234567890abcdef", "trace_ref": "REF", + "count": "2", "average": 5.5, "lowest": "2" + }]}); + Mock::given(method("POST")) + .and(body_string_contains("FROM lens_feedback FINAL")) + .and(query_param("param_all_teams", "0")) + .and(query_param("param_team", "feedback-team")) + .and(query_param( + "param_trace_ids", + "['1234567890abcdef1234567890abcdef']", + )) + .respond_with(ResponseTemplate::new(200).set_body_json(&result)) + .expect(1) + .mount(&store) + .await; + let server = serve(&store.uri(), true).await; + let response = http_client() + .unwrap() + .post(format!("{}/internal/read", server.url)) + .bearer_auth(SERVICE_TOKEN) + .json(&json!({ + "operation": "query", "name": "feedback_summary", "parameters": { + "all_teams": 0, "team": "feedback-team", "key_hash": "", + "trace_ids": ["1234567890abcdef1234567890abcdef"] + } + })) + .send() + .await + .unwrap(); + assert_eq!(response.status(), 200); + assert_eq!(response.json::().await.unwrap(), result); +} + +#[rstest] +#[tokio::test] +async fn feedback_rows_are_written_to_the_feedback_table() { + let store = MockServer::start().await; + Mock::given(method("POST")) + .and(query_param( + "query", + "INSERT INTO `litellm`.lens_feedback FORMAT JSONEachRow", + )) + .respond_with(ResponseTemplate::new(200)) + .expect(1) + .mount(&store) + .await; + let server = serve(&store.uri(), true).await; + let response = http_client() + .unwrap() + .post(format!("{}/internal/feedback", server.url)) + .bearer_auth(SERVICE_TOKEN) + .json(&json!([{ + "TeamId": "team", "ApiKeyHash": "", "TraceId": "1234567890abcdef1234567890abcdef", + "Author": "customer-1042", "Score": 2, "Comment": "wrong command", + "CreatedAt": "2026-10-07T21:57:01.414Z", "UpdatedAt": "2026-10-07T21:57:01.414Z", + "IsDeleted": 0 + }])) + .send() + .await + .unwrap(); + assert_eq!(response.status(), 204); +} + #[rstest] #[tokio::test] async fn ingestion_confirms_storage_and_overwrites_exporter_tenant() { @@ -297,7 +365,7 @@ async fn ingestion_key_cannot_read_or_export_gateway_records() { let store = MockServer::start().await; let server = serve(&store.uri(), true).await; let client = http_client().unwrap(); - for path in ["/internal/read", "/internal/spend"] { + for path in ["/internal/read", "/internal/spend", "/internal/feedback"] { let response = client .post(format!("{}{path}", server.url)) .bearer_auth(KEY) diff --git a/litellm-rust/crates/traces-cache/src/cache.rs b/litellm-rust/crates/traces-cache/src/cache.rs index 287160df96b..eda054e6c7d 100644 --- a/litellm-rust/crates/traces-cache/src/cache.rs +++ b/litellm-rust/crates/traces-cache/src/cache.rs @@ -328,6 +328,7 @@ mod tests { source_type: String::new(), source_url: String::new(), source_title: String::new(), + source_user: String::new(), team_id: String::new(), api_key_hash: String::new(), user_id: String::new(), diff --git a/litellm-rust/crates/traces-cache/tests/read.rs b/litellm-rust/crates/traces-cache/tests/read.rs index ade70b71800..22293db98e8 100644 --- a/litellm-rust/crates/traces-cache/tests/read.rs +++ b/litellm-rust/crates/traces-cache/tests/read.rs @@ -289,6 +289,7 @@ fn span(index: usize) -> TraceSpansRow { source_type: String::new(), source_url: String::new(), source_title: String::new(), + source_user: String::new(), team_id: "team".into(), api_key_hash: "key".into(), user_id: "user".into(), diff --git a/litellm-rust/crates/traces-cache/tests/snapshots.rs b/litellm-rust/crates/traces-cache/tests/snapshots.rs index bebff0bb1aa..4ad4e2b7545 100644 --- a/litellm-rust/crates/traces-cache/tests/snapshots.rs +++ b/litellm-rust/crates/traces-cache/tests/snapshots.rs @@ -42,6 +42,7 @@ fn row(span_id: &str, parent: &str, name: &str, kind: &str, agent: &str) -> Trac source_type: String::new(), source_url: String::new(), source_title: String::new(), + source_user: String::new(), team_id: String::new(), api_key_hash: String::new(), user_id: String::new(), diff --git a/litellm-rust/crates/traces-clickhouse/migrations/0017_lens_feedback.sql b/litellm-rust/crates/traces-clickhouse/migrations/0017_lens_feedback.sql new file mode 100644 index 00000000000..9cef629e223 --- /dev/null +++ b/litellm-rust/crates/traces-clickhouse/migrations/0017_lens_feedback.sql @@ -0,0 +1,18 @@ +CREATE TABLE IF NOT EXISTS {database}.lens_feedback +( + TeamId LowCardinality(String), + ApiKeyHash String, + TraceId String CODEC(ZSTD(1)), + Author String, + Score UInt8, + Comment String CODEC(ZSTD(3)), + CreatedAt DateTime64(3), + UpdatedAt DateTime64(3), + IsDeleted UInt8, + EngineReceivedMs UInt64 DEFAULT 0, + INDEX idx_trace_id TraceId TYPE bloom_filter(0.001) GRANULARITY 1, + CONSTRAINT score_range CHECK Score <= 10 +) +ENGINE = ReplacingMergeTree(UpdatedAt, IsDeleted) +ORDER BY (TeamId, ApiKeyHash, TraceId, Author) +SETTINGS materialize_ttl_recalculate_only = 1, non_replicated_deduplication_window = 1000 diff --git a/litellm-rust/crates/traces-clickhouse/query/lens_feedback.sql b/litellm-rust/crates/traces-clickhouse/query/lens_feedback.sql new file mode 100644 index 00000000000..11540af6ddf --- /dev/null +++ b/litellm-rust/crates/traces-clickhouse/query/lens_feedback.sql @@ -0,0 +1,12 @@ +SELECT TraceId AS trace_id, + hex(SHA256(concat(TeamId, char(0), ApiKeyHash, char(0), TraceId))) AS trace_ref, + Author AS author, Score AS score, Comment AS comment, + formatDateTime(CreatedAt, '%FT%T.%fZ', 'UTC') AS created_at, + formatDateTime(UpdatedAt, '%FT%T.%fZ', 'UTC') AS updated_at +FROM lens_feedback FINAL +WHERE TraceId = {trace_id:String} + AND hex(SHA256(concat(TeamId, char(0), ApiKeyHash, char(0), TraceId))) = {trace_ref:String} + AND ({all_teams:UInt8}=1 OR TeamId={team:String}) + AND ({key_hash:String}='' OR ApiKeyHash={key_hash:String}) + AND IsDeleted = 0 +ORDER BY UpdatedAt DESC, Author diff --git a/litellm-rust/crates/traces-clickhouse/query/lens_feedback_summary.sql b/litellm-rust/crates/traces-clickhouse/query/lens_feedback_summary.sql new file mode 100644 index 00000000000..8786ecdfa42 --- /dev/null +++ b/litellm-rust/crates/traces-clickhouse/query/lens_feedback_summary.sql @@ -0,0 +1,9 @@ +SELECT TraceId AS trace_id, + hex(SHA256(concat(TeamId, char(0), ApiKeyHash, char(0), TraceId))) AS trace_ref, + count() AS count, avg(Score) AS average, min(Score) AS lowest +FROM lens_feedback FINAL +WHERE TraceId IN {trace_ids:Array(String)} + AND ({all_teams:UInt8}=1 OR TeamId={team:String}) + AND ({key_hash:String}='' OR ApiKeyHash={key_hash:String}) + AND IsDeleted = 0 +GROUP BY TeamId, ApiKeyHash, TraceId diff --git a/litellm-rust/crates/traces-clickhouse/query/lens_feedback_target.sql b/litellm-rust/crates/traces-clickhouse/query/lens_feedback_target.sql new file mode 100644 index 00000000000..e2b7f96d90b --- /dev/null +++ b/litellm-rust/crates/traces-clickhouse/query/lens_feedback_target.sql @@ -0,0 +1,9 @@ +SELECT TeamId AS team_id, ApiKeyHash AS key_hash, + hex(SHA256(concat(TeamId, char(0), ApiKeyHash, char(0), TraceId))) AS trace_ref +FROM agent_traces_by_key +WHERE TraceId = {trace_id:String} + AND ({all_teams:UInt8}=1 OR TeamId={team:String}) + AND ({key_hash:String}='' OR ApiKeyHash={key_hash:String}) + AND ({trace_ref:String}='' OR hex(SHA256(concat(TeamId, char(0), ApiKeyHash, char(0), TraceId)))={trace_ref:String}) +GROUP BY TeamId, ApiKeyHash, TraceId +LIMIT 2 diff --git a/litellm-rust/crates/traces-clickhouse/query/list_traces.sql b/litellm-rust/crates/traces-clickhouse/query/list_traces.sql index c52adf7ef49..fa9a67c1e20 100644 --- a/litellm-rust/crates/traces-clickhouse/query/list_traces.sql +++ b/litellm-rust/crates/traces-clickhouse/query/list_traces.sql @@ -5,7 +5,6 @@ SELECT TraceId AS trace_id, ifNull(any(RootName), '') AS name, any(ServiceName) AS service, ifNull(any(RootInput), '') AS input_preview, ifNull(any(RootStatus), '') AS status, toUnixTimestamp64Milli(min(StartTs)) AS start_ms, - min(StartTs) AS trace_start, max(EndTs) AS trace_end, dateDiff('millisecond', min(StartTs), max(EndTs)) AS duration_ms, sum(SpanCount) AS span_count, sum(AgentCount) AS agent_invocations, @@ -27,7 +26,7 @@ HAVING min(StartTs) >= fromUnixTimestamp64Milli({start_ms:Int64}) ORDER BY start_ms DESC, trace_ref DESC LIMIT {limit:UInt32} ) -SELECT page.* EXCEPT (trace_start, trace_end), +SELECT page.*, identities.agent_names AS agent_names, identities.agent_count AS agent_count, identities.frameworks AS frameworks FROM page @@ -37,10 +36,8 @@ LEFT JOIN ( arraySort(groupUniqArrayIf(toString(Framework), Framework != '')) AS frameworks, uniqExactIf(if(AgentName = '', SpanName, AgentName), ObservationType = 'agent') AS agent_count FROM otel_traces - WHERE Timestamp >= (SELECT min(trace_start) FROM page) - AND Timestamp <= (SELECT max(trace_end) FROM page) + WHERE Timestamp >= fromUnixTimestamp64Milli({start_ms:Int64}) AND TraceId IN (SELECT trace_id FROM page) - AND (TeamId, ApiKeyHash, TraceId) IN (SELECT team_id, api_key_hash, trace_id FROM page) GROUP BY TeamId, ApiKeyHash, TraceId ) AS identities ON page.team_id = identities.TeamId AND page.api_key_hash = identities.ApiKeyHash diff --git a/litellm-rust/crates/traces-clickhouse/query/trace_list_span_batch.sql b/litellm-rust/crates/traces-clickhouse/query/trace_list_span_batch.sql index b764005024e..327483d7a45 100644 --- a/litellm-rust/crates/traces-clickhouse/query/trace_list_span_batch.sql +++ b/litellm-rust/crates/traces-clickhouse/query/trace_list_span_batch.sql @@ -14,7 +14,7 @@ SELECT o.TraceId AS trace_id, o.SpanAttributes['lens.original_trace_id'] AS orig coalesce(nullIf(o.SpanAttributes['gen_ai.tool.call.id'], ''), nullIf(o.SpanAttributes['tool.id'], ''), '')) AS tool_call_id, o.SpanAttributes['agent.source.type'] AS source_type, o.SpanAttributes['agent.source.url'] AS source_url, - o.SpanAttributes['agent.source.title'] AS source_title, + o.SpanAttributes['agent.source.title'] AS source_title, o.SpanAttributes['agent.source.user'] AS source_user, o.UserId AS user_id, o.TeamId AS team_id, o.ApiKeyHash AS api_key_hash FROM otel_traces AS o WHERE o.Timestamp >= fromUnixTimestamp64Milli({start_ms:Int64}) diff --git a/litellm-rust/crates/traces-clickhouse/query/trace_page_spans.sql b/litellm-rust/crates/traces-clickhouse/query/trace_page_spans.sql index 5af30920df9..dd8cdee7e80 100644 --- a/litellm-rust/crates/traces-clickhouse/query/trace_page_spans.sql +++ b/litellm-rust/crates/traces-clickhouse/query/trace_page_spans.sql @@ -13,7 +13,7 @@ SELECT o.TraceId AS trace_id, o.SpanAttributes['lens.original_trace_id'] AS orig coalesce(nullIf(o.SpanAttributes['gen_ai.tool.call.id'], ''), nullIf(o.SpanAttributes['tool.id'], ''), '')) AS tool_call_id, o.SpanAttributes['agent.source.type'] AS source_type, o.SpanAttributes['agent.source.url'] AS source_url, - o.SpanAttributes['agent.source.title'] AS source_title, + o.SpanAttributes['agent.source.title'] AS source_title, o.SpanAttributes['agent.source.user'] AS source_user, o.UserId AS user_id, o.TeamId AS team_id, o.ApiKeyHash AS api_key_hash FROM otel_traces AS o WHERE o.Timestamp >= fromUnixTimestamp64Milli({start_ms:Int64}) diff --git a/litellm-rust/crates/traces-clickhouse/query/trace_span_batch.sql b/litellm-rust/crates/traces-clickhouse/query/trace_span_batch.sql index c7d50a44544..85345cb05c1 100644 --- a/litellm-rust/crates/traces-clickhouse/query/trace_span_batch.sql +++ b/litellm-rust/crates/traces-clickhouse/query/trace_span_batch.sql @@ -14,7 +14,7 @@ SELECT o.TraceId AS trace_id, o.SpanAttributes['lens.original_trace_id'] AS orig coalesce(nullIf(o.SpanAttributes['gen_ai.tool.call.id'], ''), nullIf(o.SpanAttributes['tool.id'], ''), '')) AS tool_call_id, o.SpanAttributes['agent.source.type'] AS source_type, o.SpanAttributes['agent.source.url'] AS source_url, - o.SpanAttributes['agent.source.title'] AS source_title, + o.SpanAttributes['agent.source.title'] AS source_title, o.SpanAttributes['agent.source.user'] AS source_user, o.UserId AS user_id, o.TeamId AS team_id, o.ApiKeyHash AS api_key_hash FROM otel_traces AS o WHERE o.TraceId = {trace_id:String} diff --git a/litellm-rust/crates/traces-clickhouse/query/trace_spans.sql b/litellm-rust/crates/traces-clickhouse/query/trace_spans.sql index 974a6d050a4..a2327e57660 100644 --- a/litellm-rust/crates/traces-clickhouse/query/trace_spans.sql +++ b/litellm-rust/crates/traces-clickhouse/query/trace_spans.sql @@ -13,7 +13,7 @@ SELECT o.TraceId AS trace_id, o.SpanAttributes['lens.original_trace_id'] AS orig coalesce(nullIf(o.SpanAttributes['gen_ai.tool.call.id'], ''), nullIf(o.SpanAttributes['tool.id'], ''), '')) AS tool_call_id, o.SpanAttributes['agent.source.type'] AS source_type, o.SpanAttributes['agent.source.url'] AS source_url, - o.SpanAttributes['agent.source.title'] AS source_title, + o.SpanAttributes['agent.source.title'] AS source_title, o.SpanAttributes['agent.source.user'] AS source_user, o.UserId AS user_id, o.TeamId AS team_id, o.ApiKeyHash AS api_key_hash FROM otel_traces AS o WHERE o.TraceId = {trace_id:String} diff --git a/litellm-rust/crates/traces-clickhouse/src/insert.rs b/litellm-rust/crates/traces-clickhouse/src/insert.rs index d45cc8d53b8..ed8db188c14 100644 --- a/litellm-rust/crates/traces-clickhouse/src/insert.rs +++ b/litellm-rust/crates/traces-clickhouse/src/insert.rs @@ -33,6 +33,7 @@ pub type InsertRow = BTreeMap>; pub enum InsertTable { OtelTraces, SpendLogs, + LensFeedback, } impl InsertTable { @@ -40,6 +41,7 @@ impl InsertTable { match value { "otel_traces" => Ok(Self::OtelTraces), "spend_logs" => Ok(Self::SpendLogs), + "lens_feedback" => Ok(Self::LensFeedback), _ => Err(Error::InvalidTable), } } @@ -48,6 +50,7 @@ impl InsertTable { match self { Self::OtelTraces => "otel_traces", Self::SpendLogs => "spend_logs", + Self::LensFeedback => "lens_feedback", } } } diff --git a/litellm-rust/crates/traces-clickhouse/src/query/lens.rs b/litellm-rust/crates/traces-clickhouse/src/query/lens.rs index fcca4fe6f96..97b74de4c20 100644 --- a/litellm-rust/crates/traces-clickhouse/src/query/lens.rs +++ b/litellm-rust/crates/traces-clickhouse/src/query/lens.rs @@ -6,13 +6,16 @@ const SAMPLE_READ_LIMITS: ReadLimits = ReadLimits { ..litellm_storage_clickhouse::READ_LIMITS }; -pub const LENS_QUERIES: [litellm_traces::ReadQuery; 6] = [ +pub const LENS_QUERIES: [litellm_traces::ReadQuery; 9] = [ litellm_traces::ReadQuery::TraceAgents, litellm_traces::ReadQuery::Availability, litellm_traces::ReadQuery::Agents, litellm_traces::ReadQuery::Sample, litellm_traces::ReadQuery::Content, litellm_traces::ReadQuery::Evidence, + litellm_traces::ReadQuery::FeedbackTarget, + litellm_traces::ReadQuery::Feedback, + litellm_traces::ReadQuery::FeedbackSummary, ]; #[macro_rules_attribute::apply(wire_type)] @@ -339,3 +342,108 @@ impl Query for LensEvidence { const SQL: &'static str = include_str!("../../query/lens_evidence.sql"); } + +pub struct LensFeedbackTarget; + +#[macro_rules_attribute::apply(wire_type)] +#[derive(Debug)] +#[serde(deny_unknown_fields)] +pub struct LensFeedbackTargetParams { + #[serde(flatten)] + pub access: LensAccessParams, + pub trace_id: String, + pub trace_ref: String, +} + +#[macro_rules_attribute::apply(wire_type)] +#[derive(Debug)] +#[cfg_attr(feature = "schema", schemars(rename = "FeedbackTargetRow"))] +pub struct LensFeedbackTargetRow { + pub team_id: String, + pub key_hash: String, + pub trace_ref: String, +} + +impl Query for LensFeedbackTarget { + type Params = LensFeedbackTargetParams; + type Row = LensFeedbackTargetRow; + + const SQL: &'static str = include_str!("../../query/lens_feedback_target.sql"); +} + +pub struct LensFeedback; + +#[macro_rules_attribute::apply(wire_type)] +#[derive(Debug)] +#[serde(deny_unknown_fields)] +pub struct LensFeedbackParams { + #[serde(flatten)] + pub access: LensAccessParams, + pub trace_id: String, + pub trace_ref: String, +} + +#[macro_rules_attribute::apply(wire_type)] +#[derive(Debug)] +#[cfg_attr(feature = "schema", schemars(rename = "FeedbackRow"))] +pub struct LensFeedbackRow { + pub trace_id: String, + pub trace_ref: String, + pub author: String, + #[serde(deserialize_with = "super::number::deserialize")] + #[cfg_attr( + feature = "schema", + schemars(schema_with = "crate::wire_schema::u64_number") + )] + pub score: u64, + pub comment: String, + pub created_at: String, + pub updated_at: String, +} + +impl Query for LensFeedback { + type Params = LensFeedbackParams; + type Row = LensFeedbackRow; + + const SQL: &'static str = include_str!("../../query/lens_feedback.sql"); +} + +pub struct LensFeedbackSummary; + +#[macro_rules_attribute::apply(wire_type)] +#[derive(Debug)] +#[serde(deny_unknown_fields)] +pub struct LensFeedbackSummaryParams { + #[serde(flatten)] + pub access: LensAccessParams, + pub trace_ids: Vec, +} + +#[macro_rules_attribute::apply(wire_type)] +#[derive(Debug)] +#[cfg_attr(feature = "schema", schemars(rename = "FeedbackSummaryRow"))] +pub struct LensFeedbackSummaryRow { + pub trace_id: String, + pub trace_ref: String, + #[serde(deserialize_with = "super::number::deserialize")] + #[cfg_attr( + feature = "schema", + schemars(schema_with = "crate::wire_schema::u64_number") + )] + pub count: u64, + #[serde(deserialize_with = "super::number::deserialize")] + pub average: f64, + #[serde(deserialize_with = "super::number::deserialize")] + #[cfg_attr( + feature = "schema", + schemars(schema_with = "crate::wire_schema::u64_number") + )] + pub lowest: u64, +} + +impl Query for LensFeedbackSummary { + type Params = LensFeedbackSummaryParams; + type Row = LensFeedbackSummaryRow; + + const SQL: &'static str = include_str!("../../query/lens_feedback_summary.sql"); +} diff --git a/litellm-rust/crates/traces-clickhouse/src/query/named.rs b/litellm-rust/crates/traces-clickhouse/src/query/named.rs index 1b98ad39912..b7de28646e4 100644 --- a/litellm-rust/crates/traces-clickhouse/src/query/named.rs +++ b/litellm-rust/crates/traces-clickhouse/src/query/named.rs @@ -134,6 +134,8 @@ struct TraceSpansRowEncoding { pub source_url: String, #[serde(default)] pub source_title: String, + #[serde(default)] + pub source_user: String, pub team_id: String, pub api_key_hash: String, pub user_id: String, @@ -355,7 +357,7 @@ mod tests { quoted, ); round_trip::( - json!({"trace_id": "trace", "original_trace_id": "original", "span_id": "span", "parent_span_id": "parent", "name": "agent", "type": "agent", "wrapper_candidate": 1, "agent": "agent", "framework": "claude-agent-sdk", "status": "STATUS_CODE_ERROR", "status_message": "error", "error_truncated": 1, "start_ns": -1, "duration_ns": u64::MAX, "service": "service", "input_preview": "input", "model": "model", "input_tokens": u32::MAX, "output_tokens": 6, "litellm_request_id": "request", "call_keys": ["provider_response:request"], "call_evidence": "complete", "tool_call_id": "call", "source_type": "slack", "source_url": "https://acme.slack.com/archives/C1/p1", "source_title": "thread", "team_id": "team", "api_key_hash": "key", "user_id": "user"}), + json!({"trace_id": "trace", "original_trace_id": "original", "span_id": "span", "parent_span_id": "parent", "name": "agent", "type": "agent", "wrapper_candidate": 1, "agent": "agent", "framework": "claude-agent-sdk", "status": "STATUS_CODE_ERROR", "status_message": "error", "error_truncated": 1, "start_ns": -1, "duration_ns": u64::MAX, "service": "service", "input_preview": "input", "model": "model", "input_tokens": u32::MAX, "output_tokens": 6, "litellm_request_id": "request", "call_keys": ["provider_response:request"], "call_evidence": "complete", "tool_call_id": "call", "source_type": "slack", "source_url": "https://acme.slack.com/archives/C1/p1", "source_title": "thread", "source_user": "tin@berri.ai", "team_id": "team", "api_key_hash": "key", "user_id": "user"}), quoted, ); round_trip::( diff --git a/litellm-rust/crates/traces-clickhouse/src/schema.rs b/litellm-rust/crates/traces-clickhouse/src/schema.rs index 562dbb976c2..5eb8bf0b973 100644 --- a/litellm-rust/crates/traces-clickhouse/src/schema.rs +++ b/litellm-rust/crates/traces-clickhouse/src/schema.rs @@ -14,10 +14,11 @@ static MIGRATOR: Migrator = Migrator { ..sqlx::migrate!("./migrations") }; -const RETENTION: [(&str, &str); 3] = [ +const RETENTION: [(&str, &str); 4] = [ ("otel_traces", "toDateTime(Timestamp)"), ("agent_traces_by_key", "toDateTime(StartTs)"), ("spend_logs", "toDateTime(start_time)"), + ("lens_feedback", "toDateTime(CreatedAt)"), ]; fn validate_schema(database: &str, retention_days: u32) -> Result<(), Error> { diff --git a/litellm-rust/crates/traces-clickhouse/src/span_batches.rs b/litellm-rust/crates/traces-clickhouse/src/span_batches.rs index 5e406d67cfa..a08ede7a43f 100644 --- a/litellm-rust/crates/traces-clickhouse/src/span_batches.rs +++ b/litellm-rust/crates/traces-clickhouse/src/span_batches.rs @@ -4,7 +4,7 @@ use std::{future::Future, marker::PhantomData}; use litellm_http::Client; -use litellm_storage_clickhouse::{Query, fetch}; +use litellm_storage_clickhouse::{Query, ReadLimits, fetch}; use litellm_traces::query::named as contracts; use litellm_traces_cache::{MAX_GRAPH_BYTES, MAX_GRAPH_SPANS, StoreError}; use serde::{Serialize, de::DeserializeOwned}; @@ -14,7 +14,12 @@ use crate::{ query::named::{SpendByResponseIdsParams, SpendByResponseIdsRow, TraceSpansRow}, }; -const PAGE_SIZE: u32 = 256; +const PAGE_SIZE: u32 = 8192; +const SPAN_BATCH_READ_LIMITS: ReadLimits = ReadLimits { + result_rows: PAGE_SIZE as u64, + response_bytes: 16 * 1024 * 1024, + ..litellm_storage_clickhouse::READ_LIMITS +}; #[derive(Default)] struct ReadBudget { @@ -67,6 +72,7 @@ struct Paged(PhantomData); impl Query for Paged { type Params = Batch; type Row = K::Row; + const READ_LIMITS: ReadLimits = SPAN_BATCH_READ_LIMITS; const SQL: &'static str = K::SQL; } @@ -316,8 +322,8 @@ mod tests { } #[rstest] - #[case::fits(1000, PAGE_SIZE, &[256, 256, 256, 256])] - #[case::uniform_large_rows(1000, 100, &[256, 128, 64, 64, 64, 64, 64, 64, 64, 64, 64, 64, 64, 64, 64, 64, 64, 64])] + #[case::fits(1000, PAGE_SIZE, &[8192])] + #[case::uniform_large_rows(200, 100, &[8192, 4096, 2048, 1024, 512, 256, 128, 64, 64, 64, 64])] #[tokio::test] async fn a_rejected_page_size_is_not_retried( #[case] total: u32, @@ -334,6 +340,13 @@ mod tests { assert_eq!(table.requests.lock().unwrap().as_slice(), requests); } + #[test] + fn span_batches_use_larger_read_limits() { + let limits = as Query>::READ_LIMITS; + assert_eq!(limits.result_rows, 8192); + assert_eq!(limits.response_bytes, 16 * 1024 * 1024); + } + #[rstest] #[tokio::test] async fn a_single_oversized_row_fails_the_read() { @@ -346,7 +359,9 @@ mod tests { assert!(matches!(result, Err(StoreError::TooLarge)), "{result:?}"); assert_eq!( table.requests.lock().unwrap().as_slice(), - &[256, 128, 64, 32, 16, 8, 4, 2, 1] + &[ + 8192, 4096, 2048, 1024, 512, 256, 128, 64, 32, 16, 8, 4, 2, 1 + ] ); } } diff --git a/litellm-rust/crates/traces-clickhouse/src/sql.rs b/litellm-rust/crates/traces-clickhouse/src/sql.rs index 1f41f0f6f43..613159cf06d 100644 --- a/litellm-rust/crates/traces-clickhouse/src/sql.rs +++ b/litellm-rust/crates/traces-clickhouse/src/sql.rs @@ -37,6 +37,13 @@ pub async fn execute_named_read( ReadQuery::Sample => named_json::(client, connection, parameters).await, ReadQuery::Content => named_json::(client, connection, parameters).await, ReadQuery::Evidence => named_json::(client, connection, parameters).await, + ReadQuery::FeedbackTarget => { + named_json::(client, connection, parameters).await + } + ReadQuery::Feedback => named_json::(client, connection, parameters).await, + ReadQuery::FeedbackSummary => { + named_json::(client, connection, parameters).await + } } } diff --git a/litellm-rust/crates/traces-clickhouse/src/wire_schema.rs b/litellm-rust/crates/traces-clickhouse/src/wire_schema.rs index f7a52a9edd0..408d224f67b 100644 --- a/litellm-rust/crates/traces-clickhouse/src/wire_schema.rs +++ b/litellm-rust/crates/traces-clickhouse/src/wire_schema.rs @@ -81,6 +81,18 @@ pub fn schemas() -> BTreeMap<&'static str, Schema> { ("LensSampleParams", received::()), ("LensContentParams", received::()), ("LensEvidenceParams", received::()), + ( + "LensFeedbackTargetParams", + received::(), + ), + ("LensFeedbackParams", received::()), + ( + "LensFeedbackSummaryParams", + received::(), + ), + ("FeedbackTargetRow", received::()), + ("FeedbackRow", received::()), + ("FeedbackSummaryRow", received::()), ( "ActivityAvailability", received::(), diff --git a/litellm-rust/crates/traces-clickhouse/tests/lens_feedback.rs b/litellm-rust/crates/traces-clickhouse/tests/lens_feedback.rs new file mode 100644 index 00000000000..b2efea71823 --- /dev/null +++ b/litellm-rust/crates/traces-clickhouse/tests/lens_feedback.rs @@ -0,0 +1,408 @@ +use std::collections::BTreeMap; + +use litellm_traces_clickhouse::{ + Connection, InsertTable, Parameter, ReadQuery, ensure_schema, execute_named_read, execute_read, + insert_rows, +}; +use rstest::rstest; +use serde_json::{Value, json}; +use time::{Duration, OffsetDateTime, format_description::well_known::Rfc3339}; + +mod support; + +use support::{ClickHouseDatabase, TestResult, database}; + +struct Feedback<'a> { + team: &'a str, + key: &'a str, + trace: &'a str, + author: &'a str, + score: u8, + comment: &'a str, + edited_after_seconds: i64, + deleted: bool, +} + +const SAVED: Feedback<'static> = Feedback { + team: "team-a", + key: "key-a", + trace: "trace-1", + author: "alice", + score: 3, + comment: "missed the file", + edited_after_seconds: 0, + deleted: false, +}; + +fn now() -> TestResult { + Ok(OffsetDateTime::now_utc().replace_millisecond(0)?) +} + +fn iso(at: OffsetDateTime) -> TestResult { + Ok(at.format(&Rfc3339)?) +} + +fn row(feedback: &Feedback, created: OffsetDateTime) -> TestResult> { + let updated = created + Duration::seconds(feedback.edited_after_seconds); + Ok(serde_json::from_value(json!({ + "TeamId": feedback.team, "ApiKeyHash": feedback.key, "TraceId": feedback.trace, + "Author": feedback.author, "Score": feedback.score, "Comment": feedback.comment, + "CreatedAt": iso(created)?, "UpdatedAt": iso(updated)?, + "IsDeleted": u8::from(feedback.deleted) + }))?) +} + +async fn ready(database: &ClickHouseDatabase) -> TestResult { + let writer = Connection::writer(&database.url)?; + ensure_schema(&database.client, &writer, "trace_test", 7).await?; + Ok(writer) +} + +async fn save( + database: &ClickHouseDatabase, + writer: &Connection, + created: OffsetDateTime, + feedback: &[Feedback<'_>], +) -> TestResult { + let rows = feedback + .iter() + .map(|entry| row(entry, created)) + .collect::>>()?; + insert_rows( + &database.client, + writer, + "trace_test", + InsertTable::LensFeedback, + rows, + ) + .await?; + Ok(()) +} + +fn access(team: &str) -> BTreeMap { + BTreeMap::from([ + ( + "all_teams".into(), + Parameter::Integer(i64::from(team.is_empty())), + ), + ("team".into(), Parameter::Text(team.into())), + ("key_hash".into(), Parameter::Text(String::new())), + ]) +} + +async fn read( + database: &ClickHouseDatabase, + query: ReadQuery, + parameters: BTreeMap, +) -> TestResult> { + let reader = Connection::configured(&database.url, "trace_test", "default", "")?; + let body: Value = serde_json::from_str( + &execute_named_read(&database.client, &reader, query, ¶meters).await?, + )?; + Ok(body["data"] + .as_array() + .cloned() + .ok_or_else(|| format!("no data in {body}"))?) +} + +async fn trace_ref( + database: &ClickHouseDatabase, + team: &str, + key: &str, + trace: &str, +) -> TestResult { + let reader = Connection::configured(&database.url, "trace_test", "default", "")?; + let sql = format!( + "SELECT hex(SHA256(concat('{team}', char(0), '{key}', char(0), '{trace}'))) AS trace_ref" + ); + let response: Value = serde_json::from_str( + &execute_read(&database.client, &reader, &sql, &BTreeMap::new()).await?, + )?; + Ok(response["data"][0]["trace_ref"] + .as_str() + .unwrap_or_default() + .to_owned()) +} + +async fn feedback( + database: &ClickHouseDatabase, + team: &str, + trace: &str, + trace_ref: &str, +) -> TestResult> { + let mut parameters = access(team); + parameters.insert("trace_id".into(), Parameter::Text(trace.into())); + parameters.insert("trace_ref".into(), Parameter::Text(trace_ref.into())); + read(database, ReadQuery::Feedback, parameters).await +} + +fn timestamp(row: &Value, field: &str) -> TestResult { + Ok(OffsetDateTime::parse( + row[field].as_str().unwrap_or_default(), + &Rfc3339, + )?) +} + +#[rstest] +#[tokio::test] +async fn a_later_save_replaces_the_authors_feedback_and_keeps_other_authors( + #[future(awt)] database: TestResult, +) -> TestResult { + let database = database?; + let writer = ready(&database).await?; + let created = now()?; + save(&database, &writer, created, &[SAVED]).await?; + save( + &database, + &writer, + created, + &[ + Feedback { + score: 8, + comment: "fine after retry", + edited_after_seconds: 300, + ..SAVED + }, + Feedback { + author: "bob", + score: 10, + comment: "", + ..SAVED + }, + ], + ) + .await?; + let reference = trace_ref(&database, "team-a", "key-a", "trace-1").await?; + + let rows = feedback(&database, "", "trace-1", &reference).await?; + + let shown: Vec<(&str, u64, &str)> = rows + .iter() + .map(|row| { + ( + row["author"].as_str().unwrap_or_default(), + row["score"].as_u64().unwrap_or(99), + row["comment"].as_str().unwrap_or_default(), + ) + }) + .collect(); + assert_eq!(shown, [("alice", 8, "fine after retry"), ("bob", 10, "")]); + assert_eq!(timestamp(&rows[0], "created_at")?, created); + assert_eq!( + timestamp(&rows[0], "updated_at")?, + created + Duration::seconds(300) + ); + Ok(()) +} + +#[rstest] +#[tokio::test] +async fn a_deleted_version_hides_the_authors_feedback( + #[future(awt)] database: TestResult, +) -> TestResult { + let database = database?; + let writer = ready(&database).await?; + let created = now()?; + save( + &database, + &writer, + created, + &[ + SAVED, + Feedback { + author: "bob", + score: 6, + ..SAVED + }, + ], + ) + .await?; + save( + &database, + &writer, + created, + &[Feedback { + deleted: true, + edited_after_seconds: 540, + ..SAVED + }], + ) + .await?; + let reference = trace_ref(&database, "team-a", "key-a", "trace-1").await?; + + let authors: Vec = feedback(&database, "", "trace-1", &reference) + .await? + .into_iter() + .map(|row| row["author"].clone()) + .collect(); + + assert_eq!(authors, [json!("bob")]); + Ok(()) +} + +#[rstest] +#[tokio::test] +async fn summaries_aggregate_live_feedback_per_trace_and_respect_team_scope( + #[future(awt)] database: TestResult, +) -> TestResult { + let database = database?; + let writer = ready(&database).await?; + save( + &database, + &writer, + now()?, + &[ + Feedback { score: 2, ..SAVED }, + Feedback { + author: "bob", + score: 7, + ..SAVED + }, + Feedback { + author: "carol", + score: 0, + deleted: true, + ..SAVED + }, + Feedback { + team: "team-b", + key: "key-b", + trace: "trace-2", + score: 9, + ..SAVED + }, + ], + ) + .await?; + let summary = |team: &str| { + let mut parameters = access(team); + parameters.insert( + "trace_ids".into(), + Parameter::Strings(vec![ + "trace-1".into(), + "trace-2".into(), + "trace-unrated".into(), + ]), + ); + parameters + }; + + let everyone = read(&database, ReadQuery::FeedbackSummary, summary("")).await?; + let team_a = read(&database, ReadQuery::FeedbackSummary, summary("team-a")).await?; + + let by_trace: BTreeMap<&str, (u64, f64, u64)> = everyone + .iter() + .map(|row| { + ( + row["trace_id"].as_str().unwrap_or_default(), + ( + row["count"].as_u64().unwrap_or(0), + row["average"].as_f64().unwrap_or(-1.0), + row["lowest"].as_u64().unwrap_or(99), + ), + ) + }) + .collect(); + assert_eq!( + by_trace, + BTreeMap::from([("trace-1", (2, 4.5, 2)), ("trace-2", (1, 9.0, 9))]) + ); + assert_eq!( + team_a + .iter() + .map(|row| row["trace_id"].clone()) + .collect::>(), + [json!("trace-1")] + ); + Ok(()) +} + +#[rstest] +#[tokio::test] +async fn feedback_for_another_teams_trace_is_not_readable( + #[future(awt)] database: TestResult, +) -> TestResult { + let database = database?; + let writer = ready(&database).await?; + save(&database, &writer, now()?, &[SAVED]).await?; + let reference = trace_ref(&database, "team-a", "key-a", "trace-1").await?; + + assert!( + feedback(&database, "team-b", "trace-1", &reference) + .await? + .is_empty() + ); + assert_eq!( + feedback(&database, "team-a", "trace-1", &reference) + .await? + .len(), + 1 + ); + Ok(()) +} + +#[rstest] +#[tokio::test] +async fn the_table_rejects_scores_above_ten( + #[future(awt)] database: TestResult, +) -> TestResult { + let database = database?; + let writer = ready(&database).await?; + + let rejected = save( + &database, + &writer, + now()?, + &[Feedback { score: 11, ..SAVED }], + ) + .await; + + assert!(rejected.is_err()); + Ok(()) +} + +#[rstest] +#[tokio::test] +async fn target_resolves_team_and_key_for_a_visible_trace_only( + #[future(awt)] database: TestResult, +) -> TestResult { + let database = database?; + let writer = ready(&database).await?; + let span: BTreeMap = serde_json::from_value(json!({ + "Timestamp": now()?.unix_timestamp_nanos() as i64, "TraceId": "trace-1", "SpanId": "root", + "ParentSpanId": "", "ServiceName": "agent", "SpanName": "run", + "ResourceAttributes": {"litellm.team_id": "team-a", "litellm.api_key_hash": "key-a"} + }))?; + insert_rows( + &database.client, + &writer, + "trace_test", + InsertTable::OtelTraces, + vec![span], + ) + .await?; + let target = |team: &str| { + let mut parameters = access(team); + parameters.insert("trace_id".into(), Parameter::Text("trace-1".into())); + parameters.insert("trace_ref".into(), Parameter::Text(String::new())); + parameters + }; + + let visible = read(&database, ReadQuery::FeedbackTarget, target("team-a")).await?; + let hidden = read(&database, ReadQuery::FeedbackTarget, target("team-b")).await?; + + assert_eq!(visible.len(), 1); + assert_eq!( + ( + visible[0]["team_id"].as_str(), + visible[0]["key_hash"].as_str() + ), + (Some("team-a"), Some("key-a")) + ); + assert_eq!( + visible[0]["trace_ref"].as_str().map(str::to_owned), + Some(trace_ref(&database, "team-a", "key-a", "trace-1").await?) + ); + assert!(hidden.is_empty()); + Ok(()) +} diff --git a/litellm-rust/crates/traces-clickhouse/tests/migrations.rs b/litellm-rust/crates/traces-clickhouse/tests/migrations.rs index 9a3cd4b3050..a00f34a0002 100644 --- a/litellm-rust/crates/traces-clickhouse/tests/migrations.rs +++ b/litellm-rust/crates/traces-clickhouse/tests/migrations.rs @@ -335,10 +335,10 @@ async fn concurrent_schema_setup_succeeds( &database, "SELECT count() AS tables FROM system.tables \ WHERE database = 'trace_test' AND name IN \ - ('otel_traces', 'agent_traces_by_key', 'spend_logs')", + ('otel_traces', 'agent_traces_by_key', 'spend_logs', 'lens_feedback')", ) .await?; - assert_eq!(tables["data"][0]["tables"].as_u64(), Some(3)); + assert_eq!(tables["data"][0]["tables"].as_u64(), Some(4)); assert_eq!( migration_ledger_versions(&database).await?, migration_versions() @@ -1040,6 +1040,7 @@ async fn retention_changes_materialize_existing_rows_and_remain_idempotent( tables["data"], serde_json::json!([ {"name": "agent_traces_by_key"}, + {"name": "lens_feedback"}, {"name": "otel_traces"}, {"name": "spend_logs"} ]) @@ -1057,7 +1058,13 @@ async fn retention_changes_materialize_existing_rows_and_remain_idempotent( "start_time": old_timestamp_ms, "end_time": old_timestamp_ms + 1000 }))?; insert_rows(&database, "otel_traces", vec![span]).await?; + let old_iso = old_time.format(&time::format_description::well_known::Rfc3339)?; + let feedback = serde_json::from_value(serde_json::json!({ + "TeamId": "team-1", "ApiKeyHash": "", "TraceId": "expired", "Author": "admin", + "Score": 4, "Comment": "", "CreatedAt": old_iso, "UpdatedAt": old_iso, "IsDeleted": 0 + }))?; insert_rows(&database, "spend_logs", vec![spend]).await?; + insert_rows(&database, "lens_feedback", vec![feedback]).await?; assert_eq!(table_rows(&database, "agent_traces_by_key").await?, 1); ensure_schema(&database.client, &writer, "trace_test", 14).await?; let deadline = tokio::time::Instant::now() + Duration::from_secs(60); @@ -1087,6 +1094,8 @@ async fn retention_changes_materialize_existing_rows_and_remain_idempotent( ) .await?; execute_write(&database, "OPTIMIZE TABLE trace_test.spend_logs FINAL").await?; + execute_write(&database, "OPTIMIZE TABLE trace_test.lens_feedback FINAL").await?; + assert_eq!(table_rows(&database, "lens_feedback").await?, 0); assert_eq!(table_rows(&database, "otel_traces").await?, 0); assert_eq!(table_rows(&database, "agent_traces_by_key").await?, 0); assert_eq!(table_rows(&database, "spend_logs").await?, 0); @@ -1109,7 +1118,7 @@ async fn retention_reconciliation_updates_each_table_ttl( &database, "SELECT name, create_table_query FROM system.tables \ WHERE database = 'trace_test' AND name IN \ - ('otel_traces', 'agent_traces_by_key', 'spend_logs') ORDER BY name", + ('otel_traces', 'agent_traces_by_key', 'spend_logs', 'lens_feedback') ORDER BY name", ) .await?; let ttl_queries = ttl_queries["data"].as_array().expect("retention tables"); @@ -1118,7 +1127,12 @@ async fn retention_reconciliation_updates_each_table_ttl( .iter() .map(|row| row["name"].as_str().expect("table name")) .collect::>(), - ["agent_traces_by_key", "otel_traces", "spend_logs"] + [ + "agent_traces_by_key", + "lens_feedback", + "otel_traces", + "spend_logs" + ] ); for row in ttl_queries { let query = row["create_table_query"] @@ -1135,7 +1149,7 @@ async fn retention_reconciliation_updates_each_table_ttl( &database, "SELECT name, create_table_query FROM system.tables \ WHERE database = 'trace_test' AND name IN \ - ('otel_traces', 'agent_traces_by_key', 'spend_logs') ORDER BY name", + ('otel_traces', 'agent_traces_by_key', 'spend_logs', 'lens_feedback') ORDER BY name", ) .await?; for row in ttl_queries["data"].as_array().expect("retention tables") { diff --git a/litellm-rust/crates/traces-clickhouse/tests/queries/support.rs b/litellm-rust/crates/traces-clickhouse/tests/queries/support.rs index 7532d9aecc9..2b425bcfc1a 100644 --- a/litellm-rust/crates/traces-clickhouse/tests/queries/support.rs +++ b/litellm-rust/crates/traces-clickhouse/tests/queries/support.rs @@ -13,6 +13,7 @@ pub const DATABASE: &str = "trace_test"; pub struct SeededDatabase { pub database: ClickHouseDatabase, + #[allow(dead_code)] // dead_code: also shared with reads.rs, which uses a direct storage reader pub readers: QueryReaders, } diff --git a/litellm-rust/crates/traces-clickhouse/tests/reads.rs b/litellm-rust/crates/traces-clickhouse/tests/reads.rs index c938b6f451b..35dfafafd1d 100644 --- a/litellm-rust/crates/traces-clickhouse/tests/reads.rs +++ b/litellm-rust/crates/traces-clickhouse/tests/reads.rs @@ -3,9 +3,7 @@ use std::collections::BTreeMap; use litellm_http::Client; use litellm_traces::query::named::ReadAccessParams; use litellm_traces_cache::{ReadError, TraceReader}; -use litellm_traces_clickhouse::{ - ClickHouseTraces, Connection, InsertTable, QueryScope, insert_rows, -}; +use litellm_traces_clickhouse::{ClickHouseTraces, Connection, InsertTable, insert_rows}; use rstest::rstest; use serde_json::json; @@ -83,10 +81,7 @@ async fn list_costs_match_each_run_when_response_ids_are_reused( .collect(), ) .await?; - let connection = fixture - .readers - .connection(client, &QueryScope::All, "fixture-secret") - .await?; + let connection = Connection::reader(&fixture.database.url, DATABASE)?; let (reader, store) = make_reader(client, connection); let access = ReadAccessParams { all_teams: false, @@ -212,10 +207,7 @@ async fn large_runs_remain_complete_under_default_reader_limits( .collect::>(); insert_rows(client, &writer, DATABASE, InsertTable::SpendLogs, costs).await?; } - let connection = fixture - .readers - .connection(client, &QueryScope::All, "fixture-secret") - .await?; + let connection = Connection::reader(&fixture.database.url, DATABASE)?; let (reader, store) = make_reader(client, connection); let access = ReadAccessParams { all_teams: false, @@ -365,10 +357,7 @@ async fn cursor_pages_keep_a_tenant_scoped_snapshot_when_more_spans_arrive( ) -> TestResult { let fixture = seeded_database?; let client = &fixture.database.client; - let connection = fixture - .readers - .connection(client, &QueryScope::All, "fixture-secret") - .await?; + let connection = Connection::reader(&fixture.database.url, DATABASE)?; let (reader, store) = make_reader(client, connection.clone()); let access = ReadAccessParams { all_teams: true, @@ -528,10 +517,7 @@ async fn an_oversized_span_keeps_the_run_list_available_with_partial_totals( ) -> TestResult { let fixture = seeded_database?; let client = &fixture.database.client; - let connection = fixture - .readers - .connection(client, &QueryScope::All, "fixture-secret") - .await?; + let connection = Connection::reader(&fixture.database.url, DATABASE)?; let (reader, store) = make_reader(client, connection.clone()); let access = ReadAccessParams { all_teams: true, @@ -557,10 +543,7 @@ async fn an_oversized_span_keeps_the_run_list_available_with_partial_totals( ("TraceId".into(), json!(run.trace_id)), ("SpanId".into(), json!("oversized-child")), ("ParentSpanId".into(), json!("0101010101010101")), - ( - "SpanName".into(), - json!("x".repeat(litellm_storage_clickhouse::READ_LIMITS.response_bytes + 1)), - ), + ("SpanName".into(), json!("x".repeat(16 * 1024 * 1024 + 1))), ("ObservationType".into(), json!("tool")), ("TeamId".into(), json!("team-a")), ("ApiKeyHash".into(), json!("key-a")), @@ -699,10 +682,7 @@ async fn assigned_call_ids_require_shared_ownership_through_detail_and_batch_rea .collect(), ) .await?; - let connection = fixture - .readers - .connection(client, &QueryScope::All, "fixture-secret") - .await?; + let connection = Connection::reader(&fixture.database.url, DATABASE)?; let (reader, store) = make_reader(client, connection); let access = ReadAccessParams { all_teams: false, @@ -815,10 +795,7 @@ async fn native_cost_correlation_survives_session_grouping_and_excludes_other_ow .collect(), ) .await?; - let connection = fixture - .readers - .connection(client, &QueryScope::All, "fixture-secret") - .await?; + let connection = Connection::reader(&fixture.database.url, DATABASE)?; let (reader, store) = make_reader(client, connection); let access = ReadAccessParams { all_teams: true, diff --git a/litellm-rust/crates/traces/src/query.rs b/litellm-rust/crates/traces/src/query.rs index e2653e15b61..26eaa33d5ac 100644 --- a/litellm-rust/crates/traces/src/query.rs +++ b/litellm-rust/crates/traces/src/query.rs @@ -17,6 +17,9 @@ pub enum ReadQuery { Sample, Content, Evidence, + FeedbackTarget, + Feedback, + FeedbackSummary, } impl ReadQuery { diff --git a/litellm-rust/crates/traces/src/query/named.rs b/litellm-rust/crates/traces/src/query/named.rs index 03459645ca9..dfa7ac2f2cd 100644 --- a/litellm-rust/crates/traces/src/query/named.rs +++ b/litellm-rust/crates/traces/src/query/named.rs @@ -116,6 +116,8 @@ pub struct TraceSpansRow { pub source_url: String, #[serde(default)] pub source_title: String, + #[serde(default)] + pub source_user: String, pub team_id: String, pub api_key_hash: String, pub user_id: String, diff --git a/litellm-rust/crates/traces/src/resolve/view.rs b/litellm-rust/crates/traces/src/resolve/view.rs index 97f83c255f8..62c27676605 100644 --- a/litellm-rust/crates/traces/src/resolve/view.rs +++ b/litellm-rust/crates/traces/src/resolve/view.rs @@ -146,6 +146,7 @@ fn source(row: &TraceSpansRow) -> Option { .unwrap_or(RunSourceType::Custom), url: row.source_url.clone(), title: row.source_title.clone(), + user: row.source_user.clone(), }) } diff --git a/litellm-rust/crates/traces/src/view.rs b/litellm-rust/crates/traces/src/view.rs index 1613d96d232..729ab5ef1c3 100644 --- a/litellm-rust/crates/traces/src/view.rs +++ b/litellm-rust/crates/traces/src/view.rs @@ -88,6 +88,10 @@ pub struct RunSource { pub kind: RunSourceType, pub url: String, pub title: String, + /// Who started the conversation, e.g. the Slack user's email. + #[serde(default, skip_serializing_if = "String::is_empty")] + #[cfg_attr(feature = "schema", schemars(extend("x-python-optional" = true)))] + pub user: String, } #[macro_rules_attribute::apply(response_type)] diff --git a/litellm-rust/crates/traces/tests/captures.rs b/litellm-rust/crates/traces/tests/captures.rs index 922e28f97ac..a81bca609ea 100644 --- a/litellm-rust/crates/traces/tests/captures.rs +++ b/litellm-rust/crates/traces/tests/captures.rs @@ -210,6 +210,7 @@ fn trace_span(span: DecodedSpan) -> TraceSpansRow { source_type: String::new(), source_url: String::new(), source_title: String::new(), + source_user: String::new(), team_id: "fixture-team".into(), api_key_hash: "fixture-key".into(), user_id: "fixture-user".into(), @@ -297,6 +298,7 @@ fn unrelated_transport(call: &TraceSpansRow) -> TraceSpansRow { source_type: String::new(), source_url: String::new(), source_title: String::new(), + source_user: String::new(), team_id: call.team_id.clone(), api_key_hash: call.api_key_hash.clone(), user_id: call.user_id.clone(), diff --git a/litellm-rust/crates/traces/tests/otlp.rs b/litellm-rust/crates/traces/tests/otlp.rs index 68a25fb028b..1b817a0bef5 100644 --- a/litellm-rust/crates/traces/tests/otlp.rs +++ b/litellm-rust/crates/traces/tests/otlp.rs @@ -1650,6 +1650,9 @@ fn session_capture_joins_native_logs_and_traces_across_turns_without_changing_sp let first = litellm_traces::decode_otlp_logs(&logs.encode_to_vec(), None).unwrap(); let second = decode_otlp(&request.encode_to_vec(), None).unwrap(); assert_eq!(first[0].trace_id, second[0].trace_id); + // Lens feedback resolves session ids the same way; keep in sync with + // litellm/proxy/lens/feedback_repository.py::session_trace_id. + assert_eq!(second[0].trace_id, "5fddf060372c8501dca4f331b9da882b"); assert_eq!( first[0].attributes["lens.original_trace_id"], "01".repeat(16) diff --git a/litellm-rust/crates/traces/tests/query.rs b/litellm-rust/crates/traces/tests/query.rs index 67f422d6997..7c125b01379 100644 --- a/litellm-rust/crates/traces/tests/query.rs +++ b/litellm-rust/crates/traces/tests/query.rs @@ -14,6 +14,9 @@ use rstest::rstest; #[case::sample("sample", ReadQuery::Sample)] #[case::content("content", ReadQuery::Content)] #[case::evidence("evidence", ReadQuery::Evidence)] +#[case::feedback_target("feedback_target", ReadQuery::FeedbackTarget)] +#[case::feedback("feedback", ReadQuery::Feedback)] +#[case::feedback_summary("feedback_summary", ReadQuery::FeedbackSummary)] fn names_select_the_public_query(#[case] name: &str, #[case] query: ReadQuery) { assert_eq!(ReadQuery::parse(name).unwrap(), query); assert_eq!(query.as_ref(), name); diff --git a/litellm-rust/crates/traces/tests/query/named.rs b/litellm-rust/crates/traces/tests/query/named.rs index 25cfd7083a5..af33ee3fe85 100644 --- a/litellm-rust/crates/traces/tests/query/named.rs +++ b/litellm-rust/crates/traces/tests/query/named.rs @@ -53,7 +53,7 @@ fn result_contracts_preserve_public_field_names() { json!({"trace_id": "trace", "trace_ref": "ref", "team_id": "team", "api_key_hash": "key", "user_id": "user", "name": "agent", "service": "service", "input_preview": "input", "status": "STATUS_CODE_OK", "start_ms": -1, "duration_ms": 20, "span_count": u64::MAX, "agent_count": 1, "agent_invocations": 2, "agent_names": ["agent"], "frameworks": ["framework"], "llm_calls": 3, "tool_calls": 4, "input_tokens": 5, "output_tokens": 6, "models": ["model"], "error_count": 0, "request_ids": ["request"]}), ); round_trip::( - json!({"trace_id": "trace", "original_trace_id": "original", "span_id": "span", "parent_span_id": "parent", "name": "agent", "type": "agent", "wrapper_candidate": 1, "agent": "agent", "framework": "framework", "status": "STATUS_CODE_ERROR", "status_message": "error", "error_truncated": 1, "start_ns": -1, "duration_ns": u64::MAX, "service": "service", "input_preview": "input", "model": "model", "input_tokens": u32::MAX, "output_tokens": 6, "litellm_request_id": "request", "call_keys": ["provider_response:request"], "call_evidence": "complete", "tool_call_id": "call", "source_type": "slack", "source_url": "https://acme.slack.com/archives/C1/p1", "source_title": "thread", "team_id": "team", "api_key_hash": "key", "user_id": "user"}), + json!({"trace_id": "trace", "original_trace_id": "original", "span_id": "span", "parent_span_id": "parent", "name": "agent", "type": "agent", "wrapper_candidate": 1, "agent": "agent", "framework": "framework", "status": "STATUS_CODE_ERROR", "status_message": "error", "error_truncated": 1, "start_ns": -1, "duration_ns": u64::MAX, "service": "service", "input_preview": "input", "model": "model", "input_tokens": u32::MAX, "output_tokens": 6, "litellm_request_id": "request", "call_keys": ["provider_response:request"], "call_evidence": "complete", "tool_call_id": "call", "source_type": "slack", "source_url": "https://acme.slack.com/archives/C1/p1", "source_title": "thread", "source_user": "tin@berri.ai", "team_id": "team", "api_key_hash": "key", "user_id": "user"}), ); round_trip::( json!({"span_id": "span", "input": "input", "output": "output", "attributes": {"count": "42"}}), diff --git a/litellm-rust/crates/traces/tests/resolve.rs b/litellm-rust/crates/traces/tests/resolve.rs index baaa4c338cf..73f691cb2a6 100644 --- a/litellm-rust/crates/traces/tests/resolve.rs +++ b/litellm-rust/crates/traces/tests/resolve.rs @@ -36,6 +36,7 @@ fn row(span_id: &str, parent: &str, name: &str, kind: &str, agent: &str) -> Trac source_type: String::new(), source_url: String::new(), source_title: String::new(), + source_user: String::new(), team_id: "team".into(), api_key_hash: "key".into(), user_id: String::new(), @@ -227,6 +228,16 @@ fn summary_source_type_picks_the_app(#[case] source_type: &str, #[case] expected assert_eq!(source.map(|source| source.kind), Some(expected)); } +#[rstest] +#[case::set("tin@berri.ai")] +#[case::missing("")] +fn summary_source_carries_who_started_it(#[case] user: &str) { + let mut root = sourced(row("root", "", "agent", "agent", "agent"), THREAD, "t"); + root.source_user = user.into(); + let source = resolve_trace("t", "", &[root], &[]).unwrap().summary.source; + assert_eq!(source.map(|source| source.user), Some(user.to_owned())); +} + #[rstest] fn spans_are_offset_from_the_trace_start() { let trace = resolve_trace("t1", "", &deep_agent(1), &[]).unwrap(); diff --git a/litellm/__init__.py b/litellm/__init__.py index c0dc61e2911..3f93166ecc8 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -1667,6 +1667,24 @@ if TYPE_CHECKING: from .llms.hosted_vllm.rerank.transformation import ( HostedVLLMRerankConfig as HostedVLLMRerankConfig, ) + from .llms.perplexity.decisions.transformation import ( + PerplexityDecisionsConfig as PerplexityDecisionsConfig, + ) + from .llms.typesafe.decisions.transformation import ( + TypeSafeDecisionsConfig as TypeSafeDecisionsConfig, + ) + from .llms.openrouter.decisions.transformation import ( + OpenRouterDecisionsConfig as OpenRouterDecisionsConfig, + ) + from .llms.cloudflare.decisions.transformation import ( + CloudflareDecisionsConfig as CloudflareDecisionsConfig, + ) + from .llms.strands_decider.decisions.transformation import ( + StrandsDeciderDecisionsConfig as StrandsDeciderDecisionsConfig, + ) + from .llms.openai.decisions.transformation import ( + OpenAIDecisionsConfig as OpenAIDecisionsConfig, + ) from .llms.nvidia_nim.rerank.transformation import ( NvidiaNimRerankConfig as NvidiaNimRerankConfig, ) diff --git a/litellm/_lazy_imports_registry.py b/litellm/_lazy_imports_registry.py index 78f5e4ef04b..abefab95e9f 100644 --- a/litellm/_lazy_imports_registry.py +++ b/litellm/_lazy_imports_registry.py @@ -156,6 +156,12 @@ LLM_CONFIG_NAMES: Final = ( "ScalewayRerankConfig", "DeepinfraRerankConfig", "HostedVLLMRerankConfig", + "PerplexityDecisionsConfig", + "TypeSafeDecisionsConfig", + "OpenRouterDecisionsConfig", + "CloudflareDecisionsConfig", + "StrandsDeciderDecisionsConfig", + "OpenAIDecisionsConfig", "NvidiaNimRerankConfig", "NvidiaNimRankingConfig", "VertexAIRerankConfig", @@ -707,6 +713,15 @@ _LLM_CONFIGS_IMPORT_MAP: Final = { ".llms.hosted_vllm.rerank.transformation", "HostedVLLMRerankConfig", ), + "PerplexityDecisionsConfig": (".llms.perplexity.decisions.transformation", "PerplexityDecisionsConfig"), + "TypeSafeDecisionsConfig": (".llms.typesafe.decisions.transformation", "TypeSafeDecisionsConfig"), + "OpenRouterDecisionsConfig": (".llms.openrouter.decisions.transformation", "OpenRouterDecisionsConfig"), + "CloudflareDecisionsConfig": (".llms.cloudflare.decisions.transformation", "CloudflareDecisionsConfig"), + "StrandsDeciderDecisionsConfig": ( + ".llms.strands_decider.decisions.transformation", + "StrandsDeciderDecisionsConfig", + ), + "OpenAIDecisionsConfig": (".llms.openai.decisions.transformation", "OpenAIDecisionsConfig"), "NvidiaNimRerankConfig": ( ".llms.nvidia_nim.rerank.transformation", "NvidiaNimRerankConfig", diff --git a/litellm/batches/batch_utils.py b/litellm/batches/batch_utils.py index ba625ac5c00..6573f13be75 100644 --- a/litellm/batches/batch_utils.py +++ b/litellm/batches/batch_utils.py @@ -416,11 +416,11 @@ def _provider_output_file_id(output_file_id: str) -> str: llm_output_file_id, model-encoded ids decode to the raw provider id, raw ids pass through. """ from litellm.proxy.openai_files_endpoints.common_utils import ( - _is_base64_encoded_unified_file_id, get_original_file_id, + is_base64_encoded_unified_file_id, ) - unified_file_id: Final = _is_base64_encoded_unified_file_id(output_file_id) + unified_file_id: Final = is_base64_encoded_unified_file_id(output_file_id) if not unified_file_id: return get_original_file_id(output_file_id) try: diff --git a/litellm/completion_extras/litellm_responses_transformation/transformation.py b/litellm/completion_extras/litellm_responses_transformation/transformation.py index dc79c6d2555..142d624ef71 100644 --- a/litellm/completion_extras/litellm_responses_transformation/transformation.py +++ b/litellm/completion_extras/litellm_responses_transformation/transformation.py @@ -19,7 +19,7 @@ from openai.types.responses.response_input_param import ( from openai.types.responses.tool_choice_custom_param import ToolChoiceCustomParam from openai.types.responses.tool_choice_function_param import ToolChoiceFunctionParam from openai.types.responses.tool_param import FunctionToolParam -from pydantic import BaseModel +from pydantic import BaseModel, TypeAdapter, ValidationError import litellm from litellm import ModelResponse @@ -46,6 +46,7 @@ from litellm.types.llms.openai import ( ChatCompletionToolCallChunk, ChatCompletionToolCallFunctionChunk, ChatCompletionToolParamFunctionChunk, + PromptCacheBreakpoint, Reasoning, ResponsesAPIOptionalRequestParams, ResponsesAPIResponse, @@ -238,6 +239,19 @@ def _map_incomplete_reason_to_finish_reason(incomplete_reason: str | None) -> Li return "length" +_PROMPT_CACHE_BREAKPOINT: Final = TypeAdapter(PromptCacheBreakpoint) + + +def _prompt_cache_breakpoint_for_wire(marker: object, drop_params: bool) -> object: + if marker is None or not drop_params: + return marker + try: + return _PROMPT_CACHE_BREAKPOINT.validate_python(marker) + except ValidationError: + verbose_logger.debug("Chat provider: dropping malformed prompt_cache_breakpoint %r under drop_params", marker) + return None + + def _input_file_from_file_value(file_value: object) -> dict[str, object]: if not isinstance(file_value, dict): return {"type": "input_file"} @@ -400,9 +414,12 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): self, messages: list["AllMessageValues"], *, + drop_params: bool = False, keep_prompt_cache_breakpoints: bool = False, ) -> tuple[list[object], str | None]: - converted_input_items, instructions = self._convert_chat_completion_messages_to_responses_input(messages) + converted_input_items, instructions = self._convert_chat_completion_messages_to_responses_input( + messages, drop_params=drop_params + ) return ( converted_input_items if keep_prompt_cache_breakpoints @@ -411,7 +428,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): ) def _convert_chat_completion_messages_to_responses_input( - self, messages: list["AllMessageValues"] + self, messages: list["AllMessageValues"], *, drop_params: bool = False ) -> tuple[list[object], str | None]: input_items: Final[list[object]] = [] instructions: str | None = None @@ -452,6 +469,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): "content": self._convert_content_to_responses_format( content, role, + drop_params=drop_params, ), } ) @@ -470,6 +488,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): tool_output = self._convert_content_to_responses_format( content, "user", # Use "user" role to get input_* types + drop_params=drop_params, ) else: # Fallback: convert unexpected types to input_text @@ -497,7 +516,9 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): { "type": "message", "role": "assistant", - "content": self._convert_content_to_responses_format(content, "assistant"), + "content": self._convert_content_to_responses_format( + content, "assistant", drop_params=drop_params + ), } ) for tool_call in tool_calls: @@ -531,7 +552,9 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): { "type": "message", "role": role, - "content": self._convert_content_to_responses_format(content, cast(str, role)), + "content": self._convert_content_to_responses_format( + content, cast(str, role), drop_params=drop_params + ), } ) elif role == "assistant": @@ -647,6 +670,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): ) converted_input_items, converted_instructions = self.convert_chat_completion_messages_to_responses_api( messages, + drop_params=bool(litellm_params.get("drop_params") or litellm.drop_params), keep_prompt_cache_breakpoints=supports_prompt_cache_breakpoint, ) # OpenAI's Responses API rejects an empty input. For a system-only @@ -1126,6 +1150,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): ] | None, role: str, + drop_params: bool = False, ) -> list[dict[str, object]]: """Convert chat completion content to responses API format""" from litellm.types.llms.openai import ChatCompletionImageObject @@ -1152,7 +1177,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): if original_type == "text": converted = with_prompt_cache_breakpoint( self._convert_content_str_to_input_text(item.get("text", ""), role), - item.get("prompt_cache_breakpoint"), + _prompt_cache_breakpoint_for_wire(item.get("prompt_cache_breakpoint"), drop_params), ) result.append(converted) verbose_logger.debug("Chat provider: text -> %s", converted) @@ -1165,7 +1190,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): cast(ChatCompletionImageObject, item), role ), ), - item.get("prompt_cache_breakpoint"), + _prompt_cache_breakpoint_for_wire(item.get("prompt_cache_breakpoint"), drop_params), ) result.append(converted) verbose_logger.debug("Chat provider: image_url -> %s", converted) @@ -1181,7 +1206,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): _input_file_from_file_value( cast("ChatCompletionFileObject", item).get("file"), # cast-ok: type tag checked ), - item.get("prompt_cache_breakpoint"), + _prompt_cache_breakpoint_for_wire(item.get("prompt_cache_breakpoint"), drop_params), ) result.append(converted) verbose_logger.debug("Chat provider: file -> %s", converted) @@ -1203,7 +1228,10 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): verbose_logger.debug("Chat provider: passthrough -> %s", item) else: # Default to input_text for unknown types - converted = self._convert_content_str_to_input_text(str(item.get("text", item)), role) + converted = with_prompt_cache_breakpoint( + self._convert_content_str_to_input_text(str(item.get("text", item)), role), + _prompt_cache_breakpoint_for_wire(item.get("prompt_cache_breakpoint"), drop_params), + ) result.append(converted) verbose_logger.debug("Chat provider: unknown(%s) -> %s", original_type, converted) verbose_logger.debug("Chat provider: Final converted content: %s", result) diff --git a/litellm/constants.py b/litellm/constants.py index b09da3d6e7d..c4adc0de22f 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -64,6 +64,8 @@ AGENT_TRACING_AGENT_LIST_LIMIT: Final = get_env_int("AGENT_TRACING_AGENT_LIST_LI LENS_DATASET_MAX_CASES: Final = get_env_int("LENS_DATASET_MAX_CASES", 200) LENS_DATASET_MAX_CASE_CHARS: Final = get_env_int("LENS_DATASET_MAX_CASE_CHARS", 20_000) LENS_DATASET_TRACE_PAGE_SIZE: Final = 500 +LENS_FEEDBACK_MAX_COMMENT_CHARS: Final = 10_000 +LENS_FEEDBACK_MAX_SCORE: Final = 10 DEFAULT_S3_FLUSH_INTERVAL_SECONDS: Final = int(os.getenv("DEFAULT_S3_FLUSH_INTERVAL_SECONDS", 10)) DEFAULT_S3_BATCH_SIZE: Final = int(os.getenv("DEFAULT_S3_BATCH_SIZE", 512)) DEFAULT_S3_MAX_CONCURRENT_UPLOADS: Final = int(os.getenv("DEFAULT_S3_MAX_CONCURRENT_UPLOADS", "16")) diff --git a/litellm/cost_calculator.py b/litellm/cost_calculator.py index 4d5f882ea10..da420ab0578 100644 --- a/litellm/cost_calculator.py +++ b/litellm/cost_calculator.py @@ -9,7 +9,7 @@ from typing import TYPE_CHECKING, Any, Final, Literal, cast from httpx import Response from pydantic import BaseModel -from typing_extensions import ReadOnly, TypedDict +from typing_extensions import ReadOnly, TypedDict, assert_never import litellm import litellm._logging @@ -100,7 +100,7 @@ from litellm.llms.vertex_ai.cost_calculator import cost_router as google_cost_ro from litellm.llms.xai.cost_calculator import cost_per_token as xai_cost_per_token from litellm.responses.utils import ResponseAPILoggingUtils from litellm.types.agents import LiteLLMSendMessageResponse -from litellm.types.decisions import DecisionsResponse, DecisionsUsage +from litellm.types.decisions import DecisionsResponse, DecisionsUsage, OpenAIDecisionResponse, OpenAIDecisionUsage from litellm.types.llms.base import CachedTokensDetails, LiteLLMBaseModel from litellm.types.llms.openai import ( HttpxBinaryResponseContent, @@ -1034,6 +1034,8 @@ def get_usage_object( ), ) + if isinstance(completion_response, (DecisionsResponse, OpenAIDecisionResponse)): + return None if completion_response.usage is None else _decisions_usage(completion_response.usage) if usage_obj is None: return None if isinstance(usage_obj, Usage): @@ -1064,12 +1066,37 @@ def get_usage_object( return None +def _decisions_prompt_tokens_details(usage: DecisionsUsage | OpenAIDecisionUsage) -> PromptTokensDetailsWrapper: + match usage: + case DecisionsUsage(): + return PromptTokensDetailsWrapper( + cached_tokens=usage.cached_tokens, cache_write_tokens=usage.cache_write_tokens + ) + case OpenAIDecisionUsage(): + return PromptTokensDetailsWrapper( + cached_tokens=usage.input_tokens_details.cached_tokens, + cache_write_tokens=usage.input_tokens_details.cache_write_tokens, + ) + case _: + assert_never(usage) + + +def _decisions_usage(usage: DecisionsUsage | OpenAIDecisionUsage) -> Usage: + return Usage( + prompt_tokens=usage.input_tokens, + completion_tokens=usage.output_tokens, + total_tokens=usage.input_tokens + usage.output_tokens, + prompt_tokens_details=_decisions_prompt_tokens_details(usage), + ) + + def _is_known_usage_objects(usage_obj): """Returns True if the usage obj is a known Usage type""" return ( isinstance(usage_obj, litellm.Usage) or isinstance(usage_obj, ResponseAPIUsage) or isinstance(usage_obj, DecisionsUsage) + or isinstance(usage_obj, OpenAIDecisionUsage) or TranscriptionUsageObjectTransformation.is_transcription_usage_object(usage_obj) ) @@ -1481,11 +1508,8 @@ def completion_cost( "usage", litellm.Usage(**_usage_for_dump.model_dump()), ) - if isinstance(usage_obj, DecisionsUsage): - _usage = { - "prompt_tokens": usage_obj.input_tokens, - "completion_tokens": usage_obj.output_tokens, - } + if isinstance(usage_obj, (DecisionsUsage, OpenAIDecisionUsage)): + _usage = _decisions_usage(usage_obj).model_dump() elif usage_obj is None: _usage = {} elif isinstance(usage_obj, BaseModel): @@ -1977,7 +2001,8 @@ def response_cost_calculator( | OpenAIModerationResponse | Response | SearchResponse - | DecisionsResponse, + | DecisionsResponse + | OpenAIDecisionResponse, model: str, custom_llm_provider: str | None, call_type: Literal[ diff --git a/litellm/decisions/main.py b/litellm/decisions/main.py index 8304c52a886..25ec501fd48 100644 --- a/litellm/decisions/main.py +++ b/litellm/decisions/main.py @@ -1,114 +1,147 @@ -from collections.abc import Mapping +from collections.abc import Mapping, Sequence from dataclasses import dataclass, field -from types import MappingProxyType -from typing import Final +from typing import Final, TypeAlias import httpx from pydantic import TypeAdapter, ValidationError +from typing_extensions import assert_never import litellm from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj -from litellm.llms.base_llm.decisions.transformation import DecisionsProviderConfig -from litellm.llms.cloudflare.decisions.transformation import CLOUDFLARE_DECISIONS_ENDPOINT -from litellm.llms.custom_httpx.http_handler import get_async_httpx_client, get_httpx_client -from litellm.llms.openrouter.decisions.transformation import OPENROUTER_DECISIONS_ENDPOINT -from litellm.llms.perplexity.decisions.transformation import PERPLEXITY_DECISIONS_ENDPOINT -from litellm.llms.strands_decider.decisions.transformation import STRANDS_DECIDER_DECISIONS_ENDPOINT -from litellm.llms.typesafe.decisions.transformation import TYPESAFE_DECISIONS_ENDPOINT -from litellm.secret_managers.main import get_secret_str +from litellm.llms.base_llm.chat.transformation import BaseLLMException +from litellm.llms.base_llm.decisions.transformation import ( + BaseDecisionsConfig, + ir_to_systemone_response, + systemone_request_to_ir, +) +from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler +from litellm.llms.openai.decisions.transformation import ir_to_openai_response, openai_request_to_ir from litellm.types.decisions import ( DecisionQuestion, + DecisionsIRRequest, + DecisionsIRResponse, DecisionsJSON, - DecisionsRequest, + DecisionsRequestBody, DecisionsResponse, + OpenAIDecisionInput, + OpenAIDecisionQuestion, + OpenAIDecisionRequestBody, + OpenAIDecisionResponse, + UnsupportedDecisionsRequest, ) -from litellm.utils import client +from litellm.types.utils import LlmProviders +from litellm.utils import ProviderConfigManager, client -DECISIONS_ENDPOINTS: Final[Mapping[str, DecisionsProviderConfig]] = MappingProxyType( - { - "perplexity": PERPLEXITY_DECISIONS_ENDPOINT, - "typesafe": TYPESAFE_DECISIONS_ENDPOINT, - "openrouter": OPENROUTER_DECISIONS_ENDPOINT, - "cloudflare": CLOUDFLARE_DECISIONS_ENDPOINT, - "strands_decider": STRANDS_DECIDER_DECISIONS_ENDPOINT, - } +DecisionsQuestions: TypeAlias = ( + Mapping[str, DecisionQuestion | Mapping[str, object]] | Sequence[OpenAIDecisionQuestion | Mapping[str, object]] ) +DecisionsRequestFormat: TypeAlias = DecisionsRequestBody | OpenAIDecisionRequestBody -_DECISIONS_REQUEST_ADAPTER: Final[TypeAdapter[DecisionsRequest]] = TypeAdapter(DecisionsRequest) -_DECISIONS_PAYLOAD_ADAPTER: Final[TypeAdapter[object]] = TypeAdapter(object) -_DECISIONS_RESPONSE_ADAPTER: Final[TypeAdapter[DecisionsResponse]] = TypeAdapter(DecisionsResponse) +_SYSTEMONE_REQUEST_ADAPTER: Final[TypeAdapter[DecisionsRequestBody]] = TypeAdapter(DecisionsRequestBody) +_OPENAI_REQUEST_ADAPTER: Final[TypeAdapter[OpenAIDecisionRequestBody]] = TypeAdapter(OpenAIDecisionRequestBody) +_HANDLER: Final = BaseLLMHTTPHandler() @dataclass(frozen=True, slots=True, repr=False) -class _PreparedDecisionsRequest: - config: DecisionsProviderConfig - provider: str - upstream_model: str - url: str - api_key: str | None = field(repr=False) - headers: Mapping[str, str] = field(repr=False) +class _DecisionsCall: + model: str + requested_model: str + custom_llm_provider: str + provider_config: BaseDecisionsConfig + request: DecisionsRequestFormat + ir_request: DecisionsIRRequest body: Mapping[str, object] = field(repr=False) + api_base: str + api_key: str | None = field(repr=False) + logging_obj: LiteLLMLoggingObj | None + headers: Mapping[str, str] + timeout: float | httpx.Timeout | None -def _resolve_provider_model(model: str, custom_llm_provider: str | None) -> tuple[str, str]: - provider: Final = model.partition("/")[0] if custom_llm_provider is None else custom_llm_provider - if provider not in DECISIONS_ENDPOINTS: - supported: Final = ", ".join(DECISIONS_ENDPOINTS) +def _supported_providers() -> tuple[str, ...]: + return tuple( + provider.value + for provider in LlmProviders + if ProviderConfigManager.get_provider_decisions_config(model="", provider=provider) is not None + ) + + +def _provider_config(model: str, custom_llm_provider: str) -> BaseDecisionsConfig: + provider: Final = next((member for member in LlmProviders if member.value == custom_llm_provider), None) + provider_config: Final = ( + None if provider is None else ProviderConfigManager.get_provider_decisions_config(model, provider) + ) + if provider_config is None: + supported: Final = ", ".join(_supported_providers()) raise litellm.BadRequestError( - message=f"Unknown Decisions provider '{provider}'. Supported providers: {supported}", + message=f"Unknown Decisions provider '{custom_llm_provider}'. Supported providers: {supported}", model=model, - llm_provider=provider, + llm_provider=custom_llm_provider, ) - upstream_model: Final = model.removeprefix(f"{provider}/") + return provider_config + + +def _validate_request( + *, + state: DecisionsJSON | None, + questions: DecisionsQuestions | None, + decision_input: OpenAIDecisionInput | None, + safety_identifier: str | None, +) -> DecisionsRequestFormat: + if decision_input is None: + return _SYSTEMONE_REQUEST_ADAPTER.validate_python({"state": state, "questions": questions}) + return _OPENAI_REQUEST_ADAPTER.validate_python( + {"input": decision_input, "questions": questions, "safety_identifier": safety_identifier} + ) + + +def _ir_request(request: DecisionsRequestFormat) -> DecisionsIRRequest: + match request: + case DecisionsRequestBody(): + return systemone_request_to_ir(request) + case OpenAIDecisionRequestBody(): + return openai_request_to_ir(request) + case _: + assert_never(request) + + +def _prepare_call( + *, + model: str, + state: DecisionsJSON | None, + questions: DecisionsQuestions | None, + decision_input: OpenAIDecisionInput | None, + safety_identifier: str | None, + api_key: str | None, + api_base: str | None, + timeout: float | httpx.Timeout | None, + custom_llm_provider: str | None, + extra_headers: Mapping[str, str] | None, + kwargs: Mapping[str, object], +) -> _DecisionsCall: + upstream_model, provider, dynamic_api_key, dynamic_api_base = litellm.get_llm_provider( + model=model, + custom_llm_provider=custom_llm_provider, + api_base=api_base, + api_key=api_key, + ) + provider_config: Final = _provider_config(upstream_model, provider) + canonical_model: Final = provider_config.canonical_model(upstream_model) if not upstream_model: raise litellm.BadRequestError( message="A model name is required for the Decisions API", model=model, llm_provider=provider, ) - return provider, upstream_model - - -def _resolve_api_key( - *, - provider: str, - model: str, - endpoint: DecisionsProviderConfig, - api_key: str | None, -) -> str | None: - if api_key is not None: - return api_key - - server_api_key: Final = next( - (key for key in (get_secret_str(name) for name in endpoint.api_key_env) if key), - None, - ) - if server_api_key is None: - if not endpoint.api_key_required: - return None - raise litellm.AuthenticationError( - message=f"Missing API key for Decisions provider '{provider}'", + if state is not None and decision_input is not None: + raise litellm.BadRequestError( + message="Pass either state (System One format) or input (OpenAI format) to the Decisions API, not both", model=model, llm_provider=provider, ) - - return server_api_key - - -def _prepare_request( - *, - model: str, - state: DecisionsJSON, - questions: Mapping[str, DecisionQuestion | Mapping[str, object]], - api_key: str | None, - api_base: str | None, - custom_llm_provider: str | None, - extra_headers: Mapping[str, str] | None, -) -> _PreparedDecisionsRequest: - provider, upstream_model = _resolve_provider_model(model, custom_llm_provider) try: - validated_request: Final = _DECISIONS_REQUEST_ADAPTER.validate_python( - {"model": model, "state": state, "questions": questions} + request: Final = _validate_request( + state=state, questions=questions, decision_input=decision_input, safety_identifier=safety_identifier ) except ValidationError as error: raise litellm.BadRequestError( @@ -117,109 +150,94 @@ def _prepare_request( llm_provider=provider, ) from error - endpoint: Final = DECISIONS_ENDPOINTS[provider] - env_api_base: Final = get_secret_str(endpoint.api_base_env) - default_api_base: Final = endpoint.default_api_base() - resolved_api_base: Final = api_base or env_api_base or default_api_base + resolved_api_base: Final = provider_config.resolve_api_base(dynamic_api_base or api_base) if resolved_api_base is None: raise litellm.BadRequestError( - message=endpoint.missing_api_base_message(provider), + message=provider_config.missing_api_base_message(provider), + model=model, + llm_provider=provider, + ) + resolved_api_key: Final = provider_config.resolve_api_key(dynamic_api_key or api_key) + if resolved_api_key is None and provider_config.api_key_required: + raise litellm.AuthenticationError( + message=f"Missing API key for Decisions provider '{provider}'", model=model, llm_provider=provider, ) - resolved_api_key: Final = _resolve_api_key( - provider=provider, - model=model, - endpoint=endpoint, - api_key=api_key, + ir_request: Final = _ir_request(request) + body: Final = provider_config.transform_decisions_request( + model=canonical_model, request=ir_request, custom_llm_provider=provider ) + if isinstance(body, UnsupportedDecisionsRequest): + raise litellm.BadRequestError( + message=f"Decisions provider '{provider}' cannot serve this request: {body.reason}", + model=model, + llm_provider=provider, + ) - canonical_model: Final = endpoint.canonical_model(upstream_model) - outbound_headers: Final = MappingProxyType( - { - **{ - name: value - for name, value in (extra_headers or {}).items() - if name.lower() not in {"authorization", "content-type"} - }, - **({"Authorization": f"Bearer {resolved_api_key}"} if resolved_api_key is not None else {}), - "Content-Type": "application/json", - } - ) - body: Final = MappingProxyType( - { - "model": endpoint.request_model(canonical_model), - "state": validated_request.state, - "questions": { - name: question.model_dump(mode="json", exclude_none=True) - for name, question in validated_request.questions.items() - }, - } - ) - return _PreparedDecisionsRequest( - config=endpoint, - provider=provider, - upstream_model=canonical_model, - url=endpoint.endpoint_url(resolved_api_base, canonical_model), - api_key=resolved_api_key, - headers=outbound_headers, - body=body, - ) - - -def _log_request( - prepared: _PreparedDecisionsRequest, - kwargs: Mapping[str, object], -) -> LiteLLMLoggingObj | None: logging_obj: Final = kwargs.get("litellm_logging_obj") - if not isinstance(logging_obj, LiteLLMLoggingObj): - return None - logging_obj.update_from_kwargs( - kwargs=dict(kwargs), - model=prepared.upstream_model, - litellm_params={ - "litellm_call_id": kwargs.get("litellm_call_id"), - "api_base": prepared.url, - }, - custom_llm_provider=prepared.provider, + if isinstance(logging_obj, LiteLLMLoggingObj): + logging_obj.update_from_kwargs( + kwargs=dict(kwargs), + model=canonical_model, + litellm_params={ + "litellm_call_id": kwargs.get("litellm_call_id"), + "api_base": provider_config.get_complete_url(resolved_api_base, canonical_model), + }, + custom_llm_provider=provider, + ) + return _DecisionsCall( + model=canonical_model, + requested_model=model, + custom_llm_provider=provider, + provider_config=provider_config, + request=request, + ir_request=ir_request, + body=body, + api_base=resolved_api_base, + api_key=resolved_api_key, + logging_obj=logging_obj if isinstance(logging_obj, LiteLLMLoggingObj) else None, + headers=extra_headers or {}, + timeout=timeout, ) - request_body: Final = dict(prepared.body) - request_headers: Final = dict(prepared.headers) - logging_obj.pre_call( - input=request_body, - api_key=prepared.api_key, - model=prepared.upstream_model, - additional_args={ - "api_base": prepared.url, - "complete_input_dict": request_body, - "headers": request_headers, - }, - ) - return logging_obj -def _parse_response( - response: httpx.Response, - prepared: _PreparedDecisionsRequest, -) -> DecisionsResponse: - response.raise_for_status() - payload: Final[object] = _DECISIONS_PAYLOAD_ADAPTER.validate_json(response.content) - result: Final = _DECISIONS_RESPONSE_ADAPTER.validate_python(prepared.config.unwrap_response(payload)) - result.hidden_params.update( +def _format_response(response: DecisionsIRResponse, call: _DecisionsCall) -> DecisionsResponse | OpenAIDecisionResponse: + formatted: Final = _formatted_response(response, call) + formatted.set_hidden_params( { - "model": f"{prepared.provider}/{prepared.upstream_model}", - "custom_llm_provider": prepared.provider, - "provider_response_model": f"{prepared.provider}/{prepared.upstream_model}", + "model": f"{call.custom_llm_provider}/{call.model}", + "custom_llm_provider": call.custom_llm_provider, + "provider_response_model": f"{call.custom_llm_provider}/{call.model}", } ) - return result + return formatted -def _map_upstream_exception(error: Exception, prepared: _PreparedDecisionsRequest) -> Exception: +def _formatted_response( + response: DecisionsIRResponse, call: _DecisionsCall +) -> DecisionsResponse | OpenAIDecisionResponse: + match call.request: + case DecisionsRequestBody(): + return ir_to_systemone_response(response, call.ir_request) + case OpenAIDecisionRequestBody(): + return ir_to_openai_response(response, call.ir_request, call.requested_model) + case _: + assert_never(call.request) + + +def _map_upstream_exception(error: Exception, call: _DecisionsCall) -> Exception: + if isinstance(error, BaseLLMException) and error.status_code_is_synthesized: + provider_label: Final = f"{call.custom_llm_provider[0].upper()}{call.custom_llm_provider[1:]}Exception" + return litellm.APIConnectionError( + message=f"{provider_label} - {error.message}", + llm_provider=call.custom_llm_provider, + model=f"{call.custom_llm_provider}/{call.model}", + ) return litellm.exception_type( - model=f"{prepared.provider}/{prepared.upstream_model}", - custom_llm_provider=prepared.provider, + model=f"{call.custom_llm_provider}/{call.model}", + custom_llm_provider=call.custom_llm_provider, original_exception=error, ) @@ -227,73 +245,91 @@ def _map_upstream_exception(error: Exception, prepared: _PreparedDecisionsReques @client async def adecisions( model: str, - state: DecisionsJSON, - questions: Mapping[str, DecisionQuestion | Mapping[str, object]], + state: DecisionsJSON | None = None, + questions: DecisionsQuestions | None = None, api_key: str | None = None, api_base: str | None = None, timeout: float | httpx.Timeout | None = None, custom_llm_provider: str | None = None, extra_headers: Mapping[str, str] | None = None, + input: OpenAIDecisionInput | None = None, + safety_identifier: str | None = None, **kwargs: object, -) -> DecisionsResponse: - prepared: Final = _prepare_request( +) -> DecisionsResponse | OpenAIDecisionResponse: + call: Final = _prepare_call( model=model, state=state, questions=questions, + decision_input=input, + safety_identifier=safety_identifier, api_key=api_key, api_base=api_base, + timeout=timeout, custom_llm_provider=custom_llm_provider, extra_headers=extra_headers, + kwargs=kwargs, ) - logging_obj: Final = _log_request(prepared, kwargs) try: - handler: Final = get_async_httpx_client(llm_provider=prepared.provider) - response: Final = await handler.post( - prepared.url, - json=dict(prepared.body), - headers=dict(prepared.headers), - timeout=timeout, - logging_obj=logging_obj, + response: Final = await _HANDLER.adecisions( + model=call.model, + custom_llm_provider=call.custom_llm_provider, + logging_obj=call.logging_obj, + provider_config=call.provider_config, + request=call.ir_request, + body=call.body, + api_base=call.api_base, + api_key=call.api_key, + headers=call.headers, + timeout=call.timeout, ) - return _parse_response(response=response, prepared=prepared) except Exception as error: - raise _map_upstream_exception(error, prepared) from error + raise _map_upstream_exception(error, call) from error + return _format_response(response, call) @client def decisions( model: str, - state: DecisionsJSON, - questions: Mapping[str, DecisionQuestion | Mapping[str, object]], + state: DecisionsJSON | None = None, + questions: DecisionsQuestions | None = None, api_key: str | None = None, api_base: str | None = None, timeout: float | httpx.Timeout | None = None, custom_llm_provider: str | None = None, extra_headers: Mapping[str, str] | None = None, + input: OpenAIDecisionInput | None = None, + safety_identifier: str | None = None, **kwargs: object, -) -> DecisionsResponse: - prepared: Final = _prepare_request( +) -> DecisionsResponse | OpenAIDecisionResponse: + call: Final = _prepare_call( model=model, state=state, questions=questions, + decision_input=input, + safety_identifier=safety_identifier, api_key=api_key, api_base=api_base, + timeout=timeout, custom_llm_provider=custom_llm_provider, extra_headers=extra_headers, + kwargs=kwargs, ) - logging_obj: Final = _log_request(prepared, kwargs) try: - handler: Final = get_httpx_client() - response: Final = handler.post( - prepared.url, - json=dict(prepared.body), - headers=dict(prepared.headers), - timeout=timeout, - logging_obj=logging_obj, + response: Final = _HANDLER.decisions( + model=call.model, + custom_llm_provider=call.custom_llm_provider, + logging_obj=call.logging_obj, + provider_config=call.provider_config, + request=call.ir_request, + body=call.body, + api_base=call.api_base, + api_key=call.api_key, + headers=call.headers, + timeout=call.timeout, ) - return _parse_response(response=response, prepared=prepared) except Exception as error: - raise _map_upstream_exception(error, prepared) from error + raise _map_upstream_exception(error, call) from error + return _format_response(response, call) -__all__ = ["DECISIONS_ENDPOINTS", "adecisions", "decisions"] +__all__ = ["adecisions", "decisions"] diff --git a/litellm/google_genai/streaming_iterator.py b/litellm/google_genai/streaming_iterator.py index 27883c938db..64508ec93a6 100644 --- a/litellm/google_genai/streaming_iterator.py +++ b/litellm/google_genai/streaming_iterator.py @@ -90,7 +90,7 @@ class BaseGoogleGenAIGenerateContentStreamingIterator: end_time: Final = datetime.now() asyncio.create_task( - PassThroughStreamingHandler._route_streaming_logging_to_handler( + PassThroughStreamingHandler.route_streaming_logging_to_handler( litellm_logging_obj=self.litellm_logging_obj, passthrough_success_handler_obj=GLOBAL_PASS_THROUGH_SUCCESS_HANDLER_OBJ, url_route="/v1/generateContent", diff --git a/litellm/integrations/SlackAlerting/slack_alerting.py b/litellm/integrations/SlackAlerting/slack_alerting.py index 704a25b0457..7e9a8fe877c 100644 --- a/litellm/integrations/SlackAlerting/slack_alerting.py +++ b/litellm/integrations/SlackAlerting/slack_alerting.py @@ -1845,7 +1845,7 @@ Model Info: try: from litellm.proxy.spend_tracking.spend_management_endpoints import ( - _get_spend_report_for_time_range, + get_spend_report_for_time_range, ) # Parse the time range @@ -1862,7 +1862,7 @@ Model Info: if await self.internal_usage_cache.async_get_cache(key=_event_cache_key): return - _resp: Final = await _get_spend_report_for_time_range( + _resp: Final = await get_spend_report_for_time_range( start_date=start_date.strftime("%Y-%m-%d"), end_date=todays_date.strftime("%Y-%m-%d"), ) @@ -1909,7 +1909,7 @@ Model Info: from calendar import monthrange from litellm.proxy.spend_tracking.spend_management_endpoints import ( - _get_spend_report_for_time_range, + get_spend_report_for_time_range, ) todays_date: Final = datetime.datetime.now().date() @@ -1921,7 +1921,7 @@ Model Info: if await self.internal_usage_cache.async_get_cache(key=_event_cache_key): return - _resp: Final = await _get_spend_report_for_time_range( + _resp: Final = await get_spend_report_for_time_range( start_date=first_day_of_month.strftime("%Y-%m-%d"), end_date=last_day_of_month.strftime("%Y-%m-%d"), ) diff --git a/litellm/integrations/gcs_pubsub/pub_sub.py b/litellm/integrations/gcs_pubsub/pub_sub.py index 293be174811..c799245c215 100644 --- a/litellm/integrations/gcs_pubsub/pub_sub.py +++ b/litellm/integrations/gcs_pubsub/pub_sub.py @@ -44,9 +44,9 @@ class GcsPubSubLogger(CustomBatchLogger): topic_id (str): Pub/Sub topic ID credentials_path (str, optional): Path to Google Cloud credentials JSON file """ - from litellm.proxy.utils import _premium_user_check + from litellm.proxy.utils import premium_user_check - _premium_user_check() + premium_user_check() self.async_httpx_client = get_async_httpx_client(llm_provider=httpxSpecialProvider.LoggingCallback) @@ -107,9 +107,9 @@ class GcsPubSubLogger(CustomBatchLogger): from litellm.proxy.spend_tracking.spend_tracking_utils import ( get_logging_payload, ) - from litellm.proxy.utils import _premium_user_check + from litellm.proxy.utils import premium_user_check - _premium_user_check() + premium_user_check() try: verbose_logger.debug("PubSub: Logging - Enters logging function for model %s", kwargs) diff --git a/litellm/integrations/prometheus.py b/litellm/integrations/prometheus.py index ede75381277..e11ce2bf472 100644 --- a/litellm/integrations/prometheus.py +++ b/litellm/integrations/prometheus.py @@ -3833,7 +3833,7 @@ class PrometheusLogger(CustomLogger): """ from litellm.constants import UI_SESSION_TOKEN_TEAM_ID from litellm.proxy.management_endpoints.key_management_endpoints import ( - _list_key_helper, + list_key_helper, ) from litellm.proxy.proxy_server import prisma_client @@ -3847,7 +3847,7 @@ class PrometheusLogger(CustomLogger): list[str | UserAPIKeyAuth | LiteLLM_DeletedVerificationToken], int | None, ]: - key_list_response: Final = await _list_key_helper( + key_list_response: Final = await list_key_helper( prisma_client=prisma_client, page=page, size=page_size, diff --git a/litellm/integrations/shadow_eval_logger.py b/litellm/integrations/shadow_eval_logger.py index 2169a89e8d7..944ee722517 100644 --- a/litellm/integrations/shadow_eval_logger.py +++ b/litellm/integrations/shadow_eval_logger.py @@ -633,9 +633,9 @@ async def _key_or_team_is_over_budget(metadata: Mapping[str, object]) -> bool: from litellm.exceptions import BudgetExceededError from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.auth.auth_checks import ( - _team_max_budget_check, - _virtual_key_max_budget_check, get_team_object, + team_max_budget_check, + virtual_key_max_budget_check, ) from litellm.proxy.proxy_server import prisma_client, proxy_logging_obj, user_api_key_cache except ImportError: @@ -645,7 +645,7 @@ async def _key_or_team_is_over_budget(metadata: Mapping[str, object]) -> bool: if not isinstance(auth, UserAPIKeyAuth): return False try: - await _virtual_key_max_budget_check(valid_token=auth, proxy_logging_obj=proxy_logging_obj) + await virtual_key_max_budget_check(valid_token=auth, proxy_logging_obj=proxy_logging_obj) if auth.team_id: team: Final = await get_team_object( team_id=auth.team_id, @@ -653,7 +653,7 @@ async def _key_or_team_is_over_budget(metadata: Mapping[str, object]) -> bool: user_api_key_cache=user_api_key_cache, check_cache_only=True, ) - await _team_max_budget_check(team_object=team, valid_token=auth, proxy_logging_obj=proxy_logging_obj) + await team_max_budget_check(team_object=team, valid_token=auth, proxy_logging_obj=proxy_logging_obj) except BudgetExceededError: return True except Exception as e: # noqa: BLE001 # advisory gate: a failed read must not block sampling diff --git a/litellm/litellm_core_utils/custom_logger_registry.py b/litellm/litellm_core_utils/custom_logger_registry.py index 1d277995211..8367c9ee00e 100644 --- a/litellm/litellm_core_utils/custom_logger_registry.py +++ b/litellm/litellm_core_utils/custom_logger_registry.py @@ -53,8 +53,14 @@ from litellm.integrations.vector_store_integrations.vector_store_pre_call_hook i VectorStorePreCallHook, ) from litellm.integrations.zerobus import ZerobusLogger -from litellm.proxy.hooks.dynamic_rate_limiter import _PROXY_DynamicRateLimitHandler -from litellm.proxy.hooks.dynamic_rate_limiter_v3 import _PROXY_DynamicRateLimitHandlerV3 +from litellm.proxy.hooks.dynamic_rate_limiter import ( # noqa: F401 # legacy module exports + PROXY_DynamicRateLimitHandler, + _PROXY_DynamicRateLimitHandler, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export +) +from litellm.proxy.hooks.dynamic_rate_limiter_v3 import ( # noqa: F401 # legacy module exports + PROXY_DynamicRateLimitHandlerV3, + _PROXY_DynamicRateLimitHandlerV3, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export +) class CustomLoggerRegistry: @@ -101,8 +107,8 @@ class CustomLoggerRegistry: "pointfive": PointFiveLogger, "zerobus": ZerobusLogger, "aws_sqs": SQSLogger, - "dynamic_rate_limiter": _PROXY_DynamicRateLimitHandler, - "dynamic_rate_limiter_v3": _PROXY_DynamicRateLimitHandlerV3, + "dynamic_rate_limiter": PROXY_DynamicRateLimitHandler, + "dynamic_rate_limiter_v3": PROXY_DynamicRateLimitHandlerV3, "vector_store_pre_call_hook": VectorStorePreCallHook, "dotprompt": DotpromptManager, "bitbucket": BitBucketPromptManager, diff --git a/litellm/litellm_core_utils/exception_mapping_utils.py b/litellm/litellm_core_utils/exception_mapping_utils.py index 7fb28fb418b..b5afa962cdc 100644 --- a/litellm/litellm_core_utils/exception_mapping_utils.py +++ b/litellm/litellm_core_utils/exception_mapping_utils.py @@ -2147,7 +2147,10 @@ def _map_openrouter_exception( exception_provider: str, extra_information: str, ) -> None: - if hasattr(original_exception, "status_code"): + received_status: Final = hasattr(original_exception, "status_code") and not getattr( + original_exception, "status_code_is_synthesized", False + ) + if received_status: if original_exception.status_code == 400: raise BadRequestError( message=f"{exception_provider} - {error_str}", diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index 3a2a25a13e6..b79e7490871 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -120,7 +120,7 @@ from litellm.llms.base_llm.search.transformation import SearchResponse from litellm.responses.utils import ResponseAPILoggingUtils from litellm.types.agents import LiteLLMSendMessageResponse from litellm.types.containers.main import ContainerObject -from litellm.types.decisions import DecisionsResponse +from litellm.types.decisions import DecisionsResponse, OpenAIDecisionResponse from litellm.types.integrations.s3_v2 import S3PartitionGranularity from litellm.types.interactions import ( InteractionsAPIResponse, @@ -2650,6 +2650,7 @@ class Logging(LiteLLMLoggingBaseClass): or isinstance(logging_result, OCRResponse) # OCR or isinstance(logging_result, SearchResponse) # Search API or isinstance(logging_result, DecisionsResponse) + or isinstance(logging_result, OpenAIDecisionResponse) or ( isinstance(logging_result, InteractionsAPIResponse) and logging_result.usage is not None @@ -4969,17 +4970,17 @@ def _init_custom_logger_compatible_class( return _otel_logger elif logging_integration == "dynamic_rate_limiter": from litellm.proxy.hooks.dynamic_rate_limiter import ( - _PROXY_DynamicRateLimitHandler, + PROXY_DynamicRateLimitHandler, ) for callback in _in_memory_loggers: - if isinstance(callback, _PROXY_DynamicRateLimitHandler): + if isinstance(callback, PROXY_DynamicRateLimitHandler): return callback if internal_usage_cache is None: raise Exception(f"Internal Error: Cache cannot be empty - internal_usage_cache={internal_usage_cache}") - dynamic_rate_limiter_obj: Final = _PROXY_DynamicRateLimitHandler(internal_usage_cache=internal_usage_cache) + dynamic_rate_limiter_obj: Final = PROXY_DynamicRateLimitHandler(internal_usage_cache=internal_usage_cache) if llm_router is not None and isinstance(llm_router, litellm.Router): dynamic_rate_limiter_obj.update_variables(llm_router=llm_router) @@ -4987,17 +4988,19 @@ def _init_custom_logger_compatible_class( return dynamic_rate_limiter_obj elif logging_integration == "dynamic_rate_limiter_v3": from litellm.proxy.hooks.dynamic_rate_limiter_v3 import ( - _PROXY_DynamicRateLimitHandlerV3, + PROXY_DynamicRateLimitHandlerV3, ) for callback in _in_memory_loggers: - if isinstance(callback, _PROXY_DynamicRateLimitHandlerV3): + if isinstance(callback, PROXY_DynamicRateLimitHandlerV3): return callback if internal_usage_cache is None: raise Exception(f"Internal Error: Cache cannot be empty - internal_usage_cache={internal_usage_cache}") - dynamic_rate_limiter_obj_v3 = _PROXY_DynamicRateLimitHandlerV3(internal_usage_cache=internal_usage_cache) + dynamic_rate_limiter_obj_v3: Final = PROXY_DynamicRateLimitHandlerV3( + internal_usage_cache=internal_usage_cache + ) if llm_router is not None and isinstance(llm_router, litellm.Router): dynamic_rate_limiter_obj_v3.update_variables(llm_router=llm_router) @@ -5546,19 +5549,19 @@ def get_custom_logger_compatible_class( elif logging_integration == "dynamic_rate_limiter": from litellm.proxy.hooks.dynamic_rate_limiter import ( - _PROXY_DynamicRateLimitHandler, + PROXY_DynamicRateLimitHandler, ) for callback in _in_memory_loggers: - if isinstance(callback, _PROXY_DynamicRateLimitHandler): + if isinstance(callback, PROXY_DynamicRateLimitHandler): return callback elif logging_integration == "dynamic_rate_limiter_v3": from litellm.proxy.hooks.dynamic_rate_limiter_v3 import ( - _PROXY_DynamicRateLimitHandlerV3, + PROXY_DynamicRateLimitHandlerV3, ) for callback in _in_memory_loggers: - if isinstance(callback, _PROXY_DynamicRateLimitHandlerV3): + if isinstance(callback, PROXY_DynamicRateLimitHandlerV3): return callback elif logging_integration == "langtrace": diff --git a/litellm/litellm_core_utils/prompt_templates/common_utils.py b/litellm/litellm_core_utils/prompt_templates/common_utils.py index 0e0f51e5f8b..95d1c98b13d 100644 --- a/litellm/litellm_core_utils/prompt_templates/common_utils.py +++ b/litellm/litellm_core_utils/prompt_templates/common_utils.py @@ -565,9 +565,9 @@ def update_messages_with_model_file_ids( } """ from litellm.proxy.openai_files_endpoints.common_utils import ( - _is_base64_encoded_unified_file_id, convert_b64_uid_to_unified_uid, get_original_file_id, + is_base64_encoded_unified_file_id, is_model_embedded_id, ) @@ -603,7 +603,7 @@ def update_messages_with_model_file_ids( if model_file_id_mapping and model_id is not None else None ) - if not provider_file_id and _is_base64_encoded_unified_file_id(file_id): + if not provider_file_id and is_base64_encoded_unified_file_id(file_id): unified_file_id = convert_b64_uid_to_unified_uid(file_id) if "llm_output_file_id," in unified_file_id: provider_file_id = unified_file_id.split("llm_output_file_id,")[1].split(";")[0] @@ -634,9 +634,9 @@ def update_responses_input_with_model_file_ids( Format: {"litellm_file_id": {"model_id": "provider_file_id"}} """ from litellm.proxy.openai_files_endpoints.common_utils import ( - _is_base64_encoded_unified_file_id, convert_b64_uid_to_unified_uid, get_original_file_id, + is_base64_encoded_unified_file_id, is_model_embedded_id, ) @@ -671,7 +671,7 @@ def update_responses_input_with_model_file_ids( updated_content.append(updated_content_item) else: # Check if this is a base64-encoded unified file ID without mapping - is_unified_file_id = _is_base64_encoded_unified_file_id(file_id) + is_unified_file_id = is_base64_encoded_unified_file_id(file_id) if is_unified_file_id: # Fallback: decode unified file ID unified_file_id = convert_b64_uid_to_unified_uid(file_id) diff --git a/litellm/litellm_core_utils/prompt_templates/factory.py b/litellm/litellm_core_utils/prompt_templates/factory.py index 4580f9bd01b..3b0cf0ce238 100644 --- a/litellm/litellm_core_utils/prompt_templates/factory.py +++ b/litellm/litellm_core_utils/prompt_templates/factory.py @@ -206,36 +206,108 @@ def _handle_ollama_system_message(messages: list, prompt: str, msg_i: int) -> tu return system_content_str, msg_i +_OLLAMA_USER_ROLES: Final = frozenset({"user", "tool", "function"}) + + +def _ollama_bad_message(model: str, message: AllMessageValues, msg_i: int, detail: str) -> "litellm.BadRequestError": + return litellm.BadRequestError( + message=BAD_MESSAGE_ERROR_STR + f"the {message['role']} message at index {msg_i} {detail}", + model=model, + llm_provider="ollama", + ) + + +def _ollama_content_part(model: str, message: AllMessageValues, part: object, msg_i: int) -> tuple[str, str]: + match part: + case {"type": "text", "text": str() as text}: + return text, "" + case {"type": "text", "text": bad_text}: + raise _ollama_bad_message( + model, message, msg_i, f"has a {type(bad_text).__name__} text part; text must be a string" + ) + case {"type": "text"}: + raise _ollama_bad_message(model, message, msg_i, "has a text part with no text; text must be a string") + case {"type": "image_url", "image_url": str() as image_url}: + return "", image_url + case {"type": "image_url", "image_url": {"url": str() as image_url}}: + return "", image_url + case {"type": "image_url", "image_url": dict()}: + raise _ollama_bad_message( + model, + message, + msg_i, + "has an image_url object without a url string; image_url must be a URL string or an object with a url", + ) + case {"type": "image_url", "image_url": bad_image_url}: + raise _ollama_bad_message( + model, + message, + msg_i, + f"has a {type(bad_image_url).__name__} image_url; image_url must be a URL string or an object with a url", + ) + case {"type": "image_url"}: + raise _ollama_bad_message( + model, + message, + msg_i, + "has an image_url part with no image_url; image_url must be a URL string or an object with a url", + ) + case Mapping(): + return "", "" + case _: + raise _ollama_bad_message( + model, message, msg_i, f"has a {type(part).__name__} content part; content parts must be objects" + ) + + +def _ollama_user_message_parts( + model: str, message: AllMessageValues, msg_i: int +) -> tuple[tuple[str, ...], tuple[str, ...]]: + msg_content: Final = message.get("content") + if msg_content is None: + return (), () + if isinstance(msg_content, str): + return ((msg_content,) if msg_content else ()), () + if not isinstance(msg_content, list): + raise _ollama_bad_message( + model, + message, + msg_i, + f"has {type(msg_content).__name__} content; content must be a string or a list of content parts", + ) + parts: Final = tuple(_ollama_content_part(model, message, part, msg_i) for part in msg_content) + texts: Final = tuple(text for text, _ in parts if text) + image_urls: Final = tuple(image_url for _, image_url in parts if image_url) + return texts, image_urls + + +def _ollama_user_turn(model: str, messages: Sequence[AllMessageValues], msg_i: int) -> tuple[str, tuple[str, ...], int]: + user_run: Final = tuple( + itertools.takewhile( + lambda message: message["role"] in _OLLAMA_USER_ROLES, (messages[i] for i in range(msg_i, len(messages))) + ) + ) + user_parts: Final = tuple( + _ollama_user_message_parts(model, message, msg_i + offset) for offset, message in enumerate(user_run) + ) + user_content_str: Final = "\n".join("\n".join(texts) for texts, _ in user_parts if texts) + image_urls: Final = tuple(itertools.chain.from_iterable(urls for _, urls in user_parts)) + return user_content_str, image_urls, len(user_run) + + def ollama_pt( model: str, messages: list ) -> ( str | OllamaVisionModelObject ): # https://github.com/ollama/ollama/blob/af4cf55884ac54b9e637cd71dadfe9b7a5685877/docs/modelfile.md#template - user_message_types: Final = {"user", "tool", "function"} msg_i = 0 images: Final = [] prompt = "" while msg_i < len(messages): init_msg_i = msg_i - user_content_str = "" - ## MERGE CONSECUTIVE USER CONTENT ## - while msg_i < len(messages) and messages[msg_i]["role"] in user_message_types: - msg_content = messages[msg_i].get("content") - if msg_content: - if isinstance(msg_content, list): - for m in msg_content: - if m.get("type", "") == "image_url": - if isinstance(m["image_url"], str): - images.append(m["image_url"]) - elif isinstance(m["image_url"], dict): - images.append(m["image_url"]["url"]) - elif m.get("type", "") == "text": - user_content_str += m["text"] - else: - # Tool message content will always be a string - user_content_str += msg_content - - msg_i += 1 + user_content_str, image_urls, user_run_len = _ollama_user_turn(model, messages, msg_i) + images.extend(image_urls) + msg_i += user_run_len if user_content_str: prompt += f"### User:\n{user_content_str}\n\n" diff --git a/litellm/llms/anthropic/chat/guardrail_translation/handler.py b/litellm/llms/anthropic/chat/guardrail_translation/handler.py index 6d26dd31b5e..bb6d4030ab8 100644 --- a/litellm/llms/anthropic/chat/guardrail_translation/handler.py +++ b/litellm/llms/anthropic/chat/guardrail_translation/handler.py @@ -322,7 +322,7 @@ class AnthropicMessagesHandler(BaseTranslation): if not chunks: return None try: - return AnthropicPassthroughLoggingHandler._build_usage_only_response_from_chunks( + return AnthropicPassthroughLoggingHandler.build_usage_only_response_from_chunks( all_chunks=chunks, model=str((request_data or {}).get("model") or ""), ) @@ -1226,7 +1226,7 @@ class AnthropicMessagesHandler(BaseTranslation): has_ended: Final = self._check_streaming_has_ended(responses_so_far) if has_ended: # build the model response from the responses_so_far - built_response: Final = AnthropicPassthroughLoggingHandler._build_complete_streaming_response( + built_response: Final = AnthropicPassthroughLoggingHandler.build_complete_streaming_response( all_chunks=responses_so_far, litellm_logging_obj=cast("LiteLLMLoggingObj", litellm_logging_obj), model="", diff --git a/litellm/llms/anthropic/pass_through/adapters/streaming_iterator.py b/litellm/llms/anthropic/pass_through/adapters/streaming_iterator.py index 2b7f52e0ddc..5f6d644853a 100644 --- a/litellm/llms/anthropic/pass_through/adapters/streaming_iterator.py +++ b/litellm/llms/anthropic/pass_through/adapters/streaming_iterator.py @@ -69,6 +69,12 @@ def _error_status_and_message(exc: Exception) -> tuple[int, str]: return 500, str(exc) or "Upstream stream ended before completion" +def _provider_error(exc: Exception) -> Exception: + if isinstance(exc, MidStreamFallbackError) and exc.original_exception is not None: + return exc.original_exception + return exc + + def _mid_stream_error_sse_event(exc: Exception) -> bytes: from litellm.anthropic_interface.exceptions.exception_mapping_utils import ( anthropic_error_sse_frame, @@ -336,6 +342,7 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper): self._message_id: str = f"msg_{uuid.uuid4()}" if litellm_logging_obj is not None: litellm_logging_obj.record_streamed_anthropic_message_id(self._message_id) + self.litellm_logging_obj = litellm_logging_obj # Mapping of truncated tool names to original names (for OpenAI's 64-char limit) self.tool_name_mapping = tool_name_mapping or {} # Polyfill applied_edits on final message_delta. @@ -1041,7 +1048,25 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper): else: yield chunk except Exception as e: # noqa: BLE001 # boundary before the socket: any upstream failure becomes an Anthropic error event - verbose_logger.exception("Anthropic Adapter - mid-stream error, emitting Anthropic error event: %s", e) + verbose_logger.exception("Anthropic Adapter - mid-stream error: %s", e) + logging_obj: Final = self.litellm_logging_obj + provider_error: Final = _provider_error(e) + if logging_obj is not None and logging_obj.on_detached_stream_failure is not None: + if provider_error is e: + raise + raise provider_error from e + if logging_obj is not None: + try: + await logging_obj.dispatch_failure_handlers( + exception=provider_error, + traceback_exception=traceback.format_exc(), + prefer_async_handlers=True, + ) + except Exception as failure_handler_error: # noqa: BLE001 # a failing failure handler must not also drop the error frame + verbose_logger.exception( + "Anthropic Adapter - failure handler raised while reporting a mid-stream error: %s", + failure_handler_error, + ) yield _mid_stream_error_sse_event(e) def _increment_content_block_index(self): diff --git a/litellm/llms/anthropic/pass_through/context_management/editors/compact.py b/litellm/llms/anthropic/pass_through/context_management/editors/compact.py index 62826865894..f743320b702 100644 --- a/litellm/llms/anthropic/pass_through/context_management/editors/compact.py +++ b/litellm/llms/anthropic/pass_through/context_management/editors/compact.py @@ -228,7 +228,7 @@ async def _check_summary_model_access( try: from litellm.proxy._types import ProxyException from litellm.proxy.auth.auth_checks import ( - _can_object_call_model, + can_object_call_model, can_project_access_model, can_user_call_model, get_project_object, @@ -258,7 +258,7 @@ async def _check_summary_model_access( if not models: continue try: - _can_object_call_model( + can_object_call_model( model=summary_model, llm_router=llm_router, models=models, @@ -370,7 +370,7 @@ async def _check_summary_model_access( ) if member_allowed_models: try: - _can_object_call_model( + can_object_call_model( model=summary_model, llm_router=llm_router, models=list(member_allowed_models), @@ -558,7 +558,7 @@ async def _check_summary_model_rate_limit( limiter: Final[object] = get_proxy_hook("parallel_request_limiter") if get_proxy_hook is not None else None should_rate_limit_check: Final[_ShouldRateLimit | None] = getattr(limiter, "should_rate_limit", None) create_descriptors: Final[_CreateRateLimitDescriptors | None] = getattr( - limiter, "_create_rate_limit_descriptors", None + limiter, "create_rate_limit_descriptors", None ) add_team_descriptor: Final[_AddModelRateLimitDescriptor | None] = getattr( limiter, "_add_team_model_rate_limit_descriptor_from_metadata", None diff --git a/litellm/llms/anthropic/pass_through/messages/streaming_iterator.py b/litellm/llms/anthropic/pass_through/messages/streaming_iterator.py index 25e0f93b04c..f37e5190850 100644 --- a/litellm/llms/anthropic/pass_through/messages/streaming_iterator.py +++ b/litellm/llms/anthropic/pass_through/messages/streaming_iterator.py @@ -421,7 +421,7 @@ class BaseAnthropicMessagesStreamingIterator: if self.completion_start_time is not None: self.litellm_logging_obj.completion_start_time = self.completion_start_time self.litellm_logging_obj.model_call_details["completion_start_time"] = self.completion_start_time - logging_coroutine: Final = PassThroughStreamingHandler._route_streaming_logging_to_handler( + logging_coroutine: Final = PassThroughStreamingHandler.route_streaming_logging_to_handler( litellm_logging_obj=self.litellm_logging_obj, passthrough_success_handler_obj=GLOBAL_PASS_THROUGH_SUCCESS_HANDLER_OBJ, url_route="/v1/messages", diff --git a/litellm/llms/anthropic/wif.py b/litellm/llms/anthropic/wif.py index 2444a8f833e..b47c4f28eb2 100644 --- a/litellm/llms/anthropic/wif.py +++ b/litellm/llms/anthropic/wif.py @@ -65,6 +65,7 @@ _IDENTITY_SOURCE_PARAM: Final = "anthropic_identity_source" _IDENTITY_SOURCE_ENV: Final = "ANTHROPIC_IDENTITY_SOURCE" _IDENTITY_TOKEN_FILE_PARAM: Final = "anthropic_identity_token_file" _IDENTITY_TOKEN_PARAM: Final = "anthropic_identity_token" +_LEGACY_REF_PARAMS: Final = (_IDENTITY_TOKEN_FILE_PARAM, _IDENTITY_TOKEN_PARAM) # litellm_params key -> InternalIssuerSource/KeycloakSource field name. Every key here must # also be listed in ANTHROPIC_WIF_KWARGS_KEYS (types/workload_identity.py), which is what makes it @@ -139,7 +140,7 @@ def resolve_anthropic_wif_params(litellm_params: Mapping[str, object] | None) -> ) organization_id: Final = _config_value(litellm_params, "anthropic_organization_id", "ANTHROPIC_ORGANIZATION_ID") if federation_rule_id is None or organization_id is None: - _raise_if_identity_source_configured(litellm_params, federation_rule_id, organization_id) + _raise_if_federation_requested(litellm_params, federation_rule_id, organization_id) return None identity_source: Final = _resolve_identity_source(litellm_params) if identity_source is None: @@ -204,17 +205,15 @@ def _raise_unknown_source_kind(source_kind: str) -> NoReturn: ) -def _raise_if_identity_source_configured( +def _raise_if_federation_requested( litellm_params: Mapping[str, object] | None, federation_rule_id: str | None, organization_id: str | None ) -> None: - """A configured identity source is an explicit request to federate, so a missing rule or - organization id fails closed with the ids named, rather than silently skipping federation - and surfacing later as a missing API key.""" - source_kind: Final = _resolve_source_kind(litellm_params) - if source_kind is None: + """A configured identity source, or a token file or inline token set on the deployment, is an + explicit request to federate, so a missing rule or organization id fails closed with the ids + named, rather than silently skipping federation and surfacing later as a missing API key.""" + request: Final = _explicit_federation_request(litellm_params) + if request is None: return - if source_kind not in {kind.value for kind in AnthropicIdentitySourceKind}: - _raise_unknown_source_kind(source_kind) missing: Final = tuple( param for param, value in ( @@ -225,7 +224,7 @@ def _raise_if_identity_source_configured( ) raise litellm.AuthenticationError( message=( - f"{_IDENTITY_SOURCE_PARAM} is {source_kind!r}, but {' and '.join(missing)} " + f"{request}, but {' and '.join(missing)} " f"{'is' if len(missing) == 1 else 'are'} not set. {_MISSING_IDS_HINT}" ), llm_provider="anthropic", @@ -233,14 +232,25 @@ def _raise_if_identity_source_configured( ) +def _explicit_federation_request(litellm_params: Mapping[str, object] | None) -> str | None: + source_kind: Final = _resolve_source_kind(litellm_params) + if source_kind is None: + legacy_ref_param: Final = _legacy_ref_param(litellm_params) + return None if legacy_ref_param is None else f"{legacy_ref_param} is set" + if source_kind not in {kind.value for kind in AnthropicIdentitySourceKind}: + _raise_unknown_source_kind(source_kind) + return f"{_IDENTITY_SOURCE_PARAM} is {source_kind!r}" + + def _resolve_source_kind(litellm_params: Mapping[str, object] | None) -> str | None: param_kind: Final = _param_str(litellm_params, _IDENTITY_SOURCE_PARAM) if param_kind is not None: return param_kind - has_param_legacy_ref: Final = any( - _param_str(litellm_params, key) is not None for key in (_IDENTITY_TOKEN_FILE_PARAM, _IDENTITY_TOKEN_PARAM) - ) - return None if has_param_legacy_ref else _env_str(_IDENTITY_SOURCE_ENV) + return None if _legacy_ref_param(litellm_params) is not None else _env_str(_IDENTITY_SOURCE_ENV) + + +def _legacy_ref_param(litellm_params: Mapping[str, object] | None) -> str | None: + return next((key for key in _LEGACY_REF_PARAMS if _param_str(litellm_params, key) is not None), None) def _reject_foreign_variant_fields( diff --git a/litellm/llms/base_llm/decisions/__init__.py b/litellm/llms/base_llm/decisions/__init__.py index c18ac9b00f2..313e57e9762 100644 --- a/litellm/llms/base_llm/decisions/__init__.py +++ b/litellm/llms/base_llm/decisions/__init__.py @@ -1,3 +1,3 @@ -from .transformation import DecisionsProviderConfig, JevCompatibleDecisionsEndpoint +from .transformation import BaseDecisionsConfig -__all__ = ["DecisionsProviderConfig", "JevCompatibleDecisionsEndpoint"] +__all__ = ["BaseDecisionsConfig"] diff --git a/litellm/llms/base_llm/decisions/systemone.py b/litellm/llms/base_llm/decisions/systemone.py new file mode 100644 index 00000000000..43be31b0c30 --- /dev/null +++ b/litellm/llms/base_llm/decisions/systemone.py @@ -0,0 +1,247 @@ +"""The Jev / System One wire shape and its translation to and from the OpenAI Decisions shape. + +System One (TypeSafe, Perplexity, OpenRouter, Cloudflare Clef, Strands Decider) takes +{"model", "state", "questions": {name: question}} and answers with {"model", "answers": {name: answer}, "usage"}. +Predicates are `noul` questions, choice options are a `criteria` map, score levels are a `criteria` list. +""" + +import itertools +from collections.abc import Mapping, Sequence +from typing import Final, Literal, TypeAlias + +from pydantic import ConfigDict, TypeAdapter +from typing_extensions import assert_never + +from litellm.llms.base_llm.chat.transformation import BaseLLMException +from litellm.types.llms.base import LiteLLMPydanticObjectBase +from litellm.types.openai_decisions import ( + ChoiceAnswer, + ChoiceProbability, + ChoiceQuestion, + DecisionAnswer, + DecisionChoice, + DecisionInput, + DecisionInputMessage, + DecisionInputPart, + DecisionInputTokensDetails, + DecisionOutputTokensDetails, + DecisionQuestion, + DecisionsRequest, + DecisionsRequestBody, + DecisionsResponse, + DecisionUsage, + PredicateAnswer, + PredicateQuestion, + ScoreAnswer, + ScoreProbability, + ScoreQuestion, +) + + +class SystemOneObjectBase(LiteLLMPydanticObjectBase): + model_config = ConfigDict(extra="allow", frozen=True) + + +class SystemOneNoulAnswer(SystemOneObjectBase): + type: Literal["noul"] + noul: float + + +class SystemOneChoiceAnswer(SystemOneObjectBase): + type: Literal["choice"] + choice: str + confidence: float + probabilities: Mapping[str, float] + + +class SystemOneScoreAnswer(SystemOneObjectBase): + type: Literal["score"] + score: float + confidence: float + probabilities: Mapping[str, float] + + +SystemOneAnswer: TypeAlias = SystemOneNoulAnswer | SystemOneChoiceAnswer | SystemOneScoreAnswer + + +class SystemOneUsage(SystemOneObjectBase): + input_tokens: int = 0 + output_tokens: int = 0 + + +class SystemOneResponse(SystemOneObjectBase): + model: str | None = None + answers: Mapping[str, SystemOneAnswer] + usage: SystemOneUsage | None = None + + +SYSTEM_ONE_RESPONSE_ADAPTER: Final[TypeAdapter[SystemOneResponse]] = TypeAdapter(SystemOneResponse) + + +def _unsupported(what: str, custom_llm_provider: str) -> BaseLLMException: + return BaseLLMException( + status_code=400, + message=f"Decisions provider '{custom_llm_provider}' does not support {what}", + ) + + +def to_system_one_request(model: str, body: DecisionsRequestBody, custom_llm_provider: str) -> dict[str, object]: + keys: Final = question_keys(body.questions, custom_llm_provider) + return { + "model": model, + "state": _state(body.input, custom_llm_provider), + "questions": { + key: _question(question, custom_llm_provider) for key, question in zip(keys, body.questions, strict=True) + }, + } + + +def question_keys(questions: Sequence[DecisionQuestion], custom_llm_provider: str) -> tuple[str, ...]: + """System One keys questions and answers by name, so unnamed questions get a positional key.""" + names: Final = tuple(question.name for question in questions if question.name is not None) + if len(set(names)) != len(names): + raise BaseLLMException( + status_code=400, + message=f"Decisions provider '{custom_llm_provider}' requires a unique name per question", + ) + taken: Final = frozenset(names) + return tuple( + question.name if question.name is not None else _positional_key(index, taken) + for index, question in enumerate(questions) + ) + + +def _positional_key(index: int, taken: frozenset[str]) -> str: + candidates: Final = (f"{'_' * depth}q{index}" for depth in itertools.count()) + return next(key for key in candidates if key not in taken) + + +def _state(input_value: DecisionInput, custom_llm_provider: str) -> str: + if isinstance(input_value, str): + return input_value + return "\n".join(_message_text(message, custom_llm_provider) for message in input_value) + + +def _message_text(message: DecisionInputMessage, custom_llm_provider: str) -> str: + if isinstance(message.content, str): + return message.content + return "\n".join(_part_text(part, custom_llm_provider) for part in message.content) + + +def _part_text(part: DecisionInputPart, custom_llm_provider: str) -> str: + if part.type != "input_text": + raise _unsupported("input_image parts", custom_llm_provider) + return part.text + + +def _question(question: DecisionQuestion, custom_llm_provider: str) -> dict[str, object]: + match question: + case PredicateQuestion(): + return {"type": "noul", "instructions": question.instructions} + case ChoiceQuestion(): + return { + "type": "choice", + "instructions": question.instructions, + "criteria": _choice_criteria(question, custom_llm_provider), + } + case ScoreQuestion(): + return { + "type": "score", + "instructions": question.instructions, + "criteria": [_level_text(level.label, level.description) for level in question.levels], + } + case _: + assert_never(question) + + +def _choice_criteria(question: ChoiceQuestion, custom_llm_provider: str) -> dict[str, str | None]: + values: Final = tuple(_choice_key(choice, custom_llm_provider) for choice in question.choices) + if len(set(values)) != len(values): + raise _unsupported("repeated choice values", custom_llm_provider) + return {value: choice.description for value, choice in zip(values, question.choices, strict=True)} + + +def _choice_key(choice: DecisionChoice, custom_llm_provider: str) -> str: + if not isinstance(choice.value, str): + raise _unsupported("boolean choice values", custom_llm_provider) + return choice.value + + +def _level_text(label: str, description: str | None) -> str: + return description if description is not None else label + + +def to_decisions_response( + system_one: SystemOneResponse, + request: DecisionsRequest, + custom_llm_provider: str, +) -> DecisionsResponse: + keys: Final = question_keys(request.body.questions, custom_llm_provider) + return DecisionsResponse( + model=system_one.model if system_one.model is not None else request.model, + answers=[ + _answer(key, question, system_one.answers, custom_llm_provider) + for key, question in zip(keys, request.body.questions, strict=True) + ], + usage=_usage(system_one.usage), + ) + + +def _answer( + key: str, + question: DecisionQuestion, + answers: Mapping[str, SystemOneAnswer], + custom_llm_provider: str, +) -> DecisionAnswer: + match question, answers.get(key): + case PredicateQuestion(), SystemOneNoulAnswer() as answer: + return PredicateAnswer(type="predicate", name=question.name, probability=answer.noul) + case ChoiceQuestion(), SystemOneChoiceAnswer() as answer: + return ChoiceAnswer( + type="choice", + name=question.name, + choice=answer.choice, + probabilities=_choice_probabilities(question, answer), + confidence=answer.confidence, + ) + case ScoreQuestion(), SystemOneScoreAnswer() as answer: + return ScoreAnswer( + type="score", + name=question.name, + score=answer.score, + probabilities=_score_probabilities(question, answer), + confidence=answer.confidence, + ) + case _: + raise BaseLLMException( + status_code=500, + message=( + f"Decisions provider '{custom_llm_provider}' returned no {question.type} answer " + f"for question '{key}'" + ), + ) + + +def _choice_probabilities(question: ChoiceQuestion, answer: SystemOneChoiceAnswer) -> list[ChoiceProbability]: + return [ + ChoiceProbability(value=choice.value, probability=answer.probabilities.get(str(choice.value), 0.0)) + for choice in question.choices + ] + + +def _score_probabilities(question: ScoreQuestion, answer: SystemOneScoreAnswer) -> list[ScoreProbability]: + return [ + ScoreProbability(value=index, label=level.label, probability=answer.probabilities.get(str(index), 0.0)) + for index, level in enumerate(question.levels) + ] + + +def _usage(usage: SystemOneUsage | None) -> DecisionUsage: + counted: Final = usage if usage is not None else SystemOneUsage() + return DecisionUsage( + input_tokens=counted.input_tokens, + input_tokens_details=DecisionInputTokensDetails(cached_tokens=0, cache_write_tokens=0), + output_tokens=counted.output_tokens, + output_tokens_details=DecisionOutputTokensDetails(reasoning_tokens=0), + total_tokens=counted.input_tokens + counted.output_tokens, + ) diff --git a/litellm/llms/base_llm/decisions/transformation.py b/litellm/llms/base_llm/decisions/transformation.py index d4fcea24793..a8d14ca59a3 100644 --- a/litellm/llms/base_llm/decisions/transformation.py +++ b/litellm/llms/base_llm/decisions/transformation.py @@ -1,20 +1,322 @@ -from dataclasses import dataclass -from typing import Protocol +import json +from abc import ABC +from collections.abc import Mapping, Sequence +from types import MappingProxyType +from typing import Final + +import httpx +from pydantic import TypeAdapter, ValidationError +from typing_extensions import assert_never + +from litellm.llms.base_llm.chat.transformation import BaseLLMException +from litellm.secret_managers.main import get_secret_str +from litellm.types.decisions import ( + ChoiceAnswer, + ChoiceQuestion, + DecisionAnswer, + DecisionQuestion, + DecisionsIRAnswer, + DecisionsIRChoiceAnswer, + DecisionsIRChoiceOption, + DecisionsIRChoiceProbability, + DecisionsIRChoiceQuestion, + DecisionsIRMessages, + DecisionsIRPredicateAnswer, + DecisionsIRPredicateQuestion, + DecisionsIRQuestion, + DecisionsIRRefusal, + DecisionsIRRequest, + DecisionsIRResponse, + DecisionsIRScoreAnswer, + DecisionsIRScoreLevel, + DecisionsIRScoreProbability, + DecisionsIRScoreQuestion, + DecisionsIRState, + DecisionsIRUsage, + DecisionsJSON, + DecisionsRequestBody, + DecisionsResponse, + DecisionsUsage, + NoulAnswer, + NoulQuestion, + OpenAIDecisionInputImage, + OpenAIDecisionInputMessage, + OpenAIDecisionInputText, + ScoreAnswer, + ScoreQuestion, + UnsupportedDecisionsRequest, + systemone_choice_key, +) + +_PAYLOAD_ADAPTER: Final[TypeAdapter[object]] = TypeAdapter(object) +_SYSTEMONE_RESPONSE_ADAPTER: Final[TypeAdapter[DecisionsResponse]] = TypeAdapter(DecisionsResponse) +_RESERVED_HEADERS: Final[frozenset[str]] = frozenset({"authorization", "content-type"}) +_TEXT_ONLY: Final = UnsupportedDecisionsRequest( + reason="input_image content parts are not supported because System One providers accept text input only" +) -@dataclass(frozen=True, slots=True) -class JevCompatibleDecisionsEndpoint: - default_api_base_value: str | None - path: str - api_key_env: tuple[str, ...] - api_base_env: str +def decisions_text(value: DecisionsJSON) -> str: + return value if isinstance(value, str) else json.dumps(value) + + +def systemone_keys(questions: Sequence[DecisionsIRQuestion]) -> tuple[str, ...]: + names: Final = tuple(question.name for question in questions if question.name is not None) + if len(frozenset(names)) == len(questions): + return names + return tuple(str(index) for index in range(len(questions))) + + +def _ir_question(name: str, question: DecisionQuestion) -> DecisionsIRQuestion: + extra: Final = MappingProxyType(question.model_extra or {}) + match question: + case NoulQuestion(): + return DecisionsIRPredicateQuestion( + name=name, instructions=question.instructions, criteria=question.criteria, extra=extra + ) + case ChoiceQuestion(): + return DecisionsIRChoiceQuestion( + name=name, + instructions=question.instructions, + choices=tuple( + DecisionsIRChoiceOption(value=value, description=description) + for value, description in question.criteria.items() + ), + extra=extra, + ) + case ScoreQuestion(): + return DecisionsIRScoreQuestion( + name=name, + instructions=question.instructions, + levels=tuple( + DecisionsIRScoreLevel(label=criterion, description=None) for criterion in question.criteria + ), + extra=extra, + ) + case _: + assert_never(question) + + +def systemone_request_to_ir(request: DecisionsRequestBody) -> DecisionsIRRequest: + return DecisionsIRRequest( + input=DecisionsIRState(state=request.state), + questions=tuple(_ir_question(name, question) for name, question in request.questions.items()), + ) + + +def _has_image(message: OpenAIDecisionInputMessage) -> bool: + return not isinstance(message.content, str) and any( + isinstance(part, OpenAIDecisionInputImage) for part in message.content + ) + + +def _message_text(message: OpenAIDecisionInputMessage) -> str: + if isinstance(message.content, str): + return message.content + return "\n\n".join(part.text for part in message.content if isinstance(part, OpenAIDecisionInputText)) + + +def _systemone_state( + decision_input: DecisionsIRState | DecisionsIRMessages, +) -> DecisionsJSON | UnsupportedDecisionsRequest: + match decision_input: + case DecisionsIRState(): + return decision_input.state + case DecisionsIRMessages(): + if any(_has_image(message) for message in decision_input.messages): + return _TEXT_ONLY + return "\n\n".join(_message_text(message) for message in decision_input.messages) + case _: + assert_never(decision_input) + + +def _optional(key: str, value: object) -> Mapping[str, object]: + return {} if value is None else {key: value} + + +def _level_criterion(level: DecisionsIRScoreLevel) -> DecisionsJSON: + return level.label if level.description is None else f"{decisions_text(level.label)}: {level.description}" + + +def _systemone_question(question: DecisionsIRQuestion) -> Mapping[str, object]: + match question: + case DecisionsIRPredicateQuestion(): + return { + "type": "noul", + **_optional("instructions", question.instructions), + **_optional("criteria", question.criteria), + **question.extra, + } + case DecisionsIRChoiceQuestion(): + return { + "type": "choice", + **_optional("instructions", question.instructions), + "criteria": {systemone_choice_key(option.value): option.description for option in question.choices}, + **question.extra, + } + case DecisionsIRScoreQuestion(): + return { + "type": "score", + **_optional("instructions", question.instructions), + "criteria": [_level_criterion(level) for level in question.levels], + **question.extra, + } + case _: + assert_never(question) + + +def ir_to_systemone_request( + model: str, request: DecisionsIRRequest +) -> Mapping[str, object] | UnsupportedDecisionsRequest: + state: Final = _systemone_state(request.input) + if isinstance(state, UnsupportedDecisionsRequest): + return state + keyed_questions: Final = zip(systemone_keys(request.questions), request.questions, strict=True) + return { + "model": model, + "state": state, + "questions": {key: _systemone_question(question) for key, question in keyed_questions}, + } + + +def _ir_choice_answer(question: DecisionsIRChoiceQuestion, answer: ChoiceAnswer) -> DecisionsIRChoiceAnswer: + typed_values: Final = {systemone_choice_key(option.value): option.value for option in question.choices} + return DecisionsIRChoiceAnswer( + choice=typed_values.get(answer.choice, answer.choice), + confidence=answer.confidence, + probabilities=tuple( + DecisionsIRChoiceProbability(value=typed_values.get(key, key), probability=probability) + for key, probability in answer.probabilities.items() + ), + extra=MappingProxyType(answer.model_extra or {}), + ) + + +def _ir_score_answer(question: DecisionsIRScoreQuestion, answer: ScoreAnswer) -> DecisionsIRScoreAnswer: + return DecisionsIRScoreAnswer( + score=answer.score, + confidence=answer.confidence, + probabilities=tuple( + DecisionsIRScoreProbability( + value=index, + label=answer.legend.get(str(index), level.label), + probability=answer.probabilities.get(str(index), 0.0), + ) + for index, level in enumerate(question.levels) + ), + extra=MappingProxyType(answer.model_extra or {}), + ) + + +def _ir_answer(question: DecisionsIRQuestion, answer: DecisionAnswer | None) -> DecisionsIRAnswer: + match question, answer: + case DecisionsIRPredicateQuestion(), NoulAnswer(): + return DecisionsIRPredicateAnswer(probability=answer.noul, extra=MappingProxyType(answer.model_extra or {})) + case DecisionsIRChoiceQuestion(), ChoiceAnswer(): + return _ir_choice_answer(question, answer) + case DecisionsIRScoreQuestion(), ScoreAnswer(): + return _ir_score_answer(question, answer) + case _: + return DecisionsIRRefusal() + + +def systemone_response_to_ir(response: DecisionsResponse, request: DecisionsIRRequest) -> DecisionsIRResponse: + usage: Final = response.usage or DecisionsUsage() + keyed_questions: Final = zip(systemone_keys(request.questions), request.questions, strict=True) + return DecisionsIRResponse( + model=response.model, + answers=tuple(_ir_answer(question, response.answers.get(key)) for key, question in keyed_questions), + usage=DecisionsIRUsage( + input_tokens=usage.input_tokens, + output_tokens=usage.output_tokens, + cached_tokens=usage.cached_tokens, + cache_write_tokens=usage.cache_write_tokens, + extra=MappingProxyType(usage.model_extra or {}), + ), + extra=MappingProxyType(response.model_extra or {}), + ) + + +def parse_systemone_response(payload: object, request: DecisionsIRRequest) -> DecisionsIRResponse: + return systemone_response_to_ir(_SYSTEMONE_RESPONSE_ADAPTER.validate_python(payload), request) + + +def _systemone_answer(answer: DecisionsIRAnswer) -> DecisionAnswer | None: + match answer: + case DecisionsIRPredicateAnswer(): + return NoulAnswer.model_validate({**answer.extra, "type": "noul", "noul": answer.probability}) + case DecisionsIRChoiceAnswer(): + return ChoiceAnswer.model_validate( + { + **answer.extra, + "type": "choice", + "choice": systemone_choice_key(answer.choice), + "confidence": answer.confidence, + "probabilities": { + systemone_choice_key(item.value): item.probability for item in answer.probabilities + }, + } + ) + case DecisionsIRScoreAnswer(): + return ScoreAnswer.model_validate( + { + **answer.extra, + "type": "score", + "score": answer.score, + "confidence": answer.confidence, + "legend": {str(item.value): item.label for item in answer.probabilities}, + "probabilities": {str(item.value): item.probability for item in answer.probabilities}, + } + ) + case DecisionsIRRefusal(): + return None + case _: + assert_never(answer) + + +def ir_to_systemone_response(response: DecisionsIRResponse, request: DecisionsIRRequest) -> DecisionsResponse: + keyed_answers: Final = zip( + systemone_keys(request.questions), (_systemone_answer(answer) for answer in response.answers), strict=True + ) + return DecisionsResponse.model_validate( + { + **response.extra, + "model": response.model, + "answers": {key: answer for key, answer in keyed_answers if answer is not None}, + "usage": DecisionsUsage.model_validate( + { + **response.usage.extra, + "input_tokens": response.usage.input_tokens, + "output_tokens": response.usage.output_tokens, + "cached_tokens": response.usage.cached_tokens, + "cache_write_tokens": response.usage.cache_write_tokens, + } + ), + } + ) + + +class BaseDecisionsConfig(ABC): + path: str = "/v1/systemone" + api_key_env: tuple[str, ...] = () + api_base_env: tuple[str, ...] = () api_key_required: bool = True - def default_api_base(self) -> str | None: - return self.default_api_base_value + def get_default_api_base(self) -> str | None: + return None - def missing_api_base_message(self, provider: str) -> str: - return f"api_base is required for Decisions provider '{provider}'" + def missing_api_base_message(self, custom_llm_provider: str) -> str: + return f"api_base is required for Decisions provider '{custom_llm_provider}'" + + def resolve_api_base(self, api_base: str | None) -> str | None: + return api_base or self._first_secret(self.api_base_env) or self.get_default_api_base() + + def resolve_api_key(self, api_key: str | None) -> str | None: + return api_key or self._first_secret(self.api_key_env) + + @staticmethod + def _first_secret(names: tuple[str, ...]) -> str | None: + return next((value for value in (get_secret_str(name) for name in names) if value), None) def canonical_model(self, model: str) -> str: return model @@ -22,31 +324,50 @@ class JevCompatibleDecisionsEndpoint: def request_model(self, model: str) -> str: return model - def endpoint_url(self, api_base: str, model: str) -> str: + def validate_environment(self, headers: Mapping[str, str], model: str, api_key: str | None) -> dict[str, str]: + return { + **{name: value for name, value in headers.items() if name.lower() not in _RESERVED_HEADERS}, + **({"Authorization": f"Bearer {api_key}"} if api_key is not None else {}), + "Content-Type": "application/json", + } + + def get_complete_url(self, api_base: str, model: str) -> str: return f"{api_base.rstrip('/').removesuffix('/v1')}{self.path}" + def transform_decisions_request( + self, + model: str, + request: DecisionsIRRequest, + custom_llm_provider: str, + ) -> Mapping[str, object] | UnsupportedDecisionsRequest: + return ir_to_systemone_request(self.request_model(model), request) + def unwrap_response(self, payload: object) -> object: return payload + def parse_response(self, payload: object, request: DecisionsIRRequest) -> DecisionsIRResponse: + return parse_systemone_response(self.unwrap_response(payload), request) -class DecisionsProviderConfig(Protocol): - @property - def api_key_env(self) -> tuple[str, ...]: ... + def transform_decisions_response( + self, + model: str, + custom_llm_provider: str, + raw_response: httpx.Response, + request: DecisionsIRRequest, + ) -> DecisionsIRResponse: + payload: Final[object] = _PAYLOAD_ADAPTER.validate_json(raw_response.content) + try: + return self.parse_response(payload, request) + except ValidationError as error: + raise BaseLLMException( + status_code=500, + message=f"Decisions provider '{custom_llm_provider}' returned an unexpected response: {error}", + ) from error - @property - def api_base_env(self) -> str: ... - - @property - def api_key_required(self) -> bool: ... - - def default_api_base(self) -> str | None: ... - - def missing_api_base_message(self, provider: str) -> str: ... - - def canonical_model(self, model: str) -> str: ... - - def request_model(self, model: str) -> str: ... - - def endpoint_url(self, api_base: str, model: str) -> str: ... - - def unwrap_response(self, payload: object) -> object: ... + def get_error_class( + self, + error_message: str, + status_code: int, + headers: dict[str, str] | httpx.Headers, + ) -> BaseLLMException: + return BaseLLMException(status_code=status_code, message=error_message, headers=headers) diff --git a/litellm/llms/bedrock/batches/handler.py b/litellm/llms/bedrock/batches/handler.py index 4952b6950af..81057a1ee23 100644 --- a/litellm/llms/bedrock/batches/handler.py +++ b/litellm/llms/bedrock/batches/handler.py @@ -25,7 +25,7 @@ _BEDROCK_MIJ_STATUS_TO_OPENAI: Final[ ] = { "Submitted": "validating", "Validating": "validating", - "Scheduled": "validating", + "Scheduled": "in_progress", "InProgress": "in_progress", "Stopping": "cancelling", "Stopped": "cancelled", diff --git a/litellm/llms/cloudflare/decisions/transformation.py b/litellm/llms/cloudflare/decisions/transformation.py index 6e8b2999778..fa8c68ad3a4 100644 --- a/litellm/llms/cloudflare/decisions/transformation.py +++ b/litellm/llms/cloudflare/decisions/transformation.py @@ -1,9 +1,9 @@ from collections.abc import Mapping -from dataclasses import dataclass from typing import Final -from pydantic import TypeAdapter +from pydantic import TypeAdapter, ValidationError +from litellm.llms.base_llm.decisions.transformation import BaseDecisionsConfig from litellm.secret_managers.main import ( get_secret_str, normalize_nonempty_secret_str, @@ -12,19 +12,17 @@ from litellm.secret_managers.main import ( _RESPONSE_MAPPING_ADAPTER: Final[TypeAdapter[Mapping[str, object]]] = TypeAdapter(Mapping[str, object]) -@dataclass(frozen=True, slots=True) -class CloudflareDecisionsEndpoint: - api_key_env: tuple[str, ...] = ("CLOUDFLARE_API_KEY",) - api_base_env: str = "CLOUDFLARE_API_BASE" - api_key_required: bool = True +class CloudflareDecisionsConfig(BaseDecisionsConfig): + api_key_env = ("CLOUDFLARE_API_KEY",) + api_base_env = ("CLOUDFLARE_API_BASE",) - def default_api_base(self) -> str | None: + def get_default_api_base(self) -> str | None: account_id: Final = normalize_nonempty_secret_str(get_secret_str("CLOUDFLARE_ACCOUNT_ID")) if account_id is None: return None return f"https://api.cloudflare.com/client/v4/accounts/{account_id}/ai/run" - def missing_api_base_message(self, provider: str) -> str: + def missing_api_base_message(self, custom_llm_provider: str) -> str: return "Missing CLOUDFLARE_ACCOUNT_ID - set CLOUDFLARE_ACCOUNT_ID or pass api_base explicitly" def canonical_model(self, model: str) -> str: @@ -35,7 +33,7 @@ class CloudflareDecisionsEndpoint: def request_model(self, model: str) -> str: return model.rsplit("/", maxsplit=1)[-1] - def endpoint_url(self, api_base: str, model: str) -> str: + def get_complete_url(self, api_base: str, model: str) -> str: normalized_api_base: Final = api_base.rstrip("/") if normalized_api_base.endswith("/ai/v1"): return f"{normalized_api_base.removesuffix('/ai/v1')}/ai/run/{model}" @@ -44,15 +42,13 @@ class CloudflareDecisionsEndpoint: return f"{normalized_api_base}/ai/run/{model}" def unwrap_response(self, payload: object) -> object: - if not isinstance(payload, Mapping): + try: + response_mapping: Final = _RESPONSE_MAPPING_ADAPTER.validate_python(payload) + except ValidationError: return payload - response_mapping: Final = _RESPONSE_MAPPING_ADAPTER.validate_python(payload) if "answers" in response_mapping: - return payload - result: Final = response_mapping.get("result") - if isinstance(result, Mapping): - return result - return payload - - -CLOUDFLARE_DECISIONS_ENDPOINT: Final[CloudflareDecisionsEndpoint] = CloudflareDecisionsEndpoint() + return response_mapping + try: + return _RESPONSE_MAPPING_ADAPTER.validate_python(response_mapping.get("result")) + except ValidationError: + return response_mapping diff --git a/litellm/llms/custom_httpx/http_handler.py b/litellm/llms/custom_httpx/http_handler.py index 51f9f1e9fa4..3351d34b16f 100644 --- a/litellm/llms/custom_httpx/http_handler.py +++ b/litellm/llms/custom_httpx/http_handler.py @@ -450,6 +450,18 @@ def get_shared_realtime_ssl_context() -> bool | str | ssl.SSLContext: return _shared_realtime_ssl_context +def realtime_ssl_for_url(url: str) -> bool | str | ssl.SSLContext | None: + if url.startswith("ws://"): + return None + shared: Final = get_shared_realtime_ssl_context() + if shared is not False: + return shared + unverified: Final = ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT) + unverified.check_hostname = False + unverified.verify_mode = ssl.CERT_NONE + return unverified + + def mask_sensitive_info(error_message): # Find the start of the key parameter if isinstance(error_message, str): diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index edca129daa8..6298ee6ffa7 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -70,6 +70,7 @@ from litellm.llms.base_llm.base_model_iterator import ( from litellm.llms.base_llm.batches.transformation import BaseBatchesConfig from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMException from litellm.llms.base_llm.containers.transformation import BaseContainerConfig +from litellm.llms.base_llm.decisions.transformation import BaseDecisionsConfig from litellm.llms.base_llm.embedding.transformation import BaseEmbeddingConfig from litellm.llms.base_llm.evals.transformation import BaseEvalsAPIConfig from litellm.llms.base_llm.files.transformation import ( @@ -122,6 +123,7 @@ from litellm.types.containers.main import ( ContainerObject, DeleteContainerResult, ) +from litellm.types.decisions import DecisionsIRRequest, DecisionsIRResponse from litellm.types.files import StreamingMediaUploadConfig, TwoStepFileUploadConfig from litellm.types.integrations.custom_logger import ( NON_CODE_INTERPRETER_INTERCEPTION_INTERNAL_PREFIXES, @@ -197,7 +199,7 @@ def _rust_responses_websocket_enabled( return decision(context) is not Decision.PYTHON -from .http_handler import get_shared_realtime_ssl_context +from .http_handler import get_shared_realtime_ssl_context, realtime_ssl_for_url if TYPE_CHECKING: from aiohttp import ClientSession @@ -258,7 +260,7 @@ class _WebsocketsModule(Protocol): *, additional_headers: Mapping[str, str], max_size: int | None, - ssl: bool | str | ssl.SSLContext, + ssl: bool | str | ssl.SSLContext | None, open_timeout: float, ) -> Awaitable["ClientConnection"]: ... @@ -1481,6 +1483,100 @@ class BaseLLMHTTPHandler: request_data=request_data, ) + def _prepare_decisions_request( + self, + model: str, + logging_obj: LiteLLMLoggingObj | None, + provider_config: BaseDecisionsConfig, + body: Mapping[str, object], + api_base: str, + api_key: str | None, + headers: Mapping[str, str], + ) -> tuple[str, dict[str, str], dict[str, object]]: + outbound_headers: Final = provider_config.validate_environment(headers=headers, model=model, api_key=api_key) + url: Final = provider_config.get_complete_url(api_base=api_base, model=model) + data: Final = dict(body) + if logging_obj is not None: + logging_obj.pre_call( + input=data, + api_key=api_key, + model=model, + additional_args={"api_base": url, "complete_input_dict": data, "headers": outbound_headers}, + ) + return url, outbound_headers, data + + def decisions( + self, + model: str, + custom_llm_provider: str, + logging_obj: LiteLLMLoggingObj | None, + provider_config: BaseDecisionsConfig, + request: DecisionsIRRequest, + body: Mapping[str, object], + api_base: str, + api_key: str | None, + headers: Mapping[str, str], + timeout: float | httpx.Timeout | None, + client: HTTPHandler | None = None, + ) -> DecisionsIRResponse: + url, outbound_headers, data = self._prepare_decisions_request( + model=model, + logging_obj=logging_obj, + provider_config=provider_config, + body=body, + api_base=api_base, + api_key=api_key, + headers=headers, + ) + sync_httpx_client: Final = client if client is not None else get_httpx_client() + try: + response: Final = sync_httpx_client.post( + url, json=data, headers=outbound_headers, timeout=timeout, logging_obj=logging_obj + ) + except httpx.HTTPError as e: + raise self._handle_error(e=e, provider_config=provider_config) + return provider_config.transform_decisions_response( + model=model, custom_llm_provider=custom_llm_provider, raw_response=response, request=request + ) + + async def adecisions( + self, + model: str, + custom_llm_provider: str, + logging_obj: LiteLLMLoggingObj | None, + provider_config: BaseDecisionsConfig, + request: DecisionsIRRequest, + body: Mapping[str, object], + api_base: str, + api_key: str | None, + headers: Mapping[str, str], + timeout: float | httpx.Timeout | None, + client: AsyncHTTPHandler | None = None, + ) -> DecisionsIRResponse: + url, outbound_headers, data = self._prepare_decisions_request( + model=model, + logging_obj=logging_obj, + provider_config=provider_config, + body=body, + api_base=api_base, + api_key=api_key, + headers=headers, + ) + async_httpx_client: Final = ( + client + if client is not None + else get_async_httpx_client(llm_provider=litellm.LlmProviders(custom_llm_provider)) + ) + try: + response: Final = await async_httpx_client.post( + url, json=data, headers=outbound_headers, timeout=timeout, logging_obj=logging_obj + ) + except httpx.HTTPError as e: + raise self._handle_error(e=e, provider_config=provider_config) + return provider_config.transform_decisions_response( + model=model, custom_llm_provider=custom_llm_provider, raw_response=response, request=request + ) + def _prepare_audio_transcription_request( self, model: str, @@ -6156,6 +6252,7 @@ class BaseLLMHTTPHandler: e: Exception, provider_config: Union[ BaseConfig, + BaseDecisionsConfig, BaseRerankConfig, BaseResponsesAPIConfig, BaseImageEditConfig, @@ -6257,7 +6354,7 @@ class BaseLLMHTTPHandler: websockets_module: _WebsocketsModule, url: str, headers: dict, - ssl_context: bool | str | ssl.SSLContext, + ssl_context: bool | str | ssl.SSLContext | None, *, open_timeout: float = 8.0, max_attempts: int = 3, @@ -6329,12 +6426,7 @@ class BaseLLMHTTPHandler: ) try: - ssl_context = get_shared_realtime_ssl_context() - if url.startswith("wss://") and ssl_context is False: - # Keep TLS for wss:// while honoring SSL_VERIFY=False semantics. - ssl_context = ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT) - ssl_context.check_hostname = False - ssl_context.verify_mode = ssl.CERT_NONE + ssl_context: Final = realtime_ssl_for_url(url) provider_backend: Final = await provider_config.open_backend(url, headers) backend_ws: Final = ( provider_backend diff --git a/litellm/llms/ollama/completion/transformation.py b/litellm/llms/ollama/completion/transformation.py index 107cae1a98a..59e59a0d50f 100644 --- a/litellm/llms/ollama/completion/transformation.py +++ b/litellm/llms/ollama/completion/transformation.py @@ -39,6 +39,7 @@ from litellm.types.utils import ( ModelResponseStream, ProviderField, StreamingChoices, + generate_id, ) from ..common_utils import OllamaError, OllamaModelInfo, convert_image @@ -524,6 +525,7 @@ class OllamaConfig(BaseConfig): class OllamaTextCompletionResponseIterator(BaseModelResponseIterator): def __init__(self, streaming_response, sync_stream: bool, json_mode: bool | None = False): super().__init__(streaming_response, sync_stream, json_mode) + self.response_id: Final[str] = generate_id() self.started_reasoning_content: bool = False self.finished_reasoning_content: bool = False self.streamed_content: bool = False @@ -622,6 +624,7 @@ class OllamaTextCompletionResponseIterator(BaseModelResponseIterator): content = self._hold_json_object_start(text) return ModelResponseStream( + id=self.response_id, choices=[ StreamingChoices( index=0, @@ -641,24 +644,26 @@ class OllamaTextCompletionResponseIterator(BaseModelResponseIterator): # Return reasoning content as ModelResponseStream so UIs can render it thinking_content: Final = chunk.get("thinking") or "" return ModelResponseStream( + id=self.response_id, choices=[ StreamingChoices( index=0, delta=Delta(reasoning_content=thinking_content), ) - ] + ], ) else: # In this case, 'thinking' is not present in the chunk, chunk["done"] is false, # and chunk["response"] is falsy (None or empty string), # but Ollama is just starting to stream, so it should be processed as a normal dict return ModelResponseStream( + id=self.response_id, choices=[ StreamingChoices( index=0, delta=Delta(reasoning_content=""), ) - ] + ], ) # raise Exception(f"Unable to parse ollama chunk - {chunk}") except Exception as e: diff --git a/litellm/llms/openai/decisions/transformation.py b/litellm/llms/openai/decisions/transformation.py new file mode 100644 index 00000000000..0afbe1a434e --- /dev/null +++ b/litellm/llms/openai/decisions/transformation.py @@ -0,0 +1,320 @@ +from collections.abc import Mapping, Sequence +from typing import Final + +from pydantic import TypeAdapter +from typing_extensions import assert_never + +import litellm +from litellm.llms.base_llm.decisions.transformation import BaseDecisionsConfig, decisions_text +from litellm.secret_managers.main import get_secret_str +from litellm.types.decisions import ( + DecisionsIRAnswer, + DecisionsIRChoiceAnswer, + DecisionsIRChoiceOption, + DecisionsIRChoiceProbability, + DecisionsIRChoiceQuestion, + DecisionsIRMessages, + DecisionsIRPredicateAnswer, + DecisionsIRPredicateQuestion, + DecisionsIRQuestion, + DecisionsIRRefusal, + DecisionsIRRequest, + DecisionsIRResponse, + DecisionsIRScoreAnswer, + DecisionsIRScoreLevel, + DecisionsIRScoreProbability, + DecisionsIRScoreQuestion, + DecisionsIRState, + DecisionsIRUsage, + DecisionsJSON, + OpenAIChoiceAnswer, + OpenAIChoiceProbability, + OpenAIChoiceQuestion, + OpenAIDecisionAnswer, + OpenAIDecisionInput, + OpenAIDecisionInputTokensDetails, + OpenAIDecisionOutputTokensDetails, + OpenAIDecisionQuestion, + OpenAIDecisionRequestBody, + OpenAIDecisionResponse, + OpenAIDecisionUsage, + OpenAIPredicateAnswer, + OpenAIPredicateQuestion, + OpenAIRefusalAnswer, + OpenAIScoreAnswer, + OpenAIScoreProbability, + OpenAIScoreQuestion, + UnsupportedDecisionsRequest, + systemone_choice_key, +) + +_OPENAI_RESPONSE_ADAPTER: Final[TypeAdapter[OpenAIDecisionResponse]] = TypeAdapter(OpenAIDecisionResponse) +_SINGLE_OPTION: Final = UnsupportedDecisionsRequest( + reason="OpenAI needs at least 2 choices or levels on every choice or score question" +) + + +def _ir_input(decision_input: OpenAIDecisionInput) -> DecisionsIRState | DecisionsIRMessages: + if isinstance(decision_input, str): + return DecisionsIRState(state=decision_input) + return DecisionsIRMessages(messages=tuple(decision_input)) + + +def _ir_question(question: OpenAIDecisionQuestion) -> DecisionsIRQuestion: + match question: + case OpenAIPredicateQuestion(): + return DecisionsIRPredicateQuestion(name=question.name, instructions=question.instructions) + case OpenAIChoiceQuestion(): + return DecisionsIRChoiceQuestion( + name=question.name, + instructions=question.instructions, + choices=tuple( + DecisionsIRChoiceOption(value=option.value, description=option.description) + for option in question.choices + ), + ) + case OpenAIScoreQuestion(): + return DecisionsIRScoreQuestion( + name=question.name, + instructions=question.instructions, + levels=tuple( + DecisionsIRScoreLevel(label=level.label, description=level.description) for level in question.levels + ), + ) + case _: + assert_never(question) + + +def openai_request_to_ir(request: OpenAIDecisionRequestBody) -> DecisionsIRRequest: + return DecisionsIRRequest( + input=_ir_input(request.input), + questions=tuple(_ir_question(question) for question in request.questions), + safety_identifier=request.safety_identifier, + ) + + +def _text_field(key: str, value: DecisionsJSON | None) -> Mapping[str, str]: + return {} if value is None else {key: decisions_text(value)} + + +def _instructions(value: DecisionsJSON | None, default: str) -> str: + return default if value is None else decisions_text(value) + + +def _predicate_instructions(question: DecisionsIRPredicateQuestion) -> str: + instructions: Final = () if question.instructions is None else (decisions_text(question.instructions),) + criteria: Final = tuple( + f"Answer {answer} when: {decisions_text(rule)}" + for answer, rule in (question.criteria or {}).items() + if rule is not None + ) + return "\n\n".join((*instructions, *criteria)) + + +def _openai_question(question: DecisionsIRQuestion) -> Mapping[str, object]: + match question: + case DecisionsIRPredicateQuestion(): + return { + "type": "predicate", + **_text_field("name", question.name), + "instructions": _predicate_instructions(question) or "Is this true of the input?", + } + case DecisionsIRChoiceQuestion(): + return { + "type": "choice", + **_text_field("name", question.name), + "instructions": _instructions(question.instructions, "Which choice best fits the input?"), + "choices": [ + {"value": option.value, **_text_field("description", option.description)} + for option in question.choices + ], + } + case DecisionsIRScoreQuestion(): + return { + "type": "score", + **_text_field("name", question.name), + "instructions": _instructions(question.instructions, "Which level best fits the input?"), + "levels": [ + {"label": decisions_text(level.label), **_text_field("description", level.description)} + for level in question.levels + ], + } + case _: + assert_never(question) + + +def _openai_input(decision_input: DecisionsIRState | DecisionsIRMessages) -> str | Sequence[Mapping[str, object]]: + match decision_input: + case DecisionsIRState(): + return decisions_text(decision_input.state) + case DecisionsIRMessages(): + return [message.model_dump(mode="json", exclude_none=True) for message in decision_input.messages] + case _: + assert_never(decision_input) + + +def _has_one_option(question: DecisionsIRQuestion) -> bool: + match question: + case DecisionsIRChoiceQuestion(): + return len(question.choices) < 2 + case DecisionsIRScoreQuestion(): + return len(question.levels) < 2 + case _: + return False + + +def ir_to_openai_request(model: str, request: DecisionsIRRequest) -> Mapping[str, object] | UnsupportedDecisionsRequest: + if any(_has_one_option(question) for question in request.questions): + return _SINGLE_OPTION + return { + "model": model, + "input": _openai_input(request.input), + "questions": [_openai_question(question) for question in request.questions], + **_text_field("safety_identifier", request.safety_identifier), + } + + +def _level_label(question: DecisionsIRScoreQuestion, probability: OpenAIScoreProbability) -> DecisionsJSON: + if 0 <= probability.value < len(question.levels): + return question.levels[probability.value].label + return probability.label + + +def _ir_answer(question: DecisionsIRQuestion, answer: OpenAIDecisionAnswer | None) -> DecisionsIRAnswer: + match question, answer: + case DecisionsIRPredicateQuestion(), OpenAIPredicateAnswer(): + return DecisionsIRPredicateAnswer(probability=answer.probability) + case DecisionsIRChoiceQuestion(), OpenAIChoiceAnswer(): + return DecisionsIRChoiceAnswer( + choice=answer.choice, + confidence=answer.confidence, + probabilities=tuple( + DecisionsIRChoiceProbability(value=item.value, probability=item.probability) + for item in answer.probabilities + ), + ) + case DecisionsIRScoreQuestion(), OpenAIScoreAnswer(): + return DecisionsIRScoreAnswer( + score=answer.score, + confidence=answer.confidence, + probabilities=tuple( + DecisionsIRScoreProbability( + value=item.value, label=_level_label(question, item), probability=item.probability + ) + for item in answer.probabilities + ), + ) + case _: + return DecisionsIRRefusal() + + +def openai_response_to_ir(response: OpenAIDecisionResponse, request: DecisionsIRRequest) -> DecisionsIRResponse: + answers: Final = response.answers + return DecisionsIRResponse( + model=response.model, + answers=tuple( + _ir_answer(question, answers[index] if index < len(answers) else None) + for index, question in enumerate(request.questions) + ), + usage=DecisionsIRUsage( + input_tokens=response.usage.input_tokens, + output_tokens=response.usage.output_tokens, + cached_tokens=response.usage.input_tokens_details.cached_tokens, + cache_write_tokens=response.usage.input_tokens_details.cache_write_tokens, + reasoning_tokens=response.usage.output_tokens_details.reasoning_tokens, + ), + ) + + +def _openai_choice_answer(question: DecisionsIRChoiceQuestion, answer: DecisionsIRChoiceAnswer) -> OpenAIChoiceAnswer: + probabilities: Final = {systemone_choice_key(item.value): item.probability for item in answer.probabilities} + return OpenAIChoiceAnswer( + name=question.name, + choice=answer.choice, + probabilities=tuple( + OpenAIChoiceProbability( + value=option.value, probability=probabilities.get(systemone_choice_key(option.value), 0.0) + ) + for option in question.choices + ), + confidence=answer.confidence, + ) + + +def _openai_score_answer(question: DecisionsIRScoreQuestion, answer: DecisionsIRScoreAnswer) -> OpenAIScoreAnswer: + probabilities: Final = {item.value: item.probability for item in answer.probabilities} + return OpenAIScoreAnswer( + name=question.name, + score=answer.score, + probabilities=tuple( + OpenAIScoreProbability( + value=index, label=decisions_text(level.label), probability=probabilities.get(index, 0.0) + ) + for index, level in enumerate(question.levels) + ), + confidence=answer.confidence, + ) + + +def _openai_answer(question: DecisionsIRQuestion, answer: DecisionsIRAnswer) -> OpenAIDecisionAnswer: + match question, answer: + case DecisionsIRPredicateQuestion(), DecisionsIRPredicateAnswer(): + return OpenAIPredicateAnswer(name=question.name, probability=answer.probability) + case DecisionsIRChoiceQuestion(), DecisionsIRChoiceAnswer(): + return _openai_choice_answer(question, answer) + case DecisionsIRScoreQuestion(), DecisionsIRScoreAnswer(): + return _openai_score_answer(question, answer) + case _: + return OpenAIRefusalAnswer(name=question.name) + + +def ir_to_openai_response( + response: DecisionsIRResponse, request: DecisionsIRRequest, requested_model: str +) -> OpenAIDecisionResponse: + return OpenAIDecisionResponse( + model=response.model or requested_model, + answers=tuple( + _openai_answer(question, answer) + for question, answer in zip(request.questions, response.answers, strict=True) + ), + usage=OpenAIDecisionUsage( + input_tokens=response.usage.input_tokens, + input_tokens_details=OpenAIDecisionInputTokensDetails( + cached_tokens=response.usage.cached_tokens, + cache_write_tokens=response.usage.cache_write_tokens, + ), + output_tokens=response.usage.output_tokens, + output_tokens_details=OpenAIDecisionOutputTokensDetails(reasoning_tokens=response.usage.reasoning_tokens), + total_tokens=response.usage.input_tokens + response.usage.output_tokens, + ), + ) + + +class OpenAIDecisionsConfig(BaseDecisionsConfig): + path = "/v1/decisions" + + def get_default_api_base(self) -> str | None: + return "https://api.openai.com" + + def resolve_api_base(self, api_base: str | None) -> str | None: + return ( + api_base + or litellm.api_base + or get_secret_str("OPENAI_BASE_URL") + or get_secret_str("OPENAI_API_BASE") + or self.get_default_api_base() + ) + + def resolve_api_key(self, api_key: str | None) -> str | None: + return api_key or litellm.api_key or litellm.openai_key or get_secret_str("OPENAI_API_KEY") + + def transform_decisions_request( + self, + model: str, + request: DecisionsIRRequest, + custom_llm_provider: str, + ) -> Mapping[str, object] | UnsupportedDecisionsRequest: + return ir_to_openai_request(self.request_model(model), request) + + def parse_response(self, payload: object, request: DecisionsIRRequest) -> DecisionsIRResponse: + return openai_response_to_ir(_OPENAI_RESPONSE_ADAPTER.validate_python(payload), request) diff --git a/litellm/llms/openrouter/decisions/transformation.py b/litellm/llms/openrouter/decisions/transformation.py index 7a7466b1239..53ddce8b402 100644 --- a/litellm/llms/openrouter/decisions/transformation.py +++ b/litellm/llms/openrouter/decisions/transformation.py @@ -1,10 +1,10 @@ -from typing import Final +from litellm.llms.base_llm.decisions.transformation import BaseDecisionsConfig -from litellm.llms.base_llm.decisions.transformation import JevCompatibleDecisionsEndpoint -OPENROUTER_DECISIONS_ENDPOINT: Final[JevCompatibleDecisionsEndpoint] = JevCompatibleDecisionsEndpoint( - default_api_base_value="https://openrouter.ai/api", - path="/alpha/decisions", - api_key_env=("OPENROUTER_API_KEY",), - api_base_env="OPENROUTER_API_BASE", -) +class OpenRouterDecisionsConfig(BaseDecisionsConfig): + path = "/alpha/decisions" + api_key_env = ("OPENROUTER_API_KEY",) + api_base_env = ("OPENROUTER_API_BASE",) + + def get_default_api_base(self) -> str | None: + return "https://openrouter.ai/api" diff --git a/litellm/llms/perplexity/decisions/transformation.py b/litellm/llms/perplexity/decisions/transformation.py index 69a4753f4a3..11e2a38f06e 100644 --- a/litellm/llms/perplexity/decisions/transformation.py +++ b/litellm/llms/perplexity/decisions/transformation.py @@ -1,10 +1,10 @@ -from typing import Final +from litellm.llms.base_llm.decisions.transformation import BaseDecisionsConfig -from litellm.llms.base_llm.decisions.transformation import JevCompatibleDecisionsEndpoint -PERPLEXITY_DECISIONS_ENDPOINT: Final[JevCompatibleDecisionsEndpoint] = JevCompatibleDecisionsEndpoint( - default_api_base_value="https://api.perplexity.ai", - path="/v1/decisions", - api_key_env=("PERPLEXITYAI_API_KEY", "PERPLEXITY_API_KEY"), - api_base_env="PERPLEXITY_API_BASE", -) +class PerplexityDecisionsConfig(BaseDecisionsConfig): + path = "/v1/decisions" + api_key_env = ("PERPLEXITYAI_API_KEY", "PERPLEXITY_API_KEY") + api_base_env = ("PERPLEXITY_API_BASE",) + + def get_default_api_base(self) -> str | None: + return "https://api.perplexity.ai" diff --git a/litellm/llms/strands_decider/decisions/transformation.py b/litellm/llms/strands_decider/decisions/transformation.py index 265afb2b148..9ef966e18f2 100644 --- a/litellm/llms/strands_decider/decisions/transformation.py +++ b/litellm/llms/strands_decider/decisions/transformation.py @@ -1,11 +1,7 @@ -from typing import Final +from litellm.llms.base_llm.decisions.transformation import BaseDecisionsConfig -from litellm.llms.base_llm.decisions.transformation import JevCompatibleDecisionsEndpoint -STRANDS_DECIDER_DECISIONS_ENDPOINT: Final[JevCompatibleDecisionsEndpoint] = JevCompatibleDecisionsEndpoint( - default_api_base_value=None, - path="/v1/systemone", - api_key_env=("STRANDS_DECIDER_API_KEY",), - api_base_env="STRANDS_DECIDER_API_BASE", - api_key_required=False, -) +class StrandsDeciderDecisionsConfig(BaseDecisionsConfig): + api_key_env = ("STRANDS_DECIDER_API_KEY",) + api_base_env = ("STRANDS_DECIDER_API_BASE",) + api_key_required = False diff --git a/litellm/llms/typesafe/decisions/transformation.py b/litellm/llms/typesafe/decisions/transformation.py index 17fb24ce443..e71c526a840 100644 --- a/litellm/llms/typesafe/decisions/transformation.py +++ b/litellm/llms/typesafe/decisions/transformation.py @@ -1,10 +1,9 @@ -from typing import Final +from litellm.llms.base_llm.decisions.transformation import BaseDecisionsConfig -from litellm.llms.base_llm.decisions.transformation import JevCompatibleDecisionsEndpoint -TYPESAFE_DECISIONS_ENDPOINT: Final[JevCompatibleDecisionsEndpoint] = JevCompatibleDecisionsEndpoint( - default_api_base_value="https://api.typesafe.ai", - path="/v1/systemone", - api_key_env=("TYPESAFE_API_KEY",), - api_base_env="TYPESAFE_API_BASE", -) +class TypeSafeDecisionsConfig(BaseDecisionsConfig): + api_key_env = ("TYPESAFE_API_KEY",) + api_base_env = ("TYPESAFE_API_BASE",) + + def get_default_api_base(self) -> str | None: + return "https://api.typesafe.ai" diff --git a/litellm/llms/vertex_ai/realtime/transformation.py b/litellm/llms/vertex_ai/realtime/transformation.py index 9fed6d52f0e..c14dd535145 100644 --- a/litellm/llms/vertex_ai/realtime/transformation.py +++ b/litellm/llms/vertex_ai/realtime/transformation.py @@ -5,9 +5,14 @@ Extends GeminiRealtimeConfig but adapts the WSS URL and auth header for the Vertex AI endpoint instead of Google AI Studio. URL pattern: - wss://{location}-aiplatform.googleapis.com/ws/ + wss://{vertex host for the location}/ws/ google.cloud.aiplatform.v1.LlmBidiService/BidiGenerateContent +The host is the one ``litellm.llms.vertex_ai.common_utils.get_vertex_base_url`` +resolves for the location: ``{region}-aiplatform.googleapis.com`` for a region, +``aiplatform.{geo}.rep.googleapis.com`` for the ``us`` / ``eu`` multi-regions, +and ``aiplatform.googleapis.com`` for ``global``. + Auth: OAuth2 Bearer token (not an API key). """ @@ -21,6 +26,7 @@ from litellm.llms.vertex_ai.audio_transcription.realtime_transformation import ( VertexChirpRealtimeConfig, is_vertex_speech_to_text_model, ) +from litellm.llms.vertex_ai.common_utils import get_vertex_base_url from litellm.llms.vertex_ai.vertex_llm_base import VertexBase @@ -63,12 +69,7 @@ class VertexAIRealtimeConfig(GeminiRealtimeConfig): base = base.replace("https://", "wss://").replace("http://", "ws://") return f"{base}/ws/google.cloud.aiplatform.v1.LlmBidiService/BidiGenerateContent" - location: Final = self._location - if location == "global": - host = "aiplatform.googleapis.com" - else: - host = f"{location}-aiplatform.googleapis.com" - + host: Final = get_vertex_base_url(self._location).removeprefix("https://") return f"wss://{host}/ws/google.cloud.aiplatform.v1.LlmBidiService/BidiGenerateContent" # ------------------------------------------------------------------ diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index c2aa83a9216..f34606ce127 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -34527,7 +34527,8 @@ "supported_endpoints": [ "/v1/chat/completions", "/v1/batch", - "/v1/responses" + "/v1/responses", + "/v1/decisions" ], "supported_modalities": [ "text", @@ -59538,7 +59539,7 @@ "bedrock_converse_supports_strict_tools": false, "cache_creation_input_token_cost": 2.75e-06, "cache_creation_input_token_cost_above_1hr": 4.4e-06, - "cache_read_input_token_cost": 2.2e-07, + "cache_read_input_token_cost": 1.1e-07, "input_cost_per_token": 2.2e-06, "litellm_provider": "bedrock_mantle", "supports_tool_search": true, @@ -59573,7 +59574,7 @@ "supports_forced_tool_use": false, "thinking_always_on": true, "prompt_cache_min_tokens": 512, - "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-anthropic-claude-sonnet-5-5.html" + "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" }, "us.xai.grok-4.6": { "supports_regex_lookaround": false, @@ -66362,7 +66363,7 @@ "bedrock_output_config_effort_ceiling": "xhigh", "cache_creation_input_token_cost": 3e-06, "cache_creation_input_token_cost_above_1hr": 4.8e-06, - "cache_read_input_token_cost": 2.4e-07, + "cache_read_input_token_cost": 1.2e-07, "input_cost_per_token": 2.4e-06, "litellm_provider": "bedrock_mantle", "supports_tool_search": true, @@ -66391,7 +66392,7 @@ "supports_tool_choice": true, "supports_vision": true, "supports_xhigh_reasoning_effort": true, - "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-anthropic-claude-sonnet-5-5.html" + "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" }, "bedrock_mantle/us-gov-east-1/openai.gpt-5.4": { "litellm_provider": "bedrock_mantle", @@ -79407,7 +79408,7 @@ "bedrock_converse_supports_strict_tools": false, "cache_creation_input_token_cost": 2.5e-06, "cache_creation_input_token_cost_above_1hr": 4e-06, - "cache_read_input_token_cost": 2e-07, + "cache_read_input_token_cost": 1e-07, "input_cost_per_token": 2e-06, "litellm_provider": "bedrock_converse", "supports_tool_search": true, @@ -79477,14 +79478,14 @@ "cache_creation_input_token_cost": 2.75e-06, "input_cost_per_token": 2.2e-06, "output_cost_per_token": 1.1e-05, - "cache_read_input_token_cost": 2.2e-07, + "cache_read_input_token_cost": 1.1e-07, "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" }, "au.anthropic.claude-sonnet-5-5": { "bedrock_converse_supports_strict_tools": false, "cache_creation_input_token_cost": 2.75e-06, "cache_creation_input_token_cost_above_1hr": 4.4e-06, - "cache_read_input_token_cost": 2.2e-07, + "cache_read_input_token_cost": 1.1e-07, "input_cost_per_token": 2.2e-06, "litellm_provider": "bedrock_converse", "supports_tool_search": true, @@ -79519,7 +79520,7 @@ "supports_forced_tool_use": false, "thinking_always_on": true, "prompt_cache_min_tokens": 512, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" }, "azure_ai/claude-sonnet-5-5": { "supports_mid_conversation_system": true, @@ -79562,7 +79563,7 @@ "bedrock_output_config_effort_ceiling": "xhigh", "cache_creation_input_token_cost": 3e-06, "cache_creation_input_token_cost_above_1hr": 4.8e-06, - "cache_read_input_token_cost": 2.4e-07, + "cache_read_input_token_cost": 1.2e-07, "input_cost_per_token": 2.4e-06, "litellm_provider": "bedrock", "supports_tool_search": true, @@ -79574,6 +79575,7 @@ "supports_forced_tool_use": false, "thinking_always_on": true, "prompt_cache_min_tokens": 512, + "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json", "supports_adaptive_thinking": true, "supports_assistant_prefill": false, "supports_computer_use": true, @@ -79597,7 +79599,7 @@ "bedrock_output_config_effort_ceiling": "xhigh", "cache_creation_input_token_cost": 3e-06, "cache_creation_input_token_cost_above_1hr": 4.8e-06, - "cache_read_input_token_cost": 2.4e-07, + "cache_read_input_token_cost": 1.2e-07, "input_cost_per_token": 2.4e-06, "litellm_provider": "bedrock", "supports_tool_search": true, @@ -79609,6 +79611,7 @@ "supports_forced_tool_use": false, "thinking_always_on": true, "prompt_cache_min_tokens": 512, + "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json", "supports_adaptive_thinking": true, "supports_assistant_prefill": false, "supports_computer_use": true, @@ -79631,7 +79634,7 @@ "bedrock_converse_supports_strict_tools": false, "cache_creation_input_token_cost": 2.75e-06, "cache_creation_input_token_cost_above_1hr": 4.4e-06, - "cache_read_input_token_cost": 2.2e-07, + "cache_read_input_token_cost": 1.1e-07, "input_cost_per_token": 2.2e-06, "litellm_provider": "bedrock_converse", "supports_tool_search": true, @@ -79666,13 +79669,13 @@ "supports_forced_tool_use": false, "thinking_always_on": true, "prompt_cache_min_tokens": 512, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" }, "global.anthropic.claude-sonnet-5-5": { "bedrock_converse_supports_strict_tools": false, "cache_creation_input_token_cost": 2.5e-06, "cache_creation_input_token_cost_above_1hr": 4e-06, - "cache_read_input_token_cost": 2e-07, + "cache_read_input_token_cost": 1e-07, "input_cost_per_token": 2e-06, "litellm_provider": "bedrock_converse", "supports_tool_search": true, @@ -79707,13 +79710,13 @@ "supports_forced_tool_use": false, "thinking_always_on": true, "prompt_cache_min_tokens": 512, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" }, "jp.anthropic.claude-sonnet-5-5": { "bedrock_converse_supports_strict_tools": false, "cache_creation_input_token_cost": 2.75e-06, "cache_creation_input_token_cost_above_1hr": 4.4e-06, - "cache_read_input_token_cost": 2.2e-07, + "cache_read_input_token_cost": 1.1e-07, "input_cost_per_token": 2.2e-06, "litellm_provider": "bedrock_converse", "supports_tool_search": true, @@ -79748,7 +79751,7 @@ "supports_forced_tool_use": false, "thinking_always_on": true, "prompt_cache_min_tokens": 512, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" }, "openrouter/anthropic/claude-sonnet-5.5": { "input_cost_per_token": 2e-06, @@ -79858,7 +79861,7 @@ "bedrock_output_config_effort_ceiling": "xhigh", "cache_creation_input_token_cost": 3e-06, "cache_creation_input_token_cost_above_1hr": 4.8e-06, - "cache_read_input_token_cost": 2.4e-07, + "cache_read_input_token_cost": 1.2e-07, "input_cost_per_token": 2.4e-06, "litellm_provider": "bedrock_converse", "supports_tool_search": true, @@ -79870,7 +79873,7 @@ "supports_forced_tool_use": false, "thinking_always_on": true, "prompt_cache_min_tokens": 512, - "source": "https://aws.amazon.com/bedrock/pricing/", + "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json", "supports_adaptive_thinking": true, "supports_assistant_prefill": false, "supports_computer_use": true, @@ -79893,7 +79896,7 @@ "bedrock_converse_supports_strict_tools": false, "cache_creation_input_token_cost": 2.75e-06, "cache_creation_input_token_cost_above_1hr": 4.4e-06, - "cache_read_input_token_cost": 2.2e-07, + "cache_read_input_token_cost": 1.1e-07, "input_cost_per_token": 2.2e-06, "litellm_provider": "bedrock_converse", "supports_tool_search": true, @@ -79928,7 +79931,7 @@ "supports_forced_tool_use": false, "thinking_always_on": true, "prompt_cache_min_tokens": 512, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" }, "vertex_ai/claude-sonnet-5-5": { "regional_endpoint_uplift_multiplier": 1.1, @@ -79936,8 +79939,8 @@ "cache_creation_input_token_cost": 2.5e-06, "cache_creation_input_token_cost_above_1hr": 4e-06, "cache_creation_input_token_cost_batches": 1.25e-06, - "cache_read_input_token_cost": 2e-07, - "cache_read_input_token_cost_batches": 1e-07, + "cache_read_input_token_cost": 1e-07, + "cache_read_input_token_cost_batches": 5e-08, "input_cost_per_token": 2e-06, "input_cost_per_token_batches": 1e-06, "litellm_provider": "vertex_ai-anthropic_models", @@ -79976,8 +79979,8 @@ "cache_creation_input_token_cost": 2.5e-06, "cache_creation_input_token_cost_above_1hr": 4e-06, "cache_creation_input_token_cost_batches": 1.25e-06, - "cache_read_input_token_cost": 2e-07, - "cache_read_input_token_cost_batches": 1e-07, + "cache_read_input_token_cost": 1e-07, + "cache_read_input_token_cost_batches": 5e-08, "input_cost_per_token": 2e-06, "input_cost_per_token_batches": 1e-06, "litellm_provider": "vertex_ai-anthropic_models", diff --git a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py index 9ee904dd4be..608f695c114 100644 --- a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py +++ b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py @@ -56,12 +56,16 @@ from litellm.proxy.agent_endpoints.auth.agent_access_groups import ( from litellm.proxy.agent_endpoints.auth.agent_caller import agent_caller_auth from litellm.proxy.agent_endpoints.auth.managed_authorization import managed_agent_policy from litellm.proxy.auth.ip_address_utils import IPAddressUtils -from litellm.proxy.auth.user_api_key_auth import ( +from litellm.proxy.auth.user_api_key_auth import ( # noqa: F401 # legacy module exports _get_bearer_token_or_received_api_key, # pyright: ignore[reportPrivateUsage] # shared x-litellm-api-key parser lives with user_api_key_auth - _run_centralized_common_checks, + _run_centralized_common_checks, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + run_centralized_common_checks, user_api_key_auth, ) -from litellm.proxy.common_utils.http_parsing_utils import _read_request_body +from litellm.proxy.common_utils.http_parsing_utils import ( # noqa: F401 # legacy module exports + _read_request_body, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + read_request_body, +) from litellm.proxy.common_utils.user_api_key_cache import ( AUTH_OBJECTS_TARGET, USER_NO_MCP_PERMISSION_SENTINEL, @@ -225,7 +229,7 @@ def _agent_capped_servers( ) -def _is_mcp_admitted_user_subject(user_api_key_auth: UserAPIKeyAuth | None) -> bool: +def is_mcp_admitted_user_subject(user_api_key_auth: UserAPIKeyAuth | None) -> bool: """True when this auth is a keyless subject admitted by the gateway session / bridge user path, as opposed to a JWT or other keyless auth that merely lacks a ``team_id``. @@ -235,6 +239,9 @@ def _is_mcp_admitted_user_subject(user_api_key_auth: UserAPIKeyAuth | None) -> b return user_api_key_auth is not None and user_api_key_auth.mcp_admitted_user_subject is True +_is_mcp_admitted_user_subject: Final = is_mcp_admitted_user_subject + + def _gateway_dcr_challenge_target( route: str, mcp_servers: list[str] | None, @@ -452,7 +459,7 @@ class MCPRequestHandler: HTTPException: If headers are invalid or missing required headers """ async with global_manager().catalog.operation(): - headers: Final = MCPRequestHandler._safe_get_headers_from_scope(scope) + headers: Final = MCPRequestHandler.safe_get_headers_from_scope(scope) # Check if there is an explicit LiteLLM API key (primary header) has_explicit_litellm_key: Final = ( @@ -462,13 +469,19 @@ class MCPRequestHandler: litellm_api_key: Final = MCPRequestHandler.get_litellm_api_key_from_headers(headers) or "" # Get the old mcp_auth_header for backward compatibility - mcp_auth_header = MCPRequestHandler._get_mcp_auth_header_from_headers(headers) + mcp_auth_header = MCPRequestHandler.get_mcp_auth_header_from_headers( # rebind-ok: pre-existing rebinding on a rename-only line + headers + ) # Get the new server-specific auth headers - mcp_server_auth_headers = MCPRequestHandler._get_mcp_server_auth_headers_from_headers(headers) + mcp_server_auth_headers = MCPRequestHandler.get_mcp_server_auth_headers_from_headers( # rebind-ok: pre-existing rebinding on a rename-only line + headers + ) # Get the oauth2 headers - oauth2_headers = MCPRequestHandler._get_oauth2_headers_from_headers(headers) + oauth2_headers = MCPRequestHandler.get_oauth2_headers_from_headers( # rebind-ok: pre-existing rebinding on a rename-only line + headers + ) # Parse MCP servers from header mcp_servers_header: Final = headers.get(MCPRequestHandler.LITELLM_MCP_SERVERS_HEADER_NAME) @@ -621,7 +634,7 @@ class MCPRequestHandler: mcp_auth_header, mcp_server_auth_headers, ) = MCPRequestHandler._scrub_gateway_admission_credentials( - admitted=_is_mcp_admitted_user_subject(validated_user_api_key_auth), + admitted=is_mcp_admitted_user_subject(validated_user_api_key_auth), oauth2_headers=oauth2_headers, raw_headers=raw_headers, mcp_auth_header=mcp_auth_header, @@ -1060,7 +1073,7 @@ class MCPRequestHandler: await pre_db_read_auth_checks( request=request, - request_data=await _read_request_body(request=request), + request_data=await read_request_body(request=request), route=route, ) @@ -1351,10 +1364,10 @@ class MCPRequestHandler: admitted.budget_reservation = None try: RouteChecks.should_call_route(route=route, valid_token=admitted, request=request) - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=admitted, request=request, - request_data=await _read_request_body(request=request), + request_data=await read_request_body(request=request), route=route, ) except (HTTPException, ProxyException): @@ -1401,7 +1414,7 @@ class MCPRequestHandler: return mcp_servers_header if mcp_servers_header is not None else [] @staticmethod - def _get_mcp_auth_header_from_headers(headers: Headers) -> str | None: + def get_mcp_auth_header_from_headers(headers: Headers) -> str | None: """ Get the header passed to LiteLLM to pass to downstream MCP servers @@ -1424,8 +1437,10 @@ class MCPRequestHandler: ) return auth_header + _get_mcp_auth_header_from_headers = get_mcp_auth_header_from_headers + @staticmethod - def _get_mcp_server_auth_headers_from_headers( + def get_mcp_server_auth_headers_from_headers( headers: Headers, ) -> dict[str, dict[str, str]]: """ @@ -1478,8 +1493,10 @@ class MCPRequestHandler: return server_auth_headers + _get_mcp_server_auth_headers_from_headers = get_mcp_server_auth_headers_from_headers + @staticmethod - def _get_oauth2_headers_from_headers(headers: Headers) -> dict[str, str]: + def get_oauth2_headers_from_headers(headers: Headers) -> dict[str, str]: """ Get the oauth2 headers from the request headers. """ @@ -1489,6 +1506,8 @@ class MCPRequestHandler: oauth2_headers["Authorization"] = header_value return oauth2_headers + _get_oauth2_headers_from_headers = get_oauth2_headers_from_headers + @staticmethod def get_mcp_client_side_auth_header_name() -> str: """ @@ -1535,7 +1554,7 @@ class MCPRequestHandler: return None @staticmethod - def _safe_get_headers_from_scope(scope: Scope) -> Headers: + def safe_get_headers_from_scope(scope: Scope) -> Headers: """ Safely extract headers from ASGI scope using Starlette's Headers class which handles case insensitivity and proper header parsing. @@ -1563,6 +1582,8 @@ class MCPRequestHandler: # Return empty Headers object with empty dict return Headers({}) + _safe_get_headers_from_scope = safe_get_headers_from_scope + @staticmethod def _reject_duplicate_authorization(raw_headers: object) -> None: """Raise 400 when the raw ASGI headers carry more than one ``Authorization`` header.""" @@ -1638,7 +1659,7 @@ class MCPRequestHandler: # matters: the no_mcp_servers opt-out below reads the caller's own object_permission, so above # this branch a user's own opt-out would wrongly zero their TEAMS' grants too (each source is # independent; an opt-out silences only its own source, inside the recursive call). - if _is_mcp_admitted_user_subject(user_api_key_auth) and user_api_key_auth is not None: + if is_mcp_admitted_user_subject(user_api_key_auth) and user_api_key_auth is not None: return MCPServerAccess( server_ids=tuple(await MCPRequestHandler.resolve_admitted_subject_servers(user_api_key_auth)), ) @@ -1950,18 +1971,18 @@ class MCPRequestHandler: # already-exceeded state; ATTRIBUTION of new spend stays with the user (documented deferral). from litellm.exceptions import BudgetExceededError from litellm.proxy.auth.auth_checks import ( - _organization_max_budget_check, - _team_max_budget_check, + organization_max_budget_check, + team_max_budget_check, ) source_view: Final = MCPRequestHandler._scoped_source_auth( auth, team_id=team_id, org_id=team_obj.organization_id or auth.org_id, carry_user_grants=False ) try: - await _team_max_budget_check( + await team_max_budget_check( team_object=team_obj, valid_token=source_view, proxy_logging_obj=proxy_logging_obj ) - await _organization_max_budget_check( + await organization_max_budget_check( valid_token=source_view, team_object=team_obj, prisma_client=prisma_client, @@ -2042,14 +2063,14 @@ class MCPRequestHandler: owning ``org_id`` so the team's budget accumulates and the right org is charged. Falls back to user-level attribution (rather than guessing a team) when the tool name does not resolve to a server, reusing the manager's own tool-name lookup.""" - if not _is_mcp_admitted_user_subject(auth): + if not is_mcp_admitted_user_subject(auth): return auth try: from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( global_mcp_server_manager, ) - server: Final = global_mcp_server_manager._get_mcp_server_from_tool_name(tool_name) + server: Final = global_mcp_server_manager.get_mcp_server_from_tool_name(tool_name) if server is None: return auth source: Final = await MCPRequestHandler.attributing_source_for_server(auth, server.server_id) @@ -2330,7 +2351,7 @@ class MCPRequestHandler: # FIRST statement, mirroring get_allowed_mcp_servers: a keyless admitted subject resolves per # source and shares nothing with the single-credential prelude below. Ordering is the invariant: # sat after the prelude, a fault in a lookup the subject never uses denied tools its teams grant. - if _is_mcp_admitted_user_subject(user_api_key_auth): + if is_mcp_admitted_user_subject(user_api_key_auth): return await MCPRequestHandler.resolve_admitted_subject_tools(server_id, user_api_key_auth) # Get key and team object permissions (already loaded in main auth flow) @@ -2422,7 +2443,9 @@ class MCPRequestHandler: # not collapse to allow-all (None); key/JWT auth keeps its prior allow-all-on-error. Both # keyless_source AND the marker are needed: each source resolves through an UNMARKED auth, so # without keyless_source a fault under a source returns None and wins the union as allow-all. - deny_all = unreadable_entitlement or keyless_source or _is_mcp_admitted_user_subject(user_api_key_auth) + deny_all: Final = ( + unreadable_entitlement or keyless_source or is_mcp_admitted_user_subject(user_api_key_auth) + ) return [] if deny_all else None @staticmethod @@ -2557,7 +2580,7 @@ class MCPRequestHandler: global_mcp_server_manager, ) from litellm.proxy.auth.auth_checks import ( - _get_mcp_server_ids_from_access_groups, + get_mcp_server_ids_from_access_groups, ) from litellm.proxy.proxy_server import ( prisma_client, @@ -2565,7 +2588,7 @@ class MCPRequestHandler: user_api_key_cache, ) - raw_server_ids: Final = await _get_mcp_server_ids_from_access_groups( + raw_server_ids: Final = await get_mcp_server_ids_from_access_groups( access_group_ids=user_api_key_auth.access_group_ids or [], prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, @@ -2647,7 +2670,7 @@ class MCPRequestHandler: ) # Get MCP servers from access groups - access_group_servers: Final = await MCPRequestHandler._get_mcp_servers_from_access_groups( + access_group_servers: Final = await MCPRequestHandler.get_mcp_servers_from_access_groups( key_object_permission.mcp_access_groups or [], requires_fresh_policy=user_api_key_auth.requires_fresh_policy, ) @@ -2763,7 +2786,7 @@ class MCPRequestHandler: return set(team_access_group_servers) if SpecialMCPServerName.all_proxy_servers.value in (object_permissions.mcp_servers or []): return set(global_mcp_server_manager.get_registry().keys()) - legacy_access_group_servers: Final = await MCPRequestHandler._get_mcp_servers_from_access_groups( + legacy_access_group_servers: Final = await MCPRequestHandler.get_mcp_servers_from_access_groups( object_permissions.mcp_access_groups or [], requires_fresh_policy=requires_fresh_policy, ) @@ -2797,7 +2820,7 @@ class MCPRequestHandler: """ try: from litellm.proxy.auth.auth_checks import ( - _get_mcp_server_ids_from_access_groups, + get_mcp_server_ids_from_access_groups, get_team_object, ) from litellm.proxy.proxy_server import ( @@ -2825,7 +2848,7 @@ class MCPRequestHandler: # pinned to a single team_id, but a keyless admitted identity (no team_id) unions # across all of its teams and would otherwise inherit a blocked team's MCP grants. return [] - team_access_group_servers: Final = await _get_mcp_server_ids_from_access_groups( + team_access_group_servers: Final = await get_mcp_server_ids_from_access_groups( access_group_ids=team_obj.access_group_ids or [], prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, @@ -2968,7 +2991,7 @@ class MCPRequestHandler: # Expand names/aliases to canonical server IDs (consistent with key/team/end-user path) direct_mcp_servers = global_mcp_server_manager.expand_permission_list(object_permissions.mcp_servers or []) - access_group_servers: Final = await MCPRequestHandler._get_mcp_servers_from_access_groups( + access_group_servers: Final = await MCPRequestHandler.get_mcp_servers_from_access_groups( object_permissions.mcp_access_groups or [], requires_fresh_policy=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy), ) @@ -3073,7 +3096,7 @@ class MCPRequestHandler: direct_mcp_servers = global_mcp_server_manager.expand_permission_list(object_permission.mcp_servers or []) # Get MCP servers from access groups - access_group_servers: Final = await MCPRequestHandler._get_mcp_servers_from_access_groups( + access_group_servers: Final = await MCPRequestHandler.get_mcp_servers_from_access_groups( object_permission.mcp_access_groups or [], requires_fresh_policy=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy), ) @@ -3212,7 +3235,7 @@ class MCPRequestHandler: direct_mcp_servers = global_mcp_server_manager.expand_permission_list(object_permissions.mcp_servers or []) fresh: Final = bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy) - access_group_servers: Final = await MCPRequestHandler._get_mcp_servers_from_access_groups( + access_group_servers: Final = await MCPRequestHandler.get_mcp_servers_from_access_groups( object_permissions.mcp_access_groups or [], requires_fresh_policy=fresh, ) @@ -3318,7 +3341,7 @@ class MCPRequestHandler: return False object_permission: Final = user_api_key_auth.object_permission credential_scoped: Final = ( - not _is_mcp_admitted_user_subject(user_api_key_auth) + not is_mcp_admitted_user_subject(user_api_key_auth) and object_permission is not None and object_permission.mcp_servers is not None ) @@ -3566,7 +3589,7 @@ class MCPRequestHandler: expanded_direct_servers: Final = global_mcp_server_manager.expand_permission_list( obj_perm.mcp_servers or [] ) - access_group_servers: Final = await MCPRequestHandler._get_mcp_servers_from_access_groups( + access_group_servers: Final = await MCPRequestHandler.get_mcp_servers_from_access_groups( obj_perm.mcp_access_groups or [], requires_fresh_policy=user_api_key_auth.requires_fresh_policy, ) @@ -3696,7 +3719,7 @@ class MCPRequestHandler: return server_ids @staticmethod - async def _get_mcp_servers_from_access_groups( + async def get_mcp_servers_from_access_groups( access_groups: list[str], *, requires_fresh_policy: bool = False, @@ -3735,6 +3758,8 @@ class MCPRequestHandler: verbose_logger.warning("Failed to get MCP servers from access groups: %s", e) return [] + _get_mcp_servers_from_access_groups = get_mcp_servers_from_access_groups + @staticmethod async def get_mcp_access_groups( user_api_key_auth: UserAPIKeyAuth | None = None, @@ -3863,5 +3888,5 @@ class MCPRequestHandler: """ Extract and parse the x-mcp-access-groups header from an ASGI scope. """ - headers: Final = MCPRequestHandler._safe_get_headers_from_scope(scope) + headers: Final = MCPRequestHandler.safe_get_headers_from_scope(scope) return MCPRequestHandler.get_mcp_access_groups_from_headers(headers) diff --git a/litellm/proxy/_experimental/mcp_server/bridge_token_flow.py b/litellm/proxy/_experimental/mcp_server/bridge_token_flow.py index 77bdbd26b35..9f393f25e49 100644 --- a/litellm/proxy/_experimental/mcp_server/bridge_token_flow.py +++ b/litellm/proxy/_experimental/mcp_server/bridge_token_flow.py @@ -14,8 +14,9 @@ from typing_extensions import assert_never from litellm._logging import verbose_logger from litellm.proxy._experimental.mcp_server.oauth_utils import TOKEN_NO_CACHE_HEADERS -from litellm.proxy.common_utils.encrypt_decrypt_utils import ( - _V2_GCM_PREFIX, # pyright: ignore[reportPrivateUsage] # reuse the encrypted credential's format discriminator +from litellm.proxy.common_utils.encrypt_decrypt_utils import ( # noqa: F401 # legacy module exports + _V2_GCM_PREFIX, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + V2_GCM_PREFIX, # pyright: ignore[reportPrivateUsage] # reuse the encrypted credential's format discriminator ) from litellm.types.mcp_server.mcp_server_manager import MCPServer @@ -99,7 +100,7 @@ async def _opaque_bearer_is_gateway_credential(token: str) -> bool: user_api_key_cache, ) - if is_envelope(token) or is_refresh_envelope(token) or token.startswith(_V2_GCM_PREFIX): + if is_envelope(token) or is_refresh_envelope(token) or token.startswith(V2_GCM_PREFIX): return True try: if ExperimentalUIJWTToken.get_key_object_from_ui_hash_key(token) is not None: @@ -283,12 +284,15 @@ async def _reload_active_key_by_hash(key_hash: str) -> "_ResolvedKey | _KeyResol return _ResolvedKey(key_hash=key_hash, key=key_obj) -async def _reload_active_user_by_id(user_id: str) -> "_KeyResolutionFailure | None": +async def reload_active_user_by_id(user_id: str) -> "_KeyResolutionFailure | None": """``None`` when the user is live, else the precise failure ``load_active_user_by_id`` found.""" loaded: Final = await load_active_user_by_id(user_id) return loaded if isinstance(loaded, str) else None +_reload_active_user_by_id: Final = reload_active_user_by_id + + UserRowSource = Literal["cache", "database"] @@ -400,12 +404,12 @@ async def _revalidate_active_subject(identity: "EnvelopeIdentity") -> "_KeyResol return "no_active_key" return None case "user_id": - return await _reload_active_user_by_id(identity.subject) + return await reload_active_user_by_id(identity.subject) case _: assert_never(identity.subject_type) -async def _extract_user_id_from_request(request: Request) -> str | None: +async def extract_user_id_from_request(request: Request) -> str | None: """Resolve the caller for identity binding without granting credential-write permission.""" from litellm.proxy.auth.handle_jwt import JWTIdentity # noqa: PLC0415 # proxy import cycle @@ -415,6 +419,9 @@ async def _extract_user_id_from_request(request: Request) -> str | None: return _active_key_user_id(resolved) if resolved is not None else None +_extract_user_id_from_request: Final = extract_user_id_from_request + + async def authorize_oauth_credential_request(request: Request, server_id: str) -> str | None: from litellm.proxy._types import UserAPIKeyAuth # noqa: PLC0415 # proxy import cycle @@ -448,13 +455,13 @@ async def can_store_oauth_credential(request: Request, auth: "UserAPIKeyAuth", s ) from litellm.proxy.auth.route_checks import RouteChecks # noqa: PLC0415 # proxy import cycle from litellm.proxy.auth.user_api_key_auth import ( # noqa: PLC0415 # proxy import cycle - _run_centralized_common_checks, # pyright: ignore[reportPrivateUsage] # reuse admission policy for the credential-write action + run_centralized_common_checks, # pyright: ignore[reportPrivateUsage] # reuse admission policy for the credential-write action ) write_route: Final = f"/v1/mcp/server/{server_id}/oauth-user-credential" try: RouteChecks.is_virtual_key_allowed_to_call_route(route=write_route, valid_token=auth, request=request) - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=auth, request=request, request_data={}, @@ -646,7 +653,7 @@ class _BridgeMintReady: keys: "EnvelopeKeys" -def _bridge_mint_error_response(error: _BridgeMintError) -> JSONResponse: +def bridge_mint_error_response(error: _BridgeMintError) -> JSONResponse: """Map a bridge-mint failure value to its token-endpoint response: one place, RFC 6749 §5.2 shape (top-level ``error``, no-store headers) for every case, with a status truthful about where the failure is. The caller's request is 400, a transient gateway outage is 503, a gateway @@ -731,6 +738,9 @@ def _bridge_mint_error_response(error: _BridgeMintError) -> JSONResponse: ) +_bridge_mint_error_response: Final = bridge_mint_error_response + + def _key_resolution_failure_to_mint_error(failure: _KeyResolutionFailure) -> _BridgeMintError: """Lift an identity-resolution failure into the mint taxonomy, preserving origin so the status stays truthful: the caller's missing credential is 400, a transient DB outage is 503, and a gateway that @@ -759,7 +769,7 @@ def _upstream_rejection_to_mint_error(rejection: _UpstreamGrantRejection) -> _Br assert_never(rejection) -async def _prepare_bridge_mint( +async def prepare_bridge_mint( request: Request, mcp_server: MCPServer, bridge_identity: "_BridgeAuthorizationCode | None" = None, @@ -820,6 +830,9 @@ async def _prepare_bridge_mint( return _BridgeMintReady(identity=identity, keys=keys) +_prepare_bridge_mint: Final = prepare_bridge_mint + + @dataclass(frozen=True, slots=True) class _BridgeRefreshReady: """A validated refresh request: the identity+keys to mint the renewed pair under, the upstream refresh @@ -854,7 +867,7 @@ def _refresh_key_failure_to_mint_error(failure: _KeyResolutionFailure) -> _Bridg assert_never(failure) -async def _prepare_bridge_refresh( +async def prepare_bridge_refresh( mcp_server: MCPServer, refresh_value: str | None ) -> "_BridgeRefreshReady | _BridgeMintError": """Phase 1 for the refresh_token grant, BEFORE the upstream exchange: open the client's refresh @@ -891,7 +904,10 @@ async def _prepare_bridge_refresh( ) -def _finish_bridge_mint( +_prepare_bridge_refresh: Final = prepare_bridge_refresh + + +def finish_bridge_mint( ready: "_BridgeMintReady", mcp_server: MCPServer, token_response: object, now: datetime ) -> "JSONResponse | _BridgeMintError": """Phase 3, AFTER the upstream exchange: seal the upstream grant into the client-held access envelope @@ -933,6 +949,9 @@ def _finish_bridge_mint( return JSONResponse(body, headers=TOKEN_NO_CACHE_HEADERS) +_finish_bridge_mint: Final = finish_bridge_mint + + def _upstream_refresh_credential(token_response: object) -> "RefreshCredential | None": """Extract the upstream refresh grant from a token response, or ``None`` when there is none to seal. Each field is isinstance-checked so nothing untyped reaches the refresh envelope; ``refresh_expires_in`` diff --git a/litellm/proxy/_experimental/mcp_server/byok_oauth_endpoints.py b/litellm/proxy/_experimental/mcp_server/byok_oauth_endpoints.py index aece755e4c4..2181e92baa9 100644 --- a/litellm/proxy/_experimental/mcp_server/byok_oauth_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/byok_oauth_endpoints.py @@ -83,12 +83,15 @@ def _oauth_token_error(code: str, status: int = 400) -> JSONResponse: return JSONResponse(status_code=status, content={"error": code}, headers=TOKEN_NO_CACHE_HEADERS) -def _user_id_from_session_cookie(request: Request) -> str | None: +def user_id_from_session_cookie(request: Request) -> str | None: """Return user_id from the UI ``token`` cookie, or None if missing/invalid.""" user_id, _ = _session_identity_from_cookie(request) return user_id +_user_id_from_session_cookie: Final = user_id_from_session_cookie + + def _session_identity_from_cookie(request: Request) -> tuple[str | None, str | None]: """Return ``(user_id, session_key)`` from the UI ``token`` cookie (HS256-signed with ``master_key``), or ``(None, None)`` if missing/invalid. diff --git a/litellm/proxy/_experimental/mcp_server/catalog.py b/litellm/proxy/_experimental/mcp_server/catalog.py index 4ef82a364f1..c9da8c1fcaa 100644 --- a/litellm/proxy/_experimental/mcp_server/catalog.py +++ b/litellm/proxy/_experimental/mcp_server/catalog.py @@ -842,7 +842,7 @@ async def get_filtered_server_tools( listed_generation: Final = global_mcp_server_manager.listed_tools_generation(server.server_id) if params is None: page = ListToolsResult( - tools=await global_mcp_server_manager._get_tools_from_server( + tools=await global_mcp_server_manager.get_tools_from_server( server=server, mcp_auth_header=server_auth_header, extra_headers=extra_headers, diff --git a/litellm/proxy/_experimental/mcp_server/db.py b/litellm/proxy/_experimental/mcp_server/db.py index 3b442795c55..e2d26054c80 100644 --- a/litellm/proxy/_experimental/mcp_server/db.py +++ b/litellm/proxy/_experimental/mcp_server/db.py @@ -33,13 +33,14 @@ from litellm.proxy._types import ( NewMCPServerRequest, UpdateMCPServerRequest, ) -from litellm.proxy.common_utils.encrypt_decrypt_utils import ( +from litellm.proxy.common_utils.encrypt_decrypt_utils import ( # noqa: F401 # legacy module exports SecretMapDecodeError, - _get_salt_key, + _get_salt_key, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export decode_secret_map, decrypt_value_helper, encrypt_secret_map, encrypt_value_helper, + get_salt_key, ) from litellm.proxy.utils import PrismaClient from litellm.repositories.config_repository import ConfigRepository @@ -464,7 +465,7 @@ def _prepare_mcp_server_data( blob_value = credentials.pop(te_field, None) if blob_value is not None and te_field not in data_dict: data_dict[te_field] = blob_value - data_dict["credentials"] = encrypt_credentials(credentials=credentials, encryption_key=_get_salt_key()) + data_dict["credentials"] = encrypt_credentials(credentials=credentials, encryption_key=get_salt_key()) data_dict["credentials"] = safe_dumps( _bind_submitted_oauth_client(data_dict["credentials"], data.issuer, data.url) if not exclude_unset and data.auth_type == "oauth2" @@ -1474,7 +1475,7 @@ async def upsert_mcp_server_oauth_client_credentials( same way regardless of which store a server's client came from.""" from litellm.litellm_core_utils.safe_json_dumps import safe_dumps - encrypted: Final = encrypt_credentials(credentials=MCPCredentials(**credentials), encryption_key=_get_salt_key()) + encrypted: Final = encrypt_credentials(credentials=MCPCredentials(**credentials), encryption_key=get_salt_key()) blob: Final = safe_dumps(encrypted) await _oauth_client_table_actions(prisma_client).upsert( where={"server_id": server_id}, @@ -1613,11 +1614,14 @@ def _parse_oauth_payload(decoded: str | None) -> OAuthCredentialPayload | None: return None -def _decode_oauth_payload(stored: str) -> OAuthCredentialPayload | None: +def decode_oauth_payload(stored: str) -> OAuthCredentialPayload | None: """Return the OAuth2 payload dict held in ``stored``, else ``None``.""" return _parse_oauth_payload(_decode_user_credential(stored)) +_decode_oauth_payload: Final = decode_oauth_payload + + async def rotate_mcp_user_credentials_master_key(prisma_client: PrismaClient, new_master_key: str): """Re-encrypt every ``LiteLLM_MCPUserCredentials`` row with ``new_master_key``. @@ -1865,7 +1869,7 @@ async def get_user_oauth_credential( def _server_user_credential_item( row: "prisma_db_models.LiteLLM_MCPUserCredentials", ) -> MCPServerUserCredentialListItem: - oauth_payload: Final = _decode_oauth_payload(row.credential_b64) + oauth_payload: Final = decode_oauth_payload(row.credential_b64) if oauth_payload is None: return MCPServerUserCredentialListItem( user_id=row.user_id, @@ -1986,7 +1990,7 @@ async def purge_user_oauth_credentials_for_server( invalidate_token_cache is injectable for tests; it defaults to the manager's shared invalidate_user_oauth_token_cache, the single invalidation point for per-user tokens.""" rows: Final = await _db_find_user_credential_rows(prisma_client, {"server_id": server_id}) - oauth_rows: Final = [row for row in rows if _decode_oauth_payload(row.credential_b64) is not None] + oauth_rows: Final = [row for row in rows if decode_oauth_payload(row.credential_b64) is not None] if not oauth_rows: return 0 deleted_count: Final = await _user_credential_actions(prisma_client).delete_many( @@ -2201,7 +2205,7 @@ async def resolve_user_oauth_access_token( return None try: from litellm.proxy._experimental.mcp_server.oauth2_token_cache import ( - _compute_per_user_token_ttl, + compute_per_user_token_ttl, mcp_per_user_token_cache, ) @@ -2248,7 +2252,7 @@ async def resolve_user_oauth_access_token( access_token: Final[str] = cred["access_token"] if prefetched_creds is None: - ttl: Final = _compute_per_user_token_ttl(server, _remaining_token_seconds(cred.get("expires_at"))) + ttl: Final = compute_per_user_token_ttl(server, _remaining_token_seconds(cred.get("expires_at"))) await mcp_per_user_token_cache.set( user_id, server_id, access_token, ttl, identity_binding_proof=cred.get("identity_binding_proof") ) diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index 1f38742701e..da6ce954b1d 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -24,18 +24,24 @@ from litellm.proxy._experimental.mcp_server.auth.token_endpoint_auth import ( TokenEndpointAuthConfigError, normalize_token_endpoint_auth_method, ) -from litellm.proxy._experimental.mcp_server.bridge_token_flow import ( - _bridge_mint_error_response, +from litellm.proxy._experimental.mcp_server.bridge_token_flow import ( # noqa: F401 # legacy module exports + _bridge_mint_error_response, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export _BridgeMintReady, _BridgeRefreshReady, - _extract_user_id_from_request, - _finish_bridge_mint, - _prepare_bridge_mint, - _prepare_bridge_refresh, - _reload_active_user_by_id, + _extract_user_id_from_request, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _finish_bridge_mint, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _prepare_bridge_mint, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _prepare_bridge_refresh, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _reload_active_user_by_id, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export authorize_oauth_credential_request, + bridge_mint_error_response, can_store_oauth_credential, + extract_user_id_from_request, + finish_bridge_mint, oauth_authorization_uses_gateway_credential, + prepare_bridge_mint, + prepare_bridge_refresh, + reload_active_user_by_id, ) from litellm.proxy._experimental.mcp_server.catalog import public_catalog_operation from litellm.proxy._experimental.mcp_server.faults import ( @@ -91,7 +97,10 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import ( decrypt_value_helper, encrypt_value_helper, ) -from litellm.proxy.common_utils.http_parsing_utils import _read_request_body +from litellm.proxy.common_utils.http_parsing_utils import ( # noqa: F401 # legacy module exports + _read_request_body, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + read_request_body, +) from litellm.types.llms.base import LiteLLMBaseModel from litellm.types.mcp import MCPAuth, MCPCredentials from litellm.types.mcp_server.mcp_server_manager import MCPServer, MCPTokenEndpointAuthMethod @@ -431,10 +440,10 @@ def _session_cookie_user_id(request: Request) -> str | None: aggregate DCR flow's verbs receive the identity as a plain value instead of parsing cookies themselves.""" from litellm.proxy._experimental.mcp_server.byok_oauth_endpoints import ( # noqa: PLC0415 # circular import at module load - _user_id_from_session_cookie, + user_id_from_session_cookie, ) - return _user_id_from_session_cookie(request) + return user_id_from_session_cookie(request) def _redirect_to_litellm_login(request: Request) -> RedirectResponse: @@ -637,7 +646,7 @@ async def _store_per_user_token_server_side( client even when server-side storage fails. """ from litellm.proxy._experimental.mcp_server.oauth2_token_cache import ( # noqa: PLC0415 - _compute_per_user_token_ttl, + compute_per_user_token_ttl, mcp_per_user_token_cache, ) from litellm.proxy.utils import get_prisma_client_or_throw # noqa: PLC0415 @@ -693,7 +702,7 @@ async def _store_per_user_token_server_side( await global_mcp_server_manager.invalidate_user_oauth_token_cache(user_id, server.server_id) # Warm the Redis cache so the first subsequent MCP call is a cache hit - ttl: Final = _compute_per_user_token_ttl(server, expires_in) + ttl: Final = compute_per_user_token_ttl(server, expires_in) await mcp_per_user_token_cache.set( user_id=user_id, server_id=server.server_id, @@ -703,7 +712,7 @@ async def _store_per_user_token_server_side( ) -def _raise_if_not_oauth2(mcp_server: MCPServer) -> None: +def raise_if_not_oauth2(mcp_server: MCPServer) -> None: """Reject a server without upstream OAuth from the gateway's authorize/token/register flow. The client-forwarded token modes (``true_passthrough`` / ``oauth_delegate``) are allowed @@ -714,10 +723,10 @@ def _raise_if_not_oauth2(mcp_server: MCPServer) -> None: Authorize path with ``persist_credentials`` enabled writes nothing to the server row). """ from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415 # circular import with mcp_server_manager at module load - _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES, + UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES, ) - if mcp_server.auth_type in _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES: + if mcp_server.auth_type in UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES: return raise HTTPException( status_code=400, @@ -733,6 +742,9 @@ def _raise_if_not_oauth2(mcp_server: MCPServer) -> None: ) +_raise_if_not_oauth2: Final = raise_if_not_oauth2 + + def _endpoint_not_configured_detail( mcp_server: MCPServer, endpoint_label: str, @@ -919,7 +931,7 @@ async def _resolve_oauth_authorization_user( ) -> str | RedirectResponse: """Resolve the authorization subject without replacing denied credentials with cookie grants.""" from litellm.proxy._experimental.mcp_server.byok_oauth_endpoints import ( # noqa: PLC0415 # proxy import cycle - _user_id_from_session_cookie, + user_id_from_session_cookie, ) use_gateway_credential: Final = enforce_binding and await oauth_authorization_uses_gateway_credential(request) @@ -928,7 +940,7 @@ async def _resolve_oauth_authorization_user( ) if use_gateway_credential and request_user_id is None: return _bridge_access_denied_redirect(redirect_uri, state, mcp_server) - user_id: Final = request_user_id or _user_id_from_session_cookie(request) + user_id: Final = request_user_id or user_id_from_session_cookie(request) if user_id is None: return _redirect_to_litellm_login(request) if not await _user_can_reach_mcp_server(user_id, mcp_server.server_id): @@ -948,7 +960,7 @@ async def authorize_with_server( scope: str | None = None, ephemeral_dcr_client: "EphemeralDcrClient | None" = None, ): - _raise_if_not_oauth2(mcp_server) + raise_if_not_oauth2(mcp_server) resolved_server: Final = await _server_with_oauth_endpoints(mcp_server, _register_flow_needed_endpoint) if not oauth_client_registration_matches( resolved_server.dcr_issuer, resolved_server.dcr_server_url, resolved_server.issuer, resolved_server.url @@ -1084,7 +1096,7 @@ async def exchange_token_with_server( scope: str | None = None, client_token_endpoint_auth_method: MCPTokenEndpointAuthMethod | None = None, ): - _raise_if_not_oauth2(mcp_server) + raise_if_not_oauth2(mcp_server) if grant_type not in ("authorization_code", "refresh_token"): raise HTTPException(status_code=400, detail="Unsupported grant_type") @@ -1131,7 +1143,7 @@ async def exchange_token_with_server( raise HTTPException(status_code=400, detail=str(exc)) from exc request_user_id: Final = ( - await _extract_user_id_from_request(request) + await extract_user_id_from_request(request) if resolved_server.needs_user_oauth_token or resolved_server.oauth_identity_binding is not None else None ) @@ -1148,9 +1160,9 @@ async def exchange_token_with_server( # identity, and unwrap the real upstream refresh token BEFORE building token_data, so the exchange # sends the upstream token and never the envelope. A failure returns without touching the upstream. if is_bridge: - prepared_refresh: Final = await _prepare_bridge_refresh(resolved_server, refresh_token) + prepared_refresh: Final = await prepare_bridge_refresh(resolved_server, refresh_token) if not isinstance(prepared_refresh, _BridgeRefreshReady): - return _bridge_mint_error_response(prepared_refresh) + return bridge_mint_error_response(prepared_refresh) bridge_mint_ready = prepared_refresh.ready bridge_upstream_refresh = prepared_refresh.upstream_refresh_token bridge_upstream_scope = prepared_refresh.upstream_scope @@ -1227,9 +1239,9 @@ async def exchange_token_with_server( # Phase 1 for a bridge authorization_code mint: resolve identity (the SSO user recovered above, or # the presented litellm key) and the envelope keys BEFORE the exchange consumes the single-use code. if is_bridge: - prepared: Final = await _prepare_bridge_mint(request, resolved_server, bridge_identity) + prepared: Final = await prepare_bridge_mint(request, resolved_server, bridge_identity) if not isinstance(prepared, _BridgeMintReady): - return _bridge_mint_error_response(prepared) + return bridge_mint_error_response(prepared) bridge_mint_ready = prepared refresh_binding: Final = resolved_server.oauth_identity_binding @@ -1269,7 +1281,7 @@ async def exchange_token_with_server( "re-runs authorization_code rather than an opaque upstream error", resolved_server.server_id, ) - return _bridge_mint_error_response("invalid_refresh") + return bridge_mint_error_response("invalid_refresh") return render_token_fault(fault) token_response = response.json() @@ -1357,10 +1369,10 @@ async def exchange_token_with_server( token_response = {**token_response, "scope": refresh_request_scope} # Phase 3: seal the upstream grant into the client-held envelope; failures map through the same # OAuth-shaped response as the phase-1 preconditions. - minted: Final = _finish_bridge_mint( + minted: Final = finish_bridge_mint( bridge_mint_ready, resolved_server, token_response, datetime.now(timezone.utc) ) - return minted if isinstance(minted, JSONResponse) else _bridge_mint_error_response(minted) + return minted if isinstance(minted, JSONResponse) else bridge_mint_error_response(minted) raw_access_token: Final = token_response.get("access_token") if isinstance(token_response, dict) else None if not isinstance(raw_access_token, str) or not raw_access_token: @@ -1937,7 +1949,7 @@ async def register_client_with_server( client_redirect_uris: list[str] | None = None, client_application_type: Literal["native", "web"] | None = None, ): - _raise_if_not_oauth2(mcp_server) + raise_if_not_oauth2(mcp_server) request_base_url: Final = get_request_base_url(request) current_redirect_uri: Final = f"{request_base_url}/callback" client_facing_redirect_uris: Final = client_redirect_uris or [current_redirect_uri] @@ -2111,7 +2123,7 @@ async def authorize( mcp_server = _resolve_oauth2_server_for_root_endpoints(client_ip=client_ip) if mcp_server is None: raise HTTPException(status_code=404, detail="MCP server not found") - _raise_if_not_oauth2(mcp_server) + raise_if_not_oauth2(mcp_server) # Use server's stored client_id when caller doesn't supply one. # Raise a clear error instead of passing an empty string — an empty # client_id would silently produce a broken authorization URL. @@ -2181,7 +2193,7 @@ async def token_endpoint( code_verifier=code_verifier, refresh_token=refresh_token, master_key=master_key, - reload_user=_reload_active_user_by_id, + reload_user=reload_active_user_by_id, cache=user_api_key_cache, resource=resource, mint_proxy_credential=mint_proxy_credential, @@ -2303,7 +2315,7 @@ async def introspect_endpoint(token: str = Form(...)) -> Response: return await introspect_gateway_token( token=token, master_key=master_key, - reload_user=_reload_active_user_by_id, + reload_user=reload_active_user_by_id, cache=user_api_key_cache, ) @@ -3102,7 +3114,7 @@ async def register_client(request: Request, mcp_server_name: str | None = None): # Get the correct base URL considering X-Forwarded-* headers request_base_url: Final = get_request_base_url(request) - request_data: Final = await _read_request_body(request=request) + request_data: Final = await read_request_body(request=request) data: Final[dict] = {**request_data} client_redirect_uris: Final = client_supplied_redirect_uris(data.get("redirect_uris")) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_context.py b/litellm/proxy/_experimental/mcp_server/mcp_context.py index 11325a9f127..8ec44f65bdf 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_context.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_context.py @@ -27,16 +27,22 @@ def get_active_mcp_request_ctx() -> "ServerRequestContext | None": # Set server-side in proxy_server.py route handlers when a request arrives via # /toolset/{name}/mcp or the toolset fallback in dynamic_mcp_route. # Never populated from client-supplied headers. -_mcp_active_toolset_id: Final[ContextVar[str | None]] = ContextVar("_mcp_active_toolset_id", default=None) +mcp_active_toolset_id: Final[ContextVar[str | None]] = ContextVar("_mcp_active_toolset_id", default=None) + +_mcp_active_toolset_id: Final = mcp_active_toolset_id # Per-request merged InitializeResult.instructions; set in MCP HTTP/SSE handlers. -_mcp_gateway_initialize_instructions: Final[ContextVar[str | None]] = ContextVar( +mcp_gateway_initialize_instructions: Final[ContextVar[str | None]] = ContextVar( "_mcp_gateway_initialize_instructions", default=None ) +_mcp_gateway_initialize_instructions: Final = mcp_gateway_initialize_instructions + # Per-request scoped server name; set in MCP HTTP/SSE handlers when the path # identifies exactly one upstream server. Never populated from client-supplied headers. -_mcp_gateway_server_name: Final[ContextVar[str | None]] = ContextVar("_mcp_gateway_server_name", default=None) +mcp_gateway_server_name: Final[ContextVar[str | None]] = ContextVar("_mcp_gateway_server_name", default=None) + +_mcp_gateway_server_name: Final = mcp_gateway_server_name # Set server-side by the /mcp/proxy route. Never populated from client-supplied headers. _mcp_proxy_mode: Final[ContextVar[bool]] = ContextVar("_mcp_proxy_mode", default=False) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_debug.py b/litellm/proxy/_experimental/mcp_server/mcp_debug.py index f3fc1c9d44c..a63e930ec2b 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_debug.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_debug.py @@ -390,7 +390,7 @@ class MCPDebug: server_auth_type = server.auth_type break - scope_headers: Final = MCPRequestHandler._safe_get_headers_from_scope(scope) + scope_headers: Final = MCPRequestHandler.safe_get_headers_from_scope(scope) litellm_key: Final = MCPRequestHandler.get_litellm_api_key_from_headers(scope_headers) return MCPDebug.build_debug_headers( diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 1a31c91ad45..15145a20e33 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -76,10 +76,11 @@ from litellm.integrations.custom_guardrail import ( ) from litellm.litellm_core_utils.url_utils import SSRFError, async_safe_get from litellm.llms.custom_httpx.http_handler import get_async_httpx_client -from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( +from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( # noqa: F401 # legacy module exports MCPRequestHandler, MCPServerAccess, - _is_mcp_admitted_user_subject, + _is_mcp_admitted_user_subject, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + is_mcp_admitted_user_subject, ) from litellm.proxy._experimental.mcp_server.catalog import _configuration_identity, _DiscoveryCache, _DiscoveryKey from litellm.proxy._experimental.mcp_server.contracts import OperationContext @@ -100,10 +101,11 @@ from litellm.proxy._experimental.mcp_server.oauth2_token_cache import ( MCPPerUserTokenCache, mcp_per_user_token_cache, ) -from litellm.proxy._experimental.mcp_server.oauth_utils import ( - _redact_mcp_resource_url, +from litellm.proxy._experimental.mcp_server.oauth_utils import ( # noqa: F401 # legacy module exports + _redact_mcp_resource_url, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export canonicalize_url_identity, get_byok_www_authenticate, + redact_mcp_resource_url, ) from litellm.proxy._experimental.mcp_server.outbound_credentials import ( Error, @@ -285,12 +287,14 @@ class ListedToolsCaller: # gateway discovers from the upstream itself: interactive oauth2 and the two client-forwarded modes. # OBO/M2M endpoint discovery is decided separately via _obo_needs_endpoint_discovery. Shared by the # config-YAML and DB server loaders so the two paths cannot drift on which modes trigger discovery. -_UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES: Final[tuple[MCPAuth, ...]] = ( +UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES: Final[tuple[MCPAuth, ...]] = ( MCPAuth.oauth2, MCPAuth.true_passthrough, MCPAuth.oauth_delegate, ) +_UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES: Final = UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES + _MCP_OAUTH_DISCOVERY_ON_STARTUP_ENV: Final = "LITELLM_MCP_OAUTH_DISCOVERY_ON_STARTUP" _TRUE_ENV_VALUES: Final = frozenset(("1", "true", "yes", "on")) @@ -756,7 +760,7 @@ def _flow_endpoints_missing( # A configured exchange endpoint replaces discovery entirely; only a server that must # discover its token endpoint and still has none is unresolved. return token_exchange_endpoint is None and token_url is None - if auth_type not in _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES: + if auth_type not in UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES: return False if oauth2_flow == "client_credentials": return token_url is None @@ -926,7 +930,7 @@ def _restrict_discovery_to_corroborated_authorization_server( def _redacted_origin_list(urls: Sequence[str]) -> str: - return ", ".join(_redact_mcp_resource_url(url) or "" for url in urls) + return ", ".join(redact_mcp_resource_url(url) or "" for url in urls) def _sanitized_error_text(exc: Exception) -> str: @@ -1078,7 +1082,7 @@ def _warn_oauth_endpoints_unresolved( "(RFC 8414)", server_ref, ", ".join(unresolved), - _redact_mcp_resource_url(server_url) or "", + redact_mcp_resource_url(server_url) or "", ) return verbose_logger.warning( @@ -1107,7 +1111,7 @@ def _write_user_env_vars_cache(user_id: str, server_id: str, values: dict[str, s _user_env_vars_cache[cache_key] = (values, time.monotonic()) -def _should_strip_caller_authorization( +def should_strip_caller_authorization( mcp_server: MCPServer, raw_headers: dict[str, str] | None, user_api_key_auth: UserAPIKeyAuth | None, @@ -1171,6 +1175,9 @@ def _should_strip_caller_authorization( ) +_should_strip_caller_authorization: Final = should_strip_caller_authorization + + LITELLM_VIRTUAL_KEY_PREFIX: Final = "sk-" @@ -1272,7 +1279,7 @@ def _openapi_forwarded_extra_headers( if not mcp_server.extra_headers or not raw_headers: return None normalized_raw: Final = {str(k).lower(): v for k, v in raw_headers.items() if isinstance(k, str)} - skip_caller_authorization: Final = _should_strip_caller_authorization( + skip_caller_authorization: Final = should_strip_caller_authorization( mcp_server=mcp_server, raw_headers=raw_headers, user_api_key_auth=user_api_key_auth, @@ -1289,7 +1296,7 @@ def _openapi_forwarded_extra_headers( return forwarded or None -def _resolve_openapi_tool_auth( +def resolve_openapi_tool_auth( mcp_server: MCPServer, mcp_auth_header: str | None, mcp_server_auth_headers: Mapping[str, str | dict[str, str]] | None, @@ -1339,6 +1346,9 @@ def _resolve_openapi_tool_auth( return None, forwarded, None +_resolve_openapi_tool_auth: Final = resolve_openapi_tool_auth + + async def _resolve_byok_mcp_auth_header( mcp_server: MCPServer, user_api_key_auth: UserAPIKeyAuth | None, @@ -1385,7 +1395,7 @@ def _catalog_auth_header( return mcp_auth_header if catalog_auth_header is ... else catalog_auth_header -def _client_forwarded_authorization_headers( +def client_forwarded_authorization_headers( mcp_server: MCPServer, oauth2_headers: dict[str, str] | None, raw_headers: dict[str, str] | None, @@ -1399,7 +1409,7 @@ def _client_forwarded_authorization_headers( paths cannot drift, mirroring the ``_should_strip_caller_authorization`` split. """ extra_headers: Final = oauth2_headers.copy() if oauth2_headers else None - if extra_headers and _should_strip_caller_authorization( + if extra_headers and should_strip_caller_authorization( mcp_server=mcp_server, raw_headers=raw_headers, user_api_key_auth=user_api_key_auth, @@ -1408,6 +1418,9 @@ def _client_forwarded_authorization_headers( return extra_headers +_client_forwarded_authorization_headers: Final = client_forwarded_authorization_headers + + async def _materialize_auth_headers(auth: httpx2.Auth | None) -> dict[str, str] | None: """Extract the header a resolved ``httpx2.Auth`` would set, as a plain dict, or None. @@ -1472,7 +1485,7 @@ def _redacted_registry_dump(servers: dict[str, MCPServer]) -> dict[str, dict[str } -def _caller_authorization_fans_out( +def caller_authorization_fans_out( server: MCPServer, scope_servers: list[MCPServer] | None, ) -> bool: @@ -1489,6 +1502,9 @@ def _caller_authorization_fans_out( ) +_caller_authorization_fans_out: Final = caller_authorization_fans_out + + def _extract_upstream_auth_failure( exc: BaseException, ) -> tuple[int, str | None] | None: @@ -1986,7 +2002,7 @@ class MCPServerManager: manual_issuer: Final = _blank_to_none(server.issuer) manual_authorization_url: Final = _blank_to_none(server.authorization_url) manual_token_url: Final = _blank_to_none(server.token_url) - is_discovery_auth_type: Final = server.auth_type in _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES + is_discovery_auth_type: Final = server.auth_type in UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES use_issuer_anchor: Final = server.issuer_is_anchored obo_needs_discovery: Final = self._obo_needs_endpoint_discovery( server.auth_type, @@ -2276,7 +2292,7 @@ class MCPServerManager: if raw and str(raw).strip(): self._upstream_initialize_instructions_by_server_id[server.server_id] = str(raw).strip() - async def _ensure_upstream_initialize_instructions_cached(self, server: MCPServer) -> None: + async def ensure_upstream_initialize_instructions_cached(self, server: MCPServer) -> None: """ Open one upstream session and cache InitializeResult.instructions if missing. @@ -2326,7 +2342,7 @@ class MCPServerManager: raise_on_missing=False, ) extra_headers: dict[str, str] | None = dict(resolved_static_headers) if resolved_static_headers else None - client: Final = await self._create_mcp_client( + client: Final = await self.create_mcp_client( server=server, mcp_auth_header=None, extra_headers=extra_headers, @@ -2345,6 +2361,8 @@ class MCPServerManager: e, ) + _ensure_upstream_initialize_instructions_cached = ensure_upstream_initialize_instructions_cached + def get_registry(self) -> Mapping[str, MCPServer]: """ Get the registered MCP Servers from the registry and union with the config MCP Servers @@ -2458,7 +2476,7 @@ class MCPServerManager: 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")) - is_discovery_auth_type = auth_type in _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES + is_discovery_auth_type = auth_type in UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES obo_needs_discovery = self._obo_needs_endpoint_discovery( auth_type, server_config.get("token_exchange_endpoint"), @@ -3110,7 +3128,7 @@ class MCPServerManager: 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) - is_discovery_auth_type: Final = auth_type in _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES + is_discovery_auth_type: Final = auth_type in UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES token_exchange_endpoint: Final = mcp_server.token_exchange_endpoint or ( credentials_dict.get("token_exchange_endpoint") if credentials_dict else None ) @@ -3452,7 +3470,7 @@ class MCPServerManager: # applying this rule would hide almost every admitted user's OWN submitted servers. Their # submissions are theirs by authorship, and their scope comes from the per-source union. has_explicit_object_permission: Final = ( - not _is_mcp_admitted_user_subject(user_api_key_auth) + not is_mcp_admitted_user_subject(user_api_key_auth) and key_object_permission is not None and (key_object_permission.mcp_servers is not None) ) @@ -3470,7 +3488,7 @@ class MCPServerManager: the exception fallback, and applied AFTER every union (grants, operator-open, submitted) because the scope is a ceiling over the whole reachable set; a resolver fault therefore never widens a scoped bearer to the allow-all set.""" - if user_api_key_auth is None or not _is_mcp_admitted_user_subject(user_api_key_auth): + if user_api_key_auth is None or not is_mcp_admitted_user_subject(user_api_key_auth): return None return user_api_key_auth.mcp_session_resource_server_id @@ -3505,7 +3523,7 @@ class MCPServerManager: # rides the HUMAN, not the credential: an admin's session resolves the same registry their # dashboard shows (connect-page parity), bounded like an admin key by explicit # object_permission scope, the entitlement ceiling, and the session resource scope below. - is_admitted_subject: Final = _is_mcp_admitted_user_subject(user_api_key_auth) + is_admitted_subject: Final = is_mcp_admitted_user_subject(user_api_key_auth) # The key explicitly opted out of every MCP server. Return zero before # layering on allow_all_keys or submitted servers so the opt-out is absolute. @@ -3759,7 +3777,7 @@ class MCPServerManager: blocked = 0 for sid in server_ids: s = self.get_mcp_server_by_id(sid) - if s is not None and self._is_server_accessible_from_ip(s, client_ip): + if s is not None and self.is_server_accessible_from_ip(s, client_ip): allowed.append(sid) elif s is not None: blocked += 1 @@ -3774,7 +3792,7 @@ class MCPServerManager: if server is None: verbose_logger.warning("MCP Server %s not found", server_id) return [] - return list(await self._get_tools_from_server(server)) + return list(await self.get_tools_from_server(server)) except Exception as e: verbose_logger.warning("Failed to get tools from server %s: %s", server_id, e) return [] @@ -3812,7 +3830,7 @@ class MCPServerManager: try: tools: Final = list( - await self._get_tools_from_server( + await self.get_tools_from_server( server=server, mcp_auth_header=server_auth_header, user_api_key_auth=user_api_key_auth, @@ -3861,7 +3879,7 @@ class MCPServerManager: return None @staticmethod - def _extract_subject_token( + def extract_subject_token( oauth2_headers: Mapping[str, str] | None, raw_headers: Mapping[str, str] | None, user_api_key_auth: UserAPIKeyAuth | None, @@ -3878,6 +3896,8 @@ class MCPServerManager: return None return bearer + _extract_subject_token = extract_subject_token + def _obo_subject_token( self, server: MCPServer, @@ -3892,9 +3912,9 @@ class MCPServerManager: """ if server.auth_type != MCPAuth.oauth2_token_exchange: return None - return self._extract_subject_token(None, raw_headers, user_api_key_auth) + return self.extract_subject_token(None, raw_headers, user_api_key_auth) - def _build_stdio_env( + def build_stdio_env( self, server: MCPServer, raw_headers: Mapping[str, str] | None = None, @@ -3921,6 +3941,8 @@ class MCPServerManager: return resolved_env + _build_stdio_env = build_stdio_env + def _references_per_user_env_var(self, server: MCPServer) -> bool: """True when ``server.static_headers`` reference a per-user ``${NAME}`` env var. @@ -4112,7 +4134,7 @@ class MCPServerManager: Only OBO has a discovery challenge to raise; ID-JAG's failures are plain statuses whose body already names what the user has to do, so they map through ``raise_public`` as at egress. """ - subject_token: Final = self._extract_subject_token(oauth2_headers, raw_headers, user_api_key_auth) + subject_token: Final = self.extract_subject_token(oauth2_headers, raw_headers, user_api_key_auth) match server.auth_type: case MCPAuth.oauth2_token_exchange: if not self._extract_bearer_token(oauth2_headers, None): @@ -4140,7 +4162,7 @@ class MCPServerManager: ) raise_public(err) - async def _create_mcp_client( + async def create_mcp_client( self, server: MCPServer, mcp_auth_header: str | dict[str, str] | None = None, @@ -4182,7 +4204,9 @@ class MCPServerManager: elicitation_callback=(_create_elicitation_callback() if resolved_server.allow_elicitation else None), ) - async def _get_tools_from_server( + _create_mcp_client = create_mcp_client + + async def get_tools_from_server( self, server: MCPServer, mcp_auth_header: str | dict[str, str] | None = None, @@ -4212,6 +4236,8 @@ class MCPServerManager: ) return result.tools + _get_tools_from_server = get_tools_from_server + async def get_tools_page( self, server: MCPServer, @@ -4302,18 +4328,18 @@ class MCPServerManager: for_list_tools=True, ) - stdio_env: Final = self._build_stdio_env(server, raw_headers) + stdio_env: Final = self.build_stdio_env(server, raw_headers) # token_exchange (OBO) discovery needs the caller's token too: list it with the user's own # token (mirrors the call path), not v1's deleted client_credentials fallback. Other modes # never read the inbound bearer, so leave subject_token None to avoid forwarding it. subject_token: Final = ( - self._extract_subject_token(oauth2_headers, raw_headers, user_api_key_auth) + self.extract_subject_token(oauth2_headers, raw_headers, user_api_key_auth) if server.auth_type == MCPAuth.oauth2_token_exchange else None ) - client = await self._create_mcp_client( + client = await self.create_mcp_client( # rebind-ok: pre-existing rebinding on a rename-only line server=server, mcp_auth_header=mcp_auth_header, extra_headers=extra_headers, @@ -4458,10 +4484,10 @@ class MCPServerManager: return None auth: Final = caller.user_api_key_auth forwarded: Final = self._forwarded_header_values(server, caller.raw_headers) or None - header_env: Final = self._build_stdio_env(server, caller.raw_headers) - stdio_env: Final = None if header_env == self._build_stdio_env(server) else header_env + header_env: Final = self.build_stdio_env(server, caller.raw_headers) + stdio_env: Final = None if header_env == self.build_stdio_env(server) else header_env caller_bearer: Final = ( - self._extract_subject_token(caller.oauth2_headers, caller.raw_headers, auth) + self.extract_subject_token(caller.oauth2_headers, caller.raw_headers, auth) if _consumes_caller_authorization(server) or server.auth_type == MCPAuth.oauth2_token_exchange else None ) @@ -4608,9 +4634,9 @@ class MCPServerManager: ) or None ) - stdio_env: Final = self._build_stdio_env(server, raw_headers) + stdio_env: Final = self.build_stdio_env(server, raw_headers) subject_token: Final = self._obo_subject_token(server, raw_headers, user_api_key_auth) - client: Final = await self._create_mcp_client( + client: Final = await self.create_mcp_client( server=server, mcp_auth_header=mcp_auth_header, extra_headers=headers, @@ -4657,9 +4683,9 @@ class MCPServerManager: ) or None ) - stdio_env: Final = self._build_stdio_env(server, raw_headers) + stdio_env: Final = self.build_stdio_env(server, raw_headers) subject_token: Final = self._obo_subject_token(server, raw_headers, user_api_key_auth) - client: Final = await self._create_mcp_client( + client: Final = await self.create_mcp_client( server=server, mcp_auth_header=mcp_auth_header, extra_headers=headers, @@ -4712,9 +4738,9 @@ class MCPServerManager: ) or None ) - stdio_env: Final = self._build_stdio_env(server, raw_headers) + stdio_env: Final = self.build_stdio_env(server, raw_headers) subject_token: Final = self._obo_subject_token(server, raw_headers, user_api_key_auth) - client: Final = await self._create_mcp_client( + client: Final = await self.create_mcp_client( server=server, mcp_auth_header=mcp_auth_header, extra_headers=headers, @@ -4767,9 +4793,9 @@ class MCPServerManager: ) or None ) - stdio_env: Final = self._build_stdio_env(server, raw_headers) + stdio_env: Final = self.build_stdio_env(server, raw_headers) subject_token: Final = self._obo_subject_token(server, raw_headers, user_api_key_auth) - client: Final = await self._create_mcp_client( + client: Final = await self.create_mcp_client( server=server, mcp_auth_header=mcp_auth_header, extra_headers=headers, @@ -4820,10 +4846,10 @@ class MCPServerManager: extra_headers = {} extra_headers.update(server.static_headers) - stdio_env: Final = self._build_stdio_env(server, raw_headers) + stdio_env: Final = self.build_stdio_env(server, raw_headers) subject_token: Final = self._obo_subject_token(server, raw_headers, user_api_key_auth) - client: Final = await self._create_mcp_client( + client: Final = await self.create_mcp_client( server=server, mcp_auth_header=mcp_auth_header, extra_headers=extra_headers, @@ -4857,10 +4883,10 @@ class MCPServerManager: extra_headers = {} extra_headers.update(server.static_headers) - stdio_env: Final = self._build_stdio_env(server, raw_headers) + stdio_env: Final = self.build_stdio_env(server, raw_headers) subject_token: Final = self._obo_subject_token(server, raw_headers, user_api_key_auth) - client: Final = await self._create_mcp_client( + client: Final = await self.create_mcp_client( server=server, mcp_auth_header=mcp_auth_header, extra_headers=extra_headers, @@ -4946,7 +4972,7 @@ class MCPServerManager: "MCP OAuth endpoint discovery against %s found no authorization server metadata. Attempts: %s. " "The MCP server url may be misconfigured, or the upstream may not support OAuth discovery " "(RFC 9728 / RFC 8414)", - _redact_mcp_resource_url(server_url) or "", + redact_mcp_resource_url(server_url) or "", "; ".join(attempts) if attempts else "none recorded", ) return metadata @@ -4957,7 +4983,7 @@ class MCPServerManager: *, allow_origin_fallback: bool, ) -> tuple[MCPOAuthMetadata | None, tuple[str, ...]]: - origin: Final = _redact_mcp_resource_url(server_url) or "" + origin: Final = redact_mcp_resource_url(server_url) or "" try: client: Final = get_async_httpx_client( llm_provider=httpxSpecialProvider.MCP, @@ -5006,7 +5032,7 @@ class MCPServerManager: *, allow_origin_fallback: bool, ) -> tuple[MCPOAuthMetadata | None, tuple[str, ...]]: - origin: Final = _redact_mcp_resource_url(server_url) or "" + origin: Final = redact_mcp_resource_url(server_url) or "" verbose_logger.debug( "MCP OAuth discovery for %s received status error: %s", server_url, @@ -5910,14 +5936,14 @@ class MCPServerManager: } # Create MCP request object for processing - mcp_request_obj: Final = proxy_logging_obj._create_mcp_request_object_from_kwargs(pre_hook_kwargs) + mcp_request_obj: Final = proxy_logging_obj.create_mcp_request_object_from_kwargs(pre_hook_kwargs) # Convert to LLM format for existing guardrail compatibility. # Unified guardrails read the seeded logger off the request dict and pass it # into ``apply_guardrail``, so ``@log_guardrail_information`` bridges their # evaluations itself; the ``finally`` below covers native guardrails, which # never receive it. Same seeding the pass-through routes do. - synthetic_llm_data: Final = proxy_logging_obj._convert_mcp_to_llm_format(mcp_request_obj, pre_hook_kwargs) + synthetic_llm_data: Final = proxy_logging_obj.convert_mcp_to_llm_format(mcp_request_obj, pre_hook_kwargs) synthetic_llm_data["litellm_logging_obj"] = litellm_logging_obj try: @@ -5930,7 +5956,9 @@ class MCPServerManager: await proxy_logging_obj.enforce_mcp_server_rate_limits(user_api_key_auth, server) if modified_data: # Convert response back to MCP format and apply modifications - modified_kwargs = proxy_logging_obj._convert_mcp_hook_response_to_kwargs(modified_data, pre_hook_kwargs) + modified_kwargs: Final = proxy_logging_obj.convert_mcp_hook_response_to_kwargs( + modified_data, pre_hook_kwargs + ) if modified_kwargs.get("arguments") != arguments: hook_result["arguments"] = modified_kwargs["arguments"] if modified_kwargs.get("extra_headers"): @@ -5990,7 +6018,7 @@ class MCPServerManager: } # Seeded for the same reason as in ``pre_call_tool_check``. - synthetic_llm_data: Final = proxy_logging_obj._convert_mcp_to_llm_format(request_obj, during_hook_kwargs) + synthetic_llm_data: Final = proxy_logging_obj.convert_mcp_to_llm_format(request_obj, during_hook_kwargs) synthetic_llm_data["litellm_logging_obj"] = litellm_logging_obj # Wrapped so the bridge runs inside the task: the caller only holds the task and @@ -6063,7 +6091,7 @@ class MCPServerManager: spec: Final = to_server_spec(mcp_server) if spec is not None: await self._cred_provider.invalidate_credentials(to_subject(user_api_key_auth, subject_token), spec) - retry_client: Final = await self._create_mcp_client( + retry_client: Final = await self.create_mcp_client( server=mcp_server, mcp_auth_header=server_auth_header, extra_headers=extra_headers, @@ -6132,7 +6160,9 @@ class MCPServerManager: MCPAuth.oauth2_token_exchange, MCPAuth.oauth2_id_jag, ): - subject_token = self._extract_subject_token(oauth2_headers, raw_headers, user_api_key_auth) + subject_token = self.extract_subject_token( # rebind-ok: pre-existing rebinding on a rename-only line + oauth2_headers, raw_headers, user_api_key_auth + ) elif mcp_server.auth_type == MCPAuth.oauth2: if mcp_server.has_client_credentials: # For M2M OAuth servers, Authorization must come from token fetch. @@ -6143,18 +6173,20 @@ class MCPServerManager: # token, so drop the caller-forwarded Authorization (apply-if-absent would # otherwise let it shadow the resolved token). Delegate keeps it. Centralized # via _should_strip_caller_authorization to match _prepare_mcp_server_headers. - if extra_headers and _should_strip_caller_authorization( + if extra_headers and should_strip_caller_authorization( mcp_server=mcp_server, raw_headers=raw_headers, user_api_key_auth=user_api_key_auth, ): extra_headers = without_header(extra_headers, DEFAULT_CREDENTIAL_HEADER) elif mcp_server.is_client_forwarded_token: - extra_headers = _client_forwarded_authorization_headers( - mcp_server=mcp_server, - oauth2_headers=oauth2_headers, - raw_headers=raw_headers, - user_api_key_auth=user_api_key_auth, + extra_headers = ( # rebind-ok: pre-existing rebinding on a rename-only line + client_forwarded_authorization_headers( + mcp_server=mcp_server, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + user_api_key_auth=user_api_key_auth, + ) ) if mcp_server.extra_headers and raw_headers: @@ -6162,7 +6194,7 @@ class MCPServerManager: extra_headers = {} normalized_raw_headers: Final = {str(k).lower(): v for k, v in raw_headers.items() if isinstance(k, str)} - strip_caller_authorization: Final = _should_strip_caller_authorization( + strip_caller_authorization: Final = should_strip_caller_authorization( mcp_server=mcp_server, raw_headers=raw_headers, user_api_key_auth=user_api_key_auth, @@ -6216,9 +6248,9 @@ class MCPServerManager: if extra_headers is not None and len(extra_headers) == 0: extra_headers = None - stdio_env: Final = self._build_stdio_env(mcp_server, raw_headers) + stdio_env: Final = self.build_stdio_env(mcp_server, raw_headers) - client: Final = await self._create_mcp_client( + client: Final = await self.create_mcp_client( server=mcp_server, mcp_auth_header=server_auth_header, extra_headers=extra_headers, @@ -6347,7 +6379,9 @@ class MCPServerManager: ) -> MCPServer: """Resolve MCP server for call_tool (prefixed name, registry, fallback).""" prefixed_tool_name: Final = add_server_prefix_to_name(name, server_name) - mcp_server = self._get_mcp_server_from_tool_name(prefixed_tool_name) + mcp_server = self.get_mcp_server_from_tool_name( # rebind-ok: pre-existing rebinding on a rename-only line + prefixed_tool_name + ) resolved_by_server_name_only = False normalized_server_name: Final = normalize_server_name(server_name) @@ -6368,7 +6402,7 @@ class MCPServerManager: resolved_by_server_name_only = True break if mcp_server is None: - fallback: Final = self._get_mcp_server_from_tool_name(name) + fallback: Final = self.get_mcp_server_from_tool_name(name) if fallback is not None and (not server_name or _candidate_matches_server_name(fallback)): mcp_server = fallback if mcp_server is None: @@ -6495,7 +6529,9 @@ class MCPServerManager: subject_token: str | None = None if isinstance(spec.config, (TokenExchangeConfig, IdJagConfig)): - subject_token = self._extract_subject_token(oauth2_headers, raw_headers, user_api_key_auth) + subject_token = self.extract_subject_token( # rebind-ok: pre-existing rebinding on a rename-only line + oauth2_headers, raw_headers, user_api_key_auth + ) elif isinstance(spec.config, PassthroughConfig): inbound_token, forwarded_headers = take_forwarded_authorization(forwarded_headers) per_server_token: Final = passthrough_token_from_mcp_auth_header(mcp_auth_header) @@ -6637,7 +6673,7 @@ class MCPServerManager: server_name, ) - auth_header_value, openapi_forwarded_headers, upstream_credential = _resolve_openapi_tool_auth( + auth_header_value, openapi_forwarded_headers, upstream_credential = resolve_openapi_tool_auth( mcp_server=mcp_server, mcp_auth_header=mcp_auth_header, mcp_server_auth_headers=mcp_server_auth_headers, @@ -6655,21 +6691,21 @@ class MCPServerManager: async def _call_openapi_via_handler(): from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( - _request_auth_header, - _request_extra_headers, - _request_resolved_auth_headers, + request_auth_header, + request_extra_headers, + request_resolved_auth_headers, ) - auth_token: Final = _request_auth_header.set(auth_header_value) - extra_token: Final = _request_extra_headers.set(forwarded_headers) - resolved_token: Final = _request_resolved_auth_headers.set(resolved_auth_headers) + auth_token: Final = request_auth_header.set(auth_header_value) + extra_token: Final = request_extra_headers.set(forwarded_headers) + resolved_token: Final = request_resolved_auth_headers.set(resolved_auth_headers) try: async with self._limit_outbound_concurrency(mcp_server): return await self._call_openapi_tool_handler(mcp_server, name, arguments, wire_compat) finally: - _request_auth_header.reset(auth_token) - _request_extra_headers.reset(extra_token) - _request_resolved_auth_headers.reset(resolved_token) + request_auth_header.reset(auth_token) + request_extra_headers.reset(extra_token) + request_resolved_auth_headers.reset(resolved_token) tasks.append(asyncio.create_task(_call_openapi_via_handler())) else: @@ -6720,7 +6756,7 @@ class MCPServerManager: # Skip OAuth2 servers that rely on user-provided tokens continue try: - tools = await self._get_tools_from_server(server) + tools = await self.get_tools_from_server(server) except MCPUpstreamAuthError as e: # Pass-through servers expect a user-supplied bearer token; # at startup we have none, so an upstream 401 is normal. @@ -6741,7 +6777,7 @@ class MCPServerManager: self.tool_name_to_mcp_server_name_mapping[original_name] = server.name self.tool_name_to_mcp_server_name_mapping[tool.name] = server.name - def _get_mcp_server_from_tool_name(self, tool_name: str) -> MCPServer | None: + def get_mcp_server_from_tool_name(self, tool_name: str) -> MCPServer | None: """ Get the MCP Server from the tool name (handles both prefixed and non-prefixed names) @@ -6786,6 +6822,8 @@ class MCPServerManager: return None + _get_mcp_server_from_tool_name = get_mcp_server_from_tool_name + async def reload_servers_from_database(self): await self.catalog.reload() @@ -6809,7 +6847,7 @@ class MCPServerManager: # Fallback if proxy_server not available return {} - def _is_server_accessible_from_ip(self, server: MCPServer, client_ip: str | None) -> bool: + def is_server_accessible_from_ip(self, server: MCPServer, client_ip: str | None) -> bool: """ Check if a server is accessible from the given client IP. @@ -6830,12 +6868,14 @@ class MCPServerManager: internal_networks = IPAddressUtils.parse_internal_networks(general_settings.get("mcp_internal_ip_ranges")) return IPAddressUtils.is_internal_ip(client_ip, internal_networks) + _is_server_accessible_from_ip = is_server_accessible_from_ip + def get_mcp_server_by_id(self, server_id: str, client_ip: str | None = None) -> MCPServer | None: """Get the MCP Server from the server id.""" registry: Final = self.get_registry() for server in registry.values(): if server.server_id == server_id: - if not self._is_server_accessible_from_ip(server, client_ip): + if not self.is_server_accessible_from_ip(server, client_ip): return None return server return None @@ -6961,19 +7001,19 @@ class MCPServerManager: # Pass 1: Match by alias (highest priority) for server in registry.values(): if server.alias == server_name: - if not self._is_server_accessible_from_ip(server, client_ip): + if not self.is_server_accessible_from_ip(server, client_ip): return None return server # Pass 2: Match by server_name for server in registry.values(): if server.server_name == server_name: - if not self._is_server_accessible_from_ip(server, client_ip): + if not self.is_server_accessible_from_ip(server, client_ip): return None return server # Pass 3: Match by name (lowest priority) for server in registry.values(): if server.name == server_name: - if not self._is_server_accessible_from_ip(server, client_ip): + if not self.is_server_accessible_from_ip(server, client_ip): return None return server return None @@ -6989,7 +7029,7 @@ class MCPServerManager: registry: Final = self.get_registry() if client_ip is None: return registry - return {k: v for k, v in registry.items() if self._is_server_accessible_from_ip(v, client_ip)} + return {k: v for k, v in registry.items() if self.is_server_accessible_from_ip(v, client_ip)} def _generate_stable_server_id( self, @@ -7054,7 +7094,7 @@ class MCPServerManager: if server.spec_path: spec_status, spec_error, spec_checked_at = await self._openapi_health_probes(server.spec_path).check() - return self._build_mcp_server_table(server).model_copy( + return self.build_mcp_server_table(server).model_copy( update=MappingProxyType( { "status": spec_status, @@ -7086,7 +7126,7 @@ class MCPServerManager: raise_on_missing=False, ) extra_headers: Final = dict(resolved_static_headers) if resolved_static_headers else {} - client: Final = await self._create_mcp_client( + client: Final = await self.create_mcp_client( server=server, mcp_auth_header=None, extra_headers=extra_headers, @@ -7207,7 +7247,7 @@ class MCPServerManager: verbose_logger.warning("MCP Server %s not found in registry", server_id) continue - mcp_server_table = self._build_mcp_server_table(server) + mcp_server_table = self.build_mcp_server_table(server) list_mcp_servers.append(mcp_server_table) return list_mcp_servers @@ -7220,7 +7260,7 @@ class MCPServerManager: return None return [MCPEnvVar.model_validate(env_var) for env_var in env_vars] - def _build_mcp_server_table(self, server: MCPServer) -> LiteLLM_MCPServerTable: + def build_mcp_server_table(self, server: MCPServer) -> LiteLLM_MCPServerTable: return LiteLLM_MCPServerTable( server_id=server.server_id, is_config=self.is_config_declared_server(server.server_id) and server.server_id not in self.registry, @@ -7275,6 +7315,8 @@ class MCPServerManager: rpm=server.rpm, ) + _build_mcp_server_table = build_mcp_server_table + async def get_all_mcp_servers_unfiltered(self) -> list[LiteLLM_MCPServerTable]: """Return all MCP servers from registry without applying access controls.""" @@ -7284,7 +7326,7 @@ class MCPServerManager: servers: Final[list[LiteLLM_MCPServerTable]] = [] for server in registry.values(): - servers.append(self._build_mcp_server_table(server)) + servers.append(self.build_mcp_server_table(server)) return servers async def get_all_mcp_servers_with_health_unfiltered( diff --git a/litellm/proxy/_experimental/mcp_server/oauth2_flow_backfill.py b/litellm/proxy/_experimental/mcp_server/oauth2_flow_backfill.py index 150900e7ff2..71281b276fa 100644 --- a/litellm/proxy/_experimental/mcp_server/oauth2_flow_backfill.py +++ b/litellm/proxy/_experimental/mcp_server/oauth2_flow_backfill.py @@ -42,7 +42,11 @@ from typing import Final, Literal, Protocol from pydantic import JsonValue from litellm._logging import verbose_proxy_logger -from litellm.proxy._experimental.mcp_server.db import _decode_oauth_payload, decrypt_credentials +from litellm.proxy._experimental.mcp_server.db import ( # noqa: F401 # legacy module exports + _decode_oauth_payload, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + decode_oauth_payload, + decrypt_credentials, +) from litellm.proxy.utils import PrismaClient from litellm.types.mcp import MCPCredentials @@ -156,7 +160,7 @@ async def backfill_null_oauth2_flows(prisma_client: PrismaClient) -> dict[Backfi where={"server_id": {"in": server_ids}}, ) server_ids_with_oauth_tokens: Final[set[str]] = { - token_row.server_id for token_row in token_rows if _decode_oauth_payload(token_row.credential_b64) is not None + token_row.server_id for token_row in token_rows if decode_oauth_payload(token_row.credential_b64) is not None } classified: Final = tuple( diff --git a/litellm/proxy/_experimental/mcp_server/oauth2_token_cache.py b/litellm/proxy/_experimental/mcp_server/oauth2_token_cache.py index 080665fda8a..6d4fefcdb6c 100644 --- a/litellm/proxy/_experimental/mcp_server/oauth2_token_cache.py +++ b/litellm/proxy/_experimental/mcp_server/oauth2_token_cache.py @@ -205,7 +205,7 @@ class MCPOAuth2TokenCache(InMemoryCache): mcp_oauth2_token_cache: Final = MCPOAuth2TokenCache() -def _compute_per_user_token_ttl(server: "MCPServer", expires_in: int | None) -> int: +def compute_per_user_token_ttl(server: "MCPServer", expires_in: int | None) -> int: """Compute Redis TTL for a per-user token. Uses server.token_storage_ttl_seconds when configured, capped at the token's @@ -223,6 +223,9 @@ def _compute_per_user_token_ttl(server: "MCPServer", expires_in: int | None) -> return MCP_PER_USER_TOKEN_DEFAULT_TTL +_compute_per_user_token_ttl: Final = compute_per_user_token_ttl + + class MCPPerUserTokenCache: """Redis-backed cache for per-user OAuth2 access tokens. diff --git a/litellm/proxy/_experimental/mcp_server/oauth_utils.py b/litellm/proxy/_experimental/mcp_server/oauth_utils.py index ccbf69d0f54..e60c4ef1c30 100644 --- a/litellm/proxy/_experimental/mcp_server/oauth_utils.py +++ b/litellm/proxy/_experimental/mcp_server/oauth_utils.py @@ -84,7 +84,7 @@ def _origin_label(scheme: str, netloc: str) -> str: return f"{scheme}://{netloc}" if netloc else f"{scheme}://" -def _redact_mcp_resource_url(url: str | None) -> str | None: +def redact_mcp_resource_url(url: str | None) -> str | None: """Reduce an MCP server URL to its origin (scheme + host + port) for logging. Everything else is dropped: userinfo (``user:pass@``), the query string, the @@ -107,6 +107,9 @@ def _redact_mcp_resource_url(url: str | None) -> str | None: return urlunsplit((parts.scheme, netloc, "", "", "")) or None +_redact_mcp_resource_url: Final = redact_mcp_resource_url + + def _resolve_proxy_base_url_env() -> str | None: global _warned_invalid_proxy_base_url configured: Final = os.environ.get("PROXY_BASE_URL", "").strip() diff --git a/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py b/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py index 5b23695d06d..8ef573aa489 100644 --- a/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py +++ b/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py @@ -27,7 +27,9 @@ from litellm.proxy._experimental.mcp_server.exceptions import ( # tag-namespaced operationIds like "actions/download-job-logs-for-workflow-run" # which include '/'. Sanitize here so the same regex passes everywhere downstream. _OPENAPI_TOOL_NAME_INVALID_CHARS: Final = re.compile(r"[^a-zA-Z0-9_-]") -_OPENAPI_TOOL_NAME_MAX_LEN: Final = 128 +OPENAPI_TOOL_NAME_MAX_LEN: Final = 128 + +_OPENAPI_TOOL_NAME_MAX_LEN: Final = OPENAPI_TOOL_NAME_MAX_LEN def sanitize_openapi_tool_name(raw_name: str) -> str: @@ -41,7 +43,7 @@ def sanitize_openapi_tool_name(raw_name: str) -> str: if not raw_name: return raw_name sanitized: Final = _OPENAPI_TOOL_NAME_INVALID_CHARS.sub("_", raw_name).lower() - return sanitized[:_OPENAPI_TOOL_NAME_MAX_LEN] + return sanitized[:OPENAPI_TOOL_NAME_MAX_LEN] from litellm._logging import verbose_logger @@ -106,23 +108,31 @@ HEADERS: Final[dict[str, str]] = {} # Per-request auth header override for BYOK servers. # Set this ContextVar before calling a local tool handler to inject the user's # stored credential into the HTTP request made by the tool function closure. -_request_auth_header: contextvars.ContextVar[str | None] = contextvars.ContextVar("_request_auth_header", default=None) +request_auth_header: Final[contextvars.ContextVar[str | None]] = contextvars.ContextVar( + "_request_auth_header", default=None +) + +_request_auth_header: Final = request_auth_header # Per-request extra headers forwarded from the client request. # Populated from MCPServer.extra_headers names matched against raw request # headers in server.py before dispatching to a local/OpenAPI tool handler. -_request_extra_headers: Final[contextvars.ContextVar[dict[str, str] | None]] = contextvars.ContextVar( +request_extra_headers: Final[contextvars.ContextVar[dict[str, str] | None]] = contextvars.ContextVar( "_request_extra_headers", default=None ) +_request_extra_headers: Final = request_extra_headers + # Per-request headers carrying the gateway-resolved upstream credential # (stored per-user OAuth token, minted M2M token, exchanged OBO token). # Set from MCPServerManager.resolve_openapi_upstream_auth; authoritative # over every other Authorization source in _merge_openapi_tool_request_headers. -_request_resolved_auth_headers: Final[contextvars.ContextVar[dict[str, str] | None]] = contextvars.ContextVar( +request_resolved_auth_headers: Final[contextvars.ContextVar[dict[str, str] | None]] = contextvars.ContextVar( "_request_resolved_auth_headers", default=None ) +_request_resolved_auth_headers: Final = request_resolved_auth_headers + _request_upstream_url: Final[contextvars.ContextVar[str | None]] = contextvars.ContextVar( "_request_upstream_url", default=None ) @@ -369,7 +379,7 @@ async def _drop_credential_across_origin(request: httpx.Request) -> None: built would never be closed. """ guard: Final = credential_redirect_hook( - _request_upstream_url.get() or "", custom_credential_slot(_request_resolved_auth_headers.get()) + _request_upstream_url.get() or "", custom_credential_slot(request_resolved_auth_headers.get()) ) if guard is not None: await guard(request) @@ -382,7 +392,7 @@ def _upstream_client() -> AsyncHTTPHandler: itself, so this arm installs the same hook the MCP client uses. Both variants come from the shared cache, so a guarded call reuses its connection pool like any other. """ - if custom_credential_slot(_request_resolved_auth_headers.get()) is None: + if custom_credential_slot(request_resolved_auth_headers.get()) is None: return get_async_httpx_client(llm_provider=httpxSpecialProvider.MCP) return get_async_httpx_client( llm_provider=httpxSpecialProvider.MCP, @@ -417,20 +427,20 @@ def _merge_openapi_tool_request_headers( Header names are compared case-insensitively so different casing cannot bypass the precedence rules. """ - request_extra: Final = _request_extra_headers.get() or {} + request_extra: Final = request_extra_headers.get() or {} static: Final = static_headers or {} static_lower_names: Final = {k.lower() for k in static} effective_headers: dict[str, str] = {k: v for k, v in request_extra.items() if k.lower() not in static_lower_names} effective_headers.update(static) - override_auth: Final = _request_auth_header.get() + override_auth: Final = request_auth_header.get() if override_auth: for existing in [k for k in effective_headers if k.lower() == "authorization"]: del effective_headers[existing] effective_headers["Authorization"] = override_auth - resolved_auth_headers: Final = _request_resolved_auth_headers.get() or {} + resolved_auth_headers: Final = request_resolved_auth_headers.get() or {} for name, value in resolved_auth_headers.items(): for existing in [k for k in effective_headers if k.lower() == name.lower()]: del effective_headers[existing] @@ -622,7 +632,7 @@ def register_tools_from_openapi(spec: Mapping[str, Any], base_url: str) -> None: while unique in used_names: n += 1 suffix = f"_{n}" - unique = tool_name[: _OPENAPI_TOOL_NAME_MAX_LEN - len(suffix)] + suffix + unique = tool_name[: OPENAPI_TOOL_NAME_MAX_LEN - len(suffix)] + suffix tool_name = unique used_names.add(tool_name) diff --git a/litellm/proxy/_experimental/mcp_server/operations.py b/litellm/proxy/_experimental/mcp_server/operations.py index e226b7f3fcb..6ec5bdf6af7 100644 --- a/litellm/proxy/_experimental/mcp_server/operations.py +++ b/litellm/proxy/_experimental/mcp_server/operations.py @@ -83,23 +83,31 @@ from litellm.proxy._experimental.mcp_server.faults.list_outcomes import ( classify_list_exception, outcome_wire_value, ) -from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( +from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: F401 # legacy module exports MCPServerManager, - _caller_authorization_fans_out, - _client_forwarded_authorization_headers, - _resolve_openapi_tool_auth, - _should_strip_caller_authorization, + _caller_authorization_fans_out, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _client_forwarded_authorization_headers, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _resolve_openapi_tool_auth, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _should_strip_caller_authorization, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + caller_authorization_fans_out, + client_forwarded_authorization_headers, global_mcp_server_manager, listed_tools_caller_for, + resolve_openapi_tool_auth, + should_strip_caller_authorization, ) -from litellm.proxy._experimental.mcp_server.oauth_utils import ( - _redact_mcp_resource_url, +from litellm.proxy._experimental.mcp_server.oauth_utils import ( # noqa: F401 # legacy module exports + _redact_mcp_resource_url, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export get_byok_www_authenticate, + redact_mcp_resource_url, ) -from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( - _request_auth_header, - _request_extra_headers, - _request_resolved_auth_headers, +from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( # noqa: F401 # legacy module exports + _request_auth_header, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _request_extra_headers, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _request_resolved_auth_headers, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + request_auth_header, + request_extra_headers, + request_resolved_auth_headers, ) from litellm.proxy._experimental.mcp_server.result_conversion import ( WireCompat, @@ -523,7 +531,7 @@ async def _get_allowed_mcp_servers_from_mcp_server_names( if not server_name_matched: try: - access_group_server_ids = await MCPRequestHandler._get_mcp_servers_from_access_groups( + access_group_server_ids = await MCPRequestHandler.get_mcp_servers_from_access_groups( [server_or_group] ) # Only include servers that the user has access to @@ -882,7 +890,7 @@ def _prepare_mcp_server_headers( # x-mcp-{alias}-authorization. The decision is computed once so BOTH the forwarding branch and # the extra_headers copy loop below honor it — otherwise a server that lists Authorization in # extra_headers would re-copy the withheld bearer from raw_headers and replay it anyway. - withhold_forwarded_authorization: Final = is_client_forwarded_mode and _caller_authorization_fans_out( + withhold_forwarded_authorization: Final = is_client_forwarded_mode and caller_authorization_fans_out( server, scope_servers ) if server.auth_type == MCPAuth.oauth2: @@ -897,7 +905,7 @@ def _prepare_mcp_server_headers( # token, so drop the caller-forwarded Authorization (apply-if-absent would # otherwise let it shadow the resolved token). Delegate keeps it. Centralized # via _should_strip_caller_authorization to match _call_regular_mcp_tool. - if extra_headers and _should_strip_caller_authorization( + if extra_headers and should_strip_caller_authorization( mcp_server=server, raw_headers=raw_headers, user_api_key_auth=user_api_key_auth, @@ -905,11 +913,13 @@ def _prepare_mcp_server_headers( extra_headers = without_header(extra_headers, DEFAULT_CREDENTIAL_HEADER) elif is_client_forwarded_mode: if not withhold_forwarded_authorization: - extra_headers = _client_forwarded_authorization_headers( - mcp_server=server, - oauth2_headers=oauth2_headers, - raw_headers=raw_headers, - user_api_key_auth=user_api_key_auth, + extra_headers = ( # rebind-ok: pre-existing rebinding on a rename-only line + client_forwarded_authorization_headers( + mcp_server=server, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + user_api_key_auth=user_api_key_auth, + ) ) if server.extra_headers and raw_headers: @@ -922,7 +932,7 @@ def _prepare_mcp_server_headers( # ``MCPServerManager._call_regular_mcp_tool`` so the two # code paths cannot drift on this security-sensitive choice. # See ``_should_strip_caller_authorization`` for the rules. - strip_caller_authorization: Final = _should_strip_caller_authorization( + strip_caller_authorization: Final = should_strip_caller_authorization( mcp_server=server, raw_headers=raw_headers, user_api_key_auth=user_api_key_auth, @@ -1787,7 +1797,7 @@ def _challenge_missing_token_exchange_subject( return if all(allowed.server_id != server.server_id for allowed in allowed_mcp_servers): return - if global_mcp_server_manager._extract_subject_token(oauth2_headers, raw_headers, user_api_key_auth) is not None: + if global_mcp_server_manager.extract_subject_token(oauth2_headers, raw_headers, user_api_key_auth) is not None: return from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import ( # noqa: PLC0415 # lazy: adapter pulls MCP subgraph raise_token_exchange_challenge, @@ -1978,10 +1988,12 @@ async def _execute_mcp_tool( original_tool_name = name else: # Resolve from tool name (MCP JSON-RPC or prefixed REST tool names). - mcp_server = global_mcp_server_manager._get_mcp_server_from_tool_name(name) + mcp_server = global_mcp_server_manager.get_mcp_server_from_tool_name( # rebind-ok: pre-existing rebinding on a rename-only line + name + ) if mcp_server is None and requested_server is not None: for known_prefix in iter_known_server_prefixes(requested_server): - candidate = global_mcp_server_manager._get_mcp_server_from_tool_name( + candidate = global_mcp_server_manager.get_mcp_server_from_tool_name( add_server_prefix_to_name(name, known_prefix) ) if candidate is not None: @@ -2034,7 +2046,9 @@ async def _execute_mcp_tool( # Resolve the MCP server early so BYOK checks and credential injection # apply to ALL dispatch paths (local tool registry AND managed MCP server). if mcp_server is None: - mcp_server = global_mcp_server_manager._get_mcp_server_from_tool_name(name) + mcp_server = global_mcp_server_manager.get_mcp_server_from_tool_name( # rebind-ok: pre-existing rebinding on a rename-only line + name + ) client_auth_header: Final = mcp_auth_header if mcp_server: @@ -2131,7 +2145,7 @@ async def _execute_mcp_tool( verbose_logger.debug("Executing local registry tool: %s", name) # The credential rides ContextVars because the tool function has its # headers baked into the closure at registration time. - auth_header_value, openapi_forwarded_headers, upstream_credential = _resolve_openapi_tool_auth( + auth_header_value, openapi_forwarded_headers, upstream_credential = resolve_openapi_tool_auth( mcp_server=mcp_server, mcp_auth_header=mcp_auth_header, mcp_server_auth_headers=mcp_server_auth_headers, @@ -2150,15 +2164,15 @@ async def _execute_mcp_tool( forwarded_headers=openapi_forwarded_headers, ) - _auth_token: Final = _request_auth_header.set(auth_header_value) - _extra_token: Final = _request_extra_headers.set(forwarded_headers) - _resolved_token: Final = _request_resolved_auth_headers.set(resolved_auth_headers) + _auth_token: Final = request_auth_header.set(auth_header_value) + _extra_token: Final = request_extra_headers.set(forwarded_headers) + _resolved_token: Final = request_resolved_auth_headers.set(resolved_auth_headers) try: response = await _handle_local_mcp_tool(name, arguments, wire_compat) finally: - _request_auth_header.reset(_auth_token) - _request_extra_headers.reset(_extra_token) - _request_resolved_auth_headers.reset(_resolved_token) + request_auth_header.reset(_auth_token) + request_extra_headers.reset(_extra_token) + request_resolved_auth_headers.reset(_resolved_token) # Try managed MCP server tool (the name is bare; the prefix boundary was # already resolved above against this server's registered prefixes) @@ -2618,7 +2632,7 @@ def _get_standard_logging_mcp_tool_call( server_name: str | None, session_id: str | None = None, ) -> StandardLoggingMCPToolCall: - mcp_server: Final = global_mcp_server_manager._get_mcp_server_from_tool_name( + mcp_server: Final = global_mcp_server_manager.get_mcp_server_from_tool_name( add_server_prefix_to_name(name, server_name) if server_name else name ) namespaced_tool_name: Final = f"{server_name}/{name}" if server_name else name @@ -2632,7 +2646,7 @@ def _get_standard_logging_mcp_tool_call( namespaced_tool_name=namespaced_tool_name, mcp_session_id=session_id, mcp_auth_mode=mcp_server.auth_type, - mcp_server_resource=_redact_mcp_resource_url(mcp_server.url), + mcp_server_resource=redact_mcp_resource_url(mcp_server.url), ) else: return StandardLoggingMCPToolCall( diff --git a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py index 4a1ae19b2ca..ac543e1f041 100644 --- a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py @@ -37,7 +37,10 @@ from litellm.proxy._experimental.mcp_server.faults.list_outcomes import ( outcome_wire_value, ) from litellm.proxy._experimental.mcp_server.faults.traversal import iter_exception_tree -from litellm.proxy._experimental.mcp_server.oauth_utils import _redact_mcp_resource_url +from litellm.proxy._experimental.mcp_server.oauth_utils import ( # noqa: F401 # legacy module exports + _redact_mcp_resource_url, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + redact_mcp_resource_url, +) from litellm.proxy._experimental.mcp_server.result_conversion import WireCompat, complete_call_tool_result from litellm.proxy._experimental.mcp_server.ui_session_utils import ( acting_user_auth, @@ -64,7 +67,10 @@ if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.proxy._experimental.mcp_server.db import OAuthCredentialPayload from litellm.proxy.utils import ProxyLogging -from litellm.proxy.common_utils.http_parsing_utils import _safe_get_request_headers +from litellm.proxy.common_utils.http_parsing_utils import ( # noqa: F401 # legacy module exports + _safe_get_request_headers, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + safe_get_request_headers, +) from litellm.types.mcp import MCPAuth from litellm.types.utils import CallTypes, StandardLoggingMCPToolCall @@ -144,7 +150,7 @@ def _known_connection_error_message(exc: BaseException, url: str | None, timeout if isinstance(exc, TimeoutError): return ( "Failed to connect to MCP server: no valid MCP response received from " - f"{_redact_mcp_resource_url(url) or 'the server'} " + f"{redact_mcp_resource_url(url) or 'the server'} " f"within {timeout_seconds:.0f}s. Check that the LiteLLM proxy can reach this URL " "from its network (DNS, egress rules, firewalls) and that the server answers MCP requests." ) @@ -223,8 +229,9 @@ if MCP_AVAILABLE: from litellm.llms.litellm_proxy.skills.skill_search import ( DEFAULT_SKILL_SEARCH_TOP_K, ) - from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( - _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES, + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: F401 # legacy module exports + _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES, ListedToolsCaller, global_mcp_server_manager, ) @@ -242,8 +249,9 @@ if MCP_AVAILABLE: filter_tools_by_key_team_permissions, fire_mcp_tool_call_failure_logging, ) - from litellm.proxy._experimental.mcp_server.server import ( - _apply_toolset_scope, + from litellm.proxy._experimental.mcp_server.server import ( # noqa: F401 # legacy module exports + _apply_toolset_scope, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + apply_toolset_scope, reject_disallowed_mcp_client, ) from litellm.proxy._experimental.mcp_server.tool_catalog_guard import ( @@ -369,7 +377,7 @@ if MCP_AVAILABLE: virtual_mcp_server_auth_headers, virtual_raw_headers, ) = _extract_mcp_headers_from_request(request, MCPRequestHandler) - virtual_oauth2_headers: Final = MCPRequestHandler._get_oauth2_headers_from_headers(request.headers) + virtual_oauth2_headers: Final = MCPRequestHandler.get_oauth2_headers_from_headers(request.headers) if tool_name == MCP_TOOL_SEARCH_TOOL_NAME: return await handle_mcp_tool_search( query=tool_arguments.get("query", ""), @@ -656,7 +664,7 @@ if MCP_AVAILABLE: if ( _server is not None and _rest_client_ip is not None - and not global_mcp_server_manager._is_server_accessible_from_ip(_server, _rest_client_ip) + and not global_mcp_server_manager.is_server_accessible_from_ip(_server, _rest_client_ip) ): raise HTTPException( status_code=403, @@ -707,7 +715,7 @@ if MCP_AVAILABLE: record_listing: bool, ) -> list[MCPTool]: return list( - await global_mcp_server_manager._get_tools_from_server( + await global_mcp_server_manager.get_tools_from_server( server=server, mcp_auth_header=server_auth_header, extra_headers=extra_headers, @@ -856,7 +864,7 @@ if MCP_AVAILABLE: if ( _server is not None and rest_client_ip is not None - and not global_mcp_server_manager._is_server_accessible_from_ip(_server, rest_client_ip) + and not global_mcp_server_manager.is_server_accessible_from_ip(_server, rest_client_ip) ): raise HTTPException( status_code=403, @@ -953,7 +961,7 @@ if MCP_AVAILABLE: status_code=404, detail=f"Toolset '{toolset_name}' not found", ) - return await _apply_toolset_scope(user_api_key_dict, toolset.toolset_id) + return await apply_toolset_scope(user_api_key_dict, toolset.toolset_id) @router.get("/tools/list", dependencies=[Depends(user_api_key_auth)]) @catalog_operation(global_manager) @@ -1034,8 +1042,8 @@ if MCP_AVAILABLE: # Extract auth headers from request headers: Final = request.headers raw_headers_from_request: Final = dict(headers) - mcp_auth_header: Final = MCPRequestHandler._get_mcp_auth_header_from_headers(headers) - mcp_server_auth_headers: Final = MCPRequestHandler._get_mcp_server_auth_headers_from_headers(headers) + mcp_auth_header: Final = MCPRequestHandler.get_mcp_auth_header_from_headers(headers) + mcp_server_auth_headers: Final = MCPRequestHandler.get_mcp_server_auth_headers_from_headers(headers) auth_contexts: Final = await build_effective_auth_contexts(user_api_key_dict) @@ -1293,7 +1301,7 @@ if MCP_AVAILABLE: if target_server is not None: user_oauth_extra_headers = await _get_user_oauth_extra_headers(target_server, user_api_key_dict) caller_oauth2_headers: Final = ( - MCPRequestHandler._get_oauth2_headers_from_headers(request.headers) + MCPRequestHandler.get_oauth2_headers_from_headers(request.headers) if target_server is not None and target_server.auth_type in _CLIENT_FORWARDED_TOKEN_AUTH_TYPES else None ) @@ -1396,9 +1404,10 @@ if MCP_AVAILABLE: # /health/tools/list -> List tools from MCP server # For these routes users will dynamically pass the MCP connection params, they don't need to be on the MCP registry ######################################################## - from litellm.proxy.management_endpoints.mcp_management_endpoints import ( + from litellm.proxy.management_endpoints.mcp_management_endpoints import ( # noqa: F401 # legacy module exports NewMCPServerRequest, - _inherit_credentials_from_existing_server, + _inherit_credentials_from_existing_server, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + inherit_credentials_from_existing_server, ) def _extract_credentials( @@ -1461,7 +1470,7 @@ if MCP_AVAILABLE: saved_origin is not None and saved_origin == preview_origin ) request: Final = ( - _inherit_credentials_from_existing_server(new_mcp_server_request) if may_inherit else new_mcp_server_request + inherit_credentials_from_existing_server(new_mcp_server_request) if may_inherit else new_mcp_server_request ) mcp_auth_header: Final = ( request.credentials.get("auth_value") @@ -1472,8 +1481,8 @@ if MCP_AVAILABLE: # when the primary x-litellm-api-key header is absent, the Authorization value is the # caller's LiteLLM key, not an upstream token, and must never be forwarded upstream. oauth2_headers: Final = ( - MCPRequestHandler._get_oauth2_headers_from_headers(headers) - if request.auth_type in _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES + MCPRequestHandler.get_oauth2_headers_from_headers(headers) + if request.auth_type in UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES and headers.get(MCPRequestHandler.LITELLM_API_KEY_HEADER_NAME_PRIMARY) else None ) @@ -1564,7 +1573,7 @@ if MCP_AVAILABLE: instructions=request.instructions, ) - stdio_env: Final = global_mcp_server_manager._build_stdio_env(server_model, raw_headers) + stdio_env: Final = global_mcp_server_manager.build_stdio_env(server_model, raw_headers) # For M2M OAuth servers, drop the incoming Authorization header so that # resolve_mcp_auth can auto-fetch a token via client_credentials. @@ -1617,7 +1626,7 @@ if MCP_AVAILABLE: ) with anyio.fail_after(timeout_seconds): - client: Final = await global_mcp_server_manager._create_mcp_client( + client: Final = await global_mcp_server_manager.create_mcp_client( server=server_model, mcp_auth_header=mcp_auth_header, extra_headers=merged_headers, @@ -1648,7 +1657,7 @@ if MCP_AVAILABLE: async def _preview_openapi_tools(spec_path: str) -> dict: """Generate tool previews from an OpenAPI spec without creating a server.""" from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( - _OPENAPI_TOOL_NAME_MAX_LEN, + OPENAPI_TOOL_NAME_MAX_LEN, build_input_schema, load_openapi_spec_async, resolve_operation_params, @@ -1681,7 +1690,7 @@ if MCP_AVAILABLE: while unique in used_names: n += 1 suffix = f"_{n}" - unique = op_id[: _OPENAPI_TOOL_NAME_MAX_LEN - len(suffix)] + suffix + unique = op_id[: OPENAPI_TOOL_NAME_MAX_LEN - len(suffix)] + suffix op_id = unique used_names.add(op_id) summary = operation.get("summary", "") @@ -1738,7 +1747,7 @@ if MCP_AVAILABLE: _test_connection_operation, mcp_auth_header=staged.mcp_auth_header, oauth2_headers=staged.oauth2_headers, - raw_headers=_safe_get_request_headers(request), + raw_headers=safe_get_request_headers(request), ) @router.post("/test/tools/list", dependencies=[Depends(user_api_key_auth)]) @@ -1797,5 +1806,5 @@ if MCP_AVAILABLE: _list_tools_operation, mcp_auth_header=staged.mcp_auth_header, oauth2_headers=staged.oauth2_headers, - raw_headers=_safe_get_request_headers(request), + raw_headers=safe_get_request_headers(request), ) diff --git a/litellm/proxy/_experimental/mcp_server/sampling_handler.py b/litellm/proxy/_experimental/mcp_server/sampling_handler.py index 361d8d5ae31..e1469f2fff7 100644 --- a/litellm/proxy/_experimental/mcp_server/sampling_handler.py +++ b/litellm/proxy/_experimental/mcp_server/sampling_handler.py @@ -772,11 +772,11 @@ async def _check_model_access(model: str, user_api_key_auth: "UserAPIKeyAuth | N import litellm from litellm.proxy._types import ModelAccessDeniedProxyException from litellm.proxy.auth.auth_checks import ( - _check_team_member_model_access, can_key_call_model, can_project_access_model, can_team_access_model, can_user_call_model, + check_team_member_model_access, get_project_object, get_team_object, get_user_object, @@ -832,7 +832,7 @@ async def _check_model_access(model: str, user_api_key_auth: "UserAPIKeyAuth | N team_model_aliases=getattr(user_api_key_auth, "team_model_aliases", None), ) if _user_id and _proxy_logging_obj: - await _check_team_member_model_access( + await check_team_member_model_access( model=model, team_object=team_obj, valid_token=user_api_key_auth, diff --git a/litellm/proxy/_experimental/mcp_server/semantic_tool_filter.py b/litellm/proxy/_experimental/mcp_server/semantic_tool_filter.py index 0644238071e..2f88d8c361f 100644 --- a/litellm/proxy/_experimental/mcp_server/semantic_tool_filter.py +++ b/litellm/proxy/_experimental/mcp_server/semantic_tool_filter.py @@ -117,7 +117,7 @@ class SemanticMCPToolFilter: self.tool_router = None raise - def _extract_tool_info(self, tool) -> tuple[str, str]: + def extract_tool_info(self, tool) -> tuple[str, str]: """Extract name and description from MCP tool or OpenAI function dict.""" name: str description: str @@ -133,6 +133,8 @@ class SemanticMCPToolFilter: return name, description + _extract_tool_info = extract_tool_info + def _build_router(self, tools: list) -> None: """Build semantic router with tools (MCPTool objects or OpenAI function dicts).""" from semantic_router.routers import SemanticRouter @@ -153,7 +155,7 @@ class SemanticMCPToolFilter: self._tool_map = {} for tool in tools: - name, description = self._extract_tool_info(tool) + name, description = self.extract_tool_info(tool) self._tool_map[name] = tool routes.append( @@ -187,13 +189,13 @@ class SemanticMCPToolFilter: def _has_tools_missing_from_index(self, tools: Sequence[object]) -> 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)) + 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: Sequence[object]) -> Mapping[str, object]: """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) + for name, tool in ((self.extract_tool_info(t)[0], t) for t in tools) if name and name not in self._tool_map } @@ -228,7 +230,7 @@ class SemanticMCPToolFilter: if not missing: return - descriptions: Final = {name: self._extract_tool_info(tool)[1] for name, tool in missing.items()} + descriptions: Final = {name: self.extract_tool_info(tool)[1] for name, tool in missing.items()} routes: Final = [ Route( name=name, @@ -302,7 +304,7 @@ class SemanticMCPToolFilter: verbose_logger.warning("Semantic router could not be built from the request's tools") return available_tools - available_names: Final = [name for name in (self._extract_tool_info(t)[0] for t in available_tools) if name] + available_names: Final = [name for name in (self.extract_tool_info(t)[0] for t in available_tools) if name] if not available_names: return available_tools @@ -406,7 +408,7 @@ class SemanticMCPToolFilter: # names happen to be tail-compatible with the same incoming name. available_by_name: Final[dict[str, object]] = {} for tool in available_tools: - client_name, _ = self._extract_tool_info(tool) + client_name, _ = self.extract_tool_info(tool) if client_name and client_name not in available_by_name: available_by_name[client_name] = tool diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 4021c7ddb47..0c0300fa0e5 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -32,9 +32,10 @@ from litellm.llms.custom_httpx.http_handler import ( get_async_httpx_client, httpxSpecialProvider, ) -from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( +from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( # noqa: F401 # legacy module exports MCPRequestHandler, - _is_mcp_admitted_user_subject, + _is_mcp_admitted_user_subject, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + is_mcp_admitted_user_subject, ) from litellm.proxy._experimental.mcp_server.client_allowlist import ( MCPClientAllowlist, @@ -47,13 +48,16 @@ from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( from litellm.proxy._experimental.mcp_server.exceptions import ( MCPUpstreamAuthError, ) -from litellm.proxy._experimental.mcp_server.mcp_context import ( - _mcp_active_toolset_id, - _mcp_gateway_initialize_instructions, - _mcp_gateway_server_name, +from litellm.proxy._experimental.mcp_server.mcp_context import ( # noqa: F401 # legacy module exports + _mcp_active_toolset_id, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _mcp_gateway_initialize_instructions, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _mcp_gateway_server_name, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export _mcp_proxy_mode, # pyright: ignore[reportPrivateUsage] # server-owned request mode active_mcp_request_ctx_var, get_active_mcp_request_ctx, + mcp_active_toolset_id, + mcp_gateway_initialize_instructions, + mcp_gateway_server_name, ) from litellm.proxy._experimental.mcp_server.mcp_debug import ( MCP_AUTH_DIAGNOSTICS_SCOPE_KEY, @@ -61,9 +65,9 @@ from litellm.proxy._experimental.mcp_server.mcp_debug import ( MCPDebug, ) from litellm.proxy._experimental.mcp_server.oauth_utils import ( - _redact_mcp_resource_url, get_passthrough_www_authenticate, get_route_relative_request_path, + redact_mcp_resource_url, well_known_root_suffix, ) from litellm.proxy._experimental.mcp_server.ui_session_utils import ( @@ -98,6 +102,8 @@ if TYPE_CHECKING: from mcp.server.session import ServerSession as _McpServerSession +_redact_mcp_resource_url: Final = redact_mcp_resource_url + _STATEFUL_SESSION_IDLE_TIMEOUT_SECONDS: Final = 30 * 60 # Upper bound on concurrent stateful sessions a single caller may hold. Each # `initialize` creates a session that survives until the idle timeout, so @@ -486,6 +492,7 @@ if MCP_AVAILABLE: "mcp_get_prompt", "mcp_read_resource", "raise_denied_scoped_mcp_access", + "redact_mcp_resource_url", ) from mcp.server import Server @@ -579,10 +586,10 @@ if MCP_AVAILABLE: else base_options ) updates: Final[dict[str, str]] = {} - merged: Final = _mcp_gateway_initialize_instructions.get() + merged: Final = mcp_gateway_initialize_instructions.get() if merged is not None: updates["instructions"] = merged - scoped_server_name: Final = _mcp_gateway_server_name.get() + scoped_server_name: Final = mcp_gateway_server_name.get() if scoped_server_name is not None: updates["server_name"] = scoped_server_name return opts.model_copy(update=updates) if updates else opts @@ -1028,7 +1035,7 @@ if MCP_AVAILABLE: # cancel sibling probes or 500 the gateway initialize request. await asyncio.gather( *[ - operations.global_mcp_server_manager._ensure_upstream_initialize_instructions_cached(s) + operations.global_mcp_server_manager.ensure_upstream_initialize_instructions_cached(s) for s in allowed if s is not None ], @@ -1041,13 +1048,13 @@ if MCP_AVAILABLE: scoped_server_name = ( scoped_server.alias or scoped_server.server_name or scoped_server.name or scoped_server.server_id ) - instructions_token: Final = _mcp_gateway_initialize_instructions.set(merged) - server_name_token: Final = _mcp_gateway_server_name.set(scoped_server_name) + instructions_token: Final = mcp_gateway_initialize_instructions.set(merged) + server_name_token: Final = mcp_gateway_server_name.set(scoped_server_name) try: yield finally: - _mcp_gateway_initialize_instructions.reset(instructions_token) - _mcp_gateway_server_name.reset(server_name_token) + mcp_gateway_initialize_instructions.reset(instructions_token) + mcp_gateway_server_name.reset(server_name_token) from litellm.proxy._experimental.mcp_server.operations import ( _MCP_CREDENTIAL_REQUEST_FIELDS, @@ -1524,7 +1531,7 @@ if MCP_AVAILABLE: scope["headers"] = [(k, v) for k, v in _headers if _normalize_header_name(k) != _mcp_session_header] return False - async def _apply_toolset_scope( + async def apply_toolset_scope( user_api_key_auth: UserAPIKeyAuth, toolset_id: str, acting_user: ActingUser = acting_user_auth, @@ -1544,7 +1551,7 @@ if MCP_AVAILABLE: of its grant sources. Admins always pass. """ from litellm.proxy._types import LiteLLM_ObjectPermissionTable - from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view + from litellm.proxy.management_endpoints.common_utils import user_api_key_has_admin_view # A key scoped to no MCP servers opts out of every MCP path. Enforce it # here too, since toolset scoping replaces mcp_servers and would otherwise @@ -1558,13 +1565,13 @@ if MCP_AVAILABLE: ) acting: Final = await acting_user(user_api_key_auth) - is_admin: Final = _user_has_admin_view(acting) + is_admin: Final = user_api_key_has_admin_view(acting) if not is_admin and toolset_id not in await granted(acting): raise HTTPException( status_code=403, detail=f"API key does not have access to toolset '{toolset_id}'.", ) - if _is_mcp_admitted_user_subject(acting): + if is_mcp_admitted_user_subject(acting): resource_server_id: Final = acting.mcp_session_resource_server_id if resource_server_id is not None and resource_server_id not in ( await operations.global_mcp_server_manager.resolve_toolset_tool_permissions( @@ -1600,6 +1607,8 @@ if MCP_AVAILABLE: ) return acting.model_copy(update={"object_permission": updated_op, "mcp_toolset_id": toolset_id}) + _apply_toolset_scope: Final = apply_toolset_scope + async def _toolset_server_ids(toolset_id: str) -> set[str]: return set( await operations.global_mcp_server_manager.resolve_toolset_tool_permissions(toolset_ids=[toolset_id]) @@ -1674,7 +1683,7 @@ if MCP_AVAILABLE: if await operations.global_mcp_server_manager.has_user_oauth_token(server, user_api_key_auth): continue - if _is_mcp_admitted_user_subject(user_api_key_auth): + if is_mcp_admitted_user_subject(user_api_key_auth): raise HTTPException( status_code=401, detail="Unauthorized", @@ -2033,10 +2042,12 @@ if MCP_AVAILABLE: # Apply toolset scope if set server-side via ContextVar (set by # /toolset/{name}/mcp and /{name}/mcp route handlers in proxy_server.py). - active_toolset_id: Final = _mcp_active_toolset_id.get() + active_toolset_id: Final = mcp_active_toolset_id.get() toolset_allowed_server_ids: set[str] | None = None if active_toolset_id and user_api_key_auth is not None: - user_api_key_auth = await _apply_toolset_scope(user_api_key_auth, active_toolset_id) + user_api_key_auth = ( # rebind-ok: pre-existing rebinding on a rename-only line + await apply_toolset_scope(user_api_key_auth, active_toolset_id) + ) toolset_allowed_server_ids = await _toolset_server_ids(active_toolset_id) # https://datatracker.ietf.org/doc/html/rfc9728#name-www-authenticate-response @@ -2380,10 +2391,12 @@ if MCP_AVAILABLE: # Apply toolset scope if set server-side via ContextVar so the # downstream probe list matches the fully-authorized server set # (mirrors the streamable HTTP handler). - active_toolset_id: Final = _mcp_active_toolset_id.get() + active_toolset_id: Final = mcp_active_toolset_id.get() toolset_allowed_server_ids: set[str] | None = None if active_toolset_id and user_api_key_auth is not None: - user_api_key_auth = await _apply_toolset_scope(user_api_key_auth, active_toolset_id) + user_api_key_auth = ( # rebind-ok: pre-existing rebinding on a rename-only line + await apply_toolset_scope(user_api_key_auth, active_toolset_id) + ) toolset_allowed_server_ids = await _toolset_server_ids(active_toolset_id) # https://datatracker.ietf.org/doc/html/rfc9728#name-www-authenticate-response diff --git a/litellm/proxy/_experimental/mcp_server/server_resolution.py b/litellm/proxy/_experimental/mcp_server/server_resolution.py index e828d8d8ee6..31bcd897e72 100644 --- a/litellm/proxy/_experimental/mcp_server/server_resolution.py +++ b/litellm/proxy/_experimental/mcp_server/server_resolution.py @@ -19,9 +19,13 @@ class MCPServerRegistry(Protocol): def get_mcp_server_by_name(self, server_name: str, client_ip: str | None = None) -> MCPServer | None: ... - def _is_server_accessible_from_ip(self, server: MCPServer, client_ip: str | None) -> bool: ... + def is_server_accessible_from_ip(self, server: MCPServer, client_ip: str | None) -> bool: ... - def _build_mcp_server_table(self, server: MCPServer) -> LiteLLM_MCPServerTable: ... + _is_server_accessible_from_ip = is_server_accessible_from_ip + + def build_mcp_server_table(self, server: MCPServer) -> LiteLLM_MCPServerTable: ... + + _build_mcp_server_table = build_mcp_server_table async def get_allowed_mcp_servers(self, user_api_key_auth: UserAPIKeyAuth) -> list[str]: ... @@ -50,7 +54,7 @@ async def resolve_mcp_server( temporary_server: Final[MCPServer | None] = await temp_lookup(server_id) if temporary_server is not None: return ResolvedMCPServer( - table=manager._build_mcp_server_table(temporary_server), + table=manager.build_mcp_server_table(temporary_server), runtime=temporary_server, source="temp", ) @@ -64,12 +68,12 @@ async def resolve_mcp_server( registry_server: Final[MCPServer | None] = ( registry_candidate if registry_candidate is not None - and (id_client_ip is None or manager._is_server_accessible_from_ip(registry_candidate, id_client_ip)) + and (id_client_ip is None or manager.is_server_accessible_from_ip(registry_candidate, id_client_ip)) else None ) if registry_server is not None: return ResolvedMCPServer( - table=manager._build_mcp_server_table(registry_server), + table=manager.build_mcp_server_table(registry_server), runtime=registry_server, source="registry", ) @@ -78,7 +82,7 @@ async def resolve_mcp_server( named_server: Final[MCPServer | None] = manager.get_mcp_server_by_name(server_id, client_ip=name_client_ip) if named_server is not None: return ResolvedMCPServer( - table=manager._build_mcp_server_table(named_server), + table=manager.build_mcp_server_table(named_server), runtime=named_server, source="registry", ) diff --git a/litellm/proxy/_experimental/mcp_server/ui_session_utils.py b/litellm/proxy/_experimental/mcp_server/ui_session_utils.py index b9e25259868..9f60f0afbb6 100644 --- a/litellm/proxy/_experimental/mcp_server/ui_session_utils.py +++ b/litellm/proxy/_experimental/mcp_server/ui_session_utils.py @@ -125,9 +125,9 @@ async def acting_user_auth(user_api_key_auth: UserAPIKeyAuth) -> UserAPIKeyAuth: if not is_ui_session_credential(user_api_key_auth): return user_api_key_auth - from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view + from litellm.proxy.management_endpoints.common_utils import user_api_key_has_admin_view - if _user_has_admin_view(user_api_key_auth): + if user_api_key_has_admin_view(user_api_key_auth): return user_api_key_auth admitted: Final = await admitted_user_context(user_api_key_auth) return admitted if admitted is not None else user_api_key_auth diff --git a/litellm/proxy/_lazy_features.py b/litellm/proxy/_lazy_features.py index f018029058f..517c0916e13 100644 --- a/litellm/proxy/_lazy_features.py +++ b/litellm/proxy/_lazy_features.py @@ -267,7 +267,7 @@ LAZY_FEATURES: Final[tuple[LazyFeature, ...]] = ( LazyFeature( name="decisions", module_path="litellm.proxy.decisions_endpoints.endpoints", - path_prefixes=("/v1/decisions", "/decisions"), + path_prefixes=("/v1/decisions", "/decisions", "/v1/systemone", "/systemone"), ), LazyFeature( name="claude_code_marketplace", diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index 06f1a666a70..9322fb77615 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -9521,6 +9521,30 @@ ] } }, + "/systemone": { + "post": { + "operationId": "systemone_systemone_post", + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Systemone", + "tags": [ + "decisions" + ] + } + }, "/v1/decisions": { "post": { "operationId": "decisions_v1_decisions_post", @@ -9544,6 +9568,30 @@ "decisions" ] } + }, + "/v1/systemone": { + "post": { + "operationId": "systemone_v1_systemone_post", + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Systemone", + "tags": [ + "decisions" + ] + } } } }, diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 0ed2b38cfa6..bc27b2352a9 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -488,6 +488,8 @@ class LiteLLMRoutes(enum.Enum): "/v1/search/{search_tool_name}", "/decisions", "/v1/decisions", + "/systemone", + "/v1/systemone", # OCR "/ocr", "/v1/ocr", @@ -549,6 +551,8 @@ class LiteLLMRoutes(enum.Enum): "/lens/{lens_id}/executions/{execution_id}", "/lens/{lens_id}/cancel", "/lens/{lens_id}/findings/{finding_id}", + "/lens/feedback", + "/lens/feedback/summary", "/lens/preview/sample", "/lens/workers/register", "/lens/workers/{worker_id}", @@ -853,6 +857,7 @@ class LiteLLMRoutes(enum.Enum): "/public/mcp_hub", "/public/skill_hub", "/public/litellm_model_cost_map", + "/moyai/connect/exchange", ) ) @@ -1063,6 +1068,7 @@ class LiteLLMRoutes(enum.Enum): admin_viewer_routes = ( [ "/lens/traces/findings", + "/lens/feedback/summary", "/user/list", "/user/available_users", "/user/available_roles", @@ -2737,6 +2743,9 @@ class ScheduledJobStaggerSettings(LiteLLMPydanticObjectBase): ) +DEFAULT_RESPONSES_WEBSOCKET_SESSION_LIMIT_SECONDS: Final[float] = 3600.0 + + class ConfigGeneralSettings(LiteLLMPydanticObjectBase): """ Documents all the fields supported by `general_settings` in config.yaml @@ -3047,6 +3056,12 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase): default=None, description="Default upstream request timeout in seconds for native and custom pass-through endpoints that use pass_through_request. Defaults to 600 when unset.", ) + responses_websocket_session_limit_seconds: float = Field( + default=DEFAULT_RESPONSES_WEBSOCKET_SESSION_LIMIT_SECONDS, + ge=60, + le=7200, + description="Maximum lifetime in seconds of a Responses API WebSocket session, measured from connection accept and covering the idle wait for the first response.create frame. Defaults to 3600, matching OpenAI's documented 60-minute WebSocket connection limit. Must be between 60 and 7200 seconds.", + ) pass_through_endpoints: list[PassThroughGenericEndpoint] | None = Field( default=None, description="Set-up pass-through endpoints for provider-specific endpoints. Docs - https://docs.litellm.ai/docs/proxy/pass_through", diff --git a/litellm/proxy/agent_endpoints/auth/agent_permission_handler.py b/litellm/proxy/agent_endpoints/auth/agent_permission_handler.py index 4e8880d37c5..f671e1f9098 100644 --- a/litellm/proxy/agent_endpoints/auth/agent_permission_handler.py +++ b/litellm/proxy/agent_endpoints/auth/agent_permission_handler.py @@ -459,9 +459,9 @@ class AgentRequestHandler: """ Resolve unified access group ids to agent IDs. """ - from litellm.proxy.auth.auth_checks import _get_agent_ids_from_access_groups + from litellm.proxy.auth.auth_checks import get_agent_ids_from_access_groups - return await _get_agent_ids_from_access_groups(access_group_ids=access_group_ids, check_db_only=check_db_only) + return await get_agent_ids_from_access_groups(access_group_ids=access_group_ids, check_db_only=check_db_only) @staticmethod async def _get_agents_from_access_groups( diff --git a/litellm/proxy/agent_endpoints/auth/managed_authorization.py b/litellm/proxy/agent_endpoints/auth/managed_authorization.py index 5d8287227a5..64e50cd74e7 100644 --- a/litellm/proxy/agent_endpoints/auth/managed_authorization.py +++ b/litellm/proxy/agent_endpoints/auth/managed_authorization.py @@ -30,6 +30,7 @@ _MANAGED_MODEL_ROUTES: Final = frozenset( "moderations", "rerank", "decisions", + "systemone", "ocr", ), ) @@ -74,6 +75,7 @@ _MODEL_ROUTE_KINDS: Final[ "/audio/speech": "speech", "/rerank": "body", "/decisions": "body", + "/systemone": "body", "/messages/count_tokens": "body", ":countTokens": "path", } diff --git a/litellm/proxy/agent_endpoints/endpoints.py b/litellm/proxy/agent_endpoints/endpoints.py index fdc7f2900a9..af39c1818b2 100644 --- a/litellm/proxy/agent_endpoints/endpoints.py +++ b/litellm/proxy/agent_endpoints/endpoints.py @@ -335,9 +335,9 @@ async def _resolve_daily_activity_agent_ids( user_api_key_dict: UserAPIKeyAuth, prisma_client: PrismaClient, ) -> tuple[str, ...] | None: - from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view + from litellm.proxy.management_endpoints.common_utils import user_api_key_has_admin_view - if _user_has_admin_view(user_api_key_dict): + if user_api_key_has_admin_view(user_api_key_dict): return agent_ids permitted_agent_ids: Final = await _permitted_daily_activity_agent_ids( user_api_key_dict=user_api_key_dict, diff --git a/litellm/proxy/anthropic_endpoints/claude_code_endpoints/claude_code_marketplace.py b/litellm/proxy/anthropic_endpoints/claude_code_endpoints/claude_code_marketplace.py index 78af282941d..2f0bdd20a8c 100644 --- a/litellm/proxy/anthropic_endpoints/claude_code_endpoints/claude_code_marketplace.py +++ b/litellm/proxy/anthropic_endpoints/claude_code_endpoints/claude_code_marketplace.py @@ -72,7 +72,7 @@ class _MarketplaceEntry(TypedDict, total=False): category: object -async def _get_prisma_client() -> object: +async def get_prisma_client() -> object: """Get the prisma client from proxy_server.""" from litellm.proxy.proxy_server import prisma_client @@ -84,6 +84,9 @@ async def _get_prisma_client() -> object: return prisma_client +_get_prisma_client: Final = get_prisma_client + + @router.get( "/claude-code/marketplace.json", tags=["Claude Code Marketplace"], @@ -111,7 +114,7 @@ async def get_marketplace(request: Request, key: str | None = None): ``` """ try: - prisma_client: Final = await _get_prisma_client() + prisma_client: Final = await get_prisma_client() caller: Final[UserAPIKeyAuth | None] = ( await user_api_key_auth(request=request, api_key=f"Bearer {key}") if key else None @@ -328,7 +331,7 @@ async def register_plugin( try: _require_proxy_admin(user_api_key_dict) - prisma_client: Final = await _get_prisma_client() + prisma_client: Final = await get_prisma_client() if not re.match(r"^[a-z0-9-]+$", request.name): raise HTTPException( @@ -408,7 +411,7 @@ async def list_plugins( List of plugins with their metadata. """ try: - prisma_client: Final = await _get_prisma_client() + prisma_client: Final = await get_prisma_client() visibility: Final[SkillVisibility] = skill_visibility(user_api_key_dict) plugins: Final[Sequence[_PluginRecord]] = await ClaudeCodePluginRepository(prisma_client).table.find_many( @@ -478,7 +481,7 @@ async def get_plugin( Plugin details including source and metadata. """ try: - prisma_client: Final = await _get_prisma_client() + prisma_client: Final = await get_prisma_client() plugin: Final[_PluginRecord | None] = await ClaudeCodePluginRepository(prisma_client).table.find_unique( where={"name": plugin_name} @@ -579,7 +582,7 @@ async def update_plugin( try: _require_proxy_admin(user_api_key_dict) - prisma_client: Final = await _get_prisma_client() + prisma_client: Final = await get_prisma_client() _validate_plugin_source(request.source) @@ -646,7 +649,7 @@ async def enable_plugin( try: _require_proxy_admin(user_api_key_dict) - prisma_client: Final = await _get_prisma_client() + prisma_client: Final = await get_prisma_client() plugin: Final[_PluginRecord | None] = await ClaudeCodePluginRepository(prisma_client).table.find_unique( where={"name": plugin_name} @@ -695,7 +698,7 @@ async def disable_plugin( try: _require_proxy_admin(user_api_key_dict) - prisma_client: Final = await _get_prisma_client() + prisma_client: Final = await get_prisma_client() plugin: Final[_PluginRecord | None] = await ClaudeCodePluginRepository(prisma_client).table.find_unique( where={"name": plugin_name} @@ -744,7 +747,7 @@ async def delete_plugin( try: _require_proxy_admin(user_api_key_dict) - prisma_client: Final = await _get_prisma_client() + prisma_client: Final = await get_prisma_client() plugin: Final[_PluginRecord | None] = await ClaudeCodePluginRepository(prisma_client).table.find_unique( where={"name": plugin_name} diff --git a/litellm/proxy/anthropic_endpoints/endpoints.py b/litellm/proxy/anthropic_endpoints/endpoints.py index 2911e7801f7..f2c5d066714 100644 --- a/litellm/proxy/anthropic_endpoints/endpoints.py +++ b/litellm/proxy/anthropic_endpoints/endpoints.py @@ -30,7 +30,10 @@ from litellm.proxy.common_request_processing import ( resolve_litellm_call_id, ) from litellm.proxy.common_utils.error_body_call_id import error_body_call_id -from litellm.proxy.common_utils.http_parsing_utils import _read_request_body +from litellm.proxy.common_utils.http_parsing_utils import ( # noqa: F401 # legacy module exports + _read_request_body, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + read_request_body, +) from litellm.proxy.common_utils.openai_error_payload import ( LITELLM_CALL_ID_HEADER, error_status_code, @@ -145,7 +148,7 @@ async def anthropic_response( version, ) - data: Final = await _read_request_body(request=request) + data: Final = await read_request_body(request=request) base_llm_response_processor: Final = ProxyBaseLLMRequestProcessing(data=data) try: result: Final = await base_llm_response_processor.base_process_llm_request( @@ -319,7 +322,7 @@ async def count_tokens( litellm_call_id: Final = resolve_litellm_call_id(request.headers.get("x-litellm-call-id")) try: - request_data: Final = await _read_request_body(request=request) + request_data: Final = await read_request_body(request=request) data: Final[dict] = {**request_data} # Extract required fields diff --git a/litellm/proxy/anthropic_endpoints/gateway_endpoints.py b/litellm/proxy/anthropic_endpoints/gateway_endpoints.py index 42ed6c425ab..4c4155379cd 100644 --- a/litellm/proxy/anthropic_endpoints/gateway_endpoints.py +++ b/litellm/proxy/anthropic_endpoints/gateway_endpoints.py @@ -52,7 +52,10 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.session_token i from litellm.proxy._types import LiteLLM_UserTable, LitellmUserRoles from litellm.proxy.anthropic_endpoints.endpoints import anthropic_response, count_tokens from litellm.proxy.auth.user_api_key_auth import user_api_key_auth -from litellm.proxy.common_utils.http_parsing_utils import _safe_set_request_parsed_body +from litellm.proxy.common_utils.http_parsing_utils import ( # noqa: F401 # legacy module exports + _safe_set_request_parsed_body, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + safe_set_request_parsed_body, +) from litellm.proxy.management_endpoints.sso_helper_utils import CLI_SSO_SESSIONS_TARGET from litellm.proxy.management_endpoints.ui_sso import CliSsoTeamDetail from litellm.types.llms.base import LiteLLMBaseModel @@ -449,7 +452,7 @@ async def managed_settings(request: Request) -> Response: async def _skip_otlp_body_parsing(request: Request) -> None: - _safe_set_request_parsed_body(request=request, parsed_body={}) + safe_set_request_parsed_body(request=request, parsed_body={}) _OTLP_AUTHENTICATED: Final = (Depends(_skip_otlp_body_parsing), *_AUTHENTICATED) diff --git a/litellm/proxy/anthropic_endpoints/skills_endpoints.py b/litellm/proxy/anthropic_endpoints/skills_endpoints.py index ab513b1acc4..2c9c8efcdd8 100644 --- a/litellm/proxy/anthropic_endpoints/skills_endpoints.py +++ b/litellm/proxy/anthropic_endpoints/skills_endpoints.py @@ -170,7 +170,7 @@ async def create_skill( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -300,7 +300,7 @@ async def list_skills( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -397,7 +397,7 @@ async def get_skill( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -496,7 +496,7 @@ async def delete_skill( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index bea62f40c76..1435033e1be 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -96,9 +96,11 @@ from litellm.proxy.auth.model_access_denied import ( from litellm.proxy.auth.route_checks import RouteChecks from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import publish_auth_cache_invalidation from litellm.proxy.common_utils.cache_pydantic_utils import CacheCodec -from litellm.proxy.common_utils.http_parsing_utils import ( - _safe_get_request_headers, - _safe_get_request_query_params, +from litellm.proxy.common_utils.http_parsing_utils import ( # noqa: F401 # legacy module exports + _safe_get_request_headers, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _safe_get_request_query_params, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + safe_get_request_headers, + safe_get_request_query_params, ) from litellm.proxy.common_utils.model_listing_utils import alias_map from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time @@ -463,7 +465,7 @@ def _get_router_zero_cost_cache(llm_router: Router) -> dict[str, bool] | None: return cache if isinstance(cache, dict) else None -def _is_model_cost_zero(model: str | list[str] | None, llm_router: Router | None) -> bool: +def is_model_cost_zero(model: str | list[str] | None, llm_router: Router | None) -> bool: """ Check if a model has zero cost (no configured pricing). @@ -582,6 +584,9 @@ def _is_model_cost_zero(model: str | list[str] | None, llm_router: Router | None return True +_is_model_cost_zero: Final = is_model_cost_zero + + _NO_MODEL_INFO: Final[Mapping[str, object]] = MappingProxyType({}) _TEAM_GRANT_RELATIONS: Final[Mapping[str, object]] = MappingProxyType({"litellm_model_table": True}) @@ -974,7 +979,7 @@ def route_skips_budget_checks(route: str) -> bool: def request_skips_budget_checks(route: str, model: str | list[str] | None, llm_router: Router | None) -> bool: - return route_skips_budget_checks(route=route) or _is_model_cost_zero(model=model, llm_router=llm_router) + return route_skips_budget_checks(route=route) or is_model_cost_zero(model=model, llm_router=llm_router) async def common_checks( @@ -1017,8 +1022,8 @@ async def common_checks( _model: Final[str | list[str] | None] = get_model_from_request( request_data=request_body, route=route, - request_headers=_safe_get_request_headers(request=request), - request_query_params=_safe_get_request_query_params(request=request), + request_headers=safe_get_request_headers(request=request), + request_query_params=safe_get_request_query_params(request=request), llm_router=llm_router, request=request, team_id=valid_token.team_id if valid_token is not None else None, @@ -1073,7 +1078,7 @@ async def common_checks( except ProxyException as team_denial: if team_denial.type != ProxyErrorTypes.team_model_access_denied: raise - if not await _key_access_group_grants_model( + if not await key_access_group_grants_model( model=_model, valid_token=valid_token, team_object=team_object, @@ -1085,7 +1090,7 @@ async def common_checks( # 2.2. If team member has per-member model scope, enforce it if _model and team_object and valid_token and valid_token.user_id: with tracer.trace("litellm.proxy.auth.common_checks.check_team_member_model_access"): - await _check_team_member_model_access( + await check_team_member_model_access( model=_model, team_object=team_object, valid_token=valid_token, @@ -1122,7 +1127,7 @@ async def common_checks( managed_models: Final = (managed_policy.object_permission or MappingProxyType({})).get("models", ()) if not isinstance(managed_models, (list, tuple)) or not managed_models: raise HTTPException(403, "This agent has no model grants") - _can_object_call_model( + can_object_call_model( model=_resolve_team_alias( _model, team_model_aliases_for_auth_check(valid_token), valid_token.team_id, llm_router ), @@ -1261,7 +1266,7 @@ async def common_checks( budget_check_coros: Final = tuple( coro for coro in ( - _team_max_budget_check( + team_max_budget_check( team_object=team_object, proxy_logging_obj=proxy_logging_obj, valid_token=valid_token, @@ -1273,7 +1278,7 @@ async def common_checks( proxy_logging_obj=proxy_logging_obj, valid_token=valid_token, ), - _organization_max_budget_check( + organization_max_budget_check( valid_token=valid_token, team_object=team_object, prisma_client=prisma_client, @@ -1305,7 +1310,7 @@ async def common_checks( team_membership=loaded_team_membership, team_membership_loaded=team_membership_loaded, ), - _check_end_user_budget(end_user_obj=end_user_object, route=route) + check_end_user_budget(end_user_obj=end_user_object, route=route) if end_user_object is not None and end_user_object.litellm_budget_table is not None else None, ) @@ -1381,7 +1386,7 @@ def effective_user_role(user_role: str | None) -> LitellmUserRoles: return LitellmUserRoles.INTERNAL_USER -def _get_user_role( +def get_user_role( user_obj: LiteLLM_UserTable | None, ) -> LitellmUserRoles | None: if user_obj is None: @@ -1389,6 +1394,9 @@ def _get_user_role( return effective_user_role(user_obj.user_role) +_get_user_role: Final = get_user_role + + def _is_api_route_allowed( route: str, request: Request, @@ -1399,12 +1407,12 @@ def _is_api_route_allowed( """ - Route b/w api token check and normal token check """ - _user_role: Final = _get_user_role(user_obj=user_obj) + _user_role: Final = get_user_role(user_obj=user_obj) if valid_token is None: raise Exception("Invalid proxy server token passed. valid_token=None.") - if not _is_user_proxy_admin(user_obj=user_obj): # if non-admin + if not is_user_proxy_admin(user_obj=user_obj): # if non-admin RouteChecks.non_proxy_admin_allowed_routes_check( user_obj=user_obj, _user_role=_user_role, @@ -1416,7 +1424,7 @@ def _is_api_route_allowed( return True -def _is_user_proxy_admin(user_obj: LiteLLM_UserTable | None): +def is_user_proxy_admin(user_obj: LiteLLM_UserTable | None) -> bool: if user_obj is None: return False @@ -1426,6 +1434,9 @@ def _is_user_proxy_admin(user_obj: LiteLLM_UserTable | None): return False +_is_user_proxy_admin: Final = is_user_proxy_admin + + def _allowed_routes_check(user_route: str, allowed_routes: list) -> bool: """ Return if a user is allowed to access route. Helper function for `allowed_routes_check`. @@ -1723,7 +1734,7 @@ async def _apply_default_budget_to_end_user( return end_user_obj.model_copy(update=MappingProxyType({"litellm_budget_table": default_budget})) -async def _check_end_user_budget( +async def check_end_user_budget( end_user_obj: LiteLLM_EndUserTable, route: str, ) -> None: @@ -1765,6 +1776,9 @@ async def _check_end_user_budget( ) +_check_end_user_budget: Final = check_end_user_budget + + #: Columns whose non-null value makes an end-user row restrict something auth enforces. ``blocked`` #: is separate: it restricts when true rather than when merely set. _RESTRICTED_COLUMNS: Final = ("budget_id", "allowed_model_region", "default_model", "object_permission_id") @@ -2900,12 +2914,12 @@ async def _cache_management_object( @with_service_target(AUTH_OBJECTS_TARGET) -async def _cache_team_object( +async def cache_team_object( team_id: str, team_table: LiteLLM_TeamTableCachedObj, user_api_key_cache: UserApiKeyCache, proxy_logging_obj: ProxyLogging | None, -): +) -> None: ## CACHE REFRESH TIME! team_table.last_refreshed_at = time.time() @@ -2957,6 +2971,9 @@ async def _cache_team_object( await _invalidate_usage_cache_entry(usage_cache, alias_key, redis_shared=redis_shared, stale="team alias") +_cache_team_object: Final = cache_team_object + + @with_service_target(SPEND_COUNTERS_TARGET) async def _invalidate_usage_cache_entry( usage_cache: DualCache | None, @@ -3137,18 +3154,18 @@ async def delete_cache_team_object( await publish_auth_cache_invalidation(cache_key=key) -async def _cache_key_object( +async def cache_key_object( hashed_token: str, user_api_key_obj: UserAPIKeyAuth, user_api_key_cache: UserApiKeyCache, proxy_logging_obj: ProxyLogging | None, -): +) -> None: key: Final = hashed_token ## CACHE REFRESH TIME user_api_key_obj.last_refreshed_at = time.time() - cached_key_obj: Final = _copy_user_api_key_auth_for_cache(user_api_key_obj=user_api_key_obj) + cached_key_obj: Final = copy_user_api_key_auth_for_cache(user_api_key_obj=user_api_key_obj) await _cache_management_object( key=key, value=cached_key_obj, @@ -3158,12 +3175,15 @@ async def _cache_key_object( ) +_cache_key_object: Final = cache_key_object + + @with_service_target(AUTH_OBJECTS_TARGET) -async def _delete_cache_key_object( +async def delete_cache_key_object( hashed_token: str, user_api_key_cache: UserApiKeyCache, proxy_logging_obj: ProxyLogging | None, -): +) -> None: """ Evict one key object, best-effort, matching `delete_cache_team_object` and `delete_cache_key_objects`. @@ -3196,6 +3216,9 @@ async def _delete_cache_key_object( await publish_auth_cache_invalidation(cache_key=key) +_delete_cache_key_object: Final = delete_cache_key_object + + async def delete_cache_key_objects( hashed_tokens: Sequence[str], user_api_key_cache: UserApiKeyCache, @@ -3215,7 +3238,7 @@ async def delete_cache_key_objects( """ results: Final = await asyncio.gather( *( - _delete_cache_key_object( + delete_cache_key_object( hashed_token=hashed_token, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, @@ -3339,7 +3362,7 @@ async def _get_team_object_from_user_api_key_cache( ) # save the team object to cache - await _cache_team_object( + await cache_team_object( team_id=team_id, team_table=_response, user_api_key_cache=user_api_key_cache, @@ -3357,7 +3380,7 @@ async def _get_team_object_from_user_api_key_cache( @with_service_target(AUTH_OBJECTS_TARGET) -async def _get_team_object_from_cache( +async def get_team_object_from_cache( key: str, user_api_key_cache: UserApiKeyCache, parent_otel_span: Span | None, @@ -3370,6 +3393,9 @@ async def _get_team_object_from_cache( return decoded +_get_team_object_from_cache: Final = get_team_object_from_cache + + async def get_team_object( team_id: str, prisma_client: PrismaClient | None, @@ -3395,7 +3421,7 @@ async def get_team_object( key: Final = f"team_id:{team_id}" if not check_db_only: - cached_team_obj: Final = await _get_team_object_from_cache( + cached_team_obj: Final = await get_team_object_from_cache( key=key, user_api_key_cache=user_api_key_cache, parent_otel_span=parent_otel_span, @@ -3433,12 +3459,12 @@ async def get_team_object( @with_service_target(AUTH_OBJECTS_TARGET) -async def _cache_access_object( +async def cache_access_object( access_group_id: str, access_group_table: LiteLLM_AccessGroupTable, user_api_key_cache: UserApiKeyCache, proxy_logging_obj: ProxyLogging | None = None, -): +) -> None: key: Final = f"access_group_id:{access_group_id}" await user_api_key_cache.async_set_cache( key=key, @@ -3448,12 +3474,15 @@ async def _cache_access_object( ) +_cache_access_object: Final = cache_access_object + + @with_service_target(AUTH_OBJECTS_TARGET) -async def _delete_cache_access_object( +async def delete_cache_access_object( access_group_id: str, user_api_key_cache: UserApiKeyCache, proxy_logging_obj: ProxyLogging | None = None, -): +) -> None: key: Final = f"access_group_id:{access_group_id}" user_api_key_cache.delete_cache(key=key) @@ -3463,6 +3492,9 @@ async def _delete_cache_access_object( await proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache(key=key) +_delete_cache_access_object: Final = delete_cache_access_object + + @log_db_metrics @with_service_target(AUTH_OBJECTS_TARGET) async def get_access_object( @@ -3510,7 +3542,7 @@ async def get_access_object( _response: Final = LiteLLM_AccessGroupTable.model_validate(response.dict()) # Save to cache - await _cache_access_object( + await cache_access_object( access_group_id=access_group_id, access_group_table=_response, user_api_key_cache=user_api_key_cache, @@ -3566,7 +3598,7 @@ async def get_team_object_by_alias( # Check cache first (keyed by alias) cache_key: Final = f"team_alias:{team_alias}" - cached_team_obj: Final = await _get_team_object_from_cache( + cached_team_obj: Final = await get_team_object_from_cache( key=cache_key, user_api_key_cache=user_api_key_cache, parent_otel_span=parent_otel_span, @@ -3861,7 +3893,7 @@ class ExperimentalUIJWTToken: raise Exception(f"Invalid hash key. Hash key={hashed_token}. Decrypted token={decrypted_token}. Error: {e}") -async def _fetch_key_object_from_db_with_reconnect( +async def fetch_key_object_from_db_with_reconnect( hashed_token: str, prisma_client: PrismaClient, parent_otel_span: Span | None, @@ -3889,6 +3921,9 @@ async def _fetch_key_object_from_db_with_reconnect( ) +_fetch_key_object_from_db_with_reconnect: Final = fetch_key_object_from_db_with_reconnect + + async def _fetch_key_object_from_db_unbounded( hashed_token: str, prisma_client: PrismaClient, @@ -4031,13 +4066,13 @@ async def get_key_object( None if check_db_only else await user_api_key_cache.async_get_cache(key=key, model_type=UserAPIKeyAuth) ) if user_api_key_auth is not None: - return _copy_user_api_key_auth_for_cache(user_api_key_obj=user_api_key_auth) + return copy_user_api_key_auth_for_cache(user_api_key_obj=user_api_key_auth) if check_cache_only: raise Exception(f"Key doesn't exist in cache + check_cache_only=True. key={key}.") # else, check db - _valid_token: Final[BaseModel | None] = await _fetch_key_object_from_db_with_reconnect( + _valid_token: Final[BaseModel | None] = await fetch_key_object_from_db_with_reconnect( hashed_token=hashed_token, prisma_client=prisma_client, parent_otel_span=parent_otel_span, @@ -4079,7 +4114,7 @@ async def get_key_object( return _response # save the key object to cache - await _cache_key_object( + await cache_key_object( hashed_token=hashed_token, user_api_key_obj=_response, user_api_key_cache=user_api_key_cache, @@ -4089,7 +4124,7 @@ async def get_key_object( return _response -def _copy_user_api_key_auth_for_cache( +def copy_user_api_key_auth_for_cache( user_api_key_obj: UserAPIKeyAuth, ) -> UserAPIKeyAuth: copied_key_obj: Final = user_api_key_obj.model_copy() @@ -4100,6 +4135,9 @@ def _copy_user_api_key_auth_for_cache( return copied_key_obj +_copy_user_api_key_auth_for_cache: Final = copy_user_api_key_auth_for_cache + + @log_db_metrics @with_service_target(AUTH_OBJECTS_TARGET) async def get_object_permission( @@ -4424,7 +4462,7 @@ async def _get_resources_from_access_groups( return list(set(resources)) -async def _get_models_from_access_groups( +async def get_models_from_access_groups( access_group_ids: Sequence[str], prisma_client: DatabaseClient | None = None, user_api_key_cache: UserApiKeyCache | None = None, @@ -4443,7 +4481,10 @@ async def _get_models_from_access_groups( ) -async def _get_mcp_server_ids_from_access_groups( +_get_models_from_access_groups: Final = get_models_from_access_groups + + +async def get_mcp_server_ids_from_access_groups( access_group_ids: list[str], prisma_client: PrismaClient | None = None, user_api_key_cache: UserApiKeyCache | None = None, @@ -4464,7 +4505,10 @@ async def _get_mcp_server_ids_from_access_groups( ) -async def _get_agent_ids_from_access_groups( +_get_mcp_server_ids_from_access_groups: Final = get_mcp_server_ids_from_access_groups + + +async def get_agent_ids_from_access_groups( access_group_ids: list[str], prisma_client: PrismaClient | None = None, user_api_key_cache: UserApiKeyCache | None = None, @@ -4485,6 +4529,9 @@ async def _get_agent_ids_from_access_groups( ) +_get_agent_ids_from_access_groups: Final = get_agent_ids_from_access_groups + + def _resolve_all_team_model_sentinel_for_auth_check( models: list[str], llm_router: Router | None, @@ -4499,7 +4546,7 @@ def _resolve_all_team_model_sentinel_for_auth_check( return list(dict.fromkeys(non_sentinel_models + proxy_models)) -def _check_model_access_helper( +def check_model_access_helper( model: str, llm_router: Router | None, models: list[str], @@ -4547,6 +4594,9 @@ def _check_model_access_helper( return True +_check_model_access_helper: Final = check_model_access_helper + + def _can_object_call_model( model: str | list[str], llm_router: Router | None, @@ -4621,7 +4671,7 @@ def _can_object_call_model( ## check model access for alias + underlying model - allow if either is in allowed models for m in potential_models: - if _check_model_access_helper( + if check_model_access_helper( model=m, llm_router=llm_router, models=models, @@ -4647,6 +4697,9 @@ def _can_object_call_model( ) +can_object_call_model: Final = _can_object_call_model + + def _resolve_team_alias( model: str | list[str], team_model_aliases: Mapping[str, str] | None, @@ -4707,7 +4760,7 @@ async def _check_agent_access_group_model_access( param="model", code=status.HTTP_403_FORBIDDEN, ) - _can_object_call_model( + can_object_call_model( model=dispatched, llm_router=llm_router, models=sorted(ceiling.models), @@ -4751,7 +4804,7 @@ async def _check_agent_caller_model_access( prisma_client=prisma_client, key_model_aliases=caller_key_model_aliases, ) - await _check_team_member_model_access( + await check_team_member_model_access( model=model, team_object=caller_team, valid_token=caller_auth, @@ -5112,7 +5165,7 @@ async def can_key_call_model( """ key_models: Final = _resolve_key_models_for_auth_check(valid_token=valid_token) try: - return _can_object_call_model( + return can_object_call_model( model=model, llm_router=llm_router, models=key_models, @@ -5125,12 +5178,12 @@ async def can_key_call_model( # Fallback: check key's access_group_ids key_access_group_ids: Final = valid_token.access_group_ids or [] if key_access_group_ids: - models_from_groups: Final = await _get_models_from_access_groups( + models_from_groups: Final = await get_models_from_access_groups( access_group_ids=key_access_group_ids, prisma_client=prisma_client, ) if models_from_groups: - return _can_object_call_model( + return can_object_call_model( model=model, llm_router=llm_router, models=models_from_groups, @@ -5200,7 +5253,7 @@ async def can_key_call_resolved_model( except ProxyException as team_denial: if team_denial.type != ProxyErrorTypes.team_model_access_denied: raise - if not await _key_access_group_grants_model( + if not await key_access_group_grants_model( model=model, valid_token=valid_token, team_object=team_object, @@ -5210,7 +5263,7 @@ async def can_key_call_resolved_model( raise if valid_token.user_id is not None and team_object_from_lookup: - await _check_team_member_model_access( + await check_team_member_model_access( model=model, team_object=team_object, valid_token=valid_token, @@ -5265,7 +5318,7 @@ def can_org_access_model( Returns True if the team can access a specific model. """ - return _can_object_call_model( + return can_object_call_model( model=model, llm_router=llm_router, models=org_object.models if org_object else [], @@ -5289,7 +5342,7 @@ async def can_team_access_model( 2. If not allowed natively, falls back to access_group_ids on the team """ try: - return _can_object_call_model( + return can_object_call_model( model=model, llm_router=llm_router, models=team_object.models if team_object else [], @@ -5302,12 +5355,12 @@ async def can_team_access_model( # Fallback: check team's access_group_ids team_access_group_ids: Final = (team_object.access_group_ids or []) if team_object else [] if team_access_group_ids: - models_from_groups: Final = await _get_models_from_access_groups( + models_from_groups: Final = await get_models_from_access_groups( access_group_ids=team_access_group_ids, prisma_client=prisma_client, ) if models_from_groups: - return _can_object_call_model( + return can_object_call_model( model=model, llm_router=llm_router, models=list(dict.fromkeys([*(team_object.models if team_object else []), *models_from_groups])), @@ -5367,7 +5420,7 @@ async def get_authorized_resources_from_key_access_groups( return list(set(authorized_resources)) -async def _key_access_group_grants_model( +async def key_access_group_grants_model( model: str | list[str], valid_token: UserAPIKeyAuth | None, team_object: LiteLLM_TeamTable | None, @@ -5387,7 +5440,7 @@ async def _key_access_group_grants_model( if not authorized_models: return False try: - _can_object_call_model( + can_object_call_model( model=model, llm_router=llm_router, models=authorized_models, @@ -5401,6 +5454,9 @@ async def _key_access_group_grants_model( return False +_key_access_group_grants_model: Final = key_access_group_grants_model + + def can_project_access_model( model: str | list[str], project_object: LiteLLM_ProjectTable, @@ -5412,7 +5468,7 @@ def can_project_access_model( Raises ProxyException if access is denied. """ - return _can_object_call_model( + return can_object_call_model( model=model, llm_router=llm_router, models=project_object.models if project_object else [], @@ -5437,7 +5493,7 @@ def can_customer_access_model( ) if team_target != name and name in (end_user_object.models or ()): return - _can_object_call_model( + can_object_call_model( model=team_target, llm_router=llm_router, models=end_user_object.models, @@ -5472,7 +5528,7 @@ async def can_user_call_model( code=status.HTTP_403_FORBIDDEN, ) - return _can_object_call_model( + return can_object_call_model( model=model, llm_router=llm_router, models=user_object.models, @@ -5764,11 +5820,11 @@ def _apply_budget_exceeded_throttle(valid_token: UserAPIKeyAuth) -> bool: return True -async def _virtual_key_max_budget_check( +async def virtual_key_max_budget_check( valid_token: UserAPIKeyAuth, proxy_logging_obj: ProxyLogging, user_obj: LiteLLM_UserTable | None = None, -): +) -> None: """ Raises: BudgetExceededError if the token is over it's max budget. @@ -5844,6 +5900,9 @@ async def _virtual_key_max_budget_check( ) +_virtual_key_max_budget_check: Final = virtual_key_max_budget_check + + async def _virtual_key_multi_budget_check( valid_token: UserAPIKeyAuth, ): @@ -5888,11 +5947,11 @@ async def _virtual_key_multi_budget_check( ) -async def _virtual_key_soft_budget_check( +async def virtual_key_soft_budget_check( valid_token: UserAPIKeyAuth, proxy_logging_obj: ProxyLogging, user_obj: LiteLLM_UserTable | None = None, -): +) -> None: """ Triggers a budget alert if the token is over it's soft budget. @@ -5927,6 +5986,9 @@ async def _virtual_key_soft_budget_check( ) +_virtual_key_soft_budget_check: Final = virtual_key_soft_budget_check + + def _parse_email_list(raw: str | Sequence[object] | None) -> list[str]: """Parse emails from a list or comma-separated string.""" if isinstance(raw, list): @@ -5968,11 +6030,11 @@ def _merge_budget_alert_email_configs( } -async def _virtual_key_max_budget_alert_check( +async def virtual_key_max_budget_alert_check( valid_token: UserAPIKeyAuth, proxy_logging_obj: ProxyLogging, user_obj: LiteLLM_UserTable | None = None, -): +) -> None: """ Triggers a budget alert if the token has reached EMAIL_BUDGET_ALERT_MAX_SPEND_ALERT_PERCENTAGE (default 80%) of its max budget. @@ -6050,6 +6112,9 @@ async def _virtual_key_max_budget_alert_check( ) +_virtual_key_max_budget_alert_check: Final = virtual_key_max_budget_alert_check + + TEAM_MEMBER_MAX_BUDGET_ALERT_EMAILS_KEY: Final = "team_member_max_budget_alert_emails" _TEAM_MEMBER_ALERT_CONFIG_ADAPTER: Final[TypeAdapter[Mapping[str, object]]] = TypeAdapter(Mapping[str, object]) @@ -6074,7 +6139,7 @@ def _valid_alert_threshold_config(raw_config: object) -> Mapping[str, str | Sequ ) -def _team_member_max_budget_alert_check( +def team_member_max_budget_alert_check( team_id: str, team_alias: str | None, team_metadata: Mapping[str, object] | None, @@ -6108,6 +6173,9 @@ def _team_member_max_budget_alert_check( asyncio.create_task(proxy_logging_obj.budget_alerts(type="max_budget_alert", user_info=call_info)) +_team_member_max_budget_alert_check: Final = team_member_max_budget_alert_check + + async def _check_team_member_budget( team_object: LiteLLM_TeamTable | None, user_object: LiteLLM_UserTable | None, @@ -6176,7 +6244,7 @@ async def _check_team_member_budget( if not math.isfinite(team_member_budget): return - _team_member_max_budget_alert_check( + team_member_max_budget_alert_check( team_id=team_object.team_id, team_alias=team_object.team_alias, team_metadata=team_object.metadata, @@ -6198,7 +6266,7 @@ async def _check_team_member_budget( ) -async def _check_team_member_model_access( +async def check_team_member_model_access( model: str | list[str], team_object: LiteLLM_TeamTable, valid_token: UserAPIKeyAuth, @@ -6238,7 +6306,7 @@ async def _check_team_member_model_access( member_allowed_models: Final[list[str]] = loaded_membership.litellm_budget_table.allowed_models try: - _can_object_call_model( + can_object_call_model( model=model, llm_router=llm_router, models=member_allowed_models, @@ -6260,11 +6328,14 @@ async def _check_team_member_model_access( ) -async def _team_max_budget_check( +_check_team_member_model_access: Final = check_team_member_model_access + + +async def team_max_budget_check( team_object: LiteLLM_TeamTable | None, valid_token: UserAPIKeyAuth | None, proxy_logging_obj: ProxyLogging, -): +) -> None: """ Check if the team is over it's max budget. @@ -6310,6 +6381,9 @@ async def _team_max_budget_check( ) +_team_max_budget_check: Final = team_max_budget_check + + async def _team_multi_budget_check( team_object: LiteLLM_TeamTable | None, ): @@ -6588,13 +6662,13 @@ async def delete_cached_project_object( ) -async def _organization_max_budget_check( +async def organization_max_budget_check( valid_token: UserAPIKeyAuth | None, team_object: LiteLLM_TeamTable | None, prisma_client: PrismaClient | None, user_api_key_cache: UserApiKeyCache, proxy_logging_obj: ProxyLogging, -): +) -> None: """ Check if the organization is over its max budget. @@ -6687,6 +6761,9 @@ async def _organization_max_budget_check( ) +_organization_max_budget_check: Final = organization_max_budget_check + + async def _tag_max_budget_check( request_body: dict, prisma_client: PrismaClient | None, diff --git a/litellm/proxy/auth/auth_checks_organization.py b/litellm/proxy/auth/auth_checks_organization.py index 37c025f0b2c..6820ded0eba 100644 --- a/litellm/proxy/auth/auth_checks_organization.py +++ b/litellm/proxy/auth/auth_checks_organization.py @@ -133,7 +133,7 @@ def get_user_organization_info( return _user_organizations, _user_organization_role_mapping -def _user_is_org_admin( +def user_is_org_admin( request_data: dict, user_object: LiteLLM_UserTable | None = None, ) -> bool: @@ -173,6 +173,9 @@ def _user_is_org_admin( return all(org_id in admin_org_ids for org_id in candidate_org_ids) +_user_is_org_admin: Final = user_is_org_admin + + TEAM_ORG_CONTEXT_ROUTES: Final = frozenset({"/team/update"}) # The RESTful update route carries the team id in the path. Match on the route # template so the sibling /team/ routes (which share the single-segment diff --git a/litellm/proxy/auth/auth_exception_handler.py b/litellm/proxy/auth/auth_exception_handler.py index a59a12d6807..0266131fc33 100644 --- a/litellm/proxy/auth/auth_exception_handler.py +++ b/litellm/proxy/auth/auth_exception_handler.py @@ -20,8 +20,9 @@ from litellm.proxy._types import ( ProxyException, UserAPIKeyAuth, ) -from litellm.proxy.auth.auth_utils import ( - _get_request_ip_address, +from litellm.proxy.auth.auth_utils import ( # noqa: F401 # legacy module exports + _get_request_ip_address, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + get_request_ip_address, is_invalid_virtual_key_error, mark_invalid_virtual_key_error, normalize_request_route, @@ -125,7 +126,7 @@ def _identity_log_suffix(resolved_identity: UserAPIKeyAuth | None) -> str: class UserAPIKeyAuthExceptionHandler: @staticmethod - async def _handle_authentication_error( + async def handle_authentication_error( e: Exception, request: Request, request_data: dict[str, object], @@ -180,7 +181,7 @@ class UserAPIKeyAuthExceptionHandler: ) else: # raise the exception to the caller - requester_ip: Final = _get_request_ip_address( + requester_ip: Final = get_request_ip_address( request=request, use_x_forwarded_for=general_settings.get("use_x_forwarded_for") is True, ) @@ -258,3 +259,5 @@ class UserAPIKeyAuthExceptionHandler: extra=log_extra, ) raise final_exception + + _handle_authentication_error = handle_authentication_error diff --git a/litellm/proxy/auth/auth_utils.py b/litellm/proxy/auth/auth_utils.py index e49fa6e9acd..edd6028f07c 100644 --- a/litellm/proxy/auth/auth_utils.py +++ b/litellm/proxy/auth/auth_utils.py @@ -77,7 +77,7 @@ def mark_invalid_virtual_key_error(exception: ProxyException, is_invalid_virtual return marked_exception -def _get_request_ip_address(request: Request, use_x_forwarded_for: bool | None = False) -> str | None: +def get_request_ip_address(request: Request, use_x_forwarded_for: bool | None = False) -> str | None: client_ip = None if use_x_forwarded_for is True and "x-forwarded-for" in request.headers: client_ip = request.headers["x-forwarded-for"] @@ -89,6 +89,9 @@ def _get_request_ip_address(request: Request, use_x_forwarded_for: bool | None = return client_ip +_get_request_ip_address: Final = get_request_ip_address + + def _check_valid_ip( allowed_ips: list[str] | None, request: Request, @@ -101,7 +104,7 @@ def _check_valid_ip( return True, None # if general_settings.get("use_x_forwarded_for") is True then use x-forwarded-for - client_ip: Final = _get_request_ip_address(request=request, use_x_forwarded_for=use_x_forwarded_for) + client_ip: Final = get_request_ip_address(request=request, use_x_forwarded_for=use_x_forwarded_for) # Check if IP address is allowed if client_ip not in allowed_ips: @@ -1884,14 +1887,14 @@ def _extract_models_from_managed_resource_id( try: from litellm.proxy.openai_files_endpoints.common_utils import ( - _is_base64_encoded_unified_file_id, decode_model_from_file_id, get_model_id_from_unified_batch_id, get_models_from_unified_file_id, + is_base64_encoded_unified_file_id, ) _append_model_candidates(candidates=candidates, value=decode_model_from_file_id(resource_id)) - unified_file_id: Final = _is_base64_encoded_unified_file_id(resource_id) + unified_file_id: Final = is_base64_encoded_unified_file_id(resource_id) if unified_file_id: _append_model_candidates( candidates=candidates, @@ -2180,7 +2183,7 @@ def _router_model_from_azure_route(route: str, llm_router: Router | None) -> str def _model_from_bedrock_route(route: str) -> str | None: from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( - _extract_model_from_bedrock_endpoint, + extract_model_from_bedrock_endpoint, is_bedrock_count_tokens_endpoint, ) @@ -2188,7 +2191,7 @@ def _model_from_bedrock_route(route: str) -> str | None: if is_bedrock_count_tokens_endpoint(bedrock_endpoint): return None try: - return _extract_model_from_bedrock_endpoint(bedrock_endpoint) + return extract_model_from_bedrock_endpoint(bedrock_endpoint) except ValueError: return None diff --git a/litellm/proxy/auth/authorization.py b/litellm/proxy/auth/authorization.py index cd0a7acff0a..bb498c33763 100644 --- a/litellm/proxy/auth/authorization.py +++ b/litellm/proxy/auth/authorization.py @@ -38,10 +38,10 @@ async def resolve_owned_read_scope( def can_read_team_logs(auth: UserAPIKeyAuth, team: LiteLLM_TeamTable) -> bool: from litellm.proxy.management.teams.authz import is_team_admin from litellm.proxy.management_endpoints.common_utils import ( - _team_member_has_permission, # pyright: ignore[reportPrivateUsage] # reuse existing team permission policy + team_member_has_permission, # pyright: ignore[reportPrivateUsage] # reuse existing team permission policy ) - return is_team_admin(user_api_key_dict=auth, team_obj=team) or _team_member_has_permission( + return is_team_admin(user_api_key_dict=auth, team_obj=team) or team_member_has_permission( user_api_key_dict=auth, team_obj=team, permission=KeyManagementRoutes.SPEND_LOGS.value, diff --git a/litellm/proxy/auth/fallback_budget.py b/litellm/proxy/auth/fallback_budget.py index 9029f9d6996..5aa28a08bb8 100644 --- a/litellm/proxy/auth/fallback_budget.py +++ b/litellm/proxy/auth/fallback_budget.py @@ -39,8 +39,9 @@ from pydantic import ValidationError from litellm._logging import verbose_proxy_logger from litellm.proxy._types import UserAPIKeyAuth -from litellm.proxy.auth.auth_checks import ( - _is_model_cost_zero, # pyright: ignore[reportPrivateUsage] # the zero-cost predicate the auth-time budget checks use; no public equivalent +from litellm.proxy.auth.auth_checks import ( # noqa: F401 # legacy module exports + _is_model_cost_zero, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + is_model_cost_zero, # pyright: ignore[reportPrivateUsage] # the zero-cost predicate the auth-time budget checks use; no public equivalent ) from litellm.router import Router from litellm.types.llms.base import LiteLLMBaseModel @@ -109,7 +110,7 @@ async def is_token_within_budget_for_model(*, model: str, valid_token: UserAPIKe A zero-cost fallback target is always allowed: refusing it would deny a request on spend some other model accrued, which is the same reasoning behind the auth-time bypass. """ - if _is_model_cost_zero(model=model, llm_router=llm_router): + if is_model_cost_zero(model=model, llm_router=llm_router): return True key_budget: Final = valid_token.max_budget diff --git a/litellm/proxy/auth/ip_address_utils.py b/litellm/proxy/auth/ip_address_utils.py index c2614b85016..5d32cc7eb4c 100644 --- a/litellm/proxy/auth/ip_address_utils.py +++ b/litellm/proxy/auth/ip_address_utils.py @@ -16,7 +16,10 @@ from fastapi import Request from pydantic import TypeAdapter, ValidationError from litellm._logging import verbose_proxy_logger -from litellm.proxy.auth.auth_utils import _get_request_ip_address +from litellm.proxy.auth.auth_utils import ( # noqa: F401 # legacy module exports + _get_request_ip_address, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + get_request_ip_address, +) # One-shot warning so operators upgrading from the prior "always trust X-Forwarded-*" # behaviour see an actionable message in their logs the first time it triggers. @@ -368,4 +371,4 @@ class IPAddressUtils: return client_ip case _HopCountUnset(): pass - return _get_request_ip_address(request, use_x_forwarded_for=use_xff) + return get_request_ip_address(request, use_x_forwarded_for=use_xff) diff --git a/litellm/proxy/auth/model_checks.py b/litellm/proxy/auth/model_checks.py index f48f5dd7162..133e95eadba 100644 --- a/litellm/proxy/auth/model_checks.py +++ b/litellm/proxy/auth/model_checks.py @@ -13,7 +13,7 @@ from litellm.router import Router from litellm.router_utils.fallback_event_handlers import get_fallback_model_group from litellm.types.router import CredentialLiteLLMParams, LiteLLM_Params from litellm.types.utils import LlmProviders -from litellm.utils import get_valid_models +from litellm.utils import ProviderConfigManager, get_valid_models _CREDENTIAL_LITELLM_PARAM_FIELDS = set(CredentialLiteLLMParams.model_fields) @@ -45,7 +45,14 @@ def get_provider_models(provider: str, litellm_params: LiteLLM_Params | None = N if provider in litellm.models_by_provider: provider_models: Final = get_valid_models(custom_llm_provider=provider, litellm_params=litellm_params) return provider_models - return None + + try: + llm_provider: Final = LlmProviders(provider) + except ValueError: + return None + if ProviderConfigManager.get_provider_model_info(model=None, provider=llm_provider) is None: + return None + return get_valid_models(custom_llm_provider=provider, litellm_params=litellm_params) def _get_models_from_access_groups( diff --git a/litellm/proxy/auth/resolvers/store.py b/litellm/proxy/auth/resolvers/store.py index 13e24387551..0c8d05e0b94 100644 --- a/litellm/proxy/auth/resolvers/store.py +++ b/litellm/proxy/auth/resolvers/store.py @@ -8,10 +8,13 @@ from pydantic import BaseModel from litellm._internal_context import with_service_target from litellm._logging import verbose_proxy_logger from litellm.proxy._types import UserAPIKeyAuth -from litellm.proxy.auth.auth_checks import ( - _cache_key_object, - _copy_user_api_key_auth_for_cache, - _fetch_key_object_from_db_with_reconnect, +from litellm.proxy.auth.auth_checks import ( # noqa: F401 # legacy module exports + _cache_key_object, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _copy_user_api_key_auth_for_cache, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _fetch_key_object_from_db_with_reconnect, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + cache_key_object, + copy_user_api_key_auth_for_cache, + fetch_key_object_from_db_with_reconnect, get_object_permission, ) from litellm.proxy.auth.auth_method import AuthMethod @@ -81,7 +84,7 @@ class IdentityStore: network: NetworkContext | None = None, ) -> Principal: key: Final = await self._resolve_key(hashed_token) - return self._principal_from_key( + return self.principal_from_key( key, auth_method=auth_method, network=network, @@ -108,12 +111,12 @@ class IdentityStore: cached: Final = await self._cache.async_get_cache(key=hashed_token, model_type=UserAPIKeyAuth) if cached is not None: - return _copy_user_api_key_auth_for_cache(user_api_key_obj=cached) + return copy_user_api_key_auth_for_cache(user_api_key_obj=cached) if self._check_cache_only: raise KeyNotInCacheError(hashed_token) - from_db: Final[BaseModel | None] = await _fetch_key_object_from_db_with_reconnect( + from_db: Final[BaseModel | None] = await fetch_key_object_from_db_with_reconnect( hashed_token=hashed_token, prisma_client=self._prisma, parent_otel_span=self._parent_otel_span, @@ -140,7 +143,7 @@ class IdentityStore: e, ) - await _cache_key_object( + await cache_key_object( hashed_token=hashed_token, user_api_key_obj=key, user_api_key_cache=self._cache, @@ -149,7 +152,7 @@ class IdentityStore: return key @staticmethod - def _principal_from_key( + def principal_from_key( key: UserAPIKeyAuth, *, auth_method: AuthMethod, @@ -192,3 +195,5 @@ class IdentityStore: network=network or NetworkContext(), source_key=key, ) + + _principal_from_key = principal_from_key diff --git a/litellm/proxy/auth/route_checks.py b/litellm/proxy/auth/route_checks.py index 4a4913fd65b..2c1643c7c42 100644 --- a/litellm/proxy/auth/route_checks.py +++ b/litellm/proxy/auth/route_checks.py @@ -15,7 +15,10 @@ from litellm.proxy._types import ( UserAPIKeyAuth, ) -from .auth_checks_organization import _user_is_org_admin +from .auth_checks_organization import ( # noqa: F401 # legacy module exports + _user_is_org_admin, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + user_is_org_admin, +) # Management write routes denied to PROXY_ADMIN_VIEW_ONLY. Adding a new write # endpoint to a management router REQUIRES adding it here too — the surrounding @@ -128,7 +131,7 @@ class RouteChecks: allowed_route in _AUTH_ENFORCED_PASS_THROUGH_ROUTE_GROUPS and RouteChecks.is_auth_enforced_pass_through_route( route=route, - method=RouteChecks._get_request_method(request=request), + method=RouteChecks.get_request_method(request=request), ) ): if RouteChecks.check_passthrough_route_access(route=route, user_api_key_dict=valid_token): @@ -141,7 +144,7 @@ class RouteChecks: # For llm_api_routes, also check registered pass-through endpoints ################################################ if allowed_route == "llm_api_routes": - if route == "/auto_router/session" and RouteChecks._get_request_method(request) == "GET": + if route == "/auto_router/session" and RouteChecks.get_request_method(request) == "GET": return True from litellm.proxy.pass_through_endpoints.pass_through_endpoints import ( @@ -151,7 +154,7 @@ class RouteChecks: if InitPassThroughEndpointHelpers.is_registered_pass_through_route(route=route): if RouteChecks.is_auth_enforced_pass_through_route( route=route, - method=RouteChecks._get_request_method(request=request), + method=RouteChecks.get_request_method(request=request), ): if RouteChecks.check_passthrough_route_access( route=route, user_api_key_dict=valid_token @@ -277,7 +280,7 @@ class RouteChecks: if RouteChecks.is_auth_enforced_pass_through_route( route=route, - method=RouteChecks._get_request_method(request=request), + method=RouteChecks.get_request_method(request=request), ): RouteChecks._require_auth_pass_through_access( route=route, @@ -329,7 +332,7 @@ class RouteChecks: elif ( _user_role == LitellmUserRoles.INTERNAL_USER.value and RouteChecks.check_route_access(route=route, allowed_routes=LiteLLMRoutes.internal_user_routes.value) - or _user_is_org_admin(request_data=request_data, user_object=user_obj) + or user_is_org_admin(request_data=request_data, user_object=user_obj) and RouteChecks.check_route_access(route=route, allowed_routes=LiteLLMRoutes.org_admin_allowed_routes.value) or _user_role == LitellmUserRoles.INTERNAL_USER_VIEW_ONLY.value and RouteChecks.check_route_access( @@ -423,7 +426,7 @@ class RouteChecks: if RouteChecks._route_matches_pattern(route=route, pattern=openai_route): return True # Check for wildcard patterns like "/containers/*" - if RouteChecks._is_wildcard_pattern(pattern=openai_route): + if RouteChecks.is_wildcard_pattern(pattern=openai_route): if RouteChecks.route_matches_wildcard_pattern(route=route, pattern=openai_route): return True @@ -549,12 +552,14 @@ class RouteChecks: return False @staticmethod - def _is_wildcard_pattern(pattern: str) -> bool: + def is_wildcard_pattern(pattern: str) -> bool: """ Check if pattern is a wildcard pattern """ return pattern.endswith("*") + _is_wildcard_pattern = is_wildcard_pattern + @staticmethod def route_matches_wildcard_pattern(route: str, pattern: str) -> bool: """ @@ -635,7 +640,7 @@ class RouteChecks: if any( RouteChecks.route_matches_wildcard_pattern(route=route, pattern=allowed_route) for allowed_route in allowed_routes - if RouteChecks._is_wildcard_pattern(pattern=allowed_route) + if RouteChecks.is_wildcard_pattern(pattern=allowed_route) ): return True @@ -653,7 +658,7 @@ class RouteChecks: return False @staticmethod - def _get_request_method(request: Request | None) -> str | None: + def get_request_method(request: Request | None) -> str | None: if request is None: return None @@ -666,6 +671,8 @@ class RouteChecks: return method.upper() + _get_request_method = get_request_method + @staticmethod def is_auth_enforced_pass_through_route(route: str, method: str | None = None) -> bool: """ @@ -829,7 +836,7 @@ class RouteChecks: return False @staticmethod - def _is_assistants_api_request(request: Request) -> bool: + def is_assistants_api_request(request: Request) -> bool: """ Returns True if `thread` or `assistant` is in the request path @@ -847,6 +854,8 @@ class RouteChecks: return True return False + _is_assistants_api_request = is_assistants_api_request + @staticmethod def is_generate_content_route(route: str) -> bool: """ diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index af487ad19af..dda72928f26 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -41,22 +41,26 @@ from litellm.litellm_core_utils.dd_tracing import tracer from litellm.litellm_core_utils.dot_notation_indexing import get_nested_value from litellm.proxy._types import * from litellm.proxy.agent_endpoints.auth.agent_caller import agent_caller_from_headers -from litellm.proxy.auth.auth_checks import ( +from litellm.proxy.auth.auth_checks import ( # noqa: F401 # legacy module exports ExperimentalUIJWTToken, TeamNotFoundError, - _cache_key_object, - _can_object_call_model, - _check_end_user_budget, - _delete_cache_key_object, - _get_user_role, - _is_model_cost_zero, - _is_user_proxy_admin, - _team_member_max_budget_alert_check, - _virtual_key_max_budget_alert_check, - _virtual_key_max_budget_check, - _virtual_key_soft_budget_check, + _cache_key_object, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _can_object_call_model, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _check_end_user_budget, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _delete_cache_key_object, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _get_user_role, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _is_model_cost_zero, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _is_user_proxy_admin, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _team_member_max_budget_alert_check, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _virtual_key_max_budget_alert_check, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _virtual_key_max_budget_check, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _virtual_key_soft_budget_check, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + cache_key_object, can_key_call_model, + can_object_call_model, + check_end_user_budget, common_checks, + delete_cache_key_object, get_end_user_object, get_jwt_key_mapping_object, get_key_end_user_budget_id, @@ -66,11 +70,18 @@ from litellm.proxy.auth.auth_checks import ( get_team_membership, get_team_object, get_user_object, + get_user_role, + is_model_cost_zero, + is_user_proxy_admin, is_valid_fallback_model, jwt_key_mapping_cache_key, key_model_aliases_for_auth_check, resolve_and_validate_end_user_id, resolve_default_end_user_budget, + team_member_max_budget_alert_check, + virtual_key_max_budget_alert_check, + virtual_key_max_budget_check, + virtual_key_soft_budget_check, ) from litellm.proxy.auth.auth_exception_handler import UserAPIKeyAuthExceptionHandler from litellm.proxy.auth.auth_method import AuthMethod @@ -113,18 +124,25 @@ from litellm.proxy.auth.route_checks import RouteChecks from litellm.proxy.auth.team_grants import team_grants from litellm.proxy.auth.trusted_proxy_utils import get_trusted_proxy_cidrs from litellm.proxy.common_utils.cache_coordinator import EventDrivenCacheCoordinator -from litellm.proxy.common_utils.http_parsing_utils import ( - _read_request_body, - _safe_get_request_headers, - _safe_get_request_query_params, - _safe_set_request_parsed_body, +from litellm.proxy.common_utils.http_parsing_utils import ( # noqa: F401 # legacy module exports + _read_request_body, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _safe_get_request_headers, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _safe_get_request_query_params, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _safe_set_request_parsed_body, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export is_opaque_audio_pass_through_request, populate_request_with_path_params, read_raw_json_body, + read_request_body, rewrite_request_model, + safe_get_request_headers, + safe_get_request_query_params, + safe_set_request_parsed_body, ) from litellm.proxy.common_utils.model_listing_utils import claude_code_requested_group -from litellm.proxy.common_utils.realtime_utils import _realtime_request_body +from litellm.proxy.common_utils.realtime_utils import ( # noqa: F401 # legacy module exports + _realtime_request_body, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + realtime_request_body, +) from litellm.proxy.common_utils.user_api_key_cache import ( UserApiKeyCache, end_user_cache_key, @@ -222,8 +240,8 @@ def _get_model_from_request_context( return get_model_from_request( request_data=request_data, route=route, - request_headers=_safe_get_request_headers(request=request), - request_query_params=_safe_get_request_query_params(request=request), + request_headers=safe_get_request_headers(request=request), + request_query_params=safe_get_request_query_params(request=request), llm_router=llm_router, request=request, team_id=team_id, @@ -477,7 +495,7 @@ async def _check_key_model_budget_with_fallback( llm_router=llm_router, ) if valid_token.team_models: - _can_object_call_model( + can_object_call_model( model=fallback_model, llm_router=llm_router, models=valid_token.team_models, @@ -489,9 +507,9 @@ async def _check_key_model_budget_with_fallback( except ProxyException: raise e request_data["model"] = fallback_model - _safe_set_request_parsed_body(request=request, parsed_body=request_data) - request._json = request_data - request._body = orjson.dumps(request_data) + safe_set_request_parsed_body(request=request, parsed_body=request_data) + request._json = request_data # pyright: ignore[reportPrivateUsage] # Starlette JSON cache + request._body = orjson.dumps(request_data) # pyright: ignore[reportPrivateUsage] # Starlette body cache path_params: Final = request.scope.get("path_params") if isinstance(path_params, dict) and "model" in path_params: path_params["model"] = fallback_model @@ -591,9 +609,9 @@ def _should_route_jwt_to_oauth2_override(token: str, jwt_handler: JWTHandler) -> return False -def _get_bearer_token( +def get_bearer_token( api_key: str, -): +) -> str: if api_key.startswith("Bearer "): # ensure Bearer token passed in api_key = api_key.replace("Bearer ", "") # extract the token elif api_key.startswith("Basic "): @@ -619,6 +637,9 @@ def _get_bearer_token( return api_key +_get_bearer_token: Final = get_bearer_token + + def _apply_budget_limits_to_end_user_params( end_user_params: dict, budget_info: LiteLLM_BudgetTable, @@ -675,10 +696,10 @@ async def user_api_key_auth_websocket_for_model(websocket: WebSocket, model: str synthetic_scope[key] = ws_scope[key] request: Final = Request(scope=synthetic_scope) - request._url = websocket.url + request._url = websocket.url # pyright: ignore[reportPrivateUsage] # Starlette WebSocket URL storage async def return_body(): - return _realtime_request_body(model) + return realtime_request_body(model) request.body = return_body @@ -739,7 +760,7 @@ def update_valid_token_with_end_user_params(valid_token: UserAPIKeyAuth, end_use _global_spend_coordinator: Final = EventDrivenCacheCoordinator(log_prefix="[GLOBAL SPEND]") -async def _fetch_global_spend_with_event_coordination( +async def fetch_global_spend_with_event_coordination( cache_key: str, user_api_key_cache: UserApiKeyCache, prisma_client: PrismaClient, @@ -766,6 +787,9 @@ async def _fetch_global_spend_with_event_coordination( ) +_fetch_global_spend_with_event_coordination: Final = fetch_global_spend_with_event_coordination + + async def get_global_proxy_spend( litellm_proxy_admin_name: str, user_api_key_cache: UserApiKeyCache, @@ -777,10 +801,12 @@ async def get_global_proxy_spend( if litellm.max_budget > 0 and prisma_client is not None: # user set proxy max budget # Use event-driven coordination to prevent cache stampede cache_key: Final = GLOBAL_PROXY_SPEND_CACHE_KEY - global_proxy_spend = await _fetch_global_spend_with_event_coordination( - cache_key=cache_key, - user_api_key_cache=user_api_key_cache, - prisma_client=prisma_client, + global_proxy_spend = ( # rebind-ok: pre-existing rebinding on a rename-only line + await fetch_global_spend_with_event_coordination( + cache_key=cache_key, + user_api_key_cache=user_api_key_cache, + prisma_client=prisma_client, + ) ) if global_proxy_spend is not None: user_info: Final = CallInfo( @@ -824,7 +850,7 @@ def get_api_key( """ from litellm.proxy.auth.route_checks import RouteChecks from litellm.proxy.common_utils.http_parsing_utils import ( - _safe_get_request_query_params, + safe_get_request_query_params, ) api_key = api_key @@ -834,7 +860,7 @@ def get_api_key( api_key = _get_bearer_token_or_received_api_key(custom_litellm_key_header) elif isinstance(api_key, str) and len(api_key) > 0: passed_in_key = api_key - api_key = _get_bearer_token(api_key=api_key) + api_key = get_bearer_token(api_key=api_key) elif isinstance(azure_api_key_header, str): passed_in_key = azure_api_key_header api_key = azure_api_key_header @@ -850,9 +876,9 @@ def get_api_key( elif ( RouteChecks.is_generate_content_route(route=route) and request is not None - and _safe_get_request_query_params(request).get("key") + and safe_get_request_query_params(request).get("key") ): - google_auth_key: Final[str] = _safe_get_request_query_params(request).get("key") or "" + google_auth_key: Final[str] = safe_get_request_query_params(request).get("key") or "" passed_in_key = google_auth_key api_key = google_auth_key elif pass_through_endpoints is not None: @@ -1388,7 +1414,7 @@ def _ensure_parent_otel_span_on_request_state(request: Request) -> None: return parent_otel_span: Final = open_telemetry_logger.create_litellm_proxy_request_started_span( start_time=start_time, - headers=_safe_get_request_headers(request), + headers=safe_get_request_headers(request), ) # Under V2 the FastAPI instrumentor stamps http.route / url.path on the server # span; only the legacy logger needs these set explicitly. @@ -1413,12 +1439,12 @@ async def _read_request_body_deferring_parse_failure( """ if is_opaque_audio_pass_through_request( route=get_request_route(request=request), - content_type=_safe_get_request_headers(request=request).get("content-type", ""), + content_type=safe_get_request_headers(request=request).get("content-type", ""), ): - _safe_set_request_parsed_body(request=request, parsed_body={}) + safe_set_request_parsed_body(request=request, parsed_body={}) return {}, None try: - parsed_body: Final = await _read_request_body(request=request) + parsed_body: Final = await read_request_body(request=request) except ProxyException as parse_exception: return {}, parse_exception return populate_request_with_path_params(request_data=parsed_body, request=request), None @@ -1481,7 +1507,7 @@ async def _refresh_session_token_grants( { **valid_token.model_dump(exclude_none=True), **team_grants(team_object, team_membership, user_object.user_id), - "user_role": _get_user_role(user_object), + "user_role": get_user_role(user_object), "models": () if team_object is not None else user_models(user_object), } ) @@ -1515,7 +1541,7 @@ async def _resolve_object_permission_for_unresolvable_team( ) -async def _user_api_key_auth_builder( +async def user_api_key_auth_builder( request: Request, api_key: str, azure_api_key_header: str, @@ -1765,8 +1791,8 @@ async def _user_api_key_auth_builder( user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, parent_otel_span=parent_otel_span, - request_headers=_safe_get_request_headers(request), - request_method=RouteChecks._get_request_method(request=request), + request_headers=safe_get_request_headers(request), + request_method=RouteChecks.get_request_method(request=request), ) is_proxy_admin: Final = result["is_proxy_admin"] @@ -1862,9 +1888,11 @@ async def _user_api_key_auth_builder( ) skip_budget_checks = False if model is not None and llm_router is not None: - from litellm.proxy.auth.auth_checks import _is_model_cost_zero + from litellm.proxy.auth.auth_checks import is_model_cost_zero - skip_budget_checks = _is_model_cost_zero(model=model, llm_router=llm_router) + skip_budget_checks = ( # rebind-ok: pre-existing rebinding on a rename-only line + is_model_cost_zero(model=model, llm_router=llm_router) + ) if skip_budget_checks: verbose_proxy_logger.info("Skipping all budget checks for zero-cost model: %s", model) @@ -1950,7 +1978,7 @@ async def _user_api_key_auth_builder( _end_user_object = None end_user_params: Final = {} - raw_end_user_id: Final = get_end_user_id_from_request_body(request_data, _safe_get_request_headers(request)) + raw_end_user_id: Final = get_end_user_id_from_request_body(request_data, safe_get_request_headers(request)) end_user_id = await resolve_and_validate_end_user_id( raw_end_user_id=raw_end_user_id, prisma_client=prisma_client, @@ -2070,7 +2098,7 @@ async def _user_api_key_auth_builder( if expiry_time.tzinfo is None or expiry_time.tzinfo.utcoffset(expiry_time) is None: expiry_time = expiry_time.replace(tzinfo=timezone.utc) if expiry_time < current_time: - await _delete_cache_key_object( + await delete_cache_key_object( hashed_token=hash_token(api_key), user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, @@ -2151,7 +2179,7 @@ async def _user_api_key_auth_builder( start_time=start_time, ) asyncio.create_task( - _cache_key_object( + cache_key_object( hashed_token=hash_token(master_key), user_api_key_obj=_user_api_key_obj, user_api_key_cache=user_api_key_cache, @@ -2248,7 +2276,7 @@ async def _user_api_key_auth_builder( _end_user_object=_end_user_object, ) except Exception as e: - return await UserAPIKeyAuthExceptionHandler._handle_authentication_error( + return await UserAPIKeyAuthExceptionHandler.handle_authentication_error( e=e, request=request, request_data=request_data, @@ -2259,6 +2287,9 @@ async def _user_api_key_auth_builder( ) +_user_api_key_auth_builder: Final = user_api_key_auth_builder + + async def validate_resolved_virtual_key( # noqa: C901 # Preserve ordering of existing shared authorization checks request: Request, request_data: dict[str, object], @@ -2303,7 +2334,7 @@ async def validate_resolved_virtual_key( # noqa: C901 # Preserve ordering of e ## base case ## key is disabled if valid_token.blocked is True: raise Exception("Key is blocked. Update via `/key/unblock` if you're an admin.") - await _enforce_key_and_fallback_model_access( + await enforce_key_and_fallback_model_access( valid_token=valid_token, request_data=request_data, route=route, @@ -2359,9 +2390,11 @@ async def validate_resolved_virtual_key( # noqa: C901 # Preserve ordering of e ) skip_budget_checks = False if model is not None and llm_router is not None: - from litellm.proxy.auth.auth_checks import _is_model_cost_zero + from litellm.proxy.auth.auth_checks import is_model_cost_zero - skip_budget_checks = _is_model_cost_zero(model=model, llm_router=llm_router) + skip_budget_checks = is_model_cost_zero( # rebind-ok: pre-existing rebinding on a rename-only line + model=model, llm_router=llm_router + ) if skip_budget_checks: verbose_proxy_logger.info("Skipping all budget checks for zero-cost model: %s", model) @@ -2412,7 +2445,7 @@ async def validate_resolved_virtual_key( # noqa: C901 # Preserve ordering of e if team_member_spend >= team_member_budget: # common_checks sends this alert on requests that get past here, so only the # request rejected here sends it from the builder. - _team_member_max_budget_alert_check( + team_member_max_budget_alert_check( team_id=_team_id, team_alias=valid_token.team_alias, team_metadata=valid_token.team_metadata, @@ -2461,7 +2494,7 @@ async def validate_resolved_virtual_key( # noqa: C901 # Preserve ordering of e # Check 4. Max Budget Alert Check (runs before budget enforcement # so multi-threshold 100% alerts fire on the request that crosses # max_budget, before BudgetExceededError is raised below) - await _virtual_key_max_budget_alert_check( + await virtual_key_max_budget_alert_check( valid_token=valid_token, proxy_logging_obj=proxy_logging_obj, user_obj=user_obj, @@ -2469,14 +2502,14 @@ async def validate_resolved_virtual_key( # noqa: C901 # Preserve ordering of e # Check 5. Token Spend is under budget if RouteChecks.is_llm_api_route(route=route): - await _virtual_key_max_budget_check( + await virtual_key_max_budget_check( valid_token=valid_token, proxy_logging_obj=proxy_logging_obj, user_obj=user_obj, ) # Check 6. Soft Budget Check - await _virtual_key_soft_budget_check( + await virtual_key_soft_budget_check( valid_token=valid_token, proxy_logging_obj=proxy_logging_obj, user_obj=user_obj, @@ -2613,7 +2646,7 @@ async def validate_resolved_virtual_key( # noqa: C901 # Preserve ordering of e if litellm.max_budget > 0 and prisma_client is not None: # user set proxy max budget cache_key: Final = GLOBAL_PROXY_SPEND_CACHE_KEY with tracer.trace("litellm.proxy.auth.get_global_proxy_spend"): - global_proxy_spend = await _fetch_global_spend_with_event_coordination( + global_proxy_spend = await fetch_global_spend_with_event_coordination( # rebind-ok: pre-existing rebinding on a rename-only line cache_key=cache_key, user_api_key_cache=user_api_key_cache, prisma_client=prisma_client, @@ -2803,7 +2836,7 @@ def is_no_auth_dev_mode(master_key: str | None, general_settings: Mapping[str, o @tracer.wrap() -async def _run_centralized_common_checks( +async def run_centralized_common_checks( user_api_key_auth_obj: UserAPIKeyAuth, request: Request, request_data: dict[str, object], @@ -2884,7 +2917,7 @@ async def _run_centralized_common_checks( key_end_user_budget_id: Final = get_key_end_user_budget_id(user_api_key_auth_obj.metadata) end_user_id = user_api_key_auth_obj.end_user_id if end_user_id is None: - raw_end_user_id: Final = get_end_user_id_from_request_body(request_data, _safe_get_request_headers(request)) + raw_end_user_id: Final = get_end_user_id_from_request_body(request_data, safe_get_request_headers(request)) end_user_id = await resolve_and_validate_end_user_id( raw_end_user_id=raw_end_user_id, prisma_client=prisma_client, @@ -3164,6 +3197,9 @@ async def _run_centralized_common_checks( release_spend_counter_batch() +_run_centralized_common_checks: Final = run_centralized_common_checks + + async def _noop_none() -> None: """Sentinel coroutine for asyncio.gather when a fetch is unnecessary (e.g. token has no team_id). Keeps the result tuple positional.""" @@ -3271,7 +3307,7 @@ def _should_skip_budget_checks( team_id=team_id, ) if model is not None and llm_router is not None: - return _is_model_cost_zero(model=model, llm_router=llm_router) + return is_model_cost_zero(model=model, llm_router=llm_router) return False @@ -3290,7 +3326,7 @@ def _resolve_request_principal(request: Request, valid_token: UserAPIKeyAuth) -> TrustedProxyConfig(use_forwarded_for=bool(cidrs), trusted_proxy_cidrs=cidrs), ) auth_method: Final = AuthMethod.BEARER_JWT if valid_token.jwt_claims else AuthMethod.API_KEY - return IdentityStore._principal_from_key( + return IdentityStore.principal_from_key( valid_token, auth_method=auth_method, network=network, @@ -3362,7 +3398,7 @@ async def _authorize_authenticated_request( billable=request_data.get("method") in (None, "message/send", "message/stream", "SendMessage", "SendStreamingMessage"), ) - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=user_api_key_auth_obj, request=request, request_data=authorized_data, @@ -3370,7 +3406,7 @@ async def _authorize_authenticated_request( force_virtual_key_checks=force_virtual_key_checks, ) except Exception as e: - return await UserAPIKeyAuthExceptionHandler._handle_authentication_error( + return await UserAPIKeyAuthExceptionHandler.handle_authentication_error( e=e, request=request, request_data=request_data, @@ -3393,7 +3429,7 @@ async def _authorize_authenticated_request( user_api_key_cache, ) - raw_end_user_id: Final = get_end_user_id_from_request_body(request_data, _safe_get_request_headers(request)) + raw_end_user_id: Final = get_end_user_id_from_request_body(request_data, safe_get_request_headers(request)) if raw_end_user_id is not None: resolved_end_user_id: Final = await resolve_and_validate_end_user_id( raw_end_user_id=raw_end_user_id, @@ -3475,7 +3511,7 @@ def _seed_request_destinations(user_api_key_dict: UserAPIKeyAuth, request: Reque set_request_destinations( deliverable_destinations( - resolve_tenant_otel_destinations(user_api_key_dict, _safe_get_request_headers(request)), + resolve_tenant_otel_destinations(user_api_key_dict, safe_get_request_headers(request)), fan_out_provider(), ) ) @@ -3518,7 +3554,7 @@ async def user_api_key_auth( spend_counter_batch_scope(_spend_counter_redis_cache()), ): try: - user_api_key_auth_obj: Final = await _user_api_key_auth_builder( + user_api_key_auth_obj: Final = await user_api_key_auth_builder( request=request, api_key=api_key, azure_api_key_header=azure_api_key_header, @@ -3537,7 +3573,7 @@ async def user_api_key_auth( raise user_api_key_auth_obj.budget_reservation = None user_api_key_auth_obj.agent_caller = agent_caller_from_headers( - _safe_get_request_headers(request), user_api_key_auth_obj + safe_get_request_headers(request), user_api_key_auth_obj ) _seed_request_destinations(user_api_key_auth_obj, request) @@ -3611,7 +3647,7 @@ async def _return_user_api_key_auth_obj( ) ) - retrieved_user_role: Final = user_role or _get_user_role(user_obj=user_obj) or LitellmUserRoles.INTERNAL_USER + retrieved_user_role: Final = user_role or get_user_role(user_obj=user_obj) or LitellmUserRoles.INTERNAL_USER user_api_key_kwargs: Final = { "api_key": api_key, @@ -3628,7 +3664,7 @@ async def _return_user_api_key_auth_obj( user_max_budget=getattr(user_obj, "max_budget", None), user_model_max_budget=getattr(user_obj, "model_max_budget", None), ) - if user_obj is not None and _is_user_proxy_admin(user_obj=user_obj): + if user_obj is not None and is_user_proxy_admin(user_obj=user_obj): user_api_key_kwargs.update( user_role=LitellmUserRoles.PROXY_ADMIN, ) @@ -3659,7 +3695,7 @@ def get_api_key_from_custom_header(request: Request, custom_litellm_key_header_n ) custom_api_key: Final = _headers.get(custom_litellm_key_header_name) if custom_api_key: - api_key = _get_bearer_token(api_key=custom_api_key) + api_key = get_bearer_token(api_key=custom_api_key) # rebind-ok: pre-existing rebinding on a rename-only line verbose_proxy_logger.debug( "Found custom API key using header: %s, setting api_key=%s", custom_litellm_key_header_name, @@ -3756,7 +3792,7 @@ async def _lookup_end_user_and_apply_budget( return valid_token, end_user_object -async def _enforce_key_and_fallback_model_access( +async def enforce_key_and_fallback_model_access( *, valid_token: UserAPIKeyAuth, request_data: dict, @@ -3816,6 +3852,9 @@ async def _enforce_key_and_fallback_model_access( ) +_enforce_key_and_fallback_model_access: Final = enforce_key_and_fallback_model_access + + async def _run_post_custom_auth_checks( valid_token: UserAPIKeyAuth, request: Request, @@ -3849,7 +3888,7 @@ async def _run_post_custom_auth_checks( # custom_auth_run_common_checks is set. Enforce it here on that path # so an over-budget end user can't keep making requests. if end_user_object is not None and not general_settings.get("custom_auth_run_common_checks", False): - await _check_end_user_budget(end_user_obj=end_user_object, route=route) + await check_end_user_budget(end_user_obj=end_user_object, route=route) # 2. Check token expiry if valid_token.expires is not None: @@ -3869,7 +3908,7 @@ async def _run_post_custom_auth_checks( ) if general_settings.get("custom_auth_run_common_checks", False): - await _enforce_key_and_fallback_model_access( + await enforce_key_and_fallback_model_access( valid_token=valid_token, request_data=request_data, route=route, @@ -3892,7 +3931,7 @@ async def _run_post_custom_auth_checks( # every budget check for these; this path did not, so the same request could # be refused under custom auth and served under the other two. skip_budget_checks: Final = ( - _is_model_cost_zero(model=current_model, llm_router=llm_router) + is_model_cost_zero(model=current_model, llm_router=llm_router) if current_model is not None and llm_router is not None else False ) diff --git a/litellm/proxy/batches_endpoints/endpoints.py b/litellm/proxy/batches_endpoints/endpoints.py index 485ea137081..880e3fa2072 100644 --- a/litellm/proxy/batches_endpoints/endpoints.py +++ b/litellm/proxy/batches_endpoints/endpoints.py @@ -36,14 +36,17 @@ from litellm.proxy.common_request_processing import ( request_litellm_call_id, ) from litellm.proxy.common_utils.callback_utils import sanitize_openai_provider_metadata -from litellm.proxy.common_utils.http_parsing_utils import _read_request_body +from litellm.proxy.common_utils.http_parsing_utils import ( # noqa: F401 # legacy module exports + _read_request_body, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + read_request_body, +) from litellm.proxy.common_utils.openai_endpoint_utils import ( get_custom_llm_provider_from_request_headers, get_custom_llm_provider_from_request_query, ) -from litellm.proxy.openai_files_endpoints.common_utils import ( +from litellm.proxy.openai_files_endpoints.common_utils import ( # noqa: F401 # legacy module exports BATCH_CREATE_HIDDEN_PARAM, - _is_base64_encoded_unified_file_id, + _is_base64_encoded_unified_file_id, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export add_deployment_model_info, add_internal_model_credentials, apply_team_provider_credentials, @@ -59,6 +62,7 @@ from litellm.proxy.openai_files_endpoints.common_utils import ( get_model_id_from_unified_batch_id, get_models_from_unified_file_id, get_original_file_id, + is_base64_encoded_unified_file_id, is_litellm_executed_batch, prepare_data_with_credentials, update_batch_in_database, @@ -269,7 +273,7 @@ async def create_batch( data: dict = {} try: - data = await _read_request_body(request=request) + data = await read_request_body(request=request) # rebind-ok: pre-existing rebinding on a rename-only line verbose_proxy_logger.debug( "Request received by LiteLLM:\n%s", json.dumps(data, indent=4), @@ -341,7 +345,9 @@ async def create_batch( model_from_file_id = None if input_file_id: model_from_file_id = decode_model_from_file_id(input_file_id) - unified_file_id = _is_base64_encoded_unified_file_id(input_file_id) + unified_file_id = ( # rebind-ok: pre-existing rebinding on a rename-only line + is_base64_encoded_unified_file_id(input_file_id) + ) # SCENARIO 1: File ID is encoded with model info if model_from_file_id is not None and input_file_id: @@ -587,7 +593,7 @@ async def retrieve_batch( ) data = cast(dict, _retrieve_batch_request) - unified_batch_id: Final = _is_base64_encoded_unified_file_id(batch_id) + unified_batch_id: Final = is_base64_encoded_unified_file_id(batch_id) base_llm_response_processor: Final = ProxyBaseLLMRequestProcessing(data=data) ( @@ -891,7 +897,7 @@ async def list_batches( ) # Include original request and headers in the data - data = await _read_request_body(request=request) + data = await read_request_body(request=request) # rebind-ok: pre-existing rebinding on a rename-only line base_llm_response_processor: Final = ProxyBaseLLMRequestProcessing(data=data) ( data, @@ -1084,7 +1090,7 @@ async def cancel_batch( ) data = cast(dict, _cancel_batch_request) - unified_batch_id: Final = _is_base64_encoded_unified_file_id(batch_id) + unified_batch_id: Final = is_base64_encoded_unified_file_id(batch_id) base_llm_response_processor: Final = ProxyBaseLLMRequestProcessing(data=data) ( diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index fcc50e57fd5..317203cc2d5 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -28,6 +28,7 @@ from fastapi import HTTPException, Request, status from fastapi.responses import JSONResponse, Response, StreamingResponse from pydantic import BaseModel, TypeAdapter, ValidationError from starlette.types import Receive, Scope, Send +from typing_extensions import Never import litellm from litellm._logging import redact_internal_details_from_client_message, verbose_proxy_logger @@ -118,7 +119,11 @@ from litellm.proxy.native_compaction import with_proxy_compaction_executor from litellm.proxy.route_llm_request import ( route_request, ) -from litellm.proxy.utils import ProxyLogging, _check_and_merge_model_level_guardrails +from litellm.proxy.utils import ( # noqa: F401 # legacy module exports + ProxyLogging, + _check_and_merge_model_level_guardrails, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + check_and_merge_model_level_guardrails, +) from litellm.router import Router from litellm.router_utils.add_retry_fallback_headers import get_hidden_params_dict from litellm.router_utils.common_utils import resolve_model_group_alias @@ -290,13 +295,16 @@ def resolve_litellm_call_id(client_call_id: str | None) -> str: return str(uuid.uuid4()) -def _should_return_raw_model_name(request_data: dict[str, object]) -> bool: +def should_return_raw_model_name(request_data: dict[str, object]) -> bool: return any( isinstance(metadata, dict) and metadata.get(RETURN_RAW_MODEL_NAME_METADATA_KEY) is True for metadata in (request_data.get("metadata"), request_data.get("litellm_metadata")) ) +_should_return_raw_model_name: Final = should_return_raw_model_name + + def _apply_client_disconnect_metadata(target_metadata: dict[str, object] | None) -> None: if target_metadata is None: return @@ -1302,7 +1310,7 @@ async def open_sse_before_first_byte( ) -def _is_azure_model_router_request(model: str, hidden_params: Mapping[str, object] | None = None) -> bool: +def is_azure_model_router_request(model: str, hidden_params: Mapping[str, object] | None = None) -> bool: """ Check if a request went down the Azure Model Router route. @@ -1323,6 +1331,9 @@ def _is_azure_model_router_request(model: str, hidden_params: Mapping[str, objec return AzureFoundryModelInfo.is_model_router_call(model=model, hidden_params=hidden_params) +_is_azure_model_router_request: Final = is_azure_model_router_request + + def _override_openai_response_model( *, response_obj: object, @@ -1388,7 +1399,7 @@ def _override_openai_response_model( return # Check if this is an Azure Model Router request - if so, preserve the actual model used - if _is_azure_model_router_request(requested_model, hidden_params): + if is_azure_model_router_request(requested_model, hidden_params): verbose_proxy_logger.debug( "%s: Azure Model Router detected - preserving actual model used from response instead of overriding to router model.", log_context, @@ -2074,9 +2085,9 @@ class ProxyBaseLLMRequestProcessing: # Store queue time in metadata after add_litellm_data_to_request to ensure it's preserved if queue_time_seconds is not None: - from litellm.proxy.litellm_pre_call_utils import _get_metadata_variable_name + from litellm.proxy.litellm_pre_call_utils import get_metadata_variable_name - _metadata_variable_name: Final = _get_metadata_variable_name(request) + _metadata_variable_name: Final = get_metadata_variable_name(request) if _metadata_variable_name not in self.data: self.data[_metadata_variable_name] = {} if not isinstance(self.data[_metadata_variable_name], dict): @@ -2200,11 +2211,11 @@ class ProxyBaseLLMRequestProcessing: merged_for_requested: Final = ( self.data if rate_limited_model is None - else _check_and_merge_model_level_guardrails( + else check_and_merge_model_level_guardrails( data=self.data, llm_router=llm_router, trust_client_model_info=False, model_alias=rate_limited_model ) ) - self.data = _check_and_merge_model_level_guardrails( + self.data = check_and_merge_model_level_guardrails( data=merged_for_requested, llm_router=llm_router, trust_client_model_info=False, @@ -2803,7 +2814,7 @@ class ProxyBaseLLMRequestProcessing: proxy_logging_obj=proxy_logging_obj, request=request, restamp_model=( - None if _should_return_raw_model_name(self.data) else requested_model_from_client + None if should_return_raw_model_name(self.data) else requested_model_from_client ), ) selected_data_generator = wrap_sse_stream_with_keepalive_pings( @@ -2942,7 +2953,7 @@ class ProxyBaseLLMRequestProcessing: response_obj=response, requested_model=requested_model_from_client, log_context=f"litellm_call_id={logging_obj.litellm_call_id}", - return_raw_model_name=_should_return_raw_model_name(self.data), + return_raw_model_name=should_return_raw_model_name(self.data), ) fastapi_response.headers.update( @@ -3235,9 +3246,9 @@ class ProxyBaseLLMRequestProcessing: because should_run_guardrail treats it as matching every hook. """ from litellm.proxy.proxy_server import llm_router - from litellm.proxy.utils import _check_and_merge_model_level_guardrails + from litellm.proxy.utils import check_and_merge_model_level_guardrails - guardrail_data: Final = _check_and_merge_model_level_guardrails(data=self.data, llm_router=llm_router) + guardrail_data: Final = check_and_merge_model_level_guardrails(data=self.data, llm_router=llm_router) for cb in litellm.callbacks: if not isinstance(cb, CustomGuardrail): continue @@ -3558,11 +3569,13 @@ class ProxyBaseLLMRequestProcessing: try: from litellm.proxy.proxy_server import llm_router as _global_llm_router from litellm.proxy.utils import ( - _check_and_merge_model_level_guardrails, + check_and_merge_model_level_guardrails, stream_gated_guardrail_names, ) - guardrail_data = _check_and_merge_model_level_guardrails(data=captured_data, llm_router=_global_llm_router) + guardrail_data: Final = check_and_merge_model_level_guardrails( + data=captured_data, llm_router=_global_llm_router + ) stream_gated: Final = stream_gated_guardrail_names(captured_data, captured_user_api_key_dict) for cb in litellm.callbacks: if not isinstance(cb, CustomGuardrail): @@ -3636,13 +3649,13 @@ class ProxyBaseLLMRequestProcessing: if isinstance(e, RouterRateLimitError) and e.cooldown_time > 0: headers["retry-after"] = str(math.ceil(e.cooldown_time)) - async def _handle_llm_api_exception( + async def handle_llm_api_exception( self, e: Exception, user_api_key_dict: UserAPIKeyAuth, proxy_logging_obj: ProxyLogging, version: str | None = None, - ): + ) -> Never: """Raises ProxyException (OpenAI API compatible) if an exception is raised""" log_llm_api_exception(e, self.litellm_call_id) # Allow callbacks to transform the error response @@ -3787,6 +3800,8 @@ class ProxyBaseLLMRequestProcessing: headers=safe_headers, ) + _handle_llm_api_exception = handle_llm_api_exception + ######################################################### # Proxy Level Streaming Data Generator ######################################################### @@ -3820,7 +3835,7 @@ class ProxyBaseLLMRequestProcessing: return serialize @staticmethod - async def _finalize_streaming_generator_cleanup( + async def finalize_streaming_generator_cleanup( request: Request | None, request_data: dict, response: Any, @@ -3840,7 +3855,7 @@ class ProxyBaseLLMRequestProcessing: ) if recorded_client_disconnect: deferred_stream_logging_armed: Final = _deferred_stream_logging_is_armed(request_data) - ProxyLogging._fire_deferred_stream_logging(request_data) + ProxyLogging.fire_deferred_stream_logging(request_data) # A disconnect-time success event (the deferred-guardrail flush # above, or the partial-spend billing below) releases the # request's max_parallel_requests slot through the limiter's @@ -3858,7 +3873,7 @@ class ProxyBaseLLMRequestProcessing: and proxy_logging_obj is not None and user_api_key_dict is not None ): - await proxy_logging_obj._arelease_max_parallel_requests_on_disconnect(user_api_key_dict) + await proxy_logging_obj.arelease_max_parallel_requests_on_disconnect(user_api_key_dict) if hasattr(response, "aclose"): try: @@ -3877,6 +3892,8 @@ class ProxyBaseLLMRequestProcessing: ): await logging_obj.invalidate_baseline_cache_estimate("incomplete_response", completed=True) + _finalize_streaming_generator_cleanup = finalize_streaming_generator_cleanup + @staticmethod async def async_streaming_data_generator( response: object, @@ -3913,7 +3930,7 @@ class ProxyBaseLLMRequestProcessing: # consumed, and cost injection is a no-op -- so the per-chunk coroutine # await, response-string materialization, and cost-injection call are # pure overhead on the streaming hot path (the default config). - caps: Final = ProxyLogging._callback_capabilities() + caps: Final = ProxyLogging.callback_capabilities() cost_injection_enabled: Final = bool(getattr(litellm, "include_cost_in_streaming_usage", False)) fast_path = not caps.has_streaming_chunk_override and not caps.has_guardrail and not cost_injection_enabled debug_enabled: Final = verbose_proxy_logger.isEnabledFor(logging.DEBUG) @@ -3956,7 +3973,7 @@ class ProxyBaseLLMRequestProcessing: str_so_far += str(chunk.get("content", "")) model_name = request_data.get("model", "") - chunk = ProxyBaseLLMRequestProcessing._process_chunk_with_cost_injection( + chunk = ProxyBaseLLMRequestProcessing.process_chunk_with_cost_injection( chunk, model_name, request_data.get("litellm_logging_obj") ) @@ -4022,7 +4039,7 @@ class ProxyBaseLLMRequestProcessing: seal: Final = "" if seal_open_frame is None else seal_open_frame(recent_tail) yield seal + error_frame if seal else error_frame finally: - await ProxyBaseLLMRequestProcessing._finalize_streaming_generator_cleanup( + await ProxyBaseLLMRequestProcessing.finalize_streaming_generator_cleanup( request=request, request_data=request_data, response=response, @@ -4071,18 +4088,18 @@ class ProxyBaseLLMRequestProcessing: @overload @staticmethod - def _process_chunk_with_cost_injection( + def process_chunk_with_cost_injection( chunk: bytes, model_name: str, litellm_logging_obj: LiteLLMLoggingObj | None = None ) -> bytes: ... @overload @staticmethod - def _process_chunk_with_cost_injection( + def process_chunk_with_cost_injection( chunk: object, model_name: str, litellm_logging_obj: LiteLLMLoggingObj | None = None ) -> object: ... @staticmethod - def _process_chunk_with_cost_injection( + def process_chunk_with_cost_injection( chunk: object, model_name: str, litellm_logging_obj: LiteLLMLoggingObj | None = None ) -> object: """ @@ -4131,6 +4148,8 @@ class ProxyBaseLLMRequestProcessing: return chunk + _process_chunk_with_cost_injection = process_chunk_with_cost_injection + @staticmethod def _inject_cost_into_sse_frame_str( frame_str: str, model_name: str, litellm_logging_obj: LiteLLMLoggingObj | None = None diff --git a/litellm/proxy/common_utils/auth_cache_invalidation_pubsub.py b/litellm/proxy/common_utils/auth_cache_invalidation_pubsub.py index 519e604783e..2d874f7b473 100644 --- a/litellm/proxy/common_utils/auth_cache_invalidation_pubsub.py +++ b/litellm/proxy/common_utils/auth_cache_invalidation_pubsub.py @@ -6,10 +6,12 @@ from typing import TYPE_CHECKING, Final from litellm._internal_context import with_service_target from litellm._logging import verbose_proxy_logger -from litellm.proxy.common_utils.config_sync_pubsub import ( - _ConfigSyncPubSub, - _pubsub_capable_client, +from litellm.proxy.common_utils.config_sync_pubsub import ( # noqa: F401 # legacy module exports + ConfigSyncPubSub, + _ConfigSyncPubSub, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _pubsub_capable_client, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export coordination_redis_cache, + pubsub_capable_client, ) from litellm.proxy.common_utils.user_api_key_cache import AUTH_OBJECTS_TARGET @@ -75,7 +77,7 @@ def _message_from_data(data: object) -> _CacheInvalidationMessage | None: async def _publish_to_redis(redis_cache: "RedisCache", cache_key: str, message: str) -> None: try: - client: Final = _pubsub_capable_client(redis_cache) + client: Final = pubsub_capable_client(redis_cache) async with _in_flight_publishes: await client.publish(auth_cache_invalidation_channel(redis_cache), message) except Exception as e: # noqa: BLE001 # best-effort publish; mutations must never fail on redis errors @@ -181,7 +183,7 @@ class AuthCacheInvalidationSubscriber: backoff_seconds = _BACKOFF_INITIAL_SECONDS # rebind-ok: exponential backoff accumulator across reconnects while True: try: - client = _pubsub_capable_client(self._redis_cache) + client = pubsub_capable_client(self._redis_cache) pubsub = client.pubsub() try: await pubsub.subscribe(auth_cache_invalidation_channel(self._redis_cache)) @@ -200,7 +202,7 @@ class AuthCacheInvalidationSubscriber: await asyncio.sleep(backoff_seconds) backoff_seconds = min(backoff_seconds * 2, _BACKOFF_MAX_SECONDS) - async def _consume(self, pubsub: _ConfigSyncPubSub) -> None: + async def _consume(self, pubsub: ConfigSyncPubSub) -> None: while True: message = await pubsub.get_message(ignore_subscribe_messages=True, timeout=_POLL_TIMEOUT_SECONDS) if message is None: @@ -222,7 +224,7 @@ class AuthCacheInvalidationSubscriber: additional_cache.delete_cache(parsed.cache_key) @staticmethod - async def _close_pubsub(pubsub: _ConfigSyncPubSub) -> None: + async def _close_pubsub(pubsub: ConfigSyncPubSub) -> None: try: await pubsub.aclose() except Exception as e: # noqa: BLE001 # best-effort close of a possibly-broken connection diff --git a/litellm/proxy/common_utils/cache_aware_routing.py b/litellm/proxy/common_utils/cache_aware_routing.py index d0685218c87..d5ae9aa9f66 100644 --- a/litellm/proxy/common_utils/cache_aware_routing.py +++ b/litellm/proxy/common_utils/cache_aware_routing.py @@ -243,7 +243,7 @@ async def _choose_cached_model( from litellm.proxy import proxy_server from litellm.proxy.common_utils.prompt_cache_prediction import has_request_transforms from litellm.proxy.hooks.parallel_request_limiter_v3 import ( - _PROXY_MaxParallelRequestsHandler_v3, # pyright: ignore[reportPrivateUsage] # use the configured proxy limiter's shared capacity owner + PROXY_MaxParallelRequestsHandler_v3, # pyright: ignore[reportPrivateUsage] # use the configured proxy limiter's shared capacity owner ) from litellm.router_strategy.complexity_router.context_compaction import compaction_pending @@ -289,7 +289,7 @@ async def _choose_cached_model( if body is None: return None limiter: Final = proxy_server.proxy_logging_obj.get_proxy_hook("parallel_request_limiter") - if not isinstance(limiter, _PROXY_MaxParallelRequestsHandler_v3): + if not isinstance(limiter, PROXY_MaxParallelRequestsHandler_v3): return None def counter_for_model(model_name: str) -> TokenCounter: diff --git a/litellm/proxy/common_utils/callback_utils.py b/litellm/proxy/common_utils/callback_utils.py index 9ab39fa9303..4fa7709e5c7 100644 --- a/litellm/proxy/common_utils/callback_utils.py +++ b/litellm/proxy/common_utils/callback_utils.py @@ -181,7 +181,7 @@ def initialize_callbacks_on_proxy( imported_list.append(callback) elif isinstance(callback, str) and callback == "presidio": from litellm.proxy.guardrails.guardrail_hooks.presidio import ( - _OPTIONAL_PresidioPIIMasking, + OPTIONAL_PresidioPIIMasking, ) presidio_logging_only: bool | None = litellm_settings.get("presidio_logging_only", None) @@ -196,7 +196,7 @@ def initialize_callbacks_on_proxy( "logging_only": presidio_logging_only, **_presidio_params, } - pii_masking_object = _OPTIONAL_PresidioPIIMasking(**params) + pii_masking_object = OPTIONAL_PresidioPIIMasking(**params) imported_list.append(pii_masking_object) elif isinstance(callback, str) and callback == "llamaguard_moderations": try: @@ -324,7 +324,7 @@ def initialize_callbacks_on_proxy( imported_list.append(banned_keywords_obj) elif isinstance(callback, str) and callback == "detect_prompt_injection": from litellm.proxy.hooks.prompt_injection_detection import ( - _OPTIONAL_PromptInjectionDetection, + OPTIONAL_PromptInjectionDetection, ) prompt_injection_params = None @@ -332,20 +332,20 @@ def initialize_callbacks_on_proxy( prompt_injection_params_in_config = litellm_settings["prompt_injection_params"] prompt_injection_params = LiteLLMPromptInjectionParams(**prompt_injection_params_in_config) - prompt_injection_detection_obj = _OPTIONAL_PromptInjectionDetection( + prompt_injection_detection_obj = OPTIONAL_PromptInjectionDetection( prompt_injection_params=prompt_injection_params, ) imported_list.append(prompt_injection_detection_obj) elif isinstance(callback, str) and callback == "batch_redis_requests": from litellm.proxy.hooks.batch_redis_get import ( - _PROXY_BatchRedisRequests, + PROXY_BatchRedisRequests, ) - batch_redis_obj = _PROXY_BatchRedisRequests() + batch_redis_obj = PROXY_BatchRedisRequests() imported_list.append(batch_redis_obj) elif isinstance(callback, str) and callback == "azure_content_safety": from litellm.proxy.hooks.azure_content_safety import ( - _PROXY_AzureContentSafety, + PROXY_AzureContentSafety, ) azure_content_safety_params = litellm_settings["azure_content_safety_params"] @@ -353,7 +353,7 @@ def initialize_callbacks_on_proxy( if v is not None and isinstance(v, str) and v.startswith("os.environ/"): azure_content_safety_params[k] = get_secret(v) - azure_content_safety_obj = _PROXY_AzureContentSafety( + azure_content_safety_obj = PROXY_AzureContentSafety( **azure_content_safety_params, ) imported_list.append(azure_content_safety_obj) diff --git a/litellm/proxy/common_utils/config_sync_pubsub.py b/litellm/proxy/common_utils/config_sync_pubsub.py index b4ebb5fa876..579db290190 100644 --- a/litellm/proxy/common_utils/config_sync_pubsub.py +++ b/litellm/proxy/common_utils/config_sync_pubsub.py @@ -21,10 +21,13 @@ class _ConfigSyncPubSub(Protocol): def aclose(self) -> Awaitable[object]: ... +ConfigSyncPubSub = _ConfigSyncPubSub + + class _ConfigSyncPubSubClient(Protocol): def publish(self, channel: str, message: str) -> Awaitable[int]: ... - def pubsub(self) -> _ConfigSyncPubSub: ... + def pubsub(self) -> ConfigSyncPubSub: ... CONFIG_SYNC_CHANNEL: Final = "litellm_proxy.config_change" @@ -82,13 +85,16 @@ def config_sync_channel(redis_cache: "RedisCache") -> str: return f"{redis_cache.namespace}:{CONFIG_SYNC_CHANNEL}" -def _pubsub_capable_client(redis_cache: "RedisCache") -> _ConfigSyncPubSubClient: +def pubsub_capable_client(redis_cache: "RedisCache") -> _ConfigSyncPubSubClient: return cast( # cast-ok: protocol view of the pub/sub-capable async redis client _ConfigSyncPubSubClient, redis_cache.init_pubsub_client(), # pyright: ignore[reportUnknownMemberType] # redis generics ) +_pubsub_capable_client: Final = pubsub_capable_client + + @dataclass(frozen=True, slots=True) class _ConfigChangeMessage: object_type: str @@ -102,7 +108,7 @@ async def publish_config_change(redis_cache: "RedisCache | None", object_type: s if redis_cache is None: return try: - client: Final = _pubsub_capable_client(redis_cache) + client: Final = pubsub_capable_client(redis_cache) await client.publish(config_sync_channel(redis_cache), _config_change_message_json(object_type)) except Exception as e: # noqa: BLE001 # best-effort publish; writes must never fail on redis errors verbose_proxy_logger.warning("config sync publish for %s failed: %s", object_type, e) @@ -222,7 +228,7 @@ class ConfigSyncSubscriber: backoff_seconds = self._backoff_initial_seconds while True: try: - client = _pubsub_capable_client(self._redis_cache) + client = pubsub_capable_client(self._redis_cache) pubsub = client.pubsub() try: await pubsub.subscribe(config_sync_channel(self._redis_cache)) @@ -241,7 +247,7 @@ class ConfigSyncSubscriber: await self._sleep(backoff_seconds) backoff_seconds = min(backoff_seconds * 2, self._backoff_max_seconds) - async def _consume(self, pubsub: _ConfigSyncPubSub) -> None: + async def _consume(self, pubsub: ConfigSyncPubSub) -> None: while True: message = await pubsub.get_message(ignore_subscribe_messages=True, timeout=_POLL_TIMEOUT_SECONDS) if message is None: @@ -265,7 +271,7 @@ class ConfigSyncSubscriber: await self._sleep(seconds_until_next_resync) @staticmethod - async def _drain_pending(pubsub: _ConfigSyncPubSub) -> None: + async def _drain_pending(pubsub: ConfigSyncPubSub) -> None: while await pubsub.get_message(ignore_subscribe_messages=True, timeout=0) is not None: pass @@ -277,7 +283,7 @@ class ConfigSyncSubscriber: verbose_proxy_logger.warning("config sync resync callback failed: %s", e) @staticmethod - async def _close_pubsub(pubsub: _ConfigSyncPubSub) -> None: + async def _close_pubsub(pubsub: ConfigSyncPubSub) -> None: try: await pubsub.aclose() except Exception as e: # noqa: BLE001 # best-effort close of a possibly-broken connection diff --git a/litellm/proxy/common_utils/encrypt_decrypt_utils.py b/litellm/proxy/common_utils/encrypt_decrypt_utils.py index ae7240b8a7f..e5f6bc5dda8 100644 --- a/litellm/proxy/common_utils/encrypt_decrypt_utils.py +++ b/litellm/proxy/common_utils/encrypt_decrypt_utils.py @@ -12,18 +12,24 @@ from litellm._logging import verbose_proxy_logger # Legacy XSalsa20-Poly1305 (nacl) values carry no marker; the colon in the # prefix can never appear in base64url(nacl output), so the prefix check is an # unambiguous discriminator between the two formats on read. -_V2_GCM_PREFIX: Final = "v2:gcm:" +V2_GCM_PREFIX: Final = "v2:gcm:" + +_V2_GCM_PREFIX: Final = V2_GCM_PREFIX # general_settings key selecting the at-rest encryption algorithm for new writes. # Default preserves the legacy algorithm so existing deployments are byte-for-byte # unchanged until they explicitly opt in. Decrypt is always format-detecting, so # flipping this flag forward (or back) never strands previously-written data. -_ENCRYPTION_ALGORITHM_SETTING: Final = "encryption_algorithm" -_ALGO_AES_GCM: Final = "aes-256-gcm" +ENCRYPTION_ALGORITHM_SETTING: Final = "encryption_algorithm" + +_ENCRYPTION_ALGORITHM_SETTING: Final = ENCRYPTION_ALGORITHM_SETTING +ALGO_AES_GCM: Final = "aes-256-gcm" + +_ALGO_AES_GCM: Final = ALGO_AES_GCM _ALGO_XSALSA20: Final = "xsalsa20-poly1305" -def _get_salt_key(): +def get_salt_key() -> str | None: from litellm.proxy.proxy_server import master_key salt_key = os.getenv("LITELLM_SALT_KEY", None) @@ -34,6 +40,9 @@ def _get_salt_key(): return salt_key +_get_salt_key: Final = get_salt_key + + def _get_encryption_algorithm() -> str: """ Resolve the configured at-rest encryption algorithm for *new writes*. @@ -45,14 +54,14 @@ def _get_encryption_algorithm() -> str: try: from litellm.proxy.proxy_server import general_settings - algo: Final = general_settings.get(_ENCRYPTION_ALGORITHM_SETTING, _ALGO_XSALSA20) + algo: Final = general_settings.get(ENCRYPTION_ALGORITHM_SETTING, _ALGO_XSALSA20) except Exception: # general_settings may not be importable in some contexts (e.g. SDK-only # use of these helpers). Fall back to the legacy algorithm. return _ALGO_XSALSA20 - if isinstance(algo, str) and algo.lower() == _ALGO_AES_GCM: - return _ALGO_AES_GCM + if isinstance(algo, str) and algo.lower() == ALGO_AES_GCM: + return ALGO_AES_GCM return _ALGO_XSALSA20 @@ -92,18 +101,18 @@ def _open_aes_gcm(sealed: bytes, signing_key: str, aad: bytes | None) -> str: def _encrypt_aes_gcm(value: str, signing_key: str) -> str: """Encrypt under AES-256-GCM and return the versioned ``v2:gcm:`` string.""" sealed: Final = _seal_aes_gcm(value=value, signing_key=signing_key, aad=None) - return _V2_GCM_PREFIX + base64.urlsafe_b64encode(sealed).decode("utf-8") + return V2_GCM_PREFIX + base64.urlsafe_b64encode(sealed).decode("utf-8") def _decrypt_aes_gcm(value: str, signing_key: str) -> str: """Decrypt a versioned ``v2:gcm:`` string produced by :func:`_encrypt_aes_gcm`.""" - sealed: Final = base64.urlsafe_b64decode(value[len(_V2_GCM_PREFIX) :]) + sealed: Final = base64.urlsafe_b64decode(value[len(V2_GCM_PREFIX) :]) return _open_aes_gcm(sealed=sealed, signing_key=signing_key, aad=None) def encrypt_bearer_token(value: str, prefix: str) -> str: """AES-256-GCM as unpadded base64url behind ``prefix``, which is also the AAD so a token can't change kind.""" - salt_key: Final = _get_salt_key() + salt_key: Final = get_salt_key() if not isinstance(salt_key, str): raise ValueError("Set LITELLM_SALT_KEY or a master key to mint bearer tokens") sealed: Final = _seal_aes_gcm(value=value, signing_key=salt_key, aad=prefix.encode("utf-8")) @@ -112,7 +121,7 @@ def encrypt_bearer_token(value: str, prefix: str) -> str: def decrypt_bearer_token(token: str, prefix: str) -> str | None: """None unless ``token`` came from :func:`encrypt_bearer_token` with the same ``prefix``.""" - salt_key: Final = _get_salt_key() + salt_key: Final = get_salt_key() if not isinstance(salt_key, str) or not token.startswith(prefix): return None encoded: Final = token.removeprefix(prefix) @@ -124,11 +133,11 @@ def decrypt_bearer_token(token: str, prefix: str) -> str | None: def encrypt_value_helper(value: str, new_encryption_key: str | None = None): - signing_key: Final = new_encryption_key or _get_salt_key() + signing_key: Final = new_encryption_key or get_salt_key() try: if isinstance(value, str): - if _get_encryption_algorithm() == _ALGO_AES_GCM: + if _get_encryption_algorithm() == ALGO_AES_GCM: # AES path: the v2:gcm: output is already a base64url string, so it # is returned directly with no extra base64 wrapper. return _encrypt_aes_gcm(value=value, signing_key=cast(str, signing_key)) @@ -160,7 +169,7 @@ def _legacy_ciphertext_bytes(value: str) -> bytes: def _decrypt_with_signing_key(value: str, signing_key: str) -> str: # Versioned AES-256-GCM values are detected before any base64 decode. # The prefix is the algorithm tag the legacy nacl format never carried. - if value.startswith(_V2_GCM_PREFIX): + if value.startswith(V2_GCM_PREFIX): return _decrypt_aes_gcm(value=value, signing_key=signing_key) return decrypt_value(value=_legacy_ciphertext_bytes(value), signing_key=signing_key) @@ -171,7 +180,7 @@ def decrypt_if_encrypted_with(value: str, signing_key: str) -> str | None: try: # base64 decoding skips characters outside its alphabet, so "" and "*" decode to no bytes, # which decrypt_value reads as an empty plaintext under any key. - decodes_to_nothing: Final = not value.startswith(_V2_GCM_PREFIX) and not _legacy_ciphertext_bytes(value) + decodes_to_nothing: Final = not value.startswith(V2_GCM_PREFIX) and not _legacy_ciphertext_bytes(value) return None if decodes_to_nothing else _decrypt_with_signing_key(value=value, signing_key=signing_key) except Exception: # noqa: BLE001 # base64, nacl and AES-GCM each raise their own "not a ciphertext" type return None @@ -183,7 +192,7 @@ def decrypt_value_helper( exception_type: Literal["debug", "error"] = "error", return_original_value: bool = False, ) -> str | None: - signing_key: Final = _get_salt_key() + signing_key: Final = get_salt_key() try: if isinstance(value, str): diff --git a/litellm/proxy/common_utils/http_parsing_utils.py b/litellm/proxy/common_utils/http_parsing_utils.py index d1f34f7933f..8e8779ec193 100644 --- a/litellm/proxy/common_utils/http_parsing_utils.py +++ b/litellm/proxy/common_utils/http_parsing_utils.py @@ -186,7 +186,7 @@ def is_otlp_trace_request(request: Request) -> bool: return request.method == "POST" and get_route_path(request.scope) in {"/v1/traces", "/v1/logs"} -async def _read_request_body(request: Request | None) -> dict: +async def read_request_body(request: Request | None) -> dict: """ Safely read the request body and parse it as JSON. @@ -208,7 +208,7 @@ async def _read_request_body(request: Request | None) -> dict: if _cached_request_body is not None: return _cached_request_body - _request_headers: Final[dict] = _safe_get_request_headers(request=request) + _request_headers: Final[dict] = safe_get_request_headers(request=request) content_type: Final = _request_headers.get("content-type", "") if _normalize_media_type(content_type) in _BINARY_CONTENT_TYPES: @@ -285,7 +285,7 @@ async def _read_request_body(request: Request | None) -> dict: ) # Cache the parsed result - _safe_set_request_parsed_body(request=request, parsed_body=parsed_body) + safe_set_request_parsed_body(request=request, parsed_body=parsed_body) return parsed_body except (json.JSONDecodeError, orjson.JSONDecodeError, ProxyException) as e: @@ -298,6 +298,9 @@ async def _read_request_body(request: Request | None) -> dict: return {} +_read_request_body: Final = read_request_body + + def is_opaque_audio_pass_through_request(route: str, content_type: str) -> bool: """Azure Speech bodies (raw audio, multipart uploads) are forwarded byte for byte, so auth must not consume them.""" media_type: Final = _normalize_media_type(content_type) @@ -309,7 +312,7 @@ def is_opaque_audio_pass_through_request(route: str, content_type: str) -> bool: async def read_raw_json_body(request: Request | None) -> bytes | None: if request is None or _safe_get_request_parsed_body(request=request) is None: return None - content_type: Final = _safe_get_request_headers(request=request).get("content-type", "") + content_type: Final = safe_get_request_headers(request=request).get("content-type", "") if _is_form_content_type(content_type): return None try: @@ -334,7 +337,7 @@ def get_client_requested_model(request: Request | None) -> str | None: return model if isinstance(model, str) else None -def _safe_get_request_query_params(request: Request | None) -> dict: +def safe_get_request_query_params(request: Request | None) -> dict: if request is None: return {} try: @@ -346,7 +349,10 @@ def _safe_get_request_query_params(request: Request | None) -> dict: return {} -def _safe_set_request_parsed_body( +_safe_get_request_query_params: Final = safe_get_request_query_params + + +def safe_set_request_parsed_body( request: Request | None, parsed_body: dict, ) -> None: @@ -358,6 +364,9 @@ def _safe_set_request_parsed_body( verbose_proxy_logger.debug("Unexpected error setting request parsed body - %s", e) +_safe_set_request_parsed_body: Final = safe_set_request_parsed_body + + def rewrite_request_model( request_data: dict[str, object], request: Request | None, @@ -371,12 +380,12 @@ def rewrite_request_model( return cached_body: Final = _safe_get_request_parsed_body(request=request) body: Final = {**cached_body, "model": model} if cached_body is not None else request_data - _safe_set_request_parsed_body(request=request, parsed_body=body) - request._json = body - request._body = orjson.dumps(body) + safe_set_request_parsed_body(request=request, parsed_body=body) + request._json = body # pyright: ignore[reportPrivateUsage] # Starlette JSON cache + request._body = orjson.dumps(body) # pyright: ignore[reportPrivateUsage] # Starlette body cache -def _safe_get_request_headers(request: Request | None) -> dict: +def safe_get_request_headers(request: Request | None) -> dict: """ [Non-Blocking] Safely get the request headers. Caches the result on request.state to avoid re-creating dict(request.headers) per call. @@ -405,6 +414,9 @@ def _safe_get_request_headers(request: Request | None) -> dict: return headers +_safe_get_request_headers: Final = safe_get_request_headers + + def check_file_size_under_limit( request_data: dict, file: UploadFile, @@ -537,7 +549,7 @@ async def get_request_body(request: Request) -> dict[str, Any]: if request.method == "POST": content_type: Final = request.headers.get("content-type", "") if is_json_content_type(content_type): - return await _read_request_body(request) + return await read_request_body(request) elif _is_form_content_type(content_type): return await get_form_data(request) else: @@ -685,7 +697,7 @@ def populate_request_with_path_params(request_data: dict, request: Request) -> d dict: Updated request_data with path parameters and query parameters added """ # Add query parameters to request_data (for GET requests, etc.) - query_params: Final = _safe_get_request_query_params(request) + query_params: Final = safe_get_request_query_params(request) if query_params: for key, value in query_params.items(): # Don't overwrite existing values from request body diff --git a/litellm/proxy/common_utils/key_rotation_manager.py b/litellm/proxy/common_utils/key_rotation_manager.py index 352d024e20e..0ec196da6f5 100644 --- a/litellm/proxy/common_utils/key_rotation_manager.py +++ b/litellm/proxy/common_utils/key_rotation_manager.py @@ -20,8 +20,9 @@ from litellm.proxy._types import ( RegenerateKeyRequest, ) from litellm.proxy.hooks.key_management_event_hooks import KeyManagementEventHooks -from litellm.proxy.management_endpoints.key_management_endpoints import ( - _calculate_key_rotation_time, +from litellm.proxy.management_endpoints.key_management_endpoints import ( # noqa: F401 # legacy module exports + _calculate_key_rotation_time, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + calculate_key_rotation_time, regenerate_key_fn, ) from litellm.proxy.utils import PrismaClient @@ -187,7 +188,7 @@ class KeyRotationManager: if isinstance(response, GenerateKeyResponse) and response.token_id and key.rotation_interval: # Calculate next rotation time using helper function now: Final = datetime.now(timezone.utc) - next_rotation_time: Final = _calculate_key_rotation_time(key.rotation_interval) + next_rotation_time: Final = calculate_key_rotation_time(key.rotation_interval) await VerificationTokenRepository(self.prisma_client).table.update( where={"token": response.token_id}, data={ diff --git a/litellm/proxy/common_utils/openai_endpoint_utils.py b/litellm/proxy/common_utils/openai_endpoint_utils.py index f85bbf3b380..ab44d32f368 100644 --- a/litellm/proxy/common_utils/openai_endpoint_utils.py +++ b/litellm/proxy/common_utils/openai_endpoint_utils.py @@ -7,7 +7,10 @@ from typing import Final from fastapi import Request from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker -from litellm.proxy.common_utils.http_parsing_utils import _read_request_body +from litellm.proxy.common_utils.http_parsing_utils import ( # noqa: F401 # legacy module exports + _read_request_body, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + read_request_body, +) SENSITIVE_DATA_MASKER: Final = SensitiveDataMasker() @@ -55,7 +58,7 @@ async def get_custom_llm_provider_from_request_body(request: Request) -> str | N Safely reads the request body """ - request_body: Final[dict] = await _read_request_body(request=request) or {} + request_body: Final[dict] = await read_request_body(request=request) or {} if "custom_llm_provider" in request_body: return request_body["custom_llm_provider"] return None diff --git a/litellm/proxy/common_utils/openapi_schema_compat.py b/litellm/proxy/common_utils/openapi_schema_compat.py index 881b1fcc615..903e6c1cb24 100644 --- a/litellm/proxy/common_utils/openapi_schema_compat.py +++ b/litellm/proxy/common_utils/openapi_schema_compat.py @@ -40,7 +40,9 @@ def get_openapi_schema_with_compat( from pydantic_core import core_schema # Store original method - original_unknown_type_schema: Final = GenerateSchema._unknown_type_schema + original_unknown_type_schema: Final = ( + GenerateSchema._unknown_type_schema # pyright: ignore[reportPrivateUsage] # Pydantic schema internals + ) def patched_unknown_type_schema(self, obj): """Patch to handle openai.Timeout and other non-serializable types""" diff --git a/litellm/proxy/common_utils/rbac_utils.py b/litellm/proxy/common_utils/rbac_utils.py index 7c78d9a2470..b260fe127e4 100644 --- a/litellm/proxy/common_utils/rbac_utils.py +++ b/litellm/proxy/common_utils/rbac_utils.py @@ -51,10 +51,10 @@ async def check_feature_access_for_user( # Feature is disabled. Check if team/org admins are exempted. if general_settings.get(allow_team_admins_flag, False): from litellm.proxy.management_endpoints.common_utils import ( - _user_has_admin_privileges, + user_has_admin_privileges, ) - is_admin: Final = await _user_has_admin_privileges( + is_admin: Final = await user_has_admin_privileges( user_api_key_dict=user_api_key_dict, prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, diff --git a/litellm/proxy/common_utils/realtime_utils.py b/litellm/proxy/common_utils/realtime_utils.py index ff039754555..c641748843e 100644 --- a/litellm/proxy/common_utils/realtime_utils.py +++ b/litellm/proxy/common_utils/realtime_utils.py @@ -1,12 +1,16 @@ from functools import lru_cache +from typing import Final from litellm.constants import _REALTIME_BODY_CACHE_SIZE @lru_cache(maxsize=_REALTIME_BODY_CACHE_SIZE) -def _realtime_request_body(model: str | None) -> bytes: +def realtime_request_body(model: str | None) -> bytes: """ Generate the realtime websocket request body. Cached with LRU semantics to avoid repeated string formatting work while keeping memory usage bounded. """ return f'{{"model": "{model or ""}"}}'.encode() + + +_realtime_request_body: Final = realtime_request_body diff --git a/litellm/proxy/container_endpoints/endpoints.py b/litellm/proxy/container_endpoints/endpoints.py index 3142ea62b24..e11b0de3501 100644 --- a/litellm/proxy/container_endpoints/endpoints.py +++ b/litellm/proxy/container_endpoints/endpoints.py @@ -9,7 +9,10 @@ from litellm._logging import verbose_proxy_logger from litellm.proxy._types import * from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth, user_api_key_auth from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing -from litellm.proxy.common_utils.http_parsing_utils import _read_request_body +from litellm.proxy.common_utils.http_parsing_utils import ( # noqa: F401 # legacy module exports + _read_request_body, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + read_request_body, +) from litellm.proxy.common_utils.openai_endpoint_utils import ( get_custom_llm_provider_from_request_body, get_custom_llm_provider_from_request_headers, @@ -89,7 +92,7 @@ async def create_container( ) # Read request body - data: Final = await _read_request_body(request=request) + data: Final = await read_request_body(request=request) # Extract custom_llm_provider using priority chain # Priority: headers > query params > request body > default @@ -125,7 +128,7 @@ async def create_container( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -245,7 +248,7 @@ async def list_containers( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -360,7 +363,7 @@ async def retrieve_container( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -466,7 +469,7 @@ async def delete_container( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, diff --git a/litellm/proxy/container_endpoints/handler_factory.py b/litellm/proxy/container_endpoints/handler_factory.py index 3989bdacef1..35bd69ef0bd 100644 --- a/litellm/proxy/container_endpoints/handler_factory.py +++ b/litellm/proxy/container_endpoints/handler_factory.py @@ -260,7 +260,7 @@ async def _process_binary_request( ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -349,7 +349,7 @@ async def _process_multipart_upload_request( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -438,7 +438,7 @@ async def _process_request( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, diff --git a/litellm/proxy/custom_hooks/custom_ui_sso_hook.py b/litellm/proxy/custom_hooks/custom_ui_sso_hook.py index 53002680829..9220022a648 100644 --- a/litellm/proxy/custom_hooks/custom_ui_sso_hook.py +++ b/litellm/proxy/custom_hooks/custom_ui_sso_hook.py @@ -5,7 +5,10 @@ from fastapi_sso.sso.base import OpenID from litellm._logging import verbose_logger from litellm.integrations.custom_logger import CustomLogger -from litellm.proxy.common_utils.http_parsing_utils import _safe_get_request_headers +from litellm.proxy.common_utils.http_parsing_utils import ( # noqa: F401 # legacy module exports + _safe_get_request_headers, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + safe_get_request_headers, +) class CustomSSOLoginHandler(CustomLogger): @@ -22,7 +25,7 @@ class CustomSSOLoginHandler(CustomLogger): self, request: Request, ) -> OpenID: - request_headers_dict: Final = _safe_get_request_headers(request) + request_headers_dict: Final = safe_get_request_headers(request) verbose_logger.debug("inside custom ui sso sign in hook...") return OpenID( id=request_headers_dict.get("x-litellm-user-id") or "123", diff --git a/litellm/proxy/db/db_span.py b/litellm/proxy/db/db_span.py index 0cbd7c8db10..06b1a6aa11b 100644 --- a/litellm/proxy/db/db_span.py +++ b/litellm/proxy/db/db_span.py @@ -22,7 +22,12 @@ from typing import Final, TypeVar from litellm._logging import verbose_proxy_logger from litellm._service_logger import ServiceLogging, ServiceTypes -from litellm.proxy.db.log_db_metrics import _is_exception_related_to_db, claim_db_io, db_io_claimed +from litellm.proxy.db.log_db_metrics import ( # noqa: F401 # legacy module exports + _is_exception_related_to_db, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + claim_db_io, + db_io_claimed, + is_exception_related_to_db, +) _T = TypeVar("_T") @@ -75,7 +80,7 @@ async def db_span(call_type: str, table: str | None, operation: str | None = Non try: yield except Exception as e: - if service_logging is not None and _is_exception_related_to_db(e): + if service_logging is not None and is_exception_related_to_db(e): await _emit_failure(service_logging, call_type, event_metadata, start_time, e) raise if service_logging is None or not witness.touched: diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index 56b1b3176c6..370644675ce 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -1546,7 +1546,7 @@ class DBSpendUpdateWriter: else: - Regular flow of this method """ - if RedisUpdateBuffer._should_commit_spend_updates_to_redis(): + if RedisUpdateBuffer.should_commit_spend_updates_to_redis(): await self._commit_spend_updates_to_db_with_redis( prisma_client=prisma_client, n_retry_times=n_retry_times, @@ -1885,12 +1885,12 @@ class DBSpendUpdateWriter: ################## Tool Registry Upserts ################## await self._flush_tool_discovery_queue(prisma_client=prisma_client) - async def _commit_daily_tag_spend_to_db( + async def commit_daily_tag_spend_to_db( self, prisma_client: PrismaClient, n_retry_times: int, proxy_logging_obj: ProxyLogging, - ): + ) -> None: """ Commit only tag spend updates to database. This is called by a separate scheduler job at a longer interval. @@ -1904,12 +1904,14 @@ class DBSpendUpdateWriter: proxy_logging_obj=proxy_logging_obj, ) - async def _commit_daily_tag_spend_to_db_with_redis( + _commit_daily_tag_spend_to_db = commit_daily_tag_spend_to_db + + async def commit_daily_tag_spend_to_db_with_redis( self, prisma_client: PrismaClient, n_retry_times: int, proxy_logging_obj: ProxyLogging, - ): + ) -> None: """ Commit daily tag spend updates using Redis buffering. @@ -1942,6 +1944,8 @@ class DBSpendUpdateWriter: cronjob_id=DB_DAILY_TAG_SPEND_UPDATE_JOB_NAME, ) + _commit_daily_tag_spend_to_db_with_redis = commit_daily_tag_spend_to_db_with_redis + @staticmethod async def _commit_window_spend_updates( prisma_client: PrismaClient, @@ -2027,7 +2031,7 @@ class DBSpendUpdateWriter: verbose_proxy_logger.debug("_flush_tool_discovery_queue error (non-blocking): %s", e) @staticmethod - async def _handle_spend_update_failure( + async def handle_spend_update_failure( e: Exception, attempt: int, n_retry_times: int, @@ -2038,7 +2042,7 @@ class DBSpendUpdateWriter: ``lock_timeout`` (55P03), else re-raise. All three roll the transaction back before any increment applied, so re-sending the same batch cannot double-count.""" from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler - from litellm.proxy.utils import _raise_failed_update_spend_exception + from litellm.proxy.utils import raise_failed_update_spend_exception is_retryable = ( isinstance(e, DB_RETRY_SAFE_ERROR_TYPES) @@ -2046,7 +2050,7 @@ class DBSpendUpdateWriter: or PrismaDBExceptionHandler.is_lock_timeout_error(e) ) if not is_retryable or attempt >= n_retry_times: - _raise_failed_update_spend_exception(e=e, start_time=start_time, proxy_logging_obj=proxy_logging_obj) + raise_failed_update_spend_exception(e=e, start_time=start_time, proxy_logging_obj=proxy_logging_obj) verbose_proxy_logger.warning( "Retrying spend update after retryable DB error (attempt %s/%s): %s", attempt + 1, @@ -2055,6 +2059,8 @@ class DBSpendUpdateWriter: ) await asyncio.sleep(random.uniform(2**attempt, 2 ** (attempt + 1))) + _handle_spend_update_failure = handle_spend_update_failure + async def _commit_spend_updates_to_db( self, prisma_client: PrismaClient, @@ -2088,7 +2094,7 @@ class DBSpendUpdateWriter: ) break except Exception as e: - await self._handle_spend_update_failure( + await self.handle_spend_update_failure( e=e, attempt=i, n_retry_times=n_retry_times, @@ -2132,7 +2138,7 @@ class DBSpendUpdateWriter: ) break except Exception as e: - await self._handle_spend_update_failure( + await self.handle_spend_update_failure( e=e, attempt=i, n_retry_times=n_retry_times, @@ -2162,7 +2168,7 @@ class DBSpendUpdateWriter: ) break except Exception as e: - await self._handle_spend_update_failure( + await self.handle_spend_update_failure( e=e, attempt=i, n_retry_times=n_retry_times, @@ -2192,7 +2198,7 @@ class DBSpendUpdateWriter: # Transaction succeeded, break out of retry loop break except Exception as e: - await self._handle_spend_update_failure( + await self.handle_spend_update_failure( e=e, attempt=i, n_retry_times=n_retry_times, @@ -2234,7 +2240,7 @@ class DBSpendUpdateWriter: ) break except Exception as e: - await self._handle_spend_update_failure( + await self.handle_spend_update_failure( e=e, attempt=i, n_retry_times=n_retry_times, @@ -2262,7 +2268,7 @@ class DBSpendUpdateWriter: ) break except Exception as e: - await self._handle_spend_update_failure( + await self.handle_spend_update_failure( e=e, attempt=i, n_retry_times=n_retry_times, @@ -2389,7 +2395,7 @@ class DBSpendUpdateWriter: ) break except Exception as e: - await DBSpendUpdateWriter._handle_spend_update_failure( + await DBSpendUpdateWriter.handle_spend_update_failure( e=e, attempt=i, n_retry_times=n_retry_times, @@ -2489,7 +2495,7 @@ class DBSpendUpdateWriter: """ Generic function to update daily spend for any entity type (user, team, org, tag, end_user, agent) """ - from litellm.proxy.utils import _raise_failed_update_spend_exception + from litellm.proxy.utils import raise_failed_update_spend_exception verbose_proxy_logger.debug( "Daily %s Spend transactions: %s", entity_type.capitalize(), len(daily_spend_transactions) @@ -2589,7 +2595,7 @@ class DBSpendUpdateWriter: if not is_retryable: raise if i >= n_retry_times: - _raise_failed_update_spend_exception( + raise_failed_update_spend_exception( e=e, start_time=start_time, proxy_logging_obj=proxy_logging_obj, @@ -2604,7 +2610,7 @@ class DBSpendUpdateWriter: ) except Exception as e: - _raise_failed_update_spend_exception(e=e, start_time=start_time, proxy_logging_obj=proxy_logging_obj) + raise_failed_update_spend_exception(e=e, start_time=start_time, proxy_logging_obj=proxy_logging_obj) @staticmethod async def update_daily_user_spend( diff --git a/litellm/proxy/db/db_transaction_queue/redis_update_buffer.py b/litellm/proxy/db/db_transaction_queue/redis_update_buffer.py index 3d24b0a0612..898aee493df 100644 --- a/litellm/proxy/db/db_transaction_queue/redis_update_buffer.py +++ b/litellm/proxy/db/db_transaction_queue/redis_update_buffer.py @@ -147,7 +147,7 @@ class RedisUpdateBuffer: self.redis_cache = redis_cache @staticmethod - def _should_commit_spend_updates_to_redis() -> bool: + def should_commit_spend_updates_to_redis() -> bool: """ Checks if the Pod should commit spend updates to Redis @@ -163,6 +163,8 @@ class RedisUpdateBuffer: return False return _use_redis_transaction_buffer + _should_commit_spend_updates_to_redis = should_commit_spend_updates_to_redis + @with_service_target(SPEND_QUEUE_TARGET) async def _store_transactions_in_redis( self, @@ -556,7 +558,7 @@ class RedisUpdateBuffer: max_rows: int = REDIS_SPEND_LOGS_BUFFER_MAX_ROWS, ) -> bool: """Park spend-log rows in Redis so they outlive this pod, dropping the oldest past ``max_rows``.""" - if self.redis_cache is None or len(rows) == 0 or not self._should_commit_spend_updates_to_redis(): + if self.redis_cache is None or len(rows) == 0 or not self.should_commit_spend_updates_to_redis(): return False try: buffer_size: Final = await self.redis_cache.async_rpush_and_trim( @@ -582,7 +584,7 @@ class RedisUpdateBuffer: @with_service_target(SPEND_QUEUE_TARGET) async def get_spend_logs_from_redis_buffer(self, limit: int) -> tuple[dict[str, object], ...]: """Atomically take up to ``limit`` parked spend-log rows out of Redis.""" - if self.redis_cache is None or not self._should_commit_spend_updates_to_redis(): + if self.redis_cache is None or not self.should_commit_spend_updates_to_redis(): return () popped: Final[str | list[str] | None] = await self.redis_cache.async_lpop( key=REDIS_SPEND_LOGS_BUFFER_KEY, diff --git a/litellm/proxy/db/log_db_metrics.py b/litellm/proxy/db/log_db_metrics.py index 988f93fb962..027dbe8e9fb 100644 --- a/litellm/proxy/db/log_db_metrics.py +++ b/litellm/proxy/db/log_db_metrics.py @@ -175,7 +175,7 @@ def log_db_metrics(func): return wrapper -def _is_exception_related_to_db(e: Exception) -> bool: +def is_exception_related_to_db(e: Exception) -> bool: """ Returns True if the exception is related to the DB """ @@ -186,6 +186,9 @@ def _is_exception_related_to_db(e: Exception) -> bool: return isinstance(e, (PrismaError, httpx.TransportError)) +_is_exception_related_to_db: Final = is_exception_related_to_db + + async def _handle_logging_db_exception( e: Exception, func: Callable, @@ -198,7 +201,7 @@ async def _handle_logging_db_exception( from litellm.proxy.proxy_server import proxy_logging_obj # don't log this as a DB Service Failure, if the DB did not raise an exception - if _is_exception_related_to_db(e) is not True: + if is_exception_related_to_db(e) is not True: return False try: diff --git a/litellm/proxy/decisions_endpoints/endpoints.py b/litellm/proxy/decisions_endpoints/endpoints.py index 7dba64ee791..f2891c184bf 100644 --- a/litellm/proxy/decisions_endpoints/endpoints.py +++ b/litellm/proxy/decisions_endpoints/endpoints.py @@ -1,3 +1,4 @@ +from collections.abc import Mapping from typing import Annotated, Final from fastapi import APIRouter, Depends, Request, Response @@ -7,33 +8,50 @@ from pydantic import TypeAdapter, ValidationError from litellm.exceptions import BadRequestError from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth, user_api_key_auth from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing -from litellm.types.decisions import DecisionsRequestBody +from litellm.types.decisions import DecisionsRequestBody, OpenAIDecisionRequestBody router: Final = APIRouter() _REQUEST_DATA_ADAPTER: Final[TypeAdapter[dict[str, object]]] = TypeAdapter(dict[str, object]) _DECISIONS_REQUEST_BODY_ADAPTER: Final[TypeAdapter[DecisionsRequestBody]] = TypeAdapter(DecisionsRequestBody) +_OPENAI_DECISION_REQUEST_BODY_ADAPTER: Final[TypeAdapter[OpenAIDecisionRequestBody]] = TypeAdapter( + OpenAIDecisionRequestBody +) _GENERAL_SETTINGS_ADAPTER: Final[TypeAdapter[dict[str, object]]] = TypeAdapter(dict[str, object]) _OPTIONAL_STRING_ADAPTER: Final[TypeAdapter[str | None]] = TypeAdapter(str | None) _OPTIONAL_FLOAT_ADAPTER: Final[TypeAdapter[float | None]] = TypeAdapter(float | None) -@router.post( - "/v1/decisions", - dependencies=[Depends(user_api_key_auth)], - response_class=ORJSONResponse, # pyright: ignore[reportDeprecated] # required endpoint contract - tags=["decisions"], -) -@router.post( - "/decisions", - dependencies=[Depends(user_api_key_auth)], - response_class=ORJSONResponse, # pyright: ignore[reportDeprecated] # required endpoint contract - tags=["decisions"], -) -async def decisions( +async def _invalid_request( + raw_data: Mapping[str, object], error: ValidationError, user_api_key_dict: UserAPIKeyAuth +) -> Exception: + from litellm.proxy.proxy_server import proxy_logging_obj, version + + return await ProxyBaseLLMRequestProcessing(data=dict(raw_data)).handle_llm_api_exception( + e=BadRequestError( + message=f"Invalid Decisions request: {error}", + model=str(raw_data.get("model", "")), + llm_provider="", + ), + user_api_key_dict=user_api_key_dict, + proxy_logging_obj=proxy_logging_obj, + version=version, + ) + + +async def _request_data(request: Request, user_api_key_dict: UserAPIKeyAuth) -> dict[str, object]: + body: Final = await request.body() + try: + return _REQUEST_DATA_ADAPTER.validate_json(body) + except ValidationError as error: + raise await _invalid_request(raw_data={}, error=error, user_api_key_dict=user_api_key_dict) + + +async def _process_decisions( request: Request, fastapi_response: Response, - user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)], -): + user_api_key_dict: UserAPIKeyAuth, + body_adapter: TypeAdapter[DecisionsRequestBody] | TypeAdapter[OpenAIDecisionRequestBody], +) -> object: from litellm.proxy.proxy_server import ( general_settings as proxy_general_settings, ) @@ -55,14 +73,17 @@ async def decisions( user_temperature as proxy_user_temperature, ) - data: Final = _REQUEST_DATA_ADAPTER.validate_json(await request.body()) + data: Final = await _request_data(request, user_api_key_dict) + try: + body_adapter.validate_python(data) + except ValidationError as error: + raise await _invalid_request(raw_data=data, error=error, user_api_key_dict=user_api_key_dict) general_settings: Final = _GENERAL_SETTINGS_ADAPTER.validate_python(proxy_general_settings) user_api_base: Final = _OPTIONAL_STRING_ADAPTER.validate_python(proxy_user_api_base) user_model: Final = _OPTIONAL_STRING_ADAPTER.validate_python(proxy_user_model) user_temperature: Final = _OPTIONAL_FLOAT_ADAPTER.validate_python(proxy_user_temperature) processor: Final = ProxyBaseLLMRequestProcessing(data=data) try: - _DECISIONS_REQUEST_BODY_ADAPTER.validate_python(data) return await processor.base_process_llm_request( request=request, fastapi_response=fastapi_response, @@ -81,22 +102,60 @@ async def decisions( user_api_base=user_api_base, version=version, ) - except ValidationError as error: - bad_request_error: Final = BadRequestError( - message=f"Invalid Decisions request: {error}", - model=str(data.get("model", "")), - llm_provider="", - ) - raise await processor._handle_llm_api_exception( - e=bad_request_error, - user_api_key_dict=user_api_key_dict, - proxy_logging_obj=proxy_logging_obj, - version=version, - ) except Exception as error: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=error, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, version=version, ) + + +@router.post( + "/v1/systemone", + dependencies=[Depends(user_api_key_auth)], + response_class=ORJSONResponse, # pyright: ignore[reportDeprecated] # required endpoint contract + tags=["decisions"], +) +@router.post( + "/systemone", + dependencies=[Depends(user_api_key_auth)], + response_class=ORJSONResponse, # pyright: ignore[reportDeprecated] # required endpoint contract + tags=["decisions"], +) +async def systemone( + request: Request, + fastapi_response: Response, + user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)], +): + return await _process_decisions( + request=request, + fastapi_response=fastapi_response, + user_api_key_dict=user_api_key_dict, + body_adapter=_DECISIONS_REQUEST_BODY_ADAPTER, + ) + + +@router.post( + "/v1/decisions", + dependencies=[Depends(user_api_key_auth)], + response_class=ORJSONResponse, # pyright: ignore[reportDeprecated] # required endpoint contract + tags=["decisions"], +) +@router.post( + "/decisions", + dependencies=[Depends(user_api_key_auth)], + response_class=ORJSONResponse, # pyright: ignore[reportDeprecated] # required endpoint contract + tags=["decisions"], +) +async def decisions( + request: Request, + fastapi_response: Response, + user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)], +): + return await _process_decisions( + request=request, + fastapi_response=fastapi_response, + user_api_key_dict=user_api_key_dict, + body_adapter=_OPENAI_DECISION_REQUEST_BODY_ADAPTER, + ) diff --git a/litellm/proxy/fine_tuning_endpoints/endpoints.py b/litellm/proxy/fine_tuning_endpoints/endpoints.py index 886b8da1454..d015112cf3a 100644 --- a/litellm/proxy/fine_tuning_endpoints/endpoints.py +++ b/litellm/proxy/fine_tuning_endpoints/endpoints.py @@ -16,8 +16,9 @@ from litellm.litellm_core_utils.hidden_params import set_hidden_param from litellm.proxy._types import * from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing -from litellm.proxy.openai_files_endpoints.common_utils import ( - _is_base64_encoded_unified_file_id, +from litellm.proxy.openai_files_endpoints.common_utils import ( # noqa: F401 # legacy module exports + _is_base64_encoded_unified_file_id, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + is_base64_encoded_unified_file_id, validate_managed_id_requirement, ) from litellm.proxy.utils import handle_exception_on_proxy @@ -150,7 +151,9 @@ async def create_fine_tuning_job( ) response: LiteLLMFineTuningJob | None = None if training_file: - unified_file_id = _is_base64_encoded_unified_file_id(training_file) + unified_file_id = ( # rebind-ok: pre-existing rebinding on a rename-only line + is_base64_encoded_unified_file_id(training_file) + ) ## IF SO, Route based on that if unified_file_id: """ """ @@ -292,7 +295,9 @@ async def retrieve_fine_tuning_job( unified_finetuning_job_id: str | Literal[False] = False response: LiteLLMFineTuningJob | None = None if fine_tuning_job_id: - unified_finetuning_job_id = _is_base64_encoded_unified_file_id(fine_tuning_job_id) + unified_finetuning_job_id = ( # rebind-ok: pre-existing rebinding on a rename-only line + is_base64_encoded_unified_file_id(fine_tuning_job_id) + ) if unified_finetuning_job_id: if llm_router is None: raise HTTPException( @@ -565,7 +570,9 @@ async def cancel_fine_tuning_job( unified_finetuning_job_id: str | Literal[False] = False response: LiteLLMFineTuningJob | None = None if fine_tuning_job_id: - unified_finetuning_job_id = _is_base64_encoded_unified_file_id(fine_tuning_job_id) + unified_finetuning_job_id = ( # rebind-ok: pre-existing rebinding on a rename-only line + is_base64_encoded_unified_file_id(fine_tuning_job_id) + ) if unified_finetuning_job_id: if llm_router is None: raise HTTPException( diff --git a/litellm/proxy/google_endpoints/agents_endpoints.py b/litellm/proxy/google_endpoints/agents_endpoints.py index 59ef45c7817..00b23d846c2 100644 --- a/litellm/proxy/google_endpoints/agents_endpoints.py +++ b/litellm/proxy/google_endpoints/agents_endpoints.py @@ -23,9 +23,11 @@ from fastapi.responses import ORJSONResponse from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing -from litellm.proxy.common_utils.http_parsing_utils import ( - _read_request_body, - _safe_get_request_query_params, +from litellm.proxy.common_utils.http_parsing_utils import ( # noqa: F401 # legacy module exports + _read_request_body, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _safe_get_request_query_params, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + read_request_body, + safe_get_request_query_params, ) router: Final = APIRouter(tags=["gemini managed agents"]) @@ -92,7 +94,7 @@ def _merge_query_params_into_data(data: dict, request: Request) -> dict: headers. Use the ``litellm_params_template`` JSON body field on POST requests, or the JSON-encoded query parameter above for GET/DELETE. """ - query_params: Final = _safe_get_request_query_params(request) + query_params: Final = safe_get_request_query_params(request) if not query_params: return data @@ -172,7 +174,7 @@ async def create_gemini_agent( ``` """ srv: Final = _proxy_server_imports() - data: Final = await _read_request_body(request=request) + data: Final = await read_request_body(request=request) # Merge litellm_params_template (e.g. custom_llm_provider, api_key) into the request litellm_params_template: Final = data.pop("litellm_params_template", None) or {} if isinstance(litellm_params_template, dict): @@ -203,7 +205,7 @@ async def create_gemini_agent( version=srv["version"], ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=srv["proxy_logging_obj"], @@ -260,7 +262,7 @@ async def list_gemini_agents( version=srv["version"], ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=srv["proxy_logging_obj"], @@ -318,7 +320,7 @@ async def get_gemini_agent( version=srv["version"], ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=srv["proxy_logging_obj"], @@ -376,7 +378,7 @@ async def delete_gemini_agent( version=srv["version"], ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=srv["proxy_logging_obj"], @@ -434,7 +436,7 @@ async def list_gemini_agent_versions( version=srv["version"], ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=srv["proxy_logging_obj"], diff --git a/litellm/proxy/google_endpoints/endpoints.py b/litellm/proxy/google_endpoints/endpoints.py index dd5dc66d82b..87d2ee07818 100644 --- a/litellm/proxy/google_endpoints/endpoints.py +++ b/litellm/proxy/google_endpoints/endpoints.py @@ -6,7 +6,10 @@ from fastapi.responses import ORJSONResponse from litellm.proxy._types import * from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth, user_api_key_auth from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing -from litellm.proxy.common_utils.http_parsing_utils import _read_request_body +from litellm.proxy.common_utils.http_parsing_utils import ( # noqa: F401 # legacy module exports + _read_request_body, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + read_request_body, +) from litellm.types.llms.vertex_ai import TokenCountDetailsResponse router: Final = APIRouter( @@ -42,7 +45,7 @@ async def google_generate_content( version, ) - data: Final = await _read_request_body(request=request) + data: Final = await read_request_body(request=request) if "model" not in data: data["model"] = model_name @@ -67,7 +70,7 @@ async def google_generate_content( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -103,7 +106,7 @@ async def google_stream_generate_content( version, ) - data: Final = await _read_request_body(request=request) + data: Final = await read_request_body(request=request) if "model" not in data: data["model"] = model_name data["stream"] = True @@ -132,7 +135,7 @@ async def google_stream_generate_content( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -166,10 +169,10 @@ async def google_count_tokens(request: Request, model_name: str): ``` """ from litellm.google_genai.adapters.transformation import GoogleGenAIAdapter - from litellm.proxy.common_utils.http_parsing_utils import _read_request_body + from litellm.proxy.common_utils.http_parsing_utils import read_request_body from litellm.proxy.proxy_server import token_counter as internal_token_counter - data: Final = await _read_request_body(request=request) + data: Final = await read_request_body(request=request) contents: Final = data.get("contents", []) # Create TokenCountRequest for the internal endpoint from litellm.proxy._types import TokenCountRequest @@ -268,7 +271,7 @@ async def create_interaction( version, ) - data: Final = await _read_request_body(request=request) + data: Final = await read_request_body(request=request) # Default to gemini provider for interactions if "custom_llm_provider" not in data: @@ -295,7 +298,7 @@ async def create_interaction( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -363,7 +366,7 @@ async def get_interaction( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -431,7 +434,7 @@ async def delete_interaction( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -499,7 +502,7 @@ async def cancel_interaction( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, diff --git a/litellm/proxy/guardrails/guardrail_endpoints.py b/litellm/proxy/guardrails/guardrail_endpoints.py index 1c9639fa9dc..8e1209401ab 100644 --- a/litellm/proxy/guardrails/guardrail_endpoints.py +++ b/litellm/proxy/guardrails/guardrail_endpoints.py @@ -43,7 +43,10 @@ from litellm.proxy.guardrails.guardrail_registry import ( parse_tolerant_litellm_params, ) from litellm.proxy.guardrails.usage_endpoints import router as guardrails_usage_router -from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view +from litellm.proxy.management_endpoints.common_utils import ( # noqa: F401 # legacy module exports + _user_has_admin_view, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + user_api_key_has_admin_view, +) from litellm.repositories.prisma_protocols import TableActions from litellm.repositories.table_repositories import GuardrailsRepository from litellm.types.guardrails import ( @@ -247,7 +250,7 @@ async def list_guardrails_v2( from litellm.proxy.guardrails.guardrail_registry import IN_MEMORY_GUARDRAIL_HANDLER from litellm.proxy.proxy_server import prisma_client - is_admin: Final = _user_has_admin_view(user_api_key_dict) + is_admin: Final = user_api_key_has_admin_view(user_api_key_dict) try: guardrails = ( @@ -942,7 +945,7 @@ async def list_guardrail_submissions( # Admin Viewer follows the read-parity rule: see all submissions like a # Proxy Admin would (no writes — registration / approval still gated # elsewhere by their own per-action checks). - is_admin: Final = _user_has_admin_view(user_api_key_dict) + is_admin: Final = user_api_key_has_admin_view(user_api_key_dict) visible_team_ids: list[str] | None = None if not is_admin: visible_team_ids = await _get_user_team_ids(user_api_key_dict) @@ -1021,7 +1024,7 @@ async def get_guardrail_submission( if prisma_client is None: raise HTTPException(status_code=500, detail="Prisma client not initialized") - is_admin: Final = _user_has_admin_view(user_api_key_dict) + is_admin: Final = user_api_key_has_admin_view(user_api_key_dict) try: row: Final = await _guardrails_table(prisma_client).find_unique(where={"guardrail_id": guardrail_id}) diff --git a/litellm/proxy/guardrails/guardrail_hooks/azure/base.py b/litellm/proxy/guardrails/guardrail_hooks/azure/base.py index b69028594b2..5f57969c310 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/azure/base.py +++ b/litellm/proxy/guardrails/guardrail_hooks/azure/base.py @@ -26,7 +26,9 @@ AZURE_CONTENT_SAFETY_TEXT_RECORD_LENGTH: Final = 1000 AZURE_CONTENT_SAFETY_DEFAULT_API_VERSION: Final = "2024-09-01" JAVELIN_API_VERSION_STORED_BY_OLDER_RELEASES: Final = "v1" -_RESPONSES_API_CALL_TYPES: Final = frozenset({CallTypes.responses, CallTypes.aresponses}) +RESPONSES_API_CALL_TYPES: Final = frozenset({CallTypes.responses, CallTypes.aresponses}) + +_RESPONSES_API_CALL_TYPES: Final = RESPONSES_API_CALL_TYPES def resolve_content_safety_api_version(configured: str | None) -> str: @@ -155,7 +157,7 @@ class AzureGuardrailBase: return get_last_user_message(messages) def get_user_prompt_from_request(self, data: Mapping[str, object], call_type: CallTypesLiteral) -> str | None: - if call_type in _RESPONSES_API_CALL_TYPES: + if call_type in RESPONSES_API_CALL_TYPES: responses_input: Final = data.get("input") if not isinstance(responses_input, (str, list)): return None diff --git a/litellm/proxy/guardrails/guardrail_hooks/azure/prompt_shield.py b/litellm/proxy/guardrails/guardrail_hooks/azure/prompt_shield.py index 9a5303776ab..cdb7a4856ff 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/azure/prompt_shield.py +++ b/litellm/proxy/guardrails/guardrail_hooks/azure/prompt_shield.py @@ -27,7 +27,12 @@ from litellm.types.utils import ( GuardrailTracingDetail, ) -from .base import _RESPONSES_API_CALL_TYPES, AZURE_CONTENT_SAFETY_TEXT_RECORD_LENGTH, AzureGuardrailBase +from .base import ( # noqa: F401 # legacy module exports + _RESPONSES_API_CALL_TYPES, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + AZURE_CONTENT_SAFETY_TEXT_RECORD_LENGTH, + RESPONSES_API_CALL_TYPES, + AzureGuardrailBase, +) if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj @@ -249,7 +254,7 @@ class AzureContentSafetyPromptShieldGuardrail(AzureGuardrailBase, CustomGuardrai "Azure Prompt Shield: Running pre-call prompt scan, on call_type: %s", call_type, ) - if call_type not in _RESPONSES_API_CALL_TYPES and data.get("messages") is None: + if call_type not in RESPONSES_API_CALL_TYPES and data.get("messages") is None: verbose_proxy_logger.warning("Azure Prompt Shield: not running guardrail. No messages in data") return data user_prompt: Final = self.get_user_prompt_from_request(data, call_type) diff --git a/litellm/proxy/guardrails/guardrail_hooks/azure/text_moderation.py b/litellm/proxy/guardrails/guardrail_hooks/azure/text_moderation.py index d9147cfb62b..27f36f2724b 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/azure/text_moderation.py +++ b/litellm/proxy/guardrails/guardrail_hooks/azure/text_moderation.py @@ -16,7 +16,11 @@ from litellm.proxy._types import UserAPIKeyAuth from litellm.types.guardrails import GuardrailEventHooks from litellm.types.utils import CallTypesLiteral, GenericGuardrailAPIInputs, LLMResponseTypes -from .base import _RESPONSES_API_CALL_TYPES, AzureGuardrailBase +from .base import ( # noqa: F401 # legacy module exports + _RESPONSES_API_CALL_TYPES, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + RESPONSES_API_CALL_TYPES, + AzureGuardrailBase, +) if TYPE_CHECKING: from litellm.caching.caching import DualCache @@ -231,7 +235,7 @@ class AzureContentSafetyTextModerationGuardrail(AzureGuardrailBase, CustomGuardr "Azure Text Moderation: Running pre-call prompt scan, on call_type: %s", call_type, ) - if call_type not in _RESPONSES_API_CALL_TYPES and data.get("messages") is None: + if call_type not in RESPONSES_API_CALL_TYPES and data.get("messages") is None: verbose_proxy_logger.warning("Azure Text Moderation: not running guardrail. No messages in data") return data user_prompt: Final = self.get_user_prompt_from_request(data, call_type) diff --git a/litellm/proxy/guardrails/guardrail_hooks/cisco_ai_defense/cisco_ai_defense.py b/litellm/proxy/guardrails/guardrail_hooks/cisco_ai_defense/cisco_ai_defense.py index 1e24b2aecbc..ab43d32e260 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/cisco_ai_defense/cisco_ai_defense.py +++ b/litellm/proxy/guardrails/guardrail_hooks/cisco_ai_defense/cisco_ai_defense.py @@ -52,7 +52,10 @@ from litellm.types.utils import ( TextCompletionResponse, ) -from .cisco_ai_defense_mcp import _CiscoAIDefenseMcpMixin +from .cisco_ai_defense_mcp import ( # noqa: F401 # legacy module exports + CiscoAIDefenseMcpMixin, + _CiscoAIDefenseMcpMixin, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export +) if TYPE_CHECKING: from litellm.types.proxy.guardrails.guardrail_hooks.base import ( @@ -116,7 +119,7 @@ class _CustomGuardrailOptions(TypedDict, total=False, extra_items=object): """Base-class constructor options this guardrail forwards untouched to CustomGuardrail.""" -class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail): +class CiscoAIDefenseGuardrail(CiscoAIDefenseMcpMixin, CustomGuardrail): """ Cisco AI Defense guardrail integration. diff --git a/litellm/proxy/guardrails/guardrail_hooks/cisco_ai_defense/cisco_ai_defense_mcp.py b/litellm/proxy/guardrails/guardrail_hooks/cisco_ai_defense/cisco_ai_defense_mcp.py index 67ef05fc324..d7cc61ab511 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/cisco_ai_defense/cisco_ai_defense_mcp.py +++ b/litellm/proxy/guardrails/guardrail_hooks/cisco_ai_defense/cisco_ai_defense_mcp.py @@ -218,7 +218,7 @@ class _CiscoAIDefenseMcpMixin: inner: Final[object | None] = getattr(response_obj, "mcp_tool_call_response", None) if inner is not None: - if _CiscoAIDefenseMcpMixin._replace_mcp_tool_response(inner, replacement_obj): + if CiscoAIDefenseMcpMixin._replace_mcp_tool_response(inner, replacement_obj): return True try: setattr(response_obj, "mcp_tool_call_response", replacement) @@ -229,7 +229,7 @@ class _CiscoAIDefenseMcpMixin: content: Final = getattr(response_obj, "content", None) if isinstance(content, list): content[:] = replacement - structured_replacement: Final = _CiscoAIDefenseMcpMixin._replacement_structured_content(replacement) + structured_replacement: Final = CiscoAIDefenseMcpMixin._replacement_structured_content(replacement) if hasattr(response_obj, "structured_content"): try: setattr(response_obj, "structured_content", structured_replacement) @@ -250,12 +250,12 @@ class _CiscoAIDefenseMcpMixin: result: Final = response_obj.get("result") if isinstance(result, dict): result["content"] = replacement - result["structuredContent"] = _CiscoAIDefenseMcpMixin._replacement_structured_content(replacement) + result["structuredContent"] = CiscoAIDefenseMcpMixin._replacement_structured_content(replacement) result["isError"] = True return True response_obj["result"] = { "content": replacement, - "structuredContent": _CiscoAIDefenseMcpMixin._replacement_structured_content(replacement), + "structuredContent": CiscoAIDefenseMcpMixin._replacement_structured_content(replacement), "isError": True, } return True @@ -472,7 +472,7 @@ class _CiscoAIDefenseMcpMixin: return { "jsonrpc": "2.0", "id": response.get("id") or "litellm-mcp", - "result": _CiscoAIDefenseMcpMixin._build_mcp_result(content=content, source=response), + "result": CiscoAIDefenseMcpMixin._build_mcp_result(content=content, source=response), } if isinstance(response, list): if response and all( @@ -484,7 +484,7 @@ class _CiscoAIDefenseMcpMixin: return { "jsonrpc": "2.0", "id": "litellm-mcp", - "result": _CiscoAIDefenseMcpMixin._build_mcp_result( + "result": CiscoAIDefenseMcpMixin._build_mcp_result( content=inner_content, source=response_fields ), } @@ -493,7 +493,7 @@ class _CiscoAIDefenseMcpMixin: return { "jsonrpc": "2.0", "id": "litellm-mcp", - "result": _CiscoAIDefenseMcpMixin._build_mcp_result(content=response), + "result": CiscoAIDefenseMcpMixin._build_mcp_result(content=response), } model_dump: Final = getattr(response, "model_dump", None) if callable(model_dump): @@ -502,13 +502,13 @@ class _CiscoAIDefenseMcpMixin: except TypeError: dumped = model_dump() if isinstance(dumped, dict): - return _CiscoAIDefenseMcpMixin._normalize_mcp_response(dumped) + return CiscoAIDefenseMcpMixin._normalize_mcp_response(dumped) content = getattr(response, "content", None) if isinstance(content, list): return { "jsonrpc": "2.0", "id": "litellm-mcp", - "result": _CiscoAIDefenseMcpMixin._build_mcp_result(content=content, source=response), + "result": CiscoAIDefenseMcpMixin._build_mcp_result(content=content, source=response), } return None @@ -536,9 +536,9 @@ class _CiscoAIDefenseMcpMixin: inner: Final[object | None] = getattr(response_obj, "mcp_tool_call_response", None) if inner is not None: - return _CiscoAIDefenseMcpMixin._set_mcp_tool_response_text(inner, text) + return CiscoAIDefenseMcpMixin._set_mcp_tool_response_text(inner, text) - content_list: Final = _CiscoAIDefenseMcpMixin._coerce_to_content_list(response_obj) + content_list: Final = CiscoAIDefenseMcpMixin._coerce_to_content_list(response_obj) replaced = False if isinstance(content_list, list): @@ -586,7 +586,7 @@ class _CiscoAIDefenseMcpMixin: return None inner: Final[object | None] = getattr(response_obj, "mcp_tool_call_response", None) if inner is not None: - return _CiscoAIDefenseMcpMixin._coerce_to_content_list(inner) + return CiscoAIDefenseMcpMixin._coerce_to_content_list(inner) content: Final = getattr(response_obj, "content", None) if isinstance(content, list): return content @@ -643,3 +643,6 @@ class _CiscoAIDefenseMcpMixin: if isinstance(direct, dict) and direct: return dict(direct) return None + + +CiscoAIDefenseMcpMixin = _CiscoAIDefenseMcpMixin diff --git a/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/competitor_intent/airline.py b/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/competitor_intent/airline.py index 96365b24410..5e2b85250d2 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/competitor_intent/airline.py +++ b/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/competitor_intent/airline.py @@ -13,11 +13,14 @@ import json from pathlib import Path from typing import Any, Final -from litellm.proxy.guardrails.guardrail_hooks.litellm_content_filter.competitor_intent.base import ( +from litellm.proxy.guardrails.guardrail_hooks.litellm_content_filter.competitor_intent.base import ( # noqa: F401 # legacy module exports BaseCompetitorIntentChecker, - _compile_marker, - _count_signals, - _word_boundary_match, + _compile_marker, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _count_signals, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _word_boundary_match, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + compile_marker, + count_signals, + word_boundary_match, ) # Location/travel context: prepositions, travel verbs, booking nouns, entry/geo nouns. @@ -170,8 +173,8 @@ class AirlineCompetitorIntentChecker(BaseCompetitorIntentChecker): self._other_meaning_signals = list(merged.get("other_meaning_signals") or []) self._competitor_signals = list(merged.get("competitor_signals") or []) self._other_meaning_anchors = list(merged.get("other_meaning_anchors") or []) - self._explicit_competitor_marker = _compile_marker(merged.get("explicit_competitor_marker")) - self._explicit_other_meaning_marker = _compile_marker(merged.get("explicit_other_meaning_marker")) + self._explicit_competitor_marker = compile_marker(merged.get("explicit_competitor_marker")) + self._explicit_other_meaning_marker = compile_marker(merged.get("explicit_other_meaning_marker")) def _classify_ambiguous(self, text: str, token: str) -> tuple[str, float]: """Other meaning vs competitor using airline signals and explicit markers.""" @@ -179,21 +182,25 @@ class AirlineCompetitorIntentChecker(BaseCompetitorIntentChecker): if ( self._explicit_competitor_marker and self._explicit_competitor_marker.search(text_lower) - and _word_boundary_match(text_lower, token.lower()) + and word_boundary_match(text_lower, token.lower()) ): return "COMPETITOR", 0.85 if self._explicit_other_meaning_marker and self._explicit_other_meaning_marker.search(text_lower): return "OTHER_MEANING", 0.85 # Operational-only: baggage/lounge/check-in/refund with no comparison → product query - has_comparison: Final = _count_signals(text_lower, AIRLINE_COMPARISON_SIGNALS) > 0 - operational_count: Final = _count_signals(text_lower, AIRLINE_OPERATIONAL_SIGNALS) + has_comparison: Final = count_signals(text_lower, AIRLINE_COMPARISON_SIGNALS) > 0 + operational_count: Final = count_signals(text_lower, AIRLINE_OPERATIONAL_SIGNALS) if not has_comparison and operational_count > 0: return "OTHER_MEANING", 0.85 # Score: location/travel context vs airline context (no place-name list) - other_count = _count_signals(text_lower, self._other_meaning_signals) + other_count = count_signals( # rebind-ok: pre-existing rebinding on a rename-only line + text_lower, self._other_meaning_signals + ) if self._other_meaning_anchors: - other_count += _count_signals(text_lower, self._other_meaning_anchors) - comp_count: Final = _count_signals(text_lower, self._competitor_signals) + other_count += count_signals( # rebind-ok: pre-existing rebinding on a rename-only line + text_lower, self._other_meaning_anchors + ) + comp_count: Final = count_signals(text_lower, self._competitor_signals) total: Final = other_count + comp_count if total == 0: return "OTHER_MEANING", 0.5 diff --git a/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/competitor_intent/base.py b/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/competitor_intent/base.py index 246e56441e5..eb801a412e0 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/competitor_intent/base.py +++ b/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/competitor_intent/base.py @@ -31,17 +31,23 @@ def normalize(text: str) -> str: return re.sub(r"\s+", " ", t) -def _word_boundary_match(text: str, token: str) -> bool: +def word_boundary_match(text: str, token: str) -> bool: """True if token appears as a word in text.""" return bool(re.search(r"\b" + re.escape(token) + r"\b", text)) -def _count_signals(text: str, patterns: list[str]) -> int: +_word_boundary_match: Final = word_boundary_match + + +def count_signals(text: str, patterns: list[str]) -> int: """Count how many of the patterns appear in text.""" return sum(1 for p in patterns if re.search(p, text, re.IGNORECASE)) -def _compile_marker(pattern: str | None) -> Pattern[str] | None: +_count_signals: Final = count_signals + + +def compile_marker(pattern: str | None) -> Pattern[str] | None: """Compile optional regex string to a pattern.""" if not pattern or not pattern.strip(): return None @@ -51,6 +57,9 @@ def _compile_marker(pattern: str | None) -> Pattern[str] | None: return None +_compile_marker: Final = compile_marker + + def text_for_entity_matching(text: str) -> str: """Letters-only variant for entity matching (e.g. split punctuation).""" t: Final = re.sub(r"[^\w\s]", " ", text) @@ -117,7 +126,7 @@ class BaseCompetitorIntentChecker: found: Final[list[tuple[str, str, bool]]] = [] seen: Final[set[tuple[str, str]]] = set() for token in self._competitor_tokens: - if not _word_boundary_match(normalized, token): + if not word_boundary_match(normalized, token): continue canonical = self.competitor_canonical.get(token, token) key = (token, canonical) @@ -139,7 +148,7 @@ class BaseCompetitorIntentChecker: } for b in self.brand_self: - if _word_boundary_match(normalized, b): + if word_boundary_match(normalized, b): entities["brand_self"].append(b) evidence.append({"type": "entity", "key": "brand_self", "value": b, "match": b}) diff --git a/litellm/proxy/guardrails/guardrail_hooks/mcp_end_user_permission/mcp_end_user_permission.py b/litellm/proxy/guardrails/guardrail_hooks/mcp_end_user_permission/mcp_end_user_permission.py index 9124a98ac36..d7ce3438fca 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/mcp_end_user_permission/mcp_end_user_permission.py +++ b/litellm/proxy/guardrails/guardrail_hooks/mcp_end_user_permission/mcp_end_user_permission.py @@ -193,7 +193,7 @@ class MCPEndUserPermissionGuardrail(CustomGuardrail): MCPRequestHandler, ) - access_group_servers: Final = await MCPRequestHandler._get_mcp_servers_from_access_groups(mcp_access_groups) + access_group_servers: Final = await MCPRequestHandler.get_mcp_servers_from_access_groups(mcp_access_groups) return list(set(direct_mcp_servers + access_group_servers)) diff --git a/litellm/proxy/guardrails/guardrail_hooks/presidio.py b/litellm/proxy/guardrails/guardrail_hooks/presidio.py index e143704a086..893d43ab0ec 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/presidio.py +++ b/litellm/proxy/guardrails/guardrail_hooks/presidio.py @@ -19,7 +19,7 @@ from datetime import datetime from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Protocol, cast import aiohttp -from pydantic import ConfigDict, TypeAdapter, with_config +from pydantic import ConfigDict, JsonValue, TypeAdapter, with_config from typing_extensions import NotRequired, ReadOnly, TypedDict import litellm @@ -279,7 +279,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): self, presidio_analyzer_api_base: str | None = None, presidio_anonymizer_api_base: str | None = None, - ): + ) -> None: self.presidio_analyzer_api_base: str | None = presidio_analyzer_api_base or get_secret( "PRESIDIO_ANALYZER_API_BASE", None ) @@ -922,7 +922,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): def raise_exception_if_blocked_entities_detected( self, analyze_results: list[PresidioAnalyzeResponseItem] | _PresidioAnonymizeResponse - ): + ) -> None: """ Raise an exception if blocked entities are detected """ @@ -1022,7 +1022,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): cache: DualCache, data: dict, call_type: str, - ): + ) -> dict[str, object]: """ - Check if request turned off pii - Check if user allowed to turn off pii (key permissions -> 'allow_pii_controls') @@ -1212,10 +1212,10 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): async def async_post_call_success_hook( self, - data: dict, + data: dict[str, object], user_api_key_dict: UserAPIKeyAuth, response: ModelResponse | EmbeddingResponse | ImageResponse, - ): + ) -> dict[str, JsonValue] | ModelResponse | EmbeddingResponse | ImageResponse: """ Output parse the response object to replace the masked tokens with user sent values """ @@ -1546,7 +1546,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): and delta.get("type") == "text_delta" and isinstance(delta.get("text"), str) ): - unmasked = _OPTIONAL_PresidioPIIMasking._unmask_pii_text(delta["text"], pii_tokens) + unmasked = OPTIONAL_PresidioPIIMasking._unmask_pii_text(delta["text"], pii_tokens) if unmasked != delta["text"]: event["delta"]["text"] = unmasked line = "data: " + json.dumps(event, ensure_ascii=False) @@ -1707,7 +1707,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): return None - def print_verbose(self, print_statement): + def print_verbose(self, print_statement) -> None: try: verbose_proxy_logger.debug(print_statement) if litellm.set_verbose: @@ -1771,3 +1771,6 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): self.presidio_analyze_chunk_size_bytes = self._coerce_analyze_chunk_size( litellm_params.presidio_analyze_chunk_size_bytes ) + + +OPTIONAL_PresidioPIIMasking = _OPTIONAL_PresidioPIIMasking diff --git a/litellm/proxy/guardrails/guardrail_initializers.py b/litellm/proxy/guardrails/guardrail_initializers.py index bacbd728d89..192d55b0b91 100644 --- a/litellm/proxy/guardrails/guardrail_initializers.py +++ b/litellm/proxy/guardrails/guardrail_initializers.py @@ -134,7 +134,7 @@ def _presidio_output_mode(mode: str | list[str] | Mode, *, include_mcp: bool) -> def initialize_presidio(litellm_params: LitellmParams, guardrail: Guardrail) -> tuple[CustomGuardrail, ...]: from litellm.proxy.guardrails.guardrail_hooks.presidio import ( - _OPTIONAL_PresidioPIIMasking, + OPTIONAL_PresidioPIIMasking, ) explicit_filter_scope: Final = litellm_params.presidio_filter_scope @@ -163,7 +163,7 @@ def initialize_presidio(litellm_params: LitellmParams, guardrail: Guardrail) -> params.update(overrides) # Passed outside the heterogeneous params dict so the argument keeps # its precise int | None type. - callback: Final = _OPTIONAL_PresidioPIIMasking( + callback: Final = OPTIONAL_PresidioPIIMasking( presidio_analyze_chunk_size_bytes=litellm_params.presidio_analyze_chunk_size_bytes, **params, ) diff --git a/litellm/proxy/guardrails/guardrail_registry.py b/litellm/proxy/guardrails/guardrail_registry.py index 8a430c71326..1bfe0e5ed77 100644 --- a/litellm/proxy/guardrails/guardrail_registry.py +++ b/litellm/proxy/guardrails/guardrail_registry.py @@ -36,8 +36,9 @@ from litellm.proxy.guardrails.guardrail_hooks.grayswan import ( ) from litellm.proxy.guardrails.guardrail_hooks.lakera_ai import lakeraAI_Moderation from litellm.proxy.guardrails.guardrail_hooks.lakera_ai_v2 import LakeraAIGuardrail -from litellm.proxy.guardrails.guardrail_hooks.presidio import ( - _OPTIONAL_PresidioPIIMasking, +from litellm.proxy.guardrails.guardrail_hooks.presidio import ( # noqa: F401 # legacy module exports + OPTIONAL_PresidioPIIMasking, + _OPTIONAL_PresidioPIIMasking, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export ) from litellm.proxy.guardrails.guardrail_hooks.tool_permission import ( ToolPermissionGuardrail, @@ -226,7 +227,7 @@ guardrail_class_registry: Final[dict[str, type[CustomGuardrail]]] = { SupportedGuardrailIntegrations.GRAYSWAN.value: GraySwanGuardrail, SupportedGuardrailIntegrations.LAKERA.value: lakeraAI_Moderation, SupportedGuardrailIntegrations.LAKERA_V2.value: LakeraAIGuardrail, - SupportedGuardrailIntegrations.PRESIDIO.value: _OPTIONAL_PresidioPIIMasking, + SupportedGuardrailIntegrations.PRESIDIO.value: OPTIONAL_PresidioPIIMasking, SupportedGuardrailIntegrations.TOOL_PERMISSION.value: ToolPermissionGuardrail, } diff --git a/litellm/proxy/health_check.py b/litellm/proxy/health_check.py index 468dc7f19bb..ce30334b8e9 100644 --- a/litellm/proxy/health_check.py +++ b/litellm/proxy/health_check.py @@ -166,7 +166,7 @@ def _get_random_llm_message(): return [{"role": "user", "content": random.choice(messages)}] -def _clean_endpoint_data(endpoint_data: dict, details: bool | None = True): +def clean_endpoint_data(endpoint_data: Mapping[str, object], details: bool | None = True) -> dict[str, object]: """ Keep only the explicitly approved, JSON-safe diagnostic fields for display to users. """ @@ -174,6 +174,9 @@ def _clean_endpoint_data(endpoint_data: dict, details: bool | None = True): return {k: v for k, v in endpoint_data.items() if k in displayed} +_clean_endpoint_data: Final = clean_endpoint_data + + def health_check_filter_kwargs_from_general_settings( general_settings: dict | None, ) -> dict: @@ -543,7 +546,9 @@ async def _run_model_health_check(model: dict): model_info, litellm_params, # any-ok: untyped router config dict ) - litellm_params = _update_litellm_params_for_health_check(model_info, litellm_params) + litellm_params = update_litellm_params_for_health_check( # rebind-ok: pre-existing rebinding on a rename-only line + model_info, litellm_params + ) timeout: Final = model_info.get("health_check_timeout") or HEALTH_CHECK_TIMEOUT_SECONDS return await run_with_timeout( @@ -649,12 +654,12 @@ async def _perform_health_check( _model_id = (model.get("model_info") or {}).get("id") if isinstance(is_healthy, dict) and "error" not in is_healthy: - cleaned = _clean_endpoint_data({**litellm_params, **is_healthy}, details) + cleaned = clean_endpoint_data({**litellm_params, **is_healthy}, details) if _model_id: cleaned["model_id"] = _model_id healthy_endpoints.append(cleaned) elif isinstance(is_healthy, dict): - cleaned = _clean_endpoint_data({**litellm_params, **is_healthy}, details) + cleaned = clean_endpoint_data({**litellm_params, **is_healthy}, details) if _model_id: cleaned["model_id"] = _model_id if "exception" in is_healthy: @@ -665,7 +670,7 @@ async def _perform_health_check( cleaned["exception_status"] = getattr(exc, "status_code", 500) unhealthy_endpoints.append(cleaned) else: - cleaned = _clean_endpoint_data(litellm_params, details) + cleaned = clean_endpoint_data(litellm_params, details) if _model_id: cleaned["model_id"] = _model_id if isinstance(is_healthy, Exception): @@ -772,7 +777,7 @@ def _resolve_health_check_max_tokens(model_info: dict, litellm_params: dict) -> return None -def _update_litellm_params_for_health_check(model_info: dict, litellm_params: dict) -> dict: +def update_litellm_params_for_health_check(model_info: dict, litellm_params: dict) -> dict: """ Update the litellm params for health check. @@ -865,6 +870,9 @@ def _update_litellm_params_for_health_check(model_info: dict, litellm_params: di return litellm_params +_update_litellm_params_for_health_check: Final = update_litellm_params_for_health_check + + async def perform_health_check( model_list: list, model: str | None = None, diff --git a/litellm/proxy/health_endpoints/_health_endpoints.py b/litellm/proxy/health_endpoints/_health_endpoints.py index ff12f75479a..4180ca1a656 100644 --- a/litellm/proxy/health_endpoints/_health_endpoints.py +++ b/litellm/proxy/health_endpoints/_health_endpoints.py @@ -53,15 +53,17 @@ from litellm.proxy.db.health_check_latest import ( query_latest_health_checks, ) from litellm.proxy.db.proxy_worker_heartbeat import count_live_proxy_workers -from litellm.proxy.health_check import ( +from litellm.proxy.health_check import ( # noqa: F401 # legacy module exports ADMIN_ONLY_HEALTH_DISPLAY_PARAMS, - _clean_endpoint_data, - _update_litellm_params_for_health_check, + _clean_endpoint_data, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _update_litellm_params_for_health_check, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + clean_endpoint_data, deployments_targeted_by_name, health_check_filter_kwargs_from_general_settings, perform_health_check, resolve_health_check_mode, run_with_timeout, + update_litellm_params_for_health_check, ) from litellm.proxy.middleware.admission_control_middleware import ( get_admission_control_stats, @@ -634,7 +636,7 @@ async def health_services_endpoint( ) -def _convert_health_check_to_dict(check) -> dict: +def convert_health_check_to_dict(check) -> dict: """Convert health check database record to dictionary format""" return { "health_check_id": check.health_check_id, @@ -652,6 +654,9 @@ def _convert_health_check_to_dict(check) -> dict: } +_convert_health_check_to_dict: Final = convert_health_check_to_dict + + def _check_prisma_client(): """Helper to check if prisma_client is available and raise appropriate error""" from litellm.proxy.proxy_server import prisma_client @@ -883,7 +888,7 @@ async def _save_health_check_results_if_changed( return all(row is not None for row in rows) -async def _save_background_health_checks_to_db( +async def save_background_health_checks_to_db( prisma_client, model_list: list, healthy_endpoints: list, @@ -941,6 +946,9 @@ async def _save_background_health_checks_to_db( return False +_save_background_health_checks_to_db: Final = save_background_health_checks_to_db + + _PROXY_ADMIN_ROLES: Final = frozenset( { LitellmUserRoles.PROXY_ADMIN.value, @@ -1353,7 +1361,7 @@ async def health_check_history_endpoint( ) # Convert to dict format for JSON response using helper function - history_data: Final = [_convert_health_check_to_dict(check) for check in history] + history_data: Final = [convert_health_check_to_dict(check) for check in history] return { "health_checks": history_data, @@ -1385,7 +1393,7 @@ async def latest_health_checks_endpoint( # Convert to dict format for JSON response using helper function checks_data: Final = { - (check.model_id if check.model_id else check.model_name): _convert_health_check_to_dict(check) + (check.model_id if check.model_id else check.model_name): convert_health_check_to_dict(check) for check in latest_checks } @@ -2234,7 +2242,7 @@ async def test_model_connection( stored_params=_OBJECT_MAPPING.validate_python(config_litellm_params), request_params=_OBJECT_MAPPING.validate_python(request_litellm_params), ) - litellm_params = _update_litellm_params_for_health_check( + litellm_params = update_litellm_params_for_health_check( model_info=dict(probe_model_info), litellm_params=litellm_params, ) @@ -2272,7 +2280,7 @@ async def test_model_connection( ) # Clean the result for display - cleaned_result: Final = _clean_endpoint_data({**litellm_params, **result}, details=True) + cleaned_result: Final = clean_endpoint_data({**litellm_params, **result}, details=True) return { "status": "error" if "error" in result else "success", diff --git a/litellm/proxy/hooks/__init__.py b/litellm/proxy/hooks/__init__.py index 0e78a0843cd..138050ccc9a 100644 --- a/litellm/proxy/hooks/__init__.py +++ b/litellm/proxy/hooks/__init__.py @@ -3,35 +3,53 @@ from typing import Final, Literal from . import * from .autorouter_baseline_cache import AutoRouterBaselineCache -from .cache_control_check import _PROXY_CacheControlCheck +from .cache_control_check import ( # noqa: F401 # backwards-compatible package export + PROXY_CacheControlCheck, + _PROXY_CacheControlCheck, # pyright: ignore[reportPrivateUsage] # backwards-compatible package export +) from .litellm_skills import SkillsInjectionHook -from .max_budget_per_session_limiter import _PROXY_MaxBudgetPerSessionHandler -from .max_iterations_limiter import _PROXY_MaxIterationsHandler -from .parallel_request_limiter import _PROXY_MaxParallelRequestsHandler -from .parallel_request_limiter_v3 import _PROXY_MaxParallelRequestsHandler_v3 +from .max_budget_per_session_limiter import ( # noqa: F401 # backwards-compatible package export + PROXY_MaxBudgetPerSessionHandler, + _PROXY_MaxBudgetPerSessionHandler, # pyright: ignore[reportPrivateUsage] # backwards-compatible package export +) +from .max_iterations_limiter import ( # noqa: F401 # backwards-compatible package export + PROXY_MaxIterationsHandler, + _PROXY_MaxIterationsHandler, # pyright: ignore[reportPrivateUsage] # backwards-compatible package export +) +from .parallel_request_limiter import ( # noqa: F401 # backwards-compatible package export + PROXY_MaxParallelRequestsHandler, + _PROXY_MaxParallelRequestsHandler, # pyright: ignore[reportPrivateUsage] # backwards-compatible package export +) +from .parallel_request_limiter_v3 import ( # noqa: F401 # backwards-compatible package export + PROXY_MaxParallelRequestsHandler_v3, + _PROXY_MaxParallelRequestsHandler_v3, # pyright: ignore[reportPrivateUsage] # backwards-compatible package export +) from .prompt_cache_prediction import PromptCacheObserver from .responses_id_security import ResponsesIDSecurity -from .sensitive_data_routing import _PROXY_SensitiveDataRoutingHandler +from .sensitive_data_routing import ( # noqa: F401 # backwards-compatible package export + PROXY_SensitiveDataRoutingHandler, + _PROXY_SensitiveDataRoutingHandler, # pyright: ignore[reportPrivateUsage] # backwards-compatible package export +) # List of all available hooks that can be enabled. # Defined before the enterprise import below so that any module re-imported # transitively through `enterprise.enterprise_hooks` can resolve `PROXY_HOOKS` # and `get_proxy_hook` from this partially-initialized module without circling. PROXY_HOOKS: Final = { - "parallel_request_limiter": _PROXY_MaxParallelRequestsHandler_v3, - "cache_control_check": _PROXY_CacheControlCheck, + "parallel_request_limiter": PROXY_MaxParallelRequestsHandler_v3, + "cache_control_check": PROXY_CacheControlCheck, "responses_id_security": ResponsesIDSecurity, "litellm_skills": SkillsInjectionHook, - "max_iterations_limiter": _PROXY_MaxIterationsHandler, - "max_budget_per_session_limiter": _PROXY_MaxBudgetPerSessionHandler, - "sensitive_data_routing": _PROXY_SensitiveDataRoutingHandler, + "max_iterations_limiter": PROXY_MaxIterationsHandler, + "max_budget_per_session_limiter": PROXY_MaxBudgetPerSessionHandler, + "sensitive_data_routing": PROXY_SensitiveDataRoutingHandler, "prompt_cache_prediction": PromptCacheObserver, "autorouter_baseline_cache": AutoRouterBaselineCache, } ## FEATURE FLAG HOOKS ## if os.getenv("LEGACY_MULTI_INSTANCE_RATE_LIMITING", "false").lower() == "true": - PROXY_HOOKS["parallel_request_limiter"] = _PROXY_MaxParallelRequestsHandler + PROXY_HOOKS["parallel_request_limiter"] = PROXY_MaxParallelRequestsHandler def get_proxy_hook( diff --git a/litellm/proxy/hooks/azure_content_safety.py b/litellm/proxy/hooks/azure_content_safety.py index ad3ec844fac..7d0db84692c 100644 --- a/litellm/proxy/hooks/azure_content_safety.py +++ b/litellm/proxy/hooks/azure_content_safety.py @@ -77,7 +77,7 @@ class _PROXY_AzureContentSafety( return result - async def test_violation(self, content: str, source: str | None = None): + async def test_violation(self, content: str, source: str | None = None) -> None: verbose_proxy_logger.debug("Testing Azure Content-Safety for: %s", content) # Construct a request @@ -115,7 +115,7 @@ class _PROXY_AzureContentSafety( cache: DualCache, data: dict, call_type: str, # "completion", "embeddings", "image_generation", "moderation" - ): + ) -> None: verbose_proxy_logger.debug("Inside Azure Content-Safety Pre-Call Hook") try: if is_text_content_call_type(call_type): @@ -135,7 +135,7 @@ class _PROXY_AzureContentSafety( data: dict, user_api_key_dict: UserAPIKeyAuth, response, - ): + ) -> None: verbose_proxy_logger.debug("Inside Azure Content-Safety Post-Call Hook") if not isinstance(response, litellm.ModelResponse): return @@ -148,10 +148,13 @@ class _PROXY_AzureContentSafety( if isinstance(content, str): await self.test_violation(content=content, source="output") - # async def async_post_call_streaming_hook( - # self, - # user_api_key_dict: UserAPIKeyAuth, - # response: str, - # ): - # verbose_proxy_logger.debug("Inside Azure Content-Safety Call-Stream Hook") - # await self.test_violation(content=response, source="output") + +PROXY_AzureContentSafety: Final = _PROXY_AzureContentSafety + +# async def async_post_call_streaming_hook( +# self, +# user_api_key_dict: UserAPIKeyAuth, +# response: str, +# ): +# verbose_proxy_logger.debug("Inside Azure Content-Safety Call-Stream Hook") +# await self.test_violation(content=response, source="output") diff --git a/litellm/proxy/hooks/batch_enqueued_tokens.py b/litellm/proxy/hooks/batch_enqueued_tokens.py index c4591eb17fe..2a1719ed4c4 100644 --- a/litellm/proxy/hooks/batch_enqueued_tokens.py +++ b/litellm/proxy/hooks/batch_enqueued_tokens.py @@ -148,12 +148,12 @@ def resolve_batch_enqueued_token_scopes( def canonical_provider_batch_id(batch_id: str) -> str: from litellm.proxy.openai_files_endpoints.common_utils import ( - _is_base64_encoded_unified_file_id, # pyright: ignore[reportPrivateUsage] # canonical unified-id decoder has no public wrapper get_batch_id_from_unified_batch_id, get_original_file_id, + is_base64_encoded_unified_file_id, # pyright: ignore[reportPrivateUsage] # canonical unified-id decoder has no public wrapper ) - decoded: Final = _is_base64_encoded_unified_file_id(batch_id) + decoded: Final = is_base64_encoded_unified_file_id(batch_id) if isinstance(decoded, str): if "llm_batch_id" in decoded or "generic_response_id" in decoded: return get_batch_id_from_unified_batch_id(decoded) diff --git a/litellm/proxy/hooks/batch_rate_limiter.py b/litellm/proxy/hooks/batch_rate_limiter.py index 55bd622651e..17f8bd277dd 100644 --- a/litellm/proxy/hooks/batch_rate_limiter.py +++ b/litellm/proxy/hooks/batch_rate_limiter.py @@ -68,15 +68,15 @@ if TYPE_CHECKING: from opentelemetry.trace import Span as _Span from litellm.caching.caching import DualCache + from litellm.proxy.hooks.parallel_request_limiter_v3 import ( + PROXY_MaxParallelRequestsHandler_v3 as _ParallelRequestLimiter, + ) from litellm.proxy.hooks.parallel_request_limiter_v3 import ( RateLimitDescriptor as _RateLimitDescriptor, ) from litellm.proxy.hooks.parallel_request_limiter_v3 import ( RateLimitStatus as _RateLimitStatus, ) - from litellm.proxy.hooks.parallel_request_limiter_v3 import ( - _PROXY_MaxParallelRequestsHandler_v3 as _ParallelRequestLimiter, - ) from litellm.proxy.utils import InternalUsageCache as _InternalUsageCache from litellm.router import Router as _Router from litellm.types.llms.openai import HttpxBinaryResponseContent @@ -164,16 +164,16 @@ class _PROXY_BatchRateLimiter(CustomLogger): return None from litellm.proxy.openai_files_endpoints.common_utils import ( - _is_base64_encoded_unified_file_id, decode_model_from_file_id, get_models_from_unified_file_id, + is_base64_encoded_unified_file_id, ) model_from_file_id: Final = decode_model_from_file_id(input_file_id) if model_from_file_id: return model_from_file_id - unified_file_id: Final = _is_base64_encoded_unified_file_id(input_file_id) + unified_file_id: Final = is_base64_encoded_unified_file_id(input_file_id) if unified_file_id: target_model_names: Final = get_models_from_unified_file_id(unified_file_id) if target_model_names: @@ -250,7 +250,7 @@ class _PROXY_BatchRateLimiter(CustomLogger): minute with the submission. The daily descriptor uses its own key so its 24h window never collides with the online limiter's counters. """ - descriptors: Final = self.parallel_request_limiter._create_rate_limit_descriptors( + descriptors: Final = self.parallel_request_limiter.create_rate_limit_descriptors( user_api_key_dict=user_api_key_dict, data=data, rpm_limit_type=None, @@ -892,13 +892,13 @@ class _PROXY_BatchRateLimiter(CustomLogger): try: # Check if this is a managed file (base64 encoded unified file ID) from litellm.proxy.openai_files_endpoints.common_utils import ( - _is_base64_encoded_unified_file_id, get_models_from_unified_file_id, + is_base64_encoded_unified_file_id, ) # Managed files require bypassing the HTTP endpoint (which runs access-check hooks) # and calling the managed files hook directly with the user's credentials. - is_managed_file: Final = _is_base64_encoded_unified_file_id(file_id) + is_managed_file: Final = is_base64_encoded_unified_file_id(file_id) # For managed files the unified file id encodes the proxy model # alias(es) the file was uploaded for; auth validates against those. target_model_names: Final = get_models_from_unified_file_id(is_managed_file) if is_managed_file else [] @@ -1039,11 +1039,11 @@ class _PROXY_BatchRateLimiter(CustomLogger): enforces on `/chat/completions` apply here. """ from litellm.proxy.auth.auth_checks import ( - _check_team_member_model_access, - _key_access_group_grants_model, can_key_call_model, can_team_access_model, + check_team_member_model_access, get_team_object, + key_access_group_grants_model, ) from litellm.proxy.proxy_server import llm_router, prisma_client, proxy_logging_obj, user_api_key_cache @@ -1092,14 +1092,14 @@ class _PROXY_BatchRateLimiter(CustomLogger): except ProxyException as team_denial: if team_denial.type != ProxyErrorTypes.team_model_access_denied: raise - if not await _key_access_group_grants_model( + if not await key_access_group_grants_model( model=model_to_check, valid_token=user_api_key_dict, team_object=team_object, llm_router=llm_router, ): raise - await _check_team_member_model_access( + await check_team_member_model_access( model=model_to_check, team_object=team_object, valid_token=user_api_key_dict, @@ -1281,3 +1281,6 @@ class _PROXY_BatchRateLimiter(CustomLogger): verbose_proxy_logger.error("Error in batch rate limiting: %s", e, exc_info=True) # Don't block the request if rate limiting fails return data + + +PROXY_BatchRateLimiter = _PROXY_BatchRateLimiter diff --git a/litellm/proxy/hooks/batch_redis_get.py b/litellm/proxy/hooks/batch_redis_get.py index 94ad6fb7cb2..94eb0bb52d1 100644 --- a/litellm/proxy/hooks/batch_redis_get.py +++ b/litellm/proxy/hooks/batch_redis_get.py @@ -25,7 +25,7 @@ class _PROXY_BatchRedisRequests(CustomLogger): self.async_get_cache ) # map the litellm 'get_cache' function to our custom function - def print_verbose(self, print_statement, debug_level: Literal["INFO", "DEBUG"] = "DEBUG"): + def print_verbose(self, print_statement, debug_level: Literal["INFO", "DEBUG"] = "DEBUG") -> None: if debug_level == "DEBUG" or debug_level == "INFO": verbose_proxy_logger.debug(print_statement) if litellm.set_verbose is True: @@ -37,7 +37,7 @@ class _PROXY_BatchRedisRequests(CustomLogger): cache: DualCache, data: dict, call_type: str, - ): + ) -> None: try: """ Get the user key @@ -83,7 +83,7 @@ class _PROXY_BatchRedisRequests(CustomLogger): ) verbose_proxy_logger.debug(traceback.format_exc()) - async def async_get_cache(self, *args, **kwargs): + async def async_get_cache(self, *args, **kwargs) -> object | None: """ - Check if the cache key is in-memory @@ -113,3 +113,6 @@ class _PROXY_BatchRedisRequests(CustomLogger): return litellm.cache.get_cache_logic(cached_result=cached_result, max_age=max_age) except Exception: return None + + +PROXY_BatchRedisRequests: Final = _PROXY_BatchRedisRequests diff --git a/litellm/proxy/hooks/cache_control_check.py b/litellm/proxy/hooks/cache_control_check.py index dab2ed0b933..22af491f2e7 100644 --- a/litellm/proxy/hooks/cache_control_check.py +++ b/litellm/proxy/hooks/cache_control_check.py @@ -24,7 +24,7 @@ class _PROXY_CacheControlCheck(CustomLogger): cache: DualCache, data: dict, call_type: str, - ): + ) -> None: try: verbose_proxy_logger.debug("Inside Cache Control Check Pre-Call Hook") allowed_cache_controls: Final = user_api_key_dict.allowed_cache_controls @@ -56,3 +56,6 @@ class _PROXY_CacheControlCheck(CustomLogger): verbose_logger.exception( "litellm.proxy.hooks.cache_control_check.py::async_pre_call_hook(): Exception occured - %s", e ) + + +PROXY_CacheControlCheck: Final = _PROXY_CacheControlCheck diff --git a/litellm/proxy/hooks/dynamic_rate_limiter.py b/litellm/proxy/hooks/dynamic_rate_limiter.py index 5c99a1bacc1..646afab73d6 100644 --- a/litellm/proxy/hooks/dynamic_rate_limiter.py +++ b/litellm/proxy/hooks/dynamic_rate_limiter.py @@ -1,7 +1,6 @@ # What is this? ## Allocates dynamic tpm/rpm quota for a project based on current traffic ## Tracks num active projects per minute - import asyncio import os from collections.abc import Callable @@ -23,7 +22,7 @@ from litellm.proxy.hooks.rate_limiter_utils import ( resolve_llm_provider_for_rate_limit, ) from litellm.types.router import ModelGroupInfo -from litellm.types.utils import CallTypesLiteral +from litellm.types.utils import CallTypesLiteral, LLMResponseTypes from litellm.utils import get_utc_datetime @@ -83,7 +82,7 @@ class _PROXY_DynamicRateLimitHandler(CustomLogger): def __init__(self, internal_usage_cache: DualCache, time_fn: Callable[[], datetime] = get_utc_datetime): self.internal_usage_cache = DynamicRateLimiterCache(cache=internal_usage_cache, time_fn=time_fn) - def update_variables(self, llm_router: Router): + def update_variables(self, llm_router: Router) -> None: self.llm_router = llm_router @with_service_target("rate_limits") @@ -241,7 +240,9 @@ class _PROXY_DynamicRateLimitHandler(CustomLogger): return None @with_service_target("rate_limits") - async def async_post_call_success_hook(self, data: dict, user_api_key_dict: UserAPIKeyAuth, response): + async def async_post_call_success_hook( + self, data: dict, user_api_key_dict: UserAPIKeyAuth, response: LLMResponseTypes + ) -> LLMResponseTypes | None: try: if isinstance(response, ModelResponse): model_id: Final = response.hidden_params["model_id"] @@ -281,3 +282,6 @@ class _PROXY_DynamicRateLimitHandler(CustomLogger): "litellm.proxy.hooks.dynamic_rate_limiter.py::async_post_call_success_hook(): Exception occured - %s", e ) return response + + +PROXY_DynamicRateLimitHandler: Final = _PROXY_DynamicRateLimitHandler diff --git a/litellm/proxy/hooks/dynamic_rate_limiter_v3.py b/litellm/proxy/hooks/dynamic_rate_limiter_v3.py index 0a078321d5e..d174d9fecb1 100644 --- a/litellm/proxy/hooks/dynamic_rate_limiter_v3.py +++ b/litellm/proxy/hooks/dynamic_rate_limiter_v3.py @@ -20,11 +20,12 @@ from litellm.proxy.common_utils.proxy_rate_limit_error import ( ProxyRateLimitError, map_v3_rate_limit_type, ) -from litellm.proxy.hooks.parallel_request_limiter_v3 import ( +from litellm.proxy.hooks.parallel_request_limiter_v3 import ( # noqa: F401 # legacy module exports + PROXY_MaxParallelRequestsHandler_v3, RateLimitDescriptor, RateLimitDescriptorRateLimitObject, RateLimitResponse, - _PROXY_MaxParallelRequestsHandler_v3, + _PROXY_MaxParallelRequestsHandler_v3, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export claim_request_stash_for_data, get_or_create_request_stash, ) @@ -38,7 +39,7 @@ from litellm.router_utils.add_retry_fallback_headers import ( response_has_hidden_params, ) from litellm.types.router import ModelGroupInfo -from litellm.types.utils import CallTypesLiteral +from litellm.types.utils import CallTypesLiteral, LLMResponseTypes if TYPE_CHECKING: from litellm.types.utils import PriorityReservationSettings @@ -90,9 +91,9 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger): time_provider: Callable[[], datetime] | None = None, ): self.internal_usage_cache = InternalUsageCache(dual_cache=internal_usage_cache) - self.v3_limiter = _PROXY_MaxParallelRequestsHandler_v3(self.internal_usage_cache, time_provider=time_provider) + self.v3_limiter = PROXY_MaxParallelRequestsHandler_v3(self.internal_usage_cache, time_provider=time_provider) - def update_variables(self, llm_router: Router): + def update_variables(self, llm_router: Router) -> None: self.llm_router = llm_router def _get_saturation_check_cache_ttl(self) -> int: @@ -659,7 +660,9 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger): return None @with_service_target("rate_limits") - async def async_post_call_success_hook(self, data: dict, user_api_key_dict: UserAPIKeyAuth, response): + async def async_post_call_success_hook( + self, data: dict, user_api_key_dict: UserAPIKeyAuth, response + ) -> LLMResponseTypes: """ Post-call hook to add rate limit headers to response. Leverages v3 limiter's post-call hook functionality. @@ -689,7 +692,7 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger): return response @with_service_target("rate_limits") - async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): + async def async_log_success_event(self, kwargs, response_obj, start_time, end_time) -> None: """ Update token usage for priority-based rate limiting after successful API calls. @@ -804,3 +807,6 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger): except Exception as e: verbose_proxy_logger.exception("Error in dynamic rate limiter success event: %s", e) + + +PROXY_DynamicRateLimitHandlerV3: Final = _PROXY_DynamicRateLimitHandlerV3 diff --git a/litellm/proxy/hooks/key_management_event_hooks.py b/litellm/proxy/hooks/key_management_event_hooks.py index a1100474671..ef93809d47b 100644 --- a/litellm/proxy/hooks/key_management_event_hooks.py +++ b/litellm/proxy/hooks/key_management_event_hooks.py @@ -21,7 +21,10 @@ from litellm.proxy._types import ( UpdateKeyRequest, UserAPIKeyAuth, ) -from litellm.proxy.utils import _hash_token_if_needed +from litellm.proxy.utils import ( # noqa: F401 # legacy module exports + _hash_token_if_needed, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + hash_token_if_needed, +) from litellm.secret_managers.base_secret_manager import BaseSecretManager if TYPE_CHECKING: @@ -140,7 +143,7 @@ class KeyManagementEventHooks: ), changed_by_api_key=user_api_key_dict.api_key, table_name=LitellmTableNames.KEY_TABLE_NAME, - object_id=_hash_token_if_needed(data.key), + object_id=hash_token_if_needed(data.key), action="updated", updated_values=json.dumps(updated_fields, default=str), before_value=json.dumps(existing_key_row.json(exclude_none=True), default=str), diff --git a/litellm/proxy/hooks/max_budget_per_session_limiter.py b/litellm/proxy/hooks/max_budget_per_session_limiter.py index 2ba42c43dcc..cf3025d8c7b 100644 --- a/litellm/proxy/hooks/max_budget_per_session_limiter.py +++ b/litellm/proxy/hooks/max_budget_per_session_limiter.py @@ -130,7 +130,7 @@ class _PROXY_MaxBudgetPerSessionHandler(CustomLogger): return None @with_service_target("session_budgets") - async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): + async def async_log_success_event(self, kwargs, response_obj, start_time, end_time) -> None: """ After a successful LLM call, increment the session spend by the response cost. """ @@ -271,3 +271,6 @@ class _PROXY_MaxBudgetPerSessionHandler(CustomLogger): local_only=True, ) return new_value + + +PROXY_MaxBudgetPerSessionHandler: Final = _PROXY_MaxBudgetPerSessionHandler diff --git a/litellm/proxy/hooks/max_iterations_limiter.py b/litellm/proxy/hooks/max_iterations_limiter.py index efcafc1b6b0..5b7d7d04e24 100644 --- a/litellm/proxy/hooks/max_iterations_limiter.py +++ b/litellm/proxy/hooks/max_iterations_limiter.py @@ -221,3 +221,6 @@ class _PROXY_MaxIterationsHandler(CustomLogger): local_only=True, ) return new_value + + +PROXY_MaxIterationsHandler: Final = _PROXY_MaxIterationsHandler diff --git a/litellm/proxy/hooks/mcp_semantic_filter/hook.py b/litellm/proxy/hooks/mcp_semantic_filter/hook.py index 15c28d64d05..d33811d3dea 100644 --- a/litellm/proxy/hooks/mcp_semantic_filter/hook.py +++ b/litellm/proxy/hooks/mcp_semantic_filter/hook.py @@ -169,7 +169,7 @@ class SemanticToolFilterHook(CustomLogger): def _selected_tool_names(self, filtered_tools: Sequence[object]) -> list[str]: """Names of the semantically selected tools, as produced by the MCP expansion.""" - names: Final = (self.filter._extract_tool_info(tool)[0] for tool in filtered_tools) + names: Final = (self.filter.extract_tool_info(tool)[0] for tool in filtered_tools) return [name for name in names if name] @staticmethod @@ -219,7 +219,7 @@ class SemanticToolFilterHook(CustomLogger): return False if isinstance(tool, dict) and tool.get("type") == "function" and isinstance(tool.get("name"), str): return False - name, _ = self.filter._extract_tool_info(tool) + name, _ = self.filter.extract_tool_info(tool) return bool(name) and name in self.filter._tool_map def _get_metadata_variable_name(self, data: dict) -> str: @@ -397,14 +397,14 @@ class SemanticToolFilterHook(CustomLogger): filtered_mcp_names: Final[set[str]] = set() for t in filtered_mcp_tools: - name, _ = self.filter._extract_tool_info(t) + name, _ = self.filter.extract_tool_info(t) if name: filtered_mcp_names.add(name) filtered_tools: Final[list[object]] = [] for i, t in enumerate(tools): if i in mcp_indices: - name, _ = self.filter._extract_tool_info(t) + name, _ = self.filter.extract_tool_info(t) if name in filtered_mcp_names: filtered_tools.append(t) else: diff --git a/litellm/proxy/hooks/model_max_budget_limiter.py b/litellm/proxy/hooks/model_max_budget_limiter.py index e810b98f336..8bca51784f2 100644 --- a/litellm/proxy/hooks/model_max_budget_limiter.py +++ b/litellm/proxy/hooks/model_max_budget_limiter.py @@ -493,7 +493,7 @@ class _PROXY_VirtualKeyModelMaxBudgetLimiter(RouterBudgetLimiting): return healthy_deployments @with_service_target("model_budgets") - async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): + async def async_log_success_event(self, kwargs, response_obj, start_time, end_time) -> None: """ Track spend for virtual key + model in DualCache @@ -627,3 +627,6 @@ class _PROXY_VirtualKeyModelMaxBudgetLimiter(RouterBudgetLimiting): key=marker_key, value=1, ttl=ttl_seconds, refresh_ttl=True ) return polls == 1 + + +PROXY_VirtualKeyModelMaxBudgetLimiter = _PROXY_VirtualKeyModelMaxBudgetLimiter diff --git a/litellm/proxy/hooks/parallel_request_limiter.py b/litellm/proxy/hooks/parallel_request_limiter.py index facae4cbb81..0f884387673 100644 --- a/litellm/proxy/hooks/parallel_request_limiter.py +++ b/litellm/proxy/hooks/parallel_request_limiter.py @@ -57,7 +57,7 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): def __init__(self, internal_usage_cache: InternalUsageCache): self.internal_usage_cache = internal_usage_cache - def print_verbose(self, print_statement): + def print_verbose(self, print_statement) -> None: try: verbose_proxy_logger.debug(print_statement) if litellm.set_verbose: @@ -253,7 +253,7 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): cache: DualCache, data: dict, call_type: str, - ): + ) -> None: self.print_verbose("Inside Max Parallel Request Pre-Call Hook") api_key: Final = LiteLLMProxyRequestSetup.get_logged_api_key(user_api_key_dict) max_parallel_requests = user_api_key_dict.max_parallel_requests @@ -494,7 +494,7 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): ) @with_service_target("rate_limits") - async def async_log_success_event(self, kwargs, response_obj: object, start_time, end_time): + async def async_log_success_event(self, kwargs, response_obj: object, start_time, end_time) -> None: from litellm.proxy.common_utils.callback_utils import ( get_model_group_from_litellm_kwargs, ) @@ -700,7 +700,7 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): self.print_verbose(e) @with_service_target("rate_limits") - async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time): + async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time) -> None: try: self.print_verbose("Inside Max Parallel Request Failure Hook") litellm_parent_otel_span: Final[Span | None] = get_parent_otel_span_from_kwargs(kwargs=kwargs) @@ -808,7 +808,7 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): return None @with_service_target("rate_limits") - async def async_post_call_success_hook(self, data: dict, user_api_key_dict: UserAPIKeyAuth, response): + async def async_post_call_success_hook(self, data: dict, user_api_key_dict: UserAPIKeyAuth, response) -> None: """ Retrieve the key's remaining rate limits. """ @@ -861,3 +861,6 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): ) return await super().async_post_call_success_hook(data, user_api_key_dict, response) + + +PROXY_MaxParallelRequestsHandler: Final = _PROXY_MaxParallelRequestsHandler diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index 1a3452d9f87..5b865f3b5c3 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -867,10 +867,10 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): if self._batch_rate_limiter is None: try: from litellm.proxy.hooks.batch_rate_limiter import ( - _PROXY_BatchRateLimiter, + PROXY_BatchRateLimiter, ) - self._batch_rate_limiter = _PROXY_BatchRateLimiter( + self._batch_rate_limiter = PROXY_BatchRateLimiter( internal_usage_cache=self.internal_usage_cache, parallel_request_limiter=self, time_provider=self._time_provider, @@ -1022,17 +1022,17 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): """ if not isinstance(data, dict): return - base_capped_floor: Final = _PROXY_MaxParallelRequestsHandler_v3.no_max_tokens_output_floor(min_configured_limit) + base_capped_floor: Final = PROXY_MaxParallelRequestsHandler_v3.no_max_tokens_output_floor(min_configured_limit) capped_floor: Final = ( max(base_capped_floor, RESPONSES_API_MIN_OUTPUT_TOKENS) if call_type in RESPONSES_API_CALL_TYPES else base_capped_floor ) baseline_floor: Final = DEFAULT_MAX_TOKENS_ESTIMATE // _TPM_FLOOR_FRACTION - is_embedding: Final = _PROXY_MaxParallelRequestsHandler_v3._is_embedding_request(data, call_type) + is_embedding: Final = PROXY_MaxParallelRequestsHandler_v3._is_embedding_request(data, call_type) if ( capped_floor >= baseline_floor - or _PROXY_MaxParallelRequestsHandler_v3._has_explicit_output_cap(data, call_type) + or PROXY_MaxParallelRequestsHandler_v3._has_explicit_output_cap(data, call_type) or is_embedding or endpoint_type == EndpointType.DECISIONS ): @@ -3188,7 +3188,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): if (limit := tag_limits.get(tag)) is not None ) - def _create_rate_limit_descriptors( + def create_rate_limit_descriptors( self, user_api_key_dict: UserAPIKeyAuth, data: dict, @@ -3337,6 +3337,8 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): return descriptors + _create_rate_limit_descriptors = create_rate_limit_descriptors + async def _check_model_has_recent_failures( self, model: str, @@ -3943,7 +3945,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): if requested_model and self._is_dynamic_rate_limiting_enabled(rpm_limit_type, tpm_limit_type) else False ) - descriptors: Final = self._create_rate_limit_descriptors( # pyright: ignore[reportUnknownMemberType] # legacy helper reads a dictionary with validated keys + descriptors: Final = self.create_rate_limit_descriptors( # pyright: ignore[reportUnknownMemberType] # legacy helper reads a dictionary with validated keys user_api_key_dict=user_api_key_dict, data=dict(data), rpm_limit_type=rpm_limit_type, @@ -4031,7 +4033,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): data: dict, call_type: str, endpoint_type: EndpointType = EndpointType.GENERIC, - ): + ) -> Exception | str | dict[str, object] | None: """ Pre-call hook to check rate limits before making the API call. Supports dynamic rate limiting based on deployment health. @@ -5048,7 +5050,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): return pipeline_operations @with_service_target("rate_limits") - async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): + async def async_log_success_event(self, kwargs, response_obj, start_time, end_time) -> None: """ Update TPM usage on successful API calls by incrementing counters using pipeline """ @@ -5166,7 +5168,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): ) @with_service_target("rate_limits") - async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time): + async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time) -> None: """ On failure: decrement max_parallel_requests and refund the upfront TPM reservation only against the scopes the reservation actually @@ -5295,7 +5297,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): await self._release_stashed_parallel_slot(get_request_stash(), None) @with_service_target("rate_limits") - async def async_post_call_success_hook(self, data: dict, user_api_key_dict: UserAPIKeyAuth, response): + async def async_post_call_success_hook(self, data: dict, user_api_key_dict: UserAPIKeyAuth, response) -> None: """ Release completed-request slots and update rate limit headers in the response. """ @@ -5476,3 +5478,6 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): except Exception as e: verbose_proxy_logger.exception("Error releasing TPM reservation on post-call failure: %s", e) return + + +PROXY_MaxParallelRequestsHandler_v3 = _PROXY_MaxParallelRequestsHandler_v3 diff --git a/litellm/proxy/hooks/prompt_injection_detection.py b/litellm/proxy/hooks/prompt_injection_detection.py index 3c2eefcc933..0fee5b061f6 100644 --- a/litellm/proxy/hooks/prompt_injection_detection.py +++ b/litellm/proxy/hooks/prompt_injection_detection.py @@ -76,7 +76,7 @@ class _OPTIONAL_PromptInjectionDetection(CustomLogger): "and start from scratch", ] - def print_verbose(self, print_statement, level: Literal["INFO", "DEBUG"] = "DEBUG"): + def print_verbose(self, print_statement, level: Literal["INFO", "DEBUG"] = "DEBUG") -> None: if level == "INFO": verbose_proxy_logger.info(print_statement) elif level == "DEBUG": @@ -85,7 +85,7 @@ class _OPTIONAL_PromptInjectionDetection(CustomLogger): if litellm.set_verbose is True: print(print_statement) # noqa: T201 - def update_environment(self, router: Router | None = None): + def update_environment(self, router: Router | None = None) -> None: self.llm_router = router if self.prompt_injection_params is not None and self.prompt_injection_params.llm_api_check is True: @@ -150,9 +150,9 @@ class _OPTIONAL_PromptInjectionDetection(CustomLogger): self, user_api_key_dict: UserAPIKeyAuth, cache: DualCache, - data: dict, + data: dict[str, object], call_type: str, # "completion", "embeddings", "image_generation", "moderation" - ): + ) -> dict[str, object] | str | None: try: """ - check if user id part of call @@ -278,3 +278,6 @@ class _OPTIONAL_PromptInjectionDetection(CustomLogger): ) return is_prompt_attack + + +OPTIONAL_PromptInjectionDetection = _OPTIONAL_PromptInjectionDetection diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index ac2eff7ae6a..15b5d5aff66 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -45,9 +45,10 @@ from litellm.proxy.spend_tracking.spend_log_error_logger import ( should_suppress_spend_log_tracebacks, spend_log_error, ) -from litellm.proxy.spend_tracking.spend_tracking_utils import ( - _sanitize_error_information_for_spend_logs, +from litellm.proxy.spend_tracking.spend_tracking_utils import ( # noqa: F401 # legacy module exports + _sanitize_error_information_for_spend_logs, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export get_request_model_access_groups, + sanitize_error_information_for_spend_logs, should_store_prompts_and_responses_in_spend_logs, ) from litellm.proxy.utils import ProxyUpdateSpend @@ -152,7 +153,7 @@ class _ProxyDBLogger(CustomLogger): original_exception: Exception, user_api_key_dict: UserAPIKeyAuth, traceback_str: str | None = None, - ): + ) -> None: try: await _release_budget_reservation(budget_reservation=user_api_key_dict.budget_reservation) except Exception: @@ -168,7 +169,7 @@ class _ProxyDBLogger(CustomLogger): request_route: Final = user_api_key_dict.request_route if ( - _ProxyDBLogger._should_track_errors_in_db() is False + ProxyDBLogger._should_track_errors_in_db() is False or request_route is not None and not ( RouteChecks.is_llm_api_route(route=request_route) or RouteChecks.is_info_route(route=request_route) @@ -197,11 +198,11 @@ class _ProxyDBLogger(CustomLogger): # here because the input above is constructed non-None. _error_information = cast( StandardLoggingPayloadErrorInformation, - _sanitize_error_information_for_spend_logs(_error_information, original_exception=original_exception), + sanitize_error_information_for_spend_logs(_error_information, original_exception=original_exception), ) _metadata["error_information"] = _error_information - _metadata = await _ProxyDBLogger._enrich_failure_metadata_unless_db_stalled( + _metadata = await ProxyDBLogger._enrich_failure_metadata_unless_db_stalled( # rebind-ok: pre-existing rebinding on a rename-only line metadata=_metadata, original_exception=original_exception ) @@ -316,7 +317,7 @@ class _ProxyDBLogger(CustomLogger): # Only fetch key details when user_id wasn't already populated (e.g. direct MCP REST calls). # Avoids a cache/DB lookup on every normal LLM request. if metadata.get("user_api_key") and not metadata.get("user_api_key_user_id"): - metadata = await _ProxyDBLogger._enrich_failure_metadata_with_key_info( # rebind-ok: enriched metadata replaces the original + metadata = await ProxyDBLogger._enrich_failure_metadata_with_key_info( # rebind-ok: enriched metadata replaces the original metadata=metadata, resolve_missing_key_identity=str(kwargs.get("call_type")) not in _CAPTURED_IDENTITY_CALL_TYPES, ) @@ -503,7 +504,7 @@ class _ProxyDBLogger(CustomLogger): ) -> dict[str, object]: if isinstance(original_exception, DBLookupDeadlineExceeded): return metadata - return await _ProxyDBLogger._enrich_failure_metadata_with_key_info(metadata=metadata) + return await ProxyDBLogger._enrich_failure_metadata_with_key_info(metadata=metadata) @staticmethod async def _enrich_failure_metadata_with_key_info(metadata: dict, resolve_missing_key_identity: bool = True) -> dict: @@ -594,6 +595,9 @@ class _ProxyDBLogger(CustomLogger): return +ProxyDBLogger: Final = _ProxyDBLogger + + def _write_spend_metadata_to_kwargs(kwargs: dict, metadata: dict) -> None: patch = {k: v for k, v in metadata.items() if (k.startswith("user_api_key") or k == "tags") and v is not None} if not patch: @@ -609,7 +613,7 @@ def _write_spend_metadata_to_kwargs(kwargs: dict, metadata: dict) -> None: async def run_spend_event(line: bytes) -> None: - await _ProxyDBLogger().run_spend_event(line) + await ProxyDBLogger().run_spend_event(line) def _is_unbilled_interaction_response(completion_response: object) -> bool: diff --git a/litellm/proxy/hooks/sensitive_data_routing.py b/litellm/proxy/hooks/sensitive_data_routing.py index 1773fc2d50a..0909f36dbe5 100644 --- a/litellm/proxy/hooks/sensitive_data_routing.py +++ b/litellm/proxy/hooks/sensitive_data_routing.py @@ -204,3 +204,6 @@ class _PROXY_SensitiveDataRoutingHandler(CustomLogger): data["metadata"] = metadata return data + + +PROXY_SensitiveDataRoutingHandler: Final = _PROXY_SensitiveDataRoutingHandler diff --git a/litellm/proxy/image_endpoints/endpoints.py b/litellm/proxy/image_endpoints/endpoints.py index cd64363f3d3..28fb9baf2aa 100644 --- a/litellm/proxy/image_endpoints/endpoints.py +++ b/litellm/proxy/image_endpoints/endpoints.py @@ -280,11 +280,11 @@ async def image_edit_api( # The validation will be done at the model level if image is truly required from litellm.proxy.proxy_server import ( - _read_request_body, general_settings, llm_router, proxy_config, proxy_logging_obj, + read_request_body, select_data_generator, user_api_base, user_max_tokens, @@ -298,7 +298,7 @@ async def image_edit_api( # Read request body and convert UploadFiles to BytesIO ######################################################### form_fields: Final = coerce_numeric_form_fields( - parsed_body=await _read_request_body(request=request), + parsed_body=await read_request_body(request=request), numeric_fields=IMAGE_EDIT_NUMERIC_FORM_FIELDS, ) data: Final = { @@ -345,7 +345,7 @@ async def image_edit_api( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, diff --git a/litellm/proxy/lens/endpoints.py b/litellm/proxy/lens/endpoints.py index 746fba26316..8e83a1ad6f9 100644 --- a/litellm/proxy/lens/endpoints.py +++ b/litellm/proxy/lens/endpoints.py @@ -194,16 +194,37 @@ async def service_connection(auth: Auth) -> ServiceConnection: public_url: Final = os.environ.get("LITELLM_LENS_PUBLIC_URL", "").rstrip("/") try: connection: Final = LensConnection.from_env() + except ValueError: + return ServiceConnection( + url=public_url, + connected=False, + status=ServiceStatus(), + configured=bool(os.environ.get("LITELLM_LENS_URL")), + release=release_tag(), + ) + try: client: Final = connection.control_client() async with client.stream( "GET", connection.endpoint("/internal/status"), headers=connection.headers, timeout=2 ) as response: if response.status_code == 200: status: Final = ServiceStatus.model_validate_json(await bounded_response(response, 16 * 1024)) - return ServiceConnection(url=public_url, connected=True, status=status) + return ServiceConnection( + url=public_url, + connected=True, + status=status, + configured=True, + release=release_tag(), + ) except (ValueError, RuntimeError, httpx.HTTPError): pass - return ServiceConnection(url=public_url, connected=False, status=ServiceStatus()) + return ServiceConnection( + url=public_url, + connected=False, + status=ServiceStatus(), + configured=True, + release=release_tag(), + ) async def credential_snapshot() -> IngestionSnapshot: diff --git a/litellm/proxy/lens/feedback_endpoints.py b/litellm/proxy/lens/feedback_endpoints.py new file mode 100644 index 00000000000..f9479873a84 --- /dev/null +++ b/litellm/proxy/lens/feedback_endpoints.py @@ -0,0 +1,124 @@ +from datetime import datetime, timezone +from typing import Annotated, Final, TypeAlias + +from fastapi import APIRouter, Depends, HTTPException, Query, Response +from pydantic import Field, model_validator + +from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth +from litellm.proxy.lens.endpoints import Auth, user_scope +from litellm.proxy.lens.feedback_models import ( + Feedback, + FeedbackInput, + TraceFeedback, + TraceFeedbackRequest, + TraceFeedbackSummary, +) +from litellm.proxy.lens.feedback_repository import ( + ClickHouseFeedbackStore, + FeedbackStore, + FeedbackWrite, + session_trace_id, +) +from litellm.proxy.lens.models import Record, Scope, TraceIdentity +from litellm.proxy.tracing_runtime import provide_storage +from litellm.rust_bridge.trace.storage import ClickHouseStorage + +router: Final = APIRouter(prefix="/lens/feedback", tags=["Lens"]) +TRACE_NOT_FOUND: Final = "Trace not found" + + +class FeedbackTarget(Record): + trace_id: str | None = Field(default=None, min_length=1, max_length=128) + session_id: str | None = Field(default=None, min_length=1, max_length=512) + trace_ref: str = Field(default="", max_length=512) + + @model_validator(mode="after") + def one_target(self) -> "FeedbackTarget": + if (self.trace_id is None) == (self.session_id is None): + raise ValueError("Pass exactly one of trace_id or session_id") + return self + + def trace(self) -> TraceIdentity: + trace_id: Final = self.trace_id if self.trace_id is not None else session_trace_id(self.session_id or "") + return TraceIdentity(trace_id=trace_id, trace_ref=self.trace_ref) + + +class FeedbackSubmission(FeedbackTarget, FeedbackInput): + pass + + +class FeedbackDeletion(FeedbackTarget): + user: str = Field(default="", max_length=256) + + +def write_scope(auth: UserAPIKeyAuth) -> Scope: + if auth.user_role == LitellmUserRoles.PROXY_ADMIN: + return Scope(all_teams=True) + if auth.user_role == LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY: + raise HTTPException(403, "Admin viewers cannot write feedback") + return Scope(team_id=auth.team_id or "", api_key_hash="" if auth.team_id else auth.token or "") + + +def feedback_store( + storage: Annotated[ClickHouseStorage | None, Depends(provide_storage)], +) -> FeedbackStore: + if storage is None: + raise HTTPException(501, "Lens feedback needs agent tracing. Configure the Lens service and LITELLM_LENS_URL.") + return ClickHouseFeedbackStore(storage) + + +def stored_now() -> datetime: + now: Final = datetime.now(timezone.utc) + return now.replace(microsecond=now.microsecond // 1000 * 1000) + + +def author(auth: UserAPIKeyAuth, user: str) -> str: + identity: Final = user or auth.user_id or auth.token + if not identity: + raise HTTPException(422, "Name the user who left this feedback") + return identity + + +Store: TypeAlias = Annotated[FeedbackStore, Depends(feedback_store)] +Now: TypeAlias = Annotated[datetime, Depends(stored_now)] +Target: TypeAlias = Annotated[FeedbackTarget, Query()] +Deletion: TypeAlias = Annotated[FeedbackDeletion, Query()] + + +@router.get("", response_model=TraceFeedback) +async def read_feedback(target: Target, auth: Auth, store: Store) -> TraceFeedback: + feedback: Final = await store.for_trace(user_scope(auth), target.trace()) + if feedback is None: + raise HTTPException(404, TRACE_NOT_FOUND) + return feedback + + +@router.put("", response_model=Feedback) +async def submit_feedback(body: FeedbackSubmission, auth: Auth, store: Store, now: Now) -> Feedback: + saved: Final = await store.upsert( + write_scope(auth), + FeedbackWrite( + trace=body.trace(), + author=author(auth, body.user), + feedback=FeedbackInput(score=body.score, comment=body.comment), + at=now, + ), + ) + if saved is None: + raise HTTPException(404, TRACE_NOT_FOUND) + return saved + + +@router.delete("", status_code=204) +async def delete_feedback(target: Deletion, auth: Auth, store: Store, now: Now) -> Response: + deleted: Final = await store.delete(write_scope(auth), target.trace(), author(auth, target.user), now) + if deleted is None: + raise HTTPException(404, TRACE_NOT_FOUND) + if not deleted: + raise HTTPException(404, "No feedback from this user on this trace") + return Response(status_code=204) + + +@router.post("/summary", response_model=tuple[TraceFeedbackSummary, ...]) +async def feedback_summary(body: TraceFeedbackRequest, auth: Auth, store: Store) -> tuple[TraceFeedbackSummary, ...]: + return await store.summaries(user_scope(auth), body.traces) diff --git a/litellm/proxy/lens/feedback_models.py b/litellm/proxy/lens/feedback_models.py new file mode 100644 index 00000000000..31429c9dadc --- /dev/null +++ b/litellm/proxy/lens/feedback_models.py @@ -0,0 +1,37 @@ +from datetime import datetime +from typing import Annotated, TypeAlias + +from pydantic import Field + +from litellm.constants import LENS_FEEDBACK_MAX_COMMENT_CHARS, LENS_FEEDBACK_MAX_SCORE +from litellm.proxy.lens.models import Record, TraceIdentity + +FeedbackScore: TypeAlias = Annotated[int, Field(ge=0, le=LENS_FEEDBACK_MAX_SCORE)] + + +class FeedbackInput(Record): + score: FeedbackScore + comment: str = Field(default="", max_length=LENS_FEEDBACK_MAX_COMMENT_CHARS) + user: str = Field(default="", max_length=256) + + +class Feedback(TraceIdentity): + score: FeedbackScore + comment: str + author: str + created_at: datetime + updated_at: datetime + + +class TraceFeedback(TraceIdentity): + feedback: tuple[Feedback, ...] + + +class TraceFeedbackSummary(TraceIdentity): + count: int = Field(ge=0) + average: float | None = Field(ge=0, le=LENS_FEEDBACK_MAX_SCORE) + lowest: FeedbackScore | None + + +class TraceFeedbackRequest(Record): + traces: tuple[TraceIdentity, ...] = Field(min_length=1, max_length=500) diff --git a/litellm/proxy/lens/feedback_repository.py b/litellm/proxy/lens/feedback_repository.py new file mode 100644 index 00000000000..9ee1af164e2 --- /dev/null +++ b/litellm/proxy/lens/feedback_repository.py @@ -0,0 +1,186 @@ +import hashlib +from collections.abc import Mapping +from datetime import datetime, timezone +from typing import Final, Protocol + +from litellm.proxy.lens.feedback_models import Feedback, FeedbackInput, TraceFeedback, TraceFeedbackSummary +from litellm.proxy.lens.models import Record, Scope, TraceIdentity +from litellm.proxy.lens.sources import access_parameters +from litellm.rust_bridge.trace.generated.models import ( + FeedbackRow, + FeedbackSummaryRow, + LensFeedbackParams, + LensFeedbackSummaryParams, + LensFeedbackTargetParams, +) +from litellm.rust_bridge.trace.queries import LENS_FEEDBACK, LENS_FEEDBACK_SUMMARY, LENS_FEEDBACK_TARGET +from litellm.rust_bridge.trace.storage import ClickHouseStorage + +FEEDBACK_TABLE: Final = "lens_feedback" + + +def session_trace_id(session_id: str) -> str: + return hashlib.sha256(f"litellm.claude.session.v1\0{session_id}".encode()).digest()[:16].hex() + + +class FeedbackWrite(Record): + trace: TraceIdentity + author: str + feedback: FeedbackInput + at: datetime + + +class FeedbackStore(Protocol): + async def for_trace(self, scope: Scope, trace: TraceIdentity) -> TraceFeedback | None: ... + async def upsert(self, scope: Scope, write: FeedbackWrite) -> Feedback | None: ... + async def delete(self, scope: Scope, trace: TraceIdentity, author: str, at: datetime) -> bool | None: ... + async def summaries(self, scope: Scope, traces: tuple[TraceIdentity, ...]) -> tuple[TraceFeedbackSummary, ...]: ... + + +class _Target(Record): + trace_id: str + trace_ref: str + team_id: str + key_hash: str + + +def _iso(at: datetime) -> str: + return at.astimezone(timezone.utc).isoformat(timespec="milliseconds").replace("+00:00", "Z") + + +def _feedback(row: FeedbackRow) -> Feedback: + return Feedback( + trace_id=row.trace_id, + trace_ref=row.trace_ref, + score=row.score, + comment=row.comment, + author=row.author, + created_at=datetime.fromisoformat(row.created_at), + updated_at=datetime.fromisoformat(row.updated_at), + ) + + +class ClickHouseFeedbackStore: + def __init__(self, storage: ClickHouseStorage) -> None: + self.storage: Final = storage + + async def _target(self, scope: Scope, trace: TraceIdentity) -> _Target | None: + rows: Final = await self.storage.query( + LENS_FEEDBACK_TARGET, + LensFeedbackTargetParams( + **access_parameters(scope).model_dump(), trace_id=trace.trace_id, trace_ref=trace.trace_ref + ), + ) + if len(rows) != 1: + return None + return _Target( + trace_id=trace.trace_id, trace_ref=rows[0].trace_ref, team_id=rows[0].team_id, key_hash=rows[0].key_hash + ) + + async def _rows(self, scope: Scope, trace: TraceIdentity) -> tuple[Feedback, ...]: + rows: Final = await self.storage.query( + LENS_FEEDBACK, + LensFeedbackParams( + **access_parameters(scope).model_dump(), trace_id=trace.trace_id, trace_ref=trace.trace_ref + ), + ) + return tuple(_feedback(row) for row in rows) + + async def for_trace(self, scope: Scope, trace: TraceIdentity) -> TraceFeedback | None: + target: Final = await self._target(scope, trace) + if target is None: + return None + identity: Final = TraceIdentity(trace_id=target.trace_id, trace_ref=target.trace_ref) + return TraceFeedback( + trace_id=identity.trace_id, trace_ref=identity.trace_ref, feedback=await self._rows(scope, identity) + ) + + async def _write(self, target: _Target, author: str, values: Mapping[str, object]) -> None: + await self.storage.insert_rows( + FEEDBACK_TABLE, + ( + { + "TeamId": target.team_id, + "ApiKeyHash": target.key_hash, + "TraceId": target.trace_id, + "Author": author, + **values, + }, + ), + ) + + async def upsert(self, scope: Scope, write: FeedbackWrite) -> Feedback | None: + target: Final = await self._target(scope, write.trace) + if target is None: + return None + identity: Final = TraceIdentity(trace_id=target.trace_id, trace_ref=target.trace_ref) + previous: Final = next((f for f in await self._rows(scope, identity) if f.author == write.author), None) + created: Final = previous.created_at if previous else write.at + await self._write( + target, + write.author, + { + "Score": write.feedback.score, + "Comment": write.feedback.comment, + "CreatedAt": _iso(created), + "UpdatedAt": _iso(write.at), + "IsDeleted": 0, + }, + ) + return Feedback( + trace_id=target.trace_id, + trace_ref=target.trace_ref, + score=write.feedback.score, + comment=write.feedback.comment, + author=write.author, + created_at=created, + updated_at=write.at, + ) + + async def delete(self, scope: Scope, trace: TraceIdentity, author: str, at: datetime) -> bool | None: + target: Final = await self._target(scope, trace) + if target is None: + return None + identity: Final = TraceIdentity(trace_id=target.trace_id, trace_ref=target.trace_ref) + previous: Final = next((f for f in await self._rows(scope, identity) if f.author == author), None) + if previous is None: + return False + await self._write( + target, + author, + { + "Score": previous.score, + "Comment": "", + "CreatedAt": _iso(previous.created_at), + "UpdatedAt": _iso(at), + "IsDeleted": 1, + }, + ) + return True + + async def summaries(self, scope: Scope, traces: tuple[TraceIdentity, ...]) -> tuple[TraceFeedbackSummary, ...]: + rows: Final = await self.storage.query( + LENS_FEEDBACK_SUMMARY, + LensFeedbackSummaryParams( + **access_parameters(scope).model_dump(), trace_ids=sorted({t.trace_id for t in traces}) + ), + ) + return tuple(summary for trace in traces for summary in _summaries(trace, rows)) + + +def _summaries(trace: TraceIdentity, rows: tuple[FeedbackSummaryRow, ...]) -> tuple[TraceFeedbackSummary, ...]: + matched: Final = tuple( + row for row in rows if row.trace_id == trace.trace_id and trace.trace_ref in ("", row.trace_ref) + ) + if not matched: + return ( + TraceFeedbackSummary( + trace_id=trace.trace_id, trace_ref=trace.trace_ref, count=0, average=None, lowest=None + ), + ) + return tuple( + TraceFeedbackSummary( + trace_id=row.trace_id, trace_ref=row.trace_ref, count=row.count, average=row.average, lowest=row.lowest + ) + for row in matched + ) diff --git a/litellm/proxy/lens/ingestion.py b/litellm/proxy/lens/ingestion.py index 50cb4d82998..c7e34353b38 100644 --- a/litellm/proxy/lens/ingestion.py +++ b/litellm/proxy/lens/ingestion.py @@ -59,6 +59,8 @@ class ServiceConnection(Record): url: str connected: bool status: ServiceStatus + configured: bool = False + release: str = "" @dataclass(frozen=True, slots=True) diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index 5394ea010cd..8efa878a08d 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -65,7 +65,10 @@ from litellm.proxy.common_utils.callback_utils import ( get_metadata_variable_name_from_kwargs, strip_callback_config, ) -from litellm.proxy.common_utils.http_parsing_utils import _safe_get_request_headers +from litellm.proxy.common_utils.http_parsing_utils import ( # noqa: F401 # legacy module exports + _safe_get_request_headers, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + safe_get_request_headers, +) from litellm.proxy.spend_tracking.carried_budget_state import carried_budget_metadata from litellm.types.integrations.anthropic_cache_control_hook import GATEWAY_INJECTED_CACHE_METADATA_KEY @@ -151,6 +154,27 @@ def _session_id_from_baggage(baggage: str) -> str | None: return None +def _caller_trace_field(data: Mapping[str, object], metadata_variable_name: str, field: str) -> str | None: + """The caller's value for a trace-control field, counted only when it is a + usable id: a non-empty string. An explicitly empty/unusable value on the + active metadata container still shadows the promoted requester value, but + neither ever counts as "the caller supplied this field" on its own, so a + numeric session id or an empty string cannot suppress the W3C header + fallback or satisfy a missing-session-id policy.""" + active: Final = data.get(metadata_variable_name) + if isinstance(active, Mapping) and field in active: + active_map: Final = cast(Mapping[str, object], active) # cast-ok: isinstance above, free-form JSON values + active_value: Final = active_map[field] + return active_value if isinstance(active_value, str) and active_value else None + promoted: Final = metadata_variable_name == "litellm_metadata" and field in LITELLM_TRACE_CONTROL_METADATA_FIELDS + requester: Final = data.get("metadata") + if not promoted or not isinstance(requester, Mapping): + return None + requester_map: Final = cast(Mapping[str, object], requester) # cast-ok: isinstance above, free-form JSON values + requester_value: Final = requester_map.get(field) + return requester_value if isinstance(requester_value, str) and requester_value else None + + def _stampable_key_hash(user_api_key_dict: UserAPIKeyAuth) -> str | None: """Only proxy-validated keys are stamped, proven by the unforgeable via_virtual_key marker AND a known non-secret shape: the sha256 hex digest @@ -639,7 +663,7 @@ def _strip_router_reserved_metadata( ) -def _get_metadata_variable_name(request: Request) -> str: +def get_metadata_variable_name(request: Request) -> str: """ Helper to return what the "metadata" field should be called in the request data @@ -653,6 +677,9 @@ def _get_metadata_variable_name(request: Request) -> str: return metadata_variable_name_for_route(get_request_route(request)) +_get_metadata_variable_name: Final = get_metadata_variable_name + + def metadata_variable_name_for_route(route: str) -> Literal["metadata", "litellm_metadata"]: if "thread" in route or "assistant" in route: return "litellm_metadata" @@ -833,11 +860,22 @@ def apply_missing_session_id_policy( ): metadata["session_id"] = body_session_id return - if data.get("litellm_session_id") or metadata.get("session_id"): + caller_session_id: Final = _caller_trace_field(data, _metadata_variable_name, "session_id") + if caller_session_id is not None: + # Consumers that read the root field (router fallbacks, spend logs, + # sandbox reuse) otherwise see no session and mint a uuid4 per request. + if not data.get("litellm_session_id"): + data["litellm_session_id"] = caller_session_id # rebind-ok: data is an out-param + return + if data.get("litellm_session_id"): return match policy: case "generate": - session_id: Final = str(data.get("litellm_trace_id") or metadata.get("trace_id") or uuid.uuid4()) + session_id: Final = str( + data.get("litellm_trace_id") + or _caller_trace_field(data, _metadata_variable_name, "trace_id") + or uuid.uuid4() + ) data["litellm_session_id"] = session_id # rebind-ok: data is an out-param data.setdefault("litellm_trace_id", session_id) metadata["session_id"] = session_id @@ -941,7 +979,7 @@ def convert_key_logging_metadata_to_callback( return team_callback_settings_obj -def _get_validated_callback_metadata(item: dict, *, source: str) -> AddTeamCallback | None: +def get_validated_callback_metadata(item: dict, *, source: str) -> AddTeamCallback | None: try: return AddTeamCallback(**item) except (PydanticValidationError, ValueError) as e: @@ -953,6 +991,9 @@ def _get_validated_callback_metadata(item: dict, *, source: str) -> AddTeamCallb return None +_get_validated_callback_metadata: Final = get_validated_callback_metadata + + class KeyAndTeamLoggingSettings: """ Helper class to get the dynamic logging settings for the key and team @@ -971,7 +1012,7 @@ class KeyAndTeamLoggingSettings: return None -def _get_dynamic_logging_metadata( +def get_dynamic_logging_metadata( user_api_key_dict: UserAPIKeyAuth, proxy_config: ProxyConfig ) -> TeamCallbackMetadata | None: callback_settings_obj: TeamCallbackMetadata | None = None @@ -986,7 +1027,7 @@ def _get_dynamic_logging_metadata( ######################################################################################### if key_dynamic_logging_settings is not None: for item in key_dynamic_logging_settings: - callback = _get_validated_callback_metadata(item=item, source="key-level") + callback = get_validated_callback_metadata(item=item, source="key-level") if callback is None: continue callback_settings_obj = convert_key_logging_metadata_to_callback( @@ -998,7 +1039,7 @@ def _get_dynamic_logging_metadata( ######################################################################################### elif team_dynamic_logging_settings is not None: for item in team_dynamic_logging_settings: - callback = _get_validated_callback_metadata(item=item, source="team-level") + callback = get_validated_callback_metadata(item=item, source="team-level") if callback is None: continue callback_settings_obj = convert_key_logging_metadata_to_callback( @@ -1032,6 +1073,9 @@ def _get_dynamic_logging_metadata( return callback_settings_obj +_get_dynamic_logging_metadata: Final = get_dynamic_logging_metadata + + _TENANT_OTEL_PARAMS: Final = TypeAdapter(StandardCallbackDynamicParams) @@ -1123,7 +1167,7 @@ def resolve_tenant_otel_destinations( callbacks: Final = tuple( callback for item in entries - if (callback := _get_validated_callback_metadata(item=item, source="otel-destination")) is not None + if (callback := get_validated_callback_metadata(item=item, source="otel-destination")) is not None if callback.callback_name.lower() not in disabled ) return tuple( @@ -1440,7 +1484,7 @@ class LiteLLMProxyRequestSetup: """ Add headers to the LLM call by model group """ - from litellm.proxy.auth.auth_checks import _check_model_access_helper + from litellm.proxy.auth.auth_checks import check_model_access_helper from litellm.proxy.proxy_server import llm_router data_model: Final = data.get("model") @@ -1449,7 +1493,7 @@ class LiteLLMProxyRequestSetup: data_model is not None and litellm.model_group_settings is not None and litellm.model_group_settings.forward_client_headers_to_llm_api is not None - and _check_model_access_helper( + and check_model_access_helper( model=data_model, llm_router=llm_router, models=litellm.model_group_settings.forward_client_headers_to_llm_api, @@ -1584,16 +1628,28 @@ class LiteLLMProxyRequestSetup: # Last-resort fallback: the W3C standards for trace/session propagation # (https://www.w3.org/TR/trace-context/, https://www.w3.org/TR/baggage/). # Lower priority than everything above - only fires when neither the - # explicit litellm headers nor the Anthropic-metadata path found - # anything - but lets a caller's existing traceparent/baggage headers - # (from real OTel instrumentation) correlate with litellm's own logs - # instead of generating an unrelated trace_id. + # explicit litellm headers, the Anthropic-metadata path, nor the + # caller's own request metadata set the field to a DIFFERENT usable id + # - but lets a caller's existing traceparent/baggage headers (from + # real OTel instrumentation) correlate with litellm's own logs instead + # of generating an unrelated trace_id. normalized_headers: Final = MappingProxyType({k.lower(): v for k, v in headers.items() if isinstance(k, str)}) if "litellm_trace_id" not in data: traceparent: Final = normalized_headers.get("traceparent") if isinstance(traceparent, str): trace_id_from_traceparent: Final = _trace_id_from_traceparent(traceparent) - if trace_id_from_traceparent: + # The caller's metadata wins over the header fallback unless + # both carry the same id: stamping the root field then claims + # nothing the caller did not already ask for, and keeps the + # W3C-correlated root trace id instead of a generated uuid4. + caller_trace_id: Final = _caller_trace_field( + cast(Mapping[str, object], data), # cast-ok: request body is a str-keyed JSON object + _metadata_variable_name, + "trace_id", + ) + if trace_id_from_traceparent and ( + caller_trace_id is None or caller_trace_id == trace_id_from_traceparent + ): metadata_from_headers["trace_id"] = trace_id_from_traceparent data["litellm_trace_id"] = trace_id_from_traceparent # rebind-ok: data is an out-param verbose_proxy_logger.debug( @@ -1603,7 +1659,14 @@ class LiteLLMProxyRequestSetup: baggage: Final = normalized_headers.get("baggage") if isinstance(baggage, str): session_id_from_baggage: Final = _session_id_from_baggage(baggage) - if session_id_from_baggage: + caller_session_id: Final = _caller_trace_field( + cast(Mapping[str, object], data), # cast-ok: request body is a str-keyed JSON object + _metadata_variable_name, + "session_id", + ) + if session_id_from_baggage and ( + caller_session_id is None or caller_session_id == session_id_from_baggage + ): metadata_from_headers["session_id"] = session_id_from_baggage data["litellm_session_id"] = session_id_from_baggage # rebind-ok: data is an out-param verbose_proxy_logger.debug("Extracted session_id from W3C baggage header") @@ -1745,7 +1808,7 @@ class LiteLLMProxyRequestSetup: ## KEY-LEVEL SPEND LOGS / TAGS if "tags" in key_metadata and key_metadata["tags"] is not None: - data[_metadata_variable_name]["tags"] = LiteLLMProxyRequestSetup._merge_tags( + data[_metadata_variable_name]["tags"] = LiteLLMProxyRequestSetup.merge_tags( request_tags=data[_metadata_variable_name].get("tags"), tags_to_add=key_metadata["tags"], ) @@ -1793,8 +1856,8 @@ class LiteLLMProxyRequestSetup: team_spend_logs_metadata=team_metadata.get("spend_logs_metadata"), request_spend_logs_metadata=metadata.get("spend_logs_metadata"), ) - tags: Final = LiteLLMProxyRequestSetup._merge_tags( - request_tags=LiteLLMProxyRequestSetup._merge_tags( + tags: Final = LiteLLMProxyRequestSetup.merge_tags( + request_tags=LiteLLMProxyRequestSetup.merge_tags( request_tags=request_tags if isinstance(request_tags, list) else None, tags_to_add=team_tags if isinstance(team_tags, list) else None, ), @@ -1826,7 +1889,7 @@ class LiteLLMProxyRequestSetup: return {**(team_values or {}), **(request_values or {})} @staticmethod - def _merge_tags(request_tags: list | None, tags_to_add: list | None) -> list: + def merge_tags(request_tags: list | None, tags_to_add: list | None) -> list: """ Helper function to merge two lists of tags, ensuring no duplicates. @@ -1849,6 +1912,8 @@ class LiteLLMProxyRequestSetup: return final_tags + _merge_tags = merge_tags + @staticmethod def add_team_based_callbacks_from_config( team_id: str, @@ -1938,7 +2003,7 @@ class LiteLLMProxyRequestSetup: metadata: Final = _normalized_metadata_slot(request_data, _metadata_variable_name) existing_tags: Final = metadata.get("tags") - metadata["tags"] = LiteLLMProxyRequestSetup._merge_tags( + metadata["tags"] = LiteLLMProxyRequestSetup.merge_tags( request_tags=existing_tags if isinstance(existing_tags, list) else None, tags_to_add=key_tags, ) @@ -1968,7 +2033,7 @@ class LiteLLMProxyRequestSetup: # No allow_client_tags opt-in: caller-supplied tags always flow # into metadata.tags (see add_litellm_data_to_request). The pre-auth # merge mirrors that so _tag_max_budget_check sees the same tags. - headers: Final = _safe_get_request_headers(request=request) + headers: Final = safe_get_request_headers(request=request) raw_header_tags: Final = headers.get("x-litellm-tags") if not raw_header_tags: return @@ -1990,7 +2055,7 @@ class LiteLLMProxyRequestSetup: metadata: Final = _normalized_metadata_slot(request_data, _metadata_variable_name) existing_tags: Final = metadata.get("tags") - metadata["tags"] = LiteLLMProxyRequestSetup._merge_tags( + metadata["tags"] = LiteLLMProxyRequestSetup.merge_tags( request_tags=existing_tags if isinstance(existing_tags, list) else None, tags_to_add=header_tags, ) @@ -2088,7 +2153,7 @@ async def add_litellm_data_to_request( if _mk.startswith("user_api_key_"): del _user_metadata[_mk] - _raw_headers: Final[dict[str, str]] = RedactedDict(_safe_get_request_headers(request)) + _raw_headers: Final[dict[str, str]] = RedactedDict(safe_get_request_headers(request)) forward_llm_auth = False if general_settings: @@ -2162,7 +2227,7 @@ async def add_litellm_data_to_request( } safe_add_api_version_from_query_params(data, request) - _metadata_variable_name: Final = _get_metadata_variable_name(request) + _metadata_variable_name: Final = get_metadata_variable_name(request) if data.get(_metadata_variable_name, None) is None: data[_metadata_variable_name] = {} @@ -2494,7 +2559,7 @@ async def add_litellm_data_to_request( ) if tags is not None: - data[_metadata_variable_name]["tags"] = LiteLLMProxyRequestSetup._merge_tags( + data[_metadata_variable_name]["tags"] = LiteLLMProxyRequestSetup.merge_tags( request_tags=data[_metadata_variable_name].get("tags"), tags_to_add=tags, ) @@ -2506,7 +2571,7 @@ async def add_litellm_data_to_request( else None ) if _caller_body_tags: - data[_metadata_variable_name]["tags"] = LiteLLMProxyRequestSetup._merge_tags( # rebind-ok: matches file idiom + data[_metadata_variable_name]["tags"] = LiteLLMProxyRequestSetup.merge_tags( # rebind-ok: matches file idiom request_tags=data[_metadata_variable_name].get("tags"), tags_to_add=_caller_body_tags, ) @@ -2523,7 +2588,7 @@ async def add_litellm_data_to_request( ) # Team Callbacks controls - callback_settings_obj: Final = _get_dynamic_logging_metadata( + callback_settings_obj: Final = get_dynamic_logging_metadata( user_api_key_dict=user_api_key_dict, proxy_config=proxy_config ) if callback_settings_obj is not None: @@ -3004,7 +3069,7 @@ def _enforced_params_check( return True -def _add_guardrails_from_key_or_team_metadata( +def add_guardrails_from_key_or_team_metadata( key_metadata: dict | None, team_metadata: dict | None, data: dict, @@ -3024,7 +3089,7 @@ def _add_guardrails_from_key_or_team_metadata( project_metadata: The project metadata dictionary to check for guardrails """ - from litellm.proxy.utils import _premium_user_check + from litellm.proxy.utils import premium_user_check # Initialize guardrails set (avoiding duplicates) combined_guardrails: Final = set() @@ -3032,19 +3097,19 @@ def _add_guardrails_from_key_or_team_metadata( # Add key-level guardrails first if key_metadata and "guardrails" in key_metadata: if isinstance(key_metadata["guardrails"], list) and len(key_metadata["guardrails"]) > 0: - _premium_user_check() + premium_user_check() combined_guardrails.update(key_metadata["guardrails"]) # Add team-level guardrails (set automatically handles duplicates) if team_metadata and "guardrails" in team_metadata: if isinstance(team_metadata["guardrails"], list) and len(team_metadata["guardrails"]) > 0: - _premium_user_check() + premium_user_check() combined_guardrails.update(team_metadata["guardrails"]) # Add project-level guardrails (set automatically handles duplicates) if project_metadata and "guardrails" in project_metadata: if isinstance(project_metadata["guardrails"], list) and len(project_metadata["guardrails"]) > 0: - _premium_user_check() + premium_user_check() combined_guardrails.update(project_metadata["guardrails"]) # Set combined guardrails in metadata as list @@ -3052,6 +3117,9 @@ def _add_guardrails_from_key_or_team_metadata( data[metadata_variable_name]["guardrails"] = list(combined_guardrails) +_add_guardrails_from_key_or_team_metadata: Final = add_guardrails_from_key_or_team_metadata + + def _add_guardrails_from_policies_in_metadata( key_metadata: dict | None, team_metadata: dict | None, @@ -3077,7 +3145,7 @@ def _add_guardrails_from_policies_in_metadata( from litellm._logging import verbose_proxy_logger from litellm.proxy.policy_engine.policy_registry import get_policy_registry from litellm.proxy.policy_engine.policy_resolver import PolicyResolver - from litellm.proxy.utils import _premium_user_check + from litellm.proxy.utils import premium_user_check from litellm.types.proxy.policy_engine import PolicyMatchContext # Collect policy names from key and team metadata @@ -3086,19 +3154,19 @@ def _add_guardrails_from_policies_in_metadata( # Add key-level policies first if key_metadata and "policies" in key_metadata: if isinstance(key_metadata["policies"], list) and len(key_metadata["policies"]) > 0: - _premium_user_check() + premium_user_check() policy_names.update(key_metadata["policies"]) # Add team-level policies if team_metadata and "policies" in team_metadata: if isinstance(team_metadata["policies"], list) and len(team_metadata["policies"]) > 0: - _premium_user_check() + premium_user_check() policy_names.update(team_metadata["policies"]) # Add project-level policies if project_metadata and "policies" in project_metadata: if isinstance(project_metadata["policies"], list) and len(project_metadata["policies"]) > 0: - _premium_user_check() + premium_user_check() policy_names.update(project_metadata["policies"]) if not policy_names: @@ -3166,7 +3234,7 @@ def add_guardrails_from_auth_metadata( metadata_variable_name: str, ) -> None: """Resolve key, team, and project guardrails, direct and via policies, onto the request metadata.""" - _add_guardrails_from_key_or_team_metadata( + add_guardrails_from_key_or_team_metadata( key_metadata=user_api_key_dict.metadata, team_metadata=user_api_key_dict.team_metadata, project_metadata=user_api_key_dict.project_metadata, diff --git a/litellm/proxy/management_endpoints/access_group_endpoints.py b/litellm/proxy/management_endpoints/access_group_endpoints.py index 97311a0ef8a..3f05c3a407a 100644 --- a/litellm/proxy/management_endpoints/access_group_endpoints.py +++ b/litellm/proxy/management_endpoints/access_group_endpoints.py @@ -17,11 +17,15 @@ from litellm.proxy._types import ( UserAPIKeyAuth, ) from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry -from litellm.proxy.auth.auth_checks import ( - _cache_access_object, - _cache_key_object, - _cache_team_object, - _get_team_object_from_cache, +from litellm.proxy.auth.auth_checks import ( # noqa: F401 # legacy module exports + _cache_access_object, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _cache_key_object, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _cache_team_object, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _get_team_object_from_cache, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + cache_access_object, + cache_key_object, + cache_team_object, + get_team_object_from_cache, ) from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.common_utils.user_api_key_cache import AUTH_OBJECTS_TARGET @@ -166,9 +170,9 @@ def _require_proxy_admin(user_api_key_dict: UserAPIKeyAuth) -> None: def _require_admin_view(user_api_key_dict: UserAPIKeyAuth) -> None: """Admin Viewer parity: PROXY_ADMIN or PROXY_ADMIN_VIEW_ONLY may read.""" - from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view + from litellm.proxy.management_endpoints.common_utils import user_api_key_has_admin_view - if not _user_has_admin_view(user_api_key_dict): + if not user_api_key_has_admin_view(user_api_key_dict): raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, detail={"error": CommonProxyErrors.not_allowed_access.value}, @@ -303,7 +307,7 @@ async def _cache_access_group_record(record: _AccessGroupRecord) -> None: from litellm.proxy.proxy_server import proxy_logging_obj, user_api_key_cache access_group_table: Final = _record_to_access_group_table(record) - await _cache_access_object( + await cache_access_object( access_group_id=record.access_group_id, access_group_table=access_group_table, user_api_key_cache=user_api_key_cache, @@ -408,7 +412,7 @@ async def _patch_team_caches_add_access_group( ) -> None: """Patch cached team objects to include access_group_id.""" for team_id in team_ids: - cached_team = await _get_team_object_from_cache( + cached_team = await get_team_object_from_cache( key=f"team_id:{team_id}", user_api_key_cache=user_api_key_cache, parent_otel_span=None, @@ -421,7 +425,7 @@ async def _patch_team_caches_add_access_group( cached_team.access_group_ids = list(cached_team.access_group_ids) + [access_group_id] else: continue - await _cache_team_object( + await cache_team_object( team_id=team_id, team_table=cached_team, user_api_key_cache=user_api_key_cache, @@ -437,14 +441,14 @@ async def _patch_team_caches_remove_access_group( ) -> None: """Patch cached team objects to remove access_group_id.""" for team_id in team_ids: - cached_team = await _get_team_object_from_cache( + cached_team = await get_team_object_from_cache( key=f"team_id:{team_id}", user_api_key_cache=user_api_key_cache, parent_otel_span=None, ) if cached_team is not None and cached_team.access_group_ids: cached_team.access_group_ids = [ag for ag in cached_team.access_group_ids if ag != access_group_id] - await _cache_team_object( + await cache_team_object( team_id=team_id, team_table=cached_team, user_api_key_cache=user_api_key_cache, @@ -473,7 +477,7 @@ async def _patch_key_caches_add_access_group( cached_key.access_group_ids = list(cached_key.access_group_ids) + [access_group_id] else: continue - await _cache_key_object( + await cache_key_object( hashed_token=token, user_api_key_obj=cached_key, user_api_key_cache=user_api_key_cache, @@ -496,7 +500,7 @@ async def _patch_key_caches_remove_access_group( ) if cached_key is not None and cached_key.access_group_ids: cached_key.access_group_ids = [ag for ag in cached_key.access_group_ids if ag != access_group_id] - await _cache_key_object( + await cache_key_object( hashed_token=token, user_api_key_obj=cached_key, user_api_key_cache=user_api_key_cache, diff --git a/litellm/proxy/management_endpoints/auto_router_endpoints.py b/litellm/proxy/management_endpoints/auto_router_endpoints.py index 24abcdcc80c..ca9d5612b41 100644 --- a/litellm/proxy/management_endpoints/auto_router_endpoints.py +++ b/litellm/proxy/management_endpoints/auto_router_endpoints.py @@ -28,9 +28,10 @@ from litellm.proxy._types import ( ProxyException, UserAPIKeyAuth, ) -from litellm.proxy.auth.auth_checks import ( - _virtual_key_max_budget_check, +from litellm.proxy.auth.auth_checks import ( # noqa: F401 # legacy module exports + _virtual_key_max_budget_check, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export can_key_call_resolved_model, + virtual_key_max_budget_check, ) from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.db.autorouter_session_rollup import ( @@ -340,7 +341,7 @@ async def _authorize_models_this_test_can_call( ) try: - await _virtual_key_max_budget_check( + await virtual_key_max_budget_check( valid_token=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, ) @@ -558,10 +559,10 @@ async def preview_auto_router_routing( if member_team is not None and _models_this_test_can_call(resolved.complexity_router_config): from litellm.proxy.auth.user_api_key_auth import ( - _run_centralized_common_checks, # pyright: ignore[reportPrivateUsage] # reuse the serving admission policy + run_centralized_common_checks, # pyright: ignore[reportPrivateUsage] # reuse the serving admission policy ) - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=actor, request=http_request, request_data=request_data, diff --git a/litellm/proxy/management_endpoints/budget_management_endpoints.py b/litellm/proxy/management_endpoints/budget_management_endpoints.py index daa4539b90f..59e358e2084 100644 --- a/litellm/proxy/management_endpoints/budget_management_endpoints.py +++ b/litellm/proxy/management_endpoints/budget_management_endpoints.py @@ -22,8 +22,9 @@ from fastapi import APIRouter, Depends, HTTPException from litellm.proxy._types import * from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time -from litellm.proxy.management_endpoints.common_utils import ( - _user_has_admin_view, +from litellm.proxy.management_endpoints.common_utils import ( # noqa: F401 # legacy module exports + _user_has_admin_view, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + user_api_key_has_admin_view, validate_budget_duration, ) from litellm.proxy.utils import jsonify_object @@ -256,7 +257,7 @@ async def budget_settings( detail={"error": CommonProxyErrors.db_not_connected_error.value}, ) - if not _user_has_admin_view(user_api_key_dict): + if not user_api_key_has_admin_view(user_api_key_dict): raise HTTPException( status_code=400, detail={"error": f"{CommonProxyErrors.not_allowed_access.value}, your role={user_api_key_dict.user_role}"}, @@ -317,7 +318,7 @@ async def list_budget( detail={"error": CommonProxyErrors.db_not_connected_error.value}, ) - if not _user_has_admin_view(user_api_key_dict): + if not user_api_key_has_admin_view(user_api_key_dict): raise HTTPException( status_code=400, detail={"error": f"{CommonProxyErrors.not_allowed_access.value}, your role={user_api_key_dict.user_role}"}, diff --git a/litellm/proxy/management_endpoints/cache_settings_endpoints.py b/litellm/proxy/management_endpoints/cache_settings_endpoints.py index 2bd4b6b47f9..aaf5d30d2e1 100644 --- a/litellm/proxy/management_endpoints/cache_settings_endpoints.py +++ b/litellm/proxy/management_endpoints/cache_settings_endpoints.py @@ -467,7 +467,7 @@ async def get_cache_settings( if prisma_client is not None: cache_config = await _cache_config_table(prisma_client).find_unique(where={"id": "cache_config"}) if cache_config is not None and cache_config.cache_settings: - stored = proxy_config._decrypt_db_variables( + stored = proxy_config.decrypt_db_variables( # rebind-ok: pre-existing rebinding on a rename-only line variables_dict=_parse_stored_settings(cache_config.cache_settings) ) @@ -534,8 +534,10 @@ async def test_cache_connection( try: existing_row: Final = await _cache_config_table(prisma_client).find_unique(where={"id": "cache_config"}) if existing_row is not None and existing_row.cache_settings: - saved_settings = proxy_config._decrypt_db_variables( - variables_dict=_parse_stored_settings(existing_row.cache_settings) + saved_settings = ( # rebind-ok: pre-existing rebinding on a rename-only line + proxy_config.decrypt_db_variables( + variables_dict=_parse_stored_settings(existing_row.cache_settings) + ) ) except Exception: # noqa: BLE001 - a saved-settings lookup failure must not block a connection test saved_settings = {} @@ -614,7 +616,9 @@ async def update_cache_settings( saved_settings: dict[str, object] = {} if existing_row is not None and existing_row.cache_settings: before_settings = _parse_stored_settings(existing_row.cache_settings) - saved_settings = proxy_config._decrypt_db_variables(variables_dict=before_settings) + saved_settings = ( # rebind-ok: pre-existing rebinding on a rename-only line + proxy_config.decrypt_db_variables(variables_dict=before_settings) + ) action: Final[AUDIT_ACTIONS] = "updated" if existing_row is not None else "created" # Preserve stored secrets behind any redacted or omitted credential, then @@ -622,7 +626,7 @@ async def update_cache_settings( cache_settings: Final = _resolve_cache_url_precedence(_merge_over_saved(request.cache_settings, saved_settings)) # Encrypt sensitive fields (keep redis_type for storage) - encrypted_settings: Final = proxy_config._encrypt_env_variables(environment_variables=cache_settings) + encrypted_settings: Final = proxy_config.encrypt_env_variables(environment_variables=cache_settings) # Save to database await _cache_config_table(prisma_client).upsert( @@ -640,13 +644,13 @@ async def update_cache_settings( # Reinitialize cache with new settings # Decrypt for initialization - decrypted_settings: Final = proxy_config._decrypt_db_variables(variables_dict=encrypted_settings) + decrypted_settings: Final = proxy_config.decrypt_db_variables(variables_dict=encrypted_settings) # Remove redis_type if present (UI-only field, not a Cache parameter) cache_params: Final = {k: v for k, v in decrypted_settings.items() if k != "redis_type"} # Initialize cache (frontend sends type="redis", not redis_type) - proxy_config._init_cache(cache_params=cache_params) + proxy_config.init_cache(cache_params=cache_params) # Update the last cache params to avoid reinitializing unnecessarily CacheSettingsManager.update_cache_params(cache_params) diff --git a/litellm/proxy/management_endpoints/common_utils.py b/litellm/proxy/management_endpoints/common_utils.py index b1c7e319b5b..86e13a304da 100644 --- a/litellm/proxy/management_endpoints/common_utils.py +++ b/litellm/proxy/management_endpoints/common_utils.py @@ -43,7 +43,7 @@ def validate_budget_duration(budget_duration: str | None, status_code: int = 400 from litellm._logging import verbose_proxy_logger from litellm.caching import DualCache -from litellm.proxy._types import ( +from litellm.proxy._types import ( # re-exported CommonProxyErrors, KeyRequestBase, LiteLLM_ManagementEndpoint_MetadataFields, @@ -56,13 +56,14 @@ from litellm.proxy._types import ( NewProjectRequest, UpdateProjectRequest, UserAPIKeyAuth, -) -from litellm.proxy._types import ( # noqa: F401 re-exported - user_api_key_has_admin_view as _user_has_admin_view, + user_api_key_has_admin_view, ) from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time from litellm.proxy.management.teams.authz import is_team_admin -from litellm.proxy.utils import _premium_user_check +from litellm.proxy.utils import ( # noqa: F401 # legacy module exports + _premium_user_check, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + premium_user_check, +) from litellm.repositories.team_repository import TeamRepository from litellm.types.utils import BudgetConfig @@ -70,6 +71,8 @@ if TYPE_CHECKING: from litellm.proxy._types import NewProjectRequest, UpdateProjectRequest from litellm.proxy.utils import PrismaClient, ProxyLogging +_user_has_admin_view: Final = user_api_key_has_admin_view + # TODO: drop once the litellm-enterprise pin moves past 0.1.71, which imports this name _is_user_team_admin: Final = is_team_admin @@ -160,7 +163,7 @@ def _passthrough_routes_permission_error(field: str, entity: str) -> HTTPExcepti ) -def _check_passthrough_routes_caller_permission( +def check_passthrough_routes_caller_permission( data: BaseModel | None, user_api_key_dict: UserAPIKeyAuth, *, @@ -178,6 +181,9 @@ def _check_passthrough_routes_caller_permission( ) +_check_passthrough_routes_caller_permission: Final = check_passthrough_routes_caller_permission + + def check_allowed_passthrough_routes_caller_permission( data: BaseModel | None, user_api_key_dict: UserAPIKeyAuth, @@ -235,7 +241,7 @@ def _metadata_changes_denied_routes(data: BaseModel, metadata: object, existing_ return metadata is None and "metadata" in data.model_fields_set and existing_denied is not None -def _check_disable_global_guardrails_caller_permission( +def check_disable_global_guardrails_caller_permission( disable_global_guardrails: bool | None, metadata: Mapping[str, object] | None, user_api_key_dict: UserAPIKeyAuth, @@ -263,7 +269,10 @@ def _check_disable_global_guardrails_caller_permission( ) -def _team_member_has_permission( +_check_disable_global_guardrails_caller_permission: Final = check_disable_global_guardrails_caller_permission + + +def team_member_has_permission( user_api_key_dict: UserAPIKeyAuth, team_obj: LiteLLM_TeamTable, permission: str, @@ -279,7 +288,10 @@ def _team_member_has_permission( return False -async def _user_has_admin_privileges( +_team_member_has_permission: Final = team_member_has_permission + + +async def user_has_admin_privileges( user_api_key_dict: UserAPIKeyAuth, prisma_client: Optional["PrismaClient"] = None, user_api_key_cache: Optional["DualCache"] = None, @@ -345,6 +357,9 @@ async def _user_has_admin_privileges( return False +_user_has_admin_privileges: Final = user_has_admin_privileges + + def _org_admin_can_invite_user( admin_user_obj: LiteLLM_UserTable, target_user_obj: LiteLLM_UserTable, @@ -487,7 +502,7 @@ async def admin_can_invite_user( return False -def _set_object_metadata_field( +def set_object_metadata_field( object_data: Union[ LiteLLM_TeamTable, KeyRequestBase, @@ -508,12 +523,15 @@ def _set_object_metadata_field( value: Value to set for the field """ if field_name in LiteLLM_ManagementEndpoint_MetadataFields_Premium and value: - _premium_user_check(field_name) + premium_user_check(field_name) object_data.metadata = object_data.metadata or {} object_data.metadata[field_name] = value +_set_object_metadata_field: Final = set_object_metadata_field + + _TEAM_MEMBER_BUDGET_LIMIT_FIELDS: Final = ( "max_budget", "soft_budget", @@ -573,7 +591,7 @@ def _has_meaningful_budget_limit(budget_values: Mapping[str, object]) -> bool: return any(_is_set_budget_value(budget_values.get(field)) for field in _TEAM_MEMBER_BUDGET_LIMIT_FIELDS) -async def _upsert_budget_and_membership( +async def upsert_budget_and_membership( tx, *, team_id: str, @@ -583,7 +601,7 @@ async def _upsert_budget_and_membership( budget_patch: Mapping[str, object], team_default_budget_id: str | None = None, shared_budget_ids: frozenset[str] | None = None, -): +) -> None: """ Apply a merge-patch of per-member budget fields to a team membership. @@ -694,6 +712,9 @@ async def _upsert_budget_and_membership( ) +_upsert_budget_and_membership: Final = upsert_budget_and_membership + + def _update_metadata_field(updated_kv: dict, field_name: str) -> None: """ Helper function to update metadata fields that require premium user checks in the update endpoint @@ -708,7 +729,7 @@ def _update_metadata_field(updated_kv: dict, field_name: str) -> None: # only for a truthy value. The falsy value is still persisted below so a # previously-set field can be cleared. if updated_kv.get(field_name): - _premium_user_check() + premium_user_check() if field_name in updated_kv and updated_kv[field_name] is not None: # remove field from updated_kv @@ -730,7 +751,7 @@ def _has_non_empty_value(value: object) -> bool: return True -def _update_metadata_fields(updated_kv: dict) -> None: +def update_metadata_fields(updated_kv: dict) -> None: """ Helper function to update all metadata fields (both premium and standard). @@ -744,3 +765,6 @@ def _update_metadata_fields(updated_kv: dict) -> None: for field in LiteLLM_ManagementEndpoint_MetadataFields: if field in updated_kv and updated_kv[field] is not None: _update_metadata_field(updated_kv=updated_kv, field_name=field) + + +_update_metadata_fields: Final = update_metadata_fields diff --git a/litellm/proxy/management_endpoints/config_override_endpoints.py b/litellm/proxy/management_endpoints/config_override_endpoints.py index 9186145a78c..2a094b1fe66 100644 --- a/litellm/proxy/management_endpoints/config_override_endpoints.py +++ b/litellm/proxy/management_endpoints/config_override_endpoints.py @@ -191,7 +191,7 @@ def _mask_sensitive_fields(data: Mapping[str, object], sensitive_fields: set[str return masked -def _get_current_env_values(env_var_mapping: dict[str, str]) -> dict[str, str | None]: +def get_current_env_values(env_var_mapping: dict[str, str]) -> dict[str, str | None]: """Read current env var values as fallback when no DB record exists.""" values: Final = {} for field_name, env_var_name in env_var_mapping.items(): @@ -200,6 +200,9 @@ def _get_current_env_values(env_var_mapping: dict[str, str]) -> dict[str, str | return values +_get_current_env_values: Final = get_current_env_values + + class _JsonSchemaField(TypedDict, total=False): type: ReadOnly[str] anyOf: ReadOnly[Sequence["_JsonSchemaField"]] @@ -232,14 +235,17 @@ def _build_field_schema(model_class: type[BaseModel]) -> dict[str, object]: } -def _parse_config_value(raw: str | Mapping[str, object]) -> dict[str, object]: +def parse_config_value(raw: str | Mapping[str, object]) -> dict[str, object]: """Parse a config_value from DB (may be JSON string or dict).""" if isinstance(raw, str): return safe_json_loads(raw, default={}) return dict(raw) -def _set_env_vars( +_parse_config_value: Final = parse_config_value + + +def set_env_vars( config_data: Mapping[str, object], env_var_mapping: Mapping[str, str] = HASHICORP_ENV_VAR_MAPPING, ) -> None: @@ -252,24 +258,30 @@ def _set_env_vars( os.environ.pop(env_var_name, None) -def _clear_hashicorp_vault_state(proxy_config: "ProxyConfig") -> None: +_set_env_vars: Final = set_env_vars + + +def clear_hashicorp_vault_state(proxy_config: "ProxyConfig") -> None: """Clear all Hashicorp Vault state: env vars, secret manager, and change-detection cache.""" - _set_env_vars({}) + set_env_vars({}) if litellm._key_management_system == KeyManagementSystem.HASHICORP_VAULT: litellm.secret_manager_client = None litellm._key_management_system = None proxy_config._last_hashicorp_vault_config = None # pyright: ignore[reportPrivateUsage] # proxy-internal change-detection cache +_clear_hashicorp_vault_state: Final = clear_hashicorp_vault_state + + def _snapshot_cyberark_boot_env(proxy_config: "ProxyConfig") -> None: """Capture deployment-provided CYBERARK_* env vars once, before the first DB-driven overwrite.""" if proxy_config._cyberark_boot_env is None: # pyright: ignore[reportPrivateUsage] # proxy-internal boot snapshot - proxy_config._cyberark_boot_env = _get_current_env_values(CYBERARK_ENV_VAR_MAPPING) # pyright: ignore[reportPrivateUsage] # proxy-internal boot snapshot + proxy_config._cyberark_boot_env = get_current_env_values(CYBERARK_ENV_VAR_MAPPING) # pyright: ignore[reportPrivateUsage] # proxy-internal boot snapshot def _restore_cyberark_runtime(proxy_config: "ProxyConfig", env_values: Mapping[str, str | None]) -> None: """Restore CYBERARK_* env vars and reinitialize (or drop) the secret manager to match them.""" - _set_env_vars(env_values, CYBERARK_ENV_VAR_MAPPING) + set_env_vars(env_values, CYBERARK_ENV_VAR_MAPPING) if env_values.get("cyberark_api_base"): try: proxy_config.initialize_secret_manager(key_management_system="cyberark") @@ -367,15 +379,19 @@ async def update_hashicorp_vault_config( existing_decrypted: dict[str, object] | None = None env_values: dict[str, str | None] = {} if existing_record is not None and existing_record.config_value is not None: - existing_data: Final = _parse_config_value(existing_record.config_value) - existing_decrypted = proxy_config._decrypt_db_variables(existing_data) + existing_data: Final = parse_config_value(existing_record.config_value) + existing_decrypted = ( # rebind-ok: pre-existing rebinding on a rename-only line + proxy_config.decrypt_db_variables(existing_data) + ) for field in HASHICORP_ENV_VAR_MAPPING: if field not in config_data and existing_decrypted.get(field): config_data[field] = existing_decrypted[field] else: # No DB record (or DB record with null config_value) — merge from # current env vars instead. - env_values = _get_current_env_values(HASHICORP_ENV_VAR_MAPPING) + env_values = get_current_env_values( # rebind-ok: pre-existing rebinding on a rename-only line + HASHICORP_ENV_VAR_MAPPING + ) for field in HASHICORP_ENV_VAR_MAPPING: if field not in config_data and env_values.get(field): config_data[field] = env_values[field] @@ -404,15 +420,15 @@ async def update_hashicorp_vault_config( ) # Snapshot current env vars so we can restore on failure - previous_env: Final = _get_current_env_values(HASHICORP_ENV_VAR_MAPPING) + previous_env: Final = get_current_env_values(HASHICORP_ENV_VAR_MAPPING) # Set env vars and verify the secret manager can initialize before persisting - _set_env_vars(config_data) + set_env_vars(config_data) try: proxy_config.initialize_secret_manager(key_management_system="hashicorp_vault") except Exception as e: - _set_env_vars(previous_env) + set_env_vars(previous_env) verbose_proxy_logger.exception("Error reinitializing Hashicorp Vault secret manager: %s", str(e)) raise HTTPException( status_code=500, @@ -420,7 +436,7 @@ async def update_hashicorp_vault_config( ) # Only persist to DB after successful init - encrypted_data: Final = proxy_config._encrypt_env_variables(config_data) + encrypted_data: Final = proxy_config.encrypt_env_variables(config_data) config_value: Final = safe_dumps(encrypted_data) await _config_overrides_table(prisma_client).upsert( where={"config_type": "hashicorp_vault"}, @@ -474,11 +490,11 @@ async def get_hashicorp_vault_config( Get current Hashicorp Vault configuration. Returns decrypted values from DB, or falls back to current env vars. """ - from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view + from litellm.proxy.management_endpoints.common_utils import user_api_key_has_admin_view from litellm.proxy.proxy_server import prisma_client, proxy_config # Admin Viewer follows the read-parity rule. - if not _user_has_admin_view(user_api_key_dict): + if not user_api_key_has_admin_view(user_api_key_dict): raise HTTPException( status_code=403, detail="Only admin users can view config overrides", @@ -498,10 +514,10 @@ async def get_hashicorp_vault_config( ) if db_record is not None and db_record.config_value is not None: - config_data: Final = _parse_config_value(db_record.config_value) + config_data: Final = parse_config_value(db_record.config_value) # Decrypt then mask sensitive fields so plaintext secrets are never sent to the UI - decrypted_data: Final[Mapping[str, object]] = proxy_config._decrypt_db_variables(config_data) + decrypted_data: Final[Mapping[str, object]] = proxy_config.decrypt_db_variables(config_data) masked_data: Final = _mask_sensitive_fields(decrypted_data, HASHICORP_SENSITIVE_FIELDS) return ConfigOverrideSettingsResponse( @@ -511,7 +527,7 @@ async def get_hashicorp_vault_config( ) # Fallback to env vars — also mask sensitive values - env_values: Final = _get_current_env_values(HASHICORP_ENV_VAR_MAPPING) + env_values: Final = get_current_env_values(HASHICORP_ENV_VAR_MAPPING) masked_env_values: Final = _mask_sensitive_fields(env_values, HASHICORP_SENSITIVE_FIELDS) return ConfigOverrideSettingsResponse( @@ -556,7 +572,9 @@ async def delete_hashicorp_vault_config( before_config: dict[str, object] | None = None if existing_record is not None and existing_record.config_value is not None: try: - before_config = proxy_config._decrypt_db_variables(_parse_config_value(existing_record.config_value)) + before_config = ( # rebind-ok: pre-existing rebinding on a rename-only line + proxy_config.decrypt_db_variables(parse_config_value(existing_record.config_value)) + ) except Exception: before_config = None @@ -568,7 +586,7 @@ async def delete_hashicorp_vault_config( except RecordNotFoundError: verbose_proxy_logger.debug("No existing Hashicorp Vault config record to delete") - _clear_hashicorp_vault_state(proxy_config) + clear_hashicorp_vault_state(proxy_config) # Only emit audit log if a row was actually removed; an idempotent # delete on a non-existent row produces no security-relevant change. @@ -688,13 +706,13 @@ async def update_cyberark_config( existing_decrypted: dict[str, object] | None = None # rebind-ok: set when record exists env_values: dict[str, str | None] = {} # rebind-ok: populated when no DB record exists if existing_record is not None and existing_record.config_value is not None: - existing_data: Final = _parse_config_value(existing_record.config_value) + existing_data: Final = parse_config_value(existing_record.config_value) existing_decrypted = proxy_config._decrypt_db_variables(existing_data) # pyright: ignore[reportPrivateUsage] # rebind-ok: populated when a prior record decrypts for field in CYBERARK_ENV_VAR_MAPPING: if field not in config_data and existing_decrypted.get(field): config_data[field] = existing_decrypted[field] else: - env_values = _get_current_env_values(CYBERARK_ENV_VAR_MAPPING) # rebind-ok: populated when no DB record exists + env_values = get_current_env_values(CYBERARK_ENV_VAR_MAPPING) # rebind-ok: populated when no DB record exists for field in CYBERARK_ENV_VAR_MAPPING: if field not in config_data and env_values.get(field): config_data[field] = env_values[field] @@ -719,13 +737,13 @@ async def update_cyberark_config( ) _snapshot_cyberark_boot_env(proxy_config) - previous_env: Final = _get_current_env_values(CYBERARK_ENV_VAR_MAPPING) - _set_env_vars(config_data, CYBERARK_ENV_VAR_MAPPING) + previous_env: Final = get_current_env_values(CYBERARK_ENV_VAR_MAPPING) + set_env_vars(config_data, CYBERARK_ENV_VAR_MAPPING) try: proxy_config.initialize_secret_manager(key_management_system="cyberark") except Exception as e: # noqa: BLE001 # any init failure must roll back env vars - _set_env_vars(previous_env, CYBERARK_ENV_VAR_MAPPING) + set_env_vars(previous_env, CYBERARK_ENV_VAR_MAPPING) verbose_proxy_logger.exception("Error reinitializing CyberArk secret manager: %s", str(e)) raise HTTPException( status_code=500, @@ -776,11 +794,11 @@ async def get_cyberark_config( Sensitive fields are masked before leaving the server. """ from litellm.proxy.management_endpoints.common_utils import ( - _user_has_admin_view, # pyright: ignore[reportPrivateUsage] # proxy-internal helper, mirrors hashicorp endpoint usage + user_api_key_has_admin_view, # pyright: ignore[reportPrivateUsage] # proxy-internal helper, mirrors hashicorp endpoint usage ) from litellm.proxy.proxy_server import prisma_client, proxy_config - if not _user_has_admin_view(user_api_key_dict): + if not user_api_key_has_admin_view(user_api_key_dict): raise HTTPException( status_code=403, detail="Only admin users can view config overrides", @@ -797,7 +815,7 @@ async def get_cyberark_config( db_record: Final = await _config_overrides_table(prisma_client).find_unique(where={"config_type": "cyberark"}) if db_record is not None and db_record.config_value is not None: - config_data: Final = _parse_config_value(db_record.config_value) + config_data: Final = parse_config_value(db_record.config_value) decrypted_data: Final[Mapping[str, object]] = proxy_config._decrypt_db_variables(config_data) # pyright: ignore[reportPrivateUsage] # proxy-internal helper, mirrors hashicorp endpoint usage masked_data: Final = _mask_sensitive_fields(decrypted_data, CYBERARK_SENSITIVE_FIELDS) @@ -807,7 +825,7 @@ async def get_cyberark_config( field_schema=field_schema, ) - env_values: Final = _get_current_env_values(CYBERARK_ENV_VAR_MAPPING) + env_values: Final = get_current_env_values(CYBERARK_ENV_VAR_MAPPING) masked_env_values: Final = _mask_sensitive_fields(env_values, CYBERARK_SENSITIVE_FIELDS) return ConfigOverrideSettingsResponse( @@ -848,7 +866,7 @@ async def delete_cyberark_config( before_config: dict[str, object] | None = None # rebind-ok: set when decrypts if existing_record is not None and existing_record.config_value is not None: try: - before_config = proxy_config._decrypt_db_variables(_parse_config_value(existing_record.config_value)) # pyright: ignore[reportPrivateUsage] # rebind-ok: populated when the prior record decrypts + before_config = proxy_config._decrypt_db_variables(parse_config_value(existing_record.config_value)) # pyright: ignore[reportPrivateUsage] # rebind-ok: populated when the prior record decrypts except Exception: # noqa: BLE001 # undecryptable prior config must not block deletion before_config = None # rebind-ok: reset when decryption fails diff --git a/litellm/proxy/management_endpoints/coordination_redis_endpoints.py b/litellm/proxy/management_endpoints/coordination_redis_endpoints.py index fb8b2544d7d..663e22c4312 100644 --- a/litellm/proxy/management_endpoints/coordination_redis_endpoints.py +++ b/litellm/proxy/management_endpoints/coordination_redis_endpoints.py @@ -214,14 +214,14 @@ def _coordination_redis_source(settings: Mapping[str, object] | None) -> Coordin an explicit block wins, else a plain-Redis response-cache backend is borrowed, else the REDIS_* environment fallback applies. """ - from litellm.proxy.proxy_server import _environment_has_redis_connection_target + from litellm.proxy.proxy_server import environment_has_redis_connection_target if settings: return "coordination_redis" cache_backend: Final = litellm.cache.cache if litellm.cache is not None else None if isinstance(cache_backend, (RedisCache, RedisClusterCache)): return "cache_backend" - if _environment_has_redis_connection_target(): + if environment_has_redis_connection_target(): return "environment" return None @@ -416,7 +416,7 @@ async def check_coordination_redis_connection( Builds a throwaway client (never touching global state) and pings it. """ - from litellm.proxy.proxy_server import _build_redis_usage_cache + from litellm.proxy.proxy_server import build_redis_usage_cache _enforce_proxy_admin(user_api_key_dict) @@ -426,7 +426,9 @@ async def check_coordination_redis_connection( redis_cache: RedisCache | None = None try: - redis_cache = _build_redis_usage_cache(params.model_dump(exclude_none=True)) + redis_cache = build_redis_usage_cache( # rebind-ok: pre-existing rebinding on a rename-only line + params.model_dump(exclude_none=True) + ) await asyncio.wait_for(redis_cache.ping(), timeout=_PING_TIMEOUT_SECONDS) return CoordinationRedisTestResponse(status="healthy") except asyncio.TimeoutError: diff --git a/litellm/proxy/management_endpoints/credential_migration.py b/litellm/proxy/management_endpoints/credential_migration.py index f9c09128f66..3caaf5ba1af 100644 --- a/litellm/proxy/management_endpoints/credential_migration.py +++ b/litellm/proxy/management_endpoints/credential_migration.py @@ -40,15 +40,19 @@ from litellm.proxy.db.db_span import db_span if TYPE_CHECKING: from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.utils import PrismaClient -from litellm.proxy.common_utils.encrypt_decrypt_utils import ( - _ALGO_AES_GCM, - _ENCRYPTION_ALGORITHM_SETTING, - _V2_GCM_PREFIX, +from litellm.proxy.common_utils.encrypt_decrypt_utils import ( # noqa: F401 # legacy module exports + _ALGO_AES_GCM, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _ENCRYPTION_ALGORITHM_SETTING, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _V2_GCM_PREFIX, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + ALGO_AES_GCM, + ENCRYPTION_ALGORITHM_SETTING, + V2_GCM_PREFIX, SecretMapDecodeError, - _get_salt_key, + _get_salt_key, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export decode_secret_map, decrypt_value_helper, encrypt_value_helper, + get_salt_key, ) ValueClass = Literal["migrated", "legacy", "plaintext", "undecryptable", "not-a-string"] @@ -126,7 +130,7 @@ class MigrationReport: def is_migrated(value: object) -> bool: """True if ``value`` is already an AES-256-GCM (``v2:gcm:``) ciphertext.""" - return isinstance(value, str) and value.startswith(_V2_GCM_PREFIX) + return isinstance(value, str) and value.startswith(V2_GCM_PREFIX) def classify_value(value: object, key: str = "scan") -> ValueClass: @@ -145,7 +149,7 @@ def classify_value(value: object, key: str = "scan") -> ValueClass: return "not-a-string" if value == "": return "plaintext" - if value.startswith(_V2_GCM_PREFIX): + if value.startswith(V2_GCM_PREFIX): return "migrated" decrypted: Final = decrypt_value_helper(value=value, key=key, exception_type="debug", return_original_value=False) if decrypted is None: @@ -164,7 +168,7 @@ def reencrypt_value(value: object, key: str = "migrate") -> object: """ if not isinstance(value, str) or value == "": return value - if value.startswith(_V2_GCM_PREFIX): + if value.startswith(V2_GCM_PREFIX): return value # idempotent: already migrated decrypted: Final = decrypt_value_helper(value=value, key=key, exception_type="debug", return_original_value=False) if decrypted is None: @@ -197,11 +201,11 @@ def _assert_aes_gate_enabled() -> None: """ from litellm.proxy.proxy_server import general_settings - algo: Final = general_settings.get(_ENCRYPTION_ALGORITHM_SETTING) - if not (isinstance(algo, str) and algo.lower() == _ALGO_AES_GCM): + algo: Final = general_settings.get(ENCRYPTION_ALGORITHM_SETTING) + if not (isinstance(algo, str) and algo.lower() == ALGO_AES_GCM): raise RuntimeError( "Encryption migration requires general_settings.encryption_algorithm: " - f"'{_ALGO_AES_GCM}'. Current value: {algo!r}. Set it before migrating " + f"'{ALGO_AES_GCM}'. Current value: {algo!r}. Set it before migrating " "so re-encrypted values are written in the AES-256-GCM format." ) @@ -436,13 +440,13 @@ def _classify_callback_value(value: object) -> ValueClass: even when run with the AES write gate off. """ from litellm.proxy.common_utils.callback_utils import ( - _CALLBACK_VAR_ENCRYPTED_PREFIX, + CALLBACK_VAR_ENCRYPTED_PREFIX, ) if not isinstance(value, str): return "not-a-string" inner = value - inner = inner.removeprefix(_CALLBACK_VAR_ENCRYPTED_PREFIX) + inner = inner.removeprefix(CALLBACK_VAR_ENCRYPTED_PREFIX) # rebind-ok: pre-existing rebinding on a rename-only line return classify_value(inner, key="callback") @@ -595,17 +599,17 @@ async def _migrate_covered_tables(prisma_client: object, user_api_key_dict: obje already-v2 / scanned figures. Returns one report per covered location. """ from litellm.proxy.management_endpoints.key_management_endpoints import ( - _rotate_master_key, + rotate_master_key, ) pre: Final = {r.location: r for r in await _scan_covered_tables(prisma_client)} - current_key: Final = _get_salt_key() + current_key: Final = get_salt_key() if current_key is None: raise RuntimeError( "Cannot migrate covered tables: no salt key / master key is set. Set LITELLM_SALT_KEY before migrating." ) - await _rotate_master_key( + await rotate_master_key( prisma_client=cast("PrismaClient", prisma_client), user_api_key_dict=cast("UserAPIKeyAuth", user_api_key_dict), current_master_key=current_key, diff --git a/litellm/proxy/management_endpoints/customer_endpoints.py b/litellm/proxy/management_endpoints/customer_endpoints.py index 650e743b027..d604c58d1f5 100644 --- a/litellm/proxy/management_endpoints/customer_endpoints.py +++ b/litellm/proxy/management_endpoints/customer_endpoints.py @@ -37,9 +37,10 @@ from litellm.proxy.common_utils.user_api_key_cache import ( ) from litellm.proxy.management_endpoints.common_daily_activity import get_daily_activity from litellm.proxy.management_endpoints.common_utils import validate_budget_duration -from litellm.proxy.management_helpers.object_permission_utils import ( - _set_object_permission, +from litellm.proxy.management_helpers.object_permission_utils import ( # noqa: F401 # legacy module exports + _set_object_permission, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export handle_update_object_permission_common, + set_object_permission, ) from litellm.proxy.utils import handle_exception_on_proxy from litellm.repositories.budget_repository import BudgetRepository @@ -469,7 +470,7 @@ async def new_end_user( ## Handle Object Permission - MCP Servers, Vector Stores etc. new_end_user_obj = _STR_OBJECT_DICT.validate_python( - await _set_object_permission( + await set_object_permission( data_json=new_end_user_obj, prisma_client=prisma_client, ) diff --git a/litellm/proxy/management_endpoints/internal_user_endpoints.py b/litellm/proxy/management_endpoints/internal_user_endpoints.py index 1953370be39..2374f7d7aae 100644 --- a/litellm/proxy/management_endpoints/internal_user_endpoints.py +++ b/litellm/proxy/management_endpoints/internal_user_endpoints.py @@ -62,20 +62,23 @@ from litellm.proxy.management_endpoints.common_daily_activity import ( get_daily_activity, raise_public, ) -from litellm.proxy.management_endpoints.common_utils import ( - _user_has_admin_view, +from litellm.proxy.management_endpoints.common_utils import ( # noqa: F401 # legacy module exports + _user_has_admin_view, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export require_caller_user_id_for_non_admin, + user_api_key_has_admin_view, validate_budget_duration, validate_finite_spend, ) -from litellm.proxy.management_endpoints.key_management_endpoints import ( - _check_permissions_caller_permission, +from litellm.proxy.management_endpoints.key_management_endpoints import ( # noqa: F401 # legacy module exports + _check_permissions_caller_permission, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + check_permissions_caller_permission, generate_key_helper_fn, prepare_metadata_fields, ) -from litellm.proxy.management_helpers.object_permission_utils import ( - _set_object_permission, +from litellm.proxy.management_helpers.object_permission_utils import ( # noqa: F401 # legacy module exports + _set_object_permission, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export handle_update_object_permission_common, + set_object_permission, ) from litellm.proxy.management_helpers.utils import management_endpoint_wrapper from litellm.proxy.utils import handle_exception_on_proxy, hash_password @@ -592,7 +595,7 @@ async def new_user( if data.auto_create_key and isinstance(user_api_key_dict, UserAPIKeyAuth): enforce_batch_limits_are_admin_only(data, None, user_api_key_dict, "key") - _check_permissions_caller_permission( + check_permissions_caller_permission( data=data, user_api_key_dict=user_api_key_dict, ) @@ -602,7 +605,9 @@ async def new_user( # Persist the requested grants as their own row and link it, mirroring key/team creation. # generate_key_helper_fn only forwards object_permission_id, so without this the entitlement # the caller sent would be dropped on the floor. - data_json = await _set_object_permission(data_json=data_json, prisma_client=prisma_client) + data_json = await set_object_permission( # rebind-ok: pre-existing rebinding on a rename-only line + data_json=data_json, prisma_client=prisma_client + ) data_json.pop("password", None) teams = data.teams if teams is None: @@ -788,7 +793,7 @@ def _enforce_user_info_access(user_id: str | None, user_api_key_dict: UserAPIKey # Admin-view roles (PROXY_ADMIN and PROXY_ADMIN_VIEW_ONLY) bypass # ownership, mirroring the `/user/info` carve-out that # `RouteChecks.non_proxy_admin_allowed_routes_check` applies upstream. - if _user_has_admin_view(user_api_key_dict): + if user_api_key_has_admin_view(user_api_key_dict): return if user_id == user_api_key_dict.user_id: return @@ -969,7 +974,7 @@ async def user_info( raise Exception( "Database not connected. Connect a database to your proxy - https://docs.litellm.ai/docs/simple_proxy#managing-auth---virtual-keys" ) - if user_id is None and _user_has_admin_view(user_api_key_dict): + if user_id is None and user_api_key_has_admin_view(user_api_key_dict): return await _get_user_info_for_proxy_admin(user_api_key_dict=user_api_key_dict) elif user_id is None: user_id = user_api_key_dict.user_id @@ -1045,7 +1050,7 @@ async def _check_user_info_v2_access( ) # Rule 1: Proxy admins — fetch and return the target row directly - if _user_has_admin_view(user_api_key_dict): + if user_api_key_has_admin_view(user_api_key_dict): return await _fetch_target_user() # Rule 2: Self-lookup @@ -1423,9 +1428,9 @@ async def _invalidate_user_spend_counter_if_changed( and not safely subscriptable). """ if non_default_values.get("spend") is not None: - from litellm.proxy.proxy_server import _invalidate_spend_counter + from litellm.proxy.proxy_server import invalidate_spend_counter - await _invalidate_spend_counter(counter_key=f"spend:user:{non_default_values['user_id']}") + await invalidate_spend_counter(counter_key=f"spend:user:{non_default_values['user_id']}") def _clears_object_permission(user_request: UpdateUserRequest) -> bool: @@ -1488,7 +1493,7 @@ async def _update_single_user_helper( if not user_request.user_id and not user_request.user_email: raise ValueError("Either user_id or user_email must be provided") - _check_permissions_caller_permission( + check_permissions_caller_permission( data=user_request, user_api_key_dict=user_api_key_dict, ) @@ -2156,7 +2161,7 @@ async def _authorize_user_list_request( - Org admins: returns comma-separated org IDs scoped to their allowed orgs. - Others: raises 403. """ - if _user_has_admin_view(user_api_key_dict): + if user_api_key_has_admin_view(user_api_key_dict): return organization_ids if user_api_key_dict.user_id is None: @@ -2430,7 +2435,7 @@ async def delete_user( - user_ids: List[str] - The list of user id's to be deleted. """ from litellm.proxy.management_endpoints.team_endpoints import ( - _cleanup_members_with_roles, + cleanup_members_with_roles, ) from litellm.proxy.management_helpers.audit_logs import ( get_audit_log_changed_by, @@ -2546,7 +2551,7 @@ async def delete_user( ).table.find_many(where={"team_id": {"in": user_row.teams}}) teams_to_update: list[tuple[str, str]] = [] for team in fetch_all_teams: - removed_team_members, new_team_members = _cleanup_members_with_roles( + removed_team_members, new_team_members = cleanup_members_with_roles( existing_team_row=LiteLLM_TeamTable.model_validate(team.model_dump()), data=TeamMemberDeleteRequest( team_id=team.team_id, @@ -2680,7 +2685,7 @@ async def _resolve_org_filter_for_user_search( if not ui_settings.get("scope_user_search_to_org", False): return None # flag OFF — no filtering - if _user_has_admin_view(user_api_key_dict): + if user_api_key_has_admin_view(user_api_key_dict): return None # proxy admin — see everything # Try to resolve org admin memberships @@ -2875,7 +2880,7 @@ async def ui_view_users( def resolve_user_daily_activity_entity_ids( *, user_id: str | None, user_api_key_dict: UserAPIKeyAuth ) -> tuple[str, ...] | None | ScopeDenied: - if _user_has_admin_view(user_api_key_dict): + if user_api_key_has_admin_view(user_api_key_dict): return (user_id,) if user_id is not None else None caller_user_id: Final = require_caller_user_id_for_non_admin(user_api_key_dict) diff --git a/litellm/proxy/management_endpoints/jwt_key_mapping_endpoints.py b/litellm/proxy/management_endpoints/jwt_key_mapping_endpoints.py index 292cec1346d..531fbdef2bc 100644 --- a/litellm/proxy/management_endpoints/jwt_key_mapping_endpoints.py +++ b/litellm/proxy/management_endpoints/jwt_key_mapping_endpoints.py @@ -17,7 +17,10 @@ from litellm.proxy._types import ( from litellm.proxy.auth.auth_checks import jwt_key_mapping_cache_key from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import evict_and_broadcast -from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view +from litellm.proxy.management_endpoints.common_utils import ( # noqa: F401 # legacy module exports + _user_has_admin_view, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + user_api_key_has_admin_view, +) from litellm.repositories.table_repositories import JWTKeyMappingRepository router: Final = APIRouter() @@ -310,7 +313,7 @@ async def list_jwt_key_mappings( from litellm.proxy.proxy_server import prisma_client # Admin Viewer follows the read-parity rule. - if not _user_has_admin_view(user_api_key_dict): + if not user_api_key_has_admin_view(user_api_key_dict): raise HTTPException(status_code=403, detail="Only proxy admins can list JWT key mappings") if prisma_client is None: @@ -348,7 +351,7 @@ async def info_jwt_key_mapping( from litellm.proxy.proxy_server import prisma_client # Admin Viewer follows the read-parity rule. - if not _user_has_admin_view(user_api_key_dict): + if not user_api_key_has_admin_view(user_api_key_dict): raise HTTPException(status_code=403, detail="Only proxy admins can get JWT key mapping info") if prisma_client is None: diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 6063ee21250..ac9dfeb136d 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -53,9 +53,10 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.sso_assertion_s ) from litellm.proxy._types import * from litellm.proxy._types import Litellm_EntityType, LiteLLM_VerificationToken, hash_token -from litellm.proxy.auth.auth_checks import ( - _delete_cache_key_object, +from litellm.proxy.auth.auth_checks import ( # noqa: F401 # legacy module exports + _delete_cache_key_object, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export can_team_access_model, + delete_cache_key_object, get_jwt_key_mapping_cache_keys_for_token, get_key_end_user_budget_id, get_org_object, @@ -88,19 +89,25 @@ from litellm.proxy.hooks.key_management_event_hooks import KeyManagementEventHoo from litellm.proxy.hooks.model_max_budget_limiter import build_model_max_budget_usage from litellm.proxy.management.teams.authz import TEAM_ADMIN_ONLY, TEAM_OR_ORG_ADMIN, is_team_admin from litellm.proxy.management.teams.dependencies import get_team_access -from litellm.proxy.management_endpoints.common_utils import ( - _check_disable_global_guardrails_caller_permission, - _check_passthrough_routes_caller_permission, - _set_object_metadata_field, - _team_member_has_permission, - _user_has_admin_view, +from litellm.proxy.management_endpoints.common_utils import ( # noqa: F401 # legacy module exports + _check_disable_global_guardrails_caller_permission, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _check_passthrough_routes_caller_permission, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _set_object_metadata_field, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _team_member_has_permission, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _user_has_admin_view, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export check_allowed_passthrough_routes_caller_permission, check_denied_passthrough_routes_caller_permission, + check_disable_global_guardrails_caller_permission, + check_passthrough_routes_caller_permission, + set_object_metadata_field, + team_member_has_permission, + user_api_key_has_admin_view, validate_budget_duration, validate_finite_spend, ) -from litellm.proxy.management_endpoints.model_management_endpoints import ( - _add_model_to_db, +from litellm.proxy.management_endpoints.model_management_endpoints import ( # noqa: F401 # legacy module exports + _add_model_to_db, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + add_model_to_db, ) from litellm.proxy.management_endpoints.router_weights import validate_router_settings_weights from litellm.proxy.management_endpoints.team_admin_field_permissions import ( @@ -114,11 +121,12 @@ from litellm.proxy.management_helpers.access_group_key_sync import ( sync_key_update_access_group_membership, ) from litellm.proxy.management_helpers.key_settings_audit import with_settings_updated_at -from litellm.proxy.management_helpers.object_permission_utils import ( - _set_object_permission, +from litellm.proxy.management_helpers.object_permission_utils import ( # noqa: F401 # legacy module exports + _set_object_permission, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export attach_object_permission_to_dict, handle_update_object_permission_common, invalidate_cached_object_permissions, + set_object_permission, validate_key_mcp_servers_against_team, validate_key_search_tools_against_team, validate_key_vector_stores_against_team, @@ -130,12 +138,16 @@ from litellm.proxy.management_helpers.utils import management_endpoint_wrapper from litellm.proxy.search_endpoints.search_tool_registry import rotate_search_tools_master_key from litellm.proxy.spend_tracking.budget_reservation import get_budget_window_start from litellm.proxy.spend_tracking.spend_counter_batch import SPEND_COUNTERS_TARGET -from litellm.proxy.spend_tracking.spend_tracking_utils import _is_master_key -from litellm.proxy.utils import ( +from litellm.proxy.spend_tracking.spend_tracking_utils import ( # noqa: F401 # legacy module exports + _is_master_key, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + is_master_key, +) +from litellm.proxy.utils import ( # noqa: F401 # legacy module exports PrismaClient, ProxyLogging, - _hash_token_if_needed, + _hash_token_if_needed, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export handle_exception_on_proxy, + hash_token_if_needed, is_valid_api_key, ) from litellm.repositories.base_repository import BaseRepository @@ -526,7 +538,7 @@ def _get_user_in_team(team_table: LiteLLM_TeamTableCachedObj, user_id: str | Non return None -def _get_caller_team_role( +def get_caller_team_role( team_table: LiteLLM_TeamTableCachedObj, user_api_key_dict: UserAPIKeyAuth, ) -> Literal["admin", "user"] | None: @@ -536,7 +548,10 @@ def _get_caller_team_role( return None if member is None else member.role -def _calculate_key_rotation_time(rotation_interval: str) -> datetime: +_get_caller_team_role: Final = get_caller_team_role + + +def calculate_key_rotation_time(rotation_interval: str) -> datetime: """ Helper function to calculate the next rotation time for a key based on the rotation interval. @@ -551,6 +566,9 @@ def _calculate_key_rotation_time(rotation_interval: str) -> datetime: return now + timedelta(seconds=interval_seconds) +_calculate_key_rotation_time: Final = calculate_key_rotation_time + + def _set_key_rotation_fields( data: dict, auto_rotate: bool, @@ -583,7 +601,7 @@ def _set_key_rotation_fields( { "auto_rotate": auto_rotate, "rotation_interval": rotation_interval, - "key_rotation_at": _calculate_key_rotation_time(rotation_interval), + "key_rotation_at": calculate_key_rotation_time(rotation_interval), } ) @@ -630,7 +648,7 @@ def _team_key_operation_team_member_check( detail=f"User={assigned_user_id} not assigned to team={team_table.team_id}", ) - caller_team_role: Final = _get_caller_team_role(team_table=team_table, user_api_key_dict=user_api_key_dict) + caller_team_role: Final = get_caller_team_role(team_table=team_table, user_api_key_dict=user_api_key_dict) is_admin: Final = ( user_api_key_dict.user_role is not None and user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value @@ -1007,7 +1025,7 @@ def _enforce_allowed_routes_update_permission( ) -def _check_permissions_caller_permission( +def check_permissions_caller_permission( data: GenerateRequestBase, user_api_key_dict: UserAPIKeyAuth, ) -> None: @@ -1029,6 +1047,9 @@ def _check_permissions_caller_permission( ) +_check_permissions_caller_permission: Final = check_permissions_caller_permission + + def _check_budget_limits_delegation_ceiling( budget_limits: list[BudgetLimitEntry] | None, delegation_ceiling: float | None, @@ -1324,11 +1345,11 @@ async def _common_key_generation_helper( is_ui_session_team_key=is_ui_session_team_key, team_table=team_table, ) - _check_permissions_caller_permission( + check_permissions_caller_permission( data=data, user_api_key_dict=user_api_key_dict, ) - _check_disable_global_guardrails_caller_permission( + check_disable_global_guardrails_caller_permission( data.disable_global_guardrails, _requested_metadata, user_api_key_dict, @@ -1379,7 +1400,7 @@ async def _common_key_generation_helper( # Set Management Endpoint Metadata Fields for field in LiteLLM_ManagementEndpoint_MetadataFields_Premium: if getattr(data, field, None) is not None: - _set_object_metadata_field( + set_object_metadata_field( object_data=data, field_name=field, value=getattr(data, field), @@ -1388,7 +1409,7 @@ async def _common_key_generation_helper( for field in LiteLLM_ManagementEndpoint_MetadataFields: if getattr(data, field, None) is not None: - _set_object_metadata_field( + set_object_metadata_field( object_data=data, field_name=field, value=getattr(data, field), @@ -1480,7 +1501,7 @@ async def _common_key_generation_helper( for _op_field, _op_default_value in _default_object_permission.items(): _caller_object_permission.setdefault(_op_field, _op_default_value) - data_json = await _set_object_permission( + data_json = await set_object_permission( # rebind-ok: pre-existing rebinding on a rename-only line data_json=data_json, prisma_client=prisma_client, ) @@ -1744,7 +1765,7 @@ async def _check_team_key_limits( # Exclude the key being updated to avoid double-counting its limits. # data.key may be a raw key (sk-...) or a pre-hashed token_id. if isinstance(data, UpdateKeyRequest) and data.key is not None: - hashed_key: Final = _hash_token_if_needed(data.key) + hashed_key: Final = hash_token_if_needed(data.key) keys = [key for key in keys if key.token != hashed_key] check_team_key_model_specific_limits( keys=keys, @@ -1932,7 +1953,7 @@ async def _check_org_key_limits( # Exclude the key being updated to avoid double-counting its limits. # data.key may be a raw key (sk-...) or a pre-hashed token_id. if isinstance(data, UpdateKeyRequest) and data.key is not None: - hashed_key: Final = _hash_token_if_needed(data.key) + hashed_key: Final = hash_token_if_needed(data.key) keys = [key for key in keys if key.token != hashed_key] check_org_key_model_specific_limits( keys=keys, @@ -2088,7 +2109,7 @@ async def generate_key_fn( user_api_key_dict=user_api_key_dict, allowed_routes_was_provided="allowed_routes" in data.model_fields_set, ) - _check_passthrough_routes_caller_permission( + check_passthrough_routes_caller_permission( data=data, user_api_key_dict=user_api_key_dict, ) @@ -2262,7 +2283,7 @@ async def generate_service_account_key_fn( user_api_key_dict=user_api_key_dict, allowed_routes_was_provided="allowed_routes" in data.model_fields_set, ) - _check_passthrough_routes_caller_permission( + check_passthrough_routes_caller_permission( data=data, user_api_key_dict=user_api_key_dict, ) @@ -2368,10 +2389,10 @@ def prepare_metadata_fields(data: BaseModel, non_default_values: dict, existing_ else: casted_metadata[k] = v if k in LiteLLM_ManagementEndpoint_MetadataFields_Premium: - from litellm.proxy.utils import _premium_user_check + from litellm.proxy.utils import premium_user_check if v: - _premium_user_check(k) + premium_user_check(k) casted_metadata[k] = v except Exception as e: @@ -2441,7 +2462,7 @@ async def _update_key_row_with_soft_budget( existing_key_row: LiteLLM_VerificationToken, changed_by: str, ) -> _KeyUpdateResult: - hashed_token: Final = _hash_token_if_needed(key) + hashed_token: Final = hash_token_if_needed(key) key_where: Final[_KeyRowWhere] = {"token": hashed_token} tx: _KeyUpdateTx async with prisma_client.tx() as tx: @@ -2508,7 +2529,7 @@ async def prepare_key_update_data( # Set Management Endpoint Metadata Fields for field in LiteLLM_ManagementEndpoint_MetadataFields_Premium: if getattr(data, field, None) is not None: - _set_object_metadata_field( + set_object_metadata_field( object_data=data, field_name=field, value=getattr(data, field), @@ -2645,7 +2666,7 @@ async def _get_and_validate_existing_key( ) if token is not None: - hashed_token: Final = _hash_token_if_needed(token=token) + hashed_token: Final = hash_token_if_needed(token=token) existing_key_row: Final[LiteLLM_VerificationToken | None] = await _prisma_table( VerificationTokenRepository(prisma_client) @@ -2742,7 +2763,7 @@ async def _process_single_key_update( # Validate max_budget _validate_max_budget(update_key_request.max_budget) - _check_permissions_caller_permission( + check_permissions_caller_permission( data=update_key_request, user_api_key_dict=user_api_key_dict, ) @@ -2754,7 +2775,7 @@ async def _process_single_key_update( prisma_client=prisma_client, ) - _check_disable_global_guardrails_caller_permission( + check_disable_global_guardrails_caller_permission( update_key_request.disable_global_guardrails, update_key_request.metadata, # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # request models declare `metadata` as bare dict user_api_key_dict, @@ -2885,8 +2906,8 @@ async def _process_single_key_update( ), user_api_key_cache=user_api_key_cache, ) - await _delete_cache_key_object( - hashed_token=_hash_token_if_needed(key_request.key), + await delete_cache_key_object( + hashed_token=hash_token_if_needed(key_request.key), user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, ) @@ -2895,7 +2916,7 @@ async def _process_single_key_update( # authenticating against the access groups it just lost. await sync_key_update_access_group_membership( prisma_client=prisma_client, - key_token=_hash_token_if_needed(_resolve_token_to_update(data=key_request, existing_key_row=existing_key_row)), + key_token=hash_token_if_needed(_resolve_token_to_update(data=key_request, existing_key_row=existing_key_row)), data=key_request, existing_key_row=existing_key_row, ) @@ -3100,11 +3121,11 @@ async def _validate_update_key_data( user_api_key_dict=user_api_key_dict, ) check_allowed_passthrough_routes_caller_permission(data, user_api_key_dict) - _check_permissions_caller_permission( + check_permissions_caller_permission( data=data, user_api_key_dict=user_api_key_dict, ) - _check_disable_global_guardrails_caller_permission( + check_disable_global_guardrails_caller_permission( data.disable_global_guardrails, data.metadata, # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # request models declare `metadata` as bare dict user_api_key_dict, @@ -3570,8 +3591,8 @@ async def update_key_fn( ), user_api_key_cache=user_api_key_cache, ) - await _delete_cache_key_object( - hashed_token=_hash_token_if_needed(key), + await delete_cache_key_object( + hashed_token=hash_token_if_needed(key), user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, ) @@ -3580,7 +3601,7 @@ async def update_key_fn( # authenticating against the access groups it just lost. await sync_key_update_access_group_membership( prisma_client=prisma_client, - key_token=_hash_token_if_needed(key), + key_token=hash_token_if_needed(key), data=data, existing_key_row=existing_key_row, ) @@ -3588,7 +3609,7 @@ async def update_key_fn( if data.spend is not None: from litellm.proxy.proxy_server import spend_counter_cache - counter_key: Final = f"spend:key:{_hash_token_if_needed(key)}" + counter_key: Final = f"spend:key:{hash_token_if_needed(key)}" spend_counter_cache.in_memory_cache.set_cache(key=counter_key, value=data.spend, ttl=60) if spend_counter_cache.redis_cache is not None: try: @@ -3924,7 +3945,7 @@ async def bulk_update_team_keys( hashed_key_ids: Final = [] seen_hashes: Final = set() for k in data.key_ids: - h = _hash_token_if_needed(k) + h = hash_token_if_needed(k) if h in seen_hashes: continue seen_hashes.add(h) @@ -3970,7 +3991,7 @@ async def bulk_update_team_keys( failed_updates: Final[list[FailedKeyUpdate]] = [] for token in requested_tokens: - db_token = _hash_token_if_needed(token) + db_token = hash_token_if_needed(token) try: if db_token not in existing_by_token: raise HTTPException( @@ -4440,7 +4461,7 @@ async def info_key_fn( key = key or user_api_key_dict.api_key hashed_key: str | None = key if key is not None: - hashed_key = _hash_token_if_needed(token=key) + hashed_key = hash_token_if_needed(token=key) # rebind-ok: pre-existing rebinding on a rename-only line live_key_info: Final = await _prisma_table(VerificationTokenRepository(prisma_client)).find_unique( where={"token": hashed_key}, include={"litellm_budget_table": True}, @@ -5012,7 +5033,7 @@ async def delete_verification_tokens( failed_tokens: list = [] try: if prisma_client: - hashed_tokens: Final[list[str]] = [_hash_token_if_needed(token=key) for key in tokens] + hashed_tokens: Final[list[str]] = [hash_token_if_needed(token=key) for key in tokens] tokens = hashed_tokens _keys_being_deleted: Final[list[LiteLLM_VerificationToken]] = cast( # cast-ok: find_many returns a list "list[LiteLLM_VerificationToken]", @@ -5044,7 +5065,7 @@ async def delete_verification_tokens( status_code=status.HTTP_403_FORBIDDEN, detail={"error": "You are not authorized to delete this key"}, ) - await _persist_deleted_verification_tokens( + await persist_deleted_verification_tokens( keys=authorized_keys, prisma_client=prisma_client, user_api_key_dict=user_api_key_dict, @@ -5188,7 +5209,7 @@ async def _save_deleted_verification_token_records( await _deleted_verification_token_table(prisma_client).create_many(data=records) -async def _persist_deleted_verification_tokens( +async def persist_deleted_verification_tokens( keys: Sequence[LiteLLM_VerificationToken], prisma_client: PrismaClient, user_api_key_dict: UserAPIKeyAuth, @@ -5208,6 +5229,9 @@ async def _persist_deleted_verification_tokens( ) +_persist_deleted_verification_tokens: Final = persist_deleted_verification_tokens + + async def delete_key_aliases( key_aliases: list[str], user_api_key_cache: UserApiKeyCache, @@ -5228,7 +5252,7 @@ async def delete_key_aliases( ) -async def _rotate_master_key( +async def rotate_master_key( prisma_client: PrismaClient, user_api_key_dict: UserAPIKeyAuth, current_master_key: str, @@ -5266,7 +5290,7 @@ async def _rotate_master_key( reencrypted for model in decrypted_models if ( - reencrypted := await _add_model_to_db( + reencrypted := await add_model_to_db( model_params=Deployment(**model), user_api_key_dict=user_api_key_dict, prisma_client=prisma_client, @@ -5299,10 +5323,10 @@ async def _rotate_master_key( environment_variables_dict = _env_vars_param_value(c) if environment_variables_dict: - decrypted_env_vars: Final = proxy_config._decrypt_and_set_db_env_variables( + decrypted_env_vars: Final = proxy_config.decrypt_and_set_db_env_variables( environment_variables=dict[str, str](environment_variables_dict) ) - encrypted_env_vars: Final = proxy_config._encrypt_env_variables( + encrypted_env_vars: Final = proxy_config.encrypt_env_variables( environment_variables=decrypted_env_vars, new_encryption_key=new_master_key, ) @@ -5399,6 +5423,9 @@ async def _rotate_master_key( verbose_proxy_logger.debug("Successfully re-encrypted %s credentials with new master key", len(credentials)) +_rotate_master_key: Final = rotate_master_key + + def _require_proxy_admin(user_api_key_dict: UserAPIKeyAuth) -> None: from litellm.proxy._types import CommonProxyErrors @@ -5669,7 +5696,7 @@ async def _execute_virtual_key_regeneration( prisma_client=prisma_client, ) - await _persist_deleted_verification_tokens( + await persist_deleted_verification_tokens( keys=[key_in_db], prisma_client=prisma_client, user_api_key_dict=user_api_key_dict, @@ -5702,8 +5729,8 @@ async def _execute_virtual_key_regeneration( user_api_key_cache=user_api_key_cache, ) if hashed_api_key or key: - await _delete_cache_key_object( - hashed_token=_hash_token_if_needed(key), + await delete_cache_key_object( + hashed_token=hash_token_if_needed(key), user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, ) @@ -5739,7 +5766,7 @@ def _check_regenerate_guardrail_opt_out( ) -> None: if data is None: return - _check_disable_global_guardrails_caller_permission( + check_disable_global_guardrails_caller_permission( data.disable_global_guardrails, data.metadata, # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # request models declare `metadata` as bare dict user_api_key_dict, @@ -5836,7 +5863,7 @@ async def regenerate_key_fn( allowed_routes_was_provided="allowed_routes" in data.model_fields_set, ) check_allowed_passthrough_routes_caller_permission(data, user_api_key_dict) - _check_permissions_caller_permission( + check_permissions_caller_permission( data=data, user_api_key_dict=user_api_key_dict, ) @@ -5864,7 +5891,7 @@ async def regenerate_key_fn( is_master_key_regeneration: Final = ( data is not None and data.new_master_key is not None - and _is_master_key(api_key=regenerate_target_key, _master_key=master_key) + and is_master_key(api_key=regenerate_target_key, _master_key=master_key) ) if ( @@ -5892,7 +5919,7 @@ async def regenerate_key_fn( detail={"error": "DB not connected. prisma_client is None"}, ) - _is_master_key_valid: Final = _is_master_key(api_key=key, _master_key=master_key) + _is_master_key_valid: Final = is_master_key(api_key=key, _master_key=master_key) if master_key is not None and data and _is_master_key_valid: if data.new_master_key is None: @@ -5900,7 +5927,7 @@ async def regenerate_key_fn( status_code=status.HTTP_400_BAD_REQUEST, detail={"error": "New master key is required."}, ) - await _rotate_master_key( + await rotate_master_key( prisma_client=prisma_client, user_api_key_dict=user_api_key_dict, current_master_key=master_key, @@ -6292,7 +6319,7 @@ async def reset_key_spend_fn( # a later write would re-fetch and re-cache the pre-write row, pinning # that pod to the stale budget_limits/spend for the rest of its own # cache TTL even though the DB is already correct. - await _delete_cache_key_object( + await delete_cache_key_object( hashed_token=hashed_api_key, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, @@ -6324,7 +6351,7 @@ async def validate_key_list_check( key_hash: str | None, prisma_client: PrismaClient, ) -> LiteLLM_UserTable | None: - if _user_has_admin_view(user_api_key_dict): + if user_api_key_has_admin_view(user_api_key_dict): return None if user_api_key_dict.user_id is None: @@ -6451,7 +6478,7 @@ def _get_team_ids_with_key_list_permission_from_objects( team.team_id for team in team_objects if not is_team_admin(user_api_key_dict=user_api_key_dict, team_obj=team) - and _team_member_has_permission( + and team_member_has_permission( user_api_key_dict=user_api_key_dict, team_obj=team, permission=KeyManagementRoutes.KEY_LIST.value, @@ -6672,7 +6699,7 @@ async def list_keys( if not user_id and not is_proxy_admin: user_id = user_api_key_dict.user_id - response: Final = await _list_key_helper( + response: Final = await list_key_helper( prisma_client=prisma_client, page=page, size=size, @@ -7074,7 +7101,7 @@ def _build_key_filter_conditions( return combined_where -async def _list_key_helper( +async def list_key_helper( prisma_client: PrismaClient, page: int, size: int, @@ -7259,6 +7286,9 @@ async def _list_key_helper( ) +_list_key_helper: Final = list_key_helper + + def _get_condition_to_filter_out_ui_session_tokens() -> Mapping[str, object]: """ Condition to filter out UI session tokens @@ -7427,7 +7457,7 @@ async def block_key( ) ## UPDATE KEY CACHE - invalidate so next read re-fetches from DB - await _delete_cache_key_object( + await delete_cache_key_object( hashed_token=hashed_token, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, @@ -7541,7 +7571,7 @@ async def unblock_key( ) ## UPDATE KEY CACHE - invalidate so next read re-fetches from DB - await _delete_cache_key_object( + await delete_cache_key_object( hashed_token=hashed_token, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, diff --git a/litellm/proxy/management_endpoints/management_v1/spend_logs.py b/litellm/proxy/management_endpoints/management_v1/spend_logs.py index b122840004c..b93c05c555f 100644 --- a/litellm/proxy/management_endpoints/management_v1/spend_logs.py +++ b/litellm/proxy/management_endpoints/management_v1/spend_logs.py @@ -48,9 +48,9 @@ async def _spend_log_scope_clause( applies, so a dropdown can never offer a value from a row the caller could not open. """ - from litellm.proxy.spend_tracking.spend_management_endpoints import _is_admin_view_safe, read_scope_sql + from litellm.proxy.spend_tracking.spend_management_endpoints import is_admin_view_safe, read_scope_sql - if _is_admin_view_safe(user_api_key_dict=user_api_key_dict): + if is_admin_view_safe(user_api_key_dict=user_api_key_dict): return None, () scope: Final = await resolve_owned_read_scope( user_api_key_dict.user_id, partial(log_team_lookup, user_api_key_dict) diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index 82022639095..5f65bc498ad 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -176,13 +176,14 @@ if MCP_AVAILABLE: store_user_oauth_credential, update_mcp_server, ) - from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( - _raise_if_not_oauth2, + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( # noqa: F401 # legacy module exports + _raise_if_not_oauth2, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export authorize_with_server, client_supplied_application_type, client_supplied_redirect_uris, exchange_token_with_server, get_request_base_url, + raise_if_not_oauth2, redeem_passthrough_authorization_code, register_client_with_server, resolve_ephemeral_dcr_client, @@ -225,15 +226,20 @@ if MCP_AVAILABLE: UserMCPManagementMode, is_per_server_oauth_discovery_eligible, ) - from litellm.proxy.auth.user_api_key_auth import ( - _user_api_key_auth_builder, + from litellm.proxy.auth.user_api_key_auth import ( # noqa: F401 # legacy module exports + _user_api_key_auth_builder, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export user_api_key_auth, + user_api_key_auth_builder, ) - from litellm.proxy.common_utils.http_parsing_utils import ( - _read_request_body, + from litellm.proxy.common_utils.http_parsing_utils import ( # noqa: F401 # legacy module exports + _read_request_body, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export populate_request_with_path_params, + read_request_body, + ) + from litellm.proxy.management_endpoints.common_utils import ( # noqa: F401 # legacy module exports + _user_has_admin_view, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + user_api_key_has_admin_view, ) - from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view from litellm.proxy.management_endpoints.mcp_connector_import import ( ConnectorConversionError, ConvertedConnector, @@ -913,7 +919,7 @@ if MCP_AVAILABLE: or (payload.auth_type is not None and payload.auth_type != existing.auth_type) ) - def _inherit_credentials_from_existing_server( + def inherit_credentials_from_existing_server( payload: NewMCPServerRequest, ) -> NewMCPServerRequest: if not payload.server_id: @@ -982,6 +988,8 @@ if MCP_AVAILABLE: payload_dict["credentials"] = inherited_credentials return NewMCPServerRequest.model_validate(payload_dict) + _inherit_credentials_from_existing_server: Final = inherit_credentials_from_existing_server + async def _resolve_session_server_id(payload: NewMCPServerRequest) -> str: """Decide the id an OAuth session runs under. @@ -1204,8 +1212,8 @@ if MCP_AVAILABLE: """ from litellm.proxy.auth.auth_checks import get_team_object from litellm.proxy.management_helpers.object_permission_utils import ( - _get_allow_all_keys_server_ids, - _get_team_allowed_mcp_servers, + get_allow_all_keys_server_ids, + get_team_allowed_mcp_servers, ) from litellm.proxy.proxy_server import prisma_client, user_api_key_cache @@ -1216,8 +1224,8 @@ if MCP_AVAILABLE: check_db_only=True, ) - team_server_ids: Final = await _get_team_allowed_mcp_servers(team_obj) - allow_all_server_ids: Final = _get_allow_all_keys_server_ids() + team_server_ids: Final = await get_team_allowed_mcp_servers(team_obj) + allow_all_server_ids: Final = get_allow_all_keys_server_ids() all_allowed_ids: Final = team_server_ids | allow_all_server_ids if not all_allowed_ids: @@ -1228,7 +1236,7 @@ if MCP_AVAILABLE: for server_id in all_allowed_ids: server = global_mcp_server_manager.get_mcp_server_by_id(server_id) if server is not None: - mcp_server_table = global_mcp_server_manager._build_mcp_server_table(server) + mcp_server_table = global_mcp_server_manager.build_mcp_server_table(server) servers.append(mcp_server_table) return _redact_mcp_credentials_list(servers) @@ -1314,7 +1322,7 @@ if MCP_AVAILABLE: # Only proxy admins may query another team's MCP servers. # Non-admins must belong to the requested team. sanitized_team_id: Final = team_id.strip() - is_admin: Final = _user_has_admin_view(user_api_key_dict) + is_admin: Final = user_api_key_has_admin_view(user_api_key_dict) if not is_admin: from litellm.proxy.auth.auth_checks import get_team_object from litellm.proxy.proxy_server import ( @@ -1820,7 +1828,7 @@ if MCP_AVAILABLE: from litellm.proxy.auth.ip_address_utils import IPAddressUtils client_ip: Final = IPAddressUtils.get_mcp_client_ip(request) - is_admin_view: Final = _user_has_admin_view(user_api_key_dict) + is_admin_view: Final = user_api_key_has_admin_view(user_api_key_dict) is_restricted_virtual_key: Final = _is_restricted_virtual_key_request(user_api_key_dict) resolved: Final = await resolve_mcp_server( server_id, @@ -2124,7 +2132,7 @@ if MCP_AVAILABLE: ) created_by: Final = user_api_key_dict.user_id or LITELLM_PROXY_ADMIN_NAME - payload_with_credentials: Final = _inherit_credentials_from_existing_server(payload) + payload_with_credentials: Final = inherit_credentials_from_existing_server(payload) temp_record: Final = _build_temporary_mcp_server_record( payload_with_credentials, created_by, @@ -2229,7 +2237,7 @@ if MCP_AVAILABLE: # grants must NOT bypass auth (see comment above). path_lower: Final = get_request_route(request).rstrip("/").lower() if path_lower.endswith("/token"): - body_data: Final = await _read_request_body(request=request) + body_data: Final = await read_request_body(request=request) grant_type: Final = (body_data or {}).get("grant_type", "") if grant_type != "authorization_code": # Fall through to normal LiteLLM auth (will 401 if @@ -2244,10 +2252,12 @@ if MCP_AVAILABLE: # token can be minted via the redirect alone. return UserAPIKeyAuth() - request_data = await _read_request_body(request=request) + request_data = await read_request_body( # rebind-ok: pre-existing rebinding on a rename-only line + request=request + ) request_data = populate_request_with_path_params(request_data=request_data, request=request) - return await _user_api_key_auth_builder( + return await user_api_key_auth_builder( request=request, api_key=api_key, azure_api_key_header="", @@ -2275,7 +2285,7 @@ if MCP_AVAILABLE: authorized: Final = await catalog.resolve( server_id, user_api_key_dict, - is_admin_view=_user_has_admin_view(user_api_key_dict), + is_admin_view=user_api_key_has_admin_view(user_api_key_dict), not_found_detail={"error": f"MCP server {server_id} not found"}, forbidden_detail={"error": f"Access denied to MCP server {server_id}"}, non_admin_missing="not_found", @@ -2313,7 +2323,7 @@ if MCP_AVAILABLE: scope: str | None = None, ): async with _oauth_server_operation(server_id, user_api_key_dict, request=request) as mcp_server: - _raise_if_not_oauth2(mcp_server) + raise_if_not_oauth2(mcp_server) # Use the server's stored client_id when the caller doesn't supply one stored_or_supplied_client_id: Final = mcp_server.client_id or client_id or "" ephemeral_dcr_client: Final = ( @@ -2373,7 +2383,7 @@ if MCP_AVAILABLE: scope: str | None = Form(None), ): async with _oauth_server_operation(server_id, user_api_key_dict, request=request) as mcp_server: - _raise_if_not_oauth2(mcp_server) + raise_if_not_oauth2(mcp_server) # Sealed passthrough codes exist only for the authorization_code grant. A refresh_token # grant must never open one: the minted client is unrecoverable after the single flow by # contract, so an expired browser-held token re-runs authorize instead. @@ -2425,7 +2435,7 @@ if MCP_AVAILABLE: user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): async with _oauth_server_operation(server_id, user_api_key_dict, request=request) as mcp_server: - request_data: Final = await _read_request_body(request=request) + request_data: Final = await read_request_body(request=request) data: Final[Mapping[str, object]] = {**request_data} client_redirect_uris: Final = client_supplied_redirect_uris(data.get("redirect_uris")) client_application_type: Final = client_supplied_application_type(data.get("application_type")) @@ -2748,7 +2758,7 @@ if MCP_AVAILABLE: servers: Final = {srv.server_id: srv for srv in await get_mcp_servers(prisma_client, server_ids)} allowed_server_ids: Final = ( None - if _user_has_admin_view(user_api_key_dict) + if user_api_key_has_admin_view(user_api_key_dict) else frozenset[str]().union( *[ await global_mcp_server_manager.get_allowed_mcp_servers(context) @@ -2839,7 +2849,7 @@ if MCP_AVAILABLE: authorized: Final = await catalog.resolve( server_id, user_api_key_dict, - is_admin_view=_user_has_admin_view(user_api_key_dict), + is_admin_view=user_api_key_has_admin_view(user_api_key_dict), not_found_detail={"error": f"MCP Server {server_id} not found"}, forbidden_detail={ "error": ( @@ -3317,7 +3327,7 @@ if MCP_AVAILABLE: Used by the UI to show a discovery grid when adding new MCP servers. """ # Admin Viewer follows the read-parity rule. - if not _user_has_admin_view(user_api_key_dict): + if not user_api_key_has_admin_view(user_api_key_dict): raise HTTPException( status_code=403, detail={ @@ -3373,7 +3383,7 @@ if MCP_AVAILABLE: user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): # Admin Viewer follows the read-parity rule. - if not _user_has_admin_view(user_api_key_dict): + if not user_api_key_has_admin_view(user_api_key_dict): raise HTTPException( status_code=403, detail={ @@ -3452,7 +3462,7 @@ if MCP_AVAILABLE: ): """Return toolsets the calling key is allowed to access.""" prisma_client: Final = get_prisma_client_or_throw("Database not connected. Connect a database to your proxy") - if _user_has_admin_view(user_api_key_dict): + if user_api_key_has_admin_view(user_api_key_dict): op: Final = user_api_key_dict.object_permission if op is None or not op.mcp_toolsets: return await list_mcp_toolsets(prisma_client) @@ -3472,7 +3482,7 @@ if MCP_AVAILABLE: user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): prisma_client: Final = get_prisma_client_or_throw("Database not connected. Connect a database to your proxy") - if not _user_has_admin_view(user_api_key_dict) and toolset_id not in await granted_toolset_ids( + if not user_api_key_has_admin_view(user_api_key_dict) and toolset_id not in await granted_toolset_ids( user_api_key_dict ): raise HTTPException( diff --git a/litellm/proxy/management_endpoints/model_management_endpoints.py b/litellm/proxy/management_endpoints/model_management_endpoints.py index 3c8472f660a..92bbda1ad03 100644 --- a/litellm/proxy/management_endpoints/model_management_endpoints.py +++ b/litellm/proxy/management_endpoints/model_management_endpoints.py @@ -77,9 +77,10 @@ from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache from litellm.proxy.db.routing_prisma_wrapper import WriterPinnedClient from litellm.proxy.management.teams.authz import TEAM_ADMIN_ONLY, is_team_admin from litellm.proxy.management.teams.dependencies import get_team_access -from litellm.proxy.management_endpoints.team_endpoints import ( - _refresh_cached_team, +from litellm.proxy.management_endpoints.team_endpoints import ( # noqa: F401 # legacy module exports + _refresh_cached_team, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export append_team_models, + refresh_cached_team, team_model_add, team_model_delete, ) @@ -1548,7 +1549,7 @@ async def unblock_model( #################################################################################### -async def _add_model_to_db( +async def add_model_to_db( model_params: Deployment, user_api_key_dict: UserAPIKeyAuth, prisma_client: PrismaClient, @@ -1591,7 +1592,10 @@ async def _add_model_to_db( return await table.create(data=_create_data) -async def _add_team_model_to_db( +_add_model_to_db: Final = add_model_to_db + + +async def add_team_model_to_db( model_params: Deployment, user_api_key_dict: UserAPIKeyAuth, prisma_client: PrismaClient, @@ -1625,7 +1629,7 @@ async def _add_team_model_to_db( model_params.model_name = unique_model_name ## CREATE MODEL IN DB ## - model_response: Final = await _add_model_to_db( + model_response: Final = await add_model_to_db( model_params=model_params, user_api_key_dict=user_api_key_dict, prisma_client=prisma_client, @@ -1646,6 +1650,9 @@ async def _add_team_model_to_db( return model_response +_add_team_model_to_db: Final = add_team_model_to_db + + async def _update_team_model_in_db( db_model: Deployment, patch_data: updateDeployment, @@ -1966,7 +1973,7 @@ async def _remove_unbacked_team_models( data={"models": [model for model in existing_team_row.models if model not in names_to_remove]}, include={"object_permission": True}, ) - await _refresh_cached_team( + await refresh_cached_team( team_row=updated_team_row, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, @@ -2600,9 +2607,7 @@ async def add_new_model( reload_outcome: ReconcileOutcome = ReconcileOutcome(still_desired=None, live_after=None) try: _original_litellm_model_name: Final = model_params.model_name - add_model: Final = ( - _add_model_to_db if model_params.model_info.team_id is None else _add_team_model_to_db - ) + add_model: Final = add_model_to_db if model_params.model_info.team_id is None else add_team_model_to_db model_response = await add_model( model_params=priced_model_params, user_api_key_dict=user_api_key_dict, @@ -3202,7 +3207,7 @@ async def get_auto_router_classifier_default_prompt( ) -def _deduplicate_litellm_router_models(models: list[dict]) -> list[dict]: +def deduplicate_litellm_router_models(models: list[dict]) -> list[dict]: """ Deduplicate models based on their model_info.id field. Returns a list of unique models keeping only the first occurrence of each model ID. @@ -3223,6 +3228,9 @@ def _deduplicate_litellm_router_models(models: list[dict]) -> list[dict]: return unique_models +_deduplicate_litellm_router_models: Final = deduplicate_litellm_router_models + + _JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) @@ -3457,7 +3465,7 @@ async def clear_cache() -> ReconcileOutcome: # Reload only DB models. _add_deployment_locked, not add_deployment: this # coroutine already holds MODEL_RECONCILE_LOCK and asyncio.Lock is not # reentrant, so the public wrapper would deadlock against itself. - outcome: Final = await proxy_config._add_deployment_locked( + outcome: Final = await proxy_config.add_deployment_locked( prisma_client=prisma_client, proxy_logging_obj=proxy_logging_obj ) diff --git a/litellm/proxy/management_endpoints/organization_endpoints.py b/litellm/proxy/management_endpoints/organization_endpoints.py index e6705730e6f..b8ffec639c0 100644 --- a/litellm/proxy/management_endpoints/organization_endpoints.py +++ b/litellm/proxy/management_endpoints/organization_endpoints.py @@ -48,9 +48,11 @@ from litellm.proxy.management_endpoints.budget_management_endpoints import ( update_budget, ) from litellm.proxy.management_endpoints.common_daily_activity import get_daily_activity -from litellm.proxy.management_endpoints.common_utils import ( - _set_object_metadata_field, - _user_has_admin_view, +from litellm.proxy.management_endpoints.common_utils import ( # noqa: F401 # legacy module exports + _set_object_metadata_field, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _user_has_admin_view, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + set_object_metadata_field, + user_api_key_has_admin_view, validate_budget_duration, ) from litellm.proxy.management_helpers.object_permission_utils import ( @@ -262,7 +264,7 @@ def _table( return prisma_table -async def _verify_org_access( +async def verify_org_access( organization_id: str, user_api_key_dict: UserAPIKeyAuth, prisma_client: PrismaClient, @@ -272,7 +274,7 @@ async def _verify_org_access( Raises HTTPException(403) if the caller does not have access. """ - if _user_has_admin_view(user_api_key_dict): + if user_api_key_has_admin_view(user_api_key_dict): return if not user_api_key_dict.user_id: @@ -306,6 +308,9 @@ async def _verify_org_access( ) +_verify_org_access: Final = verify_org_access + + _STR_OBJECT_DICT_ADAPTER: Final = TypeAdapter(dict[str, object], config=ConfigDict(hide_input_in_errors=True)) _BUDGET_SETTABLE_FIELDS: Final = frozenset(LiteLLM_BudgetTable.model_fields.keys()) - {"budget_id"} _ORG_COLUMN_FIELDS: Final = frozenset({"organization_alias", "models"}) @@ -537,7 +542,7 @@ async def new_organization( for field in _ORG_METADATA_FIELDS: if getattr(data, field, None) is not None: - _set_object_metadata_field( + set_object_metadata_field( object_data=organization_row, field_name=field, value=getattr(data, field), @@ -624,7 +629,7 @@ async def resolve_organization_daily_activity_scope( prisma_client: PrismaClient, user_api_key_dict: UserAPIKeyAuth, ) -> _OrganizationDailyActivityScope: - is_admin: Final = _user_has_admin_view(user_api_key_dict) + is_admin: Final = user_api_key_has_admin_view(user_api_key_dict) memberships: Final = ( await _table(OrganizationMembershipRepository(prisma_client)).find_many( where={"user_id": user_api_key_dict.user_id} @@ -750,7 +755,7 @@ async def update_organization( # IDOR guard: only proxy admins / org admins of THIS org may update # it. Without this, any authenticated key holder could rewrite # another organization's metadata, budgets, and object permissions. - await _verify_org_access( + await verify_org_access( organization_id=data.organization_id, user_api_key_dict=user_api_key_dict, prisma_client=prisma_client, @@ -927,7 +932,7 @@ async def update_organization_v2( }, ) - await _verify_org_access( + await verify_org_access( organization_id=organization_id, user_api_key_dict=user_api_key_dict, prisma_client=prisma_client, @@ -1144,7 +1149,7 @@ async def list_organization( } # if proxy admin or admin viewer - get all orgs (with optional filters) - if _user_has_admin_view(user_api_key_dict): + if user_api_key_has_admin_view(user_api_key_dict): response = await _table(OrganizationRepository(prisma_client)).find_many( where=where_conditions if where_conditions else None, include={"litellm_budget_table": True, "members": True, "teams": True}, @@ -1210,7 +1215,7 @@ async def info_organization( raise HTTPException(status_code=500, detail={"error": "No db connected"}) # Verify caller has access to this organization - await _verify_org_access( + await verify_org_access( organization_id=organization_id, user_api_key_dict=user_api_key_dict, prisma_client=prisma_client, @@ -1263,7 +1268,7 @@ async def deprecated_info_organization( # Verify caller has access to each requested organization for org_id in data.organizations: - await _verify_org_access( + await verify_org_access( organization_id=org_id, user_api_key_dict=user_api_key_dict, prisma_client=prisma_client, @@ -1339,7 +1344,7 @@ async def organization_member_add( # organization, allowed to access this endpoint" — but the code # never enforced that. Any authenticated key holder could add # members to any org. Now gated explicitly. - await _verify_org_access( + await verify_org_access( organization_id=data.organization_id, user_api_key_dict=user_api_key_dict, prisma_client=prisma_client, @@ -1453,7 +1458,7 @@ async def organization_member_update( # update member roles. The PROXY_ADMIN-target check below was # the only access control; without this, any authenticated user # could change any non-admin member's role in any org. - await _verify_org_access( + await verify_org_access( organization_id=data.organization_id, user_api_key_dict=user_api_key_dict, prisma_client=prisma_client, @@ -1602,7 +1607,7 @@ async def organization_member_delete( # IDOR guard: only proxy admins / org admins of THIS org may # delete members. Without this, any authenticated key holder # could remove any user from any org. - await _verify_org_access( + await verify_org_access( organization_id=data.organization_id, user_api_key_dict=user_api_key_dict, prisma_client=prisma_client, diff --git a/litellm/proxy/management_endpoints/policy_endpoints/__init__.py b/litellm/proxy/management_endpoints/policy_endpoints/__init__.py index 862c92bace9..6ba9872f322 100644 --- a/litellm/proxy/management_endpoints/policy_endpoints/__init__.py +++ b/litellm/proxy/management_endpoints/policy_endpoints/__init__.py @@ -9,12 +9,20 @@ are imported directly into this namespace. from litellm.proxy.management_endpoints.policy_endpoints.endpoints import * # noqa: F403 from litellm.proxy.management_endpoints.policy_endpoints.endpoints import ( # noqa: F401 - _build_all_names_per_competitor, - _build_comparison_blocked_words, - _build_competitor_guardrail_definitions, - _build_name_blocked_words, - _build_recommendation_blocked_words, - _build_refinement_prompt, - _clean_competitor_line, - _parse_variations_response, + _build_all_names_per_competitor, # pyright: ignore[reportPrivateUsage] # backwards-compatible package export + _build_comparison_blocked_words, # pyright: ignore[reportPrivateUsage] # backwards-compatible package export + _build_competitor_guardrail_definitions, # pyright: ignore[reportPrivateUsage] # backwards-compatible package export + _build_name_blocked_words, # pyright: ignore[reportPrivateUsage] # backwards-compatible package export + _build_recommendation_blocked_words, # pyright: ignore[reportPrivateUsage] # backwards-compatible package export + _build_refinement_prompt, # pyright: ignore[reportPrivateUsage] # backwards-compatible package export + _clean_competitor_line, # pyright: ignore[reportPrivateUsage] # backwards-compatible package export + _parse_variations_response, # pyright: ignore[reportPrivateUsage] # backwards-compatible package export + build_all_names_per_competitor, + build_comparison_blocked_words, + build_competitor_guardrail_definitions, + build_name_blocked_words, + build_recommendation_blocked_words, + build_refinement_prompt, + clean_competitor_line, + parse_variations_response, ) diff --git a/litellm/proxy/management_endpoints/policy_endpoints/endpoints.py b/litellm/proxy/management_endpoints/policy_endpoints/endpoints.py index c3d1020f089..3bef7c107eb 100644 --- a/litellm/proxy/management_endpoints/policy_endpoints/endpoints.py +++ b/litellm/proxy/management_endpoints/policy_endpoints/endpoints.py @@ -757,7 +757,7 @@ async def enrich_policy_template( variations_map: Final = await _generate_competitor_variations(competitors, model=model) - enriched_definitions: Final = _build_competitor_guardrail_definitions( + enriched_definitions: Final = build_competitor_guardrail_definitions( template.get("guardrailDefinitions", []), competitors, brand_name, @@ -771,7 +771,7 @@ async def enrich_policy_template( } -def _build_refinement_prompt( +def build_refinement_prompt( instruction: str, existing_competitors: list[str], brand_name: str, @@ -788,6 +788,9 @@ def _build_refinement_prompt( ) +_build_refinement_prompt: Final = build_refinement_prompt + + async def _stream_llm_competitor_names( prompt: str, model: str, @@ -817,13 +820,13 @@ async def _stream_llm_competitor_names( buffer += delta while "\n" in buffer: line, buffer = buffer.split("\n", 1) - name = _clean_competitor_line(line) + name = clean_competitor_line(line) if name and name.lower() not in existing_lower and count < MAX_COMPETITOR_NAMES: existing_lower.add(name.lower()) count += 1 yield name, False # Handle remaining buffer - name = _clean_competitor_line(buffer) + name = clean_competitor_line(buffer) # rebind-ok: pre-existing rebinding on a rename-only line if name and name.lower() not in existing_lower and count < MAX_COMPETITOR_NAMES: yield name, False @@ -843,7 +846,7 @@ async def _stream_competitor_events( for comp in competitors: yield f"data: {json.dumps({'type': 'competitor', 'name': comp})}\n\n" - refinement_prompt: Final = _build_refinement_prompt(data.instruction, competitors, brand_name) + refinement_prompt: Final = build_refinement_prompt(data.instruction, competitors, brand_name) try: async for name, _ in _stream_llm_competitor_names(refinement_prompt, model, competitors): if name: @@ -875,7 +878,7 @@ async def _stream_competitor_events( total_variations: Final = sum(len(v) for v in variations_map.values()) yield f"data: {json.dumps({'type': 'status', 'message': f'Building guardrail definitions with {total_variations} variations...'})}\n\n" - enriched_definitions: Final = _build_competitor_guardrail_definitions( + enriched_definitions: Final = build_competitor_guardrail_definitions( template.get("guardrailDefinitions", []), competitors, brand_name, @@ -916,12 +919,15 @@ async def enrich_policy_template_stream( ) -def _clean_competitor_line(line: str) -> str | None: +def clean_competitor_line(line: str) -> str | None: """Strip numbering, bullets, and whitespace from a competitor name line.""" name: Final = line.strip().strip(".-) ").strip() return name if name and len(name) > 1 else None +_clean_competitor_line: Final = clean_competitor_line + + async def _generate_competitor_variations(competitors: list, model: str = DEFAULT_COMPETITOR_DISCOVERY_MODEL) -> dict: """Generate common misspellings, abbreviations, and alternate names for each competitor.""" if not competitors: @@ -951,13 +957,13 @@ async def _generate_competitor_variations(competitors: list, model: str = DEFAUL temperature=COMPETITOR_LLM_TEMPERATURE, ) raw: Final = response.choices[0].message.content or "" - return _parse_variations_response(raw, capped) + return parse_variations_response(raw, capped) except Exception as e: verbose_proxy_logger.error("LLM competitor variation generation failed: %s", e) return {} -def _parse_variations_response(raw: str, competitors: list) -> dict[str, list[str]]: +def parse_variations_response(raw: str, competitors: list) -> dict[str, list[str]]: """Parse the LLM response for competitor variations into a name -> variations map.""" # Build a lowercase lookup for case-insensitive matching lower_to_canonical: Final = {comp.lower(): comp for comp in competitors} @@ -978,6 +984,9 @@ def _parse_variations_response(raw: str, competitors: list) -> dict[str, list[st return variations_map +_parse_variations_response: Final = parse_variations_response + + async def _discover_competitors_via_llm(prompt: str, model: str = DEFAULT_COMPETITOR_DISCOVERY_MODEL) -> list: """Call an onboarded LLM to discover competitor names.""" try: @@ -991,21 +1000,26 @@ async def _discover_competitors_via_llm(prompt: str, model: str = DEFAULT_COMPET temperature=COMPETITOR_LLM_TEMPERATURE, ) raw: Final = response.choices[0].message.content or "" - competitors = [name for line in raw.strip().split("\n") if (name := _clean_competitor_line(line)) is not None] + competitors: Final = [ + name for line in raw.strip().split("\n") if (name := clean_competitor_line(line)) is not None + ] return competitors[:MAX_COMPETITOR_NAMES] except Exception as e: verbose_proxy_logger.error("LLM competitor discovery failed: %s", e) return [] -def _build_all_names_per_competitor( +def build_all_names_per_competitor( competitors: list[str], variations_map: dict[str, list[str]] ) -> dict[str, list[str]]: """Build canonical + variation name lists for each competitor.""" return {comp: [comp] + variations_map.get(comp, []) for comp in competitors} -def _build_competitor_guardrail_definitions( +_build_all_names_per_competitor: Final = build_all_names_per_competitor + + +def build_competitor_guardrail_definitions( definitions: list, competitors: list, brand_name: str, @@ -1014,11 +1028,11 @@ def _build_competitor_guardrail_definitions( """Build enriched guardrailDefinitions with competitor names and variations populated.""" variations_map = variations_map or {} enriched: Final = copy.deepcopy(definitions) - all_names: Final = _build_all_names_per_competitor(competitors, variations_map) + all_names: Final = build_all_names_per_competitor(competitors, variations_map) - output_blocked: Final = _build_name_blocked_words(competitors, all_names) - recommendation_blocked: Final = _build_recommendation_blocked_words(competitors, all_names) - comparison_blocked: Final = _build_comparison_blocked_words(competitors, all_names, brand_name) + output_blocked: Final = build_name_blocked_words(competitors, all_names) + recommendation_blocked: Final = build_recommendation_blocked_words(competitors, all_names) + comparison_blocked: Final = build_comparison_blocked_words(competitors, all_names, brand_name) blocked_words_map: Final = { "competitor-output-blocker": output_blocked, @@ -1042,7 +1056,10 @@ def _build_competitor_guardrail_definitions( return enriched -def _build_name_blocked_words(competitors: list[str], all_names: dict[str, list[str]]) -> list[dict]: +_build_competitor_guardrail_definitions: Final = build_competitor_guardrail_definitions + + +def build_name_blocked_words(competitors: list[str], all_names: dict[str, list[str]]) -> list[dict]: """Build blocked word entries for direct competitor name mentions.""" result: Final = [] for comp in competitors: @@ -1052,7 +1069,10 @@ def _build_name_blocked_words(competitors: list[str], all_names: dict[str, list[ return result -def _build_recommendation_blocked_words(competitors: list[str], all_names: dict[str, list[str]]) -> list[dict]: +_build_name_blocked_words: Final = build_name_blocked_words + + +def build_recommendation_blocked_words(competitors: list[str], all_names: dict[str, list[str]]) -> list[dict]: """Build blocked word entries for competitor recommendations.""" result: Final = [] for comp in competitors: @@ -1068,7 +1088,10 @@ def _build_recommendation_blocked_words(competitors: list[str], all_names: dict[ return result -def _build_comparison_blocked_words( +_build_recommendation_blocked_words: Final = build_recommendation_blocked_words + + +def build_comparison_blocked_words( competitors: list[str], all_names: dict[str, list[str]], brand_name: str ) -> list[dict]: """Build blocked word entries for unfavorable competitor comparisons.""" @@ -1102,6 +1125,9 @@ def _build_comparison_blocked_words( return result +_build_comparison_blocked_words: Final = build_comparison_blocked_words + + class SuggestTemplatesRequest(LiteLLMBaseModel): attack_examples: list[str] = Field(default_factory=list) description: str = Field(default="") diff --git a/litellm/proxy/management_endpoints/prompt_cache_prediction.py b/litellm/proxy/management_endpoints/prompt_cache_prediction.py index 9103fa09893..a469f1d4f8d 100644 --- a/litellm/proxy/management_endpoints/prompt_cache_prediction.py +++ b/litellm/proxy/management_endpoints/prompt_cache_prediction.py @@ -15,12 +15,14 @@ from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.auth.auth_checks import can_key_call_resolved_model from litellm.proxy.auth.auth_utils import get_cache_prediction_deployments from litellm.proxy.auth.user_api_key_auth import user_api_key_auth -from litellm.proxy.common_utils.http_parsing_utils import ( - _read_request_body, # pyright: ignore[reportPrivateUsage, reportUnknownVariableType] # canonical parsed-body owner; validate its legacy result at the endpoint boundary +from litellm.proxy.common_utils.http_parsing_utils import ( # noqa: F401 # legacy module exports + _read_request_body, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + read_request_body, # pyright: ignore[reportPrivateUsage, reportUnknownVariableType] # canonical parsed-body owner; validate its legacy result at the endpoint boundary ) from litellm.proxy.common_utils.prompt_cache_prediction import has_request_transforms, predict_arm -from litellm.proxy.hooks.parallel_request_limiter_v3 import ( - _PROXY_MaxParallelRequestsHandler_v3, # pyright: ignore[reportPrivateUsage] # use the configured proxy limiter's shared capacity owner +from litellm.proxy.hooks.parallel_request_limiter_v3 import ( # noqa: F401 # legacy module exports + PROXY_MaxParallelRequestsHandler_v3, + _PROXY_MaxParallelRequestsHandler_v3, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export ) from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup from litellm.types.llms.base import LiteLLMBaseModel @@ -39,7 +41,7 @@ class _CallerSettings(LiteLLMBaseModel): def _capacity_counter( - limiter: _PROXY_MaxParallelRequestsHandler_v3, + limiter: PROXY_MaxParallelRequestsHandler_v3, caller: UserAPIKeyAuth, model_name: str, request_data: Mapping[str, object], @@ -114,7 +116,7 @@ async def predict_cache_cost( or not caller or unsupported_transform or unsupported_headers - or not isinstance(limiter, _PROXY_MaxParallelRequestsHandler_v3) + or not isinstance(limiter, PROXY_MaxParallelRequestsHandler_v3) ): reason: Final = ( "unsupported_provider_headers" @@ -134,7 +136,7 @@ async def predict_cache_cost( cache_rebuild_penalty=None, ) request_data: Final = _capacity_request_data( - http_request, user_api_key_dict, _REQUEST_DATA.validate_python(await _read_request_body(http_request)) + http_request, user_api_key_dict, _REQUEST_DATA.validate_python(await read_request_body(http_request)) ) stay: Final = await predict_arm( current, diff --git a/litellm/proxy/management_endpoints/scim/scim_v2.py b/litellm/proxy/management_endpoints/scim/scim_v2.py index 12e1a0841f4..04886072856 100644 --- a/litellm/proxy/management_endpoints/scim/scim_v2.py +++ b/litellm/proxy/management_endpoints/scim/scim_v2.py @@ -45,10 +45,16 @@ from litellm.proxy._types import ( TeamMemberDeleteRequest, UserAPIKeyAuth, ) -from litellm.proxy.auth.auth_checks import _delete_cache_key_object +from litellm.proxy.auth.auth_checks import ( # noqa: F401 # legacy module exports + _delete_cache_key_object, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + delete_cache_key_object, +) from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import evict_and_broadcast -from litellm.proxy.common_utils.http_parsing_utils import _safe_get_request_headers +from litellm.proxy.common_utils.http_parsing_utils import ( # noqa: F401 # legacy module exports + _safe_get_request_headers, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + safe_get_request_headers, +) from litellm.proxy.management_endpoints.internal_user_endpoints import new_user from litellm.proxy.management_endpoints.scim.scim_transformations import ( ScimTransformations, @@ -58,10 +64,11 @@ from litellm.proxy.management_endpoints.team_endpoints import ( team_member_add, team_member_delete, ) -from litellm.proxy.utils import ( +from litellm.proxy.utils import ( # noqa: F401 # legacy module exports PrismaClient, - _premium_user_check, + _premium_user_check, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export handle_exception_on_proxy, + premium_user_check, ) from litellm.repositories.table_repositories import ( InvitationLinkRepository, @@ -264,7 +271,7 @@ class GroupMemberExtractionResult(LiteLLMBaseModel): scim_router: Final = APIRouter( prefix="/scim/v2", tags=["✨ SCIM v2 (Enterprise Only)"], - dependencies=[Depends(_premium_user_check)], + dependencies=[Depends(premium_user_check)], ) SCIM_MAX_PAGE_SIZE: Final = 100 @@ -1045,7 +1052,7 @@ async def _set_user_keys_blocked(user_id: str, blocked: bool) -> int: ) for key_row in affected_keys: - await _delete_cache_key_object( + await delete_cache_key_object( hashed_token=key_row.token, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, @@ -1546,7 +1553,7 @@ async def get_service_provider_config(request: Request): "SCIM ServiceProviderConfig request: method=%s url=%s headers=%s", request.method, request.url, - _safe_get_request_headers(request), + safe_get_request_headers(request), ) meta: Final = { "resourceType": "ServiceProviderConfig", diff --git a/litellm/proxy/management_endpoints/session_endpoints.py b/litellm/proxy/management_endpoints/session_endpoints.py index 2ba84bf03e5..4a4b81f7844 100644 --- a/litellm/proxy/management_endpoints/session_endpoints.py +++ b/litellm/proxy/management_endpoints/session_endpoints.py @@ -30,8 +30,9 @@ from litellm.proxy._types import ( ) from litellm.proxy.auth.auth_checks import delete_cache_key_objects from litellm.proxy.auth.user_api_key_auth import user_api_key_auth -from litellm.proxy.management_endpoints.key_management_endpoints import ( - _persist_deleted_verification_tokens, +from litellm.proxy.management_endpoints.key_management_endpoints import ( # noqa: F401 # legacy module exports + _persist_deleted_verification_tokens, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + persist_deleted_verification_tokens, ) from litellm.repositories.verification_token_repository import ( VerificationTokenRepository, @@ -86,7 +87,7 @@ async def revoke_ui_session_keys( return 0 revoked_tokens: Final = _TOKEN_LIST.validate_python(tuple(row.token for row in revoked_rows)) - await _persist_deleted_verification_tokens( + await persist_deleted_verification_tokens( keys=revoked_rows, prisma_client=prisma_client, user_api_key_dict=user_api_key_dict, @@ -154,7 +155,7 @@ async def session_logout( caller_row: Final = cast( # cast-ok: find_unique returns a prisma row shaped like the pydantic model "LiteLLM_VerificationToken", row ) - await _persist_deleted_verification_tokens( + await persist_deleted_verification_tokens( keys=(caller_row,), prisma_client=prisma_client, user_api_key_dict=user_api_key_dict, diff --git a/litellm/proxy/management_endpoints/team_callback_endpoints.py b/litellm/proxy/management_endpoints/team_callback_endpoints.py index 471a0814ae5..91fd447fb8d 100644 --- a/litellm/proxy/management_endpoints/team_callback_endpoints.py +++ b/litellm/proxy/management_endpoints/team_callback_endpoints.py @@ -34,19 +34,24 @@ from litellm.proxy.common_utils.callback_config_validation import ( conflicting_span_scope_error, cross_entry_family_error, ) -from litellm.proxy.common_utils.callback_utils import ( - _CALLBACK_VAR_ENCRYPTED_PREFIX, +from litellm.proxy.common_utils.callback_utils import ( # noqa: F401 # legacy module exports + _CALLBACK_VAR_ENCRYPTED_PREFIX, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + CALLBACK_VAR_ENCRYPTED_PREFIX, decrypt_callback_vars, encrypt_callback_vars, is_sensitive_callback_key, ) -from litellm.proxy.litellm_pre_call_utils import ( - _get_validated_callback_metadata, +from litellm.proxy.litellm_pre_call_utils import ( # noqa: F401 # legacy module exports + _get_validated_callback_metadata, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export convert_key_logging_metadata_to_callback, + get_validated_callback_metadata, ) from litellm.proxy.management.teams.authz import TEAM_OR_ORG_ADMIN, team_access_denied from litellm.proxy.management.teams.dependencies import get_team_access -from litellm.proxy.management_endpoints.team_endpoints import _refresh_cached_team +from litellm.proxy.management_endpoints.team_endpoints import ( # noqa: F401 # legacy module exports + _refresh_cached_team, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + refresh_cached_team, +) from litellm.proxy.management_helpers.utils import management_endpoint_wrapper from litellm.repositories.team_repository import TeamRepository @@ -114,7 +119,7 @@ def _mask_sensitive_callback_vars(callbacks: TeamCallbackMetadata) -> None: return for key in tuple(callbacks.callback_vars): value = callbacks.callback_vars[key] - if is_sensitive_callback_key(key) or str(value).startswith(_CALLBACK_VAR_ENCRYPTED_PREFIX): + if is_sensitive_callback_key(key) or str(value).startswith(CALLBACK_VAR_ENCRYPTED_PREFIX): callbacks.callback_vars[key] = _CALLBACK_VARS_REDACTED @@ -147,7 +152,7 @@ def _resolve_team_callbacks(team_metadata: object) -> TeamCallbackMetadata: for entry in logging_entries if isinstance(logging_entries, list) else (): if not isinstance(entry, dict): continue - callback = _get_validated_callback_metadata(item=entry, source="team-level read") + callback = get_validated_callback_metadata(item=entry, source="team-level read") if callback is None: continue resolved = convert_key_logging_metadata_to_callback(data=callback, team_callback_settings_obj=resolved) @@ -399,7 +404,7 @@ async def add_team_callbacks( raise _callback_error(400, f"Team id = {team_id} does not exist. Please use a different team id.") # Without this a newly registered callback stays dormant for existing keys. - await _refresh_cached_team( + await refresh_cached_team( team_row=new_team_row, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, @@ -529,7 +534,7 @@ async def delete_team_callback( # Request-time callback resolution reads the cached team, so without this # the removed callback keeps firing for live keys until the cache expires. - await _refresh_cached_team( + await refresh_cached_team( team_row=updated_team, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, @@ -671,7 +676,7 @@ async def disable_team_logging( # Request-time callback resolution reads the cached team, so without this # the DB says logging is off while live keys keep sending until it expires. - await _refresh_cached_team( + await refresh_cached_team( team_row=updated_team, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index a707accefc0..61e17f134ce 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -94,10 +94,11 @@ from litellm.proxy._types import ( UpdateTeamRequest, UserAPIKeyAuth, ) -from litellm.proxy.auth.auth_checks import ( +from litellm.proxy.auth.auth_checks import ( # noqa: F401 # legacy module exports OrganizationNotFoundError, - _cache_team_object, + _cache_team_object, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export allowed_route_check_inside_route, + cache_team_object, can_org_access_model, delete_cache_key_objects, delete_cache_team_object, @@ -129,15 +130,22 @@ from litellm.proxy.management_endpoints.common_daily_activity import ( InvalidDateRange, parse_canonical_date_range, ) -from litellm.proxy.management_endpoints.common_utils import ( - _check_disable_global_guardrails_caller_permission, - _check_passthrough_routes_caller_permission, - _set_object_metadata_field, - _team_member_has_permission, - _update_metadata_fields, - _upsert_budget_and_membership, - _user_has_admin_view, +from litellm.proxy.management_endpoints.common_utils import ( # noqa: F401 # legacy module exports + _check_disable_global_guardrails_caller_permission, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _check_passthrough_routes_caller_permission, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _set_object_metadata_field, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _team_member_has_permission, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _update_metadata_fields, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _upsert_budget_and_membership, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _user_has_admin_view, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + check_disable_global_guardrails_caller_permission, + check_passthrough_routes_caller_permission, member_budget_patch, + set_object_metadata_field, + team_member_has_permission, + update_metadata_fields, + upsert_budget_and_membership, + user_api_key_has_admin_view, validate_budget_duration, validate_team_model_max_budget, ) @@ -161,11 +169,12 @@ from litellm.proxy.management_helpers.access_group_team_sync import ( reconcile_team_access_group_membership, sync_team_access_group_membership, ) -from litellm.proxy.management_helpers.object_permission_utils import ( - _set_object_permission, +from litellm.proxy.management_helpers.object_permission_utils import ( # noqa: F401 # legacy module exports + _set_object_permission, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export enforce_all_proxy_mcp_servers_grant_is_admin_only, handle_update_object_permission_common, invalidate_cached_object_permissions, + set_object_permission, ) from litellm.proxy.management_helpers.team_member_permission_checks import ( TeamMemberPermissionChecks, @@ -454,7 +463,7 @@ def _sanitize_for_log(value: object) -> str: return text.replace("\r", "").replace("\n", "") -async def _refresh_cached_team( +async def refresh_cached_team( team_row: _CacheableTeamRow, user_api_key_cache: UserApiKeyCache, proxy_logging_obj: ProxyLogging, @@ -473,7 +482,7 @@ async def _refresh_cached_team( via `model_dump()` to match the cache shape `_cache_team_object` expects. """ - await _cache_team_object( + await cache_team_object( team_id=team_row.team_id, team_table=LiteLLM_TeamTableCachedObj.model_validate(team_row.model_dump()), user_api_key_cache=user_api_key_cache, @@ -481,6 +490,9 @@ async def _refresh_cached_team( ) +_refresh_cached_team: Final = refresh_cached_team + + _GENERAL_SETTINGS: Final = TypeAdapter(dict[str, object]) @@ -596,7 +608,7 @@ class TeamMemberBudgetHandler: new_team_data_json["metadata"]["team_member_budget_id"] = team_member_budget_table.budget_id # Remove team member fields from new_team_data_json - TeamMemberBudgetHandler._clean_team_member_fields(new_team_data_json) + TeamMemberBudgetHandler.clean_team_member_fields(new_team_data_json) return new_team_data_json @@ -668,17 +680,19 @@ class TeamMemberBudgetHandler: ) # Remove team member fields from updated_kv - TeamMemberBudgetHandler._clean_team_member_fields(updated_kv) + TeamMemberBudgetHandler.clean_team_member_fields(updated_kv) return updated_kv @staticmethod - def _clean_team_member_fields(data_dict: dict) -> None: + def clean_team_member_fields(data_dict: dict) -> None: """Remove team member fields from data dictionary""" data_dict.pop("team_member_budget", None) data_dict.pop("team_member_budget_duration", None) data_dict.pop("team_member_rpm_limit", None) data_dict.pop("team_member_tpm_limit", None) + _clean_team_member_fields = clean_team_member_fields + @staticmethod async def clear_team_member_budget_fields( team_table: _TeamBudgetRow, @@ -712,7 +726,7 @@ class TeamMemberBudgetHandler: user_api_key_dict=user_api_key_dict, ) - TeamMemberBudgetHandler._clean_team_member_fields(updated_kv) + TeamMemberBudgetHandler.clean_team_member_fields(updated_kv) return updated_kv @staticmethod @@ -1603,8 +1617,8 @@ async def new_team( 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") - _check_disable_global_guardrails_caller_permission( + check_passthrough_routes_caller_permission(data, user_api_key_dict, entity="team") + check_disable_global_guardrails_caller_permission( data.disable_global_guardrails, data.metadata, # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # request models declare `metadata` as bare dict user_api_key_dict, @@ -1651,7 +1665,7 @@ async def new_team( is_proxy_admin=user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN, prisma_client=prisma_client, ) - data_json = await _set_object_permission( + data_json = await set_object_permission( # rebind-ok: pre-existing rebinding on a rename-only line data_json=data_json, prisma_client=prisma_client, ) @@ -1682,7 +1696,7 @@ async def new_team( # Set Management Endpoint Metadata Fields for field in LiteLLM_ManagementEndpoint_MetadataFields_Premium: if getattr(data, field, None) is not None: - _set_object_metadata_field( + set_object_metadata_field( object_data=complete_team_data, field_name=field, value=getattr(data, field), @@ -1690,7 +1704,7 @@ async def new_team( for field in LiteLLM_ManagementEndpoint_MetadataFields: if getattr(data, field, None) is not None: - _set_object_metadata_field( + set_object_metadata_field( object_data=complete_team_data, field_name=field, value=getattr(data, field), @@ -2285,13 +2299,13 @@ async def update_team( entity="team", ) - _check_passthrough_routes_caller_permission( + check_passthrough_routes_caller_permission( data, user_api_key_dict, entity="team", existing_metadata=_existing_team_metadata if isinstance(_existing_team_metadata, dict) else None, ) - _check_disable_global_guardrails_caller_permission( + check_disable_global_guardrails_caller_permission( data.disable_global_guardrails, data.metadata, # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # request models declare `metadata` as bare dict user_api_key_dict, @@ -2505,7 +2519,7 @@ async def update_team( explicitly_set_fields=_team_member_fields_in_request, ) else: - TeamMemberBudgetHandler._clean_team_member_fields(updated_kv) + TeamMemberBudgetHandler.clean_team_member_fields(updated_kv) # Check object permission if data.object_permission is not None: @@ -2521,7 +2535,7 @@ async def update_team( ) # update team metadata fields - _update_metadata_fields(updated_kv=updated_kv) + update_metadata_fields(updated_kv=updated_kv) if updated_kv.get("metadata") is not None: updated_kv["metadata"] = encrypt_callback_vars(updated_kv["metadata"]) @@ -2558,7 +2572,7 @@ async def update_team( object_permission_ids=(existing_team.object_permission_id, team_row.object_permission_id), user_api_key_cache=user_api_key_cache, ) - await _refresh_cached_team( + await refresh_cached_team( team_row=team_row, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, @@ -3525,7 +3539,7 @@ def _is_member_addressed_by(member: Member, data: TeamMemberDeleteRequest) -> bo ) -def _cleanup_members_with_roles( +def cleanup_members_with_roles( existing_team_row: LiteLLM_TeamTable, data: TeamMemberDeleteRequest, ) -> tuple[tuple[Member, ...], list[Member]]: @@ -3542,6 +3556,9 @@ def _cleanup_members_with_roles( return removed_team_members, new_team_members +_cleanup_members_with_roles: Final = cleanup_members_with_roles + + @router.post( "/team/member_delete", tags=["team management"], @@ -3651,7 +3668,7 @@ async def _team_member_delete( detail={"error": f"Team id={data.team_id} does not exist in db"}, ) - removed_team_members, new_team_members = _cleanup_members_with_roles( + removed_team_members, new_team_members = cleanup_members_with_roles( existing_team_row=LiteLLM_TeamTable(team_id=data.team_id, members_with_roles=fresh_members), data=data, ) @@ -3717,10 +3734,10 @@ async def _team_member_delete( if user_ids_to_delete: if keys_to_delete: from litellm.proxy.management_endpoints.key_management_endpoints import ( - _persist_deleted_verification_tokens, + persist_deleted_verification_tokens, ) - await _persist_deleted_verification_tokens( + await persist_deleted_verification_tokens( keys=keys_to_delete, prisma_client=prisma_client, user_api_key_dict=user_api_key_dict, @@ -3887,7 +3904,7 @@ async def team_member_update( if data.role is not None else None ) - await _upsert_budget_and_membership( + await upsert_budget_and_membership( tx=tx, team_id=data.team_id, user_id=received_user_id, @@ -4431,7 +4448,7 @@ async def delete_team( ## DELETE ASSOCIATED KEYS # Fetch keys before deletion to persist them from litellm.proxy.management_endpoints.key_management_endpoints import ( - _persist_deleted_verification_tokens, + persist_deleted_verification_tokens, ) keys_to_delete: Final = await _tokens_db(prisma_client).find_many(where={"team_id": {"in": data.team_ids}}) @@ -4441,7 +4458,7 @@ async def delete_team( ) if keys_to_delete: - await _persist_deleted_verification_tokens( + await persist_deleted_verification_tokens( keys=keys_to_delete, prisma_client=prisma_client, user_api_key_dict=user_api_key_dict, @@ -5545,7 +5562,7 @@ async def _enforce_list_team_v2_access( Returns the (possibly overridden) user_id, org_admin_org_ids and, for an org admin's own query, the caller's own team ids. """ - is_proxy_admin: Final = _user_has_admin_view(user_api_key_dict) + is_proxy_admin: Final = user_api_key_has_admin_view(user_api_key_dict) caller_user_id: Final = user_api_key_dict.user_id if is_proxy_admin: @@ -5809,7 +5826,7 @@ async def _authorize_and_filter_teams( - Own query (user_id matches caller): teams the user is a member of, across all orgs. - Others: 401. """ - is_proxy_admin: Final = _user_has_admin_view(user_api_key_dict) + is_proxy_admin: Final = user_api_key_has_admin_view(user_api_key_dict) is_own_query: Final = ( user_id is not None and user_api_key_dict.user_id is not None and user_api_key_dict.user_id == user_id ) @@ -6169,7 +6186,7 @@ async def append_team_models( detail={"error": f"Team not found, passed team_id={data.team_id}"}, ) - await _refresh_cached_team( + await refresh_cached_team( team_row=updated_team, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, @@ -6252,7 +6269,7 @@ async def team_model_delete( detail={"error": f"Team not found, passed team_id={data.team_id}"}, ) - await _refresh_cached_team( + await refresh_cached_team( team_row=updated_team, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, @@ -6300,7 +6317,7 @@ async def team_member_permissions( # a Proxy Admin would. Team / org admins keep their existing scope. if ( hasattr(user_api_key_dict, "user_role") - and not _user_has_admin_view(user_api_key_dict) + and not user_api_key_has_admin_view(user_api_key_dict) and not await get_team_access().allows(user_api_key_dict, complete_team_data, TEAM_OR_ORG_ADMIN) and not _is_available_team( team_id=complete_team_data.team_id, @@ -6556,7 +6573,7 @@ async def resolve_team_daily_activity_scope( if exclude_team_ids: exclude_team_ids_list = exclude_team_ids.split(",") if exclude_team_ids else None - if not _user_has_admin_view(user_api_key_dict): + if not user_api_key_has_admin_view(user_api_key_dict): user_info: Final = await get_user_object( user_id=user_api_key_dict.user_id, prisma_client=prisma_client, @@ -6598,12 +6615,12 @@ async def resolve_team_daily_activity_scope( # filtering the entire response by their own API keys (they can re- # request the admin-only teams separately to get the wider view). user_api_keys: list[str] | None = None - if not _user_has_admin_view(user_api_key_dict) and team_ids_list and team_aliases: + if not user_api_key_has_admin_view(user_api_key_dict) and team_ids_list and team_aliases: has_full_team_view = True for team_alias in team_aliases: team_obj = LiteLLM_TeamTable.model_validate(team_alias.model_dump()) is_admin = is_team_admin(user_api_key_dict=user_api_key_dict, team_obj=team_obj) - has_perm = _team_member_has_permission( + has_perm = team_member_has_permission( user_api_key_dict=user_api_key_dict, team_obj=team_obj, permission="/team/daily/activity", diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py index 605952b44f0..c515108ab11 100644 --- a/litellm/proxy/management_endpoints/ui_sso.py +++ b/litellm/proxy/management_endpoints/ui_sso.py @@ -90,8 +90,9 @@ from litellm.proxy._types import ( UserAPIKeyAuth, ) from litellm.proxy.auth.auth_checks import ExperimentalUIJWTToken, get_user_object -from litellm.proxy.auth.auth_utils import ( - _get_request_ip_address, +from litellm.proxy.auth.auth_utils import ( # noqa: F401 # legacy module exports + _get_request_ip_address, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + get_request_ip_address, has_user_setup_sso, ) from litellm.proxy.auth.handle_jwt import JWTHandler @@ -318,7 +319,7 @@ def _cli_sso_start_response_body( def _get_cli_sso_start_rate_limit_cache_key(request: Request, use_x_forwarded_for: bool | None = False) -> str: - client_ip: Final = _get_request_ip_address(request=request, use_x_forwarded_for=use_x_forwarded_for) or "unknown" + client_ip: Final = get_request_ip_address(request=request, use_x_forwarded_for=use_x_forwarded_for) or "unknown" client_ip_hash: Final = _hash_cli_sso_secret(client_ip) return f"{_CLI_SSO_START_RATE_LIMIT_CACHE_KEY_PREFIX}:{client_ip_hash}" @@ -1060,7 +1061,7 @@ async def google_login( _get_cli_sso_flow_or_raise(login_id=key, cache=cli_sso_session_cache) # Store CLI login handle in state for OAuth flow - cli_state: Final[str | None] = SSOAuthenticationHandler._get_cli_state( + cli_state: Final[str | None] = SSOAuthenticationHandler.get_cli_state( source=source, key=key, user_code=(user_code if _cli_sso_verification_uri_complete_enabled() else None), @@ -1626,7 +1627,7 @@ async def get_generic_sso_response( param="code", code=status.HTTP_400_BAD_REQUEST, ) - combined_response: Final = await SSOAuthenticationHandler._pkce_token_exchange( + combined_response: Final = await SSOAuthenticationHandler.pkce_token_exchange( authorization_code=authorization_code, code_verifier=code_verifier, client_id=generic_client_id, @@ -1664,7 +1665,7 @@ async def get_generic_sso_response( # successfully. Deleting earlier would consume the verifier on a transient # failure, forcing the user to restart the entire OAuth flow from scratch. if pkce_cache_key: - await SSOAuthenticationHandler._delete_pkce_verifier(pkce_cache_key) + await SSOAuthenticationHandler.delete_pkce_verifier(pkce_cache_key) except Exception as e: _handle_generic_sso_error( @@ -2214,7 +2215,7 @@ async def saml_callback(request: Request): relay_state: Final = post_data.get("RelayState") cp_return_to: Final[str | None] = ( relay_state - if isinstance(relay_state, str) and SSOAuthenticationHandler._validate_return_to(relay_state) + if isinstance(relay_state, str) and SSOAuthenticationHandler.validate_return_to(relay_state) else None ) @@ -2433,7 +2434,7 @@ async def cli_sso_callback( result_non_none: Final[OpenID | dict] = cast(OpenID | dict, result) try: - parsed_openid_result: Final = SSOAuthenticationHandler._get_user_email_and_id_from_result( + parsed_openid_result: Final = SSOAuthenticationHandler.get_user_email_and_id_from_result( result=result_non_none, generic_client_id=os.getenv("GENERIC_CLIENT_ID", None), ) @@ -2808,7 +2809,7 @@ def _is_same_origin_return_path(return_to: str) -> bool: @with_service_target(SSO_SESSIONS_TARGET) -async def _sso_return_to_redirect( +async def sso_return_to_redirect( return_to: str | None, jwt_token: str, redis_usage_cache, @@ -2836,7 +2837,7 @@ async def _sso_return_to_redirect( redirect_response.delete_cookie("litellm_cp_return_to") return redirect_response - if SSOAuthenticationHandler._validate_return_to(return_to): + if SSOAuthenticationHandler.validate_return_to(return_to): code: Final = secrets.token_urlsafe(32) cache_key: Final = f"login_code:{code}" cache_value: Final = {"token": jwt_token, "redirect_url": return_to} @@ -2855,6 +2856,9 @@ async def _sso_return_to_redirect( return None +_sso_return_to_redirect: Final = sso_return_to_redirect + + def set_session_token_cookie(response: Response, request: Request, jwt_token: str) -> None: """Set the ``token`` session cookie shared by every sign-in path. @@ -2884,7 +2888,7 @@ def _persist_return_to_cookie(response: Response, return_to: str | None, request if return_to is None: return try: - safe: Final = _is_same_origin_return_path(return_to) or SSOAuthenticationHandler._validate_return_to(return_to) + safe: Final = _is_same_origin_return_path(return_to) or SSOAuthenticationHandler.validate_return_to(return_to) except HTTPException: return # a non-matching absolute return_to is ignored, never blocks sign-in if safe: @@ -2904,7 +2908,7 @@ class SSOAuthenticationHandler: """ @staticmethod - def _validate_return_to(return_to: str) -> bool: + def validate_return_to(return_to: str) -> bool: """ Validate that return_to matches the configured control_plane_url origin. @@ -2934,6 +2938,8 @@ class SSOAuthenticationHandler: return True + _validate_return_to = validate_return_to + @staticmethod async def get_sso_login_redirect( redirect_url: str, @@ -3451,7 +3457,7 @@ class SSOAuthenticationHandler: return team_request @staticmethod - def _get_cli_state( + def get_cli_state( source: str | None, key: str | None, existing_key: str | None = None, @@ -3477,8 +3483,10 @@ class SSOAuthenticationHandler: else: return None + _get_cli_state = get_cli_state + @staticmethod - def _get_user_email_and_id_from_result( + def get_user_email_and_id_from_result( result: OpenID | dict | None, generic_client_id: str | None = None, ) -> ParsedOpenIDResult: @@ -3535,6 +3543,8 @@ class SSOAuthenticationHandler: user_role=user_role, ) + _get_user_email_and_id_from_result = get_user_email_and_id_from_result + @staticmethod async def get_redirect_response_from_openid( result: OpenID | dict | CustomOpenID, @@ -3564,7 +3574,7 @@ class SSOAuthenticationHandler: prisma_client: Final = get_prisma_client_or_throw("Prisma client is None, connect a database to your proxy") # User is Authe'd in - generate key for the UI to access Proxy - parsed_openid_result: Final = SSOAuthenticationHandler._get_user_email_and_id_from_result( + parsed_openid_result: Final = SSOAuthenticationHandler.get_user_email_and_id_from_result( result=result, generic_client_id=generic_client_id ) user_email: Final = parsed_openid_result.get("user_email") @@ -3729,7 +3739,7 @@ class SSOAuthenticationHandler: # Post-SSO return_to handling (the same-origin DCR round-trip and the control-plane # cross-origin code exchange) lives in one shared helper so this method stays inside the # complexity budget. None falls through to the dashboard redirect below. - return_to_redirect: Final = await _sso_return_to_redirect( + return_to_redirect: Final = await sso_return_to_redirect( return_to=return_to, jwt_token=jwt_token, redis_usage_cache=redis_usage_cache, @@ -3862,7 +3872,7 @@ class SSOAuthenticationHandler: strict_cache_miss: Final = os.getenv("PKCE_STRICT_CACHE_MISS", "false").lower() == "true" if strict_cache_miss: if empty_value_in_dict: - await SSOAuthenticationHandler._delete_pkce_verifier(cache_key) + await SSOAuthenticationHandler.delete_pkce_verifier(cache_key) raise ProxyException( message=( f"PKCE verifier for state '{state}' was found in cache but " @@ -3873,7 +3883,7 @@ class SSOAuthenticationHandler: code=status.HTTP_401_UNAUTHORIZED, ) elif cached_data is not None: - await SSOAuthenticationHandler._delete_pkce_verifier(cache_key) + await SSOAuthenticationHandler.delete_pkce_verifier(cache_key) verbose_proxy_logger.error( "PKCE verifier for state '%s' has an unrecognized format (type=%s); " "treating as a cache miss. Investigate the cached value — it may be " @@ -3916,7 +3926,7 @@ class SSOAuthenticationHandler: ) else: if cached_data is not None: - await SSOAuthenticationHandler._delete_pkce_verifier(cache_key) + await SSOAuthenticationHandler.delete_pkce_verifier(cache_key) verbose_proxy_logger.warning( "PKCE is enabled but verifier not found in cache for state '%s' " "(cache type: %s, raw data present: %s). " @@ -3928,7 +3938,7 @@ class SSOAuthenticationHandler: @staticmethod @with_service_target(SSO_SESSIONS_TARGET) - async def _delete_pkce_verifier(cache_key: str) -> None: + async def delete_pkce_verifier(cache_key: str) -> None: """Delete a single-use PKCE verifier from cache after a successful exchange. Failure is non-fatal: a leftover verifier is a minor security concern @@ -3948,6 +3958,8 @@ class SSOAuthenticationHandler: exc, ) + _delete_pkce_verifier = delete_pkce_verifier + @staticmethod def generate_pkce_params() -> tuple[str, str]: """ @@ -4032,7 +4044,7 @@ class SSOAuthenticationHandler: return token_response @staticmethod - async def _pkce_token_exchange( + async def pkce_token_exchange( authorization_code: str, code_verifier: str, client_id: str, @@ -4158,6 +4170,8 @@ class SSOAuthenticationHandler: # Case 3: field absent from token_response — leave userinfo value as-is. return merged + _pkce_token_exchange = pkce_token_exchange + @staticmethod async def _get_pkce_userinfo( access_token: str, diff --git a/litellm/proxy/management_endpoints/usage_endpoints/endpoints.py b/litellm/proxy/management_endpoints/usage_endpoints/endpoints.py index f03259cb4e9..cff5bef05a0 100644 --- a/litellm/proxy/management_endpoints/usage_endpoints/endpoints.py +++ b/litellm/proxy/management_endpoints/usage_endpoints/endpoints.py @@ -47,14 +47,14 @@ async def usage_ai_chat( The AI agent has access to tools that query aggregated daily activity data. """ from litellm.proxy.management_endpoints.common_utils import ( - _user_has_admin_view, require_caller_user_id_for_non_admin, + user_api_key_has_admin_view, ) from litellm.proxy.management_endpoints.usage_endpoints.ai_usage_chat import ( stream_usage_ai_chat, ) - is_admin: Final = _user_has_admin_view(user_api_key_dict) + is_admin: Final = user_api_key_has_admin_view(user_api_key_dict) if is_admin: user_id = user_api_key_dict.user_id else: diff --git a/litellm/proxy/management_helpers/access_group_key_sync.py b/litellm/proxy/management_helpers/access_group_key_sync.py index a5289b89cec..024967a34f6 100644 --- a/litellm/proxy/management_helpers/access_group_key_sync.py +++ b/litellm/proxy/management_helpers/access_group_key_sync.py @@ -33,8 +33,9 @@ from litellm.proxy._types import ( RegenerateKeyRequest, UpdateKeyRequest, ) -from litellm.proxy.auth.auth_checks import ( - _delete_cache_access_object, # pyright: ignore[reportPrivateUsage] # the access-group endpoints reach for this same cache primitive +from litellm.proxy.auth.auth_checks import ( # noqa: F401 # legacy module exports + _delete_cache_access_object, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + delete_cache_access_object, # pyright: ignore[reportPrivateUsage] # the access-group endpoints reach for this same cache primitive ) from litellm.proxy.db.routing_prisma_wrapper import writer_wrapper from litellm.repositories.table_repositories import AccessGroupRepository @@ -86,7 +87,7 @@ async def _invalidate_access_group_cache(access_group_id: str) -> None: """ from litellm.proxy.proxy_server import proxy_logging_obj, user_api_key_cache - await _delete_cache_access_object( + await delete_cache_access_object( access_group_id=access_group_id, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, diff --git a/litellm/proxy/management_helpers/access_group_team_sync.py b/litellm/proxy/management_helpers/access_group_team_sync.py index fbca95b9169..1205414b4fb 100644 --- a/litellm/proxy/management_helpers/access_group_team_sync.py +++ b/litellm/proxy/management_helpers/access_group_team_sync.py @@ -20,7 +20,10 @@ from typing import Final, Protocol from pydantic import TypeAdapter -from litellm.proxy.auth.auth_checks import _delete_cache_access_object +from litellm.proxy.auth.auth_checks import ( # noqa: F401 # legacy module exports + _delete_cache_access_object, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + delete_cache_access_object, +) from litellm.proxy.db.db_span import db_span from litellm.types.llms.base import LiteLLMBaseModel @@ -99,7 +102,7 @@ async def invalidate_access_group_cache(access_group_id: str) -> None: """ from litellm.proxy.proxy_server import proxy_logging_obj, user_api_key_cache - await _delete_cache_access_object( + await delete_cache_access_object( access_group_id=access_group_id, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, diff --git a/litellm/proxy/management_helpers/auto_router_permissions.py b/litellm/proxy/management_helpers/auto_router_permissions.py index f2d0aaaec78..16f6b512e0d 100644 --- a/litellm/proxy/management_helpers/auto_router_permissions.py +++ b/litellm/proxy/management_helpers/auto_router_permissions.py @@ -19,12 +19,13 @@ from litellm.proxy._types import ( LitellmUserRoles, UserAPIKeyAuth, ) -from litellm.proxy.auth.auth_checks import ( - _check_team_member_model_access, # pyright: ignore[reportPrivateUsage] # shared membership authorization owner +from litellm.proxy.auth.auth_checks import ( # noqa: F401 # legacy module exports + _check_team_member_model_access, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export can_key_call_model, can_org_access_model, can_project_access_model, can_team_access_model, + check_team_member_model_access, # pyright: ignore[reportPrivateUsage] # shared membership authorization owner ) from litellm.proxy.auth.team_grants import team_model_aliases from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper @@ -222,7 +223,7 @@ async def authorize_member_auto_router_dependencies( llm_router=llm_router, prisma_client=prisma_client, ) - await _check_team_member_model_access( + await check_team_member_model_access( model=model, team_object=team, valid_token=scoped_actor, diff --git a/litellm/proxy/management_helpers/bulk_team_member_budgets.py b/litellm/proxy/management_helpers/bulk_team_member_budgets.py index 292cb637687..08169eab761 100644 --- a/litellm/proxy/management_helpers/bulk_team_member_budgets.py +++ b/litellm/proxy/management_helpers/bulk_team_member_budgets.py @@ -25,9 +25,10 @@ from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache from litellm.proxy.db.routing_prisma_wrapper import WriterPinnedClient from litellm.proxy.management.teams.authz import TEAM_OR_ORG_ADMIN from litellm.proxy.management.teams.dependencies import get_team_access -from litellm.proxy.management_endpoints.common_utils import ( - _upsert_budget_and_membership, # pyright: ignore[reportPrivateUsage] # the single-member write, shared so the two surfaces cannot drift +from litellm.proxy.management_endpoints.common_utils import ( # noqa: F401 # legacy module exports + _upsert_budget_and_membership, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export member_budget_patch, + upsert_budget_and_membership, # pyright: ignore[reportPrivateUsage] # the single-member write, shared so the two surfaces cannot drift ) from litellm.proxy.management_helpers.audit_logs import create_object_audit_log from litellm.proxy.management_helpers.bulk_user_deletion import ( @@ -213,7 +214,7 @@ async def bulk_update_team_member_budgets( tx, frozenset(budget_id for budget_id in budget_id_of.values() if budget_id is not None) ) for index, user_id in applied: - await _upsert_budget_and_membership( + await upsert_budget_and_membership( tx=tx, team_id=team_id, user_id=user_id, diff --git a/litellm/proxy/management_helpers/bulk_user_creation.py b/litellm/proxy/management_helpers/bulk_user_creation.py index ecaec1b6260..cca37ebd7ed 100644 --- a/litellm/proxy/management_helpers/bulk_user_creation.py +++ b/litellm/proxy/management_helpers/bulk_user_creation.py @@ -42,15 +42,17 @@ from litellm.proxy.management_endpoints.internal_user_endpoints import ( _update_internal_new_user_params, # pyright: ignore[reportPrivateUsage, reportUnknownVariableType] # /user/new defaults; result validated below check_if_default_team_set, ) -from litellm.proxy.management_endpoints.key_management_endpoints import ( - _check_permissions_caller_permission, # pyright: ignore[reportPrivateUsage] # same permission check /user/new uses +from litellm.proxy.management_endpoints.key_management_endpoints import ( # noqa: F401 # legacy module exports + _check_permissions_caller_permission, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + check_permissions_caller_permission, # pyright: ignore[reportPrivateUsage] # same permission check /user/new uses generate_key_helper_fn, # pyright: ignore[reportUnknownVariableType] # legacy untyped helper; result validated by _KEY_RESPONSE metadata_json_with_limits, ) from litellm.proxy.management_endpoints.organization_endpoints import organization_member_add from litellm.proxy.management_helpers.access_group_team_sync import TEAM_ADVISORY_LOCK_SQL -from litellm.proxy.management_helpers.object_permission_utils import ( - _set_object_permission, # pyright: ignore[reportPrivateUsage, reportUnknownVariableType] # shared with /user/new; result validated below +from litellm.proxy.management_helpers.object_permission_utils import ( # noqa: F401 # legacy module exports + _set_object_permission, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + set_object_permission, # pyright: ignore[reportPrivateUsage, reportUnknownVariableType] # shared with /user/new; result validated below ) from litellm.proxy.management_helpers.utils import ( _resolve_member_budget_id, # pyright: ignore[reportPrivateUsage] # shared with /team/member_add @@ -214,7 +216,7 @@ def _row_error(item: BulkNewUserItem, user_api_key_dict: UserAPIKeyAuth) -> str ) try: validate_budget_duration(item.budget_duration) - _check_permissions_caller_permission(data=item, user_api_key_dict=user_api_key_dict) + check_permissions_caller_permission(data=item, user_api_key_dict=user_api_key_dict) if item.auto_create_key: enforce_batch_limits_are_admin_only(item, None, user_api_key_dict, "key") except Exception as exc: # noqa: BLE001 # any validation failure is reported on this row only @@ -344,7 +346,7 @@ async def _prepare_user(user: _PendingUser, prisma_client: PrismaClient) -> _Pre data: Final = {**dumped, "user_id": user.user_id} data_json: Final = _JSON_OBJECT.validate_python(_update_internal_new_user_params(data, user.request)) with_permission: Final = _JSON_OBJECT.validate_python( - await _set_object_permission(data_json=data_json, prisma_client=prisma_client) + await set_object_permission(data_json=data_json, prisma_client=prisma_client) ) return _PreparedUser(user, _USER_ROW.validate_python(with_permission)) except Exception as exc: # noqa: BLE001 # any preparation failure is reported on this row only diff --git a/litellm/proxy/management_helpers/bulk_user_deletion.py b/litellm/proxy/management_helpers/bulk_user_deletion.py index b4e4afcbc2a..c92ebce8531 100644 --- a/litellm/proxy/management_helpers/bulk_user_deletion.py +++ b/litellm/proxy/management_helpers/bulk_user_deletion.py @@ -36,8 +36,9 @@ from litellm.proxy.hooks.user_management_event_hooks import UserManagementEventH from litellm.proxy.list_api.common import PROBLEM_TYPE_BASE, ManagementProblem from litellm.proxy.management.teams.authz import TEAM_OR_ORG_ADMIN from litellm.proxy.management.teams.dependencies import get_team_access -from litellm.proxy.management_endpoints.key_management_endpoints import ( - _persist_deleted_verification_tokens, # pyright: ignore[reportPrivateUsage] # same audit path /key/delete uses +from litellm.proxy.management_endpoints.key_management_endpoints import ( # noqa: F401 # legacy module exports + _persist_deleted_verification_tokens, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + persist_deleted_verification_tokens, # pyright: ignore[reportPrivateUsage] # same audit path /key/delete uses ) from litellm.proxy.management_helpers.access_group_team_sync import TEAM_ADVISORY_LOCK_SQL from litellm.proxy.utils import PrismaClient, ProxyLogging @@ -267,7 +268,7 @@ async def _remove_members_from_team( await _user_tx_db(tx).update(where=_eq_filter("user_id", row.user_id), data=teams_data) await _membership_tx_db(tx).delete_many(where=_team_users_filter(team_id, cleanup_ids)) if keys: - await _persist_deleted_verification_tokens( + await persist_deleted_verification_tokens( keys=keys, # pyright: ignore[reportArgumentType] # generated row model carries the same columns as LiteLLM_VerificationToken prisma_client=prisma_client, user_api_key_dict=user_api_key_dict, @@ -398,7 +399,7 @@ async def _delete_user_rows( prisma_client=prisma_client, ) if keys: - await _persist_deleted_verification_tokens( + await persist_deleted_verification_tokens( keys=keys, # pyright: ignore[reportArgumentType] # generated row model carries the same columns as LiteLLM_VerificationToken prisma_client=prisma_client, user_api_key_dict=user_api_key_dict, diff --git a/litellm/proxy/management_helpers/object_permission_utils.py b/litellm/proxy/management_helpers/object_permission_utils.py index c389381dca4..86545e6e064 100644 --- a/litellm/proxy/management_helpers/object_permission_utils.py +++ b/litellm/proxy/management_helpers/object_permission_utils.py @@ -199,10 +199,10 @@ async def invalidate_cached_object_permissions( await evict_and_broadcast(cache_keys, user_api_key_cache) -async def _set_object_permission( - data_json: dict, +async def set_object_permission( + data_json: dict[str, object], prisma_client: PrismaClient | None, -): +) -> dict[str, object]: """ Creates the LiteLLM_ObjectPermissionTable record for the key/team. Handles permissions for vector stores and mcp servers. @@ -237,6 +237,9 @@ async def _set_object_permission( return data_json +_set_object_permission: Final = set_object_permission + + def _dedupe_preserving_order(values: list[str]) -> list[str]: seen: Final[set[str]] = set() result: Final[list[str]] = [] @@ -447,7 +450,7 @@ async def _resolve_team_allowed_mcp_servers( direct_servers: Final[list[str]] = team_object_permission.mcp_servers or [] if SpecialMCPServerName.all_proxy_servers.value in direct_servers: return _get_all_mcp_server_ids() - access_group_servers: Final[list[str]] = await MCPRequestHandler._get_mcp_servers_from_access_groups( + access_group_servers: Final[list[str]] = await MCPRequestHandler.get_mcp_servers_from_access_groups( team_object_permission.mcp_access_groups or [] ) raw_tool_perms = team_object_permission.mcp_tool_permissions or {} @@ -463,7 +466,7 @@ async def _resolve_team_allowed_mcp_servers( return _flatten_resolved_mcp_server_ids(resolved_servers) | unresolved_servers -def _get_allow_all_keys_server_ids() -> set[str]: +def get_allow_all_keys_server_ids() -> set[str]: """Return the set of MCP server IDs marked with allow_all_keys=True.""" from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( global_mcp_server_manager, @@ -472,6 +475,9 @@ def _get_allow_all_keys_server_ids() -> set[str]: return set(global_mcp_server_manager.get_allow_all_keys_server_ids()) +_get_allow_all_keys_server_ids: Final = get_allow_all_keys_server_ids + + def _get_all_mcp_server_ids() -> set[str]: """Return every MCP server id registered on the proxy (config + DB union).""" from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( @@ -558,7 +564,7 @@ async def _get_grandfathered_key_mcp_server_ids( ) -async def _get_team_allowed_mcp_servers( +async def get_team_allowed_mcp_servers( team_obj: Optional["LiteLLM_TeamTableCachedObj"], prisma_client: PrismaClient | None = None, ) -> set[str]: @@ -574,10 +580,10 @@ async def _get_team_allowed_mcp_servers( return set() from litellm.proxy.auth.auth_checks import ( - _get_mcp_server_ids_from_access_groups, # pyright: ignore[reportPrivateUsage] # same resolver runtime MCP auth calls + get_mcp_server_ids_from_access_groups, # pyright: ignore[reportPrivateUsage] # same resolver runtime MCP auth calls ) - access_group_servers: Final = await _get_mcp_server_ids_from_access_groups( + access_group_servers: Final = await get_mcp_server_ids_from_access_groups( access_group_ids=team_obj.access_group_ids or [], prisma_client=prisma_client, ) @@ -599,6 +605,9 @@ async def _get_team_allowed_mcp_servers( ) +_get_team_allowed_mcp_servers: Final = get_team_allowed_mcp_servers + + def _extract_requested_mcp_server_ids( object_permission: ObjectPermissionDict | None, ) -> set[str]: @@ -692,8 +701,8 @@ async def validate_key_mcp_servers_against_team( if not requested_servers and not requested_access_groups and not requested_toolsets: return object_permission - allow_all_keys_servers: Final = _get_allow_all_keys_server_ids() - team_allowed_servers: Final = await _get_team_allowed_mcp_servers( + allow_all_keys_servers: Final = get_allow_all_keys_server_ids() + team_allowed_servers: Final = await get_team_allowed_mcp_servers( team_obj=team_obj, prisma_client=prisma_client, ) diff --git a/litellm/proxy/management_helpers/team_member_permission_checks.py b/litellm/proxy/management_helpers/team_member_permission_checks.py index a076d8240c6..3e6d0e47265 100644 --- a/litellm/proxy/management_helpers/team_member_permission_checks.py +++ b/litellm/proxy/management_helpers/team_member_permission_checks.py @@ -65,7 +65,7 @@ class TeamMemberPermissionChecks: Main handler for checking if a team member can update a key """ from litellm.proxy.management_endpoints.key_management_endpoints import ( - _get_caller_team_role, + get_caller_team_role, ) # 1. Don't execute these checks if the user role is proxy admin @@ -85,7 +85,7 @@ class TeamMemberPermissionChecks: check_db_only=True, ) - caller_team_role: Final = _get_caller_team_role(team_table=team_table, user_api_key_dict=user_api_key_dict) + caller_team_role: Final = get_caller_team_role(team_table=team_table, user_api_key_dict=user_api_key_dict) # 4. Check if the team member has permissions for the endpoint has_permission: Final = TeamMemberPermissionChecks.does_team_member_have_permissions_for_endpoint( @@ -152,7 +152,7 @@ class TeamMemberPermissionChecks: from fastapi import HTTPException from litellm.proxy.management_endpoints.key_management_endpoints import ( - _get_caller_team_role, + get_caller_team_role, ) # No-op when the request does not assign any access groups. @@ -173,7 +173,7 @@ class TeamMemberPermissionChecks: ), ) - caller_team_role: Final = _get_caller_team_role(team_table=team_table, user_api_key_dict=user_api_key_dict) + caller_team_role: Final = get_caller_team_role(team_table=team_table, user_api_key_dict=user_api_key_dict) # Team admins always bypass (consistent with other member-permission checks). if caller_team_role == "admin": @@ -209,7 +209,7 @@ class TeamMemberPermissionChecks: Returns True if the user belongs to the team that the key is assigned to """ from litellm.proxy.management_endpoints.key_management_endpoints import ( - _get_caller_team_role, + get_caller_team_role, ) from litellm.proxy.proxy_server import prisma_client, user_api_key_cache @@ -223,7 +223,7 @@ class TeamMemberPermissionChecks: check_db_only=True, ) - caller_team_role: Final = _get_caller_team_role(team_table=team_table, user_api_key_dict=user_api_key_dict) + caller_team_role: Final = get_caller_team_role(team_table=team_table, user_api_key_dict=user_api_key_dict) return caller_team_role is not None @staticmethod diff --git a/litellm/proxy/management_helpers/utils.py b/litellm/proxy/management_helpers/utils.py index b8af3950859..30aaaacad04 100644 --- a/litellm/proxy/management_helpers/utils.py +++ b/litellm/proxy/management_helpers/utils.py @@ -35,7 +35,10 @@ from litellm.proxy._types import ( # key request types; user request types; tea UserAPIKeyAuth, VirtualKeyEvent, ) -from litellm.proxy.common_utils.http_parsing_utils import _read_request_body +from litellm.proxy.common_utils.http_parsing_utils import ( # noqa: F401 # legacy module exports + _read_request_body, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + read_request_body, +) from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time from litellm.proxy.common_utils.user_api_key_cache import AUTH_OBJECTS_TARGET from litellm.proxy.utils import PrismaClient, jsonify_object @@ -713,7 +716,9 @@ async def _emit_management_endpoint_otel_span( ) route = get_request_route(http_request) - request_body: dict = await _read_request_body(request=http_request) + request_body: dict = await read_request_body( # rebind-ok: pre-existing rebinding on a rename-only line + request=http_request + ) else: route = func.__name__ request_body = {} diff --git a/litellm/proxy/moyai_endpoints.py b/litellm/proxy/moyai_endpoints.py new file mode 100644 index 00000000000..a2335c0a49b --- /dev/null +++ b/litellm/proxy/moyai_endpoints.py @@ -0,0 +1,263 @@ +"""Moyai quick-connect endpoints. + +`/moyai/connect/start` hands a proxy admin a signed, single-use code pointing +at their Moyai deployment. `/moyai/connect/exchange` trades that code for a +fresh virtual key and persists the deployment as the `moyai_url` UI setting. +The signed code is the credential for the exchange, so it must stay short +lived and single use. +""" + +import base64 +import hashlib +import hmac +import json +import os +import secrets +import time +from typing import TYPE_CHECKING, Annotated, Final +from urllib.parse import urlencode, urlparse + +from fastapi import APIRouter, Depends, HTTPException, Request, status +from pydantic import BaseModel + +from litellm._internal_context import with_service_target +from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth +from litellm.proxy.auth.user_api_key_auth import user_api_key_auth +from litellm.proxy.ui_crud_endpoints.proxy_setting_endpoints import ( + UI_SETTINGS_CACHE_KEY, + UI_SETTINGS_CACHE_TTL, + _ui_settings_db, + normalize_moyai_url, +) +from litellm.proxy.utils import CONFIG_PARAMS_TARGET +from litellm.repositories.config_repository import ConfigRepository +from litellm.repositories.table_repositories import UISettingsRepository + +if TYPE_CHECKING: + from litellm.proxy.utils import PrismaClient + +router: Final = APIRouter() + +_MOYAI_CODE_TTL_SECONDS: Final = 600 +_MOYAI_NONCE_CONFIG_PREFIX: Final = "moyai_connect_nonce:" +_MOYAI_CONNECT_EXCHANGE_ROUTE: Final = "/moyai/connect/exchange" + + +class MoyaiConnectStartRequest(BaseModel): + moyai_url: str + return_to: str + + +class MoyaiConnectStartResponse(BaseModel): + connect_url: str + + +class MoyaiConnectExchangeRequest(BaseModel): + code: str + moyai_url: str + + +class MoyaiConnectExchangeResponse(BaseModel): + api_key: str + key_alias: str + api_base: str + + +def _b64url(data: bytes) -> str: + return base64.urlsafe_b64encode(data).decode().rstrip("=") + + +def _b64url_decode(data: str) -> bytes: + return base64.urlsafe_b64decode(data + "=" * (-len(data) % 4)) + + +def _origin(url: str) -> str: + parsed: Final = urlparse(url) + return f"{parsed.scheme}://{parsed.netloc}" + + +def _master_key_hmac_key(master_key: str) -> bytes: + return hashlib.sha256(master_key.encode()).digest() + + +def _gateway_url(request: Request) -> str: + if os.environ.get("PROXY_BASE_URL"): + return os.environ["PROXY_BASE_URL"].rstrip("/") + return str(request.base_url).rstrip("/") + + +_MOYAI_KEY_ALLOWED_ROUTES: Final = ["openai_routes", "anthropic_routes", "/model/info"] + + +def _sign_connect_code(master_key: str, moyai_url: str, user_id: str | None) -> str: + payload: Final = json.dumps( + { + "moyai_origin": _origin(moyai_url), + "user_id": user_id, + "exp": int(time.time()) + _MOYAI_CODE_TTL_SECONDS, + "nonce": secrets.token_urlsafe(16), + }, + separators=(",", ":"), + sort_keys=True, + ).encode() + signature: Final = hmac.new(_master_key_hmac_key(master_key), payload, hashlib.sha256).digest() + return f"{_b64url(payload)}.{_b64url(signature)}" + + +def _decode_connect_code(master_key: str, code: str) -> dict: + try: + payload_b64, signature_b64 = code.split(".", 1) + payload_raw: Final = _b64url_decode(payload_b64) + signature: Final = _b64url_decode(signature_b64) + except (ValueError, TypeError): + raise HTTPException(status_code=400, detail="Invalid Moyai connect code") + expected: Final = hmac.new(_master_key_hmac_key(master_key), payload_raw, hashlib.sha256).digest() + if not hmac.compare_digest(signature, expected): + raise HTTPException(status_code=400, detail="Invalid Moyai connect code") + try: + payload: Final = json.loads(payload_raw) + except (ValueError, TypeError): + raise HTTPException(status_code=400, detail="Invalid Moyai connect code") + if not isinstance(payload, dict) or not isinstance(payload.get("exp"), int) or payload["exp"] < int(time.time()): + raise HTTPException(status_code=400, detail="Invalid Moyai connect code") + if not isinstance(payload.get("moyai_origin"), str) or not isinstance(payload.get("nonce"), str): + raise HTTPException(status_code=400, detail="Invalid Moyai connect code") + return payload + + +@router.post( + "/moyai/connect/start", + response_model=MoyaiConnectStartResponse, + tags=["moyai"], +) +async def moyai_connect_start( + request: Request, + body: MoyaiConnectStartRequest, + user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)], +) -> MoyaiConnectStartResponse: + from litellm.proxy.proxy_server import master_key + + if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN.value: + raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Only proxy admins can connect Moyai") + + try: + moyai_url: Final = normalize_moyai_url(body.moyai_url) + except ValueError as e: + raise HTTPException(status_code=400, detail=str(e)) + if moyai_url is None: + raise HTTPException(status_code=400, detail="moyai_url is required") + + return_to_parsed: Final = urlparse(body.return_to) + if return_to_parsed.scheme not in ("http", "https") or not return_to_parsed.netloc: + raise HTTPException(status_code=400, detail="return_to must be an absolute http or https URL") + + if not master_key: + raise HTTPException( + status_code=400, + detail="Moyai quick connect needs LITELLM_MASTER_KEY set on the proxy", + ) + + code: Final = _sign_connect_code(master_key, moyai_url, user_api_key_dict.user_id) + connect_url: Final = f"{moyai_url}/connect/litellm?" + urlencode( + {"gateway_url": _gateway_url(request), "code": code, "return_to": body.return_to} + ) + return MoyaiConnectStartResponse(connect_url=connect_url) + + +async def _claim_connect_nonce(prisma_client: "PrismaClient", nonce: str, exp: int) -> None: + from prisma.errors import UniqueViolationError + + try: + await ConfigRepository(prisma_client, use_writer=True).table.create( + data={ + "param_name": f"{_MOYAI_NONCE_CONFIG_PREFIX}{nonce}", + "param_value": json.dumps({"exp": exp}), + } + ) + except UniqueViolationError: + raise HTTPException(status_code=400, detail="Invalid Moyai connect code") + + +async def _moyai_key_alias(prisma_client, moyai_url: str) -> str: + from litellm.repositories.verification_token_repository import ( + VerificationTokenRepository, + ) + + host: Final = urlparse(moyai_url).hostname or "deployment" + alias: Final = f"moyai-{host}" + rows: Final = await VerificationTokenRepository(prisma_client).find_many(where={"key_alias": alias}, take=1) + if rows: + return f"{alias}-{secrets.token_hex(2)}" + return alias + + +@with_service_target(CONFIG_PARAMS_TARGET) +async def _persist_moyai_url(prisma_client, moyai_url: str) -> None: + from litellm.proxy.proxy_server import user_api_key_cache + + existing: dict = {} + db_existing: Final = await _ui_settings_db(UISettingsRepository(prisma_client)).find_unique( + where={"id": "ui_settings"} + ) + if db_existing and db_existing.ui_settings: + raw: Final = db_existing.ui_settings + existing = json.loads(raw) if isinstance(raw, str) else dict(raw) + + ui_settings: Final = {**existing, "moyai_url": moyai_url} + await _ui_settings_db(UISettingsRepository(prisma_client)).upsert( + where={"id": "ui_settings"}, + data={ + "create": {"id": "ui_settings", "ui_settings": json.dumps(ui_settings)}, + "update": {"ui_settings": json.dumps(ui_settings)}, + }, + ) + await user_api_key_cache.async_set_cache(key=UI_SETTINGS_CACHE_KEY, value=ui_settings, ttl=UI_SETTINGS_CACHE_TTL) + + +@router.post( + _MOYAI_CONNECT_EXCHANGE_ROUTE, + response_model=MoyaiConnectExchangeResponse, + tags=["moyai"], +) +async def moyai_connect_exchange(request: Request, body: MoyaiConnectExchangeRequest) -> MoyaiConnectExchangeResponse: + from litellm.proxy.management_endpoints.key_management_endpoints import generate_key_helper_fn + from litellm.proxy.proxy_server import llm_router, master_key, prisma_client + + if not master_key: + raise HTTPException(status_code=400, detail="Invalid Moyai connect code") + + payload: Final = _decode_connect_code(master_key, body.code) + + try: + moyai_url: Final = normalize_moyai_url(body.moyai_url) + except ValueError: + raise HTTPException(status_code=400, detail="Invalid Moyai connect code") + if moyai_url is None or _origin(moyai_url) != payload["moyai_origin"]: + raise HTTPException(status_code=400, detail="Invalid Moyai connect code") + + if prisma_client is None: + raise HTTPException(status_code=400, detail="Moyai quick connect needs a database connected to the proxy") + + await _claim_connect_nonce(prisma_client, payload["nonce"], payload["exp"]) + + alias: Final = await _moyai_key_alias(prisma_client, moyai_url) + key_response: Final = await generate_key_helper_fn( + request_type="key", + key_alias=alias, + allowed_routes=_MOYAI_KEY_ALLOWED_ROUTES, + metadata={ + "created_via": "moyai_quick_connect", + "moyai_url": moyai_url, + "connected_by": payload.get("user_id"), + }, + table_name="key", + llm_router=llm_router, + ) + + await _persist_moyai_url(prisma_client, moyai_url) + + return MoyaiConnectExchangeResponse( + api_key=key_response["token"], + key_alias=alias, + api_base=_gateway_url(request), + ) diff --git a/litellm/proxy/ocr_endpoints/endpoints.py b/litellm/proxy/ocr_endpoints/endpoints.py index 81d9f8a7b43..94fe9354386 100644 --- a/litellm/proxy/ocr_endpoints/endpoints.py +++ b/litellm/proxy/ocr_endpoints/endpoints.py @@ -340,7 +340,7 @@ async def ocr( return _native_response(response, fastapi_response) or response except Exception as e: processor = ProxyBaseLLMRequestProcessing(data=data) - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, diff --git a/litellm/proxy/openai_evals_endpoints/endpoints.py b/litellm/proxy/openai_evals_endpoints/endpoints.py index abfbed5f822..aa1c977b8e7 100644 --- a/litellm/proxy/openai_evals_endpoints/endpoints.py +++ b/litellm/proxy/openai_evals_endpoints/endpoints.py @@ -107,7 +107,7 @@ async def create_eval( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -208,7 +208,7 @@ async def list_evals( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -296,7 +296,7 @@ async def get_eval( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -386,7 +386,7 @@ async def update_eval( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -474,7 +474,7 @@ async def delete_eval( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -562,7 +562,7 @@ async def cancel_eval( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -666,7 +666,7 @@ async def create_run( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -759,7 +759,7 @@ async def list_runs( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -846,7 +846,7 @@ async def get_run( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -935,7 +935,7 @@ async def cancel_run( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -1024,7 +1024,7 @@ async def delete_run( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, diff --git a/litellm/proxy/openai_files_endpoints/common_utils.py b/litellm/proxy/openai_files_endpoints/common_utils.py index 6328b900ed0..fb0ae09382e 100644 --- a/litellm/proxy/openai_files_endpoints/common_utils.py +++ b/litellm/proxy/openai_files_endpoints/common_utils.py @@ -104,7 +104,7 @@ class ManagedFileIdResolver(Protocol): ) -> Mapping[str, str]: ... -def _is_base64_encoded_unified_file_id(b64_uid: object) -> str | Literal[False]: +def is_base64_encoded_unified_file_id(b64_uid: object) -> str | Literal[False]: # Ensure b64_uid is a string and not a mock object if not isinstance(b64_uid, str): return False @@ -121,8 +121,11 @@ def _is_base64_encoded_unified_file_id(b64_uid: object) -> str | Literal[False]: return False +_is_base64_encoded_unified_file_id: Final = is_base64_encoded_unified_file_id + + def convert_b64_uid_to_unified_uid(b64_uid: str) -> str: - is_base64_unified_file_id: Final = _is_base64_encoded_unified_file_id(b64_uid) + is_base64_unified_file_id: Final = is_base64_encoded_unified_file_id(b64_uid) if is_base64_unified_file_id: return is_base64_unified_file_id else: @@ -942,10 +945,10 @@ async def extract_file_creation_params( Returns: FileCreationParams: Structured parameters extracted from the request """ - from litellm.proxy.common_utils.http_parsing_utils import _read_request_body + from litellm.proxy.common_utils.http_parsing_utils import read_request_body if request_body is None: - request_body = await _read_request_body(request=request) or {} + request_body = await read_request_body(request=request) or {} # Extract target_storage (simplified - just use form parameter) target_storage: Final = _extract_target_storage_simple(target_storage_form) @@ -1097,7 +1100,7 @@ async def validate_managed_id_requirement( if not resource_id: return - if not _is_base64_encoded_unified_file_id(resource_id): + if not is_base64_encoded_unified_file_id(resource_id): raise HTTPException( status_code=400, detail=( @@ -1150,7 +1153,7 @@ def _batch_response_model_id_candidates( ) -> tuple[str, ...]: response_id: Final = getattr(response, "id", None) decoded_response_id: Final = ( - _is_base64_encoded_unified_file_id(response_id) if isinstance(response_id, str) else False + is_base64_encoded_unified_file_id(response_id) if isinstance(response_id, str) else False ) return tuple( candidate @@ -1215,7 +1218,7 @@ async def resolve_input_file_id_to_unified(response, prisma_client) -> None: if ( hasattr(response, "input_file_id") and response.input_file_id - and not _is_base64_encoded_unified_file_id(response.input_file_id) + and not is_base64_encoded_unified_file_id(response.input_file_id) and prisma_client ): try: @@ -1238,7 +1241,7 @@ async def resolve_output_file_ids_to_unified(response, prisma_client) -> None: return for attr in ("output_file_id", "error_file_id"): raw_id = getattr(response, attr, None) - if not raw_id or _is_base64_encoded_unified_file_id(raw_id): + if not raw_id or is_base64_encoded_unified_file_id(raw_id): continue try: managed_file = await ManagedFileRepository(prisma_client).table.find_first( @@ -1307,7 +1310,7 @@ async def ensure_batch_response_managed_file_ids( for file_attr in ("output_file_id", "error_file_id"): raw_file_id = getattr(response, file_attr, None) - if not raw_file_id or _is_base64_encoded_unified_file_id(raw_file_id): + if not raw_file_id or is_base64_encoded_unified_file_id(raw_file_id): continue try: new_unified_file_id = managed_files_obj.get_unified_output_file_id( diff --git a/litellm/proxy/openai_files_endpoints/files_endpoints.py b/litellm/proxy/openai_files_endpoints/files_endpoints.py index e170e1cf894..4e31f9713d0 100644 --- a/litellm/proxy/openai_files_endpoints/files_endpoints.py +++ b/litellm/proxy/openai_files_endpoints/files_endpoints.py @@ -42,9 +42,10 @@ from litellm.proxy.batches_endpoints.litellm_executed_batches import ( resolve_litellm_executed_provider, ) from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing -from litellm.proxy.common_utils.http_parsing_utils import ( - _read_request_body, +from litellm.proxy.common_utils.http_parsing_utils import ( # noqa: F401 # legacy module exports + _read_request_body, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export extract_nested_form_metadata, + read_request_body, ) from litellm.proxy.common_utils.openai_endpoint_utils import ( get_custom_llm_provider_from_request_body, @@ -70,8 +71,8 @@ from litellm.proxy.openai_files_endpoints.batch_guardrails import ( rewrite_batch_input_file, scan_batch_input_file, ) -from litellm.proxy.openai_files_endpoints.common_utils import ( - _is_base64_encoded_unified_file_id, +from litellm.proxy.openai_files_endpoints.common_utils import ( # noqa: F401 # legacy module exports + _is_base64_encoded_unified_file_id, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export add_internal_model_credentials, apply_team_provider_credentials, authorize_model_for_key, @@ -79,6 +80,7 @@ from litellm.proxy.openai_files_endpoints.common_utils import ( extract_file_creation_params, get_authorized_credentials_for_model, handle_model_based_routing, + is_base64_encoded_unified_file_id, prepare_data_with_credentials, validate_file_list_limit, validate_managed_files_requirement, @@ -620,7 +622,7 @@ async def create_file( ) # Extract file creation parameters using utility function - request_body: Final = await _read_request_body(request=request) or {} + request_body: Final = await read_request_body(request=request) or {} file_params: Final = await extract_file_creation_params( request=request, request_body=request_body, @@ -1019,7 +1021,7 @@ async def get_file_content( ) ## check if file_id is a litellm managed file - is_base64_unified_file_id: Final = _is_base64_encoded_unified_file_id(file_id) + is_base64_unified_file_id: Final = is_base64_encoded_unified_file_id(file_id) if is_base64_unified_file_id: managed_files_obj: Final = proxy_logging_obj.get_proxy_hook("managed_files") if managed_files_obj is None: @@ -1366,7 +1368,7 @@ async def get_file( ) ## EXISTING: check if file_id is a litellm managed file - elif _is_base64_encoded_unified_file_id(file_id): + elif is_base64_encoded_unified_file_id(file_id): managed_files_obj: Final = proxy_logging_obj.get_proxy_hook("managed_files") if managed_files_obj is None: raise ProxyException( @@ -1576,7 +1578,7 @@ async def delete_file( ) ## EXISTING: check if file_id is a litellm managed file - elif _is_base64_encoded_unified_file_id(file_id): + elif is_base64_encoded_unified_file_id(file_id): managed_files_obj: Final = proxy_logging_obj.get_proxy_hook("managed_files") if managed_files_obj is None: raise ProxyException( diff --git a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py index acacfd193bc..edc9d96b298 100644 --- a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py @@ -69,21 +69,25 @@ from litellm.proxy._types import * from litellm.proxy.auth.auth_checks import enforced_model_allowlists from litellm.proxy.auth.handle_jwt import JWTHandler from litellm.proxy.auth.route_checks import RouteChecks -from litellm.proxy.auth.user_api_key_auth import ( - _get_bearer_token, +from litellm.proxy.auth.user_api_key_auth import ( # noqa: F401 # legacy module exports + _get_bearer_token, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + get_bearer_token, is_no_auth_dev_mode, user_api_key_auth, user_api_key_auth_websocket, user_api_key_auth_websocket_for_model, ) from litellm.proxy.common_request_processing import open_sse_before_first_byte -from litellm.proxy.common_utils.http_parsing_utils import ( - _read_request_body, - _safe_get_request_headers, - _safe_set_request_parsed_body, +from litellm.proxy.common_utils.http_parsing_utils import ( # noqa: F401 # legacy module exports + _read_request_body, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _safe_get_request_headers, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _safe_set_request_parsed_body, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export get_form_data, get_request_body, is_json_content_type, + read_request_body, + safe_get_request_headers, + safe_set_request_parsed_body, ) from litellm.proxy.common_utils.resource_ownership import is_proxy_admin from litellm.proxy.common_utils.sse_keepalive import ( @@ -142,7 +146,7 @@ def create_request_copy(request: Request): return { "method": request.method, "url": str(request.url), - "headers": _safe_get_request_headers(request).copy(), + "headers": safe_get_request_headers(request).copy(), "cookies": request.cookies, "query_params": dict(request.query_params), } @@ -468,7 +472,7 @@ async def fal_ai_proxy_route( status_code=401, detail="FAL_AI_API_KEY is not set and no fal_ai pass-through deployment credentials are configured", ) - if "/requests/" not in endpoint and fal_ai_passthrough_cost(endpoint, await _read_request_body(request)) is None: + if "/requests/" not in endpoint and fal_ai_passthrough_cost(endpoint, await read_request_body(request)) is None: raise HTTPException( status_code=400, detail=f"fal_ai/{endpoint} has no pricing entry for this request; only priced Fal requests can be submitted through /fal_ai", @@ -510,7 +514,7 @@ async def vllm_proxy_route( method=request.method, endpoint=endpoint, request_query_params=request.query_params, - request_headers=_safe_get_request_headers(request), + request_headers=safe_get_request_headers(request), stream=is_streaming_request, content=None, data=None, @@ -665,7 +669,7 @@ async def bespoke_proxy_route( async def _oss_decision_proxy_route( provider: OssDecisionProvider, request: Request, fastapi_response: Response, user_api_key_dict: UserAPIKeyAuth ) -> Response: - body: Final = TypeAdapter(dict[str, object]).validate_python(await _read_request_body(request)) + body: Final = TypeAdapter(dict[str, object]).validate_python(await read_request_body(request)) try: _ = validate_oss_request(provider, body) except ValueError as exc: @@ -813,7 +817,7 @@ async def milvus_proxy_route( request_body["collectionName"] = vector_store_index # Update the request object with the modified collection name - _safe_set_request_parsed_body(request, request_body) + safe_set_request_parsed_body(request, request_body) vector_store: Final = litellm.vector_store_registry.get_litellm_managed_vector_store_from_registry_by_name( vector_store_name=vector_store_name @@ -871,7 +875,7 @@ async def is_streaming_request_fn(request: Request) -> bool: if content_type and "multipart/form-data" in content_type: _request_body = await get_form_data(request) else: - _request_body = await _read_request_body(request) + _request_body = await read_request_body(request) # rebind-ok: pre-existing rebinding on a rename-only line return is_passthrough_request_streaming(_request_body) return False @@ -949,7 +953,7 @@ def is_bedrock_count_tokens_endpoint(endpoint: str) -> bool: return "count_tokens" in endpoint or "count-tokens" in endpoint -def _extract_model_from_bedrock_endpoint(endpoint: str) -> str: +def extract_model_from_bedrock_endpoint(endpoint: str) -> str: """ Extract model name from Bedrock endpoint path. @@ -1029,6 +1033,9 @@ def _extract_model_from_bedrock_endpoint(endpoint: str) -> str: ) from e +_extract_model_from_bedrock_endpoint: Final = extract_model_from_bedrock_endpoint + + async def handle_bedrock_passthrough_router_model( model: str, endpoint: str, @@ -1110,7 +1117,7 @@ async def handle_bedrock_passthrough_router_model( return result except Exception as e: # Use common exception handling - raise await base_llm_response_processor._handle_llm_api_exception( + raise await base_llm_response_processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -1228,7 +1235,7 @@ async def bedrock_llm_proxy_route( version, ) - request_body: Final = await _read_request_body(request=request) + request_body: Final = await read_request_body(request=request) if is_bedrock_count_tokens_endpoint(endpoint): return await handle_bedrock_count_tokens( @@ -1241,7 +1248,7 @@ async def bedrock_llm_proxy_route( # Extract model from endpoint path using helper try: - model: Final = _extract_model_from_bedrock_endpoint(endpoint=endpoint) + model: Final = extract_model_from_bedrock_endpoint(endpoint=endpoint) except ValueError as e: raise HTTPException( status_code=400, @@ -1307,7 +1314,7 @@ async def bedrock_llm_proxy_route( return result except Exception as e: - raise await base_llm_response_processor._handle_llm_api_exception( + raise await base_llm_response_processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -1645,7 +1652,7 @@ async def azure_speech_proxy_route( target_url: Final = base_url.copy_with( path=HttpPassThroughEndpointHelpers.join_base_and_endpoint_path(base_url, normalized_endpoint_path) ) - request_headers: Final = _safe_get_request_headers(request) + request_headers: Final = safe_get_request_headers(request) upstream_headers: Final = MappingProxyType( { header_name: header_value @@ -1957,8 +1964,10 @@ async def assemblyai_proxy_route( [Docs](https://api.assemblyai.com) """ # Set base URL based on the route - assembly_region: Final = AssemblyAIPassthroughLoggingHandler._get_assembly_region_from_url(url=str(request.url)) - base_target_url = AssemblyAIPassthroughLoggingHandler._get_assembly_base_url_from_region(region=assembly_region) + assembly_region: Final = AssemblyAIPassthroughLoggingHandler.get_assembly_region_from_url(url=str(request.url)) + base_target_url: Final = AssemblyAIPassthroughLoggingHandler.get_assembly_base_url_from_region( + region=assembly_region + ) encoded_endpoint = httpx.URL(endpoint).path # Ensure endpoint starts with '/' for proper URL construction if not encoded_endpoint.startswith("/"): @@ -2092,7 +2101,7 @@ async def _relay_router_model( method=request.method, endpoint=endpoint, request_query_params=request.query_params, - request_headers=_safe_get_request_headers(request), + request_headers=safe_get_request_headers(request), stream=is_streaming_request, content=None, data=None, @@ -2332,7 +2341,7 @@ async def azure_proxy_route( base_target_url = _optional_str(litellm_params.get("api_base")) if base_target_url is None: raise Exception(f"API base not found for {part}") - return await BaseOpenAIPassThroughHandler._base_openai_pass_through_handler( + return await BaseOpenAIPassThroughHandler.base_openai_pass_through_handler( endpoint=endpoint, request=request, fastapi_response=fastapi_response, @@ -2360,7 +2369,7 @@ async def azure_proxy_route( if azure_api_key is None: raise Exception("Required 'AZURE_API_KEY' in environment to make pass-through calls to Azure.") - return await BaseOpenAIPassThroughHandler._base_openai_pass_through_handler( + return await BaseOpenAIPassThroughHandler.base_openai_pass_through_handler( endpoint=endpoint, request=request, fastapi_response=fastapi_response, @@ -2431,7 +2440,7 @@ def get_vertex_ai_allowed_incoming_headers(request: Request) -> dict: Returns: dict: Headers dictionary with only allowed headers """ - incoming_headers: Final = _safe_get_request_headers(request) + incoming_headers: Final = safe_get_request_headers(request) headers: Final = {} for header_name in ALLOWED_VERTEX_AI_PASSTHROUGH_HEADERS: if header_name in incoming_headers: @@ -2531,7 +2540,7 @@ def _normalize_credential_value(value: str) -> str: with no recognized scheme prefix, so a bare token (or a real Google credential that carries no scheme) falls back to its own value. """ - return _get_bearer_token(value) or value + return get_bearer_token(value) or value _VERTEX_UPSTREAM_CREDENTIAL_HEADERS: Final = frozenset({"authorization", "x-goog-api-key"}) @@ -2616,7 +2625,7 @@ def _is_authenticated_caller_secret(value: str, user_api_key_dict: UserAPIKeyAut def _caller_headers_without_litellm_secrets( request: Request, user_api_key_dict: UserAPIKeyAuth, never_forwarded: frozenset[str] ) -> Mapping[str, str]: - incoming: Final = _safe_get_request_headers(request) + incoming: Final = safe_get_request_headers(request) dropped_by_name: Final = never_forwarded.union( (_MAPPED_ROUTE_CALLER_KEY_HEADER, *_operator_configured_caller_key_header_names()) ) @@ -3048,7 +3057,7 @@ async def openai_proxy_route( if openai_api_key is None: raise Exception("Required 'OPENAI_API_KEY' in environment to make pass-through calls to OpenAI.") - return await BaseOpenAIPassThroughHandler._base_openai_pass_through_handler( + return await BaseOpenAIPassThroughHandler.base_openai_pass_through_handler( endpoint=endpoint, request=request, fastapi_response=fastapi_response, @@ -3320,7 +3329,7 @@ async def deepgram_listen_websocket_route( class BaseOpenAIPassThroughHandler: @staticmethod - async def _base_openai_pass_through_handler( + async def base_openai_pass_through_handler( endpoint: str, request: Request, fastapi_response: Response, @@ -3372,12 +3381,14 @@ class BaseOpenAIPassThroughHandler: return received_value + _base_openai_pass_through_handler = base_openai_pass_through_handler + @staticmethod def _append_openai_beta_header(headers: dict, request: Request) -> dict: """ Appends the OpenAI-Beta header to the headers if the request is an OpenAI Assistants API request """ - if RouteChecks._is_assistants_api_request(request) is True and "OpenAI-Beta" not in headers: + if RouteChecks.is_assistants_api_request(request) is True and "OpenAI-Beta" not in headers: headers["OpenAI-Beta"] = "assistants=v2" return headers @@ -4023,7 +4034,7 @@ async def handle_gigachat_passthrough_router_model( is_streaming: Final = request_body.get("stream", False) # pyright: ignore[reportUnknownVariableType] # request_body is dict[Unknown, Unknown] - data: Final[dict[str, object]] = await _read_request_body(request=request) + data: Final[dict[str, object]] = await read_request_body(request=request) if user_api_key_dict is not None: auth_metadata: Final = { metadata_key: value @@ -4096,7 +4107,7 @@ async def handle_gigachat_passthrough_router_model( ) except Exception as e: # noqa: BLE001 # Safe catch-all for handle exception # Use common exception handling - raise await base_llm_response_processor._handle_llm_api_exception( + raise await base_llm_response_processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/anthropic_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/anthropic_passthrough_logging_handler.py index dfb731972ac..73d40369d30 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/anthropic_passthrough_logging_handler.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/anthropic_passthrough_logging_handler.py @@ -173,7 +173,7 @@ class AnthropicPassthroughLoggingHandler: all_chunks: Sequence[str | bytes], model: str, speed: str | None ) -> ModelResponse | None: try: - return AnthropicPassthroughLoggingHandler._build_usage_only_response_from_chunks( + return AnthropicPassthroughLoggingHandler.build_usage_only_response_from_chunks( all_chunks=all_chunks, model=model, speed=speed ) except Exception as e: # noqa: BLE001 # the usage-only fallback must never raise out of failure logging @@ -188,7 +188,7 @@ class AnthropicPassthroughLoggingHandler: speed: str | None, ) -> ModelResponse | TextCompletionResponse | None: try: - assembled: Final = AnthropicPassthroughLoggingHandler._build_complete_streaming_response( + assembled: Final = AnthropicPassthroughLoggingHandler.build_complete_streaming_response( all_chunks=all_chunks, litellm_logging_obj=litellm_logging_obj, model=model, @@ -467,7 +467,7 @@ class AnthropicPassthroughLoggingHandler: return kwargs @staticmethod - def _handle_logging_anthropic_collected_chunks( + def handle_logging_anthropic_collected_chunks( litellm_logging_obj: LiteLLMLoggingObj, passthrough_success_handler_obj: PassThroughEndpointLogging, url_route: str, @@ -514,6 +514,8 @@ class AnthropicPassthroughLoggingHandler: "kwargs": kwargs, } + _handle_logging_anthropic_collected_chunks = handle_logging_anthropic_collected_chunks + @staticmethod def _split_sse_chunk_into_events(chunk: str | bytes) -> list[str]: """ @@ -539,7 +541,7 @@ class AnthropicPassthroughLoggingHandler: return events @staticmethod - def _build_complete_streaming_response( + def build_complete_streaming_response( all_chunks: Sequence[str | bytes], litellm_logging_obj: LiteLLMLoggingObj, model: str, @@ -577,6 +579,8 @@ class AnthropicPassthroughLoggingHandler: speed=speed, ) + _build_complete_streaming_response = build_complete_streaming_response + # Anthropic SSE block/delta types that the fast path is NOT allowed to # collapse -- their presence forces the unchanged legacy path so tool # calls, thinking, citations, etc. keep byte-identical reconstruction. @@ -775,7 +779,7 @@ class AnthropicPassthroughLoggingHandler: return None @staticmethod - def _build_usage_only_response_from_chunks( + def build_usage_only_response_from_chunks( all_chunks: Sequence[str | bytes], model: str, speed: str | None = None, @@ -894,6 +898,8 @@ class AnthropicPassthroughLoggingHandler: usage=usage_obj, ) + _build_usage_only_response_from_chunks = build_usage_only_response_from_chunks + @staticmethod def batch_creation_handler( httpx_response: httpx.Response, diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/assembly_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/assembly_passthrough_logging_handler.py index 812f72faecc..c9c89870610 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/assembly_passthrough_logging_handler.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/assembly_passthrough_logging_handler.py @@ -147,7 +147,7 @@ class AssemblyAIPassthroughLoggingHandler: logging_obj.model_call_details["response_cost"] = response_cost asyncio.run( - pass_through_endpoint_logging._handle_logging( + pass_through_endpoint_logging.handle_logging( logging_obj=logging_obj, standard_logging_response_object=self._get_response_to_log(transcript_response), result=result, @@ -216,7 +216,7 @@ class AssemblyAIPassthroughLoggingHandler: """ for _ in range(self.max_polling_attempts): # 180 attempts * 10s = 30 minutes max transcript = self._get_assembly_transcript( - request_region=AssemblyAIPassthroughLoggingHandler._get_assembly_region_from_url(url=url_route), + request_region=AssemblyAIPassthroughLoggingHandler.get_assembly_region_from_url(url=url_route), transcript_id=transcript_id, ) if transcript is None: @@ -279,14 +279,16 @@ class AssemblyAIPassthroughLoggingHandler: return None @staticmethod - def _should_log_request(request_method: str) -> bool: + def should_log_request(request_method: str) -> bool: """ only POST transcription jobs are logged. litellm will POLL assembly to wait for the transcription to complete to log the complete response / cost """ return request_method == "POST" + _should_log_request = should_log_request + @staticmethod - def _get_assembly_region_from_url(url: str | None) -> Literal["eu"] | None: + def get_assembly_region_from_url(url: str | None) -> Literal["eu"] | None: """ Get the region from the URL """ @@ -296,8 +298,10 @@ class AssemblyAIPassthroughLoggingHandler: return "eu" return None + _get_assembly_region_from_url = get_assembly_region_from_url + @staticmethod - def _get_assembly_base_url_from_region(region: Literal["eu"] | None) -> str: + def get_assembly_base_url_from_region(region: Literal["eu"] | None) -> str: """ Get the base URL for the AssemblyAI API if region == "eu", return "https://api.eu.assemblyai.com" @@ -306,3 +310,5 @@ class AssemblyAIPassthroughLoggingHandler: if region == "eu": return "https://api.eu.assemblyai.com" return "https://api.assemblyai.com" + + _get_assembly_base_url_from_region = get_assembly_base_url_from_region diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/openai_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/openai_passthrough_logging_handler.py index 0fcc1ccda38..2e9d92ed309 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/openai_passthrough_logging_handler.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/openai_passthrough_logging_handler.py @@ -78,7 +78,7 @@ def _is_openai_compatible_host(hostname: str | None) -> bool: return _hostname_matches(hostname, _OPENAI_HOSTNAMES) or _hostname_matches(hostname, _AZURE_OPENAI_HOSTNAMES) -def _is_openai_compatible_url(url_route: str | None) -> bool: +def is_openai_compatible_url(url_route: str | None) -> bool: """True if the URL targets an OpenAI-compatible API surface. For the shared Azure Cognitive Services domains we additionally require an @@ -99,6 +99,9 @@ def _is_openai_compatible_url(url_route: str | None) -> bool: return False +_is_openai_compatible_url: Final = is_openai_compatible_url + + def _is_remote_high_detail_image(part: object) -> bool: if not isinstance(part, Mapping) or part.get("type") != "image_url": return False @@ -608,7 +611,7 @@ class OpenAIPassthroughLoggingHandler(BasePassthroughLoggingHandler): return None @staticmethod - def _handle_logging_openai_collected_chunks( + def handle_logging_openai_collected_chunks( litellm_logging_obj: LiteLLMLoggingObj, passthrough_success_handler_obj: PassThroughEndpointLogging, url_route: str, @@ -725,3 +728,5 @@ class OpenAIPassthroughLoggingHandler(BasePassthroughLoggingHandler): "result": None, "kwargs": {}, } + + _handle_logging_openai_collected_chunks = handle_logging_openai_collected_chunks diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_passthrough_logging_handler.py index 23b97826bd7..c7422da7e23 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_passthrough_logging_handler.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_passthrough_logging_handler.py @@ -561,7 +561,7 @@ class VertexPassthroughLoggingHandler: } @staticmethod - def _handle_logging_vertex_collected_chunks( + def handle_logging_vertex_collected_chunks( litellm_logging_obj: LiteLLMLoggingObj, passthrough_success_handler_obj: PassThroughEndpointLogging, url_route: str, @@ -616,6 +616,8 @@ class VertexPassthroughLoggingHandler: "kwargs": kwargs, } + _handle_logging_vertex_collected_chunks = handle_logging_vertex_collected_chunks + @staticmethod def _build_complete_streaming_response( all_chunks: list[str], diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index 0a0982021da..e80d3ed189b 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -91,9 +91,11 @@ from litellm.proxy.common_request_processing import ( resolve_litellm_call_id, ) from litellm.proxy.common_utils.error_body_call_id import JSON_OBJECT, error_body_call_id, with_call_id -from litellm.proxy.common_utils.http_parsing_utils import ( - _read_request_body, - _safe_get_request_headers, +from litellm.proxy.common_utils.http_parsing_utils import ( # noqa: F401 # legacy module exports + _read_request_body, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _safe_get_request_headers, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + read_request_body, + safe_get_request_headers, ) from litellm.proxy.common_utils.openai_error_payload import ( LITELLM_CALL_ID_HEADER, @@ -105,11 +107,12 @@ from litellm.proxy.common_utils.openai_error_payload import ( from litellm.proxy.common_utils.sse_keepalive import ( wrap_passthrough_sse_bytes_with_keepalive_pings, ) -from litellm.proxy.litellm_pre_call_utils import ( +from litellm.proxy.litellm_pre_call_utils import ( # noqa: F401 # legacy module exports LiteLLMProxyRequestSetup, - _get_dynamic_logging_metadata, # pyright: ignore[reportPrivateUsage] # shared proxy helper, same import style as _read_request_body above + _get_dynamic_logging_metadata, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export _key_or_team_allows_client_pricing_override, # pyright: ignore[reportPrivateUsage] # reuse the proxy's pricing trust policy _strip_client_pricing_overrides, # pyright: ignore[reportPrivateUsage] # sanitize before trusted hooks add guardrail costs + get_dynamic_logging_metadata, # pyright: ignore[reportPrivateUsage] # shared proxy helper, same import style as _read_request_body above ) from litellm.proxy.route_llm_request import ProxyModelNotFoundError from litellm.proxy.utils import normalize_route_for_root_path @@ -586,7 +589,7 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils): ) @staticmethod - def _init_kwargs_for_pass_through_endpoint( + def init_kwargs_for_pass_through_endpoint( request: Request, user_api_key_dict: UserAPIKeyAuth, passthrough_logging_payload: PassthroughStandardLoggingPayload, @@ -697,6 +700,8 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils): return kwargs + _init_kwargs_for_pass_through_endpoint = init_kwargs_for_pass_through_endpoint + @staticmethod def construct_target_url_with_subpath(base_target: str, subpath: str, include_subpath: bool | None) -> str: """ @@ -768,7 +773,7 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils): return combined @staticmethod - def _update_stream_param_based_on_request_body( + def update_stream_param_based_on_request_body( parsed_body: dict, stream: bool | None = None, ) -> bool | None: @@ -780,6 +785,8 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils): return parsed_body.get("stream", stream) return stream + _update_stream_param_based_on_request_body = update_stream_param_based_on_request_body + def _carry_guardrail_logging_info(request_data: dict, guardrail_data: dict | None) -> None: """Copy guardrail logging entries from ``guardrail_data`` onto ``request_data``. @@ -863,7 +870,7 @@ def _resolve_team_callback_wiring( otherwise reject the vars mid-request. """ try: - callback_settings_obj: Final = _get_dynamic_logging_metadata( + callback_settings_obj: Final = get_dynamic_logging_metadata( user_api_key_dict=user_api_key_dict, proxy_config=proxy_config ) if callback_settings_obj and callback_settings_obj.callback_vars: @@ -1125,7 +1132,7 @@ async def pass_through_request( url = httpx.URL(target) headers = custom_headers headers = HttpPassThroughEndpointHelpers.forward_headers_from_request( - request_headers=_safe_get_request_headers(request).copy(), + request_headers=safe_get_request_headers(request).copy(), headers=headers, forward_headers=forward_headers, ) @@ -1161,7 +1168,7 @@ async def pass_through_request( # Don't parse multipart body here - it will be handled by make_multipart_http_request _parsed_body = {} else: - _parsed_body = await _read_request_body(request) + _parsed_body = await read_request_body(request) # rebind-ok: pre-existing rebinding on a rename-only line verbose_proxy_logger.debug( "Pass through endpoint sending request to \nURL %s\nheaders: %s\nbody: %s\n", url, @@ -1258,7 +1265,7 @@ async def pass_through_request( request_method=getattr(request, "method", None), cost_per_request=cost_per_request, ) - kwargs = HttpPassThroughEndpointHelpers._init_kwargs_for_pass_through_endpoint( + kwargs = HttpPassThroughEndpointHelpers.init_kwargs_for_pass_through_endpoint( # rebind-ok: pre-existing rebinding on a rename-only line user_api_key_dict=user_api_key_dict, _parsed_body=_parsed_body, passthrough_logging_payload=passthrough_logging_payload, @@ -1432,7 +1439,7 @@ async def pass_through_request( "headers": upstream_headers, }, ) - stream = HttpPassThroughEndpointHelpers._update_stream_param_based_on_request_body( + stream = HttpPassThroughEndpointHelpers.update_stream_param_based_on_request_body( parsed_body=_parsed_body or {}, stream=stream, ) @@ -1980,7 +1987,7 @@ def _update_metadata_with_tags_in_header(request: Request, metadata: dict) -> di # Only add tags key if there are tags to add if tags_to_add: - metadata["tags"] = LiteLLMProxyRequestSetup._merge_tags( + metadata["tags"] = LiteLLMProxyRequestSetup.merge_tags( request_tags=metadata.get("tags"), tags_to_add=tags_to_add, ) @@ -2480,7 +2487,7 @@ async def websocket_passthrough_request( ) # Initialize kwargs for logging using the same pattern as HTTP passthrough - kwargs: Final = HttpPassThroughEndpointHelpers._init_kwargs_for_pass_through_endpoint( + kwargs: Final = HttpPassThroughEndpointHelpers.init_kwargs_for_pass_through_endpoint( user_api_key_dict=user_api_key_dict, _parsed_body={}, # WebSocket doesn't have a traditional request body passthrough_logging_payload=passthrough_logging_payload, diff --git a/litellm/proxy/pass_through_endpoints/passthrough_guardrails.py b/litellm/proxy/pass_through_endpoints/passthrough_guardrails.py index de9b0acf081..1c5d270542f 100644 --- a/litellm/proxy/pass_through_endpoints/passthrough_guardrails.py +++ b/litellm/proxy/pass_through_endpoints/passthrough_guardrails.py @@ -228,7 +228,7 @@ class PassthroughGuardrailHandler: Dict of guardrail names to run (format: {guardrail_name: True}), or None """ from litellm.proxy.litellm_pre_call_utils import ( - _add_guardrails_from_key_or_team_metadata, + add_guardrails_from_key_or_team_metadata, ) # Normalize config to dict format (handles both list and dict) @@ -252,7 +252,7 @@ class PassthroughGuardrailHandler: # Add org/team/key level guardrails using shared helper temp_data: Final[dict[str, Any]] = {"metadata": {}} - _add_guardrails_from_key_or_team_metadata( + add_guardrails_from_key_or_team_metadata( key_metadata=user_api_key_dict.metadata, team_metadata=user_api_key_dict.team_metadata, data=temp_data, diff --git a/litellm/proxy/pass_through_endpoints/streaming_handler.py b/litellm/proxy/pass_through_endpoints/streaming_handler.py index 8199a96ceab..8f6a688fecd 100644 --- a/litellm/proxy/pass_through_endpoints/streaming_handler.py +++ b/litellm/proxy/pass_through_endpoints/streaming_handler.py @@ -147,7 +147,7 @@ class PassThroughStreamingHandler: route_streaming_logging: RouteStreamingLogging | None = None, ): resolved_route_streaming_logging: Final[RouteStreamingLogging] = ( - route_streaming_logging or PassThroughStreamingHandler._route_streaming_logging_to_handler + route_streaming_logging or PassThroughStreamingHandler.route_streaming_logging_to_handler ) raw_bytes: Final[list[bytes]] = [] resolved_request_body: Final[dict[str, object]] = request_body or {} @@ -219,7 +219,7 @@ class PassThroughStreamingHandler: PassThroughStreamingHandler._stamp_first_chunk_if_needed(litellm_logging_obj) complete_frames, pending = split_complete_sse_frames(pending + chunk) if complete_frames: - yield ProxyBaseLLMRequestProcessing._process_chunk_with_cost_injection( + yield ProxyBaseLLMRequestProcessing.process_chunk_with_cost_injection( complete_frames, resolved_model_name, litellm_logging_obj ) if pending: @@ -275,7 +275,7 @@ class PassThroughStreamingHandler: bind_budget_reservation_to_callbacks(litellm_logging_obj.litellm_params) @staticmethod - async def _route_streaming_logging_to_handler( + async def route_streaming_logging_to_handler( litellm_logging_obj: LiteLLMLoggingObj, passthrough_success_handler_obj: PassThroughEndpointLogging, url_route: str, @@ -373,6 +373,8 @@ class PassThroughStreamingHandler: except Exception as e: verbose_proxy_logger.error("Error in _route_streaming_logging_to_handler: %s", e) + _route_streaming_logging_to_handler = route_streaming_logging_to_handler + @staticmethod def _build_passthrough_logging_result( litellm_logging_obj: LiteLLMLoggingObj, @@ -397,7 +399,7 @@ class PassThroughStreamingHandler: kwargs: dict = {} if endpoint_type == EndpointType.ANTHROPIC: anthropic_passthrough_logging_handler_result: Final = ( - AnthropicPassthroughLoggingHandler._handle_logging_anthropic_collected_chunks( + AnthropicPassthroughLoggingHandler.handle_logging_anthropic_collected_chunks( litellm_logging_obj=litellm_logging_obj, passthrough_success_handler_obj=passthrough_success_handler_obj, url_route=url_route, @@ -412,7 +414,7 @@ class PassThroughStreamingHandler: kwargs = anthropic_passthrough_logging_handler_result["kwargs"] elif endpoint_type == EndpointType.VERTEX_AI: vertex_passthrough_logging_handler_result: Final = ( - VertexPassthroughLoggingHandler._handle_logging_vertex_collected_chunks( + VertexPassthroughLoggingHandler.handle_logging_vertex_collected_chunks( litellm_logging_obj=litellm_logging_obj, passthrough_success_handler_obj=passthrough_success_handler_obj, url_route=url_route, @@ -448,7 +450,7 @@ class PassThroughStreamingHandler: ) elif endpoint_type == EndpointType.OPENAI: openai_passthrough_logging_handler_result: Final = ( - OpenAIPassthroughLoggingHandler._handle_logging_openai_collected_chunks( + OpenAIPassthroughLoggingHandler.handle_logging_openai_collected_chunks( litellm_logging_obj=litellm_logging_obj, passthrough_success_handler_obj=passthrough_success_handler_obj, url_route=url_route, diff --git a/litellm/proxy/pass_through_endpoints/success_handler.py b/litellm/proxy/pass_through_endpoints/success_handler.py index 65a199d38b6..9c6c3111209 100644 --- a/litellm/proxy/pass_through_endpoints/success_handler.py +++ b/litellm/proxy/pass_through_endpoints/success_handler.py @@ -117,9 +117,9 @@ class PassThroughEndpointLogging: @property def _log_dispatch(self) -> PassThroughLogDispatch: - return self._injected_log_dispatch if self._injected_log_dispatch is not None else self._handle_logging + return self._injected_log_dispatch if self._injected_log_dispatch is not None else self.handle_logging - async def _handle_logging( + async def handle_logging( self, logging_obj: LiteLLMLoggingObj, standard_logging_response_object: StandardPassThroughResponseObject @@ -130,7 +130,7 @@ class PassThroughEndpointLogging: end_time: datetime, cache_hit: bool, **kwargs, - ): + ) -> None: """Log pass-through success via the shared async dispatch path.""" # Always reached from pass_through_async_success_handler, which runs in # an async context. call_type is "pass_through_endpoint" here, so the @@ -148,6 +148,8 @@ class PassThroughEndpointLogging: **kwargs, ) + _handle_logging = handle_logging + def normalize_llm_passthrough_logging_payload( self, httpx_response: httpx.Response, @@ -445,7 +447,7 @@ class PassThroughEndpointLogging: ) return if self.is_assemblyai_route(url_route) and not self.is_azure_speech_route(custom_llm_provider): - if AssemblyAIPassthroughLoggingHandler._should_log_request(httpx_response.request.method) is not True: + if AssemblyAIPassthroughLoggingHandler.should_log_request(httpx_response.request.method) is not True: return self.assemblyai_passthrough_logging_handler.assemblyai_passthrough_logging_handler( httpx_response=httpx_response, @@ -606,10 +608,10 @@ class PassThroughEndpointLogging: if not url_route: return False from .llm_provider_handlers.openai_passthrough_logging_handler import ( - _is_openai_compatible_url, + is_openai_compatible_url, ) - return _is_openai_compatible_url(url_route) + return is_openai_compatible_url(url_route) def is_gemini_route(self, url_route: str, custom_llm_provider: str | None = None): """Check if the URL route is a Gemini API route.""" diff --git a/litellm/proxy/policy_engine/policy_registry.py b/litellm/proxy/policy_engine/policy_registry.py index efe7a5002e0..cfc1f481988 100644 --- a/litellm/proxy/policy_engine/policy_registry.py +++ b/litellm/proxy/policy_engine/policy_registry.py @@ -195,7 +195,7 @@ class PolicyRegistry: for policy_name, policy_data in policies_config.items(): try: - policy = self._parse_policy(policy_name, policy_data) + policy = self.parse_policy(policy_name, policy_data) self._policies[policy_name] = policy verbose_proxy_logger.debug("Loaded policy: %s", policy_name) except Exception as e: @@ -207,7 +207,7 @@ class PolicyRegistry: self._initialized = True verbose_proxy_logger.info("Loaded %s policies", len(self._policies)) - def _parse_policy(self, policy_name: str, policy_data: dict[str, Any]) -> Policy: + def parse_policy(self, policy_name: str, policy_data: dict[str, Any]) -> Policy: """ Parse a policy from raw configuration data. @@ -246,6 +246,8 @@ class PolicyRegistry: pipeline=pipeline, ) + _parse_policy = parse_policy + @staticmethod def _parse_pipeline( pipeline_data: Optional["_RawPipelineConfig"], @@ -427,7 +429,7 @@ class PolicyRegistry: created_policy: Final = await _policy_table(prisma_client).create(data=data) # Also add to in-memory registry - policy: Final = self._parse_policy( + policy: Final = self.parse_policy( policy_request.policy_name, { "inherit": policy_request.inherit, @@ -648,7 +650,7 @@ class PolicyRegistry: try: production: Final = await self.get_all_policies_from_db(prisma_client, version_status="production") db_policies: Final = { - policy_response.policy_name: self._parse_policy( + policy_response.policy_name: self.parse_policy( policy_response.policy_name, { "inherit": policy_response.inherit, @@ -679,7 +681,7 @@ class PolicyRegistry: order={"created_at": "desc"}, ) for row in non_production: - policy = self._parse_policy( + policy = self.parse_policy( row.policy_name, { "inherit": row.inherit, @@ -731,7 +733,7 @@ class PolicyRegistry: # Build a temporary in-memory map for resolution temp_policies: Final = {} for policy_response in policies: - policy = self._parse_policy( + policy = self.parse_policy( policy_response.policy_name, { "inherit": policy_response.inherit, @@ -959,7 +961,7 @@ class PolicyRegistry: # Update in-memory registry: remove old production (by name), add this one self.remove_policy(policy_name) - policy: Final = self._parse_policy( + policy: Final = self.parse_policy( policy_name, { "inherit": updated.inherit, diff --git a/litellm/proxy/policy_engine/policy_validator.py b/litellm/proxy/policy_engine/policy_validator.py index 17542ba814d..36263ebd26e 100644 --- a/litellm/proxy/policy_engine/policy_validator.py +++ b/litellm/proxy/policy_engine/policy_validator.py @@ -193,7 +193,7 @@ class PolicyValidator: # A concrete entry is one the request-time matcher compares by exact equality; # only a trailing "*" is a wildcard (RouteChecks._is_wildcard_pattern), and those # are left unvalidated since they may match zero entities today and more later. - is_pattern: Final = RouteChecks._is_wildcard_pattern + is_pattern: Final = RouteChecks.is_wildcard_pattern concrete_teams: Final = [t for t in (teams or []) if not is_pattern(pattern=t)] concrete_keys: Final = [k for k in (keys or []) if not is_pattern(pattern=k)] concrete_models: Final = [m for m in (models or []) if not is_pattern(pattern=m)] @@ -429,7 +429,7 @@ class PolicyValidator: for policy_name, policy_data in policy_config.items(): try: - policy = temp_registry._parse_policy(policy_name, policy_data) + policy = temp_registry.parse_policy(policy_name, policy_data) policies[policy_name] = policy except Exception as e: errors.append( diff --git a/litellm/proxy/proxy_cli.py b/litellm/proxy/proxy_cli.py index 53a447bc6ec..c7371f294f0 100644 --- a/litellm/proxy/proxy_cli.py +++ b/litellm/proxy/proxy_cli.py @@ -204,18 +204,22 @@ def deprecated_v2_flag_passed_on_cli() -> bool: class ProxyInitializationHelpers: @staticmethod - def _echo_litellm_version(): + def echo_litellm_version(): pkg_version: Final = importlib.metadata.version("litellm") click.echo(f"\nLiteLLM: Current Version = {pkg_version}\n") + _echo_litellm_version = echo_litellm_version + @staticmethod - def _run_health_check(host, port): + def run_health_check(host, port): print("\nLiteLLM: Health Testing models in config") response: Final = httpx.get(url=f"http://{host}:{port}/health") print(json.dumps(response.json(), indent=4)) + _run_health_check = run_health_check + @staticmethod - def _run_config_validation(config: str | None) -> None: + def run_config_validation(config: str | None) -> None: if config is None: raise click.UsageError("--validate_config requires --config ") import asyncio @@ -233,8 +237,10 @@ class ProxyInitializationHelpers: raise click.exceptions.Exit(1) from error click.echo(f"LiteLLM: config OK ({model_count} models)") + _run_config_validation = run_config_validation + @staticmethod - def _run_test_chat_completion( + def run_test_chat_completion( host: str, port: int, model: str, @@ -283,8 +289,10 @@ class ProxyInitializationHelpers: ) print(completion_response) + _run_test_chat_completion = run_test_chat_completion + @staticmethod - def _get_default_unvicorn_init_args( + def get_default_unvicorn_init_args( host: str, port: int, log_config: str | None = None, @@ -328,8 +336,10 @@ class ProxyInitializationHelpers: ) return uvicorn_args + _get_default_unvicorn_init_args = get_default_unvicorn_init_args + @staticmethod - def _apply_uvicorn_max_requests_jitter( + def apply_uvicorn_max_requests_jitter( uvicorn_args: dict, max_requests_before_restart: int | None, jitter: int, @@ -356,6 +366,8 @@ class ProxyInitializationHelpers: f"Ignoring the flag.\033[0m" ) + _apply_uvicorn_max_requests_jitter = apply_uvicorn_max_requests_jitter + @staticmethod def _get_reload_options(config_path: str | None) -> dict: """Build uvicorn reload kwargs so --reload also reacts to .env and YAML edits.""" @@ -419,7 +431,7 @@ class ProxyInitializationHelpers: return True @staticmethod - def _configure_dev_reload(uvicorn_args: dict, config_path: str | None) -> None: + def configure_dev_reload(uvicorn_args: dict, config_path: str | None) -> None: """Wire up --reload (dev only): watch *.py, the --config YAML, and .env, and signal reloaded workers to re-read .env with override so edits to existing keys actually take effect rather than staying masked by the @@ -436,8 +448,10 @@ class ProxyInitializationHelpers: "to let a shell-exported value take precedence." ) + _configure_dev_reload = configure_dev_reload + @staticmethod - def _init_hypercorn_server( + def init_hypercorn_server( app: FastAPI, host: str, port: int, @@ -469,8 +483,10 @@ class ProxyInitializationHelpers: # hypercorn serve raises a type warning when passing a fast api app - even though fast API is a valid type asyncio.run(serve(app, config)) + _init_hypercorn_server = init_hypercorn_server + @staticmethod - def _init_granian_server( + def init_granian_server( host: str, port: int, num_workers: int, @@ -519,8 +535,10 @@ class ProxyInitializationHelpers: Granian(**kwargs).serve() + _init_granian_server = init_granian_server + @staticmethod - def _run_gunicorn_server( + def run_gunicorn_server( host: str, port: int, app: FastAPI, @@ -635,8 +653,10 @@ class ProxyInitializationHelpers: start_query_engine_reaper() StandaloneApplication(app=app, options=gunicorn_options).run() # Run gunicorn + _run_gunicorn_server = run_gunicorn_server + @staticmethod - def _run_ollama_serve(): + def run_ollama_serve(): try: command: Final = ["ollama", "serve"] @@ -647,20 +667,26 @@ class ProxyInitializationHelpers: LiteLLM Warning: proxy started with `ollama` model\n`ollama serve` failed with Exception{e}. \nEnsure you run `ollama serve` """) + _run_ollama_serve = run_ollama_serve + @staticmethod - def _is_port_in_use(port): + def is_port_in_use(port): import socket with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s: return s.connect_ex(("localhost", port)) == 0 + _is_port_in_use = is_port_in_use + @staticmethod - def _get_loop_type(): + def get_loop_type(): """Helper function to determine the event loop type based on platform""" if sys.platform in ("win32", "cygwin", "cli"): return None # Let uvicorn choose the default loop on Windows return "uvloop" + _get_loop_type = get_loop_type + @staticmethod def _prometheus_callback_configured(litellm_settings: Mapping[str, object] | None) -> bool: if litellm_settings is None: @@ -676,7 +702,7 @@ class ProxyInitializationHelpers: ) @staticmethod - def _maybe_setup_prometheus_multiproc_dir( + def maybe_setup_prometheus_multiproc_dir( num_workers: int, litellm_settings: dict | None, prometheus_metrics_port: int | None = None, @@ -707,6 +733,8 @@ class ProxyInitializationHelpers: print(f"LiteLLM: {action} PROMETHEUS_MULTIPROC_DIR={multiproc_dir}") return multiproc_dir + _maybe_setup_prometheus_multiproc_dir = maybe_setup_prometheus_multiproc_dir + @click.command() @click.argument("cli_args", nargs=-1) @@ -1102,10 +1130,10 @@ def run_server( except ModuleNotFoundError as e: raise ModuleNotFoundError(f"Missing dependency {e}. Run `pip install 'litellm[proxy]'`") from e if version is True: - ProxyInitializationHelpers._echo_litellm_version() + ProxyInitializationHelpers.echo_litellm_version() return if validate_config is True: - ProxyInitializationHelpers._run_config_validation(config) + ProxyInitializationHelpers.run_config_validation(config) return if enforce_prisma_migration_check: print( @@ -1114,12 +1142,12 @@ def run_server( "when database setup fails at startup. You can safely remove it.\033[0m" ) if model and "ollama" in model and api_base is None: - ProxyInitializationHelpers._run_ollama_serve() + ProxyInitializationHelpers.run_ollama_serve() if health is True: - ProxyInitializationHelpers._run_health_check(host, port) + ProxyInitializationHelpers.run_health_check(host, port) return if test is True: - ProxyInitializationHelpers._run_test_chat_completion(host, port, model, test) + ProxyInitializationHelpers.run_test_chat_completion(host, port, model, test) return else: if headers: @@ -1491,7 +1519,7 @@ def run_server( ) sys.exit(1) export_pooled_database_url(pooled_database_url) - if port == 4000 and ProxyInitializationHelpers._is_port_in_use(port): + if port == 4000 and ProxyInitializationHelpers.is_port_in_use(port): port = random.randint(1024, 49152) if prometheus_metrics_port == port: raise click.UsageError("--prometheus_metrics_port must differ from --port") @@ -1512,7 +1540,7 @@ def run_server( return # Auto-create PROMETHEUS_MULTIPROC_DIR for multi-worker setups - prometheus_multiproc_dir: Final = ProxyInitializationHelpers._maybe_setup_prometheus_multiproc_dir( + prometheus_multiproc_dir: Final = ProxyInitializationHelpers.maybe_setup_prometheus_multiproc_dir( num_workers=num_workers, litellm_settings=litellm_settings if config else None, prometheus_metrics_port=prometheus_metrics_port, @@ -1533,7 +1561,7 @@ def run_server( ) running_uvicorn: Final = run_gunicorn is False and run_hypercorn is False - uvicorn_args: Final = ProxyInitializationHelpers._get_default_unvicorn_init_args( + uvicorn_args: Final = ProxyInitializationHelpers.get_default_unvicorn_init_args( host=host, port=port, log_config=log_config, @@ -1547,7 +1575,7 @@ def run_server( if limit_concurrency is not None: uvicorn_args["limit_concurrency"] = limit_concurrency if max_requests_before_restart_jitter is not None: - ProxyInitializationHelpers._apply_uvicorn_max_requests_jitter( + ProxyInitializationHelpers.apply_uvicorn_max_requests_jitter( uvicorn_args=uvicorn_args, max_requests_before_restart=max_requests_before_restart, jitter=max_requests_before_restart_jitter, @@ -1559,12 +1587,12 @@ def run_server( uvicorn_args["ssl_keyfile"] = ssl_keyfile_path uvicorn_args["ssl_certfile"] = ssl_certfile_path - loop_type: Final = ProxyInitializationHelpers._get_loop_type() + loop_type: Final = ProxyInitializationHelpers.get_loop_type() if loop_type: uvicorn_args["loop"] = loop_type if reload: - ProxyInitializationHelpers._configure_dev_reload(uvicorn_args, config) + ProxyInitializationHelpers.configure_dev_reload(uvicorn_args, config) if num_workers > 1: start_query_engine_reaper() @@ -1573,7 +1601,7 @@ def run_server( workers=num_workers, ) elif run_gunicorn is True: - ProxyInitializationHelpers._run_gunicorn_server( + ProxyInitializationHelpers.run_gunicorn_server( host=host, port=port, app=app, @@ -1584,7 +1612,7 @@ def run_server( max_requests_before_restart_jitter=max_requests_before_restart_jitter, ) elif run_hypercorn is True: - ProxyInitializationHelpers._init_hypercorn_server( + ProxyInitializationHelpers.init_hypercorn_server( app=app, host=host, port=port, @@ -1593,7 +1621,7 @@ def run_server( ciphers=ciphers, ) elif run_granian is True: - ProxyInitializationHelpers._init_granian_server( + ProxyInitializationHelpers.init_granian_server( host=host, port=port, num_workers=num_workers, diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index b6ce32bbd24..06db63350e2 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -133,7 +133,10 @@ from litellm.proxy.common_utils.callback_utils import ( process_callback, strip_callback_config, ) -from litellm.proxy.common_utils.realtime_utils import _realtime_request_body +from litellm.proxy.common_utils.realtime_utils import ( # noqa: F401, RUF100 # legacy module exports + _realtime_request_body, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + realtime_request_body, +) from litellm.proxy.management_helpers.auto_router_availability import AutoRouterCatalogEntry, build_auto_router_catalog from litellm.router_utils.access_windows import access_windows_config_error from litellm.router_utils.add_retry_fallback_headers import ( @@ -384,8 +387,9 @@ from litellm.proxy.auth.model_checks import ( get_team_models, ) from litellm.proxy.auth.password_policy import validate_password_not_breached, validate_password_policy -from litellm.proxy.auth.user_api_key_auth import ( - _fetch_global_spend_with_event_coordination, +from litellm.proxy.auth.user_api_key_auth import ( # noqa: F401, RUF100 # legacy module exports + _fetch_global_spend_with_event_coordination, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + fetch_global_spend_with_event_coordination, user_api_key_auth, user_api_key_auth_websocket, ) @@ -394,17 +398,19 @@ from litellm.proxy.bug_report_config import build_proxy_bug_report ## Import All Misc routes here ## from litellm.proxy.caching_routes import router as caching_router -from litellm.proxy.common_request_processing import ( +from litellm.proxy.common_request_processing import ( # noqa: F401, RUF100 # legacy module exports KNOWN_PROXY_ROUTES, ProxyBaseLLMRequestProcessing, - _is_azure_model_router_request, - _should_return_raw_model_name, + _is_azure_model_router_request, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _should_return_raw_model_name, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export close_guarded_stream, create_response, + is_azure_model_router_request, log_llm_api_exception, open_sse_before_first_byte, request_litellm_call_id, resolve_litellm_call_id, + should_return_raw_model_name, ttft_keepalive_interval, ) from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import ( @@ -437,12 +443,14 @@ from litellm.proxy.common_utils.healthy_model_filter import ( ) from litellm.proxy.common_utils.html_forms.default_credentials_hint import should_hide_default_credentials_hint from litellm.proxy.common_utils.html_forms.ui_login import build_ui_login_form -from litellm.proxy.common_utils.http_parsing_utils import ( - _read_request_body, - _safe_get_request_headers, +from litellm.proxy.common_utils.http_parsing_utils import ( # noqa: F401, RUF100 # legacy module exports + _read_request_body, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _safe_get_request_headers, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export check_file_size_under_limit, get_form_data, + read_request_body, resolve_inference_model, + safe_get_request_headers, ) from litellm.proxy.common_utils.load_config_utils import get_config_from_bucket from litellm.proxy.common_utils.model_deprecation import collect_model_deprecations @@ -570,17 +578,24 @@ from litellm.proxy.health_check import ( perform_health_check, ) from litellm.proxy.health_endpoints._health_endpoints import router as health_router -from litellm.proxy.hooks.model_max_budget_limiter import ( - _PROXY_VirtualKeyModelMaxBudgetLimiter, +from litellm.proxy.hooks.model_max_budget_limiter import ( # noqa: F401, RUF100 # legacy module exports + PROXY_VirtualKeyModelMaxBudgetLimiter, + _PROXY_VirtualKeyModelMaxBudgetLimiter, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export ) from litellm.proxy.hooks.parallel_request_limiter_v3 import fail_closed_rate_limit_enforcement_enabled -from litellm.proxy.hooks.prompt_injection_detection import ( - _OPTIONAL_PromptInjectionDetection, +from litellm.proxy.hooks.prompt_injection_detection import ( # noqa: F401, RUF100 # legacy module exports + OPTIONAL_PromptInjectionDetection, + _OPTIONAL_PromptInjectionDetection, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export +) +from litellm.proxy.hooks.proxy_track_cost_callback import ( # noqa: F401, RUF100 # legacy module exports + ProxyDBLogger, + _ProxyDBLogger, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + run_spend_event, ) -from litellm.proxy.hooks.proxy_track_cost_callback import _ProxyDBLogger, run_spend_event from litellm.proxy.image_endpoints.endpoints import router as image_router from litellm.proxy.lens.dataset_endpoints import router as lens_dataset_router from litellm.proxy.lens.endpoints import router as lens_router +from litellm.proxy.lens.feedback_endpoints import router as lens_feedback_router from litellm.proxy.lens.repository import WriterDatabase from litellm.proxy.lens.signal_repository import SignalRepository from litellm.proxy.lens.signals import ( @@ -612,10 +627,12 @@ from litellm.proxy.management_endpoints.cache_settings_endpoints import ( from litellm.proxy.management_endpoints.callback_management_endpoints import ( router as callback_management_endpoints_router, ) -from litellm.proxy.management_endpoints.common_utils import ( - _user_has_admin_privileges, - _user_has_admin_view, +from litellm.proxy.management_endpoints.common_utils import ( # noqa: F401, RUF100 # legacy module exports + _user_has_admin_privileges, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _user_has_admin_view, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export admin_can_invite_user, + user_api_key_has_admin_view, + user_has_admin_privileges, ) from litellm.proxy.management_endpoints.coordination_redis_endpoints import ( get_persisted_coordination_redis_settings, @@ -658,10 +675,13 @@ from litellm.proxy.management_endpoints.management_v1.common import MANAGEMENT_V from litellm.proxy.management_endpoints.model_access_group_management_endpoints import ( router as model_access_group_management_router, ) -from litellm.proxy.management_endpoints.model_management_endpoints import ( - _add_model_to_db, - _add_team_model_to_db, - _deduplicate_litellm_router_models, +from litellm.proxy.management_endpoints.model_management_endpoints import ( # noqa: F401, RUF100 # legacy module exports + _add_model_to_db, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _add_team_model_to_db, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _deduplicate_litellm_router_models, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + add_model_to_db, + add_team_model_to_db, + deduplicate_litellm_router_models, live_model_ids_snapshot, ) from litellm.proxy.management_endpoints.model_management_endpoints import ( @@ -770,6 +790,7 @@ from litellm.proxy.middleware.request_size_limit_middleware import ( from litellm.proxy.middleware.security_headers_middleware import ( SecurityHeadersMiddleware, ) +from litellm.proxy.moyai_endpoints import router as moyai_router from litellm.proxy.ocr_endpoints.endpoints import router as ocr_router from litellm.proxy.openai_files_endpoints.files_endpoints import ( router as openai_files_router, @@ -840,26 +861,33 @@ from litellm.proxy.ui_crud_endpoints.proxy_setting_endpoints import ( from litellm.proxy.ui_crud_endpoints.user_banner_endpoints import ( router as user_banner_endpoints_router, ) -from litellm.proxy.utils import ( +from litellm.proxy.utils import ( # noqa: F401, RUF100 # legacy module exports PrismaClient, ProxyLogging, ProxyUpdateSpend, - _cache_user_row, - _get_docs_url, - _get_openapi_url, - _get_projected_spend_over_limit, - _get_redoc_url, - _is_projected_spend_over_limit, - _is_valid_team_configs, + _cache_user_row, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _get_docs_url, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _get_openapi_url, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _get_projected_spend_over_limit, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _get_redoc_url, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _is_projected_spend_over_limit, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _is_valid_team_configs, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + cache_user_row, evict_config_param, get_config_param, get_custom_url, + get_docs_url, get_error_message_str, + get_openapi_url, + get_projected_spend_over_limit, + get_redoc_url, get_server_root_path, handle_exception_on_proxy, hash_password, hash_token, invalidate_config_param, + is_projected_spend_over_limit, + is_valid_team_configs, litellm_config_cache, migrate_passwords_to_scrypt_async, model_dump_with_preserved_fields, @@ -1071,7 +1099,9 @@ custom_swagger_message: Final = ( ) ### CUSTOM BRANDING [ENTERPRISE FEATURE] ### -_title: Final = os.getenv("DOCS_TITLE", "LiteLLM API") if premium_user else "LiteLLM API" +title: Final = os.getenv("DOCS_TITLE", "LiteLLM API") if premium_user else "LiteLLM API" + +_title: Final = title _description: Final = ( os.getenv( "DOCS_DESCRIPTION", @@ -1237,7 +1267,7 @@ class _AiohttpConnectorKwargs(TypedDict, total=False): socket_factory: Callable[[_AiohttpAddrInfo], socket.socket] -async def _initialize_shared_aiohttp_session(): +async def initialize_shared_aiohttp_session() -> "ClientSession | None": """Initialize shared aiohttp session for connection reuse with connection limits.""" try: from aiohttp import ClientSession, DummyCookieJar, TCPConnector @@ -1277,6 +1307,9 @@ async def _initialize_shared_aiohttp_session(): return None +_initialize_shared_aiohttp_session: Final = initialize_shared_aiohttp_session + + async def _connect_to_count_stored_values() -> SupportsRawQueries: client: Final = prisma_client or PrismaClient( database_url=str(get_secret("DATABASE_URL")), proxy_logging_obj=proxy_logging_obj @@ -1427,10 +1460,12 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[ProxyLifespanState # check if DATABASE_URL in environment - load from there if prisma_client is None: _db_url: Final[str | None] = get_secret("DATABASE_URL", None) - prisma_client = await ProxyStartupEvent._setup_prisma_client( - database_url=_db_url, - proxy_logging_obj=proxy_logging_obj, - user_api_key_cache=user_api_key_cache, + prisma_client = ( # rebind-ok: pre-existing rebinding on a rename-only line + await ProxyStartupEvent.setup_prisma_client( + database_url=_db_url, + proxy_logging_obj=proxy_logging_obj, + user_api_key_cache=user_api_key_cache, + ) ) await migrate_if_requested( @@ -1498,7 +1533,7 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[ProxyLifespanState ## A coordination_redis block saved from the admin UI lives in the database, ## which is only reachable once the prisma client exists. Apply it here, before ## the coordination Redis is published to its consumers below. - db_coordination_redis_cache: Final = await ProxyStartupEvent._init_coordination_redis_from_db( + db_coordination_redis_cache: Final = await ProxyStartupEvent.init_coordination_redis_from_db( litellm_settings=proxy_config.get_config_state().get("litellm_settings") or {}, llm_router=llm_router, ) @@ -1509,11 +1544,11 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[ProxyLifespanState ## when the proxy cache backend is not Redis ## transaction_buffer_redis_cache = redis_usage_cache if transaction_buffer_redis_cache is None: - transaction_buffer_redis_cache = ProxyStartupEvent._get_transaction_buffer_redis_cache( + transaction_buffer_redis_cache = ProxyStartupEvent.get_transaction_buffer_redis_cache( # rebind-ok: pre-existing rebinding on a rename-only line general_settings=general_settings ) - ProxyStartupEvent._initialize_startup_logging( + ProxyStartupEvent.initialize_startup_logging( llm_router=llm_router, proxy_logging_obj=proxy_logging_obj, redis_usage_cache=transaction_buffer_redis_cache, @@ -1552,12 +1587,12 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[ProxyLifespanState verbose_proxy_logger.debug("Skipping OTel V2 provider setup: %s", e) ## Validate use_redis_transaction_buffer requires Redis cache ## - ProxyStartupEvent._validate_redis_transaction_buffer_config( + ProxyStartupEvent.validate_redis_transaction_buffer_config( general_settings=general_settings, redis_usage_cache=transaction_buffer_redis_cache, ) - ProxyStartupEvent._warn_if_mock_testing_params_enabled(general_settings=general_settings) + ProxyStartupEvent.warn_if_mock_testing_params_enabled(general_settings=general_settings) ## SEMANTIC TOOL FILTER ## # Read litellm_settings from config for semantic filter initialization @@ -1566,7 +1601,7 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[ProxyLifespanState _config: Final = proxy_config.get_config_state() _litellm_settings: Final = _config.get("litellm_settings", {}) verbose_proxy_logger.debug("litellm_settings keys = %s", list(_litellm_settings.keys())) - await ProxyStartupEvent._initialize_semantic_tool_filter( + await ProxyStartupEvent.initialize_semantic_tool_filter( llm_router=llm_router, litellm_settings=_litellm_settings, ) @@ -1575,28 +1610,28 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[ProxyLifespanState verbose_proxy_logger.error("Semantic filter init failed: %s", e, exc_info=True) ## JWT AUTH ## - ProxyStartupEvent._initialize_jwt_auth( + ProxyStartupEvent.initialize_jwt_auth( general_settings=general_settings, prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, ) - ProxyStartupEvent._attach_router_to_prompt_injection_detectors(llm_router=llm_router) + ProxyStartupEvent.attach_router_to_prompt_injection_detectors(llm_router=llm_router) verbose_proxy_logger.debug("prisma_client: %s", prisma_client) if prisma_client is not None and litellm.max_budget > 0: - ProxyStartupEvent._add_proxy_budget_to_db() + ProxyStartupEvent.add_proxy_budget_to_db() asyncio.create_task( - ProxyStartupEvent._warm_global_spend_cache( + ProxyStartupEvent.warm_global_spend_cache( user_api_key_cache=user_api_key_cache, prisma_client=prisma_client, ) ) - ProxyStartupEvent._warn_budget_without_db( + ProxyStartupEvent.warn_budget_without_db( max_budget=litellm.max_budget, prisma_client=prisma_client, ) - ProxyStartupEvent._warn_fail_closed_rate_limits_without_redis( + ProxyStartupEvent.warn_fail_closed_rate_limits_without_redis( fail_closed_rate_limit_enforcement=fail_closed_rate_limit_enforcement_enabled(general_settings), redis_usage_cache=redis_usage_cache, ) @@ -1615,10 +1650,10 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[ProxyLifespanState else None ) if prisma_client is not None: - await ProxyStartupEvent._update_default_team_member_budget() + await ProxyStartupEvent.update_default_team_member_budget() ## SYNC UI SETTINGS ## - await ProxyStartupEvent._sync_ui_settings_to_general_settings() + await ProxyStartupEvent.sync_ui_settings_to_general_settings() # Start background health checks AFTER models are loaded and index is built if use_background_health_checks: @@ -1637,13 +1672,15 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[ProxyLifespanState asyncio.create_task(_adaptive_router_flusher_loop()) ## [Optional] Initialize dd tracer - ProxyStartupEvent._init_dd_tracer() + ProxyStartupEvent.init_dd_tracer() ## [Optional] Initialize Pyroscope continuous profiling (env: LITELLM_ENABLE_PYROSCOPE=true) - ProxyStartupEvent._init_pyroscope() + ProxyStartupEvent.init_pyroscope() ## Initialize shared aiohttp session for connection reuse - shared_aiohttp_session = await _initialize_shared_aiohttp_session() + shared_aiohttp_session = ( # rebind-ok: pre-existing rebinding on a rename-only line + await initialize_shared_aiohttp_session() + ) model_info_refresh_disabled: Final = ( "disable_model_info_refresh" in general_settings and general_settings["disable_model_info_refresh"] is True @@ -1857,10 +1894,10 @@ def ensure_unique_openapi_operation_ids( app = FastAPI( - docs_url=_get_docs_url(), - redoc_url=_get_redoc_url(), - openapi_url=_get_openapi_url(), - title=_title, + docs_url=get_docs_url(), + redoc_url=get_redoc_url(), + openapi_url=get_openapi_url(), + title=title, description=_description, version=version, root_path=server_root_path, @@ -2688,7 +2725,7 @@ def mount_swagger_ui(): mount_swagger_ui() -docs_url: Final = _get_docs_url() +docs_url: Final = get_docs_url() root_redirect_url: Final[str | None] = os.getenv("ROOT_REDIRECT_URL") if docs_url != "/" and root_redirect_url is not None: @@ -2741,7 +2778,7 @@ user_api_key_cache: UserApiKeyCache = UserApiKeyCache( ) spend_counter_cache: Final = DualCache(default_in_memory_ttl=UserAPIKeyCacheTTLEnum.in_memory_cache_ttl.value) cli_sso_session_cache: Final = DualCache(default_in_memory_ttl=CLI_SSO_SESSION_TTL_SECONDS) -model_max_budget_limiter: Final = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=spend_counter_cache) +model_max_budget_limiter: Final = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=spend_counter_cache) litellm.logging_callback_manager.add_litellm_callback(model_max_budget_limiter) redis_usage_cache: RedisCache | None = None # redis cache used for tracking spend, tpm/rpm limits polling_via_cache_enabled: Literal["all"] | list[str] | bool = False @@ -2778,7 +2815,7 @@ proxy_config_reload_interval_seconds = PROXY_CONFIG_RELOAD_INTERVAL_SECONDS litellm_master_key_hash = None disable_spend_logs = False jwt_handler: Final = JWTHandler() -prompt_injection_detection_obj: _OPTIONAL_PromptInjectionDetection | None = None +prompt_injection_detection_obj: Final[OPTIONAL_PromptInjectionDetection | None] = None store_model_in_db: bool = False open_telemetry_logger: OpenTelemetry | None = None ### GATEWAY REQUEST COUNTS (SGR) ### @@ -2793,7 +2830,7 @@ def _gateway_request_redis_buffer() -> GatewayRequestRedisBuffer | None: """Shares the spend writer's transaction-buffer Redis and pod lock when use_redis_transaction_buffer is on.""" writer: Final = proxy_logging_obj.db_spend_update_writer redis_cache: Final = writer.redis_update_buffer.redis_cache - if redis_cache is None or not writer.redis_update_buffer._should_commit_spend_updates_to_redis(): + if redis_cache is None or not writer.redis_update_buffer.should_commit_spend_updates_to_redis(): return None return GatewayRequestRedisBuffer(redis_cache=redis_cache, pod_lock_manager=writer.pod_lock_manager) @@ -2882,8 +2919,8 @@ def cost_tracking(): from litellm.integrations.shadow_eval_logger import ShadowEvalLogger spend_event_producer = build_spend_event_producer(CollectorSettings(), fallback=run_spend_event) - litellm.logging_callback_manager.add_litellm_callback(_ProxyDBLogger(spend_event_producer)) - litellm.logging_callback_manager.add_litellm_async_success_callback(_ProxyDBLogger(spend_event_producer)) + litellm.logging_callback_manager.add_litellm_callback(ProxyDBLogger(spend_event_producer)) + litellm.logging_callback_manager.add_litellm_async_success_callback(ProxyDBLogger(spend_event_producer)) litellm.logging_callback_manager.add_litellm_callback(ShadowEvalLogger()) @@ -3665,7 +3702,7 @@ async def _prepare_spend_counter_increment( 4. Increment is returned for the caller to apply via pipeline """ with service_target(SPEND_COUNTERS_TARGET): - await _ensure_spend_counter_initialized( + await ensure_spend_counter_initialized( counter_key=counter_key, source_cache_key=source_cache_key, ) @@ -3737,7 +3774,7 @@ async def _prepare_window_spend_counter_increment( return None with service_target(SPEND_COUNTERS_TARGET): - initialized: Final = await _ensure_window_spend_counter_initialized( + initialized: Final = await ensure_window_spend_counter_initialized( counter_key=counter_key, entity_type=entity_type, entity_id=entity_id, @@ -3749,10 +3786,10 @@ async def _prepare_window_spend_counter_increment( return PendingSpendIncrement(counter_key=counter_key, increment=increment) -async def _ensure_spend_counter_initialized( +async def ensure_spend_counter_initialized( counter_key: str, source_cache_key: str | list[str], -): +) -> None: is_warm: Final = await _is_spend_counter_cache_warm(counter_key=counter_key) if is_warm is False: # Shares the per-counter lock with get_current_spend. @@ -3766,7 +3803,10 @@ async def _ensure_spend_counter_initialized( # DB unavailable - fall back to in-process cache (may be stale). base_spend: Final = await _get_source_cache_base_spend(source_cache_key=source_cache_key) if base_spend > 0: - await _increment_spend_counter_cache(counter_key=counter_key, increment=base_spend) + await increment_spend_counter_cache(counter_key=counter_key, increment=base_spend) + + +_ensure_spend_counter_initialized: Final = ensure_spend_counter_initialized async def _get_source_cache_base_spend( @@ -3783,7 +3823,7 @@ async def _get_source_cache_base_spend( return 0.0 -async def _ensure_window_spend_counter_initialized( +async def ensure_window_spend_counter_initialized( counter_key: str, entity_type: str, entity_id: str, @@ -3812,6 +3852,9 @@ async def _ensure_window_spend_counter_initialized( return True +_ensure_window_spend_counter_initialized: Final = ensure_window_spend_counter_initialized + + @with_service_target(SPEND_COUNTERS_TARGET) async def _is_spend_counter_cache_warm(counter_key: str) -> bool: batched: Final = await read_batched_spend_counter(counter_key) @@ -3848,7 +3891,7 @@ async def increment_spend_counter(counter_key: str, increment: float): """Public raw-counter increment for budget domains outside the entity scopes (e.g. shadow eval's per-leg spend), sharing the primitive the entity counters use so invalidation and read semantics can never drift.""" - return await _increment_spend_counter_cache(counter_key=counter_key, increment=increment) + return await increment_spend_counter_cache(counter_key=counter_key, increment=increment) @with_service_target(SPEND_COUNTERS_TARGET) @@ -3863,7 +3906,7 @@ async def refresh_spend_counter_ttl(counter_key: str) -> bool: @with_service_target(SPEND_COUNTERS_TARGET) -async def _increment_spend_counter_cache(counter_key: str, increment: float): +async def increment_spend_counter_cache(counter_key: str, increment: float) -> float | None: if spend_counter_cache.redis_cache is not None: try: current_value: Final = await spend_counter_cache.redis_cache.async_increment( @@ -3872,7 +3915,7 @@ async def _increment_spend_counter_cache(counter_key: str, increment: float): refresh_ttl=True, ) except Exception: - await _invalidate_spend_counter(counter_key=counter_key) + await invalidate_spend_counter(counter_key=counter_key) raise spend_counter_cache.in_memory_cache.set_cache( key=counter_key, @@ -3886,8 +3929,11 @@ async def _increment_spend_counter_cache(counter_key: str, increment: float): ) +_increment_spend_counter_cache: Final = increment_spend_counter_cache + + @with_service_target(SPEND_COUNTERS_TARGET) -async def _invalidate_spend_counter(counter_key: str): +async def invalidate_spend_counter(counter_key: str) -> None: forget_spend_counter(counter_key) spend_counter_cache.in_memory_cache.delete_cache(key=counter_key) if spend_counter_cache.redis_cache is not None: @@ -3901,6 +3947,9 @@ async def _invalidate_spend_counter(counter_key: str): ) +_invalidate_spend_counter: Final = invalidate_spend_counter + + async def _apply_spend_counter_increments(pending: Sequence[PendingSpendIncrement]) -> None: if _defer_spend_counter_increments(pending): return @@ -3943,7 +3992,7 @@ def _settle_spend_counter_increment(item: PendingSpendIncrement) -> Callable[[as verbose_proxy_logger.warning( "Spend counter %s increment did not land in the post-call pipeline; invalidating it", item.counter_key ) - await _invalidate_spend_counter(counter_key=item.counter_key) + await invalidate_spend_counter(counter_key=item.counter_key) return settle @@ -3956,7 +4005,7 @@ async def increment_spend_counters_pipeline(pending: Sequence[PendingSpendIncrem try: return await run_spend_counter_pipeline(pending=pending) except Exception: - await asyncio.gather(*(_invalidate_spend_counter(counter_key=item.counter_key) for item in pending)) + await asyncio.gather(*(invalidate_spend_counter(counter_key=item.counter_key) for item in pending)) raise @@ -4105,14 +4154,14 @@ async def update_cache( existing_spend_obj.soft_budget_cooldown is False and existing_spend_obj.soft_budget is not None and ( - _is_projected_spend_over_limit( + is_projected_spend_over_limit( current_spend=new_spend, soft_budget_limit=existing_spend_obj.soft_budget, ) is True ) ): - projected_spend, projected_exceeded_date = _get_projected_spend_over_limit( + projected_spend, projected_exceeded_date = get_projected_spend_over_limit( current_spend=new_spend, soft_budget_limit=existing_spend_obj.soft_budget, ) @@ -4490,13 +4539,13 @@ def _schedule_background_health_check_db_save( import time as time_module from litellm.proxy.health_endpoints._health_endpoints import ( - _save_background_health_checks_to_db, + save_background_health_checks_to_db, ) checked_by: Final = shared_health_manager.pod_id if shared_health_manager is not None else "background_health_check" start_time: Final = time_module.time() save: Final = partial( - _save_background_health_checks_to_db, + save_background_health_checks_to_db, prisma_client, model_list, healthy_endpoints, @@ -5015,7 +5064,7 @@ def _resolve_coordination_redis_env_refs(raw_params: Mapping[str, object]) -> di } -def _build_redis_usage_cache(redis_params: Mapping[str, object]) -> RedisCache: +def build_redis_usage_cache(redis_params: Mapping[str, object]) -> RedisCache: """ Builds the proxy's coordination Redis client from resolved connection params. Cluster-mode targets (explicit `startup_nodes` or the @@ -5034,7 +5083,10 @@ def _build_redis_usage_cache(redis_params: Mapping[str, object]) -> RedisCache: return RedisCache(**non_node_params) -def _environment_has_redis_connection_target() -> bool: +_build_redis_usage_cache: Final = build_redis_usage_cache + + +def environment_has_redis_connection_target() -> bool: """ Whether the REDIS_* environment variables name a Redis to connect to (host, url, cluster nodes, or sentinel nodes). Read-only: callers that only need to @@ -5050,6 +5102,9 @@ def _environment_has_redis_connection_target() -> bool: ) +_environment_has_redis_connection_target: Final = environment_has_redis_connection_target + + def _build_redis_usage_cache_from_environment() -> RedisCache | None: """ Builds a standalone coordination Redis from REDIS_* environment variables. @@ -5061,9 +5116,9 @@ def _build_redis_usage_cache_from_environment() -> RedisCache | None: Returns None when the environment carries no connection target (host, url, cluster nodes, or sentinel nodes). """ - if not _environment_has_redis_connection_target(): + if not environment_has_redis_connection_target(): return None - return _build_redis_usage_cache(litellm._redis._redis_kwargs_from_environment()) + return build_redis_usage_cache(litellm._redis._redis_kwargs_from_environment()) def _attach_redis_usage_cache(redis_cache: RedisCache, enable_redis_auth_cache: bool) -> None: @@ -5713,7 +5768,7 @@ class ProxyConfig: environment_variables: Final = new_config.get("environment_variables") if include_env_vars and environment_variables is not None: encrypted_environment_variables: Final = ( - self._encrypt_env_variables_for_db(environment_variables=environment_variables) + self.encrypt_env_variables_for_db(environment_variables=environment_variables) if isinstance(environment_variables, dict) and environment_variables else environment_variables ) @@ -5874,7 +5929,7 @@ class ProxyConfig: existing: Final[dict] = dict(row.param_value) if row is not None and row.param_value is not None else {} to_set: Final = {k: v for k, v in updates.items() if v is not None} - encrypted: Final = self._encrypt_env_variables_for_db(environment_variables=to_set) if to_set else {} + encrypted: Final = self.encrypt_env_variables_for_db(environment_variables=to_set) if to_set else {} deleted_keys: Final = {k for k, v in updates.items() if v is None} merged: Final = {**{k: v for k, v in existing.items() if k not in deleted_keys}, **encrypted} @@ -6027,7 +6082,7 @@ class ProxyConfig: "set one of host, url, startup_nodes, or sentinel_nodes" ) - coordination_redis_cache: Final = _build_redis_usage_cache(coordination_params.model_dump(exclude_none=True)) + coordination_redis_cache: Final = build_redis_usage_cache(coordination_params.model_dump(exclude_none=True)) _attach_redis_usage_cache( coordination_redis_cache, enable_redis_auth_cache=litellm_settings.get("enable_redis_auth_cache", False) is True, @@ -6095,7 +6150,7 @@ class ProxyConfig: ) return env_coordination_redis_cache - def _init_cache( + def init_cache( self, cache_params: dict, enable_redis_auth_cache: bool = False, @@ -6141,6 +6196,8 @@ class ProxyConfig: verbose_proxy_logger.info("litellm_config_cache: no Redis configured; cluster-wide cache sharing disabled.") return resolved_usage_cache + _init_cache = init_cache + def switch_on_llm_response_caching(self): """ Enable caching on the router by setting cache_responses=True. @@ -6511,7 +6568,7 @@ class ProxyConfig: ## to pass a complete url, or set ssl=True, etc. just set it as `os.environ[REDIS_URL] = `, _redis.py checks for REDIS specific environment variables _set_redis_usage_cache( - self._init_cache( + self.init_cache( cache_params=cache_params, enable_redis_auth_cache=litellm_settings.get("enable_redis_auth_cache", False) is True, ) @@ -7730,7 +7787,7 @@ class ProxyConfig: displaced=previous.displaced + _entries_missing_from(before, after), ) - def _encrypt_env_variables(self, environment_variables: dict, new_encryption_key: str | None = None) -> dict: + def encrypt_env_variables(self, environment_variables: dict, new_encryption_key: str | None = None) -> dict: """ Encrypts a dictionary of environment variables and returns them. """ @@ -7740,7 +7797,9 @@ class ProxyConfig: encrypted_env_vars[k] = encrypted_value return encrypted_env_vars - def _decrypt_and_set_db_env_variables( + _encrypt_env_variables = encrypt_env_variables + + def decrypt_and_set_db_env_variables( self, environment_variables: dict, return_original_value: bool = False ) -> dict: """ @@ -7769,7 +7828,9 @@ class ProxyConfig: verbose_proxy_logger.error("Error setting env variable: %s - %s", k, str(e)) return decrypted_env_vars - def _decrypt_db_variables(self, variables_dict: dict) -> dict: + _decrypt_and_set_db_env_variables = decrypt_and_set_db_env_variables + + def decrypt_db_variables(self, variables_dict: dict) -> dict: """ Decrypts a dictionary of variables and returns them. """ @@ -7779,7 +7840,9 @@ class ProxyConfig: decrypted_variables[k] = decrypted_value return decrypted_variables - def _encrypt_env_variables_for_db(self, environment_variables: dict, new_encryption_key: str | None = None) -> dict: + _decrypt_db_variables = decrypt_db_variables + + def encrypt_env_variables_for_db(self, environment_variables: dict, new_encryption_key: str | None = None) -> dict: """ Idempotently encrypt environment variables for a DB write. @@ -7794,12 +7857,14 @@ class ProxyConfig: _decrypt_and_set_db_env_variables): this is a write path, and loading values into os.environ is the read path's responsibility. """ - decrypted_env_vars: Final = self._decrypt_db_variables(environment_variables) - return self._encrypt_env_variables( + decrypted_env_vars: Final = self.decrypt_db_variables(environment_variables) + return self.encrypt_env_variables( environment_variables=decrypted_env_vars, new_encryption_key=new_encryption_key, ) + _encrypt_env_variables_for_db = encrypt_env_variables_for_db + @staticmethod def _parse_router_settings_value(value: object) -> dict | None: """ @@ -8054,7 +8119,7 @@ class ProxyConfig: async def _apply_pass_through_settings(self, db_values: Mapping[str, SettingsJsonValue]) -> None: db_endpoints: Final = db_values.get("pass_through_endpoints") if isinstance(db_endpoints, list): - await self._serve_pass_through_endpoints(db_endpoints) + await self.serve_pass_through_endpoints(db_endpoints) return if "pass_through_endpoints" not in self.settings: self._publish_pass_through_endpoints(()) @@ -8064,10 +8129,12 @@ class ProxyConfig: list(db_endpoints), config_passthrough_endpoints ) - async def _serve_pass_through_endpoints(self, db_endpoints: Sequence[SettingsJsonValue]) -> None: + async def serve_pass_through_endpoints(self, db_endpoints: Sequence[SettingsJsonValue]) -> None: self._publish_pass_through_endpoints(db_endpoints) await initialize_pass_through_endpoints(pass_through_endpoints=_ENDPOINT_DICTS.validate_python(db_endpoints)) + _serve_pass_through_endpoints = serve_pass_through_endpoints + async def _apply_boolean_settings(self, db_values: Mapping[str, SettingsJsonValue]) -> None: for key in ( "store_prompts_in_spend_logs", @@ -8199,7 +8266,7 @@ class ProxyConfig: def _prepared_db_settings_values(self, section: Section, value: object) -> Mapping[str, SettingsJsonValue]: if section == "environment_variables": - decrypted: Final = self._decrypt_and_set_db_env_variables( + decrypted: Final = self.decrypt_and_set_db_env_variables( dict(_as_settings_mapping(value)), return_original_value=True ) normalized: Final = { @@ -8228,7 +8295,7 @@ class ProxyConfig: row for row in self.auto_router_db_catalog if row.model_id not in model_ids ) - async def _get_models_from_db(self, prisma_client: PrismaClient) -> Sequence[_ProxyModelRow] | None: + async def get_models_from_db(self, prisma_client: PrismaClient) -> Sequence[_ProxyModelRow] | None: """ Fetch all model deployments from the DB. @@ -8256,6 +8323,8 @@ class ProxyConfig: ) return None + _get_models_from_db = get_models_from_db + async def add_deployment( self, prisma_client: PrismaClient, @@ -8288,9 +8357,9 @@ class ProxyConfig: await sync_ui_settings_to_general_settings(prisma_client) async with MODEL_RECONCILE_LOCK: - return await self._add_deployment_locked(prisma_client=prisma_client, proxy_logging_obj=proxy_logging_obj) + return await self.add_deployment_locked(prisma_client=prisma_client, proxy_logging_obj=proxy_logging_obj) - async def _add_deployment_locked( + async def add_deployment_locked( self, prisma_client: PrismaClient, proxy_logging_obj: ProxyLogging, @@ -8317,7 +8386,7 @@ class ProxyConfig: ) load_models: Final = self._should_load_db_object(object_type="models") - new_models: Final = await self._get_models_from_db(prisma_client=prisma_client) if load_models else None + new_models: Final = await self.get_models_from_db(prisma_client=prisma_client) if load_models else None await self.get_credentials(prisma_client=prisma_client) if load_models: still_desired_ids = await self._update_llm_router( @@ -8347,6 +8416,8 @@ class ProxyConfig: live_after=None if still_desired_ids is None else live_model_ids_snapshot(), ) + _add_deployment_locked = add_deployment_locked + def start_config_sync_subscriber( self, prisma_client: PrismaClient, @@ -8450,7 +8521,7 @@ class ProxyConfig: await CacheSettingsManager.init_cache_settings_in_db(prisma_client=prisma_client, proxy_config=self) if self._should_load_db_object(object_type="semantic_filter_settings"): - await self._init_semantic_filter_settings_in_db(prisma_client=prisma_client) + await self.init_semantic_filter_settings_in_db(prisma_client=prisma_client) if self._should_load_db_object(object_type=SupportedDBObjectType.WEBSEARCH_INTERCEPTION_SETTINGS): await self.init_websearch_interception_settings_in_db(prisma_client=prisma_client) @@ -8469,7 +8540,7 @@ class ProxyConfig: db_values: Final = self._prepared_db_settings_values("litellm_settings", raw_settings) self._apply_litellm_settings_db_values(db_values) - async def _init_semantic_filter_settings_in_db(self, prisma_client: PrismaClient): + async def init_semantic_filter_settings_in_db(self, prisma_client: PrismaClient) -> None: """ Initialize MCP semantic filter settings from database. Called periodically (approximately every 10 seconds) by background task to hot-reload settings across all pods. @@ -8533,6 +8604,8 @@ class ProxyConfig: except Exception as e: verbose_proxy_logger.exception("Error initializing semantic filter settings from DB: %s", e) + _init_semantic_filter_settings_in_db = init_semantic_filter_settings_in_db + async def init_websearch_interception_settings_in_db(self, prisma_client: PrismaClient): """ Initialize web search interception settings from database. @@ -8613,7 +8686,7 @@ class ProxyConfig: sso_settings.sso_settings.pop("team_mappings", None) sso_settings.sso_settings.pop("ui_access_mode", None) uppercase_sso_settings: Final = {key.upper(): value for key, value in sso_settings.sso_settings.items()} - self._decrypt_and_set_db_env_variables(environment_variables=uppercase_sso_settings) + self.decrypt_and_set_db_env_variables(environment_variables=uppercase_sso_settings) except Exception as e: verbose_proxy_logger.exception( "litellm.proxy.proxy_server.py::ProxyConfig:_init_sso_settings_in_db - %s", e @@ -8627,10 +8700,10 @@ class ProxyConfig: """ from litellm.proxy.management_endpoints.config_override_endpoints import ( HASHICORP_ENV_VAR_MAPPING, - _clear_hashicorp_vault_state, - _get_current_env_values, - _parse_config_value, - _set_env_vars, + clear_hashicorp_vault_state, + get_current_env_values, + parse_config_value, + set_env_vars, ) try: @@ -8647,28 +8720,28 @@ class ProxyConfig: if db_record is None or db_record.config_value is None: if self._last_hashicorp_vault_config is not None: - _clear_hashicorp_vault_state(self) + clear_hashicorp_vault_state(self) return - config_data: Final = _parse_config_value(db_record.config_value) + config_data: Final = parse_config_value(db_record.config_value) # Skip reinit if config hasn't changed since last poll if self._last_hashicorp_vault_config == config_data: return # Decrypt all fields and set env vars - decrypted_data: Final = self._decrypt_db_variables(config_data) + decrypted_data: Final = self.decrypt_db_variables(config_data) # Snapshot current env vars so we can restore on failure - previous_env: Final = _get_current_env_values(HASHICORP_ENV_VAR_MAPPING) - _set_env_vars(decrypted_data) + previous_env: Final = get_current_env_values(HASHICORP_ENV_VAR_MAPPING) + set_env_vars(decrypted_data) # Reinitialize the secret manager try: self.initialize_secret_manager(key_management_system="hashicorp_vault") except Exception: # Restore previous working env vars instead of wiping all - _set_env_vars(previous_env) + set_env_vars(previous_env) raise self._last_hashicorp_vault_config = config_data.copy() @@ -8688,10 +8761,10 @@ class ProxyConfig: from litellm.proxy.management_endpoints.config_override_endpoints import ( CYBERARK_ENV_VAR_MAPPING, _clear_cyberark_state, # pyright: ignore[reportPrivateUsage] # module-internal helper shared with the endpoint module - _get_current_env_values, # pyright: ignore[reportPrivateUsage] # module-internal helper shared with the endpoint module - _parse_config_value, # pyright: ignore[reportPrivateUsage] # module-internal helper shared with the endpoint module - _set_env_vars, # pyright: ignore[reportPrivateUsage] # module-internal helper shared with the endpoint module _snapshot_cyberark_boot_env, # pyright: ignore[reportPrivateUsage] # module-internal helper shared with the endpoint module + get_current_env_values, # pyright: ignore[reportPrivateUsage] # module-internal helper shared with the endpoint module + parse_config_value, # pyright: ignore[reportPrivateUsage] # module-internal helper shared with the endpoint module + set_env_vars, # pyright: ignore[reportPrivateUsage] # module-internal helper shared with the endpoint module ) try: @@ -8711,22 +8784,22 @@ class ProxyConfig: _clear_cyberark_state(self) return - config_data: Final = _parse_config_value(db_record.config_value) + config_data: Final = parse_config_value(db_record.config_value) # Skip reinit if config hasn't changed since last poll if self._last_cyberark_config == config_data: return - decrypted_data: Final = self._decrypt_db_variables(config_data) + decrypted_data: Final = self.decrypt_db_variables(config_data) _snapshot_cyberark_boot_env(self) - previous_env: Final = _get_current_env_values(CYBERARK_ENV_VAR_MAPPING) - _set_env_vars(decrypted_data, CYBERARK_ENV_VAR_MAPPING) + previous_env: Final = get_current_env_values(CYBERARK_ENV_VAR_MAPPING) + set_env_vars(decrypted_data, CYBERARK_ENV_VAR_MAPPING) try: self.initialize_secret_manager(key_management_system="cyberark") except Exception: - _set_env_vars(previous_env, CYBERARK_ENV_VAR_MAPPING) + set_env_vars(previous_env, CYBERARK_ENV_VAR_MAPPING) raise self._last_cyberark_config = config_data.copy() @@ -9582,7 +9655,7 @@ def _restamp_streaming_chunk_model( fallback_was_attempted: bool = False, fallback_model_from_metadata: str | None = None, ) -> tuple[Any, bool]: - if _should_return_raw_model_name(request_data): + if should_return_raw_model_name(request_data): return chunk, model_mismatch_logged target_model: Final = fallback_model_from_metadata if fallback_was_attempted else requested_model_from_client @@ -9598,7 +9671,7 @@ def _restamp_streaming_chunk_model( return chunk, model_mismatch_logged # For Azure Model Router, preserve the actual model used in each chunk - if not fallback_was_attempted and _is_azure_model_router_request(requested_model_from_client): + if not fallback_was_attempted and is_azure_model_router_request(requested_model_from_client): return chunk, model_mismatch_logged # For fastest_response batch completions, preserve the winning model's name @@ -10178,7 +10251,7 @@ async def async_data_generator( # The iterator-wrap path fires deferred logging itself; fire it # here for the no-wrap fast path so non-callback deployments # still flush their post-stream logging. - ProxyLogging._fire_deferred_stream_logging(request_data) + ProxyLogging.fire_deferred_stream_logging(request_data) if raw_sse_buffer: yield (raw_sse_buffer if raw_sse_buffer.endswith(_SSE_FRAME_DELIMITERS) else raw_sse_buffer + "\n\n") @@ -10243,7 +10316,7 @@ async def async_data_generator( stream_completed = True yield f"data: {error_returned}\n\n" finally: - await ProxyBaseLLMRequestProcessing._finalize_streaming_generator_cleanup( + await ProxyBaseLLMRequestProcessing.finalize_streaming_generator_cleanup( request=request, request_data=request_data, response=response, @@ -10317,20 +10390,20 @@ def giveup(e): class ProxyStartupEvent: @staticmethod - def _attach_router_to_prompt_injection_detectors(llm_router: Router | None) -> None: - for callback in litellm.logging_callback_manager.get_custom_loggers_for_type( - _OPTIONAL_PromptInjectionDetection - ): - if isinstance(callback, _OPTIONAL_PromptInjectionDetection): + def attach_router_to_prompt_injection_detectors(llm_router: Router | None) -> None: + for callback in litellm.logging_callback_manager.get_custom_loggers_for_type(OPTIONAL_PromptInjectionDetection): + if isinstance(callback, OPTIONAL_PromptInjectionDetection): callback.update_environment(router=llm_router) + _attach_router_to_prompt_injection_detectors = attach_router_to_prompt_injection_detectors + @staticmethod async def refresh_model_info() -> None: if llm_router is not None: await llm_router.arefresh_model_info() @staticmethod - def _warn_budget_without_db(max_budget: float | None, prisma_client: PrismaClient | None) -> None: + def warn_budget_without_db(max_budget: float | None, prisma_client: PrismaClient | None) -> None: if prisma_client is not None or not max_budget or max_budget <= 0: return @@ -10342,8 +10415,10 @@ class ProxyStartupEvent: max_budget, ) + _warn_budget_without_db = warn_budget_without_db + @staticmethod - def _warn_fail_closed_rate_limits_without_redis( + def warn_fail_closed_rate_limits_without_redis( fail_closed_rate_limit_enforcement: bool, redis_usage_cache: RedisCache | None ) -> None: if redis_usage_cache is not None or not fail_closed_rate_limit_enforcement: @@ -10356,8 +10431,10 @@ class ProxyStartupEvent: "across pods and make the setting effective." ) + _warn_fail_closed_rate_limits_without_redis = warn_fail_closed_rate_limits_without_redis + @classmethod - def _initialize_startup_logging( + def initialize_startup_logging( cls, llm_router: Router | None, proxy_logging_obj: ProxyLogging, @@ -10369,8 +10446,10 @@ class ProxyStartupEvent: proxy_logging_obj.startup_event(llm_router=llm_router, redis_usage_cache=redis_usage_cache) + _initialize_startup_logging = initialize_startup_logging + @staticmethod - def _warn_if_mock_testing_params_enabled(general_settings: dict) -> None: + def warn_if_mock_testing_params_enabled(general_settings: dict) -> None: """Announce, loudly, that any caller may inject synthetic failures.""" from litellm.proxy.route_llm_request import ( GATED_MOCK_PARAM_NAMES, @@ -10400,8 +10479,10 @@ class ProxyStartupEvent: "=" * 72, ) + _warn_if_mock_testing_params_enabled = warn_if_mock_testing_params_enabled + @staticmethod - def _validate_redis_transaction_buffer_config( + def validate_redis_transaction_buffer_config( general_settings: dict, redis_usage_cache: RedisCache | None, ): @@ -10430,8 +10511,10 @@ class ProxyStartupEvent: "Redis for the transaction buffer." ) + _validate_redis_transaction_buffer_config = validate_redis_transaction_buffer_config + @staticmethod - async def _init_coordination_redis_from_db( + async def init_coordination_redis_from_db( litellm_settings: Mapping[str, object], llm_router: Router | None, ) -> RedisCache | None: @@ -10459,7 +10542,7 @@ class ProxyStartupEvent: ) return None - coordination_redis_cache: Final = _build_redis_usage_cache(coordination_params.model_dump(exclude_none=True)) + coordination_redis_cache: Final = build_redis_usage_cache(coordination_params.model_dump(exclude_none=True)) _attach_redis_usage_cache( coordination_redis_cache, enable_redis_auth_cache=litellm_settings.get("enable_redis_auth_cache", False) is True, @@ -10472,8 +10555,10 @@ class ProxyStartupEvent: ) return coordination_redis_cache + _init_coordination_redis_from_db = init_coordination_redis_from_db + @staticmethod - def _get_transaction_buffer_redis_cache( + def get_transaction_buffer_redis_cache( general_settings: dict, ) -> RedisCache | None: """ @@ -10495,8 +10580,10 @@ class ProxyStartupEvent: return _build_redis_usage_cache_from_environment() + _get_transaction_buffer_redis_cache = get_transaction_buffer_redis_cache + @classmethod - async def _initialize_semantic_tool_filter( + async def initialize_semantic_tool_filter( cls, llm_router: Router | None, litellm_settings: dict[str, Any], @@ -10528,8 +10615,10 @@ class ProxyStartupEvent: # Only warn if the feature was configured but failed to initialize verbose_proxy_logger.warning("Semantic tool filter hook was configured but failed to initialize") + _initialize_semantic_tool_filter = initialize_semantic_tool_filter + @classmethod - def _initialize_jwt_auth( + def initialize_jwt_auth( cls, general_settings: dict, prisma_client: PrismaClient | None, @@ -10562,14 +10651,18 @@ class ProxyStartupEvent: jwt_handler.bind_agent_lookup(global_agent_registry) + _initialize_jwt_auth = initialize_jwt_auth + @classmethod - def _add_proxy_budget_to_db(cls): + def add_proxy_budget_to_db(cls): """Adds a global proxy budget to db""" if litellm.budget_duration is None: raise Exception("budget_duration not set on Proxy. budget_duration is required to use max_budget.") asyncio.create_task(cls._upsert_proxy_budget_with_reset_at_backfill()) + _add_proxy_budget_to_db = add_proxy_budget_to_db + @classmethod async def _upsert_proxy_budget_with_reset_at_backfill(cls) -> None: """ @@ -10624,7 +10717,7 @@ class ProxyStartupEvent: verbose_proxy_logger.warning("Failed to backfill budget_reset_at on proxy admin row: %s", e) @classmethod - async def _warm_global_spend_cache( + async def warm_global_spend_cache( cls, user_api_key_cache: UserApiKeyCache, prisma_client: PrismaClient, @@ -10632,7 +10725,7 @@ class ProxyStartupEvent: """Warm global spend cache once at startup to reduce impact of first wave of requests.""" try: cache_key: Final = GLOBAL_PROXY_SPEND_CACHE_KEY - await _fetch_global_spend_with_event_coordination( + await fetch_global_spend_with_event_coordination( cache_key=cache_key, user_api_key_cache=user_api_key_cache, prisma_client=prisma_client, @@ -10640,8 +10733,10 @@ class ProxyStartupEvent: except Exception as e: verbose_proxy_logger.debug("Global spend cache warm-up at startup skipped or failed: %s", e) + _warm_global_spend_cache = warm_global_spend_cache + @classmethod - async def _update_default_team_member_budget(cls): + async def update_default_team_member_budget(cls): """Update the default team member budget""" if litellm.default_internal_user_params is None: return @@ -10658,8 +10753,10 @@ class ProxyStartupEvent: user_api_key_dict=UserAPIKeyAuth(token=hash_token(master_key)), ) + _update_default_team_member_budget = update_default_team_member_budget + @classmethod - async def _sync_ui_settings_to_general_settings(cls): + async def sync_ui_settings_to_general_settings(cls): """Apply the persisted UI settings to general_settings before this pod serves traffic.""" if prisma_client is None: return @@ -10667,6 +10764,8 @@ class ProxyStartupEvent: if applied: verbose_proxy_logger.info("Synced UI settings to general_settings on startup: %s", list(applied)) + _sync_ui_settings_to_general_settings = sync_ui_settings_to_general_settings + @classmethod async def _load_heuristic_v1_tuning_baselines( cls, prisma_client: PrismaClient, deployments: Sequence[Mapping[str, object]] @@ -10719,7 +10818,7 @@ class ProxyStartupEvent: cls, prisma_client: PrismaClient, llm_router: Router | None, limit: int | None ) -> Mapping[str, str] | None: """Load a complete baseline and reject a startup that exceeds the tuning quota.""" - db_models: Final = await proxy_config._get_models_from_db(prisma_client) + db_models: Final = await proxy_config.get_models_from_db(prisma_client) if db_models is None: verbose_proxy_logger.warning("Heuristic-v1 tuning baseline unavailable, gate not enforced this boot") return None @@ -10863,10 +10962,10 @@ class ProxyStartupEvent: ### MONITOR SPEND LOGS QUEUE (queue-size-based job) ### if general_settings.get("disable_spend_logs", False) is False: - from litellm.proxy.utils import _monitor_spend_logs_queue + from litellm.proxy.utils import monitor_spend_logs_queue monitor_task: Final = asyncio.create_task( - _monitor_spend_logs_queue( + monitor_spend_logs_queue( prisma_client=prisma_client, db_writer_client=db_writer_client, proxy_logging_obj=proxy_logging_obj, @@ -11519,7 +11618,7 @@ class ProxyStartupEvent: await _scheduled_fallback_stats() @classmethod - async def _setup_prisma_client( + async def setup_prisma_client( cls, database_url: str | None, proxy_logging_obj: ProxyLogging, @@ -11573,8 +11672,10 @@ class ProxyStartupEvent: ) return connected_client + _setup_prisma_client = setup_prisma_client + @classmethod - def _init_dd_tracer(cls): + def init_dd_tracer(cls): """ Initialize dd tracer - if `USE_DDTRACE=true` in .env @@ -11598,8 +11699,10 @@ class ProxyStartupEvent: prof.start() verbose_proxy_logger.debug("Datadog Profiler started......") + _init_dd_tracer = init_dd_tracer + @classmethod - def _init_pyroscope(cls): + def init_pyroscope(cls): """ Optional continuous profiling via Grafana Pyroscope. @@ -11681,6 +11784,8 @@ class ProxyStartupEvent: "Pyroscope profiling will not run. Install with: pip install pyroscope-io" ) + _init_pyroscope = init_pyroscope + #### API ENDPOINTS #### async def _names_hidden_by_listing_callbacks( @@ -11773,7 +11878,7 @@ async def model_list( create_anthropic_model_list_response, ) from litellm.proxy.management_endpoints.common_utils import ( - _user_has_admin_privileges, + user_has_admin_privileges, ) from litellm.proxy.utils import ( create_model_info_response, @@ -11813,7 +11918,9 @@ async def model_list( # Check if scope=expand is requested and user has admin privileges should_expand_scope = False if scope == "expand": - should_expand_scope = _user_has_admin_view(user_api_key_dict) or await _user_has_admin_privileges( + should_expand_scope = user_api_key_has_admin_view( # rebind-ok: pre-existing rebinding on a rename-only line + user_api_key_dict + ) or await user_has_admin_privileges( user_api_key_dict=user_api_key_dict, prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, @@ -12175,7 +12282,7 @@ async def chat_completion( """ global general_settings, user_debug, proxy_logging_obj, llm_model_list global user_temperature, user_request_timeout, user_max_tokens, user_api_base - data: Final = await _read_request_body(request=request) + data: Final = await read_request_body(request=request) if user_api_key_dict is not None: if not isinstance(data.get("metadata"), dict): # Covers both missing and JSON-string metadata (multipart / @@ -12290,7 +12397,7 @@ async def chat_completion( _chat_response.usage = _usage return _chat_response except Exception as e: - raise await base_llm_response_processor._handle_llm_api_exception( + raise await base_llm_response_processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -12336,7 +12443,7 @@ async def completion( global user_temperature, user_request_timeout, user_max_tokens, user_api_base data = {} try: - data = await _read_request_body(request=request) + data = await read_request_body(request=request) # rebind-ok: pre-existing rebinding on a rename-only line if user_api_key_dict is not None: if data.get("metadata") is None: data["metadata"] = {} @@ -12517,7 +12624,7 @@ async def embeddings( """ global proxy_logging_obj - data: Final = await _read_request_body(request=request) + data: Final = await read_request_body(request=request) base_llm_response_processor: Final = ProxyBaseLLMRequestProcessing(data=data) try: ### HANDLE TOKEN ARRAY INPUT DECODING ### @@ -12584,7 +12691,7 @@ async def embeddings( return response except Exception as e: - raise await base_llm_response_processor._handle_llm_api_exception( + raise await base_llm_response_processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -13166,10 +13273,10 @@ async def realtime_websocket_endpoint( request: Final = Request(scope=scope) - request._url = websocket.url + request._url = websocket.url # pyright: ignore[reportPrivateUsage] # Starlette WebSocket URL storage async def return_body(): - return _realtime_request_body(route_model) + return realtime_request_body(route_model) request.body = return_body @@ -14444,7 +14551,7 @@ async def non_admin_all_models( ) # de-duplicate models. Only return unique model ids - unique_models: Final = _deduplicate_litellm_router_models(models=all_models) + unique_models: Final = deduplicate_litellm_router_models(models=all_models) return unique_models @@ -14672,7 +14779,7 @@ async def _populate_team_access_on_models( """ user_teams: list[str] | Literal["*"] | None = None direct_access_models: Sequence[str] = () - if _user_has_admin_view(user_api_key_dict): + if user_api_key_has_admin_view(user_api_key_dict): user_teams = "*" direct_access_models = tuple(llm_router.get_model_ids(exclude_team_models=True)) # access to all models elif user_api_key_dict.user_id is not None: @@ -16514,7 +16621,7 @@ async def model_deprecations( return collect_model_deprecations(llm_router=llm_router, warn_within_days=warn_within_days) -def _get_model_group_info( +def get_model_group_info( llm_router: Router, all_models_str: Sequence[str], model_group: str | None ) -> list[ModelGroupInfoProxy]: model_groups: Final[list[ModelGroupInfoProxy]] = [] @@ -16548,6 +16655,9 @@ def _get_model_group_info( return model_groups +_get_model_group_info: Final = get_model_group_info + + @router.get( "/model_group/info", tags=["model management"], @@ -16756,7 +16866,7 @@ async def model_group_info( ) model_groups: Final = await append_agents_to_model_group( - model_groups=_get_model_group_info( + model_groups=get_model_group_info( llm_router=llm_router, all_models_str=listed_group_names, model_group=model_group ), user_api_key_dict=user_api_key_dict, @@ -16853,7 +16963,7 @@ async def alerting_settings( detail={"error": CommonProxyErrors.db_not_connected_error.value}, ) - if not _user_has_admin_view(user_api_key_dict): + if not user_api_key_has_admin_view(user_api_key_dict): raise HTTPException( status_code=400, detail={"error": f"{CommonProxyErrors.not_allowed_access.value}, your role={user_api_key_dict.user_role}"}, @@ -16981,7 +17091,7 @@ async def async_queue_request( data["proxy_server_request"] = { "url": str(request.url), "method": request.method, - "headers": _safe_get_request_headers(request).copy(), + "headers": safe_get_request_headers(request).copy(), "body": copy.copy(data), # use copy instead of deepcopy } @@ -17006,7 +17116,7 @@ async def async_queue_request( data["metadata"]["user_api_key"] = logged_api_key data["metadata"]["user_api_key_hash"] = logged_api_key data["metadata"]["user_api_key_metadata"] = strip_callback_config(user_api_key_dict.metadata) - _headers: Final = _safe_get_request_headers(request).copy() + _headers: Final = safe_get_request_headers(request).copy() _headers.pop("authorization", None) # do not store the original `sk-..` api key in the db data["metadata"]["headers"] = _headers data["metadata"]["user_api_key_alias"] = getattr(user_api_key_dict, "key_alias", None) @@ -17160,8 +17270,8 @@ async def login(request: Request): # _is_same_origin_return_path (strictly relative path) so it can never be an open redirect, and the # one-shot cookie is cleared after use. from litellm.proxy.management_endpoints.ui_sso import ( - _sso_return_to_redirect, set_session_token_cookie, + sso_return_to_redirect, ) # Resume through the SAME resumer the SSO callback uses, rather than a second, narrower arm. @@ -17174,7 +17284,7 @@ async def login(request: Request): cp_return_to: Final = request.cookies.get("litellm_cp_return_to") if cp_return_to: try: - resumed = await _sso_return_to_redirect( + resumed = await sso_return_to_redirect( # rebind-ok: pre-existing rebinding on a rename-only line return_to=cp_return_to, jwt_token=jwt_token, redis_usage_cache=redis_usage_cache, @@ -17934,11 +18044,14 @@ async def new_invitation(data: InvitationNew, user_api_key_dict: UserAPIKeyAuth ) # Allow proxy admins and org/team admins (admin status from DB via get_user_object) - has_access = user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN or await _user_has_admin_privileges( - user_api_key_dict=user_api_key_dict, - prisma_client=prisma_client, - user_api_key_cache=user_api_key_cache, - proxy_logging_obj=proxy_logging_obj, + has_access: Final = ( + user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN + or await user_has_admin_privileges( + user_api_key_dict=user_api_key_dict, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=proxy_logging_obj, + ) ) if not has_access: raise HTTPException( @@ -17999,7 +18112,7 @@ async def invitation_info(invitation_id: str, user_api_key_dict: UserAPIKeyAuth detail={"error": CommonProxyErrors.db_not_connected_error.value}, ) - if not _user_has_admin_view(user_api_key_dict): + if not user_api_key_has_admin_view(user_api_key_dict): raise HTTPException( status_code=400, detail={"error": f"{CommonProxyErrors.not_allowed_access.value}, your role={user_api_key_dict.user_role}"}, @@ -18106,7 +18219,7 @@ async def invitation_delete( # Proxy admins can delete any invitation; org admins only their own is_proxy_admin: Final = user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN - is_other_admin: Final = await _user_has_admin_privileges( + is_other_admin: Final = await user_has_admin_privileges( user_api_key_dict=user_api_key_dict, prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, @@ -18292,7 +18405,7 @@ async def update_config( existing = await _read_section("environment_variables") before_environment_variables: Final = copy.deepcopy(existing) existing.update( - proxy_config._encrypt_env_variables_for_db(environment_variables=config_info.environment_variables) + proxy_config.encrypt_env_variables_for_db(environment_variables=config_info.environment_variables) ) await _upsert_section("environment_variables", existing) asyncio.create_task( @@ -18556,7 +18669,7 @@ async def update_config_general_settings( proxy_config.settings.apply_db_row("general_settings", general_settings) if is_resource_list("general_settings", data.field_name): stored_endpoints: Final = general_settings.get("pass_through_endpoints") - await proxy_config._serve_pass_through_endpoints(stored_endpoints if isinstance(stored_endpoints, list) else ()) + await proxy_config.serve_pass_through_endpoints(stored_endpoints if isinstance(stored_endpoints, list) else ()) asyncio.create_task( create_config_audit_log( "general_settings", "updated", before_general_settings, general_settings, user_api_key_dict @@ -18747,7 +18860,7 @@ async def get_config_general_settings( detail={"error": CommonProxyErrors.db_not_connected_error.value}, ) - if not _user_has_admin_view(user_api_key_dict): + if not user_api_key_has_admin_view(user_api_key_dict): raise HTTPException( status_code=400, detail={"error": CommonProxyErrors.not_allowed_access.value}, @@ -18952,7 +19065,7 @@ async def get_config_list( detail={"error": CommonProxyErrors.db_not_connected_error.value}, ) - if not _user_has_admin_view(user_api_key_dict): + if not user_api_key_has_admin_view(user_api_key_dict): raise HTTPException( status_code=400, detail={"error": f"{CommonProxyErrors.not_allowed_access.value}, your role={user_api_key_dict.user_role}"}, @@ -19178,7 +19291,7 @@ async def delete_config_general_settings( proxy_config.settings.apply_db_row("general_settings", general_settings) if is_resource_list("general_settings", data.field_name): stored_endpoints: Final = general_settings.get("pass_through_endpoints") - await proxy_config._serve_pass_through_endpoints(stored_endpoints if isinstance(stored_endpoints, list) else ()) + await proxy_config.serve_pass_through_endpoints(stored_endpoints if isinstance(stored_endpoints, list) else ()) asyncio.create_task( create_config_audit_log( "general_settings", "deleted", before_general_settings, general_settings, user_api_key_dict @@ -19704,7 +19817,7 @@ async def get_model_cost_map_reload_status( Get the status of the scheduled model cost map reload job. """ # Read-only status check — admin viewers can read. - if not _user_has_admin_view(user_api_key_dict): + if not user_api_key_has_admin_view(user_api_key_dict): raise HTTPException( status_code=403, detail=f"Access denied. Admin role required. Current role: {user_api_key_dict.user_role}", @@ -19756,7 +19869,7 @@ async def get_model_cost_map_source( - model_count: number of models in the currently loaded cost map """ # Read-only source info — admin viewers can read. - if not _user_has_admin_view(user_api_key_dict): + if not user_api_key_has_admin_view(user_api_key_dict): raise HTTPException( status_code=403, detail=f"Access denied. Admin role required. Current role: {user_api_key_dict.user_role}", @@ -19978,7 +20091,7 @@ async def get_anthropic_beta_headers_reload_status( Get the status of the scheduled Anthropic beta headers reload job. """ # Read-only status — admin viewers can read. - if not _user_has_admin_view(user_api_key_dict): + if not user_api_key_has_admin_view(user_api_key_dict): raise HTTPException( status_code=403, detail=f"Access denied. Admin role required. Current role: {user_api_key_dict.user_role}", @@ -20079,7 +20192,7 @@ async def get_adaptive_router_state( which deployment it came from. """ # Read-only state — admin viewers can read. - if not _user_has_admin_view(user_api_key_dict): + if not user_api_key_has_admin_view(user_api_key_dict): raise HTTPException( status_code=403, detail={"error": CommonProxyErrors.not_allowed_access.value}, @@ -20183,6 +20296,7 @@ app.include_router(callback_management_endpoints_router) app.include_router(debugging_endpoints_router) app.include_router(rust_control_plane_router) app.include_router(ui_crud_endpoints_router) +app.include_router(moyai_router) app.include_router(user_banner_endpoints_router) app.include_router(latest_release_endpoints_router) app.include_router(team_callback_router) @@ -20194,6 +20308,7 @@ app.include_router(tag_management_router) app.include_router(workflow_management_router) app.include_router(memory_router) app.include_router(lens_dataset_router) +app.include_router(lens_feedback_router) app.include_router(lens_router) app.include_router(plugin_router) app.include_router(cost_tracking_settings_router) @@ -20245,7 +20360,7 @@ async def _stream_mcp_asgi_response(handle_fn, scope: dict, receive) -> "Streami Call an ASGI MCP handler and return a StreamingResponse so SSE/streaming works. asyncio.create_task copies the current context, so any ContextVar set before - this call (e.g. _mcp_active_toolset_id) is visible inside the handler task. + this call (e.g. mcp_active_toolset_id) is visible inside the handler task. """ from starlette.responses import StreamingResponse @@ -20385,8 +20500,8 @@ async def toolset_mcp_route(toolset_name: str, request: Request): global_mcp_server_manager, ) from litellm.proxy._experimental.mcp_server.server import ( - _mcp_active_toolset_id, handle_streamable_http_mcp, + mcp_active_toolset_id, ) if prisma_client is None: @@ -20402,11 +20517,11 @@ async def toolset_mcp_route(toolset_name: str, request: Request): scope: Final = dict(request.scope) scope["path"] = "/mcp" - token: Final = _mcp_active_toolset_id.set(toolset.toolset_id) + token: Final = mcp_active_toolset_id.set(toolset.toolset_id) try: return await _stream_mcp_asgi_response(handle_streamable_http_mcp, scope, request.receive) finally: - _mcp_active_toolset_id.reset(token) + mcp_active_toolset_id.reset(token) except HTTPException as e: raise e @@ -20492,7 +20607,7 @@ async def _is_mcp_access_group_cached(name: str) -> bool: cached: Final[object] = await user_api_key_cache.async_get_cache(key=cache_key) if cached is not None: return bool(cached) - result: Final = bool(await MCPRequestHandler._get_mcp_servers_from_access_groups([name])) + result: Final = bool(await MCPRequestHandler.get_mcp_servers_from_access_groups([name])) await user_api_key_cache.async_set_cache( key=cache_key, value=result, @@ -20548,8 +20663,8 @@ async def dynamic_mcp_route(mcp_server_name: str, request: Request): # 3. Toolset name (cached) if prisma_client is not None: from litellm.proxy._experimental.mcp_server.server import ( - _mcp_active_toolset_id, handle_streamable_http_mcp, + mcp_active_toolset_id, ) toolset: Final = await global_mcp_server_manager.get_toolset_by_name_cached(prisma_client, mcp_server_name) @@ -20557,11 +20672,11 @@ async def dynamic_mcp_route(mcp_server_name: str, request: Request): scope: Final = dict(request.scope) scope["_original_path"] = scope.get("path", "") scope["path"] = "/mcp" - token: Final = _mcp_active_toolset_id.set(toolset.toolset_id) + token: Final = mcp_active_toolset_id.set(toolset.toolset_id) try: return await _stream_mcp_asgi_response(handle_streamable_http_mcp, scope, request.receive) finally: - _mcp_active_toolset_id.reset(token) + mcp_active_toolset_id.reset(token) # 4. MCP access group tag (cached) if await _is_mcp_access_group_cached(mcp_server_name): diff --git a/litellm/proxy/public_endpoints/public_endpoints.py b/litellm/proxy/public_endpoints/public_endpoints.py index da9a3033187..6b6ee7a9b67 100644 --- a/litellm/proxy/public_endpoints/public_endpoints.py +++ b/litellm/proxy/public_endpoints/public_endpoints.py @@ -221,10 +221,10 @@ def _load_endpoints() -> list[_EndpointEntry]: async def public_model_hub(): import litellm from litellm.proxy.health_endpoints._health_endpoints import ( - _convert_health_check_to_dict, + convert_health_check_to_dict, ) from litellm.proxy.proxy_server import ( - _get_model_group_info, + get_model_group_info, llm_router, prisma_client, ) @@ -234,7 +234,7 @@ async def public_model_hub(): model_groups: list[ModelGroupInfoProxy] = [] if litellm.public_model_groups is not None: - model_groups = _get_model_group_info( + model_groups = get_model_group_info( # rebind-ok: pre-existing rebinding on a rename-only line llm_router=llm_router, all_models_str=litellm.public_model_groups, model_group=None, @@ -248,7 +248,7 @@ async def public_model_hub(): for check in latest_checks: key = check.model_id if check.model_id else check.model_name if key: - health_check_dict = _convert_health_check_to_dict(check) + health_check_dict = convert_health_check_to_dict(check) health_checks_map[key] = health_check_dict if check.model_name: health_checks_map[check.model_name] = health_check_dict @@ -322,7 +322,7 @@ async def get_mcp_servers(): async def public_skill_hub(): """Return enabled (public) Claude Code skills — no auth required.""" from litellm.proxy.anthropic_endpoints.claude_code_endpoints.claude_code_marketplace import ( - _get_prisma_client, + get_prisma_client, ) from litellm.types.proxy.claude_code_endpoints import ( ListPluginsResponse, @@ -330,7 +330,7 @@ async def public_skill_hub(): ) try: - prisma_client: Final = await _get_prisma_client() + prisma_client: Final = await get_prisma_client() plugins: Final = await _plugin_table(prisma_client).find_many(where={"enabled": True}) items: Final = [] for plugin in plugins: @@ -366,7 +366,7 @@ async def public_skill_hub(): ) async def public_model_hub_info(): import litellm - from litellm.proxy.proxy_server import _title, version + from litellm.proxy.proxy_server import title, version try: from litellm_enterprise.proxy.proxy_server import EnterpriseProxyConfig @@ -376,7 +376,7 @@ async def public_model_hub_info(): custom_docs_description = None return PublicModelHubInfo( - docs_title=_title, + docs_title=title, custom_docs_description=custom_docs_description, litellm_version=version, useful_links=litellm.public_model_groups_links, diff --git a/litellm/proxy/public_endpoints/public_v1/model_hub.py b/litellm/proxy/public_endpoints/public_v1/model_hub.py index 93971a84547..2f5eeeb3d4b 100644 --- a/litellm/proxy/public_endpoints/public_v1/model_hub.py +++ b/litellm/proxy/public_endpoints/public_v1/model_hub.py @@ -186,7 +186,7 @@ MODEL_HUB_LIST_SPEC: Final[ListSpec[ModelGroupInfoProxy, ModelGroupInfoProxy]] = def _published_rows() -> Sequence[ModelGroupInfoProxy]: from litellm.proxy.proxy_server import ( - _get_model_group_info, # pyright: ignore[reportPrivateUsage] # /public/model_hub imports it the same way + get_model_group_info, # pyright: ignore[reportPrivateUsage] # /public/model_hub imports it the same way llm_router, ) @@ -202,7 +202,7 @@ def _published_rows() -> Sequence[ModelGroupInfoProxy]: if litellm.public_model_groups is None: return () return tuple( - _get_model_group_info( + get_model_group_info( llm_router=llm_router, all_models_str=litellm.public_model_groups, model_group=None, diff --git a/litellm/proxy/rag_endpoints/endpoints.py b/litellm/proxy/rag_endpoints/endpoints.py index 0913d216678..a1a7d62899f 100644 --- a/litellm/proxy/rag_endpoints/endpoints.py +++ b/litellm/proxy/rag_endpoints/endpoints.py @@ -31,10 +31,12 @@ from litellm.proxy.common_request_processing import ( open_sse_before_first_byte, ttft_keepalive_interval, ) -from litellm.proxy.common_utils.http_parsing_utils import ( - _read_request_body, - _safe_get_request_headers, +from litellm.proxy.common_utils.http_parsing_utils import ( # noqa: F401 # legacy module exports + _read_request_body, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _safe_get_request_headers, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export get_form_data, + read_request_body, + safe_get_request_headers, ) from litellm.proxy.rag_endpoints.upload_security import ( MAX_UPLOAD_SIZE_BYTES, @@ -420,7 +422,7 @@ async def parse_rag_ingest_request( Returns: Tuple of (ingest_options, file_data, file_url, file_id) """ - headers: Final = _safe_get_request_headers(request) + headers: Final = safe_get_request_headers(request) content_type = headers.get("content-type", "") file_data: tuple[str, bytes, str] | None = None @@ -448,7 +450,7 @@ async def parse_rag_ingest_request( else: # JSON body - data: Final = await _read_request_body(request) + data: Final = await read_request_body(request) ingest_options = data.get("ingest_options", {}) file_url = data.get("file_url") file_id = data.get("file_id") @@ -770,7 +772,7 @@ async def rag_query( try: # Parse request body - data: Final = await _read_request_body(request) + data: Final = await read_request_body(request) # Extract required fields model: Final = data.get("model") diff --git a/litellm/proxy/realtime_endpoints/endpoints.py b/litellm/proxy/realtime_endpoints/endpoints.py index be2ac2ff33e..07b1abb62e4 100644 --- a/litellm/proxy/realtime_endpoints/endpoints.py +++ b/litellm/proxy/realtime_endpoints/endpoints.py @@ -16,7 +16,10 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import ( decrypt_value_helper, encrypt_value_helper, ) -from litellm.proxy.common_utils.http_parsing_utils import _read_request_body +from litellm.proxy.common_utils.http_parsing_utils import ( # noqa: F401 # legacy module exports + _read_request_body, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + read_request_body, +) from litellm.proxy.common_utils.openai_error_payload import ( error_status_code, openai_error_param, @@ -247,7 +250,7 @@ async def create_realtime_client_secret( data: dict = {} try: - body: Final = await _read_request_body(request=request) + body: Final = await read_request_body(request=request) req: Final = RealtimeClientSecretRequest(**body) model, session_data, session_type = await _prepare_client_secret_session( @@ -559,7 +562,7 @@ async def create_realtime_transcription_session( data: dict = {} try: - body: Final = await _read_request_body(request=request) + body: Final = await read_request_body(request=request) req: Final = RealtimeTranscriptionSessionRequest(**body) model: Final[str] = req.resolved_model() or "gpt-realtime-whisper" diff --git a/litellm/proxy/response_api_endpoints/endpoints.py b/litellm/proxy/response_api_endpoints/endpoints.py index 00af96af94c..e88762f0ca5 100644 --- a/litellm/proxy/response_api_endpoints/endpoints.py +++ b/litellm/proxy/response_api_endpoints/endpoints.py @@ -12,7 +12,7 @@ from fastapi import APIRouter, Depends, HTTPException, Request, Response from fastapi.responses import JSONResponse from openai.types.responses import ResponseItemList from openai.types.responses.response_create_params import ResponseInputParam -from pydantic import ConfigDict, ValidationError +from pydantic import ConfigDict, TypeAdapter, ValidationError from starlette.websockets import WebSocket, WebSocketDisconnect from typing_extensions import ReadOnly, TypedDict @@ -30,9 +30,11 @@ from litellm.proxy.auth.user_api_key_auth import ( user_api_key_auth_websocket, ) from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing, create_response -from litellm.proxy.common_utils.http_parsing_utils import ( - _read_request_body, - _safe_set_request_parsed_body, +from litellm.proxy.common_utils.http_parsing_utils import ( # noqa: F401 # legacy module exports + _read_request_body, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _safe_set_request_parsed_body, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + read_request_body, + safe_set_request_parsed_body, ) from litellm.proxy.route_llm_request import raise_if_required_body_param_missing from litellm.types.llms.base import LiteLLMBaseModel @@ -48,6 +50,7 @@ if TYPE_CHECKING: from litellm.router import Router router: Final = APIRouter() +_RESPONSES_WS_CONFIG_VALUE_ADAPTER: Final[TypeAdapter[object | None]] = TypeAdapter(object | None) _ResponseDocSchemas: TypeAlias = dict[int | str, dict[str, object]] # fastapi's responses kwarg @@ -180,12 +183,12 @@ async def _resolve_cursor_model_variant_before_auth(request: Request) -> None: from litellm.proxy.proxy_server import llm_router try: - raw_body: Final = await _read_request_body(request=request) + raw_body: Final = await read_request_body(request=request) except (json.JSONDecodeError, ProxyException): return resolved: Final = _resolve_cursor_model_variant(raw_body, llm_router) if resolved is not raw_body: - _safe_set_request_parsed_body(request=request, parsed_body=resolved) + safe_set_request_parsed_body(request=request, parsed_body=resolved) @router.post( @@ -240,7 +243,6 @@ async def responses_api( ``` """ from litellm.proxy.proxy_server import ( - _read_request_body, general_settings, llm_router, native_background_mode, @@ -248,6 +250,7 @@ async def responses_api( polling_via_cache_enabled, proxy_config, proxy_logging_obj, + read_request_body, redis_usage_cache, select_data_generator, user_api_base, @@ -259,7 +262,7 @@ async def responses_api( ) native_data_generator: Final = partial(select_data_generator, responses_stream_errors=True) - data = await _read_request_body(request=request) + data = await read_request_body(request=request) # rebind-ok: pre-existing rebinding on a rename-only line # Check if polling via cache should be used for this request from litellm.proxy.response_polling.polling_handler import ( @@ -309,7 +312,7 @@ async def responses_api( ) raise_if_required_body_param_missing(route_type="aresponses", data=data, llm_router=llm_router) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -447,7 +450,7 @@ async def responses_api( return await create_response(generator=_blocked_stream(), media_type="text/event-stream", headers={}) return build_blocked_response(e) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -543,7 +546,7 @@ async def cursor_chat_completions( from litellm.types.llms.openai import ResponsesAPIResponse from litellm.types.utils import ModelResponse - raw_body: Final = await _read_request_body(request=request) + raw_body: Final = await read_request_body(request=request) if _is_chat_completions_body(raw_body): # Genuine chat completions body (Cursor sends these for models whose BYOK it @@ -552,7 +555,7 @@ async def cursor_chat_completions( # empty messages stub alongside a real agent-mode input array normalized: Final = _normalize_tool_dialect(raw_body, to_chat=True) if normalized is not raw_body: - _safe_set_request_parsed_body(request=request, parsed_body=normalized) + safe_set_request_parsed_body(request=request, parsed_body=normalized) return await chat_completion( request=request, fastapi_response=fastapi_response, @@ -663,7 +666,7 @@ async def cursor_chat_completions( # Streaming responses are already transformed by cursor_select_data_generator return response except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -715,11 +718,11 @@ async def get_response( ``` """ from litellm.proxy.proxy_server import ( - _read_request_body, general_settings, llm_router, proxy_config, proxy_logging_obj, + read_request_body, redis_usage_cache, select_data_generator, user_api_base, @@ -756,7 +759,7 @@ async def get_response( return state # Normal provider response flow - data: Final = await _read_request_body(request=request) + data: Final = await read_request_body(request=request) data["response_id"] = response_id processor: Final = ProxyBaseLLMRequestProcessing(data=data) try: @@ -779,7 +782,7 @@ async def get_response( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -826,11 +829,11 @@ async def delete_response( ``` """ from litellm.proxy.proxy_server import ( - _read_request_body, general_settings, llm_router, proxy_config, proxy_logging_obj, + read_request_body, redis_usage_cache, select_data_generator, user_api_base, @@ -865,7 +868,7 @@ async def delete_response( raise HTTPException(status_code=500, detail="Failed to delete polling response") # Normal provider response flow - data: Final = await _read_request_body(request=request) + data: Final = await read_request_body(request=request) data["response_id"] = response_id processor: Final = ProxyBaseLLMRequestProcessing(data=data) try: @@ -888,7 +891,7 @@ async def delete_response( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -922,11 +925,11 @@ async def get_response_input_items( ): """List input items for a response.""" from litellm.proxy.proxy_server import ( - _read_request_body, general_settings, llm_router, proxy_config, proxy_logging_obj, + read_request_body, select_data_generator, user_api_base, user_max_tokens, @@ -936,7 +939,7 @@ async def get_response_input_items( version, ) - data: Final = await _read_request_body(request=request) + data: Final = await read_request_body(request=request) data["response_id"] = response_id processor: Final = ProxyBaseLLMRequestProcessing(data=data) try: @@ -959,7 +962,7 @@ async def get_response_input_items( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -1005,11 +1008,11 @@ async def compact_response( ``` """ from litellm.proxy.proxy_server import ( - _read_request_body, general_settings, llm_router, proxy_config, proxy_logging_obj, + read_request_body, select_data_generator, user_api_base, user_max_tokens, @@ -1019,7 +1022,7 @@ async def compact_response( version, ) - data: Final = await _read_request_body(request=request) + data: Final = await read_request_body(request=request) processor: Final = ProxyBaseLLMRequestProcessing(data=data) try: return await processor.base_process_llm_request( @@ -1041,7 +1044,7 @@ async def compact_response( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -1158,7 +1161,7 @@ async def responses_input_tokens( Returns: `{"object": "response.input_tokens", "input_tokens": }` """ - data: Final = await _read_request_body(request=request) + data: Final = await read_request_body(request=request) model_name: Final = data.get("model") input_value: Final = data.get("input") if not isinstance(model_name, str) or not model_name: @@ -1236,11 +1239,11 @@ async def cancel_response( ``` """ from litellm.proxy.proxy_server import ( - _read_request_body, general_settings, llm_router, proxy_config, proxy_logging_obj, + read_request_body, redis_usage_cache, select_data_generator, user_api_base, @@ -1279,7 +1282,7 @@ async def cancel_response( raise HTTPException(status_code=500, detail="Failed to cancel polling response") # Normal provider response flow - data: Final = await _read_request_body(request=request) + data: Final = await read_request_body(request=request) data["response_id"] = response_id processor: Final = ProxyBaseLLMRequestProcessing(data=data) try: @@ -1302,7 +1305,7 @@ async def cancel_response( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -1310,6 +1313,21 @@ async def cancel_response( ) +def _resolve_responses_ws_session_limit_seconds() -> float: + from litellm.proxy.proxy_server import general_settings + + field: Final = "responses_websocket_session_limit_seconds" + raw: Final = _RESPONSES_WS_CONFIG_VALUE_ADAPTER.validate_python(general_settings.get(field)) + try: + return ConfigGeneralSettings.model_validate( + {} if raw is None else {field: raw} + ).responses_websocket_session_limit_seconds + except ValidationError as e: + default: Final = DEFAULT_RESPONSES_WEBSOCKET_SESSION_LIMIT_SECONDS + verbose_proxy_logger.warning("invalid general_settings.%s=%r (%s); using default %ss", field, raw, e, default) + return default + + async def _read_ws_model_from_first_frame( websocket: WebSocket, query_model: str | None = None, @@ -1317,12 +1335,10 @@ async def _read_ws_model_from_first_frame( """Read the first WS frame and return (model, raw_message), or None on error. Sends an appropriate error frame and closes the socket before returning None. + The session-duration deadline is enforced by the caller, not here. """ try: - first_message: Final = await asyncio.wait_for(websocket.receive_text(), timeout=30) - except asyncio.TimeoutError: - await websocket.close(code=1008, reason="Timed out waiting for first message") - return None + first_message: Final = await websocket.receive_text() except WebSocketDisconnect: return None except Exception: @@ -1432,8 +1448,8 @@ async def _enforce_responses_ws_first_frame_model_auth( llm_router: "Router | None", ) -> None: from litellm.proxy.auth.user_api_key_auth import ( - _enforce_key_and_fallback_model_access, - _run_centralized_common_checks, + enforce_key_and_fallback_model_access, + run_centralized_common_checks, ) from litellm.proxy.proxy_server import ( general_settings, @@ -1452,7 +1468,7 @@ async def _enforce_responses_ws_first_frame_model_auth( return if user_custom_auth is not None and not general_settings.get("custom_auth_run_common_checks", False): return - await _enforce_key_and_fallback_model_access( + await enforce_key_and_fallback_model_access( valid_token=user_api_key_dict, request_data=request_data, route=route, @@ -1460,7 +1476,7 @@ async def _enforce_responses_ws_first_frame_model_auth( llm_model_list=llm_model_list, llm_router=llm_router, ) - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=user_api_key_dict, request=request, request_data=request_data, @@ -1468,25 +1484,11 @@ async def _enforce_responses_ws_first_frame_model_auth( ) -@router.websocket("/v1/responses") -@router.websocket("/responses") -async def responses_websocket_endpoint( +async def _responses_websocket_session( websocket: WebSocket, - model: str | None = fastapi.Query(None, description="The model to use for the responses WebSocket session."), - user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth_websocket), -): - """ - Responses API WebSocket mode endpoint. - - Keeps a persistent WebSocket connection for response.create events, - enabling lower-latency agentic workflows with many tool-call round trips. - - Follows the OpenAI split: the bearer token is validated at connection time - (before accept); the model is resolved either from the ?model= query param - or from the first response.create frame, whichever is present. - - See: https://developers.openai.com/api/docs/guides/websocket-mode/ - """ + model: str | None, + user_api_key_dict: UserAPIKeyAuth, +) -> None: from litellm.proxy.proxy_server import ( general_settings, llm_router, @@ -1501,16 +1503,6 @@ async def responses_websocket_endpoint( ) from litellm.proxy.route_llm_request import route_request - # Accept the WebSocket handshake. Key was already validated by the Depends - # above; we can safely accept regardless of whether ?model= was supplied. - requested_protocols: Final = [ - p.strip() for p in (websocket.headers.get("sec-websocket-protocol") or "").split(",") if p.strip() - ] - accept_kwargs: Final[dict] = {} - if requested_protocols: - accept_kwargs["subprotocol"] = requested_protocols[0] - await websocket.accept(**accept_kwargs) - result: Final = await _read_ws_model_from_first_frame(websocket, query_model=model) if result is None: return @@ -1531,7 +1523,7 @@ async def responses_websocket_endpoint( "headers": headers_list, } request: Final = Request(scope=scope) - request._url = websocket.url + request._url = websocket.url # pyright: ignore[reportPrivateUsage] # Starlette WebSocket URL storage _body_bytes: Final = json.dumps({"model": resolved_model}).encode() @@ -1615,3 +1607,55 @@ async def responses_websocket_endpoint( request_data=routed_data, ) await websocket.close(code=1011, reason="Internal server error") + + +@router.websocket("/v1/responses") +@router.websocket("/responses") +async def responses_websocket_endpoint( + websocket: WebSocket, + model: str | None = fastapi.Query(None, description="The model to use for the responses WebSocket session."), + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth_websocket), +): + """ + Responses API WebSocket mode endpoint. + + Keeps a persistent WebSocket connection for response.create events, + enabling lower-latency agentic workflows with many tool-call round trips. + + Follows the OpenAI split: the bearer token is validated at connection time + (before accept); the model is resolved either from the ?model= query param + or from the first response.create frame, whichever is present. + + The session is bounded by a lifetime measured from accept, configured via + general_settings.responses_websocket_session_limit_seconds (60-7200, + default 3600). There is no separate first-frame deadline, so + pre-established connections may sit idle until their first response.create. + + See: https://developers.openai.com/api/docs/guides/websocket-mode/ + """ + # Accept the WebSocket handshake. Key was already validated by the Depends + # above; we can safely accept regardless of whether ?model= was supplied. + requested_protocols: Final = [ + p.strip() for p in (websocket.headers.get("sec-websocket-protocol") or "").split(",") if p.strip() + ] + accept_kwargs: Final[dict] = {} + if requested_protocols: + accept_kwargs["subprotocol"] = requested_protocols[0] + await websocket.accept(**accept_kwargs) + + limit_seconds: Final = _resolve_responses_ws_session_limit_seconds() + session_task: Final = asyncio.ensure_future( + _responses_websocket_session(websocket=websocket, model=model, user_api_key_dict=user_api_key_dict) + ) + try: + await asyncio.wait_for(asyncio.shield(session_task), timeout=limit_seconds) + except asyncio.TimeoutError: + verbose_proxy_logger.info("Responses WebSocket closed: session duration limit reached") + session_task.cancel() + with contextlib.suppress(Exception): + await websocket.close(code=1000, reason="Session duration limit reached") + finally: + if not session_task.done(): + session_task.cancel() + with contextlib.suppress(asyncio.CancelledError, Exception): + await session_task diff --git a/litellm/proxy/route_llm_request.py b/litellm/proxy/route_llm_request.py index e97099a6fe0..09c483c47f2 100644 --- a/litellm/proxy/route_llm_request.py +++ b/litellm/proxy/route_llm_request.py @@ -245,11 +245,13 @@ def _find_missing_required_body_param( if not missing_present_params: return None candidate_litellm_params: Final = _candidate_deployment_litellm_params(data, llm_router) + router_default_litellm_params: Final = _router_default_litellm_params(route_type, data, llm_router) missing_param: Final = next( ( param for param in missing_present_params - if not any(deployment_params.get(param) is not None for deployment_params in candidate_litellm_params) + if router_default_litellm_params.get(param) is None + and not any(deployment_params.get(param) is not None for deployment_params in candidate_litellm_params) ), None, ) @@ -258,6 +260,32 @@ def _find_missing_required_body_param( return MissingBodyParam(name=missing_param, model_deployments_loaded=bool(candidate_litellm_params)) +_ROUTE_TYPES_WITHOUT_ROUTER_DEFAULTS_MERGE: Final[frozenset[str]] = frozenset( + {"asearch", "acreate_agent", "acreate_eval", "acreate_run"} +) + + +def _router_default_litellm_params( + route_type: str, + data: Mapping[str, object], + llm_router: LitellmRouter | None, +) -> Mapping[str, object]: + # Mirror exactly the defaults the dispatching router will merge at dispatch time: + # user_config requests dispatch on their own throwaway Router, and the listed route + # types (plus model-less direct dispatch) never pass through the router's merge. + user_config: Final[Mapping[str, object] | None] = ( + data.get("user_config") if isinstance(data.get("user_config"), Mapping) else None + ) + if user_config is not None: + defaults: Final[Mapping[str, object] | None] = user_config.get("default_litellm_params") + return defaults if isinstance(defaults, Mapping) else {} + model_name: Final = data.get("model") + if route_type in _ROUTE_TYPES_WITHOUT_ROUTER_DEFAULTS_MERGE or not isinstance(model_name, str) or not model_name: + return {} + router_defaults: Final[Mapping[str, object] | None] = getattr(llm_router, "default_litellm_params", None) + return router_defaults if isinstance(router_defaults, Mapping) else {} + + def _candidate_deployment_litellm_params( data: Mapping[str, object], llm_router: LitellmRouter | None, @@ -412,7 +440,9 @@ async def add_shared_session_to_data(data: dict) -> None: "SESSION REUSE: Shared aiohttp session is None after re-check, recreating..." ) try: - new_session = await proxy_server._initialize_shared_aiohttp_session() + new_session = ( # rebind-ok: pre-existing rebinding on a rename-only line + await proxy_server.initialize_shared_aiohttp_session() + ) except Exception: verbose_proxy_logger.exception("SESSION REUSE: Exception during shared session recreation") new_session = None diff --git a/litellm/proxy/search_endpoints/endpoints.py b/litellm/proxy/search_endpoints/endpoints.py index 1e37e3243ea..f5adc9be9f1 100644 --- a/litellm/proxy/search_endpoints/endpoints.py +++ b/litellm/proxy/search_endpoints/endpoints.py @@ -208,7 +208,7 @@ async def search( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, diff --git a/litellm/proxy/search_endpoints/search_tool_registry.py b/litellm/proxy/search_endpoints/search_tool_registry.py index fe119a9d44d..c8cd1100e08 100644 --- a/litellm/proxy/search_endpoints/search_tool_registry.py +++ b/litellm/proxy/search_endpoints/search_tool_registry.py @@ -12,10 +12,11 @@ from pydantic import TypeAdapter, ValidationError from litellm._logging import verbose_proxy_logger from litellm.litellm_core_utils.safe_json_dumps import safe_dumps from litellm.proxy.auth.master_key_boot_check import SALT_KEY_ENV_VAR -from litellm.proxy.common_utils.encrypt_decrypt_utils import ( - _get_salt_key, +from litellm.proxy.common_utils.encrypt_decrypt_utils import ( # noqa: F401 # legacy module exports + _get_salt_key, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export decrypt_if_encrypted_with, encrypt_value_helper, + get_salt_key, ) from litellm.proxy.db.exception_handler import call_with_db_reconnect_retry from litellm.proxy.utils import PrismaClient @@ -85,7 +86,7 @@ def encrypt_search_tool_litellm_params(litellm_params: Mapping[str, object]) -> def _search_tool_plaintext(value: str) -> str | None: - signing_key: Final = _get_salt_key() + signing_key: Final = get_salt_key() return None if signing_key is None else decrypt_if_encrypted_with(value, signing_key) diff --git a/litellm/proxy/spend_tracking/budget_reservation.py b/litellm/proxy/spend_tracking/budget_reservation.py index bf9155aed1c..4b6cc58a88e 100644 --- a/litellm/proxy/spend_tracking/budget_reservation.py +++ b/litellm/proxy/spend_tracking/budget_reservation.py @@ -438,10 +438,10 @@ async def invalidate_budget_reservation_counters( if budget_reservation is None: return - from litellm.proxy.proxy_server import _invalidate_spend_counter + from litellm.proxy.proxy_server import invalidate_spend_counter for counter_key in get_reserved_counter_keys(budget_reservation=budget_reservation): - await _invalidate_spend_counter(counter_key=counter_key) + await invalidate_spend_counter(counter_key=counter_key) async def release_or_invalidate_budget_reservation( @@ -946,19 +946,19 @@ async def _reservation_counter_loaded(counter: _BudgetCounter, fail_closed_budge async def _initialize_reservation_counter(counter: _BudgetCounter) -> None: from litellm.proxy.proxy_server import ( - _ensure_spend_counter_initialized, - _ensure_window_spend_counter_initialized, - _invalidate_spend_counter, + ensure_spend_counter_initialized, + ensure_window_spend_counter_initialized, + invalidate_spend_counter, ) try: if counter.source_cache_key is not None: - await _ensure_spend_counter_initialized( + await ensure_spend_counter_initialized( counter_key=counter.counter_key, source_cache_key=counter.source_cache_key, ) elif counter.spend_log_entity_id is not None and counter.window_start is not None: - initialized: Final = await _ensure_window_spend_counter_initialized( + initialized: Final = await ensure_window_spend_counter_initialized( counter_key=counter.counter_key, entity_type=counter.entity_type, entity_id=counter.spend_log_entity_id, @@ -980,7 +980,7 @@ async def _initialize_reservation_counter(counter: _BudgetCounter) -> None: exc_info=True, ) try: - await _invalidate_spend_counter(counter_key=counter.counter_key) + await invalidate_spend_counter(counter_key=counter.counter_key) except Exception: verbose_proxy_logger.warning( "Failed to invalidate spend counter after budget reservation failure for %s", @@ -1062,7 +1062,7 @@ async def _reserve_counters( ) -> tuple[float | None, ...] | None: """One INCRBYFLOAT pipeline reserves every counter. When it fails each counter is dropped, and one that cannot be dropped is released instead in case its increment landed, so nothing is left to release by the caller.""" - from litellm.proxy.proxy_server import _invalidate_spend_counter, run_spend_counter_pipeline + from litellm.proxy.proxy_server import invalidate_spend_counter, run_spend_counter_pipeline if not counters: return () @@ -1081,7 +1081,7 @@ async def _reserve_counters( ) for counter, entry in zip(counters, entries): try: - await _invalidate_spend_counter(counter_key=counter.counter_key) + await invalidate_spend_counter(counter_key=counter.counter_key) except Exception: verbose_proxy_logger.warning( "Failed to invalidate spend counter after budget reservation failure for %s", @@ -1190,11 +1190,11 @@ async def _reseed_reserved_entry(item: _EntryAdjustment, actual_cost: float) -> reconcile: the optimistic delta no longer applies, so reseed from the DB floor and add the settled cost, since increment_spend_counters skips reserved keys. The reconcile runs before this request's spend is enqueued to the DB, so the reseeded floor excludes it.""" - from litellm.proxy.proxy_server import _increment_spend_counter_cache, reseed_spend_counter_from_db + from litellm.proxy.proxy_server import increment_spend_counter_cache, reseed_spend_counter_from_db reseeded: Final = await reseed_spend_counter_from_db(counter_key=item.counter_key) if reseeded and actual_cost > 0: - await _increment_spend_counter_cache(counter_key=item.counter_key, increment=actual_cost) + await increment_spend_counter_cache(counter_key=item.counter_key, increment=actual_cost) async def _counter_can_apply_adjustment( @@ -1230,9 +1230,9 @@ async def _release_applied_entries_best_effort( if counter_key is None: continue try: - from litellm.proxy.proxy_server import _invalidate_spend_counter + from litellm.proxy.proxy_server import invalidate_spend_counter - await _invalidate_spend_counter(counter_key=counter_key) + await invalidate_spend_counter(counter_key=counter_key) except Exception: verbose_proxy_logger.exception( "Failed to invalidate partial budget reservation counter during exception cleanup" diff --git a/litellm/proxy/spend_tracking/cloudzero_endpoints.py b/litellm/proxy/spend_tracking/cloudzero_endpoints.py index 6c0d11a2174..66a1934b40b 100644 --- a/litellm/proxy/spend_tracking/cloudzero_endpoints.py +++ b/litellm/proxy/spend_tracking/cloudzero_endpoints.py @@ -17,7 +17,10 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import ( decrypt_value_helper, encrypt_value_helper, ) -from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view +from litellm.proxy.management_endpoints.common_utils import ( # noqa: F401 # legacy module exports + _user_has_admin_view, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + user_api_key_has_admin_view, +) from litellm.repositories.config_repository import ConfigRepository from litellm.repositories.prisma_protocols import TableActions from litellm.types.proxy.cloudzero_endpoints import ( @@ -149,7 +152,7 @@ async def get_cloudzero_settings( Only admin users (Proxy Admin or Admin Viewer) can view CloudZero settings. """ # Validation — Admin Viewer follows the read-parity rule. - if not _user_has_admin_view(user_api_key_dict): + if not user_api_key_has_admin_view(user_api_key_dict): raise HTTPException( status_code=403, detail={"error": CommonProxyErrors.not_allowed_access.value}, diff --git a/litellm/proxy/spend_tracking/spend_management_endpoints.py b/litellm/proxy/spend_tracking/spend_management_endpoints.py index 73fa375d01b..3194d390f17 100644 --- a/litellm/proxy/spend_tracking/spend_management_endpoints.py +++ b/litellm/proxy/spend_tracking/spend_management_endpoints.py @@ -449,7 +449,7 @@ async def spend_key_fn( "Database not connected. Connect a database to your proxy - https://docs.litellm.ai/docs/simple_proxy#managing-auth---virtual-keys" ) - if _is_admin_view_safe(user_api_key_dict=user_api_key_dict): + if is_admin_view_safe(user_api_key_dict=user_api_key_dict): return await prisma_client.get_data(table_name="key", query_type="find_all") caller_user_id: Final = user_api_key_dict.user_id @@ -522,7 +522,7 @@ async def spend_user_fn( "Database not connected. Connect a database to your proxy - https://docs.litellm.ai/docs/simple_proxy#managing-auth---virtual-keys" ) - if not _is_admin_view_safe(user_api_key_dict=user_api_key_dict): + if not is_admin_view_safe(user_api_key_dict=user_api_key_dict): caller_user_id: Final = user_api_key_dict.user_id if not caller_user_id: return [] @@ -1244,7 +1244,7 @@ async def get_spend_capture_rate( """ from litellm.proxy.proxy_server import prisma_client - if not _is_admin_view_safe(user_api_key_dict): + if not is_admin_view_safe(user_api_key_dict): raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Only proxy admins can read the capture rate") if prisma_client is None: raise HTTPException( @@ -1863,7 +1863,7 @@ def _resolve_spend_report_scope( viewers) may request any scope. """ if requested: - if requested != caller_value and not _is_admin_view_safe(user_api_key_dict=user_api_key_dict): + if requested != caller_value and not is_admin_view_safe(user_api_key_dict=user_api_key_dict): raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, detail=f"Not authorized to view spend for a {scope_name} other than your own", @@ -1887,7 +1887,7 @@ async def _resolve_org_spend_report_scope( Callable by proxy admins (any organization) and org admins of the target organization; every other caller is a 403 from ``_verify_org_access``. """ - from litellm.proxy.management_endpoints.organization_endpoints import _verify_org_access + from litellm.proxy.management_endpoints.organization_endpoints import verify_org_access target_org = organization_id or user_api_key_dict.org_id if target_org is None: @@ -1895,7 +1895,7 @@ async def _resolve_org_spend_report_scope( status_code=status.HTTP_400_BAD_REQUEST, detail="No organization_id associated with this API key; pass an organization_id query param", ) - await _verify_org_access( + await verify_org_access( organization_id=target_org, user_api_key_dict=user_api_key_dict, prisma_client=prisma_client, @@ -2212,10 +2212,10 @@ async def global_view_spend_tags( ) -async def _get_spend_report_for_time_range( +async def get_spend_report_for_time_range( start_date: str, end_date: str, -): +) -> tuple[Sequence[_TeamSpendRow] | None, Sequence[_TagSpendRow] | None] | None: from litellm.proxy.proxy_server import prisma_client if prisma_client is None: @@ -2274,6 +2274,9 @@ async def _get_spend_report_for_time_range( verbose_proxy_logger.error("Exception in _get_daily_spend_reports %s", e) +_get_spend_report_for_time_range: Final = get_spend_report_for_time_range + + @router.post( "/spend/calculate", tags=["Budget & Spend Tracking"], @@ -2652,7 +2655,7 @@ async def ui_view_spend_logs( ) try: - is_admin_view: Final = _is_admin_view_safe(user_api_key_dict=user_api_key_dict) + is_admin_view: Final = is_admin_view_safe(user_api_key_dict=user_api_key_dict) is_request_id_lookup: Final = request_id is not None and not is_v2 is_search_lookup: Final = search is not None search_owns_window: Final = is_search_lookup and not is_v2 @@ -3387,7 +3390,7 @@ async def ui_view_request_response_for_request_id( """ from litellm.proxy.proxy_server import prisma_client - caller_is_admin: Final = _is_admin_view_safe(user_api_key_dict=user_api_key_dict) + caller_is_admin: Final = is_admin_view_safe(user_api_key_dict=user_api_key_dict) if not caller_is_admin: if prisma_client is None: raise HTTPException( @@ -4532,7 +4535,7 @@ async def ui_view_session_spend_logs( read_scope: Final = ( AllRows() - if _is_admin_view_safe(user_api_key_dict=user_api_key_dict) + if is_admin_view_safe(user_api_key_dict=user_api_key_dict) else await _spend_log_read_scope(user_api_key_dict, log_team_lookup) if _can_user_view_spend_log(user_api_key_dict=user_api_key_dict) else OwnedRows(user_api_key_dict.user_id) @@ -4793,7 +4796,7 @@ def _span_type_sql_condition(span_type: str | None) -> str | None: return _SPAN_TYPE_SQL_CONDITIONS.get(span_type) -def _is_admin_view_safe(user_api_key_dict: UserAPIKeyAuth) -> bool: +def is_admin_view_safe(user_api_key_dict: UserAPIKeyAuth) -> bool: """ Safely determine if the current user has admin view permissions. Defaults to False on any exception. @@ -4810,6 +4813,9 @@ def _is_admin_view_safe(user_api_key_dict: UserAPIKeyAuth) -> bool: return False +_is_admin_view_safe: Final = is_admin_view_safe + + async def _can_team_member_view_log( prisma_client: PrismaClient, user_api_key_dict: UserAPIKeyAuth, diff --git a/litellm/proxy/spend_tracking/spend_tracking_utils.py b/litellm/proxy/spend_tracking/spend_tracking_utils.py index b71b834a31c..9d52f7d71bb 100644 --- a/litellm/proxy/spend_tracking/spend_tracking_utils.py +++ b/litellm/proxy/spend_tracking/spend_tracking_utils.py @@ -86,7 +86,7 @@ def _get_max_string_length_prompt_in_db() -> int: return DEFAULT_MAX_STRING_LENGTH_PROMPT_IN_DB -def _is_master_key(api_key: str | None, _master_key: str | None) -> bool: +def is_master_key(api_key: str | None, _master_key: str | None) -> bool: """ Raw-only constant-time master-key comparison. The hashed form is never considered equivalent — only the raw master-key string matches. @@ -96,6 +96,9 @@ def _is_master_key(api_key: str | None, _master_key: str | None) -> bool: return secrets.compare_digest(api_key, _master_key) +_is_master_key: Final = is_master_key + + _HASHED_JWT_RE = re.compile(r"hashed-jwt-[a-fA-F0-9]{64}") _NON_SECRET_KEY_ALIASES: Final = frozenset( { @@ -1411,7 +1414,7 @@ def _redact_prompt_fields_in_guardrail_entry( return {**redacted, "guardrail_response": preserved_stats} -def _sanitize_error_information_for_spend_logs( +def sanitize_error_information_for_spend_logs( error_information: StandardLoggingPayloadErrorInformation | None, original_exception: BaseException | None = None, ) -> StandardLoggingPayloadErrorInformation | None: @@ -1452,6 +1455,9 @@ def _sanitize_error_information_for_spend_logs( return cast(StandardLoggingPayloadErrorInformation, sanitized) +_sanitize_error_information_for_spend_logs: Final = sanitize_error_information_for_spend_logs + + def _convert_to_json_serializable_dict(obj: object, visited: set[int] | None = None, max_depth: int = 20) -> object: """ Convert object to JSON-serializable dict, handling Pydantic models safely. diff --git a/litellm/proxy/spend_tracking/vantage_endpoints.py b/litellm/proxy/spend_tracking/vantage_endpoints.py index 8ea28f42cd9..d3858f95243 100644 --- a/litellm/proxy/spend_tracking/vantage_endpoints.py +++ b/litellm/proxy/spend_tracking/vantage_endpoints.py @@ -18,7 +18,10 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import ( decrypt_value_helper, encrypt_value_helper, ) -from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view +from litellm.proxy.management_endpoints.common_utils import ( # noqa: F401 # legacy module exports + _user_has_admin_view, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + user_api_key_has_admin_view, +) from litellm.repositories.config_repository import ConfigRepository from litellm.repositories.prisma_protocols import TableActions from litellm.types.proxy.vantage_endpoints import ( @@ -160,7 +163,7 @@ async def get_vantage_settings( Only admin users (Proxy Admin or Admin Viewer) can view Vantage settings. """ # Admin Viewer follows the read-parity rule. - if not _user_has_admin_view(user_api_key_dict): + if not user_api_key_has_admin_view(user_api_key_dict): raise HTTPException( status_code=403, detail={"error": CommonProxyErrors.not_allowed_access.value}, diff --git a/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py b/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py index d307805427a..c410e66f0e3 100644 --- a/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py +++ b/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py @@ -15,7 +15,7 @@ from typing import ( from urllib.parse import urlparse from fastapi import APIRouter, Body, Depends, File, HTTPException, UploadFile -from pydantic import BaseModel, ConfigDict, JsonValue, TypeAdapter, ValidationError, create_model +from pydantic import BaseModel, ConfigDict, JsonValue, TypeAdapter, ValidationError, create_model, field_validator from pydantic.fields import FieldInfo, PydanticUndefined from typing_extensions import NotRequired, ReadOnly, TypedDict @@ -234,6 +234,20 @@ class UIThemeSettingsResponse(SettingsResponse): _TEAM_ADMIN_FIELD_ENUM: Final = tuple(sorted(SUPPORTED_TEAM_ADMIN_PERMISSIONS)) +def normalize_moyai_url(value: object) -> str | None: + if value is None: + return None + if not isinstance(value, str): + raise ValueError("moyai_url must be a string") + stripped: Final = value.strip() + if not stripped: + return None + parsed: Final = urlparse(stripped) + if parsed.scheme not in ("http", "https") or not parsed.hostname or parsed.username or parsed.password: + raise ValueError("moyai_url must be an http or https URL with a host and no credentials") + return stripped.rstrip("/") + + class UISettings(LiteLLMBaseModel): """Configuration for UI-specific flags""" @@ -330,6 +344,16 @@ class UISettings(LiteLLMBaseModel): description="If true, shows the Chat page in the UI sidebar, letting users chat with an LLM and connect their own MCP server credentials via OAuth.", ) + moyai_url: str | None = Field( + default=None, + description="URL of a connected Moyai deployment. When set, the Moyai entry in the UI navigation opens this deployment instead of the Moyai landing page.", + ) + + @field_validator("moyai_url", mode="before") + @classmethod + def _validate_moyai_url(cls, value: object) -> object: + return normalize_moyai_url(value) + team_admin_editable_team_fields: Sequence[str] = Field( default=(), description=( @@ -368,6 +392,7 @@ ALLOWED_UI_SETTINGS_FIELDS: Final = { "disable_custom_api_keys", "disable_key_generate_for_org_admin", "enable_chat_ui", + "moyai_url", TEAM_ADMIN_EDITABLE_TEAM_FIELDS_SETTING, } @@ -1237,7 +1262,9 @@ async def update_sso_settings( if isinstance(stored, str): stored = json.loads(stored) if isinstance(stored, dict): - before_sso_data = proxy_config._decrypt_db_variables(stored) + before_sso_data = ( # rebind-ok: pre-existing rebinding on a rename-only line + proxy_config.decrypt_db_variables(stored) + ) # Load existing config config: Final = await proxy_config.get_config() @@ -1261,7 +1288,7 @@ async def update_sso_settings( # Clear environment variable if value is null/empty os.environ.pop(env_var_name, None) - encrypted_sso_data: Final = proxy_config._encrypt_env_variables(environment_variables=sso_data) + encrypted_sso_data: Final = proxy_config.encrypt_env_variables(environment_variables=sso_data) # Save to dedicated SSO table await _stored_sso_settings_db(SSOConfigRepository(prisma_client)).upsert( @@ -1529,7 +1556,7 @@ async def update_mcp_semantic_filter_settings( from litellm.proxy.proxy_server import prisma_client, proxy_config if prisma_client is not None: - await proxy_config._init_semantic_filter_settings_in_db(prisma_client=prisma_client) + await proxy_config.init_semantic_filter_settings_in_db(prisma_client=prisma_client) except Exception as e: verbose_proxy_logger.warning("Failed to reinitialize MCP semantic filter settings immediately: %s", e) diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 9ed73d3f0bc..614a9b19d8a 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -47,7 +47,7 @@ from typing import ( runtime_checkable, ) -from typing_extensions import ReadOnly, TypedDict +from typing_extensions import Never, ReadOnly, TypedDict from litellm import _custom_logger_compatible_callbacks_literal from litellm.constants import ( @@ -190,7 +190,11 @@ from litellm.proxy.db.health_check_latest import ( fetch_latest_health_checks, fetch_latest_health_checks_for_models, ) -from litellm.proxy.db.log_db_metrics import _is_exception_related_to_db, log_db_metrics +from litellm.proxy.db.log_db_metrics import ( # noqa: F401, RUF100 # legacy module exports + _is_exception_related_to_db, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + is_exception_related_to_db, + log_db_metrics, +) from litellm.proxy.db.pgbouncer import database_url_is_pooled from litellm.proxy.db.prisma_client import ( PrismaWrapper, @@ -213,15 +217,21 @@ from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrai resolve_endpoint_translation, ) from litellm.proxy.hooks import PROXY_HOOKS, get_proxy_hook -from litellm.proxy.hooks.cache_control_check import _PROXY_CacheControlCheck -from litellm.proxy.hooks.parallel_request_limiter import ( - _PROXY_MaxParallelRequestsHandler, +from litellm.proxy.hooks.cache_control_check import ( # noqa: F401, RUF100 # legacy module exports + PROXY_CacheControlCheck, + _PROXY_CacheControlCheck, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export ) -from litellm.proxy.hooks.parallel_request_limiter_v3 import ( - _PROXY_MaxParallelRequestsHandler_v3, +from litellm.proxy.hooks.parallel_request_limiter import ( # noqa: F401, RUF100 # legacy module exports + PROXY_MaxParallelRequestsHandler, + _PROXY_MaxParallelRequestsHandler, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export ) -from litellm.proxy.hooks.sensitive_data_routing import ( - _PROXY_SensitiveDataRoutingHandler, +from litellm.proxy.hooks.parallel_request_limiter_v3 import ( # noqa: F401, RUF100 # legacy module exports + PROXY_MaxParallelRequestsHandler_v3, + _PROXY_MaxParallelRequestsHandler_v3, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export +) +from litellm.proxy.hooks.sensitive_data_routing import ( # noqa: F401, RUF100 # legacy module exports + PROXY_SensitiveDataRoutingHandler, + _PROXY_SensitiveDataRoutingHandler, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export ) from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup, add_guardrails_from_auth_metadata from litellm.proxy.management_helpers.key_settings_audit import with_settings_updated_at @@ -1219,8 +1229,8 @@ class ProxyLogging: self.file_usage_cache: Final = InternalUsageCache( dual_cache=DualCache(in_memory_cache=InMemoryCache(max_size_in_memory=FILE_USAGE_MAX_TRACKED_COUNTERS)) ) - self.max_parallel_request_limiter = _PROXY_MaxParallelRequestsHandler(self.internal_usage_cache) - self.cache_control_check = _PROXY_CacheControlCheck() + self.max_parallel_request_limiter = PROXY_MaxParallelRequestsHandler(self.internal_usage_cache) + self.cache_control_check = PROXY_CacheControlCheck() self.alerting: list[str] | None = None self.alerting_threshold: float = 300 # default to 5 min. threshold self.alert_types: list[AlertType] = DEFAULT_ALERT_TYPES @@ -1554,6 +1564,13 @@ class ProxyLogging: ] return synthetic_data + def convert_mcp_to_llm_format( + self, + request_obj: MCPPreCallRequestObject, + kwargs: Mapping[str, object], + ) -> dict[str, object]: + return self._convert_mcp_to_llm_format(request_obj, kwargs) + def _convert_llm_result_to_mcp_response(self, llm_result, request_obj) -> MCPPreCallResponseObject | None: """ Convert LLM guardrail result back to MCP response format. @@ -1742,7 +1759,7 @@ class ProxyLogging: } return result - def _create_mcp_request_object_from_kwargs(self, kwargs: dict) -> "MCPPreCallRequestObject": + def create_mcp_request_object_from_kwargs(self, kwargs: dict) -> "MCPPreCallRequestObject": """ Helper function to create MCPPreCallRequestObject from kwargs for standard pre_call_hook. """ @@ -1761,7 +1778,9 @@ class ProxyLogging: hidden_params=HiddenParams(), ) - def _convert_mcp_hook_response_to_kwargs(self, response_data: dict | None, original_kwargs: dict) -> dict: + _create_mcp_request_object_from_kwargs = create_mcp_request_object_from_kwargs + + def convert_mcp_hook_response_to_kwargs(self, response_data: dict | None, original_kwargs: dict) -> dict: """ Helper function to convert pre_call_hook response back to kwargs for MCP usage. @@ -1790,6 +1809,8 @@ class ProxyLogging: return modified_kwargs + _convert_mcp_hook_response_to_kwargs = convert_mcp_hook_response_to_kwargs + async def process_pre_call_hook_response(self, response, data, call_type): if isinstance(response, Exception): raise response @@ -2353,7 +2374,7 @@ class ProxyLogging: """ if request_metadata.get("_guardrail_pipelines"): return True - caps: Final = ProxyLogging._callback_capabilities() + caps: Final = ProxyLogging.callback_capabilities() if caps.has_content_enforcer: return True probe: Final = {"metadata": dict(request_metadata)} @@ -2453,7 +2474,7 @@ class ProxyLogging: # otherwise makes deep copies return the original object. needs_raw_request_snapshot: Final = any( isinstance(cb, CustomGuardrail) and cb.scan_raw_request - for cb in ProxyLogging._callback_capabilities().resolved_callbacks + for cb in ProxyLogging.callback_capabilities().resolved_callbacks ) raw_request_snapshot: Final[dict | None] = independent_snapshot(data) if needs_raw_request_snapshot else None @@ -2472,7 +2493,7 @@ class ProxyLogging: frozenset() if skip_guardrails else pipeline_managed_guardrail_names(data, "pre_call") ) - caps: Final = ProxyLogging._callback_capabilities() + caps: Final = ProxyLogging.callback_capabilities() # Skip the per-request callback walk entirely when nothing in # ``litellm.callbacks`` overrides ``async_pre_call_hook`` and no # CustomGuardrail is configured. Saves the loop overhead + @@ -2540,7 +2561,7 @@ class ProxyLogging: call_type=call_type, endpoint_type=endpoint_type, ) - if isinstance(_callback, _PROXY_MaxParallelRequestsHandler_v3) + if isinstance(_callback, PROXY_MaxParallelRequestsHandler_v3) else await _callback.async_pre_call_hook( user_api_key_dict=user_api_key_dict, cache=self.call_details["user_api_key_cache"], @@ -2690,7 +2711,7 @@ class ProxyLogging: if exc.sticky_session_routing: sensitive_routing_hook: Final = self.get_proxy_hook("sensitive_data_routing") - if isinstance(sensitive_routing_hook, _PROXY_SensitiveDataRoutingHandler): + if isinstance(sensitive_routing_hook, PROXY_SensitiveDataRoutingHandler): await sensitive_routing_hook.set_session_routing( session_id=exc.session_id, model=exc.route_to_model, @@ -2810,7 +2831,7 @@ class ProxyLogging: _callback_capabilities_cache: ClassVar[dict[tuple[int, tuple[int, ...]], "_CallbackCapabilities"]] = {} @staticmethod - def _callback_capabilities() -> "_CallbackCapabilities": + def callback_capabilities() -> "_CallbackCapabilities": """ Inspect ``litellm.callbacks`` once and answer the per-hook capability questions used to short-circuit no-op work on the chat-completions hot @@ -2912,6 +2933,8 @@ class ProxyLogging: cache[sig] = caps return caps + _callback_capabilities = callback_capabilities + @staticmethod def _stream_requires_guardrail_translation(user_api_key_dict: UserAPIKeyAuth) -> bool: route: Final = user_api_key_dict.request_route @@ -2924,18 +2947,18 @@ class ProxyLogging: @staticmethod def has_post_call_response_headers_callbacks() -> bool: - return ProxyLogging._callback_capabilities().has_post_call_response_headers + return ProxyLogging.callback_capabilities().has_post_call_response_headers @staticmethod def has_streaming_callbacks() -> bool: - caps: Final = ProxyLogging._callback_capabilities() + caps: Final = ProxyLogging.callback_capabilities() return caps.has_iterator_override or caps.has_streaming_chunk_override or caps.has_guardrail @staticmethod def has_streaming_chunk_hook_overrides() -> bool: """True iff any callback overrides ``async_post_call_streaming_hook`` (the per-chunk hook, distinct from the iterator wrapper).""" - caps: Final = ProxyLogging._callback_capabilities() + caps: Final = ProxyLogging.callback_capabilities() return caps.has_streaming_chunk_override or caps.has_guardrail def needs_iterator_wrap(self) -> bool: @@ -2943,19 +2966,19 @@ class ProxyLogging: through ``async_post_call_streaming_iterator_hook``. Instance method so tests can override the gate via ``MagicMock(spec=ProxyLogging)``. """ - return ProxyLogging._callback_capabilities().has_iterator_override + return ProxyLogging.callback_capabilities().has_iterator_override def needs_per_chunk_streaming_hook(self) -> bool: """Whether ``async_data_generator`` needs to call the per-chunk ``_apply_streaming_chunk_hooks`` for every emitted chunk. Instance method for the same reason as :py:meth:`needs_iterator_wrap`. """ - caps: Final = ProxyLogging._callback_capabilities() + caps: Final = ProxyLogging.callback_capabilities() return caps.has_streaming_chunk_override or caps.has_guardrail @staticmethod def has_during_call_guardrails() -> bool: - return ProxyLogging._callback_capabilities().has_guardrail + return ProxyLogging.callback_capabilities().has_guardrail async def during_call_hook( self, @@ -2963,7 +2986,7 @@ class ProxyLogging: user_api_key_dict: UserAPIKeyAuth | None, call_type: CallTypesLiteral, ): - caps: Final = ProxyLogging._callback_capabilities() + caps: Final = ProxyLogging.callback_capabilities() if not caps.has_guardrail and not caps.has_moderation_override: return data # Step 1: Collect all guardrail tasks to run in parallel @@ -3218,7 +3241,7 @@ class ProxyLogging: ) ) - logged_by_decorator: Final = call_type in _LOG_DB_METRICS_CALL_TYPES and _is_exception_related_to_db( + logged_by_decorator: Final = call_type in _LOG_DB_METRICS_CALL_TYPES and is_exception_related_to_db( original_exception ) if hasattr(self, "service_logging_obj") and not logged_by_decorator: @@ -3560,7 +3583,7 @@ class ProxyLogging: guardrail_callbacks, other_callbacks = _partition_post_call_callbacks() try: # Merge model-level guardrails before checking which guardrails to run - guardrail_data: Final = _check_and_merge_model_level_guardrails(data=data, llm_router=llm_router) + guardrail_data: Final = check_and_merge_model_level_guardrails(data=data, llm_router=llm_router) parallel_guardrails: Final[tuple[CustomGuardrail, ...]] = tuple( callback @@ -3720,7 +3743,7 @@ class ProxyLogging: (matching the inbound ``pre_mcp_call`` behavior) rather than being swallowed into an unguarded result. """ - caps: Final = ProxyLogging._callback_capabilities() + caps: Final = ProxyLogging.callback_capabilities() if not caps.has_guardrail: return response @@ -3773,7 +3796,7 @@ class ProxyLogging: # cached detection makes the redundant interior guard cheap, but the # guard would still iterate every code path through this function so # keep it cheap and rely on the cached capability lookup. - if not ProxyLogging._callback_capabilities().has_post_call_response_headers: + if not ProxyLogging.callback_capabilities().has_post_call_response_headers: return merged_headers try: @@ -3815,7 +3838,7 @@ class ProxyLogging: async def hidden_by_listing_callbacks( self, user_api_key_dict: UserAPIKeyAuth, model_names: Sequence[str] ) -> frozenset[str]: - filters: Final = ProxyLogging._callback_capabilities().listed_models_filters + filters: Final = ProxyLogging.callback_capabilities().listed_models_filters if not filters: return frozenset() candidates: Final = tuple(model_names) @@ -3870,7 +3893,7 @@ class ProxyLogging: # active. ``get_response_string`` walks every choice/delta on the # chunk so paying it per chunk for no-op callbacks dominated stream # CPU time even after the iterator-chain fix. - caps: Final = ProxyLogging._callback_capabilities() + caps: Final = ProxyLogging.callback_capabilities() if not caps.has_streaming_chunk_override and not caps.has_guardrail: return response @@ -3903,7 +3926,7 @@ class ProxyLogging: ## CHECK FOR MODEL-LEVEL GUARDRAILS (cached per-request) if not _guardrail_data_computed: - _cached_guardrail_data = _check_and_merge_model_level_guardrails( + _cached_guardrail_data = check_and_merge_model_level_guardrails( data=data, llm_router=llm_router ) _guardrail_data_computed = True @@ -3953,7 +3976,7 @@ class ProxyLogging: Covers: 1. /chat/completions """ - caps: Final = ProxyLogging._callback_capabilities() + caps: Final = ProxyLogging.callback_capabilities() post_call_pipelines: Final = _streamable_post_call_pipelines(request_data, user_api_key_dict) # Fast path: no real overrides. Internal proxy CustomLogger callbacks # (e.g. _PROXY_CacheControlCheck, ManagedFiles) inherit the default @@ -3968,15 +3991,15 @@ class ProxyLogging: raise except Exception as e: if not ProxyLogging._discard_deferred_stream_logging_for_failure(request_data, e): - ProxyLogging._fire_deferred_stream_logging(request_data) + ProxyLogging.fire_deferred_stream_logging(request_data) raise - ProxyLogging._fire_deferred_stream_logging(request_data) + ProxyLogging.fire_deferred_stream_logging(request_data) return from litellm.proxy.proxy_server import llm_router # Merge model-level guardrails before checking which guardrails to run - request_data = _check_and_merge_model_level_guardrails(data=request_data, llm_router=llm_router) + request_data = check_and_merge_model_level_guardrails(data=request_data, llm_router=llm_router) current_response = response stream_needs_translation: Final = ProxyLogging._stream_requires_guardrail_translation(user_api_key_dict) @@ -4052,7 +4075,7 @@ class ProxyLogging: except Exception as e: ProxyLogging._record_served_stream_output(request_data, served_chunks) if not ProxyLogging._discard_deferred_stream_logging_for_failure(request_data, e): - ProxyLogging._fire_deferred_stream_logging(request_data) + ProxyLogging.fire_deferred_stream_logging(request_data) raise # Fire deferred logging AFTER all guardrail end-of-stream blocks @@ -4060,7 +4083,7 @@ class ProxyLogging: # its end-of-stream block (inside current_response), so by the time # we reach this point the metadata is fully populated. ProxyLogging._record_served_stream_output(request_data, served_chunks) - ProxyLogging._fire_deferred_stream_logging(request_data) + ProxyLogging.fire_deferred_stream_logging(request_data) async def _pipeline_gated_stream( self, @@ -4145,7 +4168,7 @@ class ProxyLogging: record_served_output_texts(logging_obj.model_call_details, served_stream_output_texts(served_chunks)) @staticmethod - def _fire_deferred_stream_logging(request_data: dict) -> None: + def fire_deferred_stream_logging(request_data: dict) -> None: """ Fire the deferred streaming logging callback after the full streaming pipeline (including guardrail end-of-stream blocks) has completed. @@ -4166,6 +4189,8 @@ class ProxyLogging: logging_obj._deferred_stream_complete_args = None asyncio.create_task(_deferred_cb(*_args)) + _fire_deferred_stream_logging = fire_deferred_stream_logging + @staticmethod def _discard_deferred_stream_logging_for_failure(request_data: Mapping[str, object], error: Exception) -> bool: """Drop the parked success dispatch for an assembled chat stream that ends in an error @@ -4183,7 +4208,7 @@ class ProxyLogging: logging_obj.record_assembled_response_for_failure(assembled) return True - async def _arelease_max_parallel_requests_on_disconnect( + async def arelease_max_parallel_requests_on_disconnect( self, user_api_key_dict: UserAPIKeyAuth, ) -> None: @@ -4203,17 +4228,19 @@ class ProxyLogging: double-decrement under the limiter's in-memory fallback. """ limiter: Final = self.get_proxy_hook("parallel_request_limiter") - if not isinstance(limiter, _PROXY_MaxParallelRequestsHandler_v3): + if not isinstance(limiter, PROXY_MaxParallelRequestsHandler_v3): return await limiter.async_release_max_parallel_requests_on_disconnect(user_api_key_dict) + _arelease_max_parallel_requests_on_disconnect = arelease_max_parallel_requests_on_disconnect + async def enforce_mcp_server_rate_limits( self, user_api_key_dict: UserAPIKeyAuth | None, server: "MCPServer", ) -> None: limiter: Final = self.get_proxy_hook("parallel_request_limiter") - if not isinstance(limiter, _PROXY_MaxParallelRequestsHandler_v3): + if not isinstance(limiter, PROXY_MaxParallelRequestsHandler_v3): return await limiter.enforce_mcp_server_rate_limits(user_api_key_dict, server) @@ -4960,7 +4987,9 @@ class PrismaClient: # check if plain text or hash if token is not None: if isinstance(token, str): - hashed_token = _hash_token_if_needed(token=token) + hashed_token = hash_token_if_needed( # rebind-ok: pre-existing rebinding on a rename-only line + token=token + ) verbose_proxy_logger.debug("PrismaClient: find_unique for token: %s", hashed_token) if query_type == "find_unique" and hashed_token is not None: if token is None: @@ -5022,7 +5051,7 @@ class PrismaClient: if token is not None: where_filter["token"] = {} if isinstance(token, str): - token = _hash_token_if_needed(token=token) + token = hash_token_if_needed(token=token) where_filter["token"]["in"] = [token] elif isinstance(token, list): hashed_tokens: Final[list[str]] = [] @@ -5192,7 +5221,9 @@ class PrismaClient: # check if plain text or hash if token is not None: if isinstance(token, str): - hashed_token = _hash_token_if_needed(token=token) + hashed_token = hash_token_if_needed( # rebind-ok: pre-existing rebinding on a rename-only line + token=token + ) verbose_proxy_logger.debug("PrismaClient: find_unique for token: %s", hashed_token) if query_type == "find_unique": if token is None: @@ -5525,7 +5556,7 @@ class PrismaClient: if token is not None: print_verbose(f"token: [set={token is not None}]") # check if plain text or hash - token = _hash_token_if_needed(token=token) + token = hash_token_if_needed(token=token) db_data["token"] = token include_object_permission: Final[LiteLLM_VerificationTokenInclude] = {"object_permission": True} response: Final = await VerificationTokenRepository(self).table.update( @@ -5873,7 +5904,7 @@ class PrismaClient: prisma_obj: Final = self.writer_db._original_prisma if prisma_obj.is_connected() is not True: return 0 - engine: Final = prisma_obj._engine + engine: Final = prisma_obj._engine # pyright: ignore[reportPrivateUsage] # Prisma engine internals process: Final = getattr(engine, "process", None) if engine is not None else None if process is not None: pid: Final[object] = process.pid @@ -7067,7 +7098,7 @@ class PrismaClient: ### HELPER FUNCTIONS ### -async def _cache_user_row(user_id: str, cache: DualCache, db: PrismaClient): +async def cache_user_row(user_id: str, cache: DualCache, db: PrismaClient) -> None: """ Check if a user_id exists in cache, if not retrieve it. @@ -7083,6 +7114,9 @@ async def _cache_user_row(user_id: str, cache: DualCache, db: PrismaClient): cache.set_cache(key=cache_key, value=cache_value, ttl=600) # store for 10 minutes +_cache_user_row: Final = cache_user_row + + def _should_use_smtp_ssl(smtp_port: int) -> bool: """ Port 465 expects an immediate TLS handshake (implicit SSL), so a plain @@ -7259,7 +7293,7 @@ async def migrate_passwords_to_scrypt_async(prisma_client) -> str: return f"Migrated {len(plaintext_users)} plaintext passwords to scrypt" -def _hash_token_if_needed(token: str) -> str: +def hash_token_if_needed(token: str) -> str: """ Hash the token if it's a string and starts with "sk-" @@ -7271,6 +7305,9 @@ def _hash_token_if_needed(token: str) -> str: return token +_hash_token_if_needed: Final = hash_token_if_needed + + async def enqueue_spend_logs( prisma_client: PrismaClient, logs: Sequence[Mapping[str, object]], @@ -7377,7 +7414,7 @@ class ProxyUpdateSpend: break except Exception as e: - await DBSpendUpdateWriter._handle_spend_update_failure( + await DBSpendUpdateWriter.handle_spend_update_failure( e=e, attempt=i, n_retry_times=n_retry_times, @@ -7472,7 +7509,7 @@ class ProxyUpdateSpend: raise await asyncio.sleep(1 << i) except Exception as e: - _raise_failed_update_spend_exception(e=e, start_time=start_time, proxy_logging_obj=proxy_logging_obj) + raise_failed_update_spend_exception(e=e, start_time=start_time, proxy_logging_obj=proxy_logging_obj) finally: # Clean up logs_to_process only if we popped it (caller-owned otherwise) if popped_batch: @@ -7631,14 +7668,14 @@ async def update_daily_tag_spend( """ n_retry_times: Final = 3 try: - if proxy_logging_obj.db_spend_update_writer.redis_update_buffer._should_commit_spend_updates_to_redis(): - await proxy_logging_obj.db_spend_update_writer._commit_daily_tag_spend_to_db_with_redis( + if proxy_logging_obj.db_spend_update_writer.redis_update_buffer.should_commit_spend_updates_to_redis(): + await proxy_logging_obj.db_spend_update_writer.commit_daily_tag_spend_to_db_with_redis( prisma_client=prisma_client, n_retry_times=n_retry_times, proxy_logging_obj=proxy_logging_obj, ) else: - await proxy_logging_obj.db_spend_update_writer._commit_daily_tag_spend_to_db( + await proxy_logging_obj.db_spend_update_writer.commit_daily_tag_spend_to_db( prisma_client=prisma_client, n_retry_times=n_retry_times, proxy_logging_obj=proxy_logging_obj, @@ -7875,11 +7912,11 @@ async def _park_remaining_spend_logs(prisma_client: PrismaClient, proxy_logging_ ) -async def _monitor_spend_logs_queue( +async def monitor_spend_logs_queue( prisma_client: PrismaClient, db_writer_client: AsyncHTTPHandler | None, proxy_logging_obj: ProxyLogging, -): +) -> Never: """ Background task that monitors the spend_log_transactions queue size and triggers processing when the threshold is reached. @@ -7951,6 +7988,9 @@ async def _monitor_spend_logs_queue( await asyncio.sleep(current_interval) +_monitor_spend_logs_queue: Final = monitor_spend_logs_queue + + MAX_SPEND_LOG_ISOLATION_FAILURES_PER_BATCH: Final = 256 @@ -8024,7 +8064,7 @@ async def _create_spend_logs_with_poison_isolation( return await _create_spend_logs_with_poison_isolation(repo, rows[mid:], remaining) -def _raise_failed_update_spend_exception(e: Exception, start_time: float, proxy_logging_obj: ProxyLogging): +def raise_failed_update_spend_exception(e: Exception, start_time: float, proxy_logging_obj: ProxyLogging) -> Never: """ Raise an exception for failed update spend logs @@ -8048,13 +8088,16 @@ def _raise_failed_update_spend_exception(e: Exception, start_time: float, proxy_ raise e +_raise_failed_update_spend_exception: Final = raise_failed_update_spend_exception + + def _get_month_end_date(today: date) -> date: if today.month == 12: return date(today.year + 1, 1, 1) - timedelta(days=1) return date(today.year, today.month + 1, 1) - timedelta(days=1) -def _is_projected_spend_over_limit(current_spend: float, soft_budget_limit: float | None): +def is_projected_spend_over_limit(current_spend: float, soft_budget_limit: float | None) -> bool: if soft_budget_limit is None: # If there's no limit, we can't exceed it. return False @@ -8081,7 +8124,10 @@ def _is_projected_spend_over_limit(current_spend: float, soft_budget_limit: floa return False -def _get_projected_spend_over_limit(current_spend: float, soft_budget_limit: float | None) -> tuple | None: +_is_projected_spend_over_limit: Final = is_projected_spend_over_limit + + +def get_projected_spend_over_limit(current_spend: float, soft_budget_limit: float | None) -> tuple | None: if soft_budget_limit is None: return None @@ -8113,7 +8159,10 @@ def _get_projected_spend_over_limit(current_spend: float, soft_budget_limit: flo return None -def _is_valid_team_configs(team_id=None, team_config=None, request_data=None): +_get_projected_spend_over_limit: Final = get_projected_spend_over_limit + + +def is_valid_team_configs(team_id=None, team_config=None, request_data=None) -> None: if team_id is None or team_config is None or request_data is None: return # check if valid model called for team @@ -8127,11 +8176,14 @@ def _is_valid_team_configs(team_id=None, team_config=None, request_data=None): return +_is_valid_team_configs: Final = is_valid_team_configs + + def _to_ns(dt): return int(dt.timestamp() * 1e9) -def _check_and_merge_model_level_guardrails( +def check_and_merge_model_level_guardrails( data: dict, llm_router: Router | None, trust_client_model_info: bool = True, @@ -8218,6 +8270,9 @@ def _check_and_merge_model_level_guardrails( return _merge_guardrails_with_existing(data, model_level_guardrails) +_check_and_merge_model_level_guardrails: Final = check_and_merge_model_level_guardrails + + def _merge_guardrails_with_existing(data: dict, model_level_guardrails: object) -> dict: """ Merge model-level guardrails with any existing guardrails in the request data. @@ -8266,7 +8321,7 @@ def get_error_message_str(e: Exception) -> str: return error_message -def _get_redoc_url() -> str | None: +def get_redoc_url() -> str | None: """ Get the Redoc URL from the environment variables. @@ -8283,7 +8338,10 @@ def _get_redoc_url() -> str | None: return "/redoc" -def _get_docs_url() -> str | None: +_get_redoc_url: Final = get_redoc_url + + +def get_docs_url() -> str | None: """ Get the docs (Swagger UI) URL from the environment variables. @@ -8300,7 +8358,10 @@ def _get_docs_url() -> str | None: return "/" -def _get_openapi_url() -> str | None: +_get_docs_url: Final = get_docs_url + + +def get_openapi_url() -> str | None: """ Get the OpenAPI JSON URL from the environment variables. @@ -8317,6 +8378,9 @@ def _get_openapi_url() -> str | None: return "/openapi.json" +_get_openapi_url: Final = get_openapi_url + + def _recreate_writer_on_read_only_transaction(prisma_client: "PrismaClient | None") -> None: if prisma_client is None: return @@ -8378,7 +8442,9 @@ def require_enterprise_license(feature: str | None = None) -> None: ) -_premium_user_check: Final = require_enterprise_license +premium_user_check: Final = require_enterprise_license + +_premium_user_check: Final = premium_user_check def is_known_model(model: str | None, llm_router: Router | None) -> bool: @@ -8607,11 +8673,11 @@ async def _get_access_group_models( proxy_logging_obj: Optional["ProxyLogging"], ) -> tuple[str, ...]: from litellm.proxy.auth.auth_checks import ( - _get_models_from_access_groups, get_authorized_resources_from_key_access_groups, + get_models_from_access_groups, ) - team_group_models: Final = await _get_models_from_access_groups( + team_group_models: Final = await get_models_from_access_groups( access_group_ids=(team_object.access_group_ids or ()) if team_object is not None else (), prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, diff --git a/litellm/proxy/vector_store_endpoints/endpoints.py b/litellm/proxy/vector_store_endpoints/endpoints.py index c2203bf7f67..22721dc26bc 100644 --- a/litellm/proxy/vector_store_endpoints/endpoints.py +++ b/litellm/proxy/vector_store_endpoints/endpoints.py @@ -118,11 +118,11 @@ async def vector_store_search( https://platform.openai.com/docs/api-reference/vector-stores/search """ from litellm.proxy.proxy_server import ( - _read_request_body, general_settings, llm_router, proxy_config, proxy_logging_obj, + read_request_body, select_data_generator, user_api_base, user_max_tokens, @@ -132,7 +132,7 @@ async def vector_store_search( version, ) - data = await _read_request_body(request=request) + data = await read_request_body(request=request) # rebind-ok: pre-existing rebinding on a rename-only line reject_caller_embedding_selection_params(payload=data, source="the search request body") data["vector_store_id"] = vector_store_id @@ -168,7 +168,7 @@ async def vector_store_search( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -198,11 +198,11 @@ async def vector_store_create( ``` """ from litellm.proxy.proxy_server import ( - _read_request_body, general_settings, llm_router, proxy_config, proxy_logging_obj, + read_request_body, select_data_generator, user_api_base, user_max_tokens, @@ -212,7 +212,7 @@ async def vector_store_create( version, ) - data: Final = await _read_request_body(request=request) + data: Final = await read_request_body(request=request) # Check for target_model_names parameter target_model_names: Final = data.pop("target_model_names", None) @@ -275,7 +275,7 @@ async def vector_store_create( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -338,7 +338,7 @@ async def vector_store_retrieve( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -408,7 +408,7 @@ async def vector_store_list( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -431,11 +431,11 @@ async def vector_store_update( https://platform.openai.com/docs/api-reference/vector-stores/modify """ from litellm.proxy.proxy_server import ( - _read_request_body, general_settings, llm_router, proxy_config, proxy_logging_obj, + read_request_body, select_data_generator, user_api_base, user_max_tokens, @@ -445,7 +445,7 @@ async def vector_store_update( version, ) - data = await _read_request_body(request=request) + data = await read_request_body(request=request) # rebind-ok: pre-existing rebinding on a rename-only line if "vector_store_id" not in data: data["vector_store_id"] = vector_store_id @@ -474,7 +474,7 @@ async def vector_store_update( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -537,7 +537,7 @@ async def vector_store_delete( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, diff --git a/litellm/proxy/vector_store_files_endpoints/endpoints.py b/litellm/proxy/vector_store_files_endpoints/endpoints.py index e44aa022334..96eda99425d 100644 --- a/litellm/proxy/vector_store_files_endpoints/endpoints.py +++ b/litellm/proxy/vector_store_files_endpoints/endpoints.py @@ -564,11 +564,11 @@ async def vector_store_file_create( user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): from litellm.proxy.proxy_server import ( - _read_request_body, general_settings, llm_router, proxy_config, proxy_logging_obj, + read_request_body, select_data_generator, user_api_base, user_max_tokens, @@ -578,7 +578,7 @@ async def vector_store_file_create( version, ) - data = await _read_request_body(request=request) + data = await read_request_body(request=request) # rebind-ok: pre-existing rebinding on a rename-only line data["vector_store_id"] = vector_store_id managed_vector_store: Final = await assert_user_can_access_vector_store_id( vector_store_id=vector_store_id, @@ -644,7 +644,7 @@ async def vector_store_file_create( return response except Exception as e: # noqa: BLE001 - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -751,7 +751,7 @@ async def vector_store_file_list( user_api_key_dict=user_api_key_dict, ) except Exception as e: # noqa: BLE001 - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -858,7 +858,7 @@ async def vector_store_file_retrieve( return response except Exception as e: # noqa: BLE001 - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -968,7 +968,7 @@ async def vector_store_file_content( return response except Exception as e: # noqa: BLE001 - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -996,11 +996,11 @@ async def vector_store_file_update( user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): from litellm.proxy.proxy_server import ( - _read_request_body, general_settings, llm_router, proxy_config, proxy_logging_obj, + read_request_body, select_data_generator, user_api_base, user_max_tokens, @@ -1010,7 +1010,7 @@ async def vector_store_file_update( version, ) - data = await _read_request_body(request=request) + data = await read_request_body(request=request) # rebind-ok: pre-existing rebinding on a rename-only line data["vector_store_id"] = vector_store_id data["file_id"] = file_id managed_vector_store: Final = await assert_user_can_access_vector_store_id( @@ -1075,7 +1075,7 @@ async def vector_store_file_update( return response except Exception as e: # noqa: BLE001 - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -1182,7 +1182,7 @@ async def vector_store_file_delete( return response except Exception as e: # noqa: BLE001 - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, diff --git a/litellm/proxy/vertex_ai_endpoints/langfuse_endpoints.py b/litellm/proxy/vertex_ai_endpoints/langfuse_endpoints.py index 523b669280e..724fb128a9a 100644 --- a/litellm/proxy/vertex_ai_endpoints/langfuse_endpoints.py +++ b/litellm/proxy/vertex_ai_endpoints/langfuse_endpoints.py @@ -21,8 +21,14 @@ import litellm from litellm.litellm_core_utils.url_utils import SSRFError, validate_url from litellm.proxy._types import * from litellm.proxy.auth.user_api_key_auth import user_api_key_auth -from litellm.proxy.common_utils.http_parsing_utils import _safe_get_request_headers -from litellm.proxy.litellm_pre_call_utils import _get_dynamic_logging_metadata +from litellm.proxy.common_utils.http_parsing_utils import ( # noqa: F401 # legacy module exports + _safe_get_request_headers, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + safe_get_request_headers, +) +from litellm.proxy.litellm_pre_call_utils import ( # noqa: F401 # legacy module exports + _get_dynamic_logging_metadata, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + get_dynamic_logging_metadata, +) from litellm.proxy.pass_through_endpoints.pass_through_endpoints import ( create_pass_through_route, ) @@ -36,7 +42,7 @@ def create_request_copy(request: Request): return { "method": request.method, "url": str(request.url), - "headers": _safe_get_request_headers(request).copy(), + "headers": safe_get_request_headers(request).copy(), "cookies": request.cookies, "query_params": dict(request.query_params), } @@ -175,7 +181,7 @@ async def langfuse_proxy_route( user_api_key_dict: Final = await user_api_key_auth(request=request, api_key=f"Bearer {api_key}") - callback_settings_obj: Final[TeamCallbackMetadata | None] = _get_dynamic_logging_metadata( + callback_settings_obj: Final[TeamCallbackMetadata | None] = get_dynamic_logging_metadata( user_api_key_dict=user_api_key_dict, proxy_config=proxy_config ) diff --git a/litellm/proxy/video_endpoints/endpoints.py b/litellm/proxy/video_endpoints/endpoints.py index 9175d067920..2813333f113 100644 --- a/litellm/proxy/video_endpoints/endpoints.py +++ b/litellm/proxy/video_endpoints/endpoints.py @@ -9,7 +9,10 @@ from starlette.datastructures import UploadFile as StarletteUploadFile from litellm.proxy._types import * from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth, user_api_key_auth from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing -from litellm.proxy.common_utils.http_parsing_utils import _read_request_body +from litellm.proxy.common_utils.http_parsing_utils import ( # noqa: F401 # legacy module exports + _read_request_body, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + read_request_body, +) from litellm.proxy.common_utils.openai_endpoint_utils import ( get_custom_llm_provider_from_request_body, get_custom_llm_provider_from_request_headers, @@ -80,7 +83,7 @@ async def video_generation( ) # Read request body - data: Final = await _read_request_body(request=request) + data: Final = await read_request_body(request=request) if input_reference is not None: input_reference_file: Final = await batch_to_bytesio([input_reference]) if input_reference_file: @@ -108,7 +111,7 @@ async def video_generation( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -195,7 +198,7 @@ async def video_list( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -295,7 +298,7 @@ async def video_status( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -402,7 +405,7 @@ async def video_content( headers={"Content-Disposition": f"attachment; filename=video_{video_id}.mp4"}, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -458,7 +461,7 @@ async def video_remix( version, ) - data: Final = await _read_request_body(request=request) + data: Final = await read_request_body(request=request) data["video_id"] = video_id decoded: Final = decode_video_id_with_provider(video_id) @@ -503,7 +506,7 @@ async def video_remix( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -560,7 +563,7 @@ async def video_create_character( version, ) - data: Final = await _read_request_body(request=request) + data: Final = await read_request_body(request=request) video_file: Final = await batch_to_bytesio([video]) if video_file: data["video"] = video_file[0] @@ -608,7 +611,7 @@ async def video_create_character( ) return response except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -714,7 +717,7 @@ async def video_get_character( ) return response except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -767,7 +770,7 @@ async def video_edit( version, ) - data: Final = await _read_request_body(request=request) + data: Final = await read_request_body(request=request) uploaded_video: Final = data.pop("video", None) if isinstance(uploaded_video, StarletteUploadFile): video_files: Final = await batch_to_bytesio((uploaded_video,)) @@ -816,7 +819,7 @@ async def video_edit( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -871,7 +874,7 @@ async def video_extension( version, ) - data: Final = await _read_request_body(request=request) + data: Final = await read_request_body(request=request) data["video_id"] = video_reference_to_id(data.pop("video", None)) decoded: Final = decode_video_id_with_provider(data["video_id"]) @@ -913,7 +916,7 @@ async def video_extension( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, diff --git a/litellm/realtime_api/main.py b/litellm/realtime_api/main.py index 83cacf6c4d9..0e79cbfda13 100644 --- a/litellm/realtime_api/main.py +++ b/litellm/realtime_api/main.py @@ -36,7 +36,7 @@ from ..litellm_core_utils.litellm_logging import Logging as LiteLLMLogging from ..llms.azure.common_utils import get_azure_ad_token from ..llms.azure.realtime.handler import AzureOpenAIRealtime, azure_realtime_protocol_for_client from ..llms.bedrock.realtime.handler import BedrockRealtime -from ..llms.custom_httpx.http_handler import get_shared_realtime_ssl_context +from ..llms.custom_httpx.http_handler import realtime_ssl_for_url from ..llms.openai.realtime.handler import OpenAIRealtime from ..llms.vertex_ai.audio_transcription.realtime_transformation import is_vertex_speech_to_text_model from ..llms.vertex_ai.realtime.transformation import VertexAIRealtimeConfig, vertex_realtime_config @@ -721,23 +721,21 @@ async def realtime_health_check( location=resolved_location, ) url = vertex_realtime_config.get_complete_url(api_base=resolved_api_base, model=model) - vertex_ssl_context: Final = get_shared_realtime_ssl_context() headers: Final = vertex_realtime_config.validate_environment(headers={}, model=model, api_key=None) async with websockets.connect( url, additional_headers=headers, max_size=REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES, - ssl=vertex_ssl_context, + ssl=realtime_ssl_for_url(url), ): return True else: raise ValueError(f"Unsupported model: {model}") - ssl_context: Final = get_shared_realtime_ssl_context() async with websockets.connect( url, additional_headers=auth_headers, max_size=REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES, - ssl=ssl_context, + ssl=realtime_ssl_for_url(url), ): return True diff --git a/litellm/responses/litellm_completion_transformation/streaming_iterator.py b/litellm/responses/litellm_completion_transformation/streaming_iterator.py index 7286563f36f..693061d0a3a 100644 --- a/litellm/responses/litellm_completion_transformation/streaming_iterator.py +++ b/litellm/responses/litellm_completion_transformation/streaming_iterator.py @@ -1,6 +1,7 @@ import time import uuid from collections.abc import Sequence +from itertools import filterfalse from typing import Any, Final, cast import litellm @@ -81,6 +82,19 @@ def _delta_has_signed_thinking_block(delta: object) -> bool: return any(isinstance(b, dict) and (b.get("signature") or b.get("data")) for b in blocks) +def _delta_carries_output(delta: ChatCompletionDelta) -> bool: + fields: Final = ( + delta.content, + getattr(delta, "reasoning_content", None), + delta.tool_calls, + delta.function_call, + getattr(delta, "annotations", None), + getattr(delta, "images", None), + getattr(delta, "audio", None), + ) + return any(fields) or _delta_has_signed_thinking_block(delta) + + class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): """ Async iterator for processing streaming responses from the Responses API. @@ -119,7 +133,8 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): self.completed_response = None self.final_text: str = "" self._cached_item_id: str | None = None - self._message_output_index: int = 0 + self._message_output_index: int | None = None + self._reasoning_output_index: int | None = None self._cached_response_id: str | None = None self._buffered_chunk: ModelResponseStream | None = None self._upstream_exhausted: bool = False @@ -130,7 +145,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): self._tool_item_id_by_call_id: dict[str, str] = {} self._tool_call_id_by_index: dict[int, str] = {} self._ambiguous_tool_call_indexes: set[int] = set() - self._next_tool_output_index: int = 1 # output_index=0 reserved for the message item + self._next_output_index: int = 0 self._final_tool_events_queued: bool = False self._sequence_number: int = 0 self._cached_reasoning_item_id: str | None = None @@ -155,11 +170,39 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): existing: Final = self._tool_output_index_by_call_id.get(call_id) if existing is not None: return existing - idx: Final = self._next_tool_output_index - self._next_tool_output_index += 1 + idx: Final = self._allocate_output_index() self._tool_output_index_by_call_id[call_id] = idx return idx + def _allocate_output_index(self) -> int: + idx: Final = self._next_output_index + self._next_output_index += 1 + return idx + + def _message_index(self) -> int: + if self._message_output_index is None: + self._message_output_index = self._allocate_output_index() + return self._message_output_index + + def _reasoning_index(self) -> int: + if self._reasoning_output_index is None: + self._reasoning_output_index = self._allocate_output_index() + return self._reasoning_output_index + + def _streamed_output_index(self, item: object) -> int | None: + match getattr(item, "type", None): + case "message": + return self._message_output_index + case "reasoning": + return self._reasoning_output_index + case _: + call_id: Final = getattr(item, "call_id", None) or str(getattr(item, "id", "")).removeprefix("ws_") + return self._tool_output_index_by_call_id.get(str(call_id)) + + def _streamed_output_position(self, item: object) -> tuple[bool, int]: + index: Final = self._streamed_output_index(item) + return (index is None, index or 0) + def _normalize_tool_call_index(self, tool_call: object) -> int | None: idx_raw: Final = tool_call.get("index") if isinstance(tool_call, dict) else getattr(tool_call, "index", None) if idx_raw is None: @@ -572,7 +615,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): self._sequence_number += 1 event: Final = OutputItemAddedEvent( type=ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED, - output_index=self._message_output_index, + output_index=self._message_index(), item=BaseLiteLLMOpenAIResponseObject( **{ "id": self._cached_item_id, @@ -594,7 +637,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): event: Final = ContentPartAddedEvent( type=ResponsesAPIStreamEvents.CONTENT_PART_ADDED, item_id=self._cached_item_id, - output_index=self._message_output_index, + output_index=self._message_index(), content_index=0, part=BaseLiteLLMOpenAIResponseObject(**{"type": "output_text", "text": "", "annotations": []}), ) @@ -606,15 +649,10 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): self._cached_item_id = f"msg_{uuid.uuid4()}" self.sent_message_item_added_event = True self.sent_content_part_added_event = True - if self._cached_reasoning_item_id is not None: - self._message_output_index = self._next_tool_output_index - self._next_tool_output_index += 1 - else: - self._message_output_index = 0 self._sequence_number += 1 event: Final = OutputItemAddedEvent( type=ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED, - output_index=self._message_output_index, + output_index=self._message_index(), item=BaseLiteLLMOpenAIResponseObject( **{ "id": self._cached_item_id, @@ -699,7 +737,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): return ReasoningSummaryTextDoneEvent( type=ResponsesAPIStreamEvents.REASONING_SUMMARY_TEXT_DONE, item_id=reasoning_item_id, - output_index=0, + output_index=self._reasoning_index(), sequence_number=sequence_number, summary_index=0, text=reasoning_content, @@ -730,7 +768,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): return ReasoningSummaryPartDoneEvent( type=ResponsesAPIStreamEvents.REASONING_SUMMARY_PART_DONE, item_id=reasoning_item_id, - output_index=0, + output_index=self._reasoning_index(), sequence_number=sequence_number, summary_index=0, part=BaseLiteLLMOpenAIResponseObject( @@ -748,7 +786,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): return OutputTextDoneEvent( type=ResponsesAPIStreamEvents.OUTPUT_TEXT_DONE, item_id=self._cached_item_id, - output_index=self._message_output_index, + output_index=self._message_index(), content_index=0, text=getattr(litellm_complete_object.choices[0].message, "content", "") or "", ) @@ -775,7 +813,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): return ContentPartDoneEvent( type=ResponsesAPIStreamEvents.CONTENT_PART_DONE, item_id=self._cached_item_id, - output_index=self._message_output_index, + output_index=self._message_index(), content_index=0, part=part, ) @@ -794,7 +832,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): ) return OutputItemDoneEvent( type=ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE, - output_index=self._message_output_index, + output_index=self._message_index(), sequence_number=1, item=BaseLiteLLMOpenAIResponseObject( **{ @@ -841,7 +879,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): """ return OutputItemDoneEvent( type=ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE, - output_index=0, + output_index=self._reasoning_index(), sequence_number=sequence_number, item=BaseLiteLLMOpenAIResponseObject( **{ @@ -939,6 +977,8 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): if not chunk.choices: return delta: Final = chunk.choices[0].delta + if chunk.choices[0].finish_reason is None and not _delta_carries_output(delta): + return self._sequence_number += 1 self.sent_output_item_added_event = True @@ -952,7 +992,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): event = OutputItemAddedEvent( type=ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED, - output_index=0, + output_index=self._reasoning_index(), item=BaseLiteLLMOpenAIResponseObject( **{ "id": self._cached_reasoning_item_id, @@ -1176,7 +1216,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): event = OutputTextAnnotationAddedEvent( type=ResponsesAPIStreamEvents.OUTPUT_TEXT_ANNOTATION_ADDED, item_id=item_id, - output_index=self._message_output_index, + output_index=self._message_index(), content_index=0, annotation_index=idx, annotation=annotation_dict, @@ -1196,7 +1236,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): return ReasoningSummaryTextDeltaEvent( type=ResponsesAPIStreamEvents.REASONING_SUMMARY_TEXT_DELTA, item_id=self._cached_reasoning_item_id, - output_index=0, + output_index=self._reasoning_index(), delta=reasoning_content, ) @@ -1209,7 +1249,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): text_delta_event: Final = OutputTextDeltaEvent( type=ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA, item_id=item_id, - output_index=self._message_output_index, + output_index=self._message_index(), content_index=0, delta=delta_content, ) @@ -1266,7 +1306,13 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): "reasoning", self._cached_reasoning_item_id, ) - return reasoning_aligned + streamed_items: Final = filterfalse(self._is_unstreamed_empty_message, reasoning_aligned) + return tuple(sorted(streamed_items, key=self._streamed_output_position)) + + def _is_unstreamed_empty_message(self, item: object) -> bool: + if getattr(item, "type", None) != "message" or self._message_output_index is not None: + return False + return not any(getattr(part, "text", None) for part in getattr(item, "content", None) or ()) def _emit_terminal_response_event( self, litellm_model_response: ModelResponse diff --git a/litellm/responses/mcp/litellm_proxy_mcp_handler.py b/litellm/responses/mcp/litellm_proxy_mcp_handler.py index 3293851fd9f..d3e3509b8f0 100644 --- a/litellm/responses/mcp/litellm_proxy_mcp_handler.py +++ b/litellm/responses/mcp/litellm_proxy_mcp_handler.py @@ -298,10 +298,10 @@ class LiteLLM_Proxy_MCP_Handler: granted_toolset_ids, ) from litellm.proxy.management_endpoints.common_utils import ( - _user_has_admin_view, + user_api_key_has_admin_view, ) - if not _user_has_admin_view(user_api_key_auth) and toolset.toolset_id not in ( + if not user_api_key_has_admin_view(user_api_key_auth) and toolset.toolset_id not in ( await (granted_toolsets or granted_toolset_ids)(user_api_key_auth) ): verbose_logger.debug("Key does not have access to toolset '%s', skipping.", name) @@ -764,7 +764,7 @@ class LiteLLM_Proxy_MCP_Handler: mcp_server = global_mcp_server_manager.get_mcp_server_by_name( server_name - ) or global_mcp_server_manager._get_mcp_server_from_tool_name(tool_name) + ) or global_mcp_server_manager.get_mcp_server_from_tool_name(tool_name) resolved_tool_name = ( _resolve_display_name_to_original(tool_name, [mcp_server]) if mcp_server else tool_name ) diff --git a/litellm/responses/mcp/mcp_streaming_iterator.py b/litellm/responses/mcp/mcp_streaming_iterator.py index 830c2e094bc..75f272c0f79 100644 --- a/litellm/responses/mcp/mcp_streaming_iterator.py +++ b/litellm/responses/mcp/mcp_streaming_iterator.py @@ -376,9 +376,9 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): if raw_headers_from_request: headers_obj: Final = Headers(raw_headers_from_request) - self.mcp_auth_header = MCPRequestHandler._get_mcp_auth_header_from_headers(headers_obj) - self.mcp_server_auth_headers = MCPRequestHandler._get_mcp_server_auth_headers_from_headers(headers_obj) - self.oauth2_headers = MCPRequestHandler._get_oauth2_headers_from_headers(headers_obj) + self.mcp_auth_header = MCPRequestHandler.get_mcp_auth_header_from_headers(headers_obj) + self.mcp_server_auth_headers = MCPRequestHandler.get_mcp_server_auth_headers_from_headers(headers_obj) + self.oauth2_headers = MCPRequestHandler.get_oauth2_headers_from_headers(headers_obj) # Also check if headers are provided in tools array (from request body) tools: Final[Sequence[object] | None] = self.original_request_params.get("tools") @@ -389,7 +389,7 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): if tool_headers and isinstance(tool_headers, dict): # Merge tool headers into mcp_server_auth_headers headers_obj_from_tool = Headers(tool_headers) - tool_mcp_server_auth_headers = MCPRequestHandler._get_mcp_server_auth_headers_from_headers( + tool_mcp_server_auth_headers = MCPRequestHandler.get_mcp_server_auth_headers_from_headers( headers_obj_from_tool ) diff --git a/litellm/responses/utils.py b/litellm/responses/utils.py index 1b1d166b689..bc77b21b54b 100644 --- a/litellm/responses/utils.py +++ b/litellm/responses/utils.py @@ -1121,9 +1121,15 @@ class ResponsesAPIRequestUtils: if raw_headers_from_request: headers_obj: Final = Headers(raw_headers_from_request) - mcp_auth_header = MCPRequestHandler._get_mcp_auth_header_from_headers(headers_obj) - mcp_server_auth_headers = MCPRequestHandler._get_mcp_server_auth_headers_from_headers(headers_obj) - oauth2_headers = MCPRequestHandler._get_oauth2_headers_from_headers(headers_obj) + mcp_auth_header = ( # rebind-ok: pre-existing rebinding on a rename-only line + MCPRequestHandler.get_mcp_auth_header_from_headers(headers_obj) + ) + mcp_server_auth_headers = ( # rebind-ok: pre-existing rebinding on a rename-only line + MCPRequestHandler.get_mcp_server_auth_headers_from_headers(headers_obj) + ) + oauth2_headers = ( # rebind-ok: pre-existing rebinding on a rename-only line + MCPRequestHandler.get_oauth2_headers_from_headers(headers_obj) + ) if tools: for tool in tools: @@ -1133,7 +1139,7 @@ class ResponsesAPIRequestUtils: # Merge tool headers into mcp_server_auth_headers # Extract server-specific headers from tool.headers headers_obj_from_tool = Headers(tool_headers) - tool_mcp_server_auth_headers = MCPRequestHandler._get_mcp_server_auth_headers_from_headers( + tool_mcp_server_auth_headers = MCPRequestHandler.get_mcp_server_auth_headers_from_headers( headers_obj_from_tool ) if tool_mcp_server_auth_headers: diff --git a/litellm/router.py b/litellm/router.py index 56c79b22283..d03e085cf1e 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -4534,61 +4534,21 @@ class Router: stream=False, **kwargs, ): - parent_otel_span: Final = get_parent_otel_span_from_kwargs(kwargs) - ### FLOW ITEM ### - _request_id: Final = str(uuid.uuid4()) - item: Final = FlowItem( - priority=priority, # 👈 SET PRIORITY FOR REQUEST - request_id=_request_id, # 👈 SET REQUEST ID - model_name=model, # 👈 SAME as 'Router' + await self._wait_for_scheduler_turn( + model=model, priority=priority, parent_otel_span=get_parent_otel_span_from_kwargs(kwargs) ) - ### [fin] ### - - ## ADDS REQUEST TO QUEUE ## - await self.scheduler.add_request(request=item) - - ## POLL QUEUE - end_time: Final = time.monotonic() + self.timeout - curr_time = time.monotonic() - poll_interval: Final = self.scheduler.polling_interval # poll every 3ms - make_request = False - - while curr_time < end_time: - _healthy_deployments, _ = await self._async_get_healthy_deployments( - model=model, parent_otel_span=parent_otel_span - ) - make_request = await self.scheduler.poll( ## POLL QUEUE ## - returns 'True' if there's healthy deployments OR if request is at top of queue - id=item.request_id, - model_name=item.model_name, - health_deployments=_healthy_deployments, - ) - if make_request: ## IF TRUE -> MAKE REQUEST - break - else: ## ELSE -> loop till default_timeout - await asyncio.sleep(poll_interval) - curr_time = time.monotonic() - - if make_request: - try: - _response: Final = await self.acompletion(model=model, messages=messages, stream=stream, **kwargs) - response_hidden_params: Final = get_hidden_params(_response) - if response_hidden_params is not None: - additional_headers: Final = cast( # cast-ok: router headers are stored as a mutable mapping - dict[str, object], response_hidden_params.setdefault("additional_headers", {}) - ) - additional_headers.update({"x-litellm-request-prioritization-used": True}) - return _response - except Exception as e: - setattr(e, "priority", priority) - raise e - else: - # Clean up the request from the scheduler queue also before raising the timeout exception - await self.scheduler.remove_request(request_id=item.request_id, model_name=item.model_name) - raise litellm.Timeout( - message="Request timed out while polling queue", - model=model, - llm_provider="openai", - ) + try: + _response: Final = await self.acompletion(model=model, messages=messages, stream=stream, **kwargs) + response_hidden_params: Final = get_hidden_params(_response) + if response_hidden_params is not None: + additional_headers: Final = cast( # cast-ok: router headers are stored as a mutable mapping + dict[str, object], response_hidden_params.setdefault("additional_headers", {}) + ) + additional_headers.update({"x-litellm-request-prioritization-used": True}) + return _response + except Exception as e: + setattr(e, "priority", priority) + raise e async def _schedule_factory( self, @@ -4598,61 +4558,32 @@ class Router: args: tuple[object, ...], kwargs: dict[str, object], ): - parent_otel_span: Final = get_parent_otel_span_from_kwargs(kwargs) - ### FLOW ITEM ### - _request_id: Final = str(uuid.uuid4()) - item: Final = FlowItem( - priority=priority, # 👈 SET PRIORITY FOR REQUEST - request_id=_request_id, # 👈 SET REQUEST ID - model_name=model, # 👈 SAME as 'Router' + await self._wait_for_scheduler_turn( + model=model, priority=priority, parent_otel_span=get_parent_otel_span_from_kwargs(kwargs) ) - ### [fin] ### + try: + _response: Final = await original_function(*args, **kwargs) + response_hidden_params: Final = get_hidden_params(_response) + if response_hidden_params is not None: + additional_headers: Final = cast( # cast-ok: router headers are stored as a mutable mapping + dict[str, object], response_hidden_params.setdefault("additional_headers", {}) + ) + additional_headers.update({"x-litellm-request-prioritization-used": True}) + return _response + except Exception as e: + setattr(e, "priority", priority) + raise e - ## ADDS REQUEST TO QUEUE ## - await self.scheduler.add_request(request=item) + async def _wait_for_scheduler_turn(self, model: str, priority: int, parent_otel_span: Span | None) -> None: + async def healthy_deployments() -> Sequence[object]: + deployments, _ = await self._async_get_healthy_deployments(model=model, parent_otel_span=parent_otel_span) + return deployments - ## POLL QUEUE - end_time: Final = time.monotonic() + self.timeout - curr_time = time.monotonic() - poll_interval: Final = self.scheduler.polling_interval # poll every 3ms - make_request = False - - while curr_time < end_time: - _healthy_deployments, _ = await self._async_get_healthy_deployments( - model=model, parent_otel_span=parent_otel_span - ) - make_request = await self.scheduler.poll( ## POLL QUEUE ## - returns 'True' if there's healthy deployments OR if request is at top of queue - id=item.request_id, - model_name=item.model_name, - health_deployments=_healthy_deployments, - ) - if make_request: ## IF TRUE -> MAKE REQUEST - break - else: ## ELSE -> loop till default_timeout - await asyncio.sleep(poll_interval) - curr_time = time.monotonic() - - if make_request: - try: - _response: Final = await original_function(*args, **kwargs) - response_hidden_params: Final = get_hidden_params(_response) - if response_hidden_params is not None: - additional_headers: Final = cast( # cast-ok: router headers are stored as a mutable mapping - dict[str, object], response_hidden_params.setdefault("additional_headers", {}) - ) - additional_headers.update({"x-litellm-request-prioritization-used": True}) - return _response - except Exception as e: - setattr(e, "priority", priority) - raise e - else: - # Clean up the request from the scheduler queue also before raising the timeout exception - await self.scheduler.remove_request(request_id=item.request_id, model_name=item.model_name) - raise litellm.Timeout( - message="Request timed out while polling queue", - model=model, - llm_provider="openai", - ) + await self.scheduler.wait_for_turn( + request=FlowItem(priority=priority, request_id=str(uuid.uuid4()), model_name=model), + timeout=self.timeout, + get_healthy_deployments=healthy_deployments, + ) def _is_prompt_management_model(self, model: str) -> bool: model_list: Final = self.get_model_list(model_name=model) diff --git a/litellm/rust_bridge/trace/generated/models.py b/litellm/rust_bridge/trace/generated/models.py index 7a748102991..c9600cc545a 100644 --- a/litellm/rust_bridge/trace/generated/models.py +++ b/litellm/rust_bridge/trace/generated/models.py @@ -192,6 +192,141 @@ class ExecutionRow(LiteLLMBaseModel): selection_key: str = "" +Score: TypeAlias = Annotated[ + int, + Field( + ..., + ge=0, + json_schema_extra={ + "x-python-normalized": { + "type": "int", + "minimum": 0, + "maximum": 18446744073709551615, + } + }, + le=18446744073709551615, + ), +] + + +Score1: TypeAlias = Annotated[ + str, + Field( + ..., + json_schema_extra={ + "x-python-normalized": { + "type": "int", + "minimum": 0, + "maximum": 18446744073709551615, + } + }, + pattern="^(?:0|[1-9][0-9]{0,18}|1[0-7][0-9]{18}|18[0-3][0-9]{17}|184[0-3][0-9]{16}|1844[0-5][0-9]{15}|18446[0-6][0-9]{14}|184467[0-3][0-9]{13}|1844674[0-3][0-9]{12}|184467440[0-6][0-9]{10}|1844674407[0-2][0-9]{9}|18446744073[0-6][0-9]{8}|1844674407370[0-8][0-9]{6}|18446744073709[0-4][0-9]{5}|184467440737095[0-4][0-9]{4}|1844674407370955[0-0][0-9]{3}|18446744073709551[0-5][0-9]{2}|184467440737095516[0-0][0-9]{1}|1844674407370955161[0-4][0-9]{0}|18446744073709551615)$", + ), +] + + +class FeedbackRow(LiteLLMBaseModel): + model_config = ConfigDict( + frozen=True, + ) + + trace_id: str + trace_ref: str + author: str + score: int = Field(..., ge=0, le=18446744073709551615) + comment: str + created_at: str + updated_at: str + + +Count2: TypeAlias = Annotated[ + int, + Field( + ..., + ge=0, + json_schema_extra={ + "x-python-normalized": { + "type": "int", + "minimum": 0, + "maximum": 18446744073709551615, + } + }, + le=18446744073709551615, + ), +] + + +Count3: TypeAlias = Annotated[ + str, + Field( + ..., + json_schema_extra={ + "x-python-normalized": { + "type": "int", + "minimum": 0, + "maximum": 18446744073709551615, + } + }, + pattern="^(?:0|[1-9][0-9]{0,18}|1[0-7][0-9]{18}|18[0-3][0-9]{17}|184[0-3][0-9]{16}|1844[0-5][0-9]{15}|18446[0-6][0-9]{14}|184467[0-3][0-9]{13}|1844674[0-3][0-9]{12}|184467440[0-6][0-9]{10}|1844674407[0-2][0-9]{9}|18446744073[0-6][0-9]{8}|1844674407370[0-8][0-9]{6}|18446744073709[0-4][0-9]{5}|184467440737095[0-4][0-9]{4}|1844674407370955[0-0][0-9]{3}|18446744073709551[0-5][0-9]{2}|184467440737095516[0-0][0-9]{1}|1844674407370955161[0-4][0-9]{0}|18446744073709551615)$", + ), +] + + +Lowest: TypeAlias = Annotated[ + int, + Field( + ..., + ge=0, + json_schema_extra={ + "x-python-normalized": { + "type": "int", + "minimum": 0, + "maximum": 18446744073709551615, + } + }, + le=18446744073709551615, + ), +] + + +Lowest1: TypeAlias = Annotated[ + str, + Field( + ..., + json_schema_extra={ + "x-python-normalized": { + "type": "int", + "minimum": 0, + "maximum": 18446744073709551615, + } + }, + pattern="^(?:0|[1-9][0-9]{0,18}|1[0-7][0-9]{18}|18[0-3][0-9]{17}|184[0-3][0-9]{16}|1844[0-5][0-9]{15}|18446[0-6][0-9]{14}|184467[0-3][0-9]{13}|1844674[0-3][0-9]{12}|184467440[0-6][0-9]{10}|1844674407[0-2][0-9]{9}|18446744073[0-6][0-9]{8}|1844674407370[0-8][0-9]{6}|18446744073709[0-4][0-9]{5}|184467440737095[0-4][0-9]{4}|1844674407370955[0-0][0-9]{3}|18446744073709551[0-5][0-9]{2}|184467440737095516[0-0][0-9]{1}|1844674407370955161[0-4][0-9]{0}|18446744073709551615)$", + ), +] + + +class FeedbackSummaryRow(LiteLLMBaseModel): + model_config = ConfigDict( + frozen=True, + ) + + trace_id: str + trace_ref: str + count: int = Field(..., ge=0, le=18446744073709551615) + average: float + lowest: int = Field(..., ge=0, le=18446744073709551615) + + +class FeedbackTargetRow(LiteLLMBaseModel): + model_config = ConfigDict( + frozen=True, + ) + + team_id: str + key_hash: str + trace_ref: str + + class LensAccessParams(LiteLLMBaseModel): model_config = ConfigDict( extra="forbid", @@ -239,6 +374,44 @@ class LensEvidenceParams(LiteLLMBaseModel): quote: str +class LensFeedbackParams(LiteLLMBaseModel): + model_config = ConfigDict( + extra="forbid", + frozen=True, + ) + + all_teams: Literal[0, 1] + team: str + key_hash: str + trace_id: str + trace_ref: str + + +class LensFeedbackSummaryParams(LiteLLMBaseModel): + model_config = ConfigDict( + extra="forbid", + frozen=True, + ) + + all_teams: Literal[0, 1] + team: str + key_hash: str + trace_ids: tuple[str, ...] + + +class LensFeedbackTargetParams(LiteLLMBaseModel): + model_config = ConfigDict( + extra="forbid", + frozen=True, + ) + + all_teams: Literal[0, 1] + team: str + key_hash: str + trace_id: str + trace_ref: str + + ExecutionSource: TypeAlias = Literal["traces", "requests", "both"] @@ -555,9 +728,15 @@ TraceWireModels: TypeAlias = Annotated[ | AgentRow | CountRow | ExecutionRow + | FeedbackRow + | FeedbackSummaryRow + | FeedbackTargetRow | LensAccessParams | LensContentParams | LensEvidenceParams + | LensFeedbackParams + | LensFeedbackSummaryParams + | LensFeedbackTargetParams | LensSampleParams | PartRow | TraceAgentRow diff --git a/litellm/rust_bridge/trace/generated/types.py b/litellm/rust_bridge/trace/generated/types.py index 4234cc5e6fa..16f22bc0b95 100644 --- a/litellm/rust_bridge/trace/generated/types.py +++ b/litellm/rust_bridge/trace/generated/types.py @@ -90,7 +90,17 @@ class TraceScope(typing_extensions.TypedDict): team_ids: ReadOnly[tuple[str, ...]] -ReadQueryName: TypeAlias = Literal["trace_agents", "availability", "agents", "sample", "content", "evidence"] +ReadQueryName: TypeAlias = Literal[ + "trace_agents", + "availability", + "agents", + "sample", + "content", + "evidence", + "feedback_target", + "feedback", + "feedback_summary", +] class UIFields(typing_extensions.TypedDict): @@ -109,6 +119,7 @@ class RunSource(typing_extensions.TypedDict): type: ReadOnly[RunSourceType] url: ReadOnly[str] title: ReadOnly[str] + user: ReadOnly[NotRequired[str]] class Span(typing_extensions.TypedDict): diff --git a/litellm/rust_bridge/trace/queries.py b/litellm/rust_bridge/trace/queries.py index 8d1ee978b2b..5bfc2303057 100644 --- a/litellm/rust_bridge/trace/queries.py +++ b/litellm/rust_bridge/trace/queries.py @@ -11,9 +11,15 @@ from .generated.models import ( AgentRow, CountRow, ExecutionRow, + FeedbackRow, + FeedbackSummaryRow, + FeedbackTargetRow, LensAccessParams, LensContentParams, LensEvidenceParams, + LensFeedbackParams, + LensFeedbackSummaryParams, + LensFeedbackTargetParams, LensSampleParams, PartRow, TraceAgentRow, @@ -74,3 +80,12 @@ LENS_CONTENT: Final[ReadQuery[LensContentParams, PartRow]] = ReadQuery( LENS_EVIDENCE: Final[ReadQuery[LensEvidenceParams, CountRow]] = ReadQuery( "evidence", LensEvidenceParams, TypeAdapter(QueryResponse[CountRow]) ) +LENS_FEEDBACK_TARGET: Final[ReadQuery[LensFeedbackTargetParams, FeedbackTargetRow]] = ReadQuery( + "feedback_target", LensFeedbackTargetParams, TypeAdapter(QueryResponse[FeedbackTargetRow]) +) +LENS_FEEDBACK: Final[ReadQuery[LensFeedbackParams, FeedbackRow]] = ReadQuery( + "feedback", LensFeedbackParams, TypeAdapter(QueryResponse[FeedbackRow]) +) +LENS_FEEDBACK_SUMMARY: Final[ReadQuery[LensFeedbackSummaryParams, FeedbackSummaryRow]] = ReadQuery( + "feedback_summary", LensFeedbackSummaryParams, TypeAdapter(QueryResponse[FeedbackSummaryRow]) +) diff --git a/litellm/scheduler.py b/litellm/scheduler.py index c88cc4ce0c0..73196d24e4d 100644 --- a/litellm/scheduler.py +++ b/litellm/scheduler.py @@ -1,14 +1,22 @@ +import asyncio import enum import heapq -from typing import Final +import time +from collections.abc import Awaitable, Callable, Sequence +from typing import Final, TypeAlias + +from pydantic import TypeAdapter from litellm import print_verbose from litellm._internal_context import with_service_target from litellm.caching.caching import DualCache, RedisCache from litellm.constants import DEFAULT_IN_MEMORY_TTL, DEFAULT_POLLING_INTERVAL +from litellm.exceptions import Timeout from litellm.types.llms.base import LiteLLMBaseModel SCHEDULER_QUEUE_TARGET: Final = "scheduler_queue" +QueueEntry: TypeAlias = tuple[int, str] +_QUEUE_ENTRIES: Final = TypeAdapter(list[QueueEntry]) class SchedulerCacheKeys(enum.Enum): @@ -33,7 +41,7 @@ class Scheduler: """ polling_interval: float or null - frequency of polling queue. Default is 3ms. """ - self.queue: list = [] + self.queue: list[QueueEntry] = [] default_in_memory_ttl: float | None = None if redis_cache is not None: # if redis-cache available frequently poll that instead of using in-memory. @@ -51,7 +59,7 @@ class Scheduler: # save the queue await self.save_queue(queue=queue, model_name=request.model_name) - async def poll(self, id: str, model_name: str, health_deployments: list) -> bool: + async def poll(self, request: FlowItem, health_deployments: Sequence[object]) -> bool: """ Return if request can be processed. @@ -62,30 +70,48 @@ class Scheduler: - False: * If no healthy deployments available * AND request not at the top of queue + + A request the queue no longer holds (its cache key expired or a concurrent writer erased the entry) + is put back at its priority so it keeps its place in the order instead of failing or jumping ahead """ - queue: Final = await self.get_queue(model_name=model_name) - if not queue: - raise Exception(f"Incorrectly setup. Queue is invalid. Queue={queue}") - - # ------------ - # Setup values - # ------------ - print_verbose(f"len(health_deployments): {len(health_deployments)}") - if len(health_deployments) == 0: - print_verbose(f"queue: {queue}, seeking id={id}") - # Check if the id is at the top of the heap - if queue[0][1] == id: - # Remove the item from the queue - heapq.heappop(queue) - await self.save_queue(queue=queue, model_name=model_name) - print_verbose(f"Popped id: {id}") - return True - else: - return False + if len(health_deployments) > 0: + return True + queue: Final = await self.get_queue(model_name=request.model_name) + entry: Final = (request.priority, request.request_id) + print_verbose(f"queue: {queue}, seeking {entry}") + if entry not in queue: + print_verbose(f"queue no longer holds {entry}, re-enqueueing it") + heapq.heappush(queue, entry) + if queue[0] != entry: + await self.save_queue(queue=queue, model_name=request.model_name) + return False + elif queue[0] != entry: + return False + + heapq.heappop(queue) + await self.save_queue(queue=queue, model_name=request.model_name) + print_verbose(f"Popped id: {request.request_id}") return True + async def wait_for_turn( + self, + request: FlowItem, + timeout: float, + get_healthy_deployments: Callable[[], Awaitable[Sequence[object]]], + ) -> None: + try: + await self.add_request(request=request) + end_time: Final = time.monotonic() + timeout + while time.monotonic() < end_time: + if await self.poll(request=request, health_deployments=await get_healthy_deployments()): + return + await asyncio.sleep(self.polling_interval) + finally: + await asyncio.shield(self.remove_request(request_id=request.request_id, model_name=request.model_name)) + raise Timeout(message="Request timed out while polling queue", model=request.model_name, llm_provider="openai") + async def remove_request(self, request_id: str, model_name: str) -> None: """ Remove a specific request from the priority queue for a model. @@ -118,21 +144,23 @@ class Scheduler: return self.queue @with_service_target(SCHEDULER_QUEUE_TARGET) - async def get_queue(self, model_name: str) -> list: + async def get_queue(self, model_name: str) -> list[QueueEntry]: """ - Return a queue for that specific model group + Return a queue for that specific model group. + + Redis hands the queue back as JSON lists, so every entry is validated into the + (priority, request_id) tuple the heap operations compare against. """ if self.cache is not None: _cache_key: Final = f"{SchedulerCacheKeys.queue.value}:{model_name}" response: Final = await self.cache.async_get_cache(key=_cache_key) - if response is None or not isinstance(response, list): + if not isinstance(response, list): return [] - elif isinstance(response, list): - return response + return _QUEUE_ENTRIES.validate_python(response) return self.queue @with_service_target(SCHEDULER_QUEUE_TARGET) - async def save_queue(self, queue: list, model_name: str) -> None: + async def save_queue(self, queue: list[QueueEntry], model_name: str) -> None: """ Save the updated queue of the model group """ diff --git a/litellm/tracing/remote.py b/litellm/tracing/remote.py index 5d71557b8b2..187300d6e23 100644 --- a/litellm/tracing/remote.py +++ b/litellm/tracing/remote.py @@ -16,6 +16,7 @@ from litellm.rust_bridge.trace.generated.types import QueryScope, ReadQueryName, MAX_RESPONSE_BYTES: Final = 64 * 1024 * 1024 _JSON: Final[TypeAdapter[JsonValue]] = TypeAdapter(JsonValue) +_INSERT_PATHS: Final[Mapping[str, str]] = {"spend_logs": "/internal/spend", "lens_feedback": "/internal/feedback"} @dataclass(frozen=True, slots=True, repr=False) @@ -125,9 +126,10 @@ class RemoteTraceStore: return _ReadFailure.INVALID_RESPONSE async def insert_rows(self, table: str, rows: Sequence[Mapping[str, object]]) -> None: - if table != "spend_logs": - raise ValueError("Lens only accepts gateway request records on this endpoint") - response: Final = await self.client.post("/internal/spend", json=tuple(dict(row) for row in rows)) + path: Final = _INSERT_PATHS.get(table) + if path is None: + raise ValueError("Lens only accepts gateway request records and feedback on this endpoint") + response: Final = await self.client.post(path, json=tuple(dict(row) for row in rows)) response.raise_for_status() async def ingest( diff --git a/litellm/tracing/types.py b/litellm/tracing/types.py index 809f264d371..43490f285c2 100644 --- a/litellm/tracing/types.py +++ b/litellm/tracing/types.py @@ -49,7 +49,7 @@ class SpendLogPayload(TypedDict, total=False): cache_hit: ReadOnly[bool | None] session_id: ReadOnly[str | None] trace_id: ReadOnly[str | None] - request_tags: ReadOnly[Sequence[str] | None] + request_tags: ReadOnly[Sequence[object] | None] messages: ReadOnly[object] response: ReadOnly[object] diff --git a/litellm/types/decisions.py b/litellm/types/decisions.py index a752ad88116..3e0dab8add5 100644 --- a/litellm/types/decisions.py +++ b/litellm/types/decisions.py @@ -1,5 +1,7 @@ from collections.abc import Mapping, Sequence -from typing import Annotated, Literal, TypeAlias +from dataclasses import dataclass, field +from types import MappingProxyType +from typing import Annotated, Final, Literal, TypeAlias from pydantic import ConfigDict, Field, PrivateAttr, model_validator, with_config from typing_extensions import ReadOnly, Required, TypedDict @@ -8,6 +10,7 @@ from litellm.types.llms.base import LiteLLMPydanticObjectBase DecisionsJSON: TypeAlias = str | Mapping[str, object] | Sequence[object] NoulCriteria: TypeAlias = Mapping[Literal["true", "false"], DecisionsJSON | None] +MAX_DECISION_QUESTIONS: Final = 128 class NoulQuestion(LiteLLMPydanticObjectBase): @@ -47,7 +50,7 @@ DecisionQuestion: TypeAlias = Annotated[ DecisionQuestionMap: TypeAlias = Annotated[ Mapping[Annotated[str, Field(min_length=1)], DecisionQuestion], - Field(min_length=1, max_length=128), + Field(min_length=1, max_length=MAX_DECISION_QUESTIONS), ] @@ -109,19 +112,339 @@ DecisionAnswer: TypeAlias = Annotated[ class DecisionsUsage(LiteLLMPydanticObjectBase): input_tokens: int = 0 output_tokens: int = 0 + cached_tokens: Annotated[int, Field(exclude=True)] = 0 + cache_write_tokens: Annotated[int, Field(exclude=True)] = 0 model_config = ConfigDict(extra="allow", frozen=True) -class DecisionsResponse(LiteLLMPydanticObjectBase): +class _HiddenParamsResponse(LiteLLMPydanticObjectBase): + _hidden_params: dict[str, object] = PrivateAttr(default_factory=dict) + + @property + def hidden_params(self) -> dict[str, object]: + return self._hidden_params + + def set_hidden_params(self, params: Mapping[str, object]) -> None: + self._hidden_params.update(params) + + +class DecisionsResponse(_HiddenParamsResponse): model: str | None = None answers: Mapping[str, DecisionAnswer] usage: DecisionsUsage | None = None model_config = ConfigDict(extra="allow", frozen=True) - _hidden_params: dict[str, object] = PrivateAttr(default_factory=dict) - @property - def hidden_params(self) -> dict[str, object]: - return self._hidden_params +class OpenAIDecisionInputText(LiteLLMPydanticObjectBase): + type: Literal["input_text"] + text: str + + model_config = ConfigDict(extra="forbid", frozen=True) + + +class OpenAIDecisionInputImage(LiteLLMPydanticObjectBase): + type: Literal["input_image"] + image_url: str + detail: str | None = None + + model_config = ConfigDict(extra="forbid", frozen=True) + + +OpenAIDecisionContentPart: TypeAlias = Annotated[ + OpenAIDecisionInputText | OpenAIDecisionInputImage, + Field(discriminator="type"), +] + + +class OpenAIDecisionInputMessage(LiteLLMPydanticObjectBase): + role: Literal["user"] = "user" + type: Literal["message"] = "message" + content: str | Sequence[OpenAIDecisionContentPart] + + model_config = ConfigDict(extra="forbid", frozen=True) + + +OpenAIDecisionInput: TypeAlias = str | Sequence[OpenAIDecisionInputMessage] + + +class OpenAIPredicateQuestion(LiteLLMPydanticObjectBase): + type: Literal["predicate"] + name: str | None = None + instructions: str + + model_config = ConfigDict(extra="forbid", frozen=True) + + +class OpenAIChoiceOption(LiteLLMPydanticObjectBase): + value: str | bool + description: str | None = None + + model_config = ConfigDict(extra="forbid", frozen=True) + + +def systemone_choice_key(value: str | bool) -> str: + if isinstance(value, bool): + return "true" if value else "false" + return value + + +class OpenAIChoiceQuestion(LiteLLMPydanticObjectBase): + type: Literal["choice"] + name: str | None = None + instructions: str + choices: Annotated[Sequence[OpenAIChoiceOption], Field(min_length=2, max_length=255)] + + model_config = ConfigDict(extra="forbid", frozen=True) + + @model_validator(mode="after") + def require_unique_systemone_keys(self) -> "OpenAIChoiceQuestion": + keys: Final = frozenset(systemone_choice_key(option.value) for option in self.choices) + if len(keys) != len(self.choices): + raise ValueError("Choice values must be unique, and a boolean cannot share its text with a string choice") + return self + + +class OpenAIScoreLevel(LiteLLMPydanticObjectBase): + label: str + description: str | None = None + + model_config = ConfigDict(extra="forbid", frozen=True) + + +class OpenAIScoreQuestion(LiteLLMPydanticObjectBase): + type: Literal["score"] + name: str | None = None + instructions: str + levels: Annotated[Sequence[OpenAIScoreLevel], Field(min_length=2, max_length=10)] + + model_config = ConfigDict(extra="forbid", frozen=True) + + +OpenAIDecisionQuestion: TypeAlias = Annotated[ + OpenAIPredicateQuestion | OpenAIChoiceQuestion | OpenAIScoreQuestion, + Field(discriminator="type"), +] + + +class OpenAIDecisionRequestBody(LiteLLMPydanticObjectBase): + input: OpenAIDecisionInput + questions: Annotated[Sequence[OpenAIDecisionQuestion], Field(min_length=1, max_length=MAX_DECISION_QUESTIONS)] + safety_identifier: str | None = None + + model_config = ConfigDict(extra="allow", frozen=True) + + +class OpenAIPredicateAnswer(LiteLLMPydanticObjectBase): + type: Literal["predicate"] = "predicate" + name: str | None + probability: float + + model_config = ConfigDict(frozen=True) + + +class OpenAIChoiceProbability(LiteLLMPydanticObjectBase): + value: str | bool + probability: float + + model_config = ConfigDict(frozen=True) + + +class OpenAIChoiceAnswer(LiteLLMPydanticObjectBase): + type: Literal["choice"] = "choice" + name: str | None + choice: str | bool + probabilities: tuple[OpenAIChoiceProbability, ...] + confidence: float + + model_config = ConfigDict(frozen=True) + + +class OpenAIScoreProbability(LiteLLMPydanticObjectBase): + value: int + label: str + probability: float + + model_config = ConfigDict(frozen=True) + + +class OpenAIScoreAnswer(LiteLLMPydanticObjectBase): + type: Literal["score"] = "score" + name: str | None + score: float + probabilities: tuple[OpenAIScoreProbability, ...] + confidence: float + + model_config = ConfigDict(frozen=True) + + +class OpenAIRefusalAnswer(LiteLLMPydanticObjectBase): + type: Literal["refusal"] = "refusal" + name: str | None + + model_config = ConfigDict(frozen=True) + + +OpenAIDecisionAnswer: TypeAlias = Annotated[ + OpenAIPredicateAnswer | OpenAIChoiceAnswer | OpenAIScoreAnswer | OpenAIRefusalAnswer, + Field(discriminator="type"), +] + + +class OpenAIDecisionInputTokensDetails(LiteLLMPydanticObjectBase): + cached_tokens: int = 0 + cache_write_tokens: int = 0 + + model_config = ConfigDict(frozen=True) + + +class OpenAIDecisionOutputTokensDetails(LiteLLMPydanticObjectBase): + reasoning_tokens: int = 0 + + model_config = ConfigDict(frozen=True) + + +class OpenAIDecisionUsage(LiteLLMPydanticObjectBase): + input_tokens: int + input_tokens_details: OpenAIDecisionInputTokensDetails = OpenAIDecisionInputTokensDetails() + output_tokens: int + output_tokens_details: OpenAIDecisionOutputTokensDetails = OpenAIDecisionOutputTokensDetails() + total_tokens: int + + model_config = ConfigDict(frozen=True) + + +class OpenAIDecisionResponse(_HiddenParamsResponse): + model: str + answers: tuple[OpenAIDecisionAnswer, ...] + usage: OpenAIDecisionUsage + + model_config = ConfigDict(extra="allow", frozen=True) + + +_NO_EXTRA: Final[Mapping[str, object]] = MappingProxyType({}) + + +@dataclass(frozen=True, slots=True) +class DecisionsIRState: + state: DecisionsJSON + + +@dataclass(frozen=True, slots=True) +class DecisionsIRMessages: + messages: tuple[OpenAIDecisionInputMessage, ...] + + +@dataclass(frozen=True, slots=True) +class DecisionsIRPredicateQuestion: + name: str | None + instructions: DecisionsJSON | None + criteria: NoulCriteria | None = None + extra: Mapping[str, object] = field(default_factory=lambda: _NO_EXTRA) + + +@dataclass(frozen=True, slots=True) +class DecisionsIRChoiceOption: + value: str | bool + description: DecisionsJSON | None + + +@dataclass(frozen=True, slots=True) +class DecisionsIRChoiceQuestion: + name: str | None + instructions: DecisionsJSON | None + choices: tuple[DecisionsIRChoiceOption, ...] + extra: Mapping[str, object] = field(default_factory=lambda: _NO_EXTRA) + + +@dataclass(frozen=True, slots=True) +class DecisionsIRScoreLevel: + label: DecisionsJSON + description: str | None + + +@dataclass(frozen=True, slots=True) +class DecisionsIRScoreQuestion: + name: str | None + instructions: DecisionsJSON | None + levels: tuple[DecisionsIRScoreLevel, ...] + extra: Mapping[str, object] = field(default_factory=lambda: _NO_EXTRA) + + +DecisionsIRQuestion: TypeAlias = DecisionsIRPredicateQuestion | DecisionsIRChoiceQuestion | DecisionsIRScoreQuestion + + +@dataclass(frozen=True, slots=True) +class DecisionsIRRequest: + input: DecisionsIRState | DecisionsIRMessages + questions: tuple[DecisionsIRQuestion, ...] + safety_identifier: str | None = None + + +@dataclass(frozen=True, slots=True) +class DecisionsIRPredicateAnswer: + probability: float + extra: Mapping[str, object] = field(default_factory=lambda: _NO_EXTRA) + + +@dataclass(frozen=True, slots=True) +class DecisionsIRChoiceProbability: + value: str | bool + probability: float + + +@dataclass(frozen=True, slots=True) +class DecisionsIRChoiceAnswer: + choice: str | bool + confidence: float + probabilities: tuple[DecisionsIRChoiceProbability, ...] + extra: Mapping[str, object] = field(default_factory=lambda: _NO_EXTRA) + + +@dataclass(frozen=True, slots=True) +class DecisionsIRScoreProbability: + value: int + label: DecisionsJSON + probability: float + + +@dataclass(frozen=True, slots=True) +class DecisionsIRScoreAnswer: + score: float + confidence: float + probabilities: tuple[DecisionsIRScoreProbability, ...] + extra: Mapping[str, object] = field(default_factory=lambda: _NO_EXTRA) + + +@dataclass(frozen=True, slots=True) +class DecisionsIRRefusal: + pass + + +DecisionsIRAnswer: TypeAlias = ( + DecisionsIRPredicateAnswer | DecisionsIRChoiceAnswer | DecisionsIRScoreAnswer | DecisionsIRRefusal +) + + +@dataclass(frozen=True, slots=True) +class DecisionsIRUsage: + input_tokens: int = 0 + output_tokens: int = 0 + cached_tokens: int = 0 + cache_write_tokens: int = 0 + reasoning_tokens: int = 0 + extra: Mapping[str, object] = field(default_factory=lambda: _NO_EXTRA) + + +@dataclass(frozen=True, slots=True) +class DecisionsIRResponse: + model: str | None + answers: tuple[DecisionsIRAnswer, ...] + usage: DecisionsIRUsage + extra: Mapping[str, object] = field(default_factory=lambda: _NO_EXTRA) + + +@dataclass(frozen=True, slots=True) +class UnsupportedDecisionsRequest: + reason: str diff --git a/litellm/types/llms/openai.py b/litellm/types/llms/openai.py index cf9fbfe2866..9dc53cf2adf 100644 --- a/litellm/types/llms/openai.py +++ b/litellm/types/llms/openai.py @@ -652,6 +652,7 @@ class ChatCompletionCachedContent(TypedDict): class PromptCacheBreakpoint(TypedDict): mode: ReadOnly[Literal["explicit"]] + ttl: NotRequired[ReadOnly[Literal["30m"]]] class PromptCacheOptions(TypedDict, total=False): diff --git a/litellm/types/openai_decisions.py b/litellm/types/openai_decisions.py new file mode 100644 index 00000000000..074c0a803f3 --- /dev/null +++ b/litellm/types/openai_decisions.py @@ -0,0 +1,162 @@ +from collections.abc import Mapping, Sequence +from dataclasses import dataclass +from typing import Annotated, Literal, TypeAlias + +from pydantic import ConfigDict, Field, PrivateAttr, StrictBool, StrictStr + +from litellm.types.llms.base import LiteLLMPydanticObjectBase + +ChoiceValue: TypeAlias = StrictStr | StrictBool + + +class DecisionsObjectBase(LiteLLMPydanticObjectBase): + model_config = ConfigDict(extra="allow", frozen=True) + + +class DecisionInputText(DecisionsObjectBase): + type: Literal["input_text"] + text: str + + +class DecisionInputImage(DecisionsObjectBase): + type: Literal["input_image"] + image_url: str + detail: Literal["low", "high", "auto", "original"] | None = None + + +DecisionInputPart: TypeAlias = Annotated[DecisionInputText | DecisionInputImage, Field(discriminator="type")] + + +class DecisionInputMessage(DecisionsObjectBase): + role: Literal["user"] + content: str | Sequence[DecisionInputPart] + type: Literal["message"] | None = None + + +DecisionInput: TypeAlias = str | Sequence[DecisionInputMessage] + + +class DecisionChoice(DecisionsObjectBase): + value: ChoiceValue + description: str | None = None + + +class DecisionLevel(DecisionsObjectBase): + label: str + description: str | None = None + + +class PredicateQuestion(DecisionsObjectBase): + type: Literal["predicate"] + instructions: str + name: str | None = None + + +class ChoiceQuestion(DecisionsObjectBase): + type: Literal["choice"] + instructions: str + choices: Sequence[DecisionChoice] + name: str | None = None + + +class ScoreQuestion(DecisionsObjectBase): + type: Literal["score"] + instructions: str + levels: Sequence[DecisionLevel] + name: str | None = None + + +DecisionQuestion: TypeAlias = Annotated[ + PredicateQuestion | ChoiceQuestion | ScoreQuestion, + Field(discriminator="type"), +] + +DecisionQuestions: TypeAlias = Sequence[DecisionQuestion] + + +class DecisionsRequestBody(DecisionsObjectBase): + input: DecisionInput + questions: DecisionQuestions + safety_identifier: str | None = None + + +@dataclass(frozen=True, slots=True) +class DecisionsRequest: + model: str + body: DecisionsRequestBody + + +class PredicateAnswer(DecisionsObjectBase): + type: Literal["predicate"] + name: str | None = None + probability: float + + +class ChoiceProbability(DecisionsObjectBase): + value: ChoiceValue + probability: float + + +class ChoiceAnswer(DecisionsObjectBase): + type: Literal["choice"] + name: str | None = None + choice: ChoiceValue + probabilities: Sequence[ChoiceProbability] + confidence: float + + +class ScoreProbability(DecisionsObjectBase): + value: int + label: str + probability: float + + +class ScoreAnswer(DecisionsObjectBase): + type: Literal["score"] + name: str | None = None + score: float + probabilities: Sequence[ScoreProbability] + confidence: float + + +class RefusalAnswer(DecisionsObjectBase): + type: Literal["refusal"] + name: str | None = None + + +DecisionAnswer: TypeAlias = Annotated[ + PredicateAnswer | ChoiceAnswer | ScoreAnswer | RefusalAnswer, + Field(discriminator="type"), +] + + +class DecisionInputTokensDetails(DecisionsObjectBase): + cached_tokens: int + cache_write_tokens: int + + +class DecisionOutputTokensDetails(DecisionsObjectBase): + reasoning_tokens: int + + +class DecisionUsage(DecisionsObjectBase): + input_tokens: int + input_tokens_details: DecisionInputTokensDetails + output_tokens: int + output_tokens_details: DecisionOutputTokensDetails + total_tokens: int + + +class DecisionsResponse(DecisionsObjectBase): + model: str + answers: Sequence[DecisionAnswer] + usage: DecisionUsage + + _hidden_params: dict[str, object] = PrivateAttr(default_factory=dict) + + @property + def hidden_params(self) -> dict[str, object]: # mutable-ok: API requires mutation + return self._hidden_params + + def set_hidden_params(self, params: Mapping[str, object]) -> None: + self._hidden_params.update(params) diff --git a/litellm/types/utils.py b/litellm/types/utils.py index 91525226ef4..bb7d2e3e4ba 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -788,6 +788,8 @@ API_ROUTE_TO_CALL_TYPES: Final[Mapping[str, Sequence[CallTypes]]] = { "/v1/search": [CallTypes.asearch, CallTypes.search], "/decisions": [CallTypes.adecisions, CallTypes.decisions], "/v1/decisions": [CallTypes.adecisions, CallTypes.decisions], + "/systemone": [CallTypes.adecisions, CallTypes.decisions], + "/v1/systemone": [CallTypes.adecisions, CallTypes.decisions], # Batches "/batches": [CallTypes.acreate_batch, CallTypes.create_batch], "/v1/batches": [CallTypes.acreate_batch, CallTypes.create_batch], diff --git a/litellm/utils.py b/litellm/utils.py index 63599815569..036b8c5b20a 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -379,6 +379,7 @@ if TYPE_CHECKING: # Type stubs for lazy-loaded config classes and types from litellm.llms.base_llm.batches.transformation import BaseBatchesConfig from litellm.llms.base_llm.containers.transformation import BaseContainerConfig + from litellm.llms.base_llm.decisions.transformation import BaseDecisionsConfig from litellm.llms.base_llm.embedding.transformation import BaseEmbeddingConfig from litellm.llms.base_llm.files.transformation import BaseFilesConfig from litellm.llms.base_llm.google_genai.transformation import ( @@ -1218,8 +1219,8 @@ def function_setup( else search_query ) elif call_type in (CallTypes.decisions.value, CallTypes.adecisions.value): - decisions_state: Final = args[1] if len(args) > 1 else kwargs.get("state", "") - messages = decisions_state if isinstance(decisions_state, str) else json.dumps(decisions_state) + decisions_state: Final = args[1] if len(args) > 1 else kwargs.get("state") or kwargs.get("input") or "" + messages = decisions_state if isinstance(decisions_state, str) else json.dumps(decisions_state, default=str) elif call_type in (CallTypes.image_edit.value, CallTypes.aimage_edit.value): messages = args[1] if len(args) > 1 else kwargs.get("prompt") elif call_type in (CallTypes.ocr.value, CallTypes.aocr.value): @@ -8961,6 +8962,22 @@ class ProviderConfigManager: return get_dashscope_family_rerank_config(provider.value) return litellm.CohereRerankConfig() + @staticmethod + def get_provider_decisions_config(model: str, provider: LlmProviders) -> BaseDecisionsConfig | None: + if provider == LlmProviders.PERPLEXITY: + return litellm.PerplexityDecisionsConfig() + if provider == LlmProviders.TYPESAFE: + return litellm.TypeSafeDecisionsConfig() + if provider == LlmProviders.OPENROUTER: + return litellm.OpenRouterDecisionsConfig() + if provider == LlmProviders.CLOUDFLARE: + return litellm.CloudflareDecisionsConfig() + if provider == LlmProviders.STRANDS_DECIDER: + return litellm.StrandsDeciderDecisionsConfig() + if provider == LlmProviders.OPENAI: + return litellm.OpenAIDecisionsConfig() + return None + @staticmethod def get_provider_anthropic_messages_config( model: str, diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index c2aa83a9216..f34606ce127 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -34527,7 +34527,8 @@ "supported_endpoints": [ "/v1/chat/completions", "/v1/batch", - "/v1/responses" + "/v1/responses", + "/v1/decisions" ], "supported_modalities": [ "text", @@ -59538,7 +59539,7 @@ "bedrock_converse_supports_strict_tools": false, "cache_creation_input_token_cost": 2.75e-06, "cache_creation_input_token_cost_above_1hr": 4.4e-06, - "cache_read_input_token_cost": 2.2e-07, + "cache_read_input_token_cost": 1.1e-07, "input_cost_per_token": 2.2e-06, "litellm_provider": "bedrock_mantle", "supports_tool_search": true, @@ -59573,7 +59574,7 @@ "supports_forced_tool_use": false, "thinking_always_on": true, "prompt_cache_min_tokens": 512, - "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-anthropic-claude-sonnet-5-5.html" + "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" }, "us.xai.grok-4.6": { "supports_regex_lookaround": false, @@ -66362,7 +66363,7 @@ "bedrock_output_config_effort_ceiling": "xhigh", "cache_creation_input_token_cost": 3e-06, "cache_creation_input_token_cost_above_1hr": 4.8e-06, - "cache_read_input_token_cost": 2.4e-07, + "cache_read_input_token_cost": 1.2e-07, "input_cost_per_token": 2.4e-06, "litellm_provider": "bedrock_mantle", "supports_tool_search": true, @@ -66391,7 +66392,7 @@ "supports_tool_choice": true, "supports_vision": true, "supports_xhigh_reasoning_effort": true, - "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-anthropic-claude-sonnet-5-5.html" + "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" }, "bedrock_mantle/us-gov-east-1/openai.gpt-5.4": { "litellm_provider": "bedrock_mantle", @@ -79407,7 +79408,7 @@ "bedrock_converse_supports_strict_tools": false, "cache_creation_input_token_cost": 2.5e-06, "cache_creation_input_token_cost_above_1hr": 4e-06, - "cache_read_input_token_cost": 2e-07, + "cache_read_input_token_cost": 1e-07, "input_cost_per_token": 2e-06, "litellm_provider": "bedrock_converse", "supports_tool_search": true, @@ -79477,14 +79478,14 @@ "cache_creation_input_token_cost": 2.75e-06, "input_cost_per_token": 2.2e-06, "output_cost_per_token": 1.1e-05, - "cache_read_input_token_cost": 2.2e-07, + "cache_read_input_token_cost": 1.1e-07, "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" }, "au.anthropic.claude-sonnet-5-5": { "bedrock_converse_supports_strict_tools": false, "cache_creation_input_token_cost": 2.75e-06, "cache_creation_input_token_cost_above_1hr": 4.4e-06, - "cache_read_input_token_cost": 2.2e-07, + "cache_read_input_token_cost": 1.1e-07, "input_cost_per_token": 2.2e-06, "litellm_provider": "bedrock_converse", "supports_tool_search": true, @@ -79519,7 +79520,7 @@ "supports_forced_tool_use": false, "thinking_always_on": true, "prompt_cache_min_tokens": 512, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" }, "azure_ai/claude-sonnet-5-5": { "supports_mid_conversation_system": true, @@ -79562,7 +79563,7 @@ "bedrock_output_config_effort_ceiling": "xhigh", "cache_creation_input_token_cost": 3e-06, "cache_creation_input_token_cost_above_1hr": 4.8e-06, - "cache_read_input_token_cost": 2.4e-07, + "cache_read_input_token_cost": 1.2e-07, "input_cost_per_token": 2.4e-06, "litellm_provider": "bedrock", "supports_tool_search": true, @@ -79574,6 +79575,7 @@ "supports_forced_tool_use": false, "thinking_always_on": true, "prompt_cache_min_tokens": 512, + "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json", "supports_adaptive_thinking": true, "supports_assistant_prefill": false, "supports_computer_use": true, @@ -79597,7 +79599,7 @@ "bedrock_output_config_effort_ceiling": "xhigh", "cache_creation_input_token_cost": 3e-06, "cache_creation_input_token_cost_above_1hr": 4.8e-06, - "cache_read_input_token_cost": 2.4e-07, + "cache_read_input_token_cost": 1.2e-07, "input_cost_per_token": 2.4e-06, "litellm_provider": "bedrock", "supports_tool_search": true, @@ -79609,6 +79611,7 @@ "supports_forced_tool_use": false, "thinking_always_on": true, "prompt_cache_min_tokens": 512, + "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json", "supports_adaptive_thinking": true, "supports_assistant_prefill": false, "supports_computer_use": true, @@ -79631,7 +79634,7 @@ "bedrock_converse_supports_strict_tools": false, "cache_creation_input_token_cost": 2.75e-06, "cache_creation_input_token_cost_above_1hr": 4.4e-06, - "cache_read_input_token_cost": 2.2e-07, + "cache_read_input_token_cost": 1.1e-07, "input_cost_per_token": 2.2e-06, "litellm_provider": "bedrock_converse", "supports_tool_search": true, @@ -79666,13 +79669,13 @@ "supports_forced_tool_use": false, "thinking_always_on": true, "prompt_cache_min_tokens": 512, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" }, "global.anthropic.claude-sonnet-5-5": { "bedrock_converse_supports_strict_tools": false, "cache_creation_input_token_cost": 2.5e-06, "cache_creation_input_token_cost_above_1hr": 4e-06, - "cache_read_input_token_cost": 2e-07, + "cache_read_input_token_cost": 1e-07, "input_cost_per_token": 2e-06, "litellm_provider": "bedrock_converse", "supports_tool_search": true, @@ -79707,13 +79710,13 @@ "supports_forced_tool_use": false, "thinking_always_on": true, "prompt_cache_min_tokens": 512, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" }, "jp.anthropic.claude-sonnet-5-5": { "bedrock_converse_supports_strict_tools": false, "cache_creation_input_token_cost": 2.75e-06, "cache_creation_input_token_cost_above_1hr": 4.4e-06, - "cache_read_input_token_cost": 2.2e-07, + "cache_read_input_token_cost": 1.1e-07, "input_cost_per_token": 2.2e-06, "litellm_provider": "bedrock_converse", "supports_tool_search": true, @@ -79748,7 +79751,7 @@ "supports_forced_tool_use": false, "thinking_always_on": true, "prompt_cache_min_tokens": 512, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" }, "openrouter/anthropic/claude-sonnet-5.5": { "input_cost_per_token": 2e-06, @@ -79858,7 +79861,7 @@ "bedrock_output_config_effort_ceiling": "xhigh", "cache_creation_input_token_cost": 3e-06, "cache_creation_input_token_cost_above_1hr": 4.8e-06, - "cache_read_input_token_cost": 2.4e-07, + "cache_read_input_token_cost": 1.2e-07, "input_cost_per_token": 2.4e-06, "litellm_provider": "bedrock_converse", "supports_tool_search": true, @@ -79870,7 +79873,7 @@ "supports_forced_tool_use": false, "thinking_always_on": true, "prompt_cache_min_tokens": 512, - "source": "https://aws.amazon.com/bedrock/pricing/", + "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json", "supports_adaptive_thinking": true, "supports_assistant_prefill": false, "supports_computer_use": true, @@ -79893,7 +79896,7 @@ "bedrock_converse_supports_strict_tools": false, "cache_creation_input_token_cost": 2.75e-06, "cache_creation_input_token_cost_above_1hr": 4.4e-06, - "cache_read_input_token_cost": 2.2e-07, + "cache_read_input_token_cost": 1.1e-07, "input_cost_per_token": 2.2e-06, "litellm_provider": "bedrock_converse", "supports_tool_search": true, @@ -79928,7 +79931,7 @@ "supports_forced_tool_use": false, "thinking_always_on": true, "prompt_cache_min_tokens": 512, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" }, "vertex_ai/claude-sonnet-5-5": { "regional_endpoint_uplift_multiplier": 1.1, @@ -79936,8 +79939,8 @@ "cache_creation_input_token_cost": 2.5e-06, "cache_creation_input_token_cost_above_1hr": 4e-06, "cache_creation_input_token_cost_batches": 1.25e-06, - "cache_read_input_token_cost": 2e-07, - "cache_read_input_token_cost_batches": 1e-07, + "cache_read_input_token_cost": 1e-07, + "cache_read_input_token_cost_batches": 5e-08, "input_cost_per_token": 2e-06, "input_cost_per_token_batches": 1e-06, "litellm_provider": "vertex_ai-anthropic_models", @@ -79976,8 +79979,8 @@ "cache_creation_input_token_cost": 2.5e-06, "cache_creation_input_token_cost_above_1hr": 4e-06, "cache_creation_input_token_cost_batches": 1.25e-06, - "cache_read_input_token_cost": 2e-07, - "cache_read_input_token_cost_batches": 1e-07, + "cache_read_input_token_cost": 1e-07, + "cache_read_input_token_cost_batches": 5e-08, "input_cost_per_token": 2e-06, "input_cost_per_token_batches": 1e-06, "litellm_provider": "vertex_ai-anthropic_models", diff --git a/scripts/trace_codegen/schemas/traces-clickhouse/FeedbackRow.json b/scripts/trace_codegen/schemas/traces-clickhouse/FeedbackRow.json new file mode 100644 index 00000000000..a107aea47e3 --- /dev/null +++ b/scripts/trace_codegen/schemas/traces-clickhouse/FeedbackRow.json @@ -0,0 +1,53 @@ +{ + "$schema": "https://json-schema.org/draft/2020-12/schema", + "properties": { + "author": { + "type": "string" + }, + "comment": { + "type": "string" + }, + "created_at": { + "type": "string" + }, + "score": { + "anyOf": [ + { + "format": "uint64", + "maximum": 18446744073709551615, + "minimum": 0, + "type": "integer" + }, + { + "pattern": "^(?:0|[1-9][0-9]{0,18}|1[0-7][0-9]{18}|18[0-3][0-9]{17}|184[0-3][0-9]{16}|1844[0-5][0-9]{15}|18446[0-6][0-9]{14}|184467[0-3][0-9]{13}|1844674[0-3][0-9]{12}|184467440[0-6][0-9]{10}|1844674407[0-2][0-9]{9}|18446744073[0-6][0-9]{8}|1844674407370[0-8][0-9]{6}|18446744073709[0-4][0-9]{5}|184467440737095[0-4][0-9]{4}|1844674407370955[0-0][0-9]{3}|18446744073709551[0-5][0-9]{2}|184467440737095516[0-0][0-9]{1}|1844674407370955161[0-4][0-9]{0}|18446744073709551615)$", + "type": "string" + } + ], + "x-python-normalized": { + "maximum": 18446744073709551615, + "minimum": 0, + "type": "int" + } + }, + "trace_id": { + "type": "string" + }, + "trace_ref": { + "type": "string" + }, + "updated_at": { + "type": "string" + } + }, + "required": [ + "trace_id", + "trace_ref", + "author", + "score", + "comment", + "created_at", + "updated_at" + ], + "title": "FeedbackRow", + "type": "object" +} diff --git a/scripts/trace_codegen/schemas/traces-clickhouse/FeedbackSummaryRow.json b/scripts/trace_codegen/schemas/traces-clickhouse/FeedbackSummaryRow.json new file mode 100644 index 00000000000..53f0e7e306a --- /dev/null +++ b/scripts/trace_codegen/schemas/traces-clickhouse/FeedbackSummaryRow.json @@ -0,0 +1,62 @@ +{ + "$schema": "https://json-schema.org/draft/2020-12/schema", + "properties": { + "average": { + "format": "double", + "type": "number" + }, + "count": { + "anyOf": [ + { + "format": "uint64", + "maximum": 18446744073709551615, + "minimum": 0, + "type": "integer" + }, + { + "pattern": "^(?:0|[1-9][0-9]{0,18}|1[0-7][0-9]{18}|18[0-3][0-9]{17}|184[0-3][0-9]{16}|1844[0-5][0-9]{15}|18446[0-6][0-9]{14}|184467[0-3][0-9]{13}|1844674[0-3][0-9]{12}|184467440[0-6][0-9]{10}|1844674407[0-2][0-9]{9}|18446744073[0-6][0-9]{8}|1844674407370[0-8][0-9]{6}|18446744073709[0-4][0-9]{5}|184467440737095[0-4][0-9]{4}|1844674407370955[0-0][0-9]{3}|18446744073709551[0-5][0-9]{2}|184467440737095516[0-0][0-9]{1}|1844674407370955161[0-4][0-9]{0}|18446744073709551615)$", + "type": "string" + } + ], + "x-python-normalized": { + "maximum": 18446744073709551615, + "minimum": 0, + "type": "int" + } + }, + "lowest": { + "anyOf": [ + { + "format": "uint64", + "maximum": 18446744073709551615, + "minimum": 0, + "type": "integer" + }, + { + "pattern": "^(?:0|[1-9][0-9]{0,18}|1[0-7][0-9]{18}|18[0-3][0-9]{17}|184[0-3][0-9]{16}|1844[0-5][0-9]{15}|18446[0-6][0-9]{14}|184467[0-3][0-9]{13}|1844674[0-3][0-9]{12}|184467440[0-6][0-9]{10}|1844674407[0-2][0-9]{9}|18446744073[0-6][0-9]{8}|1844674407370[0-8][0-9]{6}|18446744073709[0-4][0-9]{5}|184467440737095[0-4][0-9]{4}|1844674407370955[0-0][0-9]{3}|18446744073709551[0-5][0-9]{2}|184467440737095516[0-0][0-9]{1}|1844674407370955161[0-4][0-9]{0}|18446744073709551615)$", + "type": "string" + } + ], + "x-python-normalized": { + "maximum": 18446744073709551615, + "minimum": 0, + "type": "int" + } + }, + "trace_id": { + "type": "string" + }, + "trace_ref": { + "type": "string" + } + }, + "required": [ + "trace_id", + "trace_ref", + "count", + "average", + "lowest" + ], + "title": "FeedbackSummaryRow", + "type": "object" +} diff --git a/scripts/trace_codegen/schemas/traces-clickhouse/FeedbackTargetRow.json b/scripts/trace_codegen/schemas/traces-clickhouse/FeedbackTargetRow.json new file mode 100644 index 00000000000..17b9adb2d68 --- /dev/null +++ b/scripts/trace_codegen/schemas/traces-clickhouse/FeedbackTargetRow.json @@ -0,0 +1,21 @@ +{ + "$schema": "https://json-schema.org/draft/2020-12/schema", + "properties": { + "key_hash": { + "type": "string" + }, + "team_id": { + "type": "string" + }, + "trace_ref": { + "type": "string" + } + }, + "required": [ + "team_id", + "key_hash", + "trace_ref" + ], + "title": "FeedbackTargetRow", + "type": "object" +} diff --git a/scripts/trace_codegen/schemas/traces-clickhouse/LensFeedbackParams.json b/scripts/trace_codegen/schemas/traces-clickhouse/LensFeedbackParams.json new file mode 100644 index 00000000000..63794197547 --- /dev/null +++ b/scripts/trace_codegen/schemas/traces-clickhouse/LensFeedbackParams.json @@ -0,0 +1,34 @@ +{ + "$schema": "https://json-schema.org/draft/2020-12/schema", + "additionalProperties": false, + "properties": { + "all_teams": { + "enum": [ + 0, + 1 + ], + "type": "integer" + }, + "key_hash": { + "type": "string" + }, + "team": { + "type": "string" + }, + "trace_id": { + "type": "string" + }, + "trace_ref": { + "type": "string" + } + }, + "required": [ + "all_teams", + "team", + "key_hash", + "trace_id", + "trace_ref" + ], + "title": "LensFeedbackParams", + "type": "object" +} diff --git a/scripts/trace_codegen/schemas/traces-clickhouse/LensFeedbackSummaryParams.json b/scripts/trace_codegen/schemas/traces-clickhouse/LensFeedbackSummaryParams.json new file mode 100644 index 00000000000..dde3eb83612 --- /dev/null +++ b/scripts/trace_codegen/schemas/traces-clickhouse/LensFeedbackSummaryParams.json @@ -0,0 +1,33 @@ +{ + "$schema": "https://json-schema.org/draft/2020-12/schema", + "additionalProperties": false, + "properties": { + "all_teams": { + "enum": [ + 0, + 1 + ], + "type": "integer" + }, + "key_hash": { + "type": "string" + }, + "team": { + "type": "string" + }, + "trace_ids": { + "items": { + "type": "string" + }, + "type": "array" + } + }, + "required": [ + "all_teams", + "team", + "key_hash", + "trace_ids" + ], + "title": "LensFeedbackSummaryParams", + "type": "object" +} diff --git a/scripts/trace_codegen/schemas/traces-clickhouse/LensFeedbackTargetParams.json b/scripts/trace_codegen/schemas/traces-clickhouse/LensFeedbackTargetParams.json new file mode 100644 index 00000000000..4e5e475ce27 --- /dev/null +++ b/scripts/trace_codegen/schemas/traces-clickhouse/LensFeedbackTargetParams.json @@ -0,0 +1,34 @@ +{ + "$schema": "https://json-schema.org/draft/2020-12/schema", + "additionalProperties": false, + "properties": { + "all_teams": { + "enum": [ + 0, + 1 + ], + "type": "integer" + }, + "key_hash": { + "type": "string" + }, + "team": { + "type": "string" + }, + "trace_id": { + "type": "string" + }, + "trace_ref": { + "type": "string" + } + }, + "required": [ + "all_teams", + "team", + "key_hash", + "trace_id", + "trace_ref" + ], + "title": "LensFeedbackTargetParams", + "type": "object" +} diff --git a/scripts/trace_codegen/schemas/traces-clickhouse/ReadQueryName.json b/scripts/trace_codegen/schemas/traces-clickhouse/ReadQueryName.json index 1a3fb1151f0..33d66e64a35 100644 --- a/scripts/trace_codegen/schemas/traces-clickhouse/ReadQueryName.json +++ b/scripts/trace_codegen/schemas/traces-clickhouse/ReadQueryName.json @@ -6,7 +6,10 @@ "agents", "sample", "content", - "evidence" + "evidence", + "feedback_target", + "feedback", + "feedback_summary" ], "title": "ReadQueryName", "type": "string" diff --git a/scripts/trace_codegen/schemas/traces/Trace.json b/scripts/trace_codegen/schemas/traces/Trace.json index 3fcbf654101..b8b893607dc 100644 --- a/scripts/trace_codegen/schemas/traces/Trace.json +++ b/scripts/trace_codegen/schemas/traces/Trace.json @@ -71,6 +71,11 @@ }, "url": { "type": "string" + }, + "user": { + "description": "Who started the conversation, e.g. the Slack user's email.", + "type": "string", + "x-python-optional": true } }, "required": [ diff --git a/scripts/trace_codegen/schemas/traces/TracePage.json b/scripts/trace_codegen/schemas/traces/TracePage.json index 9d437bd60a1..47f48a676de 100644 --- a/scripts/trace_codegen/schemas/traces/TracePage.json +++ b/scripts/trace_codegen/schemas/traces/TracePage.json @@ -11,6 +11,11 @@ }, "url": { "type": "string" + }, + "user": { + "description": "Who started the conversation, e.g. the Slack user's email.", + "type": "string", + "x-python-optional": true } }, "required": [ diff --git a/tests/batches_tests/test_batch_rate_limits.py b/tests/batches_tests/test_batch_rate_limits.py index 06beed6b195..87d35e11c74 100644 --- a/tests/batches_tests/test_batch_rate_limits.py +++ b/tests/batches_tests/test_batch_rate_limits.py @@ -15,19 +15,19 @@ from litellm import DualCache from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.hooks.batch_rate_limiter import ( BatchFileUsage, - _PROXY_BatchRateLimiter, + PROXY_BatchRateLimiter, ) from litellm.proxy.hooks.parallel_request_limiter_v3 import ( - _PROXY_MaxParallelRequestsHandler_v3, + PROXY_MaxParallelRequestsHandler_v3, ) from litellm.proxy.utils import InternalUsageCache -def _build_batch_limiter() -> _PROXY_BatchRateLimiter: +def _build_batch_limiter() -> PROXY_BatchRateLimiter: internal_usage_cache = InternalUsageCache(dual_cache=DualCache()) - return _PROXY_BatchRateLimiter( + return PROXY_BatchRateLimiter( internal_usage_cache=internal_usage_cache, - parallel_request_limiter=_PROXY_MaxParallelRequestsHandler_v3( + parallel_request_limiter=PROXY_MaxParallelRequestsHandler_v3( internal_usage_cache=internal_usage_cache ), ) @@ -136,7 +136,7 @@ async def test_batch_rate_limit_single_file(tmp_path): # Setup: Create internal usage cache and rate limiter dual_cache = DualCache() internal_usage_cache = InternalUsageCache(dual_cache=dual_cache) - rate_limiter = _PROXY_MaxParallelRequestsHandler_v3( + rate_limiter = PROXY_MaxParallelRequestsHandler_v3( internal_usage_cache=internal_usage_cache ) @@ -198,7 +198,7 @@ async def test_batch_rate_limit_single_file(tmp_path): # Reset cache for clean test dual_cache = DualCache() internal_usage_cache = InternalUsageCache(dual_cache=dual_cache) - rate_limiter = _PROXY_MaxParallelRequestsHandler_v3( + rate_limiter = PROXY_MaxParallelRequestsHandler_v3( internal_usage_cache=internal_usage_cache ) batch_limiter = rate_limiter._get_batch_rate_limiter() @@ -278,7 +278,7 @@ async def test_batch_rate_limit_multiple_requests(tmp_path): # Setup: Create internal usage cache and rate limiter dual_cache = DualCache() internal_usage_cache = InternalUsageCache(dual_cache=dual_cache) - rate_limiter = _PROXY_MaxParallelRequestsHandler_v3( + rate_limiter = PROXY_MaxParallelRequestsHandler_v3( internal_usage_cache=internal_usage_cache ) @@ -431,7 +431,7 @@ async def test_batch_rate_limiter_with_managed_files(tmp_path): # Setup: Create internal usage cache and rate limiter dual_cache = DualCache() internal_usage_cache = InternalUsageCache(dual_cache=dual_cache) - rate_limiter = _PROXY_MaxParallelRequestsHandler_v3( + rate_limiter = PROXY_MaxParallelRequestsHandler_v3( internal_usage_cache=internal_usage_cache ) @@ -662,7 +662,7 @@ async def test_batch_rate_limiter_managed_files_regression(): # Setup: Create batch rate limiter dual_cache = DualCache() internal_usage_cache = InternalUsageCache(dual_cache=dual_cache) - rate_limiter = _PROXY_MaxParallelRequestsHandler_v3( + rate_limiter = PROXY_MaxParallelRequestsHandler_v3( internal_usage_cache=internal_usage_cache ) batch_limiter = rate_limiter._get_batch_rate_limiter() @@ -685,10 +685,10 @@ async def test_batch_rate_limiter_managed_files_regression(): # Test 1: Verify managed file detection print("\n1. Verifying managed file detection...") from litellm.proxy.openai_files_endpoints.common_utils import ( - _is_base64_encoded_unified_file_id, + is_base64_encoded_unified_file_id, ) - is_managed = _is_base64_encoded_unified_file_id(managed_file_id) + is_managed = is_base64_encoded_unified_file_id(managed_file_id) assert is_managed, "Managed file should be detected correctly" print(" ✓ Managed file detected") diff --git a/tests/code_coverage_tests/router_code_coverage.py b/tests/code_coverage_tests/router_code_coverage.py index 07521fef965..cbc09edc357 100644 --- a/tests/code_coverage_tests/router_code_coverage.py +++ b/tests/code_coverage_tests/router_code_coverage.py @@ -89,6 +89,7 @@ ignored_function_names = [ "_resolve_claude_code_session_router", # Tested through Claude Code session routing in test_router.py "_get_claude_code_session_router_binding", # Tested through the two-worker session routing test in test_router.py "_apply_updated_routing_strategy_args", # Tested via update_settings in test_lowest_latency.py (file lacks "router" in name) + "_wait_for_scheduler_turn", # Tested through prioritized acompletion and atext_completion in test_router.py "arm_routing_read_prefetch", # Tested in tests/unit/caching/test_request_redis_batch_pre_call.py (file lacks "router" in name) "_configured_model_info", # Tested through get_configured_service_tiers in test_router.py "_routable_deployments", # Tested through get_configured_service_tiers and get_routable_upstream_model in test_router.py diff --git a/tests/code_coverage_tests/unbounded_in_baseline.txt b/tests/code_coverage_tests/unbounded_in_baseline.txt index 129cf1b7797..77c7c362f39 100644 --- a/tests/code_coverage_tests/unbounded_in_baseline.txt +++ b/tests/code_coverage_tests/unbounded_in_baseline.txt @@ -38,7 +38,7 @@ litellm/proxy/management_endpoints/auto_router_endpoints.py start_shadow_eval pr litellm/proxy/management_endpoints/auto_router_endpoints.py start_shadow_eval prisma user_id.in `list(data.user_ids)` 0 litellm/proxy/management_endpoints/budget_management_endpoints.py info_budget prisma budget_id.in `data.budgets` 0 litellm/proxy/management_endpoints/common_utils.py _team_admin_can_invite_user prisma team_id.in `admin_user_obj.teams` 0 -litellm/proxy/management_endpoints/common_utils.py _user_has_admin_privileges prisma team_id.in `user_obj.teams` 0 +litellm/proxy/management_endpoints/common_utils.py user_has_admin_privileges prisma team_id.in `user_obj.teams` 0 litellm/proxy/management_endpoints/customer_endpoints.py delete_end_user prisma user_id.in `data.user_ids` 0 litellm/proxy/management_endpoints/customer_endpoints.py delete_end_user prisma user_id.in `data.user_ids` 1 litellm/proxy/management_endpoints/internal_user_endpoints.py _check_user_info_v2_access prisma team_id.in `caller_user.teams` 0 @@ -61,7 +61,7 @@ litellm/proxy/management_endpoints/key_management_endpoints.py _build_key_filter litellm/proxy/management_endpoints/key_management_endpoints.py _build_key_filter_conditions prisma team_id.in `member_only_team_ids` 0 litellm/proxy/management_endpoints/key_management_endpoints.py _build_key_filter_conditions prisma team_id.in `member_team_ids` 0 litellm/proxy/management_endpoints/key_management_endpoints.py _fetch_user_team_objects prisma team_id.in `complete_user_info.teams` 0 -litellm/proxy/management_endpoints/key_management_endpoints.py _list_key_helper prisma user_id.in `all_ids` 0 +litellm/proxy/management_endpoints/key_management_endpoints.py list_key_helper prisma user_id.in `all_ids` 0 litellm/proxy/management_endpoints/key_management_endpoints.py bulk_update_team_keys prisma token.in `hashed_key_ids` 0 litellm/proxy/management_endpoints/key_management_endpoints.py delete_key_aliases prisma key_alias.in `key_aliases` 0 litellm/proxy/management_endpoints/key_management_endpoints.py delete_verification_tokens prisma token.in `hashed_tokens` 0 diff --git a/tests/guardrails_tests/test_presidio_pii.py b/tests/guardrails_tests/test_presidio_pii.py index d8a1885b833..595cc95fa9f 100644 --- a/tests/guardrails_tests/test_presidio_pii.py +++ b/tests/guardrails_tests/test_presidio_pii.py @@ -5,7 +5,7 @@ from unittest.mock import patch import litellm from litellm.proxy.guardrails.guardrail_hooks.presidio import ( - _OPTIONAL_PresidioPIIMasking, + OPTIONAL_PresidioPIIMasking, PresidioPerRequestConfig, ) from litellm.types.guardrails import PiiEntityType, PiiAction @@ -24,7 +24,7 @@ async def test_presidio_with_blocked_entities(): PiiEntityType.EMAIL_ADDRESS: PiiAction.MASK, # This entity should be masked } - presidio_guardrail = _OPTIONAL_PresidioPIIMasking( + presidio_guardrail = OPTIONAL_PresidioPIIMasking( pii_entities_config=pii_entities_config, presidio_analyzer_api_base=os.environ.get("PRESIDIO_ANALYZER_API_BASE"), presidio_anonymizer_api_base=os.environ.get("PRESIDIO_ANONYMIZER_API_BASE"), @@ -64,7 +64,7 @@ async def test_presidio_pre_call_hook_with_blocked_entities(): PiiEntityType.EMAIL_ADDRESS: PiiAction.MASK, # This entity should be masked } - presidio_guardrail = _OPTIONAL_PresidioPIIMasking( + presidio_guardrail = OPTIONAL_PresidioPIIMasking( pii_entities_config=pii_entities_config, presidio_analyzer_api_base=os.environ.get("PRESIDIO_ANALYZER_API_BASE"), presidio_anonymizer_api_base=os.environ.get("PRESIDIO_ANONYMIZER_API_BASE"), @@ -195,10 +195,10 @@ async def test_presidio_pii_masking_logging_output_only_logged_response_guardrai assert len(litellm.guardrail_name_config_map) == 1 - pii_masking_obj: Optional[_OPTIONAL_PresidioPIIMasking] = None + pii_masking_obj: Optional[OPTIONAL_PresidioPIIMasking] = None for callback in litellm.callbacks: print(f"CALLBACK: {callback}") - if isinstance(callback, _OPTIONAL_PresidioPIIMasking): + if isinstance(callback, OPTIONAL_PresidioPIIMasking): pii_masking_obj = callback assert pii_masking_obj is not None diff --git a/tests/integration/_support/prompt_cache_breakpoint.py b/tests/integration/_support/prompt_cache_breakpoint.py new file mode 100644 index 00000000000..ce3b97b86be --- /dev/null +++ b/tests/integration/_support/prompt_cache_breakpoint.py @@ -0,0 +1,210 @@ +from __future__ import annotations + +import os +import re +import uuid +from collections.abc import Iterator, Mapping, Sequence +from contextlib import contextmanager +from dataclasses import dataclass +from pathlib import Path +from typing import Final, Literal, TypeAlias, assert_never +from urllib.parse import urlsplit + +import psutil +import psycopg +from integration._support import responses_vendor as rv +from integration._support.client import eventually, object_value, string_value +from integration._support.database import ROWS +from integration._support.openai_wire import answering_model_discovery, responses_reply +from integration._support.wire import Reply, Request, Wire +from psycopg.rows import DictRow, dict_row +from pydantic import JsonValue + +MODEL: Final = "openai/responses/gpt-6.1-sol" +EXPLICIT: Final[Mapping[str, JsonValue]] = {"mode": "explicit"} +EXPLICIT_30M: Final[Mapping[str, JsonValue]] = {"mode": "explicit", "ttl": "30m"} +NO_CACHE: Final[Mapping[str, JsonValue]] = {"cache": {"no-cache": True}} +INJECTION: Final[Mapping[str, JsonValue]] = { + "cache_control_injection_points": [{"location": "message", "role": "system"}], + "prompt_cache_options": {"mode": "explicit"}, +} +_SCRIPTED_FAILURE: Final = re.compile(r"fail-(\d{3})") +_MINTED_RESPONSE: Final = re.compile(r"^resp_([0-9a-f]{32})-[0-9a-f]{32}$") +_STARTED_WORKER: Final = re.compile(r"Started server process \[(\d+)\]") + +Kind: TypeAlias = Literal["text", "image_url", "file", "input_audio"] +KINDS: Final[tuple[Kind, ...]] = ("text", "image_url", "file", "input_audio") +WIRE_TYPE: Final[Mapping[Kind, str]] = { + "text": "input_text", + "image_url": "input_image", + "file": "input_file", + "input_audio": "input_text", +} + + +def _scripted(request: Request) -> Reply: + body: Final = rv.JSON_OBJECT.validate_json(request.body) + text: Final = request.body.decode() + marker: Final = rv.newest_marker(text) + failure: Final = _SCRIPTED_FAILURE.search(text) + if failure is not None: + return rv.error(int(failure.group(1)), f"scripted {failure.group(1)} marker-{marker}", "scripted_failure") + return responses_reply( + f"resp_{marker or uuid.uuid4().hex}-{uuid.uuid4().hex}", + string_value(body["model"]), + rv.answer(marker), + stream=body.get("stream") is True, + ) + + +respond: Final = answering_model_discovery(_scripted) + + +def response_marker(identity: str) -> str | None: + minted: Final = tuple( + found for candidate in rv.response_identities(identity) if (found := _MINTED_RESPONSE.match(candidate)) + ) + return minted[0].group(1) if minted else None + + +def answers(identity: str, marker: str) -> bool: + return response_marker(identity) == marker + + +def prompt(marker: str) -> str: + return f"Say marker-{marker}" + + +def text(value: str) -> dict[str, JsonValue]: + return {"type": "text", "text": value} + + +def marked(block: Mapping[str, JsonValue], marker: JsonValue) -> dict[str, JsonValue]: + return {**block, "prompt_cache_breakpoint": marker} + + +def block(kind: Kind, value: str) -> dict[str, JsonValue]: + match kind: + case "text": + return text(value) + case "image_url": + return {"type": "image_url", "image_url": {"url": "https://example.com/breakpoint.png"}} + case "file": + return {"type": "file", "file": {"file_id": "file-breakpoint"}} + case "input_audio": + return {"type": "input_audio", "input_audio": {"data": "Zm9v", "format": "wav"}} + case _: + assert_never(kind) + + +def drained_posts(wire: Wire) -> tuple[Request, ...]: + return tuple(request for request in wire.drain() if request.method == "POST") + + +def with_marker(posts: Sequence[Request], marker: str) -> tuple[Request, ...]: + return tuple(request for request in posts if f"marker-{marker}" in request.body.decode()) + + +def posted(wire: Wire, marker: str) -> Request: + matching: Final = with_marker(drained_posts(wire), marker) + assert len(matching) == 1, [request.body for request in matching] + (request,) = matching + assert request.target == "/v1/responses", request.target + return request + + +def body_of(request: Request) -> dict[str, JsonValue]: + return rv.JSON_OBJECT.validate_json(request.body) + + +def input_items(request: Request) -> list[dict[str, JsonValue]]: + return rv.ITEMS.validate_python(body_of(request)["input"]) + + +def content_of(items: Sequence[Mapping[str, JsonValue]], role: str) -> list[dict[str, JsonValue]]: + messages: Final = tuple(item for item in items if item.get("type") == "message" and item.get("role") == role) + assert len(messages) == 1, items + return rv.ITEMS.validate_python(messages[0]["content"]) + + +def single_block(items: Sequence[Mapping[str, JsonValue]], role: str) -> dict[str, JsonValue]: + blocks: Final = content_of(items, role) + assert len(blocks) == 1, blocks + return blocks[0] + + +def instruction_block(items: Sequence[Mapping[str, JsonValue]]) -> dict[str, JsonValue]: + messages: Final = tuple( + item for item in items if item.get("type") == "message" and item.get("role") in ("system", "developer") + ) + assert len(messages) == 1, items + blocks: Final = rv.ITEMS.validate_python(messages[0]["content"]) + assert len(blocks) == 1, blocks + return blocks[0] + + +def function_output(items: Sequence[Mapping[str, JsonValue]], call_id: str) -> list[dict[str, JsonValue]]: + outputs: Final = tuple( + item for item in items if item.get("type") == "function_call_output" and item.get("call_id") == call_id + ) + assert len(outputs) == 1, items + return rv.ITEMS.validate_python(outputs[0]["output"]) + + +def assert_marker(block_on_wire: Mapping[str, JsonValue], expected: JsonValue) -> None: + if expected is None: + assert "prompt_cache_breakpoint" not in block_on_wire, block_on_wire + return + assert block_on_wire.get("prompt_cache_breakpoint") == expected, block_on_wire + + +@dataclass(frozen=True, slots=True) +class SpendLogs: + connection: psycopg.Connection[DictRow] + + def rows_for(self, model: str) -> list[dict[str, JsonValue]]: + cursor: Final = self.connection.execute( + 'SELECT litellm_call_id, request_id, status FROM "LiteLLM_SpendLogs" WHERE model_group = %s', (model,) + ) + return ROWS.validate_python(cursor.fetchall()) + + def landed( + self, model: str, call_id: str, marker: str | None, *, status: str = "success", seconds: float = 70 + ) -> dict[str, JsonValue]: + rows: Final = eventually( + lambda: self.rows_for(model), + lambda found: any(row["litellm_call_id"] == call_id for row in found), + seconds=seconds, + ) + matching: Final = tuple(row for row in rows if row["litellm_call_id"] == call_id) + assert len(matching) == 1, rows + (row,) = matching + assert row["status"] == status, row + assert marker is None or answers(string_value(row["request_id"]), marker), (row, marker) + return row + + +@contextmanager +def spend_logs() -> Iterator[SpendLogs]: + with psycopg.connect(os.environ["DATABASE_URL"], row_factory=dict_row, autocommit=True) as connection: + connection.execute("SET default_transaction_read_only = on") + yield SpendLogs(connection) + + +def model_id(entries: Sequence[JsonValue], model: str) -> str: + matching: Final = tuple(entry for entry in entries if object_value(entry).get("model_name") == model) + assert len(matching) == 1, entries + return string_value(object_value(object_value(matching[0])["model_info"])["id"]) + + +def started_worker_pids(log: Path) -> tuple[int, ...]: + return tuple(int(found.group(1)) for found in _STARTED_WORKER.finditer(log.read_text())) + + +def open_upstream_connections(pid: int, upstream: str) -> int: + port: Final = urlsplit(upstream).port + return sum( + 1 + for connection in psutil.Process(pid).net_connections(kind="tcp") + if connection.status == psutil.CONN_ESTABLISHED and connection.raddr and connection.raddr.port == port + ) diff --git a/tests/integration/_support/upstream.py b/tests/integration/_support/upstream.py index ae87768cd06..0c7dab90fcc 100644 --- a/tests/integration/_support/upstream.py +++ b/tests/integration/_support/upstream.py @@ -440,6 +440,37 @@ class Provider: while pending: await websocket.send_json(_rendered_realtime_event(pending.popleft(), scenario_id)) + async def gemini_live(self, websocket: WebSocket) -> None: + authorization: Final = websocket.headers.get("authorization", "") + scenario_id: Final = websocket.headers.get("x-goog-user-project", "") + self.observations.put( + Observation( + websocket.url.path, + authorization, + {"host": websocket.headers.get("host", "")}, + "WEBSOCKET", + scenario_id, + ) + ) + response: Final = self.scenario_store.get(scenario_id) + if not isinstance(response, RealtimeResponse): + await websocket.close(code=4404) + return + await websocket.accept() + pending: Final = deque(response.events) + async for message in websocket.iter_json(): + payload: Final = JSON_OBJECT.validate_python(message) + self.observations.put( + Observation(websocket.url.path, authorization, payload, "WEBSOCKET_FRAME", scenario_id) + ) + if "setup" in payload: + await websocket.send_json({"setupComplete": {}}) + continue + if not _GEMINI_LIVE_TRIGGERS.intersection(payload): + continue + while pending: + await websocket.send_json(_rendered_realtime_event(pending.popleft(), scenario_id)) + @staticmethod def _response(response: StoredResponse, scenario_id: str) -> Response: unique_id: Final = f"{scenario_id}-{uuid.uuid4().hex[:8]}" @@ -539,6 +570,7 @@ class Provider: WebSocketRoute("/openai/v1/realtime", self.realtime), WebSocketRoute("/openai/realtime", self.realtime), WebSocketRoute("/v1/asr/realtime", self.muse_realtime), + WebSocketRoute(GEMINI_LIVE_PATH, self.gemini_live), ] ) @@ -562,6 +594,8 @@ def _interaction_body(interaction_id: str, state: InteractionState) -> dict[str, _REALTIME_TRIGGERS: Final = frozenset({"response.create", "input_audio_buffer.commit"}) +_GEMINI_LIVE_TRIGGERS: Final = frozenset({"clientContent", "realtimeInput"}) +GEMINI_LIVE_PATH: Final = "/ws/google.cloud.aiplatform.v1.LlmBidiService/BidiGenerateContent" _TRANSCRIPTION_UPDATE_REFUSED: Final = "Passing a realtime session update to a transcription session is not allowed." _REALTIME_UPDATE_REFUSED: Final = "Passing a transcription session update to a realtime session is not allowed." _NESTED_TURN_DETECTION_TYPE: Final = "session.audio.input.turn_detection.type" diff --git a/tests/integration/compatibility/test_missing_body_param_status.py b/tests/integration/compatibility/test_missing_body_param_status.py index 3e76ebaeda5..facf8007910 100644 --- a/tests/integration/compatibility/test_missing_body_param_status.py +++ b/tests/integration/compatibility/test_missing_body_param_status.py @@ -1025,6 +1025,99 @@ def test_anthropic_messages_uses_deployment_max_tokens_default(gateway: Gateway) assert outbound.get("max_tokens") == 32, provider_requests +def test_anthropic_messages_uses_router_wide_max_tokens_default(gateway: Gateway, tmp_path: Path) -> None: + with gateway.scenario() as scenario: + identity: Final = f"router-default-{uuid.uuid4().hex}" + handle: Final = register_scenario(identity, _response("anthropic_messages")) + scenario.cleanups.callback(delete_scenario, handle) + config: Final = tmp_path / "router-default.yaml" + config.write_text( + json.dumps( + { + "model_list": [ + { + "model_name": "router-default-anthropic", + "litellm_params": { + "model": "anthropic/claude-haiku-4-5", + "api_base": handle.api_base(), + "api_key": identity, + }, + } + ], + "router_settings": {"default_litellm_params": {"max_tokens": 32}}, + } + ), + encoding="utf-8", + ) + with owned_proxy_process(gateway, tmp_path, {}, config=config) as owned: + candidate: Final = Gateway(owned.gateway.client, owned.gateway.key, gateway.upstream_url) + response: Final = _post( + candidate, + "/v1/messages", + {"model": "router-default-anthropic", "messages": [{"role": "user", "content": "router default"}]}, + ) + assert response.status_code == 200, response.text + _assert_scripted_response("anthropic_messages", JSON_OBJECT.validate_python(response.json())) + observations: Final = _Observations(gateway.upstream_url) + captured: Final = eventually( + observations.read, + lambda _items: any(item.get("method") == "POST" for item in observations.for_scenario(identity)), + seconds=20, + ) + provider_requests: Final = tuple( + item for item in observations.for_scenario(identity) if item.get("method") == "POST" + ) + assert len(provider_requests) == 1, captured + outbound: Final = object_value(provider_requests[0]["body"]) + assert outbound.get("max_tokens") == 32, provider_requests + + +def test_rerank_uses_router_wide_documents_default(gateway: Gateway, tmp_path: Path) -> None: + with gateway.scenario() as scenario: + identity: Final = f"router-default-rerank-{uuid.uuid4().hex}" + handle: Final = register_scenario(identity, _response("arerank")) + scenario.cleanups.callback(delete_scenario, handle) + config: Final = tmp_path / "router-default-rerank.yaml" + config.write_text( + json.dumps( + { + "model_list": [ + { + "model_name": "router-default-rerank", + "litellm_params": { + "model": "cohere/rerank-v4.0", + "api_base": handle.api_base(), + "api_key": identity, + }, + } + ], + "router_settings": {"default_litellm_params": {"documents": ["router default document"]}}, + } + ), + encoding="utf-8", + ) + with owned_proxy_process(gateway, tmp_path, {}, config=config) as owned: + candidate: Final = Gateway(owned.gateway.client, owned.gateway.key, gateway.upstream_url) + response: Final = _post( + candidate, + "/rerank", + {"model": "router-default-rerank", "query": "which document?"}, + ) + assert response.status_code == 200, response.text + observations: Final = _Observations(gateway.upstream_url) + captured: Final = eventually( + observations.read, + lambda _items: any(item.get("method") == "POST" for item in observations.for_scenario(identity)), + seconds=20, + ) + provider_requests: Final = tuple( + item for item in observations.for_scenario(identity) if item.get("method") == "POST" + ) + assert len(provider_requests) == 1, captured + outbound: Final = object_value(provider_requests[0]["body"]) + assert outbound.get("documents") == ["router default document"], provider_requests + + def test_anthropic_messages_explicit_null_reaches_upstream(gateway: Gateway) -> None: with gateway.scenario() as scenario: model, identity, _handle = _register( diff --git a/tests/integration/cost_calculation/cost_tracking_case.py b/tests/integration/cost_calculation/cost_tracking_case.py index 466fdc46555..8c9a53262b6 100644 --- a/tests/integration/cost_calculation/cost_tracking_case.py +++ b/tests/integration/cost_calculation/cost_tracking_case.py @@ -250,7 +250,7 @@ class CostTrackingTestCase(BaseModel): "/v1/audio/speech", "/v1/images/generations", "/v1/images/edits", - "/v1/decisions", + "/v1/systemone", ] | Annotated[str, Field(pattern=r"^/(gemini|anthropic|bedrock)/")] ) = "/v1/chat/completions" diff --git a/tests/integration/cost_calculation/cost_tracking_cases.json b/tests/integration/cost_calculation/cost_tracking_cases.json index 94571f2a63c..81beff68c7c 100644 --- a/tests/integration/cost_calculation/cost_tracking_cases.json +++ b/tests/integration/cost_calculation/cost_tracking_cases.json @@ -31207,7 +31207,7 @@ "name": "perplexity/pplx-decider-v1-27b-decisions", "covers": "quota_management.spend_tracking.decisions_costs", "model": "perplexity/pplx-decider-v1-27b", - "endpoint": "/v1/decisions", + "endpoint": "/v1/systemone", "request": { "model": "$MODEL", "state": { diff --git a/tests/integration/cost_calculation/test_cost_tracking.py b/tests/integration/cost_calculation/test_cost_tracking.py index 379b2c13f5b..97dd7e494f1 100644 --- a/tests/integration/cost_calculation/test_cost_tracking.py +++ b/tests/integration/cost_calculation/test_cost_tracking.py @@ -268,7 +268,7 @@ def test_case_bills_expected_cost(gateway: Gateway, case: CostTrackingTestCase) assert row.spend == 0, f"{case.name}: failure spend was {row.spend}" return assert response.is_success, f"{case.name}: proxy returned {response.status_code}: {response.text[:400]}" - if case.endpoint == "/v1/decisions": + if case.endpoint == "/v1/systemone": observed: Final = JSON_OBJECT.validate_json( httpx.get(f"{gateway.upstream_url}/__observations", timeout=5, trust_env=False).content ) diff --git a/tests/integration/management/test_credential_federation_values.py b/tests/integration/management/test_credential_federation_values.py index 3f812c87c59..e17fb1154d8 100644 --- a/tests/integration/management/test_credential_federation_values.py +++ b/tests/integration/management/test_credential_federation_values.py @@ -1,11 +1,13 @@ from __future__ import annotations +import asyncio import json import uuid from collections.abc import Callable, Iterator, Mapping, Sequence from concurrent.futures import ThreadPoolExecutor from dataclasses import dataclass from functools import partial +from itertools import product from pathlib import Path from types import MappingProxyType from typing import Final, TypeVar @@ -43,7 +45,7 @@ from tests.integration._support.client import ( string_value, ) from tests.integration._support.database import read_rows -from tests.integration._support.process import OwnedProxy, owned_proxy_process +from tests.integration._support.process import OwnedProxy, graceful_stop_seconds, owned_proxy_process from tests.integration._support.wire import Reply, Request, Wire, wire_server T = TypeVar("T") @@ -562,7 +564,13 @@ class _Secrets: ) -_REMOVED_ENVIRONMENT: Final = ("ANTHROPIC_IDENTITY_TOKEN_FILE", "ANTHROPIC_API_KEY", "ANTHROPIC_API_BASE") +_REMOVED_ENVIRONMENT: Final = ( + "ANTHROPIC_IDENTITY_TOKEN_FILE", + "ANTHROPIC_API_KEY", + "ANTHROPIC_API_BASE", + "ANTHROPIC_FEDERATION_RULE_ID", + "ANTHROPIC_ORGANIZATION_ID", +) def _secrets(directory: Path) -> _Secrets: @@ -817,3 +825,388 @@ def test_worker_kill_mid_credential_burst(gateway: Gateway, tmp_path: Path) -> N for name in landed: _converged(survivor, name, _shape("token_file", f"fdrl-{name}")) _stable(partial(_chat_outcome, survivor, model), lambda outcome: outcome == (200, True)) + + +_FAIL_CLOSED_HINT: Final = "Settings > Workload identity" +_REFUSAL_CLIENTS: Final = ( + "chat", + "chat_stream", + "chat_async", + "messages", + "messages_stream", + "responses", + "responses_stream", +) +_OWNED_PROXY_CELL_SECONDS: Final = 2 * graceful_stop_seconds() + 120 + + +@dataclass(frozen=True, slots=True) +class _Misconfiguration: + reference: str + missing: tuple[str, ...] + blank: bool + expected: str + + def values(self, rig: FederationRig, tag: str) -> dict[str, JsonValue]: + token: Final = ( + str(rig.secrets.token_file) + if self.reference == "anthropic_identity_token_file" + else f"oidc/env/{_IDENTITY_TOKEN_VARIABLE}" + ) + ids: Final = {"anthropic_federation_rule_id": f"fdrl-{tag}", "anthropic_organization_id": f"org-{tag}"} + kept: Final = {key: value for key, value in ids.items() if key not in self.missing} + blanked: Final = {key: "" for key in self.missing} if self.blank else {} + return {self.reference: token, **kept, **blanked} + + +_MISCONFIGURED: Final[Mapping[str, _Misconfiguration]] = MappingProxyType( + { + "token_file_without_org": _Misconfiguration( + "anthropic_identity_token_file", + ("anthropic_organization_id",), + False, + "anthropic_identity_token_file is set, but anthropic_organization_id is not set", + ), + "token_file_without_rule": _Misconfiguration( + "anthropic_identity_token_file", + ("anthropic_federation_rule_id",), + False, + "anthropic_identity_token_file is set, but anthropic_federation_rule_id is not set", + ), + "inline_token_without_org": _Misconfiguration( + "anthropic_identity_token", + ("anthropic_organization_id",), + False, + "anthropic_identity_token is set, but anthropic_organization_id is not set", + ), + "inline_token_without_rule": _Misconfiguration( + "anthropic_identity_token", + ("anthropic_federation_rule_id",), + False, + "anthropic_identity_token is set, but anthropic_federation_rule_id is not set", + ), + "token_file_without_ids": _Misconfiguration( + "anthropic_identity_token_file", + ("anthropic_federation_rule_id", "anthropic_organization_id"), + False, + "anthropic_identity_token_file is set, but anthropic_federation_rule_id and anthropic_organization_id" + " are not set", + ), + "token_file_blank_org": _Misconfiguration( + "anthropic_identity_token_file", + ("anthropic_organization_id",), + True, + "anthropic_identity_token_file is set, but anthropic_organization_id is not set", + ), + } +) + + +@dataclass(frozen=True, slots=True) +class _Misconfigured: + tag: str + credential: str + model: str + shape: _Misconfiguration + + +def _misconfigured(rig: FederationRig, scenario: Scenario, shape: _Misconfiguration) -> _Misconfigured: + tag: Final = uuid.uuid4().hex + credential: Final = _create(rig.owned.gateway, scenario, shape.values(rig, tag)) + model: Final = _federated_deployment(rig.owned.gateway, scenario, credential, rig.peer.wire.url) + return _Misconfigured(tag=tag, credential=credential, model=model, shape=shape) + + +async def _async_chat(base_url: str, key: str, model: str, marker: str) -> None: + await openai.AsyncOpenAI(base_url=base_url + "/v1", api_key=key, max_retries=0).chat.completions.create( + model=model, messages=[{"role": "user", "content": prompt(marker)}] + ) + + +def _attempt(rig: FederationRig, client: str, model: str, marker: str) -> None: + base_url: Final = str(rig.owned.gateway.client.base_url) + key: Final = rig.owned.gateway.key + chat: Final = openai.OpenAI(base_url=base_url + "/v1", api_key=key, max_retries=0) + messages: Final = anthropic.Anthropic(base_url=base_url, api_key=key, max_retries=0) + match client: + case "chat": + chat.chat.completions.create(model=model, messages=[{"role": "user", "content": prompt(marker)}]) + case "chat_stream": + for _ in chat.chat.completions.create( + model=model, messages=[{"role": "user", "content": prompt(marker)}], stream=True + ): + pass + case "chat_async": + asyncio.run(_async_chat(base_url, key, model, marker)) + case "messages": + messages.messages.create(model=model, max_tokens=64, messages=[{"role": "user", "content": prompt(marker)}]) + case "messages_stream": + with messages.messages.stream( + model=model, max_tokens=64, messages=[{"role": "user", "content": prompt(marker)}] + ) as stream: + stream.get_final_message() + case "responses": + chat.responses.create(model=model, input=prompt(marker)) + case "responses_stream": + for _ in chat.responses.create(model=model, input=prompt(marker), stream=True): + pass + case _: + pytest.fail(f"unknown client {client!r}") + + +def _refusal(rig: FederationRig, client: str, model: str, marker: str) -> tuple[int, str]: + try: + _attempt(rig, client, model, marker) + except (openai.APIStatusError, anthropic.APIStatusError) as error: + return error.status_code, str(error) + pytest.fail(f"{client} succeeded against the misconfigured deployment {model}") + + +def _assert_names_missing_id(message: str, shape: _Misconfiguration) -> None: + assert shape.expected in message, message + assert _FAIL_CLOSED_HINT in message, message + + +def _assert_fails_closed(outcome: tuple[int, str], shape: _Misconfiguration) -> None: + status, message = outcome + assert status == 401, message + _assert_names_missing_id(message, shape) + + +def _assert_peer_untouched(rig: FederationRig, *fragments: str) -> None: + bodies: Final = tuple(request.body.decode() for request in rig.peer.requests()) + touched: Final = tuple(body for body in bodies if any(fragment in body for fragment in fragments)) + assert touched == (), touched + + +@pytest.mark.timeout(240) +@pytest.mark.parametrize("client", _REFUSAL_CLIENTS) +@pytest.mark.parametrize("shape", tuple(_MISCONFIGURED)) +def test_legacy_reference_without_ids_fails_closed_before_any_exchange( + federation: FederationRig, shape: str, client: str +) -> None: + marker: Final = uuid.uuid4().hex + with federation.owned.gateway.scenario() as scenario: + misconfigured: Final = _misconfigured(federation, scenario, _MISCONFIGURED[shape]) + _deployments_visible(federation.owned.gateway, (misconfigured.model,)) + _assert_fails_closed(_refusal(federation, client, misconfigured.model, marker), misconfigured.shape) + _assert_peer_untouched(federation, misconfigured.tag, marker) + + +@pytest.mark.timeout(240) +def test_file_upload_fails_closed_naming_the_missing_id(federation: FederationRig) -> None: + note: Final = uuid.uuid4().hex + with federation.owned.gateway.scenario() as scenario: + misconfigured: Final = _misconfigured(federation, scenario, _MISCONFIGURED["token_file_without_org"]) + _deployments_visible(federation.owned.gateway, (misconfigured.model,)) + response: Final = federation.owned.gateway.request_multipart( + "/v1/files", + {"purpose": "user_data", "model": misconfigured.model}, + {"file": ("notes.jsonl", json.dumps({"note": note}).encode(), "application/jsonl")}, + ) + _assert_fails_closed((response.status_code, response.text), misconfigured.shape) + _assert_peer_untouched(federation, misconfigured.tag, note) + + +@pytest.mark.timeout(240) +def test_skills_listing_fails_closed_naming_the_missing_id(federation: FederationRig) -> None: + with federation.owned.gateway.scenario() as scenario: + misconfigured: Final = _misconfigured(federation, scenario, _MISCONFIGURED["inline_token_without_rule"]) + _deployments_visible(federation.owned.gateway, (misconfigured.model,)) + response: Final = federation.owned.gateway.request( + "GET", "/v1/skills", params={"beta": "true", "model": misconfigured.model}, headers=_CLOSE + ) + _assert_fails_closed((response.status_code, response.text), misconfigured.shape) + upstream: Final = federation.peer.requests() + listed: Final = tuple(request for request in upstream if request.target.startswith("/v1/skills")) + assert listed == (), listed + _assert_peer_untouched(federation, misconfigured.tag) + + +@pytest.mark.timeout(240) +def test_health_report_names_the_missing_id(federation: FederationRig) -> None: + with federation.owned.gateway.scenario() as scenario: + misconfigured: Final = _misconfigured(federation, scenario, _MISCONFIGURED["token_file_without_rule"]) + _deployments_visible(federation.owned.gateway, (misconfigured.model,)) + response: Final = federation.owned.gateway.request( + "GET", "/health", params={"model": misconfigured.model}, headers=_CLOSE + ) + assert response.status_code == 503, response.text + unhealthy: Final = JSON_OBJECT.validate_json(response.content)["unhealthy_endpoints"] + assert isinstance(unhealthy, list) and len(unhealthy) == 1, response.text + _assert_names_missing_id(string_value(object_value(unhealthy[0])["error"]), misconfigured.shape) + _assert_peer_untouched(federation, misconfigured.tag) + + +@pytest.mark.timeout(240) +def test_connection_probe_names_the_missing_id(federation: FederationRig) -> None: + with federation.owned.gateway.scenario() as scenario: + misconfigured: Final = _misconfigured(federation, scenario, _MISCONFIGURED["inline_token_without_org"]) + _deployments_visible(federation.owned.gateway, (misconfigured.model,)) + response: Final = federation.owned.gateway.request( + "POST", + "/health/test_connection", + { + "litellm_params": { + "model": _MODEL, + "api_base": federation.peer.wire.url, + "litellm_credential_name": misconfigured.credential, + }, + "mode": "chat", + }, + headers=_CLOSE, + ) + assert response.status_code == 200, response.text + assert JSON_OBJECT.validate_json(response.content)["status"] == "error", response.text + _assert_names_missing_id(response.text, misconfigured.shape) + _assert_peer_untouched(federation, misconfigured.tag) + + +@pytest.mark.timeout(240) +def test_static_key_beside_a_stray_token_reference_still_wins(federation: FederationRig) -> None: + tag: Final = uuid.uuid4().hex + marker: Final = uuid.uuid4().hex + api_key: Final = f"sk-ant-api03-{tag}" + owned: Final = federation.owned + with owned.gateway.scenario() as scenario: + credential: Final = _create( + owned.gateway, + scenario, + { + "api_key": api_key, + "anthropic_identity_token_file": str(federation.secrets.token_file), + "anthropic_federation_rule_id": f"fdrl-{tag}", + }, + ) + model: Final = _federated_deployment(owned.gateway, scenario, credential, federation.peer.wire.url) + _deployments_visible(owned.gateway, (model,)) + _call(federation, "chat", model, marker) + upstream: Final = federation.peer.requests() + sent: Final = tuple( + request for request in upstream if request.target == "/v1/messages" and marker in request.body.decode() + ) + assert len(sent) == 1, upstream + assert sent[0].headers.get("x-api-key") == api_key, sent[0].headers + assert "authorization" not in sent[0].headers, sent[0].headers + _assert_peer_untouched(federation, tag) + + +@pytest.mark.timeout(240) +def test_blank_token_file_reference_is_unset_and_the_ambient_token_federates(federation: FederationRig) -> None: + rule_id: Final = f"fdrl-blank-{uuid.uuid4().hex}" + marker: Final = uuid.uuid4().hex + owned: Final = federation.owned + with owned.gateway.scenario() as scenario: + credential: Final = _create(owned.gateway, scenario, {**_ids(rule_id), "anthropic_identity_token_file": ""}) + model: Final = _federated_deployment(owned.gateway, scenario, credential, federation.peer.wire.url) + _deployments_visible(owned.gateway, (model,)) + _call(federation, "chat", model, marker) + grants: Final = tuple( + JSON_OBJECT.validate_json(request.body) + for request in federation.peer.requests() + if request.target == "/v1/oauth/token" + ) + mine: Final = tuple(grant for grant in grants if grant["federation_rule_id"] == rule_id) + assert mine, grants + assert all(grant["assertion"] == federation.secrets.environment_token for grant in mine), mine + + +def test_unauthenticated_call_is_refused_before_the_credential_is_read(federation: FederationRig) -> None: + owned: Final = federation.owned + marker: Final = uuid.uuid4().hex + with owned.gateway.scenario() as scenario: + entry: Final = _misconfigured(federation, scenario, _MISCONFIGURED["token_file_without_org"]) + _deployments_visible(owned.gateway, (entry.model,)) + refused: Final = owned.gateway.request( + "POST", + "/v1/chat/completions", + {"model": entry.model, "messages": [{"role": "user", "content": prompt(marker)}]}, + key=f"sk-not-a-key-{marker}", + headers=_CLOSE, + ) + assert refused.status_code == 401, refused.text + assert _FAIL_CLOSED_HINT not in refused.text, refused.text + _assert_fails_closed(_refusal(federation, "chat", entry.model, uuid.uuid4().hex), entry.shape) + _assert_peer_untouched(federation, entry.tag, marker) + + +@pytest.mark.timeout(_OWNED_PROXY_CELL_SECONDS) +def test_environment_organization_id_completes_a_legacy_reference(gateway: Gateway, tmp_path: Path) -> None: + directory: Final = tmp_path.resolve() + secrets: Final = _secrets(directory) + tag: Final = uuid.uuid4().hex + organization: Final = f"org-env-{tag}" + rule_id: Final = f"fdrl-{tag}" + marker: Final = uuid.uuid4().hex + with wire_server(_federation_peer) as wire: + with owned_proxy_process( + gateway, + directory, + {**secrets.overrides(), "ANTHROPIC_ORGANIZATION_ID": organization}, + remove_environment=_REMOVED_ENVIRONMENT, + workers=2, + ) as owned: + with owned.gateway.scenario() as scenario: + credential: Final = _create( + owned.gateway, + scenario, + {"anthropic_identity_token_file": str(secrets.token_file), "anthropic_federation_rule_id": rule_id}, + ) + model: Final = _federated_deployment(owned.gateway, scenario, credential, wire.url) + _deployments_visible(owned.gateway, (model,)) + response: Final = _chat(owned.gateway, model, marker) + assert response.status_code == 200, response.text + upstream: Final = wire.drain() + grants: Final = tuple( + JSON_OBJECT.validate_json(request.body) + for request in upstream + if request.target == "/v1/oauth/token" + ) + mine: Final = tuple(grant for grant in grants if grant["federation_rule_id"] == rule_id) + assert mine, upstream + assert all(grant["organization_id"] == organization for grant in mine), mine + assert all(grant["assertion"] == secrets.token_file.read_text() for grant in mine), mine + + +@pytest.mark.timeout(240) +def test_mixed_burst_fails_closed_without_disturbing_healthy_federation(federation: FederationRig) -> None: + owned: Final = federation.owned + shapes: Final = tuple(_MISCONFIGURED.values()) + healthy: Final = tuple(product(SOURCES, CLIENTS)) + with owned.gateway.scenario() as scenario: + misconfigured: Final = tuple( + _misconfigured(federation, scenario, shapes[index % len(shapes)]) for index in range(12) + ) + _deployments_visible(owned.gateway, tuple(entry.model for entry in misconfigured)) + control: Final = scenario.model() + serial_markers: Final = tuple(uuid.uuid4().hex for _ in _REFUSAL_CLIENTS) + for client, marker in zip(_REFUSAL_CLIENTS, serial_markers, strict=True): + _assert_fails_closed(_refusal(federation, client, misconfigured[0].model, marker), misconfigured[0].shape) + refused_markers: Final = tuple(uuid.uuid4().hex for _ in misconfigured) + healthy_markers: Final = tuple(uuid.uuid4().hex for _ in healthy) + + def refuse(index: int) -> tuple[int, str]: + client: Final = _REFUSAL_CLIENTS[index % len(_REFUSAL_CLIENTS)] + return _refusal(federation, client, misconfigured[index].model, refused_markers[index]) + + def federate(index: int) -> None: + source, client = healthy[index] + _call(federation, client, federation.deployments[source], healthy_markers[index]) + + with ThreadPoolExecutor(max_workers=16) as pool: + refusals: Final = tuple(pool.submit(refuse, index) for index in range(len(misconfigured))) + federations: Final = tuple(pool.submit(federate, index) for index in range(len(healthy))) + controls: Final = tuple(pool.submit(_chat_outcome, owned.gateway, control) for _ in range(4)) + outcomes: Final = tuple(future.result() for future in refusals) + for future in federations: + future.result() + assert tuple(future.result()[0] for future in controls) == (200,) * 4 + for entry, outcome in zip(misconfigured, outcomes, strict=True): + _assert_fails_closed(outcome, entry.shape) + sent: Final = tuple( + request.body.decode() for request in federation.peer.requests() if request.target == "/v1/messages" + ) + for marker in healthy_markers: + assert sum(marker in body for body in sent) == 1, marker + _assert_peer_untouched(federation, *serial_markers, *refused_markers, *(entry.tag for entry in misconfigured)) + assert owned.gateway.request("GET", "/health/liveliness").status_code == 200 diff --git a/tests/integration/messages_endpoint/chat_bridge/test_chat_bridge_mid_stream_failure_bookkeeping_wire.py b/tests/integration/messages_endpoint/chat_bridge/test_chat_bridge_mid_stream_failure_bookkeeping_wire.py new file mode 100644 index 00000000000..92a0a53529b --- /dev/null +++ b/tests/integration/messages_endpoint/chat_bridge/test_chat_bridge_mid_stream_failure_bookkeeping_wire.py @@ -0,0 +1,362 @@ +import asyncio +import json +import uuid +from collections import deque +from collections.abc import Callable, Mapping +from dataclasses import dataclass, replace +from typing import Final, Literal + +import anthropic +import httpx +import pytest +from anthropic.types import Message, MessageParam +from integration._support.anthropic_sse import ( + ANTHROPIC_ERROR_TYPES, + SseEvent, + delta_text, + dropping_reply, + error_type, + event_types, + parse_sse, + stream_reply, + user_prompt, +) +from integration._support.client import Gateway, eventually, object_value, string_value +from integration._support.database import read_rows +from integration._support.openai_wire import answering_model_discovery, chat_stream, openai_error, posted_targets +from integration._support.wire import Reply, Request, wire_server +from pydantic import JsonValue, TypeAdapter + +_BACKEND: Final = "gpt-4o-mini" +_PROVIDER_MODEL: Final = f"hosted_vllm/{_BACKEND}" +_PROVIDER_KEY: Final = "integration-provider-key" +_UPSTREAM_TARGET: Final = "/v1/chat/completions" +_TEXT: Final = "Hello" +_CALL_ID_HEADER: Final = "x-litellm-call-id" +_ROWS: Final = ( + "SELECT metadata->>'litellm_call_id' AS call_id, request_id, status, cache_hit, metadata " + 'FROM "LiteLLM_SpendLogs" WHERE model_group=%s' +) +_ROW_SECONDS: Final = 60 +_SLOW_PAUSE: Final = 1.0 +_BURST_PER_OUTCOME: Final = 8 + +Outcome = Literal[ + "drop_after_content", "error_frame_after_content", "slow_drop_after_content", "succeeds", "rejected_before_any_body" +] +_OUTCOME: Final = TypeAdapter(Outcome) +_FAILING_STREAMS: Final[tuple[Outcome, ...]] = ("drop_after_content", "error_frame_after_content") +_BURST: Final[tuple[Outcome, ...]] = ( + "slow_drop_after_content", + "error_frame_after_content", + "succeeds", +) * _BURST_PER_OUTCOME + + +@dataclass(frozen=True, slots=True) +class _Call: + outcome: Outcome + marker: str + prompt_tag: str + call_id: str + + @property + def prompt(self) -> str: + return f"{self.outcome}:{self.marker}:{self.prompt_tag}" + + @property + def streams(self) -> bool: + return self.outcome != "rejected_before_any_body" + + +@dataclass(frozen=True, slots=True) +class _SdkFailure: + text: str + error: anthropic.APIStatusError + + +def _call(outcome: Outcome, marker: str) -> _Call: + tag: Final = uuid.uuid4().hex + return _Call(outcome, marker, tag, tag) + + +def _marker() -> str: + return "bridge-mid-stream-" + uuid.uuid4().hex + + +def _error_frame(status: int) -> bytes: + error: Final = {"message": f"scripted mid-stream {status}", "type": "server_error", "code": status} + return b"data: " + json.dumps({"error": error}).encode() + b"\n\n" + + +def _reply(outcome: Outcome, chunks: tuple[bytes, bytes, bytes]) -> Reply: + match outcome: + case "drop_after_content": + return dropping_reply(chunks, abort_after=2) + case "error_frame_after_content": + return stream_reply((chunks[0], chunks[1], _error_frame(500) + b"data: [DONE]\n\n")) + case "slow_drop_after_content": + return stream_reply(chunks, abort_after=2, pause=_SLOW_PAUSE) + case "succeeds": + return stream_reply(chunks) + case "rejected_before_any_body": + return openai_error(500) + + +def _upstream(marker: str) -> Callable[[Request], Reply]: + def respond(request: Request) -> Reply: + assert (request.method, request.target) == ("POST", _UPSTREAM_TARGET), request + assert request.headers["authorization"] == f"Bearer {_PROVIDER_KEY}", request.headers + body: Final = object_value(json.loads(request.body)) + assert body["model"] == _BACKEND, body + outcome, scripted_marker, prompt_tag = user_prompt(body).split(":") + assert scripted_marker == marker, body + scripted: Final = _OUTCOME.validate_python(outcome) + assert body.get("stream", False) is (scripted != "rejected_before_any_body"), body + return _reply(scripted, chat_stream(f"chunk-{prompt_tag}", _BACKEND, _TEXT)) + + return answering_model_discovery(respond) + + +def _body(model: str, call: _Call) -> dict[str, JsonValue]: + return { + "model": model, + "max_tokens": 16, + "stream": call.streams, + "messages": [{"role": "user", "content": call.prompt}], + } + + +def _headers(call: _Call) -> dict[str, str]: + return {_CALL_ID_HEADER: call.call_id} + + +def _stream(gateway: Gateway, model: str, call: _Call) -> tuple[SseEvent, ...]: + response: Final = gateway.request("POST", "/v1/messages", _body(model, call), headers=_headers(call)) + assert response.status_code == 200, response.text + assert response.headers[_CALL_ID_HEADER] == call.call_id, response.headers + return parse_sse(response.text) + + +async def _stream_concurrently(client: httpx.AsyncClient, key: str, model: str, call: _Call) -> tuple[SseEvent, ...]: + response: Final = await client.post( + "/v1/messages", json=_body(model, call), headers={"Authorization": f"Bearer {key}", **_headers(call)} + ) + assert response.status_code == 200, response.text + assert response.headers[_CALL_ID_HEADER] == call.call_id, response.headers + return parse_sse(response.text) + + +def _leave_after_the_first_content_delta(gateway: Gateway, model: str, call: _Call) -> str: + headers: Final = {"Authorization": f"Bearer {gateway.key}", **_headers(call)} + with gateway.client.stream("POST", "/v1/messages", json=_body(model, call), headers=headers) as response: + assert response.status_code == 200, response + assert response.headers[_CALL_ID_HEADER] == call.call_id, response.headers + return next( + line for line in response.iter_lines() if line.startswith("data:") and '"content_block_delta"' in line + ) + + +def _call_id(row: dict[str, JsonValue]) -> str: + return string_value(row["call_id"]) + + +def _rows(model: str, call_ids: frozenset[str]) -> dict[str, dict[str, JsonValue]]: + rows: Final = eventually( + lambda: read_rows(_ROWS, (model,)), + lambda rows: call_ids.issubset(_call_id(row) for row in rows), + seconds=_ROW_SECONDS, + ) + by_call_id: Final = {_call_id(row): row for row in rows} + assert len(by_call_id) == len(rows), rows + return by_call_id + + +def _assert_failed_after_content(events: tuple[SseEvent, ...]) -> None: + types: Final = event_types(events) + assert types[0] == "message_start", events + assert delta_text(events) == _TEXT, events + assert types[-1] == "error" and "message_stop" not in types, events + + +def _assert_completed(events: tuple[SseEvent, ...]) -> None: + types: Final = event_types(events) + assert types[0] == "message_start" and types[-1] == "message_stop", events + assert "error" not in types, events + assert delta_text(events) == _TEXT, events + + +def _assert_outcome(outcome: Outcome, events: tuple[SseEvent, ...]) -> None: + if outcome == "succeeds": + _assert_completed(events) + return + _assert_failed_after_content(events) + + +def _expected_status(outcome: Outcome) -> str: + return "success" if outcome == "succeeds" else "failure" + + +def _assert_failure_row(row: Mapping[str, JsonValue], anthropic_error_type: str | None) -> None: + assert row["status"] == "failure", row + error: Final = object_value(object_value(row["metadata"])["error_information"]) + assert error["error_class"] != "MidStreamFallbackError", error + assert ANTHROPIC_ERROR_TYPES[int(str(error["error_code"]))] == anthropic_error_type, (error, anthropic_error_type) + + +def _snapshot_text(snapshot: Message) -> str: + return "".join(block.text for block in snapshot.content if block.type == "text") + + +def _sdk_messages(call: _Call) -> list[MessageParam]: + return [{"role": "user", "content": call.prompt}] + + +def _sdk_error_type(error: anthropic.APIStatusError) -> str: + return string_value(object_value(object_value(error.body)["error"])["type"]) + + +def _sdk_sync_failure(client: anthropic.Anthropic, model: str, call: _Call) -> _SdkFailure: + with client.messages.stream( + model=model, max_tokens=16, messages=_sdk_messages(call), extra_headers=_headers(call) + ) as stream: + assert stream.response.headers[_CALL_ID_HEADER] == call.call_id, stream.response.headers + try: + deque(stream, maxlen=0) + except anthropic.APIStatusError as error: + return _SdkFailure(_snapshot_text(stream.current_message_snapshot), error) + pytest.fail(f"{call.call_id}: the bridged stream ended without the scripted mid-stream failure") + + +async def _sdk_async_failure(client: anthropic.AsyncAnthropic, model: str, call: _Call) -> _SdkFailure: + async with client.messages.stream( + model=model, max_tokens=16, messages=_sdk_messages(call), extra_headers=_headers(call) + ) as stream: + assert stream.response.headers[_CALL_ID_HEADER] == call.call_id, stream.response.headers + try: + async for _ in stream: + pass + except anthropic.APIStatusError as error: + return _SdkFailure(_snapshot_text(stream.current_message_snapshot), error) + pytest.fail(f"{call.call_id}: the bridged stream ended without the scripted mid-stream failure") + + +def _assert_sdk_failure_bookkept(failure: _SdkFailure, model: str, call: _Call) -> None: + assert failure.text == _TEXT, failure + rows: Final = _rows(model, frozenset({call.call_id})) + assert rows.keys() == {call.call_id}, rows + _assert_failure_row(rows[call.call_id], _sdk_error_type(failure.error)) + + +@pytest.mark.parametrize("outcome", _FAILING_STREAMS) +def test_bridged_stream_failing_after_content_writes_one_failure_row(gateway: Gateway, outcome: Outcome) -> None: + marker: Final = _marker() + with wire_server(_upstream(marker)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=_PROVIDER_MODEL, api_base=wire.url + "/v1") + call: Final = _call(outcome, marker) + events: Final = _stream(gateway, model, call) + _assert_failed_after_content(events) + assert posted_targets(wire) == (_UPSTREAM_TARGET,) + rows: Final = _rows(model, frozenset({call.call_id})) + assert rows.keys() == {call.call_id}, rows + _assert_failure_row(rows[call.call_id], error_type(events)) + + +def test_anthropic_sdk_sync_stream_failing_after_content_raises_and_writes_one_failure_row(gateway: Gateway) -> None: + marker: Final = _marker() + with wire_server(_upstream(marker)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=_PROVIDER_MODEL, api_base=wire.url + "/v1") + call: Final = _call("drop_after_content", marker) + client: Final = anthropic.Anthropic( + base_url=str(gateway.client.base_url), api_key=gateway.key, max_retries=0, timeout=60 + ) + failure: Final = _sdk_sync_failure(client, model, call) + assert posted_targets(wire) == (_UPSTREAM_TARGET,) + _assert_sdk_failure_bookkept(failure, model, call) + + +async def test_anthropic_sdk_async_stream_failing_after_content_raises_and_writes_one_failure_row( + gateway: Gateway, +) -> None: + marker: Final = _marker() + with wire_server(_upstream(marker)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=_PROVIDER_MODEL, api_base=wire.url + "/v1") + call: Final = _call("error_frame_after_content", marker) + client: Final = anthropic.AsyncAnthropic( + base_url=str(gateway.client.base_url), api_key=gateway.key, max_retries=0, timeout=60 + ) + failure: Final = await _sdk_async_failure(client, model, call) + assert posted_targets(wire) == (_UPSTREAM_TARGET,) + _assert_sdk_failure_bookkept(failure, model, call) + + +def test_bridged_request_rejected_before_any_body_keeps_its_error_and_failure_row(gateway: Gateway) -> None: + marker: Final = _marker() + with wire_server(_upstream(marker)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=_PROVIDER_MODEL, api_base=wire.url + "/v1") + call: Final = _call("rejected_before_any_body", marker) + response: Final = gateway.request("POST", "/v1/messages", _body(model, call), headers=_headers(call)) + assert response.status_code == 500, response.text + error: Final = object_value(object_value(json.loads(response.text))["error"]) + assert error["type"] == ANTHROPIC_ERROR_TYPES[500], response.text + assert posted_targets(wire) == (_UPSTREAM_TARGET,) + rows: Final = _rows(model, frozenset({call.call_id})) + assert rows.keys() == {call.call_id}, rows + _assert_failure_row(rows[call.call_id], string_value(error["type"])) + + +def test_bridged_stream_served_from_the_response_cache_keeps_its_success_bookkeeping(gateway: Gateway) -> None: + marker: Final = _marker() + with wire_server(_upstream(marker)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=_PROVIDER_MODEL, api_base=wire.url + "/v1") + first: Final = _call("succeeds", marker) + _assert_completed(_stream(gateway, model, first)) + miss: Final = _rows(model, frozenset({first.call_id}))[first.call_id] + assert (miss["status"], miss["cache_hit"] != "True") == ("success", True), miss + twin: Final = replace(first, call_id=uuid.uuid4().hex) + _assert_completed(_stream(gateway, model, twin)) + assert posted_targets(wire) == (_UPSTREAM_TARGET,) + rows: Final = _rows(model, frozenset({twin.call_id})) + hit: Final = rows[twin.call_id] + assert (hit["status"], hit["cache_hit"]) == ("success", "True"), hit + assert string_value(hit["request_id"]) != string_value(miss["request_id"]), rows + assert set(rows) == {first.call_id, twin.call_id}, rows + + +def test_client_leaving_a_bridged_stream_before_its_failure_leaves_the_proxy_serving(gateway: Gateway) -> None: + marker: Final = _marker() + with wire_server(_upstream(marker)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=_PROVIDER_MODEL, api_base=wire.url + "/v1") + abandoned: Final = _call("slow_drop_after_content", marker) + first_delta: Final = _leave_after_the_first_content_delta(gateway, model, abandoned) + assert _TEXT in first_delta, first_delta + follow_up: Final = _call("succeeds", marker) + _assert_completed(_stream(gateway, model, follow_up)) + assert posted_targets(wire) == (_UPSTREAM_TARGET,) * 2 + rows: Final = _rows(model, frozenset({follow_up.call_id})) + assert rows[follow_up.call_id]["status"] == "success", rows + assert set(rows) <= {abandoned.call_id, follow_up.call_id}, rows + assert gateway.request("GET", "/health/liveliness").status_code == 200 + + +async def test_concurrent_bridged_streams_failing_after_content_each_land_one_row(gateway: Gateway) -> None: + marker: Final = _marker() + with wire_server(_upstream(marker)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=_PROVIDER_MODEL, api_base=wire.url + "/v1") + calls: Final = tuple(_call(outcome, marker) for outcome in _BURST) + async with httpx.AsyncClient(base_url=str(gateway.client.base_url), timeout=60, trust_env=False) as client: + streams: Final = await asyncio.gather( + *(_stream_concurrently(client, gateway.key, model, call) for call in calls) + ) + for call, events in zip(calls, streams, strict=True): + _assert_outcome(call.outcome, events) + assert posted_targets(wire) == (_UPSTREAM_TARGET,) * len(calls) + rows: Final = _rows(model, frozenset(call.call_id for call in calls)) + assert {call_id: row["status"] for call_id, row in rows.items()} == { + call.call_id: _expected_status(call.outcome) for call in calls + }, rows + for call, events in zip(calls, streams, strict=True): + if call.outcome != "succeeds": + _assert_failure_row(rows[call.call_id], error_type(events)) + assert gateway.request("GET", "/health/liveliness").status_code == 200 + _assert_completed(_stream(gateway, model, _call("succeeds", marker))) diff --git a/tests/integration/observability/test_langfuse_delivery.py b/tests/integration/observability/test_langfuse_delivery.py index 13be5a95887..00eee082c33 100644 --- a/tests/integration/observability/test_langfuse_delivery.py +++ b/tests/integration/observability/test_langfuse_delivery.py @@ -1,15 +1,33 @@ +import asyncio import base64 import json +import signal +import threading import time import uuid -from collections.abc import Sequence +from collections.abc import Iterator, Mapping, Sequence +from concurrent.futures import ThreadPoolExecutor +from dataclasses import dataclass from pathlib import Path from typing import Final +import anthropic +import httpx +import openai +import psutil +import pytest import yaml -from integration._support.client import Gateway, eventually, object_value, string_value +from _s3_v2_support import _chat_stream_frames, _responses_stream_frames +from integration._support.client import ( + Gateway, + Scenario, + eventually, + gateway_from_environment, + object_value, + string_value, +) from integration._support.database import read_rows, scratch_database -from integration._support.process import owned_proxy +from integration._support.process import group_members, owned_proxy, owned_proxy_process from integration._support.wire import Reply, Request, Wire, wire_server from opentelemetry.proto.collector.trace.v1.trace_service_pb2 import ExportTraceServiceRequest from opentelemetry.proto.common.v1.common_pb2 import KeyValue @@ -49,7 +67,20 @@ def _completion(text: str) -> Reply: def _projects() -> Reply: - return Reply(body=json.dumps({"data": [{"id": "integration-project", "name": "integration"}]}).encode()) + return Reply( + body=json.dumps( + { + "data": [ + { + "id": "integration-project", + "name": "integration", + "organization": {"id": "integration-org", "name": "integration"}, + "metadata": {}, + } + ] + } + ).encode() + ) def _text_prompt(name: str) -> Reply: @@ -68,15 +99,20 @@ def _text_prompt(name: str) -> Reply: ) -def _langfuse_config(tmp_path: Path) -> Path: +def _langfuse_config(tmp_path: Path, general_settings: Mapping[str, JsonValue] | None = None) -> Path: config: Final = _PROXY_CONFIG.validate_python(yaml.safe_load(STOCK_CONFIG.read_text())) settings: Final = { **_SETTINGS.validate_python(config["litellm_settings"]), "success_callback": ["langfuse"], "failure_callback": ["langfuse"], } - path: Final = tmp_path / "langfuse.yaml" - path.write_text(yaml.safe_dump({**config, "litellm_settings": settings})) + general: Final = { + **_SETTINGS.validate_python(config["general_settings"]), + **(general_settings or {}), + } + name: Final = "langfuse.yaml" if general_settings is None else "langfuse-merged.yaml" + path: Final = tmp_path / name + path.write_text(yaml.safe_dump({**config, "litellm_settings": settings, "general_settings": general})) return path @@ -367,3 +403,1494 @@ def test_prompt_fetch_encodes_the_name_retries_a_5xx_once_and_keeps_langfuse_hea assert leak not in json.dumps(dict(failure.headers)) assert "set-cookie" not in failure.headers and "x-upstream-internal" not in failure.headers assert sum(1 for target in seen_prompt_gets if target.startswith(PROMPTS_PATH + missing_prompt)) == 1 + + +def _responses_result(identity: str) -> Reply: + return Reply( + body=json.dumps( + { + "id": identity, + "object": "response", + "created_at": 1, + "status": "completed", + "model": "gpt-4o-mini", + "output": [ + { + "id": "msg_" + identity, + "type": "message", + "role": "assistant", + "status": "completed", + "content": [{"type": "output_text", "text": "integration answer", "annotations": []}], + } + ], + "usage": {"input_tokens": 7, "output_tokens": 2, "total_tokens": 9}, + } + ).encode() + ) + + +def _trace_body(kind: str, model: str, marker: str, metadata: Mapping[str, JsonValue] | None) -> dict[str, JsonValue]: + metadata_field: Final[dict[str, JsonValue]] = {} if metadata is None else {"metadata": dict(metadata)} + if kind == "responses": + return {"model": model, "input": marker + "-question", **metadata_field} + if kind == "messages": + return { + "model": model, + "max_tokens": 16, + "messages": [{"role": "user", "content": marker + "-question"}], + **metadata_field, + } + return { + "model": model, + "messages": [{"role": "user", "content": marker + "-question"}], + "cache": {"no-cache": True}, + **metadata_field, + } + + +def _w3c_headers(header_trace: str, baggage_session: str | None) -> dict[str, str]: + baggage_field: Final = {} if baggage_session is None else {"baggage": f"session.id={baggage_session}"} + return {"traceparent": f"00-{header_trace}-00f067aa0ba902b7-01", **baggage_field} + + +def _await_span(received: list[Request], destination: Wire, call_id: str) -> Span: + def exported() -> tuple[Span, ...]: + received.extend(destination.drain()) + return tuple( + span + for span in _spans(received) + if _attribute(span.attributes, "langfuse.observation.metadata.litellm_call_id") == call_id + ) + + spans: Final = eventually(exported, lambda values: len(values) == 1, seconds=20) + return spans[0] + + +@pytest.mark.parametrize( + ("endpoint", "kind", "metadata_mode", "expected_trace", "expected_session", "expected_target"), + ( + pytest.param( + "/v1/chat/completions", "chat", "both", "caller", "caller", "/v1/chat/completions", id="chat_caller_ids" + ), + pytest.param( + "/v1/responses", "responses", "both", "caller", "caller", "/v1/responses", id="responses_caller_ids" + ), + pytest.param( + "/v1/messages", "messages", "both", "caller", "caller", "/v1/responses", id="messages_caller_ids" + ), + pytest.param( + "/v1/chat/completions", "chat", "none", "header", "baggage", "/v1/chat/completions", id="chat_header_ids" + ), + pytest.param( + "/v1/chat/completions", + "chat", + "trace", + "caller", + "baggage", + "/v1/chat/completions", + id="chat_caller_trace_header_session", + ), + pytest.param( + "/v1/chat/completions", + "chat", + "empty_session", + "header", + "baggage", + "/v1/chat/completions", + id="chat_empty_session_header_ids", + ), + ), +) +def test_langfuse_trace_and_session_prefer_caller_metadata_over_w3c_headers( + gateway: Gateway, + tmp_path: Path, + endpoint: str, + kind: str, + metadata_mode: str, + expected_trace: str, + expected_session: str, + expected_target: str, +) -> None: + marker: Final = "w3c" + uuid.uuid4().hex + provider_secret: Final = "synthetic-provider-secret-" + marker + header_trace: Final = uuid.uuid4().hex + baggage_session: Final = "baggage-" + marker + caller_trace: Final = uuid.uuid4().hex + caller_session: Final = f"my-session-id-{marker}" + metadata: Final = { + "both": {"trace_id": caller_trace, "session_id": caller_session}, + "trace": {"trace_id": caller_trace}, + "empty_session": {"session_id": ""}, + "none": None, + }[metadata_mode] + expected_trace_value: Final = {"caller": caller_trace, "header": header_trace}[expected_trace] + expected_session_value: Final = {"caller": caller_session, "baggage": baggage_session}[expected_session] + upstream_targets: Final[list[str]] = [] # mutable-ok: records which upstream endpoint each call hit + + def upstream(request: Request) -> Reply: + if not request.body: + return Reply(status=404) + assert request.headers["authorization"] == f"Bearer {provider_secret}" + upstream_targets.append(request.target) + if request.target == "/v1/responses": + return _responses_result("resp-" + marker) + assert request.target == "/v1/chat/completions", request.target + return _completion(marker + "-answer") + + def langfuse(request: Request) -> Reply: + if request.method == "GET" and request.target.startswith(PROJECTS_PATH): + return _projects() + return Reply(body=b"", content_type="application/x-protobuf") + + with ( + wire_server(upstream) as provider, + wire_server(langfuse) as destination, + owned_proxy( + gateway, tmp_path, _langfuse_environment(destination), config=_langfuse_config(tmp_path) + ) as candidate, + candidate.scenario() as scenario, + ): + model: Final = scenario.model(api_base=provider.url + "/v1", api_key=provider_secret) + response: Final = candidate.request( + "POST", + endpoint, + _trace_body(kind, model, marker, metadata), + headers=_w3c_headers(header_trace, baggage_session), + ) + assert response.status_code == 200, response.text + assert upstream_targets == [expected_target], upstream_targets + received: Final[list[Request]] = [] # mutable-ok: drain() consumes the queue, later polls keep earlier ones + span: Final = _await_span(received, destination, response.headers["x-litellm-call-id"]) + assert span.trace_id.hex() == expected_trace_value + assert _attribute(span.attributes, "session.id") == expected_session_value + + +def test_missing_session_id_reject_accepts_caller_metadata_and_baggage_fallback( + gateway: Gateway, tmp_path: Path +) -> None: + marker: Final = "reject" + uuid.uuid4().hex + provider_secret: Final = "synthetic-provider-secret-" + marker + upstream_targets: Final[list[str]] = [] # mutable-ok: records which upstream endpoint each call hit + + def upstream(request: Request) -> Reply: + if not request.body: + return Reply(status=404) + assert request.headers["authorization"] == f"Bearer {provider_secret}" + upstream_targets.append(request.target) + if request.target == "/v1/responses": + return _responses_result("resp-" + marker) + assert request.target == "/v1/chat/completions", request.target + return _completion(marker + "-answer") + + def langfuse(request: Request) -> Reply: + if request.method == "GET" and request.target.startswith(PROJECTS_PATH): + return _projects() + return Reply(body=b"", content_type="application/x-protobuf") + + with ( + wire_server(upstream) as provider, + wire_server(langfuse) as destination, + owned_proxy( + gateway, + tmp_path, + _langfuse_environment(destination), + config=_langfuse_config(tmp_path, {"missing_session_id": "reject"}), + ) as candidate, + candidate.scenario() as scenario, + ): + model: Final = scenario.model(api_base=provider.url + "/v1", api_key=provider_secret) + received: Final[list[Request]] = [] # mutable-ok: drain() consumes the queue, later polls keep earlier ones + + caller_session: Final = f"my-session-id-{marker}-r1" + header_trace: Final = uuid.uuid4().hex + first: Final = candidate.request( + "POST", + "/v1/responses", + {"model": model, "input": marker + "-r1", "metadata": {"session_id": caller_session}}, + headers=_w3c_headers(header_trace, None), + ) + assert first.status_code == 200, first.text + first_span: Final = _await_span(received, destination, first.headers["x-litellm-call-id"]) + assert _attribute(first_span.attributes, "session.id") == caller_session + + baggage_session: Final = "baggage-" + marker + "-r2" + second: Final = candidate.request( + "POST", + "/v1/chat/completions", + { + "model": model, + "messages": [{"role": "user", "content": marker + "-r2"}], + "metadata": {"session_id": ""}, + }, + headers=_w3c_headers(uuid.uuid4().hex, baggage_session), + ) + assert second.status_code == 200, second.text + second_span: Final = _await_span(received, destination, second.headers["x-litellm-call-id"]) + assert _attribute(second_span.attributes, "session.id") == baggage_session + + third: Final = candidate.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": marker + "-r3"}]}, + headers=_w3c_headers(uuid.uuid4().hex, None), + ) + assert third.status_code == 400, third.text + assert upstream_targets == ["/v1/responses", "/v1/chat/completions"], upstream_targets + + +@pytest.mark.parametrize( + ("endpoint", "kind", "expected_target"), + ( + pytest.param("/v1/responses", "responses", "/v1/responses", id="responses_caller_trace"), + pytest.param("/v1/messages", "messages", "/v1/responses", id="messages_caller_trace"), + ), +) +def test_missing_session_id_generate_derives_session_from_caller_trace( + gateway: Gateway, tmp_path: Path, endpoint: str, kind: str, expected_target: str +) -> None: + marker: Final = "gen" + uuid.uuid4().hex + provider_secret: Final = "synthetic-provider-secret-" + marker + caller_trace: Final = uuid.uuid4().hex + upstream_targets: Final[list[str]] = [] # mutable-ok: records which upstream endpoint each call hit + + def upstream(request: Request) -> Reply: + if not request.body: + return Reply(status=404) + assert request.headers["authorization"] == f"Bearer {provider_secret}" + upstream_targets.append(request.target) + if request.target == "/v1/responses": + return _responses_result("resp-" + marker) + assert request.target == "/v1/chat/completions", request.target + return _completion(marker + "-answer") + + def langfuse(request: Request) -> Reply: + if request.method == "GET" and request.target.startswith(PROJECTS_PATH): + return _projects() + return Reply(body=b"", content_type="application/x-protobuf") + + with ( + wire_server(upstream) as provider, + wire_server(langfuse) as destination, + owned_proxy( + gateway, + tmp_path, + _langfuse_environment(destination), + config=_langfuse_config(tmp_path, {"missing_session_id": "generate"}), + ) as candidate, + candidate.scenario() as scenario, + ): + model: Final = scenario.model(api_base=provider.url + "/v1", api_key=provider_secret) + response: Final = candidate.request( + "POST", + endpoint, + _trace_body(kind, model, marker, {"trace_id": caller_trace}), + headers=_w3c_headers(uuid.uuid4().hex, None), + ) + assert response.status_code == 200, response.text + assert upstream_targets == [expected_target], upstream_targets + received: Final[list[Request]] = [] # mutable-ok: drain() consumes the queue, later polls keep earlier ones + span: Final = _await_span(received, destination, response.headers["x-litellm-call-id"]) + assert _attribute(span.attributes, "session.id") == caller_trace + + +_AUDIT_ENDPOINTS: Final = ( + pytest.param("/v1/chat/completions", "chat", "/v1/chat/completions", id="chat"), + pytest.param("/v1/responses", "responses", "/v1/responses", id="responses"), + pytest.param("/v1/messages", "messages", "/v1/responses", id="messages"), +) + + +def _audit_upstream(provider_secret: str, marker: str, targets: list[str]): + def upstream(request: Request) -> Reply: + if not request.body: + return Reply(status=404) + assert request.headers["authorization"] == f"Bearer {provider_secret}" + targets.append(request.target) + body: Final = TypeAdapter(dict[str, JsonValue]).validate_json(request.body) + index: Final = len(targets) + if request.target == "/v1/responses": + if body.get("stream"): + return Reply( + content_type="text/event-stream", chunks=_responses_stream_frames(f"resp-{marker}-{index}") + ) + return _responses_result(f"resp-{marker}-{index}") + assert request.target == "/v1/chat/completions", request.target + if body.get("stream"): + return Reply( + content_type="text/event-stream", chunks=_chat_stream_frames(f"chatcmpl-{marker}-{index}") + ) + return _completion(f"{marker}-{index}-answer") + + return upstream + + +def _audit_sink(): + def langfuse(request: Request) -> Reply: + if request.method == "GET" and request.target.startswith(PROJECTS_PATH): + return _projects() + return Reply(body=b"", content_type="application/x-protobuf") + + return langfuse + + +@dataclass(frozen=True, slots=True) +class _AuditRig: + candidate: Gateway + scenario: Scenario + destination: Wire + + +@pytest.fixture(scope="module") +def audit_rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[_AuditRig]: + directory: Final = tmp_path_factory.mktemp("audit") + with ( + gateway_from_environment() as gateway, + wire_server(_audit_sink()) as destination, + owned_proxy( + gateway, directory, _langfuse_environment(destination), config=_langfuse_config(directory) + ) as candidate, + candidate.scenario() as scenario, + ): + yield _AuditRig(candidate, scenario, destination) + + +@pytest.fixture(scope="module") +def audit_generate_rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[_AuditRig]: + directory: Final = tmp_path_factory.mktemp("audit-generate") + with ( + gateway_from_environment() as gateway, + wire_server(_audit_sink()) as destination, + owned_proxy( + gateway, + directory, + _langfuse_environment(destination), + config=_langfuse_config(directory, {"missing_session_id": "generate"}), + ) as candidate, + candidate.scenario() as scenario, + ): + yield _AuditRig(candidate, scenario, destination) + + +@pytest.fixture(scope="module") +def audit_reject_rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[_AuditRig]: + directory: Final = tmp_path_factory.mktemp("audit-reject") + with ( + gateway_from_environment() as gateway, + wire_server(_audit_sink()) as destination, + owned_proxy( + gateway, + directory, + _langfuse_environment(destination), + config=_langfuse_config(directory, {"missing_session_id": "reject"}), + ) as candidate, + candidate.scenario() as scenario, + ): + yield _AuditRig(candidate, scenario, destination) + + +@pytest.fixture(scope="module") +def audit_omit_rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[_AuditRig]: + directory: Final = tmp_path_factory.mktemp("audit-omit") + with ( + gateway_from_environment() as gateway, + wire_server(_audit_sink()) as destination, + owned_proxy( + gateway, + directory, + _langfuse_environment(destination), + config=_langfuse_config(directory, {"missing_session_id": "omit"}), + ) as candidate, + candidate.scenario() as scenario, + ): + yield _AuditRig(candidate, scenario, destination) + + +@pytest.fixture(scope="module") +def audit_otel_rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[_AuditRig]: + directory: Final = tmp_path_factory.mktemp("audit-otel") + + def otlp(request: Request) -> Reply: + return Reply() + + with ( + gateway_from_environment() as gateway, + wire_server(otlp) as destination, + owned_proxy( + gateway, + directory, + {"LITELLM_OTEL_V2": "1", "OTEL_BSP_SCHEDULE_DELAY": "300"}, + config=_otel_config(directory, destination.url), + ) as candidate, + candidate.scenario() as scenario, + ): + yield _AuditRig(candidate, scenario, destination) + + +def _await_spend_row(call_id: str) -> dict[str, JsonValue]: + rows: Final = eventually( + lambda: read_rows( + 'SELECT session_id, status, metadata, cache_hit FROM "LiteLLM_SpendLogs" WHERE litellm_call_id=%s', + (call_id,), + ), + lambda values: len(values) == 1, + seconds=240, + ) + return rows[0] + + +def _assert_call( + response: httpx.Response, + received: list[Request], + destination: Wire, + targets: list[str], + expected_target: str, + expected_trace: str | None, + expected_session: str | None, +) -> None: + assert targets == [expected_target], targets + _assert_span_spend(response, received, destination, expected_trace, expected_session) + + +def _assert_span_spend( + response: httpx.Response, + received: list[Request], + destination: Wire, + expected_trace: str | None, + expected_session: str | None, +) -> None: + call_id: Final = response.headers["x-litellm-call-id"] + span: Final = _await_span(received, destination, call_id) + if expected_trace is not None: + assert span.trace_id.hex() == expected_trace, f"call {call_id}: trace id" + assert _attribute(span.attributes, "session.id") == expected_session, f"call {call_id}: session.id" + row: Final = _await_spend_row(call_id) + assert row["session_id"] == expected_session, f"call {call_id}: spend session {row}" + + +@pytest.mark.parametrize(("endpoint", "kind", "expected_target"), _AUDIT_ENDPOINTS) +@pytest.mark.parametrize("metadata_mode", ("both", "none", "trace", "session")) +def test_audit_caller_metadata_wins_over_w3c_per_field( + audit_rig: _AuditRig, endpoint: str, kind: str, expected_target: str, metadata_mode: str +) -> None: + marker: Final = "audit" + uuid.uuid4().hex + provider_secret: Final = "synthetic-provider-secret-" + marker + header_trace: Final = uuid.uuid4().hex + baggage_session: Final = "baggage-" + marker + caller_trace: Final = uuid.uuid4().hex + caller_session: Final = f"my-session-id-{marker}" + metadata: Final = { + "both": {"trace_id": caller_trace, "session_id": caller_session}, + "trace": {"trace_id": caller_trace}, + "session": {"session_id": caller_session}, + "none": None, + }[metadata_mode] + expected_trace: Final = caller_trace if metadata_mode in ("both", "trace") else header_trace + expected_session: Final = caller_session if metadata_mode in ("both", "session") else baggage_session + targets: Final[list[str]] = [] # mutable-ok: the upstream records every hit + + with wire_server(_audit_upstream(provider_secret, marker, targets)) as provider: + model: Final = audit_rig.scenario.model(api_base=provider.url + "/v1", api_key=provider_secret) + response: Final = audit_rig.candidate.request( + "POST", + endpoint, + _trace_body(kind, model, marker, metadata), + headers=_w3c_headers(header_trace, baggage_session), + ) + assert response.status_code == 200, response.text + received: Final[list[Request]] = [] # mutable-ok: drain() consumes the queue, later polls keep earlier ones + _assert_call( + response, received, audit_rig.destination, targets, expected_target, expected_trace, expected_session + ) + + +@pytest.mark.parametrize(("endpoint", "kind", "expected_target"), _AUDIT_ENDPOINTS) +def test_audit_caller_metadata_wins_on_streamed_calls( + audit_rig: _AuditRig, endpoint: str, kind: str, expected_target: str +) -> None: + marker: Final = "auditstream" + uuid.uuid4().hex + provider_secret: Final = "synthetic-provider-secret-" + marker + header_trace: Final = uuid.uuid4().hex + baggage_session: Final = "baggage-" + marker + caller_trace: Final = uuid.uuid4().hex + caller_session: Final = f"my-session-id-{marker}" + targets: Final[list[str]] = [] # mutable-ok: the upstream records every hit + + with wire_server(_audit_upstream(provider_secret, marker, targets)) as provider: + model: Final = audit_rig.scenario.model(api_base=provider.url + "/v1", api_key=provider_secret) + response: Final = audit_rig.candidate.request( + "POST", + endpoint, + { + **_trace_body(kind, model, marker, {"trace_id": caller_trace, "session_id": caller_session}), + "stream": True, + }, + headers=_w3c_headers(header_trace, baggage_session), + ) + assert response.status_code == 200, response.text + received: Final[list[Request]] = [] # mutable-ok: drain() consumes the queue, later polls keep earlier ones + _assert_call(response, received, audit_rig.destination, targets, expected_target, caller_trace, caller_session) + + +def test_audit_caller_metadata_wins_through_official_sdk_clients(audit_rig: _AuditRig) -> None: + marker: Final = "auditsdk" + uuid.uuid4().hex + provider_secret: Final = "synthetic-provider-secret-" + marker + targets: Final[list[str]] = [] # mutable-ok: the upstream records every hit + + with wire_server(_audit_upstream(provider_secret, marker, targets)) as provider: + candidate: Final = audit_rig.candidate + destination: Final = audit_rig.destination + model: Final = audit_rig.scenario.model(api_base=provider.url + "/v1", api_key=provider_secret) + base_url: Final = str(candidate.client.base_url).rstrip("/") + received: Final[list[Request]] = [] # mutable-ok: drain() consumes the queue, later polls keep earlier ones + + def drive(client_kind: str) -> tuple[str, str, httpx.Headers]: + caller_trace: Final = uuid.uuid4().hex + caller_session: Final = f"{client_kind}-session-{marker}" + body_metadata: Final = { + "trace_id": caller_trace, + "session_id": caller_session, + "generation_name": marker + "-" + client_kind, + } + headers: Final = _w3c_headers(uuid.uuid4().hex, "baggage-" + marker + "-" + client_kind) + if client_kind == "chat_openai_sync": + return caller_trace, caller_session, ( + openai.OpenAI(base_url=f"{base_url}/v1", api_key=candidate.key) + .chat.completions.with_raw_response.create( + model=model, + messages=[{"role": "user", "content": marker + "-chat"}], + extra_body={"metadata": body_metadata, "cache": {"no-cache": True}}, + extra_headers=headers, + ) + .headers + ) + if client_kind == "responses_openai_async": + + async def responses_call() -> httpx.Headers: + answer: Final = await openai.AsyncOpenAI( + base_url=f"{base_url}/v1", api_key=candidate.key + ).responses.with_raw_response.create( + model=model, + input=marker + "-responses", + extra_body={"metadata": body_metadata, "cache": {"no-cache": True}}, + extra_headers=headers, + ) + return answer.headers + + return caller_trace, caller_session, asyncio.run(responses_call()) + if client_kind == "messages_anthropic_sync": + return caller_trace, caller_session, ( + anthropic.Anthropic(base_url=base_url, api_key=candidate.key) + .messages.with_raw_response.create( + model=model, + max_tokens=16, + messages=[{"role": "user", "content": marker + "-messages"}], + extra_body={"metadata": body_metadata}, + extra_headers=headers, + ) + .headers + ) + + async def messages_call() -> httpx.Headers: + answer: Final = await anthropic.AsyncAnthropic( + base_url=base_url, api_key=candidate.key + ).messages.with_raw_response.create( + model=model, + max_tokens=16, + messages=[{"role": "user", "content": marker + "-messages-async"}], + extra_body={"metadata": body_metadata}, + extra_headers=headers, + ) + return answer.headers + + return caller_trace, caller_session, asyncio.run(messages_call()) + + expected: Final = { + call_headers["x-litellm-call-id"]: (caller_trace, caller_session, client_kind) + for client_kind, (caller_trace, caller_session, call_headers) in ( + (kind, drive(kind)) + for kind in ( + "chat_openai_sync", + "responses_openai_async", + "messages_anthropic_sync", + "messages_anthropic_async", + ) + ) + } + + def spans_named() -> tuple[Span, ...]: + received.extend(destination.drain()) + return tuple( + span + for span in _spans(received) + if _attribute(span.attributes, "langfuse.observation.metadata.litellm_call_id") in expected + ) + + spans: Final = eventually(spans_named, lambda values: len(values) == len(expected), seconds=60) + for span in spans: + call_id: Final = str( + _attribute(span.attributes, "langfuse.observation.metadata.litellm_call_id") + ) + caller_trace, caller_session, client_kind = expected[call_id] + assert span.trace_id.hex() == caller_trace, f"{client_kind}: trace id" + assert _attribute(span.attributes, "session.id") == caller_session, f"{client_kind}: session.id" + row: Final = _await_spend_row(call_id) + assert row["session_id"] == caller_session, f"{client_kind} {call_id}: spend session {row}" + + +def _otel_config(tmp_path: Path, sink_url: str) -> Path: + config: Final = _PROXY_CONFIG.validate_python(yaml.safe_load(STOCK_CONFIG.read_text())) + settings: Final = {**_SETTINGS.validate_python(config["litellm_settings"]), "callbacks": ["otel"]} + path: Final = tmp_path / "otel.yaml" + path.write_text( + yaml.safe_dump( + { + **config, + "litellm_settings": settings, + "callback_settings": { + "otel": {"exporter": "http/json", "endpoint": sink_url, "mapper_names": ["genai"]} + }, + } + ) + ) + return path + + +_OTEL_SPAN: Final = TypeAdapter(dict[str, JsonValue]) + + +def _otel_spans(batches: Sequence[Request]) -> tuple[dict[str, JsonValue], ...]: + spans: list[dict[str, JsonValue]] = [] # mutable-ok: flattens nested OTLP batches into a tuple + for batch in batches: + if not batch.target.endswith("/v1/traces"): + continue + payload: Final = TypeAdapter(JsonValue).validate_json(batch.body) + envelopes: Final = payload if isinstance(payload, list) else [payload] + for envelope in envelopes: + for resource in TypeAdapter(list[JsonValue]).validate_python( + object_value(envelope)["resourceSpans"] + ): + for scope in object_value(resource)["scopeSpans"]: + spans.extend(TypeAdapter(list[JsonValue]).validate_python(object_value(scope)["spans"])) + return tuple(_OTEL_SPAN.validate_python(span) for span in spans) + + +def _otel_attribute(span: Mapping[str, JsonValue], key: str) -> str | None: + for attribute in TypeAdapter(list[JsonValue]).validate_python(span.get("attributes") or []): + entry: Final = object_value(attribute) + if entry["key"] == key: + value: Final = object_value(entry["value"]) + raw: Final = value.get("stringValue") + return str(raw) if raw is not None else None + return None + + +@pytest.mark.parametrize(("endpoint", "kind", "expected_target"), _AUDIT_ENDPOINTS[:2]) +def test_audit_otel_span_carries_caller_ids( + audit_otel_rig: _AuditRig, endpoint: str, kind: str, expected_target: str +) -> None: + marker: Final = "auditotel" + uuid.uuid4().hex + provider_secret: Final = "synthetic-provider-secret-" + marker + header_trace: Final = uuid.uuid4().hex + baggage_session: Final = "baggage-" + marker + caller_trace: Final = uuid.uuid4().hex + caller_session: Final = f"my-session-id-{marker}" + targets: Final[list[str]] = [] # mutable-ok: the upstream records every hit + + with wire_server(_audit_upstream(provider_secret, marker, targets)) as provider: + destination: Final = audit_otel_rig.destination + model: Final = audit_otel_rig.scenario.model(api_base=provider.url + "/v1", api_key=provider_secret) + response: Final = audit_otel_rig.candidate.request( + "POST", + endpoint, + _trace_body(kind, model, marker, {"trace_id": caller_trace, "session_id": caller_session}), + headers=_w3c_headers(header_trace, baggage_session), + ) + assert response.status_code == 200, response.text + assert targets == [expected_target], targets + response_id: Final = string_value(object_value(response.json())["id"]) + received: Final[list[Request]] = [] # mutable-ok: drain() consumes the queue, later polls keep earlier ones + + def exported() -> tuple[dict[str, JsonValue], ...]: + received.extend(destination.drain()) + return tuple( + span + for span in _otel_spans(received) + if _otel_attribute(span, "gen_ai.response.id") == response_id + ) + + spans: Final = eventually(exported, lambda values: len(values) == 1, seconds=60) + span: Final = spans[0] + print( + f"H7 record: otel span trace={span['traceId']} header={header_trace} " + f"session.id={_otel_attribute(span, 'session.id')} " + f"gen_ai.conversation.id={_otel_attribute(span, 'gen_ai.conversation.id')}" + ) + assert span["traceId"] == header_trace, ( + f"otel span for {response_id}: trace id must be the ambient W3C header trace" + ) + assert _otel_attribute(span, "gen_ai.conversation.id") == caller_session, ( + f"otel span conversation id for {response_id}" + ) + + +@pytest.mark.parametrize(("endpoint", "kind", "expected_target"), _AUDIT_ENDPOINTS) +@pytest.mark.parametrize( + ("case", "metadata_mode", "headers_mode", "expected_session"), + ( + pytest.param("g1", "trace", "traceparent", "caller_trace", id="g1_caller_trace_and_header"), + pytest.param("g3", "none", "traceparent", "header_trace", id="g3_header_only"), + pytest.param("g5", "session", "traceparent", "caller_session", id="g5_caller_session"), + ), +) +def test_audit_generate_policy_derives_session_from_caller_trace( + audit_generate_rig: _AuditRig, + endpoint: str, + kind: str, + expected_target: str, + case: str, + metadata_mode: str, + headers_mode: str, + expected_session: str, +) -> None: + marker: Final = "auditgen" + uuid.uuid4().hex + provider_secret: Final = "synthetic-provider-secret-" + marker + header_trace: Final = uuid.uuid4().hex + caller_trace: Final = uuid.uuid4().hex + caller_session: Final = f"my-session-id-{marker}" + metadata: Final = {"trace": {"trace_id": caller_trace}, "session": {"session_id": caller_session}, "none": None}[ + metadata_mode + ] + expected: Final = {"caller_trace": caller_trace, "header_trace": header_trace, "caller_session": caller_session}[ + expected_session + ] + targets: Final[list[str]] = [] # mutable-ok: the upstream records every hit + + with wire_server(_audit_upstream(provider_secret, marker, targets)) as provider: + model: Final = audit_generate_rig.scenario.model( + api_base=provider.url + "/v1", api_key=provider_secret + ) + response: Final = audit_generate_rig.candidate.request( + "POST", + endpoint, + _trace_body(kind, model, marker, metadata), + headers={"traceparent": f"00-{header_trace}-00f067aa0ba902b7-01"}, + ) + assert response.status_code == 200, response.text + received: Final[list[Request]] = [] # mutable-ok: drain() consumes the queue, later polls keep earlier ones + call_id: Final = response.headers["x-litellm-call-id"] + assert targets == [expected_target], targets + span: Final = _await_span(received, audit_generate_rig.destination, call_id) + assert _attribute(span.attributes, "session.id") == expected, f"call {call_id}: session.id" + row: Final = _await_spend_row(call_id) + assert row["session_id"] == expected, f"call {call_id}: spend session {row}" + + + +@pytest.mark.parametrize(("endpoint", "kind", "expected_target"), _AUDIT_ENDPOINTS) +def test_audit_generate_policy_is_stable_across_repeated_caller_trace( + audit_generate_rig: _AuditRig, endpoint: str, kind: str, expected_target: str +) -> None: + marker: Final = "auditgen2" + uuid.uuid4().hex + provider_secret: Final = "synthetic-provider-secret-" + marker + caller_trace: Final = uuid.uuid4().hex + targets: Final[list[str]] = [] # mutable-ok: the upstream records every hit + + with wire_server(_audit_upstream(provider_secret, marker, targets)) as provider: + model: Final = audit_generate_rig.scenario.model( + api_base=provider.url + "/v1", api_key=provider_secret + ) + responses: Final = tuple( + audit_generate_rig.candidate.request( + "POST", + endpoint, + _trace_body(kind, model, f"{marker}-{attempt}", {"trace_id": caller_trace}), + ) + for attempt in ("first", "second") + ) + sessions: Final[list[str | None]] = [] # mutable-ok: collects the two observed sessions in order + for attempt, response in zip(("first", "second"), responses): + assert response.status_code == 200, f"{attempt}: {response.text}" + row: Final = _await_spend_row(response.headers["x-litellm-call-id"]) + sessions.append(string_value(row["session_id"])) + assert sessions == [caller_trace, caller_trace], ( + f"generate must derive both sessions from the caller trace {caller_trace}" + ) + + +@pytest.mark.parametrize(("endpoint", "kind", "expected_target"), _AUDIT_ENDPOINTS) +def test_audit_generate_policy_fresh_session_without_any_ids( + audit_generate_rig: _AuditRig, endpoint: str, kind: str, expected_target: str +) -> None: + marker: Final = "auditgen4" + uuid.uuid4().hex + provider_secret: Final = "synthetic-provider-secret-" + marker + targets: Final[list[str]] = [] # mutable-ok: the upstream records every hit + + with wire_server(_audit_upstream(provider_secret, marker, targets)) as provider: + model: Final = audit_generate_rig.scenario.model( + api_base=provider.url + "/v1", api_key=provider_secret + ) + sessions: Final[list[str | None]] = [] # mutable-ok: collects the two observed sessions in order + for attempt in ("first", "second"): + response: Final = audit_generate_rig.candidate.request( + "POST", endpoint, _trace_body(kind, model, f"{marker}-{attempt}", None) + ) + assert response.status_code == 200, response.text + session: Final = _await_spend_row(response.headers["x-litellm-call-id"])["session_id"] + assert session, f"{attempt}: generated session id must be non-empty" + sessions.append(string_value(session)) + assert sessions[0] != sessions[1], f"two id-less calls must not share a session: {sessions}" + assert targets and set(targets) == {expected_target}, targets + + +@pytest.mark.parametrize(("endpoint", "kind", "expected_target"), _AUDIT_ENDPOINTS) +@pytest.mark.parametrize( + ("case", "metadata", "baggage", "expected_status", "expected_session"), + ( + pytest.param( + "r1", "caller_session", None, 200, "caller_session", id="r1_caller_session_no_baggage" + ), + pytest.param("r2", "empty_session", "baggage", 200, "baggage", id="r2_empty_session_baggage"), + pytest.param("r3", "none", None, 400, None, id="r3_nothing_rejected"), + ), +) +def test_audit_reject_policy( + audit_reject_rig: _AuditRig, + endpoint: str, + kind: str, + expected_target: str, + case: str, + metadata: str, + baggage: str, + expected_status: int, + expected_session: str | None, +) -> None: + marker: Final = "auditreject" + uuid.uuid4().hex + provider_secret: Final = "synthetic-provider-secret-" + marker + baggage_session: Final = "baggage-" + marker + caller_session: Final = f"my-session-id-{marker}" + body_metadata: Final = { + "caller_session": {"session_id": caller_session}, + "empty_session": {"session_id": ""}, + "none": None, + }[metadata] + expected: Final = {"caller_session": caller_session, "baggage": baggage_session}[expected_session] if expected_session else None + targets: Final[list[str]] = [] # mutable-ok: the upstream records every hit + + with wire_server(_audit_upstream(provider_secret, marker, targets)) as provider: + model: Final = audit_reject_rig.scenario.model(api_base=provider.url + "/v1", api_key=provider_secret) + response: Final = audit_reject_rig.candidate.request( + "POST", + endpoint, + _trace_body(kind, model, marker, body_metadata), + headers=_w3c_headers(uuid.uuid4().hex, baggage_session if baggage else None), + ) + assert response.status_code == expected_status, response.text + if expected_status != 200: + assert targets == [], targets + rejected_rows: Final = eventually( + lambda: read_rows( + 'SELECT request_id, status FROM "LiteLLM_SpendLogs" WHERE litellm_call_id=%s', + (response.headers["x-litellm-call-id"],), + ), + lambda values: len(values) == 1, + seconds=240, + ) + assert rejected_rows[0]["status"] == "failure", ( + f"rejected call {response.headers['x-litellm-call-id']} must write exactly one failure spend row: {rejected_rows}" + ) + return + received: Final[list[Request]] = [] # mutable-ok: drain() consumes the queue, later polls keep earlier ones + _assert_call(response, received, audit_reject_rig.destination, targets, expected_target, None, expected) + + +@pytest.mark.parametrize(("endpoint", "kind", "expected_target"), _AUDIT_ENDPOINTS) +def test_audit_omit_policy_records_no_session(audit_omit_rig: _AuditRig, endpoint: str, kind: str, expected_target: str) -> None: + marker: Final = "auditomit" + uuid.uuid4().hex + provider_secret: Final = "synthetic-provider-secret-" + marker + targets: Final[list[str]] = [] # mutable-ok: the upstream records every hit + + with wire_server(_audit_upstream(provider_secret, marker, targets)) as provider: + model: Final = audit_omit_rig.scenario.model(api_base=provider.url + "/v1", api_key=provider_secret) + response: Final = audit_omit_rig.candidate.request( + "POST", endpoint, _trace_body(kind, model, marker, None) + ) + assert response.status_code == 200, response.text + call_id: Final = response.headers["x-litellm-call-id"] + assert targets == [expected_target], targets + received: Final[list[Request]] = [] # mutable-ok: drain() consumes the queue, later polls keep earlier ones + span: Final = _await_span(received, audit_omit_rig.destination, call_id) + row: Final = _await_spend_row(call_id) + span_session: Final = _attribute(span.attributes, "session.id") + assert row["session_id"] == span_session, ( + f"call {call_id}: spend session {row['session_id']!r} must match span session {span_session!r}" + ) + + +@pytest.mark.parametrize(("endpoint", "kind", "expected_target"), _AUDIT_ENDPOINTS) +@pytest.mark.parametrize("bad_value", (123, ["x"]), ids=["int", "list"]) +def test_audit_non_string_caller_ids_fall_back_to_w3c( + audit_rig: _AuditRig, endpoint: str, kind: str, expected_target: str, bad_value: JsonValue +) -> None: + marker: Final = "auditbad" + uuid.uuid4().hex + provider_secret: Final = "synthetic-provider-secret-" + marker + header_trace: Final = uuid.uuid4().hex + baggage_session: Final = "baggage-" + marker + targets: Final[list[str]] = [] # mutable-ok: the upstream records every hit + + with wire_server(_audit_upstream(provider_secret, marker, targets)) as provider: + model: Final = audit_rig.scenario.model(api_base=provider.url + "/v1", api_key=provider_secret) + response: Final = audit_rig.candidate.request( + "POST", + endpoint, + _trace_body(kind, model, marker, {"trace_id": bad_value, "session_id": bad_value}), + headers=_w3c_headers(header_trace, baggage_session), + ) + assert response.status_code == 200, response.text + received: Final[list[Request]] = [] # mutable-ok: drain() consumes the queue, later polls keep earlier ones + _assert_call(response, received, audit_rig.destination, targets, expected_target, header_trace, baggage_session) + + +@pytest.mark.parametrize(("endpoint", "kind", "expected_target"), _AUDIT_ENDPOINTS) +@pytest.mark.parametrize( + "metadata", + ({"trace_id": "", "session_id": ""}, {"trace_id": None, "session_id": None}), + ids=["empty", "null"], +) +def test_audit_empty_and_null_caller_ids_fall_back_to_w3c( + audit_rig: _AuditRig, endpoint: str, kind: str, expected_target: str, metadata: dict[str, object] +) -> None: + marker: Final = "auditempty" + uuid.uuid4().hex + provider_secret: Final = "synthetic-provider-secret-" + marker + header_trace: Final = uuid.uuid4().hex + baggage_session: Final = "baggage-" + marker + targets: Final[list[str]] = [] # mutable-ok: the upstream records every hit + + with wire_server(_audit_upstream(provider_secret, marker, targets)) as provider: + model: Final = audit_rig.scenario.model(api_base=provider.url + "/v1", api_key=provider_secret) + response: Final = audit_rig.candidate.request( + "POST", + endpoint, + _trace_body(kind, model, marker, dict(metadata)), + headers=_w3c_headers(header_trace, baggage_session), + ) + assert response.status_code == 200, response.text + received: Final[list[Request]] = [] # mutable-ok: drain() consumes the queue, later polls keep earlier ones + _assert_call(response, received, audit_rig.destination, targets, expected_target, header_trace, baggage_session) + + +@pytest.mark.parametrize(("endpoint", "kind", "expected_target"), _AUDIT_ENDPOINTS[:2]) +def test_audit_five_kilobyte_caller_ids_win_verbatim( + audit_rig: _AuditRig, endpoint: str, kind: str, expected_target: str +) -> None: + marker: Final = "auditbig" + uuid.uuid4().hex + provider_secret: Final = "synthetic-provider-secret-" + marker + header_trace: Final = uuid.uuid4().hex + baggage_session: Final = "baggage-" + marker + caller_trace: Final = "T" * 5120 + caller_session: Final = "S" * 5120 + targets: Final[list[str]] = [] # mutable-ok: the upstream records every hit + + with wire_server(_audit_upstream(provider_secret, marker, targets)) as provider: + model: Final = audit_rig.scenario.model(api_base=provider.url + "/v1", api_key=provider_secret) + response: Final = audit_rig.candidate.request( + "POST", + endpoint, + _trace_body(kind, model, marker, {"trace_id": caller_trace, "session_id": caller_session}), + headers=_w3c_headers(header_trace, baggage_session), + ) + assert response.status_code == 200, response.text + call_id: Final = response.headers["x-litellm-call-id"] + assert targets == [expected_target], targets + received: Final[list[Request]] = [] # mutable-ok: drain() consumes the queue, later polls keep earlier ones + span: Final = _await_span(received, audit_rig.destination, call_id) + assert span.trace_id.hex() != header_trace, f"call {call_id}: caller trace must beat the W3C header" + assert _attribute(span.attributes, "session.id") == caller_session, f"call {call_id}: session.id" + assert _await_spend_row(call_id)["session_id"] == caller_session, f"call {call_id}: spend session" + + +@pytest.mark.parametrize(("endpoint", "kind", "expected_target"), _AUDIT_ENDPOINTS) +def test_audit_identical_requests_twice_log_per_call( + audit_rig: _AuditRig, endpoint: str, kind: str, expected_target: str +) -> None: + marker: Final = "auditdup" + uuid.uuid4().hex + provider_secret: Final = "synthetic-provider-secret-" + marker + header_trace: Final = uuid.uuid4().hex + baggage_session: Final = "baggage-" + marker + caller_trace: Final = uuid.uuid4().hex + caller_session: Final = f"my-session-id-{marker}" + targets: Final[list[str]] = [] # mutable-ok: the upstream records every hit + + with wire_server(_audit_upstream(provider_secret, marker, targets)) as provider: + model: Final = audit_rig.scenario.model(api_base=provider.url + "/v1", api_key=provider_secret) + received: Final[list[Request]] = [] # mutable-ok: drain() consumes the queue, later polls keep earlier ones + for attempt in ("first", "second"): + response: Final = audit_rig.candidate.request( + "POST", + endpoint, + _trace_body( + kind, model, marker, {"trace_id": caller_trace, "session_id": caller_session} + ), + headers=_w3c_headers(header_trace, baggage_session), + ) + assert response.status_code == 200, f"{attempt}: {response.text}" + call_id: Final = response.headers["x-litellm-call-id"] + span: Final = _await_span(received, audit_rig.destination, call_id) + assert span.trace_id.hex() == caller_trace, f"{attempt} {call_id}: trace id" + assert _attribute(span.attributes, "session.id") == caller_session, f"{attempt} {call_id}: session.id" + assert _await_spend_row(call_id)["session_id"] == caller_session, f"{attempt} {call_id}: spend" + assert targets and set(targets) == {expected_target}, targets + + +def test_audit_malformed_w3c_headers_are_ignored(audit_rig: _AuditRig) -> None: + marker: Final = "auditmal" + uuid.uuid4().hex + provider_secret: Final = "synthetic-provider-secret-" + marker + targets: Final[list[str]] = [] # mutable-ok: the upstream records every hit + + with wire_server(_audit_upstream(provider_secret, marker, targets)) as provider: + model: Final = audit_rig.scenario.model(api_base=provider.url + "/v1", api_key=provider_secret) + response: Final = audit_rig.candidate.request( + "POST", + "/v1/chat/completions", + _trace_body("chat", model, marker, None), + headers={"traceparent": "00-zz-00f067aa0ba902b7-01", "baggage": "not-a-session-key"}, + ) + assert response.status_code == 200, response.text + assert targets == ["/v1/chat/completions"], targets + received: Final[list[Request]] = [] # mutable-ok: drain() consumes the queue, later polls keep earlier ones + span: Final = _await_span(received, audit_rig.destination, response.headers["x-litellm-call-id"]) + assert len(span.trace_id.hex()) == 32 and "zz" not in span.trace_id.hex(), span.trace_id.hex() + row: Final = _await_spend_row(response.headers["x-litellm-call-id"]) + span_session: Final = _attribute(span.attributes, "session.id") + assert span_session in (None, row["session_id"]), ( + f"span session {span_session!r} diverges from spend session {row['session_id']!r}" + ) + assert row["session_id"], row + + +def test_audit_unauthenticated_call_leaves_no_spend_row(audit_rig: _AuditRig) -> None: + marker: Final = "auditunauth" + uuid.uuid4().hex + provider_secret: Final = "synthetic-provider-secret-" + marker + targets: Final[list[str]] = [] # mutable-ok: the upstream records every hit + + with wire_server(_audit_upstream(provider_secret, marker, targets)) as provider: + candidate: Final = audit_rig.candidate + model: Final = audit_rig.scenario.model(api_base=provider.url + "/v1", api_key=provider_secret) + denied: Final = candidate.request( + "POST", + "/v1/chat/completions", + _trace_body( + "chat", + model, + marker, + {"trace_id": uuid.uuid4().hex, "session_id": f"my-session-id-{marker}"}, + ), + headers=_w3c_headers(uuid.uuid4().hex, "baggage-" + marker), + key="sk-wrong-key", + ) + assert denied.status_code == 401, denied.text + assert targets == [], targets + control: Final = candidate.request( + "POST", "/v1/chat/completions", _trace_body("chat", model, marker + "-control", None) + ) + assert control.status_code == 200, control.text + _await_spend_row(control.headers["x-litellm-call-id"]) + assert ( + read_rows( + 'SELECT request_id FROM "LiteLLM_SpendLogs" WHERE litellm_call_id=%s', + (denied.headers.get("x-litellm-call-id") or "",), + ) + == [] + ), "an unauthenticated call must not write a spend row" + + +@pytest.mark.parametrize(("endpoint", "kind", "expected_target"), _AUDIT_ENDPOINTS[:2]) +@pytest.mark.parametrize("upstream_status", (500, 401), ids=["upstream_500", "upstream_401"]) +def test_audit_upstream_error_still_logs_caller_session( + audit_rig: _AuditRig, + endpoint: str, + kind: str, + expected_target: str, + upstream_status: int, +) -> None: + marker: Final = "auditerr" + uuid.uuid4().hex + provider_secret: Final = "synthetic-provider-secret-" + marker + caller_trace: Final = uuid.uuid4().hex + caller_session: Final = f"my-session-id-{marker}" + targets: Final[list[str]] = [] # mutable-ok: the upstream records every hit + + def upstream(request: Request) -> Reply: + if request.body: + targets.append(request.target) + return Reply(status=upstream_status, body=b'{"error": {"message": "scripted upstream failure"}}') + + with wire_server(upstream) as provider: + model: Final = audit_rig.scenario.model(api_base=provider.url + "/v1", api_key=provider_secret) + response: Final = audit_rig.candidate.request( + "POST", + endpoint, + _trace_body(kind, model, marker, {"trace_id": caller_trace, "session_id": caller_session}), + headers=_w3c_headers(uuid.uuid4().hex, "baggage-" + marker), + ) + assert response.status_code == upstream_status, response.text + assert targets and set(targets) == {expected_target}, targets + call_id: Final = response.headers["x-litellm-call-id"] + row: Final = _await_spend_row(call_id) + assert row["session_id"] == caller_session, f"call {call_id}: failure spend session {row}" + + +def test_audit_sink_rejection_does_not_break_the_caller(gateway: Gateway, tmp_path: Path) -> None: + marker: Final = "auditreject" + uuid.uuid4().hex + provider_secret: Final = "synthetic-provider-secret-" + marker + targets: Final[list[str]] = [] # mutable-ok: the upstream records every hit + attempts: Final[list[int]] = [] # mutable-ok: sink status sequence counter + + def langfuse(request: Request) -> Reply: + if request.method == "GET" and request.target.startswith(PROJECTS_PATH): + return _projects() + if request.target == TRACES_PATH: + attempts.append(1) + if len(attempts) == 1: + return Reply(status=403) + if len(attempts) == 2: + return Reply(status=404) + return Reply(body=b"", content_type="application/x-protobuf") + + with ( + wire_server(_audit_upstream(provider_secret, marker, targets)) as provider, + wire_server(langfuse) as destination, + owned_proxy( + gateway, tmp_path, _langfuse_environment(destination), config=_langfuse_config(tmp_path) + ) as candidate, + candidate.scenario() as scenario, + ): + model: Final = scenario.model(api_base=provider.url + "/v1", api_key=provider_secret) + received: Final[list[Request]] = [] # mutable-ok: drain() consumes the queue, later polls keep earlier ones + for attempt in ("first", "second", "third"): + response: Final = candidate.request( + "POST", + "/v1/chat/completions", + _trace_body( + "chat", + model, + f"{marker}-{attempt}", + {"trace_id": uuid.uuid4().hex, "session_id": f"my-session-id-{marker}-{attempt}"}, + ), + ) + assert response.status_code == 200, f"{attempt}: {response.text}" + _await_spend_row(response.headers["x-litellm-call-id"]) + assert targets == ["/v1/chat/completions"] * 3, targets + + +@pytest.mark.parametrize(("endpoint", "kind", "expected_target"), _AUDIT_ENDPOINTS[:2]) +def test_audit_string_metadata_body_does_not_crash(audit_rig: _AuditRig, endpoint: str, kind: str, expected_target: str) -> None: + marker: Final = "auditstr" + uuid.uuid4().hex + provider_secret: Final = "synthetic-provider-secret-" + marker + targets: Final[list[str]] = [] # mutable-ok: the upstream records every hit + + with wire_server(_audit_upstream(provider_secret, marker, targets)) as provider: + candidate: Final = audit_rig.candidate + model: Final = audit_rig.scenario.model(api_base=provider.url + "/v1", api_key=provider_secret) + body: Final = {**_trace_body(kind, model, marker, None), "metadata": "x"} + response: Final = candidate.request( + "POST", endpoint, body, headers=_w3c_headers(uuid.uuid4().hex, "baggage-" + marker) + ) + assert response.status_code < 500, f"metadata string must not crash the proxy: {response.status_code} {response.text}" + follow_up: Final = candidate.request( + "POST", "/v1/chat/completions", _trace_body("chat", model, marker + "-follow", None) + ) + assert follow_up.status_code == 200, follow_up.text + + +@pytest.mark.parametrize(("endpoint", "kind", "expected_target"), _AUDIT_ENDPOINTS[:2]) +def test_audit_key_metadata_session_still_honoured(audit_rig: _AuditRig, endpoint: str, kind: str, expected_target: str) -> None: + marker: Final = "auditkey" + uuid.uuid4().hex + provider_secret: Final = "synthetic-provider-secret-" + marker + key_session: Final = f"key-session-{marker}" + targets: Final[list[str]] = [] # mutable-ok: the upstream records every hit + + with wire_server(_audit_upstream(provider_secret, marker, targets)) as provider: + model: Final = audit_rig.scenario.model(api_base=provider.url + "/v1", api_key=provider_secret) + token: Final = audit_rig.scenario.key(metadata={"session_id": key_session}) + response: Final = audit_rig.candidate.request( + "POST", + endpoint, + _trace_body(kind, model, marker, None), + headers=_w3c_headers(uuid.uuid4().hex, "baggage-" + marker), + key=token, + ) + assert response.status_code == 200, response.text + call_id: Final = response.headers["x-litellm-call-id"] + assert targets == [expected_target], targets + received: Final[list[Request]] = [] # mutable-ok: drain() consumes the queue, later polls keep earlier ones + span: Final = _await_span(received, audit_rig.destination, call_id) + row: Final = _await_spend_row(call_id) + assert _attribute(span.attributes, "session.id") == row["session_id"], ( + f"call {call_id}: spend session {row['session_id']!r} must match span session" + ) + + +@pytest.mark.parametrize(("endpoint", "kind", "expected_target"), _AUDIT_ENDPOINTS[:2]) +@pytest.mark.parametrize("metadata_mode", ("both", "none"), ids=["caller_ids", "no_metadata"]) +def test_audit_cache_hit_call_keeps_winning_ids( + audit_rig: _AuditRig, endpoint: str, kind: str, expected_target: str, metadata_mode: str +) -> None: + marker: Final = "auditcache" + uuid.uuid4().hex + provider_secret: Final = "synthetic-provider-secret-" + marker + header_trace: Final = uuid.uuid4().hex + baggage_session: Final = "baggage-" + marker + caller_trace: Final = uuid.uuid4().hex + caller_session: Final = f"my-session-id-{marker}" + metadata: Final = ( + {"trace_id": caller_trace, "session_id": caller_session} if metadata_mode == "both" else None + ) + expected_trace: Final = caller_trace if metadata_mode == "both" else header_trace + expected_session: Final = caller_session if metadata_mode == "both" else baggage_session + targets: Final[list[str]] = [] # mutable-ok: the upstream records every hit + + with wire_server(_audit_upstream(provider_secret, marker, targets)) as provider: + candidate: Final = audit_rig.candidate + model: Final = audit_rig.scenario.model(api_base=provider.url + "/v1", api_key=provider_secret) + received: Final[list[Request]] = [] # mutable-ok: drain() consumes the queue, later polls keep earlier ones + body: Final = {key: value for key, value in _trace_body(kind, model, marker, metadata).items() if key != "cache"} + headers: Final = _w3c_headers(header_trace, baggage_session) + first: Final = candidate.request("POST", endpoint, body, headers=headers) + assert first.status_code == 200, first.text + second: Final = candidate.request("POST", endpoint, body, headers=headers) + assert second.status_code == 200, second.text + assert second.headers.get("x-litellm-cache-key"), ( + f"second identical call must be a cache hit: {dict(second.headers)}" + ) + call_id: Final = second.headers["x-litellm-call-id"] + assert targets == [expected_target], targets + span: Final = _await_span(received, audit_rig.destination, call_id) + if metadata_mode == "both": + assert span.trace_id.hex() == expected_trace, f"cache hit {call_id}: trace id" + assert _attribute(span.attributes, "session.id") == expected_session, ( + f"cache hit {call_id}: session.id" + ) + else: + span_session: Final = _attribute(span.attributes, "session.id") + print( + f"E2 record: cache-hit span trace={span.trace_id.hex()} session={span_session!r} " + f"header trace={header_trace} baggage session={baggage_session!r}" + ) + assert len(span.trace_id.hex()) == 32, span.trace_id.hex() + assert _await_spend_row(call_id)["session_id"] == expected_session, f"cache hit {call_id}: spend" + + +def test_audit_concurrent_requests_each_keep_their_caller_ids(audit_rig: _AuditRig) -> None: + marker: Final = "auditconc" + uuid.uuid4().hex + provider_secret: Final = "synthetic-provider-secret-" + marker + header_trace: Final = uuid.uuid4().hex + baggage_session: Final = "baggage-" + marker + targets: Final[list[str]] = [] # mutable-ok: the upstream records every hit + + with wire_server(_audit_upstream(provider_secret, marker, targets)) as provider: + candidate: Final = audit_rig.candidate + model: Final = audit_rig.scenario.model(api_base=provider.url + "/v1", api_key=provider_secret) + jobs: Final = tuple( + (index, endpoint, kind, uuid.uuid4().hex, f"my-session-id-{marker}-{index}") + for index, (endpoint, kind, _) in tuple( + enumerate(tuple(row.values for row in _AUDIT_ENDPOINTS) * 4) + )[:10] + ) + + def fire(job: tuple[int, str, str, str, str]) -> tuple[str, str, httpx.Response]: + index, endpoint, kind, caller_trace, caller_session = job + response: Final = candidate.request( + "POST", + endpoint, + _trace_body( + kind, model, f"{marker}-{index}", {"trace_id": caller_trace, "session_id": caller_session} + ), + headers=_w3c_headers(header_trace, baggage_session), + ) + return caller_trace, caller_session, response + + with ThreadPoolExecutor(max_workers=10) as pool: + answered: Final = tuple(pool.map(fire, jobs)) + received: Final[list[Request]] = [] # mutable-ok: drain() consumes the queue, later polls keep earlier ones + assert len(targets) == 10, targets + for caller_trace, caller_session, response in answered: + assert response.status_code == 200, response.text + _assert_span_spend(response, received, audit_rig.destination, caller_trace, caller_session) + + +def test_audit_sink_outage_mid_burst_loses_no_spend_rows(gateway: Gateway, tmp_path: Path) -> None: + marker: Final = "auditoutage" + uuid.uuid4().hex + provider_secret: Final = "synthetic-provider-secret-" + marker + header_trace: Final = uuid.uuid4().hex + baggage_session: Final = "baggage-" + marker + targets: Final[list[str]] = [] # mutable-ok: the upstream records every hit + outage: Final = threading.Event() + + def langfuse(request: Request) -> Reply: + if outage.is_set(): + return Reply(status=503) + if request.method == "GET" and request.target.startswith(PROJECTS_PATH): + return _projects() + return Reply(body=b"", content_type="application/x-protobuf") + + with ( + wire_server(_audit_upstream(provider_secret, marker, targets)) as provider, + wire_server(langfuse) as destination, + owned_proxy( + gateway, tmp_path, _langfuse_environment(destination), config=_langfuse_config(tmp_path) + ) as candidate, + candidate.scenario() as scenario, + ): + model: Final = scenario.model(api_base=provider.url + "/v1", api_key=provider_secret) + jobs: Final = tuple( + (index, endpoint, kind, uuid.uuid4().hex, f"my-session-id-{marker}-{index}") + for index, (endpoint, kind, _) in enumerate(tuple(row.values for row in _AUDIT_ENDPOINTS) * 10) + ) + + def fire(job: tuple[int, str, str, str, str]) -> tuple[str, str, httpx.Response]: + index, endpoint, kind, caller_trace, caller_session = job + stream: Final = index % 3 == 0 + response: Final = candidate.request( + "POST", + endpoint, + { + **_trace_body( + kind, + model, + f"{marker}-{index}", + {"trace_id": caller_trace, "session_id": caller_session}, + ), + "stream": stream, + }, + headers=_w3c_headers(header_trace, baggage_session), + ) + return caller_trace, caller_session, response + + outage.set() + with ThreadPoolExecutor(max_workers=30) as pool: + first_wave: Final = tuple(pool.map(fire, jobs[:10])) + unhealthy: Final = candidate.request("GET", "/health/services?service=langfuse") + assert unhealthy.status_code != 200 or "unhealthy" in unhealthy.text, ( + f"langfuse must report unhealthy while the sink 503s: {unhealthy.status_code} {unhealthy.text}" + ) + outage.clear() + healthy: Final = candidate.request("GET", "/health/services?service=langfuse") + assert healthy.status_code == 200, healthy.text + with ThreadPoolExecutor(max_workers=30) as pool: + answered: Final = first_wave + tuple(pool.map(fire, jobs[10:])) + for _, _, response in answered: + assert response.status_code == 200, response.text + expected_by_call: Final = { + response.headers["x-litellm-call-id"]: caller_session + for caller_session, response in ((session, res) for _, session, res in answered) + } + call_ids: Final = sorted(expected_by_call) + received: Final[list[Request]] = [] # mutable-ok: drain() consumes the queue, later polls keep earlier ones + delivered: Final = eventually( + lambda: ( + received.extend(destination.drain()) or tuple( + span + for span in _spans(received) + if _attribute(span.attributes, "langfuse.observation.metadata.litellm_call_id") + in expected_by_call + ) + ), + lambda spans: len(spans) >= len(answered), + seconds=60, + ) + spans_per_call: Final = { + call_id: sum( + 1 + for span in delivered + if _attribute(span.attributes, "langfuse.observation.metadata.litellm_call_id") == call_id + ) + for call_id in call_ids + } + assert sorted(spans_per_call.values()) == [1] * len(answered), ( + f"each call id must arrive on exactly one span: {spans_per_call}" + ) + for span in delivered: + span_call: Final = str( + _attribute(span.attributes, "langfuse.observation.metadata.litellm_call_id") + ) + assert _attribute(span.attributes, "session.id") == expected_by_call[span_call], ( + f"call {span_call}: session.id" + ) + spend_rows: Final = eventually( + lambda: read_rows( + 'SELECT session_id, litellm_call_id FROM "LiteLLM_SpendLogs" WHERE litellm_call_id = ANY(%s)', + (call_ids,), + ), + lambda values: len(values) == len(answered), + seconds=150, + ) + for row in spend_rows: + row_call: Final = string_value(row["litellm_call_id"]) + assert row["session_id"] == expected_by_call[row_call], f"call {row_call}: spend {row}" + + +def test_audit_surviving_worker_keeps_serving_after_kill(gateway: Gateway, tmp_path: Path) -> None: + marker: Final = "auditworker" + uuid.uuid4().hex + provider_secret: Final = "synthetic-provider-secret-" + marker + header_trace: Final = uuid.uuid4().hex + baggage_session: Final = "baggage-" + marker + targets: Final[list[str]] = [] # mutable-ok: the upstream records every hit + + with ( + wire_server(_audit_upstream(provider_secret, marker, targets)) as provider, + wire_server(_audit_sink()) as destination, + owned_proxy_process( + gateway, + tmp_path, + _langfuse_environment(destination), + config=_langfuse_config(tmp_path), + workers=2, + ) as owned, + owned.gateway.scenario() as scenario, + ): + candidate: Final = owned.gateway + model: Final = scenario.model(api_base=provider.url + "/v1", api_key=provider_secret) + workers: Final = tuple( + member for member in group_members(owned.process.pid) if member.pid != owned.process.pid + ) + assert workers, "expected worker processes under the owned proxy" + workers[0].send_signal(signal.SIGKILL) + psutil.wait_procs([workers[0]], timeout=5) + + def fire(index: int) -> tuple[str, httpx.Response]: + caller_trace: Final = uuid.uuid4().hex + caller_session: Final = f"my-session-id-{marker}-{index}" + response: Final = candidate.request( + "POST", + "/v1/chat/completions", + _trace_body( + "chat", + model, + f"{marker}-{index}", + {"trace_id": caller_trace, "session_id": caller_session}, + ), + headers=_w3c_headers(header_trace, baggage_session), + ) + return caller_session, response + + with ThreadPoolExecutor(max_workers=10) as pool: + answered: Final = tuple(pool.map(fire, range(10))) + received: Final[list[Request]] = [] # mutable-ok: drain() consumes the queue, later polls keep earlier ones + for caller_session, response in answered: + assert response.status_code == 200, response.text + call_id: Final = response.headers["x-litellm-call-id"] + span: Final = _await_span(received, destination, call_id) + assert _attribute(span.attributes, "session.id") == caller_session, f"{call_id}: session.id" + expected_by_call: Final = { + response.headers["x-litellm-call-id"]: caller_session + for caller_session, response in answered + } + spend_rows: Final = eventually( + lambda: read_rows( + 'SELECT session_id, litellm_call_id FROM "LiteLLM_SpendLogs" WHERE litellm_call_id = ANY(%s)', + (list(expected_by_call),), + ), + lambda values: len(values) == len(answered), + seconds=150, + ) + for row in spend_rows: + row_call: Final = string_value(row["litellm_call_id"]) + assert row["session_id"] == expected_by_call[row_call], f"call {row_call}: spend {row}" diff --git a/tests/integration/providers/test_bedrock_batch_retrieve_status_wire.py b/tests/integration/providers/test_bedrock_batch_retrieve_status_wire.py new file mode 100644 index 00000000000..151ca2017d5 --- /dev/null +++ b/tests/integration/providers/test_bedrock_batch_retrieve_status_wire.py @@ -0,0 +1,112 @@ +import json +import urllib.parse +import uuid +from collections.abc import Mapping +from itertools import count +from pathlib import Path +from typing import Final + +import pytest +from integration._support.client import Gateway +from integration._support.process import owned_proxy +from integration._support.wire import Reply, Request, wire_server +from integration.providers.test_bedrock_batch_blank_s3_env_wire import ( + BUCKET, + JOB_ARN_PREFIX, + MODEL_ID, + REGION, + ROLE_ARN, + _tls_context, + bedrock_tunnel, + s3_peer, +) + +LIFECYCLE: Final = ("Submitted", "Validating", "Scheduled", "InProgress") +EXPECTED: Final = ("validating", "validating", "in_progress", "in_progress") + + +def create_peer(job_arn: str, request: Request) -> Reply: + if request.method == "POST" and request.target == "/model-invocation-job": + return Reply(body=json.dumps({"jobArn": job_arn}).encode()) + return Reply(status=404, body=b'{"message": "not scripted"}') + + +def get_peer(calls: count, job: Mapping[str, object]) -> Reply: + index: Final = min(next(calls), len(LIFECYCLE) - 1) + body: Final = { + **job, + "status": LIFECYCLE[index], + "submitTime": "2026-10-01T00:00:00Z", + "lastModifiedTime": "2026-10-01T00:01:00Z", + } + return Reply(body=json.dumps(body).encode()) + + +@pytest.mark.timeout(180) +def test_bedrock_batch_retrieve_reports_lifecycle_status_in_order(gateway: Gateway, tmp_path: Path) -> None: + job_arn: Final = JOB_ARN_PREFIX + uuid.uuid4().hex + job: Final = { + "jobArn": job_arn, + "modelId": MODEL_ID, + "inputDataConfig": {"s3InputDataConfig": {"s3Uri": f"s3://{BUCKET}/input.jsonl"}}, + "outputDataConfig": {"s3OutputDataConfig": {"s3Uri": f"s3://{BUCKET}/out/"}}, + } + calls: Final = count() + environment: Final = { + "SSL_VERIFY": "False", + "AWS_EC2_METADATA_DISABLED": "true", + "HTTP_PROXY": "", + "NO_PROXY": "127.0.0.1,localhost", + } + with ( + wire_server(s3_peer) as s3, + wire_server(lambda request: create_peer(job_arn, request), tls=_tls_context(tmp_path)) as bedrock, + wire_server(lambda request: get_peer(calls, job)) as jobs, + bedrock_tunnel(bedrock) as tunnel, + owned_proxy( + gateway, + tmp_path, + {**environment, "HTTPS_PROXY": tunnel.url, "AWS_ENDPOINT_URL_BEDROCK": jobs.url}, + ) as candidate, + candidate.scenario() as scenario, + ): + model: Final = scenario.model( + model=f"bedrock/{MODEL_ID}", + api_key=None, + api_base=None, + aws_access_key_id="AKIAIOSFODNN7EXAMPLE", + aws_secret_access_key="wJalrXUtnFEMI/K7MDENG/bPxRfiCYEXAMPLEKEY", + aws_region_name=REGION, + s3_bucket_name=BUCKET, + s3_endpoint_url=s3.url, + aws_batch_role_arn=ROLE_ARN, + ) + line: Final = { + "custom_id": "req-1", + "method": "POST", + "url": "/v1/chat/completions", + "body": {"model": model, "messages": [{"role": "user", "content": "ping"}], "max_tokens": 8}, + } + uploaded: Final = candidate.request_multipart( + "/v1/files", + {"purpose": "batch", "target_model_names": model}, + {"file": ("in.jsonl", (json.dumps(line) + "\n").encode(), "application/jsonl")}, + ) + assert uploaded.status_code == 200, uploaded.text + created: Final = candidate.request( + "POST", + "/v1/batches", + {"input_file_id": uploaded.json()["id"], "endpoint": "/v1/chat/completions", "completion_window": "24h"}, + ) + assert created.status_code == 200, created.text + batch_id: Final = created.json()["id"] + + observed: Final = tuple(candidate.get(f"/v1/batches/{batch_id}")["status"] for _ in range(len(LIFECYCLE))) + assert observed == EXPECTED, observed + + job_id: Final = job_arn.rsplit("/", 1)[-1] + drained: Final = jobs.drain() + gets: Final = tuple( + request for request in drained if request.method == "GET" and job_id in urllib.parse.unquote(request.target) + ) + assert len(gets) == len(LIFECYCLE), [(r.method, r.target) for r in drained] diff --git a/tests/integration/providers/test_decisions_chaos.py b/tests/integration/providers/test_decisions_chaos.py index 2e88cdbbc7f..63b4a482e92 100644 --- a/tests/integration/providers/test_decisions_chaos.py +++ b/tests/integration/providers/test_decisions_chaos.py @@ -27,7 +27,7 @@ _API_KEY: Final = "synthetic-decisions-key" _JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) _STARTED_WORKER: Final = re.compile(r"Started server process \[(\d+)\]") _QUESTIONS: Final[dict[str, JsonValue]] = {"fine": {"type": "noul", "instructions": "Is the state fine?"}} -_ROUTES: Final = ("/v1/decisions", "/decisions") +_ROUTES: Final = ("/v1/systemone", "/systemone") @dataclass(frozen=True, slots=True) @@ -221,7 +221,7 @@ async def test_worker_sigkill_mid_burst_leaves_the_sibling_serving_the_default_m release.set() served: Final = await burst assert len(served) == held_by[survivor_pid], (held_by, len(served)) - follow_up: Final = _Call(route="/decisions", marker=f"ok-{uuid.uuid4().hex}", fail=False) + follow_up: Final = _Call(route="/systemone", marker=f"ok-{uuid.uuid4().hex}", fail=False) (answered,) = await _burst(base_url, candidate.key, None, (follow_up,)) await asyncio.to_thread( eventually, diff --git a/tests/integration/providers/test_decisions_wire.py b/tests/integration/providers/test_decisions_wire.py index 9878559cb07..84b31088d60 100644 --- a/tests/integration/providers/test_decisions_wire.py +++ b/tests/integration/providers/test_decisions_wire.py @@ -25,6 +25,7 @@ _PASS_THROUGH_MODEL: Final = "gpt-6-luna" _PASS_THROUGH_AUTHORIZATION: Final = "Bearer customer-held-upstream-key" _PASS_THROUGH_NEIGHBOUR: Final = "decisions-beside-a-pass-through" _JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) +_QUESTION_MAPPINGS: Final = TypeAdapter(dict[str, dict[str, object]]) _USAGE: Final[dict[str, JsonValue]] = {"input_tokens": 367, "output_tokens": 3} _STATE: Final[dict[str, JsonValue]] = {"ticket": "The export job hangs at 99%", "component": "billing"} _QUESTIONS: Final[dict[str, JsonValue]] = { @@ -32,6 +33,7 @@ _QUESTIONS: Final[dict[str, JsonValue]] = { "severity": {"type": "choice", "criteria": {"low": "cosmetic", "high": "blocks users"}, "weight": 2}, "confidence": {"type": "score", "instructions": "How sure are you?", "criteria": ["unsure", "sure"]}, } +_SDK_QUESTIONS: Final = _QUESTION_MAPPINGS.validate_python(_QUESTIONS) _ANSWERS: Final[dict[str, JsonValue]] = { "defect": {"type": "noul", "noul": 0.93}, "severity": {"type": "choice", "choice": "high", "confidence": 0.8, "probabilities": {"low": 0.2, "high": 0.8}}, @@ -103,6 +105,11 @@ _PROVIDERS: Final = ( ), ) _PERPLEXITY: Final = _PROVIDERS[0] +_OPENROUTER: Final = _PROVIDERS[2] +_OPENROUTER_CHAT_MODEL: Final = "openrouter/openai/gpt-5-mini" +_UNSUPPORTED_PROVIDER_MODEL: Final = "openai/gpt-6-luna" +_CONNECTION_ERROR: Final = "litellm.APIConnectionError" +_GENERIC_API_ERROR: Final = "litellm.APIError" _INVALID_BODIES: Final[tuple[tuple[str, dict[str, JsonValue]], ...]] = ( ("missing questions", {"state": _STATE}), ("missing state", {"questions": _QUESTIONS}), @@ -151,7 +158,7 @@ def _deployment(scenario: Scenario, handle: ScenarioHandle, provider: _Provider) def _decide(gateway: Gateway, model: str, *, key: str | None = None, **extra: JsonValue) -> httpx.Response: return gateway.request( - "POST", "/v1/decisions", {"model": model, "state": _STATE, "questions": _QUESTIONS, **extra}, key=key + "POST", "/v1/systemone", {"model": model, "state": _STATE, "questions": _QUESTIONS, **extra}, key=key ) @@ -177,6 +184,12 @@ def _spend_row(call_id: str) -> dict[str, JsonValue]: return rows[0] +def _assert_connection_error(response: httpx.Response) -> None: + assert 500 <= response.status_code < 600, response.text + assert _CONNECTION_ERROR in response.text, response.text + assert _GENERIC_API_ERROR not in response.text, response.text + + def _free_closed_port() -> int: with socket.socket() as probe: probe.bind(("127.0.0.1", 0)) @@ -257,10 +270,10 @@ async def test_sdk_sync_and_async_clients_send_the_same_request(gateway: Gateway with gateway.scenario() as scenario: handle: Final = _register(scenario, _answer_body(provider)) synchronous: Final = litellm.decisions( - model=provider.model, state=_STATE, questions=_QUESTIONS, api_base=handle.api_base(), api_key=_API_KEY + model=provider.model, state=_STATE, questions=_SDK_QUESTIONS, api_base=handle.api_base(), api_key=_API_KEY ) asynchronous: Final = await litellm.adecisions( - model=provider.model, state=_STATE, questions=_QUESTIONS, api_base=handle.api_base(), api_key=_API_KEY + model=provider.model, state=_STATE, questions=_SDK_QUESTIONS, api_base=handle.api_base(), api_key=_API_KEY ) for response in (synchronous, asynchronous): assert response.model_dump(mode="json") == { @@ -297,7 +310,7 @@ def test_invalid_bodies_are_refused_at_the_gateway_without_an_upstream_call(gate handle: Final = _register(scenario, _answer_body(_PERPLEXITY)) model: Final = _deployment(scenario, handle, _PERPLEXITY) for label, body in _INVALID_BODIES: - response: Final = gateway.request("POST", "/v1/decisions", {"model": model, **body}) + response: Final = gateway.request("POST", "/v1/systemone", {"model": model, **body}) assert response.status_code == 400, (label, response.text) assert "Invalid Decisions request" in response.text, (label, response.text) assert _upstream_calls(gateway, handle) == [] @@ -316,7 +329,7 @@ def test_key_checks_match_chat(gateway: Gateway) -> None: handle: Final = _register(scenario, _answer_body(_PERPLEXITY)) model: Final = _deployment(scenario, handle, _PERPLEXITY) anonymous: Final = gateway.client.post( - "/v1/decisions", json={"model": model, "state": _STATE, "questions": _QUESTIONS} + "/v1/systemone", json={"model": model, "state": _STATE, "questions": _QUESTIONS} ) assert anonymous.status_code == 401, anonymous.text restricted: Final = scenario.key(models=[f"other-{uuid.uuid4().hex}"]) @@ -382,7 +395,7 @@ def test_a_deployment_opted_into_client_api_base_sends_decisions_and_chat_to_the assert _calls_to(observed, configured) == [] -def test_a_config_pass_through_at_v1_decisions_keeps_answering_and_the_native_api_serves_decisions( +def test_a_config_pass_through_at_v1_decisions_keeps_answering_and_the_native_api_serves_system_one( gateway: Gateway, tmp_path: Path ) -> None: with gateway.scenario() as scenario: @@ -394,9 +407,11 @@ def test_a_config_pass_through_at_v1_decisions_keeps_answering_and_the_native_ap tmp_path, f"{pass_through_target.api_base()}/v1/decisions", native_target.api_base() ) with owned_proxy_process(gateway, tmp_path, {}, config=config) as owned: - through: Final = _decide(owned.gateway, _PASS_THROUGH_MODEL) + through: Final = owned.gateway.request( + "POST", "/v1/decisions", {"model": _PASS_THROUGH_MODEL, "state": _STATE, "questions": _QUESTIONS} + ) native: Final = owned.gateway.request( - "POST", "/decisions", {"model": _PASS_THROUGH_NEIGHBOUR, "state": _STATE, "questions": _QUESTIONS} + "POST", "/systemone", {"model": _PASS_THROUGH_NEIGHBOUR, "state": _STATE, "questions": _QUESTIONS} ) assert through.status_code == 200, through.text assert through.json() == {"model": _PASS_THROUGH_MODEL, "answers": _ANSWERS, "usage": _USAGE} @@ -441,16 +456,75 @@ def test_upstream_success_without_answers_is_a_gateway_side_server_error(gateway assert _spend_row(response.headers["x-litellm-call-id"])["status"] == "failure" -def test_unreachable_upstream_fails_only_its_own_deployment(gateway: Gateway) -> None: +@pytest.mark.parametrize("provider", _PROVIDERS, ids=lambda provider: provider.name) +def test_unreachable_upstream_fails_only_its_own_deployment(gateway: Gateway, provider: _Provider) -> None: with gateway.scenario() as scenario: - handle: Final = _register(scenario, _answer_body(_PERPLEXITY)) - healthy: Final = _deployment(scenario, handle, _PERPLEXITY) + handle: Final = _register(scenario, _answer_body(provider)) + healthy: Final = _deployment(scenario, handle, provider) dead: Final = scenario.model( - model=_PERPLEXITY.model, api_base=f"http://127.0.0.1:{_free_closed_port()}", api_key=_API_KEY + model=provider.model, api_base=f"http://127.0.0.1:{_free_closed_port()}", api_key=provider.api_key ) failed: Final = _decide(gateway, dead, num_retries=0) - assert 500 <= failed.status_code < 600, failed.text + _assert_connection_error(failed) assert _spend_row(failed.headers["x-litellm-call-id"])["status"] == "failure" served: Final = _decide(gateway, healthy) assert served.status_code == 200, served.text assert len(_upstream_calls(gateway, handle)) == 1 + + +def test_an_unreachable_openrouter_deployment_reports_a_connection_error_on_chat_embeddings_and_decisions( + gateway: Gateway, +) -> None: + with gateway.scenario() as scenario: + dead_api_base: Final = f"http://127.0.0.1:{_free_closed_port()}" + chat_model: Final = scenario.model(model=_OPENROUTER_CHAT_MODEL, api_base=dead_api_base, api_key=_API_KEY) + decisions_model: Final = scenario.model(model=_OPENROUTER.model, api_base=dead_api_base, api_key=_API_KEY) + chat: Final = _chat(gateway, chat_model, num_retries=0) + streamed: Final = _chat(gateway, chat_model, num_retries=0, stream=True) + embeddings: Final = gateway.request( + "POST", "/v1/embeddings", {"model": chat_model, "input": "hi", "num_retries": 0} + ) + decisions: Final = _decide(gateway, decisions_model, num_retries=0) + for response in (chat, streamed, embeddings, decisions): + _assert_connection_error(response) + for response in (chat, decisions): + assert _spend_row(response.headers["x-litellm-call-id"])["status"] == "failure" + + +async def test_sdk_openrouter_connection_failures_raise_a_connection_error(gateway: Gateway) -> None: + dead_api_base: Final = f"http://127.0.0.1:{_free_closed_port()}" + with pytest.raises(litellm.APIConnectionError): + litellm.completion( + model=_OPENROUTER_CHAT_MODEL, + messages=[{"role": "user", "content": "hi"}], + api_base=dead_api_base, + api_key=_API_KEY, + ) + with pytest.raises(litellm.APIConnectionError): + await litellm.acompletion( + model=_OPENROUTER_CHAT_MODEL, + messages=[{"role": "user", "content": "hi"}], + api_base=dead_api_base, + api_key=_API_KEY, + ) + with pytest.raises(litellm.APIConnectionError): + litellm.decisions( + model=_OPENROUTER.model, state=_STATE, questions=_SDK_QUESTIONS, api_base=dead_api_base, api_key=_API_KEY + ) + with pytest.raises(litellm.APIConnectionError): + await litellm.adecisions( + model=_OPENROUTER.model, state=_STATE, questions=_SDK_QUESTIONS, api_base=dead_api_base, api_key=_API_KEY + ) + + +def test_a_deployment_whose_provider_has_no_decisions_support_is_refused_naming_every_supported_provider( + gateway: Gateway, +) -> None: + with gateway.scenario() as scenario: + handle: Final = _register(scenario, _answer_body(_PERPLEXITY)) + model: Final = scenario.model(model=_UNSUPPORTED_PROVIDER_MODEL, api_base=handle.api_base(), api_key=_API_KEY) + response: Final = _decide(gateway, model) + assert response.status_code == 400, response.text + for provider in _PROVIDERS: + assert provider.name in response.text, response.text + assert _upstream_calls(gateway, handle) == [] diff --git a/tests/integration/providers/test_ollama_prompt_tools_wire.py b/tests/integration/providers/test_ollama_prompt_tools_wire.py index 115fdc65e33..f3aa18cb366 100644 --- a/tests/integration/providers/test_ollama_prompt_tools_wire.py +++ b/tests/integration/providers/test_ollama_prompt_tools_wire.py @@ -1,3 +1,4 @@ +import asyncio import itertools import json import uuid @@ -8,10 +9,11 @@ from typing import Final import anthropic import openai import pytest -from openai.types.chat import ChatCompletionChunk +from openai.types import Completion +from openai.types.chat import ChatCompletion, ChatCompletionChunk from openai.types.chat.chat_completion_chunk import Choice as ChunkChoice from openai.types.chat.chat_completion_chunk import ChoiceDeltaToolCall -from integration._support.client import Gateway, eventually +from integration._support.client import Gateway, eventually, object_value from integration._support.database import read_rows from integration._support.wire import Reply, Request, Wire, wire_server from pydantic import JsonValue, TypeAdapter @@ -343,6 +345,7 @@ async def test_async_openai_sdk_stream_answers_the_tool_result_in_plain_text(gat extra_body={"cache": _NO_CACHE}, ) chunks: Final = [chunk async for chunk in stream] + assert {chunk.id for chunk in chunks} == {chunks[0].id} choices: Final = tuple(_stream_choices(chunks)) assert "".join(choice.delta.content or "" for choice in choices) == _ANSWER assert tuple(_delta_tool_calls(choices)) == () @@ -497,6 +500,7 @@ async def test_async_openai_sdk_responses_stream_emits_the_function_call_item(ga assert json.loads(item.arguments) == _ARGUMENTS completed: Final = [event for event in events if event.type == "response.completed"] assert len(completed) == 1 + assert [item.type for item in completed[0].response.output] == ["function_call"], completed[0].response.output final_calls: Final = [item for item in completed[0].response.output if item.type == "function_call"] assert [(item.name, json.loads(item.arguments)) for item in final_calls] == [("get_weather", _ARGUMENTS)] assert completed[0].response.output_text == "" @@ -640,8 +644,8 @@ def test_ollama_server_error_on_the_tool_result_turn_does_not_take_the_deploymen pytest.param("r" * 5120, "r" * 5120, id="5kb-string-forwarded-intact"), pytest.param( [{"type": "text", "text": "Paris: 22 degrees"}, {"type": "text", "text": "clear skies"}], - "Paris: 22 degreesclear skies", - id="text-parts-joined", + "Paris: 22 degrees\nclear skies", + id="text-parts-joined-by-newline", ), ], ) @@ -668,6 +672,7 @@ def test_the_same_tool_result_twice_is_forwarded_twice_under_one_instruction(gat prompt: Final = _prompt_of(_only_generate(wire)) _assert_instructed_once(prompt, "get_weather") assert prompt.count(_RESULT) == 2, prompt + assert f"### User:\n{_RESULT}\n{_RESULT}\n\n" in prompt, prompt assert prompt.count("### User:") == 2 and prompt.count("### Assistant:") == 1, prompt @@ -707,15 +712,16 @@ def test_a_non_function_json_answer_is_returned_as_text(gateway: Gateway) -> Non assert _spend_row(completion.id) == _billed(model) -def test_int_tool_result_content_fails_in_the_response_body_and_leaves_the_deployment_serving(gateway: Gateway) -> None: +def test_int_tool_result_content_is_a_400_naming_the_field_and_leaves_the_deployment_serving(gateway: Gateway) -> None: with _ollama_server(lambda _: _generate_reply(_ANSWER)) as wire, gateway.scenario() as scenario: model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) code, text = _post( gateway, "/v1/chat/completions", {"model": model, "messages": _second_turn(22), "tools": [_WEATHER_TOOL]} ) - assert code >= 400, text + assert code == 400, text error: Final = _JSON_OBJECT.validate_json(text)["error"] - assert isinstance(error, dict) and isinstance(error["message"], str) and error["message"], text + assert isinstance(error, dict) and isinstance(error["message"], str), text + assert "content" in error["message"] and "tool message" in error["message"], text assert _generate_calls(wire) == () payload: Final = _post_chat(gateway, model, _second_turn(), tools=[_WEATHER_TOOL]) assert json.dumps(payload["choices"]).count(_ANSWER) == 1, payload @@ -794,3 +800,601 @@ def _deployments(gateway: Gateway) -> list[JsonValue]: data: Final = gateway.get("/model/info")["data"] assert isinstance(data, list) return data + + +_FOLLOW_UP: Final = "Is it windy there too?" +_TIME_CALL_ID: Final = "call_prompt_tools_2" +_TIME_RESULT: Final = "Paris: 14:05 local time" +_THINKING: Final = "The tool result already answers the question, so reply in plain text." +_ANTHROPIC_TIME_TOOL: Final[dict[str, JsonValue]] = { + "name": "get_time", + "description": "Local time for a city", + "input_schema": _PARAMETERS, +} +_RESPONSES_TIME_TOOL: Final[dict[str, JsonValue]] = { + "type": "function", + "name": "get_time", + "description": "Local time for a city", + "parameters": _PARAMETERS, +} +_AUDIO_PART: Final[dict[str, JsonValue]] = {"type": "input_audio", "input_audio": {"data": "UklGRg==", "format": "wav"}} + + +def _ndjson_line(**fields: JsonValue) -> bytes: + return json.dumps({"model": _BACKEND, "created_at": "2026-10-07T00:00:00Z", **fields}).encode() + b"\n" + + +def _ndjson_reply(lines: Sequence[bytes]) -> Reply: + final: Final = _ndjson_line( + response="", done=True, done_reason="stop", prompt_eval_count=30, eval_count=12 + ) + return Reply(content_type="application/x-ndjson", chunks=(*lines, final)) + + +def _thinking_reply(thinking: str) -> Reply: + pieces: Final = tuple(thinking[index : index + 7] for index in range(0, len(thinking), 7)) + return _ndjson_reply( + ( + _ndjson_line(response="", done=False), + *(_ndjson_line(response="", thinking=piece, done=False) for piece in pieces), + ) + ) + + +def _completion_texts(chunks: Sequence[Completion]) -> Iterator[str]: + for chunk in chunks: + yield from (choice.text or "" for choice in chunk.choices) + + +def _raw_stream_frames(gateway: Gateway, path: str, body: dict[str, JsonValue]) -> tuple[str, ...]: + with gateway.client.stream( + "POST", path, json={**body, "cache": _NO_CACHE}, headers={"Authorization": f"Bearer {gateway.key}"} + ) as response: + assert response.status_code == 200, response.read() + return tuple(line.removeprefix("data: ") for line in response.iter_lines() if line.startswith("data: ")) + + +def _raw_choice_contents(choices: JsonValue) -> Iterator[str]: + assert isinstance(choices, list) + for choice in choices: + content: Final = object_value(object_value(choice)["delta"]).get("content") + if isinstance(content, str): + yield content + + +def _raw_delta_contents(payloads: Sequence[dict[str, JsonValue]]) -> Iterator[str]: + for payload in payloads: + yield from _raw_choice_contents(payload["choices"]) + + +def _cache_rows(model: str) -> list[dict[str, JsonValue]]: + rows: Final = eventually( + lambda: read_rows( + 'SELECT request_id, cache_hit, status FROM "LiteLLM_SpendLogs" WHERE model_group=%s', (model,) + ), + lambda found: len(found) >= 2, + seconds=70, + ) + assert len(rows) == 2, rows + return rows + + +def _two_call_turn() -> list[dict[str, JsonValue]]: + return [ + {"role": "user", "content": _QUESTION}, + { + "role": "assistant", + "content": [ + {"type": "tool_use", "id": _CALL_ID, "name": "get_weather", "input": _ARGUMENTS}, + {"type": "tool_use", "id": _TIME_CALL_ID, "name": "get_time", "input": _ARGUMENTS}, + ], + }, + { + "role": "user", + "content": [ + {"type": "tool_result", "tool_use_id": _CALL_ID, "content": _RESULT}, + {"type": "tool_result", "tool_use_id": _TIME_CALL_ID, "content": _TIME_RESULT}, + ], + }, + ] + + +def _anthropic_result_with_text() -> list[dict[str, JsonValue]]: + return [ + *_anthropic_second_turn()[:2], + { + "role": "user", + "content": [ + {"type": "tool_result", "tool_use_id": _CALL_ID, "content": _RESULT}, + {"type": "text", "text": _FOLLOW_UP}, + ], + }, + ] + + +def _anthropic_int_result() -> list[dict[str, JsonValue]]: + return [ + *_anthropic_second_turn()[:2], + {"role": "user", "content": [{"type": "tool_result", "tool_use_id": _CALL_ID, "content": 22}]}, + ] + + +def _responses_two_outputs() -> list[dict[str, JsonValue]]: + return [ + {"role": "user", "content": _QUESTION}, + {"type": "function_call", "call_id": _CALL_ID, "name": "get_weather", "arguments": json.dumps(_ARGUMENTS)}, + {"type": "function_call", "call_id": _TIME_CALL_ID, "name": "get_time", "arguments": json.dumps(_ARGUMENTS)}, + {"type": "function_call_output", "call_id": _CALL_ID, "output": _RESULT}, + {"type": "function_call_output", "call_id": _TIME_CALL_ID, "output": _TIME_RESULT}, + ] + + +def _responses_int_output() -> list[dict[str, JsonValue]]: + return [*_responses_second_turn()[:2], {"type": "function_call_output", "call_id": _CALL_ID, "output": 22}] + + +def _error_message(text: str) -> str: + error: Final = _JSON_OBJECT.validate_json(text)["error"] + assert isinstance(error, dict) and isinstance(error["message"], str), text + return error["message"] + + +async def _chat_attempt(client: openai.AsyncOpenAI, model: str, messages: list[dict[str, JsonValue]]) -> ChatCompletion: + return await client.chat.completions.create( + model=model, + messages=messages, # pyright: ignore[reportArgumentType] # plain JSON messages + tools=[_WEATHER_TOOL], # pyright: ignore[reportArgumentType] # plain JSON tool + extra_body={"cache": _NO_CACHE}, + ) + + +def test_openai_sdk_completions_stream_keeps_one_id(gateway: Gateway) -> None: + with _ollama_server(lambda _: _streamed_reply(_ANSWER)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + stream: Final = _openai_client(gateway).completions.create( + model=model, prompt=_QUESTION, stream=True, extra_body={"cache": _NO_CACHE} + ) + chunks: Final = list(stream) + assert {chunk.id for chunk in chunks} == {chunks[0].id}, [chunk.id for chunk in chunks] + assert "".join(_completion_texts(chunks)) == _ANSWER + body: Final = _only_generate(wire) + assert body["stream"] is True + assert body["model"] == _BACKEND + assert body["prompt"] == _QUESTION, body + assert _spend_row(chunks[0].id) == _billed(model) + + +def test_raw_chat_stream_frames_share_one_id_and_end_with_done(gateway: Gateway) -> None: + with _ollama_server(lambda _: _streamed_reply(_ANSWER)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + frames: Final = _raw_stream_frames( + gateway, + "/v1/chat/completions", + {"model": model, "messages": _second_turn(), "tools": [_WEATHER_TOOL], "stream": True}, + ) + assert frames[-1] == "[DONE]", frames + payloads: Final = [_JSON_OBJECT.validate_json(frame) for frame in frames[:-1]] + assert {payload["id"] for payload in payloads} == {payloads[0]["id"]}, [payload["id"] for payload in payloads] + assert "".join(_raw_delta_contents(payloads)) == _ANSWER + body: Final = _only_generate(wire) + assert body["stream"] is True + _assert_tool_turn(_prompt_of(body)) + identity: Final = payloads[0]["id"] + assert isinstance(identity, str) + assert _spend_row(identity) == _billed(model) + + +async def test_cached_stream_replay_keeps_one_id_and_reaches_ollama_once(gateway: Gateway) -> None: + messages: Final[list[dict[str, JsonValue]]] = [{"role": "user", "content": f"{_QUESTION} ({uuid.uuid4().hex})"}] + with _ollama_server(lambda _: _streamed_reply(_ANSWER)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + client: Final = _async_openai_client(gateway) + first: Final = [ + chunk + async for chunk in await client.chat.completions.create( + model=model, + messages=messages, # pyright: ignore[reportArgumentType] # plain JSON messages + stream=True, + ) + ] + assert {chunk.id for chunk in first} == {first[0].id}, [chunk.id for chunk in first] + assert "".join(choice.delta.content or "" for choice in _stream_choices(first)) == _ANSWER + assert _spend_row(first[0].id) == _billed(model) + second: Final = [ + chunk + async for chunk in await client.chat.completions.create( + model=model, + messages=messages, # pyright: ignore[reportArgumentType] # plain JSON messages + stream=True, + ) + ] + assert {chunk.id for chunk in second} == {second[0].id}, [chunk.id for chunk in second] + assert "".join(choice.delta.content or "" for choice in _stream_choices(second)) == _ANSWER + assert len(_generate_calls(wire)) == 1 + rows: Final = _cache_rows(model) + hits: Final = [row for row in rows if row["cache_hit"] == "True"] + assert len(hits) == 1, rows + hit_id: Final = hits[0]["request_id"] + assert isinstance(hit_id, str) and hit_id.startswith(second[0].id), rows + assert [row["request_id"] for row in rows if row["cache_hit"] != "True"] == [first[0].id], rows + assert {row["status"] for row in rows} == {"success"}, rows + + +def test_consecutive_user_messages_are_joined_by_a_newline(gateway: Gateway) -> None: + with _ollama_server(lambda _: _generate_reply(_CALL_JSON)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + messages: Final[list[dict[str, JsonValue]]] = [*_first_turn(), {"role": "user", "content": _FOLLOW_UP}] + _post_chat(gateway, model, messages, tools=[_WEATHER_TOOL]) + prompt: Final = _prompt_of(_only_generate(wire)) + _assert_instructed_once(prompt, "get_weather") + assert f"### User:\n{_QUESTION}\n{_FOLLOW_UP}\n\n" in prompt, prompt + assert prompt.count("### User:") == 1, prompt + + +def test_a_tool_result_followed_by_user_text_is_joined_by_a_newline(gateway: Gateway) -> None: + with _ollama_server(lambda _: _generate_reply(_ANSWER)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + messages: Final[list[dict[str, JsonValue]]] = [*_second_turn(), {"role": "user", "content": _FOLLOW_UP}] + payload: Final = _post_chat(gateway, model, messages, tools=[_WEATHER_TOOL]) + assert json.dumps(payload["choices"]).count(_ANSWER) == 1, payload + prompt: Final = _prompt_of(_only_generate(wire)) + _assert_instructed_once(prompt, "get_weather") + assert ( + f"### User:\n{_QUESTION}\n\n### Assistant:\n{_CALL_JSON}\n\n### User:\n{_RESULT}\n{_FOLLOW_UP}\n\n" in prompt + ), prompt + assert prompt.count("### User:") == 2, prompt + + +def test_anthropic_sdk_two_tool_results_in_one_turn_are_separated(gateway: Gateway) -> None: + with _ollama_server(lambda _: _generate_reply(_ANSWER)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + message: Final = _anthropic_client(gateway).messages.create( + model=model, + max_tokens=64, + messages=_two_call_turn(), # pyright: ignore[reportArgumentType] # plain JSON messages + tools=[_ANTHROPIC_TOOL, _ANTHROPIC_TIME_TOOL], # pyright: ignore[reportArgumentType] # plain JSON tools + extra_body={"cache": _NO_CACHE}, + ) + assert message.stop_reason == "end_turn" + assert [(block.type, getattr(block, "text", None)) for block in message.content] == [("text", _ANSWER)] + prompt: Final = _prompt_of(_only_generate(wire)) + _assert_instructed_once(prompt, "get_weather", "get_time") + assert f"### User:\n{_RESULT}\n{_TIME_RESULT}\n\n" in prompt, prompt + assert prompt.count("### User:") == 2 and prompt.count("### Assistant:") == 1, prompt + assistant: Final = prompt.split("### Assistant:\n", 1)[1].split("### User:", 1)[0] + assert "get_weather" in assistant and "get_time" in assistant, assistant + assert _spend_row(message.id) == _billed(model) + + +async def test_async_anthropic_sdk_stream_separates_a_tool_result_from_user_text_in_one_turn( + gateway: Gateway, +) -> None: + with _ollama_server(lambda _: _streamed_reply(_ANSWER)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + async with _async_anthropic_client(gateway).messages.stream( + model=model, + max_tokens=64, + messages=_anthropic_result_with_text(), # pyright: ignore[reportArgumentType] # plain JSON messages + tools=[_ANTHROPIC_TOOL], # pyright: ignore[reportArgumentType] # plain JSON tool + extra_body={"cache": _NO_CACHE}, + ) as stream: + texts: Final = [event.text async for event in stream if event.type == "text"] + final: Final = await stream.get_final_message() + assert "".join(texts) == _ANSWER + assert final.stop_reason == "end_turn" + body: Final = _only_generate(wire) + assert body["stream"] is True + prompt: Final = _prompt_of(body) + _assert_instructed_once(prompt, "get_weather") + assert ( + f"### User:\n{_QUESTION}\n\n### Assistant:\n{_CALL_JSON}\n\n### User:\n{_RESULT}\n{_FOLLOW_UP}\n\n" in prompt + ), prompt + assert _spend_row(final.id) == _billed(model) + + +def test_anthropic_sdk_int_tool_result_content_is_dropped_before_the_prompt(gateway: Gateway) -> None: + with _ollama_server(lambda _: _generate_reply(_ANSWER)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + message: Final = _anthropic_client(gateway).messages.create( + model=model, + max_tokens=64, + messages=_anthropic_int_result(), # pyright: ignore[reportArgumentType] # plain JSON messages + tools=[_ANTHROPIC_TOOL], # pyright: ignore[reportArgumentType] # plain JSON tool + extra_body={"cache": _NO_CACHE}, + ) + assert [(block.type, getattr(block, "text", None)) for block in message.content] == [("text", _ANSWER)] + prompt: Final = _prompt_of(_only_generate(wire)) + _assert_instructed_once(prompt, "get_weather") + assert f"### User:\n{_QUESTION}\n\n### Assistant:\n{_CALL_JSON}\n\n### System:\n" in prompt, prompt + assert "22" not in prompt.split("### Assistant:\n", 1)[1], prompt + assert _spend_row(message.id) == _billed(model) + + +def test_openai_sdk_responses_two_function_outputs_are_separated(gateway: Gateway) -> None: + with _ollama_server(lambda _: _generate_reply(_ANSWER)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + response: Final = _openai_client(gateway).responses.create( + model=model, + input=_responses_two_outputs(), # pyright: ignore[reportArgumentType] # plain JSON input items + tools=[_RESPONSES_TOOL, _RESPONSES_TIME_TOOL], # pyright: ignore[reportArgumentType] # plain JSON tools + store=False, + extra_body={"cache": _NO_CACHE}, + ) + assert [item.type for item in response.output] == ["message"] + assert response.output_text == _ANSWER + prompt: Final = _prompt_of(_only_generate(wire)) + _assert_instructed_once(prompt, "get_weather", "get_time") + assert f"### User:\n{_RESULT}\n{_TIME_RESULT}\n\n" in prompt, prompt + assert prompt.count("### User:") == 2 and prompt.count("### Assistant:") == 1, prompt + assert _model_spend_rows(model, 1)[0]["status"] == "success" + + +def test_openai_sdk_responses_int_function_output_reaches_the_prompt_as_text(gateway: Gateway) -> None: + with _ollama_server(lambda _: _generate_reply(_ANSWER)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + response: Final = _openai_client(gateway).responses.create( + model=model, + input=_responses_int_output(), # pyright: ignore[reportArgumentType] # plain JSON input items + tools=[_RESPONSES_TOOL], # pyright: ignore[reportArgumentType] # plain JSON tool + store=False, + extra_body={"cache": _NO_CACHE}, + ) + assert response.output_text == _ANSWER + _assert_tool_turn(_prompt_of(_only_generate(wire)), "22") + assert _model_spend_rows(model, 1)[0]["status"] == "success" + + +@pytest.mark.parametrize( + ("messages", "message_ref", "detail"), + [ + pytest.param(_second_turn(22), "tool message at index 2", "has int content", id="tool-int-content"), + pytest.param( + _second_turn({"text": _RESULT}), "tool message at index 2", "has dict content", id="tool-dict-content" + ), + pytest.param( + _second_turn([_RESULT]), "tool message at index 2", "has a str content part", id="tool-string-part" + ), + pytest.param( + _second_turn([{"type": "text", "text": 22}]), + "tool message at index 2", + "has a int text part", + id="tool-int-text-part", + ), + pytest.param( + _second_turn([{"type": "text"}]), + "tool message at index 2", + "has a text part with no text", + id="tool-text-part-without-text", + ), + pytest.param( + [{"role": "user", "content": [{"type": "text", "text": _QUESTION}, {"type": "image_url", "image_url": 22}]}], + "user message at index 0", + "has a int image_url", + id="user-int-image-url", + ), + pytest.param( + [{"role": "user", "content": [{"type": "image_url", "image_url": {"detail": "high"}}]}], + "user message at index 0", + "has an image_url object without a url string", + id="user-image-url-object-without-url", + ), + pytest.param( + [{"role": "user", "content": [{"type": "image_url"}]}], + "user message at index 0", + "has an image_url part with no image_url", + id="user-image-url-part-without-image-url", + ), + ], +) +def test_malformed_content_is_a_400_naming_the_message_and_field( + gateway: Gateway, messages: list[dict[str, JsonValue]], message_ref: str, detail: str +) -> None: + with _ollama_server(lambda _: _generate_reply(_ANSWER)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + code, text = _post(gateway, "/v1/chat/completions", {"model": model, "messages": messages, "tools": [_WEATHER_TOOL]}) + assert code == 400, text + assert f"the {message_ref} {detail}" in _error_message(text), text + assert _generate_calls(wire) == () + + +@pytest.mark.parametrize( + "content", + [ + pytest.param(22, id="user-int-content"), + pytest.param({"text": _QUESTION}, id="user-dict-content"), + pytest.param(["just a string"], id="user-string-part"), + ], +) +def test_non_list_user_content_is_rejected_before_it_reaches_ollama(gateway: Gateway, content: JsonValue) -> None: + with _ollama_server(lambda _: _generate_reply(_ANSWER)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + code, text = _post( + gateway, + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": content}], "tools": [_WEATHER_TOOL]}, + ) + assert code >= 400, text + assert _error_message(text), text + assert _generate_calls(wire) == () + payload: Final = _post_chat(gateway, model, _second_turn(), tools=[_WEATHER_TOOL]) + assert json.dumps(payload["choices"]).count(_ANSWER) == 1, payload + _assert_tool_turn(_prompt_of(_only_generate(wire))) + + +def test_malformed_content_400_writes_a_failure_spend_row_and_leaves_the_deployment_serving(gateway: Gateway) -> None: + with _ollama_server(lambda _: _generate_reply(_ANSWER)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + code, text = _post( + gateway, "/v1/chat/completions", {"model": model, "messages": _second_turn(22), "tools": [_WEATHER_TOOL]} + ) + assert code == 400, text + assert "the tool message at index 2 has int content" in _error_message(text), text + assert [row["status"] for row in _model_spend_rows(model, 1)] == ["failure"] + assert _generate_calls(wire) == () + payload: Final = _post_chat(gateway, model, _second_turn(), tools=[_WEATHER_TOOL]) + assert json.dumps(payload["choices"]).count(_ANSWER) == 1, payload + _assert_tool_turn(_prompt_of(_only_generate(wire))) + assert sorted(str(row["status"]) for row in _model_spend_rows(model, 2)) == ["failure", "success"] + + +def test_malformed_content_on_an_unauthenticated_request_is_a_401_before_translation(gateway: Gateway) -> None: + with _ollama_server(lambda _: _generate_reply(_ANSWER)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + code, text = _post( + gateway, + "/v1/chat/completions", + {"model": model, "messages": _second_turn(22), "tools": [_WEATHER_TOOL]}, + key=f"sk-not-a-key-{uuid.uuid4().hex}", + ) + assert code == 401, text + assert _generate_calls(wire) == () + + +def test_a_content_part_of_an_unknown_type_is_dropped_from_the_prompt(gateway: Gateway) -> None: + with _ollama_server(lambda _: _generate_reply(_CALL_JSON)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + messages: Final[list[dict[str, JsonValue]]] = [ + {"role": "user", "content": [{"type": "text", "text": _QUESTION}, _AUDIO_PART]} + ] + _post_chat(gateway, model, messages, tools=[_WEATHER_TOOL]) + prompt: Final = _prompt_of(_only_generate(wire)) + _assert_instructed_once(prompt, "get_weather") + assert f"### User:\n{_QUESTION}\n\n" in prompt, prompt + assert "UklGRg==" not in prompt and "input_audio" not in prompt, prompt + + +@pytest.mark.parametrize( + "messages", + [ + pytest.param( + [ + { + "role": "user", + "content": [ + {"type": "text", "text": ""}, + {"type": "text", "text": _QUESTION}, + {"type": "text", "text": ""}, + ], + } + ], + id="empty-text-parts-dropped", + ), + pytest.param([{"role": "user", "content": None}, *_first_turn()], id="null-content-skipped"), + pytest.param([{"role": "user", "content": ""}, *_first_turn()], id="empty-string-skipped"), + ], +) +def test_empty_and_null_user_content_are_dropped_from_the_user_section( + gateway: Gateway, messages: list[dict[str, JsonValue]] +) -> None: + with _ollama_server(lambda _: _generate_reply(_CALL_JSON)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + _post_chat(gateway, model, messages, tools=[_WEATHER_TOOL]) + prompt: Final = _prompt_of(_only_generate(wire)) + _assert_instructed_once(prompt, "get_weather") + assert f"### User:\n{_QUESTION}\n\n" in prompt, prompt + assert prompt.count("### User:") == 1, prompt + + +async def test_async_openai_sdk_responses_stream_of_thinking_only_emits_one_reasoning_item(gateway: Gateway) -> None: + with _ollama_server(lambda _: _thinking_reply(_THINKING)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + stream: Final = await _async_openai_client(gateway).responses.create( + model=model, + input=_responses_second_turn(), # pyright: ignore[reportArgumentType] # plain JSON input items + tools=[_RESPONSES_TOOL], # pyright: ignore[reportArgumentType] # plain JSON tool + store=False, + stream=True, + extra_body={"cache": _NO_CACHE}, + ) + events: Final = [event async for event in stream] + added: Final = [ + (event.item.type, event.output_index) for event in events if event.type == "response.output_item.added" + ] + assert added == [("reasoning", 0)], [event.type for event in events] + completed: Final = [event for event in events if event.type == "response.completed"] + assert len(completed) == 1, [event.type for event in events] + assert [item.type for item in completed[0].response.output] == ["reasoning"], completed[0].response.output + assert completed[0].response.output_text == "" + body: Final = _only_generate(wire) + assert body["stream"] is True + _assert_tool_turn(_prompt_of(body)) + assert _model_spend_rows(model, 1)[0]["status"] == "success" + + +async def test_async_openai_sdk_responses_stream_with_an_empty_answer_keeps_one_message_item(gateway: Gateway) -> None: + with _ollama_server(lambda _: _ndjson_reply(())) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + stream: Final = await _async_openai_client(gateway).responses.create( + model=model, + input=_responses_second_turn(), # pyright: ignore[reportArgumentType] # plain JSON input items + tools=[_RESPONSES_TOOL], # pyright: ignore[reportArgumentType] # plain JSON tool + store=False, + stream=True, + extra_body={"cache": _NO_CACHE}, + ) + events: Final = [event async for event in stream] + added: Final = [ + (event.item.type, event.output_index) for event in events if event.type == "response.output_item.added" + ] + assert added == [("message", 0)], [event.type for event in events] + completed: Final = [event for event in events if event.type == "response.completed"] + assert len(completed) == 1, [event.type for event in events] + assert [item.type for item in completed[0].response.output] == ["message"], completed[0].response.output + assert completed[0].response.output_text == "" + body: Final = _only_generate(wire) + assert body["stream"] is True + _assert_tool_turn(_prompt_of(body)) + assert _model_spend_rows(model, 1)[0]["status"] == "success" + + +def test_openai_sdk_responses_stream_context_manager_gets_the_function_call_without_an_empty_message( + gateway: Gateway, +) -> None: + with _ollama_server(lambda _: _streamed_reply(_CALL_JSON)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + with _openai_client(gateway).responses.stream( + model=model, + input=_QUESTION, + tools=[_RESPONSES_TOOL], # pyright: ignore[reportArgumentType] # plain JSON tool + store=False, + extra_body={"cache": _NO_CACHE}, + ) as stream: + events: Final = list(stream) + final: Final = stream.get_final_response() + added: Final = [ + (event.item.type, event.output_index) for event in events if event.type == "response.output_item.added" + ] + assert added == [("function_call", 0)], [event.type for event in events] + assert [item.type for item in final.output] == ["function_call"], final.output + call: Final = final.output[0] + assert call.type == "function_call" and call.name == "get_weather" + assert json.loads(call.arguments) == _ARGUMENTS + assert final.output_text == "" + body: Final = _only_generate(wire) + assert body["stream"] is True + _assert_instructed_once(_prompt_of(body), "get_weather") + assert _model_spend_rows(model, 1)[0]["status"] == "success" + + +async def test_concurrent_valid_and_malformed_turns_are_each_answered_in_their_own_shape(gateway: Gateway) -> None: + with _ollama_server(lambda _: _generate_reply(_ANSWER)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + client: Final = _async_openai_client(gateway) + results: Final = await asyncio.gather( + *(_chat_attempt(client, model, _second_turn(22 if index % 3 == 0 else _RESULT)) for index in range(12)), + return_exceptions=True, + ) + answers: Final = [result for result in results if isinstance(result, ChatCompletion)] + rejections: Final = [result for result in results if isinstance(result, openai.BadRequestError)] + assert (len(answers), len(rejections)) == (8, 4), results + assert {answer.choices[0].message.content for answer in answers} == {_ANSWER} + for rejection in rejections: + assert "the tool message at index 2 has int content" in str(rejection), rejection + prompts: Final = [_prompt_of(_JSON_OBJECT.validate_json(request.body)) for request in _generate_calls(wire)] + assert len(prompts) == 8, prompts + for prompt in prompts: + _assert_tool_turn(prompt) + statuses: Final = sorted(str(row["status"]) for row in _model_spend_rows(model, 12)) + assert statuses == ["failure"] * 4 + ["success"] * 8, statuses + for answer in answers: + assert _spend_row(answer.id) == _billed(model) diff --git a/tests/integration/providers/test_responses_bridge_prompt_cache_breakpoint_wire.py b/tests/integration/providers/test_responses_bridge_prompt_cache_breakpoint_wire.py new file mode 100644 index 00000000000..75bd9da7795 --- /dev/null +++ b/tests/integration/providers/test_responses_bridge_prompt_cache_breakpoint_wire.py @@ -0,0 +1,510 @@ +import asyncio +import uuid +from collections.abc import Iterator, Mapping, Sequence +from dataclasses import dataclass +from typing import Final, Literal, TypeAlias + +import anthropic +import httpx +import openai +import pytest +from integration._support import prompt_cache_breakpoint as pcb +from integration._support import responses_vendor as rv +from integration._support.client import Gateway, eventually, gateway_from_environment, string_value +from integration._support.wire import Request, Wire, wire_server +from pydantic import JsonValue + +pytestmark: Final = pytest.mark.timeout(120) + +Mode: TypeAlias = Literal["on", "off"] + +_UNKNOWN_KEY: Final[Mapping[str, JsonValue]] = {"mode": "explicit", "note": "kept"} +_MALFORMED: Final[tuple[tuple[str, JsonValue], ...]] = ( + ("string", "yes"), + ("int", 1), + ("list", ["explicit"]), + ("empty-string", ""), + ("5kb-string", "x" * 5000), + ("empty-object", {}), + ("bogus-mode", {"mode": "bogus"}), + ("bad-ttl", {"mode": "explicit", "ttl": "1h"}), +) +_MALFORMED_IDS: Final = tuple(name for name, _ in _MALFORMED) +_MALFORMED_VALUES: Final = tuple(value for _, value in _MALFORMED) +_CASES: Final[tuple[tuple[str, Mode, JsonValue, JsonValue], ...]] = ( + ("valid-on", "on", pcb.EXPLICIT, pcb.EXPLICIT), + ("valid-off", "off", pcb.EXPLICIT, pcb.EXPLICIT), + ("malformed-on", "on", "yes", None), + ("malformed-off", "off", "yes", "yes"), +) +_CASE_IDS: Final = tuple(case[0] for case in _CASES) +_CASE_VALUES: Final = tuple(case[1:] for case in _CASES) +_ADAPTER_CASES: Final[tuple[tuple[str, Mode, JsonValue, JsonValue], ...]] = ( + ("valid-on", "on", pcb.EXPLICIT, pcb.EXPLICIT), + ("valid-off", "off", pcb.EXPLICIT, pcb.EXPLICIT), + ("malformed-on", "on", "yes", "yes"), + ("malformed-off", "off", "yes", "yes"), +) +_ADAPTER_CASE_IDS: Final = tuple(case[0] for case in _ADAPTER_CASES) +_ADAPTER_CASE_VALUES: Final = tuple(case[1:] for case in _ADAPTER_CASES) + + +@dataclass(frozen=True, slots=True) +class _Bridge: + gateway: Gateway + wire: Wire + on: str + off: str + injecting_on: str + injecting_off: str + spend: pcb.SpendLogs + + def model(self, mode: Mode) -> str: + return self.on if mode == "on" else self.off + + def injecting(self, mode: Mode) -> str: + return self.injecting_on if mode == "on" else self.injecting_off + + @property + def api_base(self) -> str: + return f"{self.wire.url}/v1" + + +@pytest.fixture(scope="module") +def bridge() -> Iterator[_Bridge]: + with ( + wire_server(pcb.respond) as wire, + gateway_from_environment() as gateway, + gateway.scenario() as scenario, + pcb.spend_logs() as spend, + ): + api_base: Final = f"{wire.url}/v1" + yield _Bridge( + gateway, + wire, + scenario.model(model=pcb.MODEL, api_base=api_base, drop_params=True), + scenario.model(model=pcb.MODEL, api_base=api_base), + scenario.model(model=pcb.MODEL, api_base=api_base, drop_params=True, **pcb.INJECTION), + scenario.model(model=pcb.MODEL, api_base=api_base, **pcb.INJECTION), + spend, + ) + + +def _v1(gateway: Gateway) -> str: + return str(gateway.client.base_url).rstrip("/") + "/v1" + + +def _user(marker: str, breakpoint: JsonValue) -> dict[str, JsonValue]: + return {"role": "user", "content": [pcb.marked(pcb.text(pcb.prompt(marker)), breakpoint)]} + + +def _chat( + bridge: _Bridge, model: str, messages: Sequence[JsonValue], *, stream: bool = False, key: str | None = None +) -> httpx.Response: + return bridge.gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": list(messages), "stream": stream, **pcb.NO_CACHE}, + key=key, + ) + + +def _completion(response: httpx.Response, marker: str) -> str: + assert response.status_code == 200, response.text + body: Final = rv.JSON_OBJECT.validate_json(response.text) + assert pcb.answers(string_value(body["id"]), marker), body + (choice,) = rv.ITEMS.validate_python(body["choices"]) + assert rv.JSON_OBJECT.validate_python(choice["message"])["content"] == rv.answer(marker), body + return response.headers["x-litellm-call-id"] + + +def _wire_body(request: Request, *, stream: bool = False) -> dict[str, JsonValue]: + body: Final = pcb.body_of(request) + assert body["model"] == "gpt-6.1-sol", body + assert (body.get("stream") is True) is stream, body + return body + + +def _user_block_on_wire(bridge: _Bridge, marker: str, *, stream: bool = False) -> dict[str, JsonValue]: + request: Final = pcb.posted(bridge.wire, marker) + block: Final = pcb.single_block(pcb.input_items(request), "user") + _wire_body(request, stream=stream) + assert block["type"] == "input_text" and block["text"] == pcb.prompt(marker), block + return block + + +def test_openai_sdk_sends_a_valid_marker_through_the_bridge(bridge: _Bridge) -> None: + marker: Final = uuid.uuid4().hex + with openai.OpenAI(api_key=bridge.gateway.key, base_url=_v1(bridge.gateway), max_retries=0) as client: + raw: Final = client.chat.completions.with_raw_response.create( + model=bridge.on, messages=[_user(marker, pcb.EXPLICIT)], extra_body=dict(pcb.NO_CACHE) + ) + completion: Final = raw.parse() + assert pcb.answers(completion.id, marker), completion + assert completion.choices[0].message.content == rv.answer(marker), completion + pcb.assert_marker(_user_block_on_wire(bridge, marker), pcb.EXPLICIT) + bridge.spend.landed(bridge.on, raw.headers["x-litellm-call-id"], marker) + + +def test_openai_sdk_stream_carries_the_system_list_marker(bridge: _Bridge) -> None: + marker: Final = uuid.uuid4().hex + system: Final[dict[str, JsonValue]] = {"role": "system", "content": [pcb.marked(pcb.text("sys"), pcb.EXPLICIT)]} + with openai.OpenAI(api_key=bridge.gateway.key, base_url=_v1(bridge.gateway), max_retries=0) as client: + raw: Final = client.chat.completions.with_raw_response.create( + model=bridge.on, + messages=[system, {"role": "user", "content": pcb.prompt(marker)}], + stream=True, + extra_body=dict(pcb.NO_CACHE), + ) + chunks: Final = tuple(raw.parse()) + assert chunks and pcb.answers(chunks[0].id, marker), chunks + streamed: Final = "".join(chunk.choices[0].delta.content or "" for chunk in chunks if chunk.choices) + assert streamed == rv.answer(marker), chunks + request: Final = pcb.posted(bridge.wire, marker) + _wire_body(request, stream=True) + (system_block,) = pcb.content_of(pcb.input_items(request), "system") + assert system_block["type"] == "input_text" and system_block["text"] == "sys", system_block + pcb.assert_marker(system_block, pcb.EXPLICIT) + bridge.spend.landed(bridge.on, raw.headers["x-litellm-call-id"], marker) + + +async def test_async_openai_sdk_keeps_the_ttl_without_drop_params(bridge: _Bridge) -> None: + marker: Final = uuid.uuid4().hex + async with openai.AsyncOpenAI(api_key=bridge.gateway.key, base_url=_v1(bridge.gateway), max_retries=0) as client: + raw: Final = await client.chat.completions.with_raw_response.create( + model=bridge.off, messages=[_user(marker, pcb.EXPLICIT_30M)], extra_body=dict(pcb.NO_CACHE) + ) + completion: Final = raw.parse() + assert pcb.answers(completion.id, marker), completion + assert completion.choices[0].message.content == rv.answer(marker), completion + pcb.assert_marker(_user_block_on_wire(bridge, marker), pcb.EXPLICIT_30M) + bridge.spend.landed(bridge.off, raw.headers["x-litellm-call-id"], marker) + + +@pytest.mark.parametrize("mode", ("on", "off")) +@pytest.mark.parametrize("breakpoint", (pcb.EXPLICIT, pcb.EXPLICIT_30M), ids=("explicit", "ttl")) +def test_valid_marker_shapes_reach_the_wire_unchanged(bridge: _Bridge, mode: Mode, breakpoint: JsonValue) -> None: + marker: Final = uuid.uuid4().hex + call_id: Final = _completion(_chat(bridge, bridge.model(mode), [_user(marker, breakpoint)]), marker) + pcb.assert_marker(_user_block_on_wire(bridge, marker), breakpoint) + bridge.spend.landed(bridge.model(mode), call_id, marker) + + +@pytest.mark.parametrize( + ("mode", "expected"), (("on", pcb.EXPLICIT), ("off", _UNKNOWN_KEY)), ids=("normalized-on", "verbatim-off") +) +def test_marker_with_an_unknown_key(bridge: _Bridge, mode: Mode, expected: JsonValue) -> None: + marker: Final = uuid.uuid4().hex + call_id: Final = _completion(_chat(bridge, bridge.model(mode), [_user(marker, _UNKNOWN_KEY)]), marker) + pcb.assert_marker(_user_block_on_wire(bridge, marker), expected) + bridge.spend.landed(bridge.model(mode), call_id, marker) + + +def _second_block_on_wire(bridge: _Bridge, marker: str) -> dict[str, JsonValue]: + request: Final = pcb.posted(bridge.wire, marker) + _wire_body(request) + first, second = pcb.content_of(pcb.input_items(request), "user") + assert first == {"type": "input_text", "text": pcb.prompt(marker)}, first + return second + + +@pytest.mark.parametrize("mode", ("on", "off")) +@pytest.mark.parametrize("kind", pcb.KINDS) +def test_valid_marker_is_carried_on_every_block_kind(bridge: _Bridge, kind: pcb.Kind, mode: Mode) -> None: + marker: Final = uuid.uuid4().hex + content: Final[list[JsonValue]] = [ + pcb.text(pcb.prompt(marker)), + pcb.marked(pcb.block(kind, "second"), pcb.EXPLICIT), + ] + call_id: Final = _completion(_chat(bridge, bridge.model(mode), [{"role": "user", "content": content}]), marker) + second: Final = _second_block_on_wire(bridge, marker) + assert second["type"] == pcb.WIRE_TYPE[kind], second + pcb.assert_marker(second, pcb.EXPLICIT) + bridge.spend.landed(bridge.model(mode), call_id, marker) + + +@pytest.mark.parametrize("kind", pcb.KINDS) +def test_malformed_marker_is_dropped_on_every_block_kind(bridge: _Bridge, kind: pcb.Kind) -> None: + marker: Final = uuid.uuid4().hex + content: Final[list[JsonValue]] = [pcb.text(pcb.prompt(marker)), pcb.marked(pcb.block(kind, "second"), "yes")] + call_id: Final = _completion(_chat(bridge, bridge.on, [{"role": "user", "content": content}]), marker) + second: Final = _second_block_on_wire(bridge, marker) + assert second["type"] == pcb.WIRE_TYPE[kind], second + pcb.assert_marker(second, None) + bridge.spend.landed(bridge.on, call_id, marker) + + +@pytest.mark.parametrize(("mode", "breakpoint", "expected"), _CASE_VALUES, ids=_CASE_IDS) +def test_tool_output_marker(bridge: _Bridge, mode: Mode, breakpoint: JsonValue, expected: JsonValue) -> None: + marker: Final = uuid.uuid4().hex + messages: Final[list[JsonValue]] = [ + {"role": "user", "content": pcb.prompt(marker)}, + { + "role": "assistant", + "content": None, + "tool_calls": [{"id": "call_1", "type": "function", "function": {"name": "lookup", "arguments": "{}"}}], + }, + {"role": "tool", "tool_call_id": "call_1", "content": [pcb.marked(pcb.text("found it"), breakpoint)]}, + ] + call_id: Final = _completion(_chat(bridge, bridge.model(mode), messages), marker) + request: Final = pcb.posted(bridge.wire, marker) + _wire_body(request) + (output,) = pcb.function_output(pcb.input_items(request), "call_1") + assert output["type"] == "input_text" and output["text"] == "found it", output + pcb.assert_marker(output, expected) + bridge.spend.landed(bridge.model(mode), call_id, marker) + + +@pytest.mark.parametrize(("mode", "breakpoint", "expected"), _CASE_VALUES, ids=_CASE_IDS) +def test_assistant_list_marker(bridge: _Bridge, mode: Mode, breakpoint: JsonValue, expected: JsonValue) -> None: + marker: Final = uuid.uuid4().hex + messages: Final[list[JsonValue]] = [ + {"role": "user", "content": pcb.prompt(marker)}, + {"role": "assistant", "content": [pcb.marked(pcb.text("earlier answer"), breakpoint)]}, + {"role": "user", "content": "and again"}, + ] + call_id: Final = _completion(_chat(bridge, bridge.model(mode), messages), marker) + request: Final = pcb.posted(bridge.wire, marker) + _wire_body(request) + (earlier,) = pcb.content_of(pcb.input_items(request), "assistant") + assert earlier["type"] == "output_text" and earlier["text"] == "earlier answer", earlier + pcb.assert_marker(earlier, expected) + bridge.spend.landed(bridge.model(mode), call_id, marker) + + +def test_injected_system_marker_survives_a_trailing_audio_block(bridge: _Bridge) -> None: + marker: Final = uuid.uuid4().hex + system: Final[dict[str, JsonValue]] = {"role": "system", "content": [pcb.text("sys"), pcb.block("input_audio", "")]} + messages: Final[list[JsonValue]] = [system, {"role": "user", "content": pcb.prompt(marker)}] + call_id: Final = _completion(_chat(bridge, bridge.injecting_on, messages), marker) + request: Final = pcb.posted(bridge.wire, marker) + body: Final = _wire_body(request) + assert body["prompt_cache_options"] == {"mode": "explicit"}, body + first, audio = pcb.content_of(pcb.input_items(request), "system") + assert first == {"type": "input_text", "text": "sys"}, first + assert audio["type"] == "input_text", audio + assert string_value(audio["text"]).startswith("{'type': 'input_audio'"), audio + pcb.assert_marker(audio, pcb.EXPLICIT) + bridge.spend.landed(bridge.injecting_on, call_id, marker) + + +def test_injected_marker_on_a_string_system_message(bridge: _Bridge) -> None: + marker: Final = uuid.uuid4().hex + messages: Final[list[JsonValue]] = [ + {"role": "system", "content": "Answer briefly"}, + {"role": "user", "content": pcb.prompt(marker)}, + ] + call_id: Final = _completion(_chat(bridge, bridge.injecting_off, messages), marker) + request: Final = pcb.posted(bridge.wire, marker) + body: Final = _wire_body(request) + assert body["prompt_cache_options"] == {"mode": "explicit"}, body + system_block: Final = pcb.single_block(pcb.input_items(request), "system") + assert system_block == {"type": "input_text", "text": "Answer briefly", "prompt_cache_breakpoint": pcb.EXPLICIT} + bridge.spend.landed(bridge.injecting_off, call_id, marker) + + +def _anthropic(bridge: _Bridge) -> anthropic.Anthropic: + return anthropic.Anthropic(base_url=str(bridge.gateway.client.base_url), api_key=bridge.gateway.key, max_retries=0) + + +def _anthropic_text(message: anthropic.types.Message, marker: str) -> None: + (content,) = message.content + assert content.type == "text" and content.text == rv.answer(marker), message + + +@pytest.mark.parametrize(("mode", "breakpoint", "expected"), _ADAPTER_CASE_VALUES, ids=_ADAPTER_CASE_IDS) +def test_anthropic_sdk_marker_on_user_text( + bridge: _Bridge, mode: Mode, breakpoint: JsonValue, expected: JsonValue +) -> None: + marker: Final = uuid.uuid4().hex + with _anthropic(bridge) as client: + raw: Final = client.messages.with_raw_response.create( + model=bridge.model(mode), + max_tokens=64, + messages=[{"role": "user", "content": [pcb.marked(pcb.text(pcb.prompt(marker)), breakpoint)]}], + extra_body=dict(pcb.NO_CACHE), + ) + _anthropic_text(raw.parse(), marker) + pcb.assert_marker(_user_block_on_wire(bridge, marker), expected) + bridge.spend.landed(bridge.model(mode), raw.headers["x-litellm-call-id"], None) + + +def test_anthropic_sdk_system_string_gets_the_injected_marker(bridge: _Bridge) -> None: + marker: Final = uuid.uuid4().hex + with _anthropic(bridge) as client: + raw: Final = client.messages.with_raw_response.create( + model=bridge.injecting_off, + max_tokens=64, + system="Answer briefly", + messages=[{"role": "user", "content": pcb.prompt(marker)}], + extra_body=dict(pcb.NO_CACHE), + ) + _anthropic_text(raw.parse(), marker) + request: Final = pcb.posted(bridge.wire, marker) + body: Final = _wire_body(request) + assert body["prompt_cache_options"] == {"mode": "explicit"}, body + instruction: Final = pcb.instruction_block(pcb.input_items(request)) + assert instruction == {"type": "input_text", "text": "Answer briefly", "prompt_cache_breakpoint": pcb.EXPLICIT} + bridge.spend.landed(bridge.injecting_off, raw.headers["x-litellm-call-id"], None) + + +def test_native_responses_request_never_enters_the_bridge(bridge: _Bridge) -> None: + marker: Final = uuid.uuid4().hex + block: Final[dict[str, JsonValue]] = {"type": "input_text", "text": pcb.prompt(marker)} + response: Final = bridge.gateway.request( + "POST", + "/v1/responses", + { + "model": bridge.on, + "input": [{"type": "message", "role": "user", "content": [pcb.marked(block, _UNKNOWN_KEY)]}], + **pcb.NO_CACHE, + }, + ) + assert response.status_code == 200, response.text + body: Final = rv.JSON_OBJECT.validate_json(response.text) + assert pcb.answers(string_value(body["id"]), marker), body + assert rv.answer(marker) in response.text, response.text + request: Final = pcb.posted(bridge.wire, marker) + _wire_body(request) + on_wire: Final = pcb.single_block(pcb.input_items(request), "user") + assert on_wire == pcb.marked(block, _UNKNOWN_KEY), on_wire + bridge.spend.landed(bridge.on, response.headers["x-litellm-call-id"], marker) + + +@pytest.mark.parametrize("breakpoint", _MALFORMED_VALUES, ids=_MALFORMED_IDS) +def test_malformed_marker_is_dropped_under_drop_params(bridge: _Bridge, breakpoint: JsonValue) -> None: + marker: Final = uuid.uuid4().hex + call_id: Final = _completion(_chat(bridge, bridge.on, [_user(marker, breakpoint)]), marker) + pcb.assert_marker(_user_block_on_wire(bridge, marker), None) + bridge.spend.landed(bridge.on, call_id, marker) + + +@pytest.mark.parametrize("breakpoint", _MALFORMED_VALUES, ids=_MALFORMED_IDS) +def test_malformed_marker_passes_verbatim_without_drop_params(bridge: _Bridge, breakpoint: JsonValue) -> None: + marker: Final = uuid.uuid4().hex + call_id: Final = _completion(_chat(bridge, bridge.off, [_user(marker, breakpoint)]), marker) + pcb.assert_marker(_user_block_on_wire(bridge, marker), breakpoint) + bridge.spend.landed(bridge.off, call_id, marker) + + +def test_two_marked_blocks_are_both_carried(bridge: _Bridge) -> None: + marker: Final = uuid.uuid4().hex + content: Final[list[JsonValue]] = [ + pcb.marked(pcb.text(pcb.prompt(marker)), pcb.EXPLICIT), + pcb.marked(pcb.text("and more"), pcb.EXPLICIT_30M), + ] + call_id: Final = _completion(_chat(bridge, bridge.on, [{"role": "user", "content": content}]), marker) + request: Final = pcb.posted(bridge.wire, marker) + _wire_body(request) + first, second = pcb.content_of(pcb.input_items(request), "user") + assert first == {"type": "input_text", "text": pcb.prompt(marker), "prompt_cache_breakpoint": pcb.EXPLICIT} + assert second == {"type": "input_text", "text": "and more", "prompt_cache_breakpoint": pcb.EXPLICIT_30M} + bridge.spend.landed(bridge.on, call_id, marker) + + +def test_wrong_key_is_refused_before_the_wire(bridge: _Bridge) -> None: + marker: Final = uuid.uuid4().hex + response: Final = _chat(bridge, bridge.on, [_user(marker, pcb.EXPLICIT)], key="sk-wrong") + assert response.status_code == 401, response.text + assert pcb.with_marker(pcb.drained_posts(bridge.wire), marker) == (), marker + + +@pytest.mark.parametrize( + ("mode", "status"), (("on", 400), ("off", 400), ("on", 401)), ids=("400-on", "400-off", "401-on") +) +def test_upstream_error_reaches_the_caller_once(bridge: _Bridge, mode: Mode, status: int) -> None: + marker: Final = uuid.uuid4().hex + failing: Final[dict[str, JsonValue]] = { + "role": "user", + "content": [pcb.marked(pcb.text(f"{pcb.prompt(marker)} fail-{status}"), pcb.EXPLICIT)], + } + response: Final = _chat(bridge, bridge.model(mode), [failing]) + assert response.status_code == status, response.text + assert f"scripted {status} marker-{marker}" in response.text, response.text + (request,) = pcb.with_marker(pcb.drained_posts(bridge.wire), marker) + pcb.assert_marker(pcb.single_block(pcb.input_items(request), "user"), pcb.EXPLICIT) + bridge.spend.landed(bridge.model(mode), response.headers["x-litellm-call-id"], None, status="failure") + follow_up: Final = uuid.uuid4().hex + call_id: Final = _completion(_chat(bridge, bridge.model(mode), [_user(follow_up, pcb.EXPLICIT)]), follow_up) + pcb.assert_marker(_user_block_on_wire(bridge, follow_up), pcb.EXPLICIT) + bridge.spend.landed(bridge.model(mode), call_id, follow_up) + + +def test_null_drop_params_on_the_deployment_means_off(bridge: _Bridge) -> None: + marker: Final = uuid.uuid4().hex + with bridge.gateway.scenario() as scenario: + model: Final = scenario.model(model=pcb.MODEL, api_base=bridge.api_base, drop_params=None) + call_id: Final = _completion(_chat(bridge, model, [_user(marker, "yes")]), marker) + pcb.assert_marker(_user_block_on_wire(bridge, marker), "yes") + bridge.spend.landed(model, call_id, marker) + + +@pytest.mark.parametrize("mode", ("on", "off")) +@pytest.mark.parametrize("shape", ("null", "missing")) +def test_null_or_missing_marker_sends_a_plain_block(bridge: _Bridge, shape: str, mode: Mode) -> None: + marker: Final = uuid.uuid4().hex + block: Final = pcb.marked(pcb.text(pcb.prompt(marker)), None) if shape == "null" else pcb.text(pcb.prompt(marker)) + call_id: Final = _completion(_chat(bridge, bridge.model(mode), [{"role": "user", "content": [block]}]), marker) + assert _user_block_on_wire(bridge, marker) == {"type": "input_text", "text": pcb.prompt(marker)} + bridge.spend.landed(bridge.model(mode), call_id, marker) + + +async def _send_marked(client: httpx.AsyncClient, key: str, model: str, marker: str) -> httpx.Response: + return await client.post( + "/v1/chat/completions", + json={"model": model, "messages": [_user(marker, "yes")], **pcb.NO_CACHE}, + headers={"Authorization": f"Bearer {key}"}, + ) + + +def _probe_marker(bridge: _Bridge, model: str) -> JsonValue: + marker: Final = uuid.uuid4().hex + _completion(_chat(bridge, model, [_user(marker, "yes")]), marker) + return _user_block_on_wire(bridge, marker).get("prompt_cache_breakpoint") + + +@pytest.mark.timeout(180) +async def test_flipping_drop_params_mid_burst_keeps_every_marked_request_answered(bridge: _Bridge) -> None: + gateway: Final = bridge.gateway + with gateway.scenario() as scenario: + model: Final = scenario.model(model=pcb.MODEL, api_base=bridge.api_base, drop_params=True) + identity: Final = pcb.model_id(gateway.get("/model/info")["data"], model) + markers: Final = tuple(uuid.uuid4().hex for _ in range(20)) + async with httpx.AsyncClient(base_url=str(gateway.client.base_url), timeout=60, trust_env=False) as client: + burst: Final = asyncio.gather(*(_send_marked(client, gateway.key, model, marker) for marker in markers)) + updated: Final = await asyncio.to_thread( + gateway.request, + "POST", + "/model/update", + { + "model_name": model, + "litellm_params": {"model": pcb.MODEL, "drop_params": False}, + "model_info": {"id": identity}, + }, + ) + responses: Final = await burst + assert updated.status_code == 200, updated.text + for marker, response in zip(markers, responses, strict=True): + _completion(response, marker) + posts: Final = pcb.drained_posts(bridge.wire) + for marker in markers: + (request,) = pcb.with_marker(posts, marker) + seen: Final = pcb.single_block(pcb.input_items(request), "user").get("prompt_cache_breakpoint") + assert seen in (None, "yes"), request.body + flipped: Final = eventually(lambda: _probe_marker(bridge, model), lambda seen: seen == "yes", seconds=70) + assert flipped == "yes" + for marker, response in zip(markers, responses, strict=True): + bridge.spend.landed(model, response.headers["x-litellm-call-id"], marker) + + +def test_three_identical_marked_requests_are_each_sent_and_logged(bridge: _Bridge) -> None: + marker: Final = uuid.uuid4().hex + responses: Final = tuple(_chat(bridge, bridge.on, [_user(marker, pcb.EXPLICIT)]) for _ in range(3)) + call_ids: Final = tuple(_completion(response, marker) for response in responses) + assert len(set(call_ids)) == 3, call_ids + posts: Final = pcb.with_marker(pcb.drained_posts(bridge.wire), marker) + assert len(posts) == 3, [request.body for request in posts] + for request in posts: + pcb.assert_marker(pcb.single_block(pcb.input_items(request), "user"), pcb.EXPLICIT) + for call_id in call_ids: + bridge.spend.landed(bridge.on, call_id, marker) diff --git a/tests/integration/providers/test_responses_bridge_prompt_cache_breakpoint_wire_chaos.py b/tests/integration/providers/test_responses_bridge_prompt_cache_breakpoint_wire_chaos.py new file mode 100644 index 00000000000..7742d4b9821 --- /dev/null +++ b/tests/integration/providers/test_responses_bridge_prompt_cache_breakpoint_wire_chaos.py @@ -0,0 +1,397 @@ +import asyncio +import dataclasses +import signal +import threading +import uuid +from collections import Counter +from collections.abc import Callable, Iterator, Mapping, Sequence +from dataclasses import dataclass +from pathlib import Path +from queue import SimpleQueue +from types import MappingProxyType +from typing import Final, Literal, TypeAlias +from urllib.parse import urlsplit + +import httpx +import psutil +import pytest +import yaml +from integration._support import prompt_cache_breakpoint as pcb +from integration._support import responses_vendor as rv +from integration._support.client import Gateway, eventually, gateway_from_environment, string_value +from integration._support.process import OwnedProxy, graceful_stop_seconds, owned_proxy_process +from integration._support.wire import Reply, Request, Wire, wire_server +from pydantic import JsonValue + +pytestmark: Final = pytest.mark.timeout(2 * graceful_stop_seconds() + 120) + +_GLOBAL_UNSET: Final = "bridge-breakpoint-global-unset" +_GLOBAL_FALSE: Final = "bridge-breakpoint-global-false" +_ENDPOINTS: Final = ("chat", "messages", "responses") + +Endpoint: TypeAlias = Literal["chat", "messages", "responses"] +_RecordProperty: TypeAlias = Callable[[str, object], None] + + +@dataclass(frozen=True, slots=True) +class _Call: + endpoint: Endpoint + stream: bool + marker: str + + +@dataclass(frozen=True, slots=True) +class _Served: + call: _Call + status: int + text: str + call_id: str + + +@dataclass(frozen=True, slots=True) +class _GlobalRig: + wire: Wire + proxy: OwnedProxy + + @property + def gateway(self) -> Gateway: + return self.proxy.gateway + + +def _global_config(directory: Path, api_base: str) -> Path: + stock: Final = rv.JSON_OBJECT.validate_python( + yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + ) + deployment: Final[Mapping[str, JsonValue]] = { + "model": pcb.MODEL, + "api_base": api_base, + "api_key": "integration-provider-key", + } + config: Final[Mapping[str, JsonValue]] = { + **stock, + "model_list": [ + {"model_name": _GLOBAL_UNSET, "litellm_params": dict(deployment)}, + {"model_name": _GLOBAL_FALSE, "litellm_params": {**deployment, "drop_params": False}}, + ], + "litellm_settings": {**rv.JSON_OBJECT.validate_python(stock["litellm_settings"]), "drop_params": True}, + "router_settings": {**rv.JSON_OBJECT.validate_python(stock.get("router_settings") or {}), "num_retries": 0}, + } + path: Final = directory / "bridge-breakpoint-global.yaml" + path.write_text(yaml.safe_dump(config)) + return path + + +@pytest.fixture(scope="module") +def global_rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[_GlobalRig]: + directory: Final = tmp_path_factory.mktemp("bridge-breakpoint-global") + with wire_server(pcb.respond) as wire, gateway_from_environment() as gateway: + config: Final = _global_config(directory, f"{wire.url}/v1") + with owned_proxy_process(gateway, directory, {}, config=config, workers=2) as owned: + yield _GlobalRig(wire, owned) + + +@pytest.fixture(scope="module") +def spend() -> Iterator[pcb.SpendLogs]: + with pcb.spend_logs() as logs: + yield logs + + +def _path(endpoint: Endpoint) -> str: + match endpoint: + case "chat": + return "/v1/chat/completions" + case "messages": + return "/v1/messages" + case "responses": + return "/v1/responses" + + +def _body(model: str, call: _Call, breakpoint: JsonValue) -> Mapping[str, JsonValue]: + common: Final[Mapping[str, JsonValue]] = {"model": model, "stream": call.stream, **pcb.NO_CACHE} + text: Final = pcb.marked(pcb.text(pcb.prompt(call.marker)), breakpoint) + match call.endpoint: + case "chat": + return {**common, "messages": [{"role": "user", "content": [text]}]} + case "messages": + return {**common, "max_tokens": 64, "messages": [{"role": "user", "content": [text]}]} + case "responses": + return { + **common, + "input": [{"type": "message", "role": "user", "content": [{**text, "type": "input_text"}]}], + } + + +def _calls(count: int, endpoints: tuple[Endpoint, ...], *, stream: bool | None = None) -> tuple[_Call, ...]: + return tuple( + _Call(endpoints[index % len(endpoints)], index % 2 == 1 if stream is None else stream, uuid.uuid4().hex) + for index in range(count) + ) + + +async def _send(client: httpx.AsyncClient, key: str, model: str, call: _Call, breakpoint: JsonValue) -> _Served: + async with client.stream( + "POST", + _path(call.endpoint), + json=_body(model, call, breakpoint), + headers={"Authorization": f"Bearer {key}", "anthropic-version": "2023-06-01"}, + ) as response: + raw: Final = await response.aread() + return _Served(call, response.status_code, raw.decode(), response.headers["x-litellm-call-id"]) + + +async def _burst( + gateway: Gateway, + model: str, + calls: tuple[_Call, ...], + *, + breakpoint: JsonValue = pcb.EXPLICIT, + tolerate_transport_errors: bool = False, +) -> tuple[_Served, ...]: + async with httpx.AsyncClient(base_url=str(gateway.client.base_url), timeout=60, trust_env=False) as client: + results: Final = await asyncio.gather( + *(_send(client, gateway.key, model, call, breakpoint) for call in calls), + return_exceptions=tolerate_transport_errors, + ) + for result in results: + assert not isinstance(result, BaseException) or isinstance(result, httpx.TransportError), repr(result) + return tuple(result for result in results if isinstance(result, _Served)) + + +def _frames(text: str) -> tuple[Mapping[str, JsonValue], ...]: + return tuple(rv.JSON_OBJECT.validate_json(line[6:]) for line in text.splitlines() if line.startswith("data: {")) + + +def _upstream_id_shown_to_caller(served: _Served) -> str | None: + if served.call.endpoint == "messages": + return None + if not served.call.stream: + return string_value(rv.JSON_OBJECT.validate_json(served.text)["id"]) + frames: Final = _frames(served.text) + if served.call.endpoint == "responses": + (completed,) = [frame for frame in frames if frame.get("type") == "response.completed"] + return string_value(rv.JSON_OBJECT.validate_python(completed["response"])["id"]) + return string_value(frames[0]["id"]) + + +def _assert_answered_in_its_own_shape(served: _Served) -> None: + assert served.status == 200, served.text + assert set(rv.MARKER.findall(served.text)) == {served.call.marker}, served.text + assert served.text.startswith(("event:", "data:")) == served.call.stream, served.text + assert served.text.startswith("{") != served.call.stream, served.text + assert ("response.completed" in served.text) == (served.call.stream and served.call.endpoint == "responses") + shown: Final = _upstream_id_shown_to_caller(served) + assert shown is None or pcb.answers(shown, served.call.marker), served.text + + +def _marked_once(posts: Sequence[Request], calls: Sequence[_Call], expected: JsonValue) -> None: + by_marker: Final = {marker: request for request in posts if (marker := rv.newest_marker(request.body.decode()))} + assert len(by_marker) == len(posts), [request.body for request in posts] + assert set(by_marker) == {call.marker for call in calls}, sorted(by_marker) + for call in calls: + block: Final = pcb.single_block(pcb.input_items(by_marker[call.marker]), "user") + assert block["type"] == "input_text" and block["text"] == pcb.prompt(call.marker), block + pcb.assert_marker(block, expected) + + +def _assert_each_lands_once( + spend: pcb.SpendLogs, model: str, failed: Sequence[_Served], served: Sequence[_Served] +) -> None: + expected: Final = len(failed) + len(served) + rows: Final = eventually(lambda: spend.rows_for(model), lambda found: len(found) >= expected, seconds=70) + by_call: Final = {string_value(row["litellm_call_id"]): row for row in rows} + assert len(by_call) == len(rows) == expected, rows + for item in failed: + assert by_call[item.call_id]["status"] == "failure", (item.call_id, rows) + for item in served: + row: Final = by_call[item.call_id] + assert row["status"] == "success", (item.call_id, row) + shown: Final = _upstream_id_shown_to_caller(item) + assert shown is None or rv.same_response(string_value(row["request_id"]), shown), (row, shown) + + +def _health(gateway: Gateway, model: str) -> Mapping[str, JsonValue]: + response: Final = gateway.request("GET", f"/health?model={model}", None) + assert response.status_code in (200, 503), response.text + return rv.JSON_OBJECT.validate_json(response.text) + + +def _free_port() -> int: + with wire_server(pcb.respond) as probe: + port: Final = urlsplit(probe.url).port + assert port is not None, probe.url + return port + + +def _chat(gateway: Gateway, model: str, marker: str, breakpoint: JsonValue) -> httpx.Response: + return gateway.request("POST", "/v1/chat/completions", dict(_body(model, _Call("chat", False, marker), breakpoint))) + + +def _completion(response: httpx.Response, marker: str) -> str: + assert response.status_code == 200, response.text + body: Final = rv.JSON_OBJECT.validate_json(response.text) + assert pcb.answers(string_value(body["id"]), marker), body + assert rv.answer(marker) in response.text, response.text + return response.headers["x-litellm-call-id"] + + +def _user_block_on_wire(wire: Wire, marker: str) -> dict[str, JsonValue]: + block: Final = pcb.single_block(pcb.input_items(pcb.posted(wire, marker)), "user") + assert block["type"] == "input_text" and block["text"] == pcb.prompt(marker), block + return block + + +@pytest.mark.parametrize("model", (_GLOBAL_UNSET, _GLOBAL_FALSE), ids=("deployment-unset", "deployment-false")) +def test_global_drop_params_drops_a_malformed_marker(global_rig: _GlobalRig, model: str, spend: pcb.SpendLogs) -> None: + marker: Final = uuid.uuid4().hex + call_id: Final = _completion(_chat(global_rig.gateway, model, marker, "yes"), marker) + pcb.assert_marker(_user_block_on_wire(global_rig.wire, marker), None) + spend.landed(model, call_id, marker) + control: Final = uuid.uuid4().hex + control_id: Final = _completion(_chat(global_rig.gateway, model, control, pcb.EXPLICIT), control) + pcb.assert_marker(_user_block_on_wire(global_rig.wire, control), pcb.EXPLICIT) + spend.landed(model, control_id, control) + + +async def test_mixed_burst_carries_every_marker_once(gateway: Gateway, spend: pcb.SpendLogs) -> None: + calls: Final = _calls(24, _ENDPOINTS) + with wire_server(pcb.respond) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=pcb.MODEL, api_base=f"{wire.url}/v1", drop_params=True) + served: Final = await _burst(gateway, model, calls) + assert len(served) == 24 + for item in served: + _assert_answered_in_its_own_shape(item) + _marked_once(pcb.drained_posts(wire), calls, pcb.EXPLICIT) + _assert_each_lands_once(spend, model, (), served) + + +async def test_upstream_outage_fails_cleanly_and_the_restarted_upstream_serves_marked_calls( + gateway: Gateway, spend: pcb.SpendLogs +) -> None: + port: Final = _free_port() + while_down: Final = _calls(12, _ENDPOINTS) + after: Final = _calls(12, _ENDPOINTS) + with gateway.scenario() as scenario: + model: Final = scenario.model(model=pcb.MODEL, api_base=f"http://127.0.0.1:{port}/v1", drop_params=True) + failed: Final = await _burst(gateway, model, while_down) + assert len(failed) == 12 + for item in failed: + assert item.status >= 500, (item.status, item.text) + assert "answer marker" not in item.text and "event:" not in item.text, item.text + down: Final = _health(gateway, model) + assert (down["healthy_count"], down["unhealthy_count"]) == (0, 1), down + with wire_server(pcb.respond, port=port) as wire: + _health(gateway, model) + probes: Final = pcb.drained_posts(wire) + assert [rv.newest_marker(request.body.decode()) for request in probes] == [None], probes + served: Final = await _burst(gateway, model, after) + assert len(served) == 12 + for item in served: + _assert_answered_in_its_own_shape(item) + _marked_once(pcb.drained_posts(wire), after, pcb.EXPLICIT) + _assert_each_lands_once(spend, model, failed, served) + + +def _slow(request: Request) -> Reply: + reply: Final = pcb.respond(request) + return dataclasses.replace(reply, pause_between_chunks=0.4) if reply.chunks else reply + + +async def test_concurrent_slow_streams_each_complete_with_one_upstream_call( + gateway: Gateway, spend: pcb.SpendLogs +) -> None: + calls: Final = _calls(6, ("chat",), stream=True) + with wire_server(_slow) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=pcb.MODEL, api_base=f"{wire.url}/v1", drop_params=True) + served: Final = await _burst(gateway, model, calls) + assert len(served) == 6 + for item in served: + _assert_answered_in_its_own_shape(item) + _marked_once(pcb.drained_posts(wire), calls, pcb.EXPLICIT) + _assert_each_lands_once(spend, model, (), served) + + +@dataclass(frozen=True, slots=True) +class _Held: + release: threading.Event + markers: SimpleQueue[str] + + def respond(self, request: Request) -> Reply: + marker: Final = rv.newest_marker(request.body.decode()) if request.method == "POST" else None + if marker is None: + return pcb.respond(request) + self.markers.put(marker) + if not self.release.wait(timeout=60): + return rv.error(504, "the burst was never released", "held") + return pcb.respond(request) + + +def _worker_pids(owned: OwnedProxy) -> tuple[int, ...]: + return eventually(lambda: pcb.started_worker_pids(owned.log), lambda pids: len(pids) == 2, seconds=30) + + +async def _hold_burst( + held: _Held, candidate: Gateway, model: str, calls: tuple[_Call, ...] +) -> asyncio.Task[tuple[_Served, ...]]: + burst: Final = asyncio.create_task(_burst(candidate, model, calls, tolerate_transport_errors=True)) + await asyncio.to_thread(eventually, held.markers.qsize, lambda size: size == len(calls), 60) + return burst + + +async def test_worker_sigkill_mid_burst_leaves_the_sibling_answering( + gateway: Gateway, tmp_path: Path, spend: pcb.SpendLogs +) -> None: + calls: Final = _calls(20, ("chat",), stream=False) + held: Final = _Held(threading.Event(), SimpleQueue()) + with wire_server(held.respond) as wire: + config: Final = _global_config(tmp_path, f"{wire.url}/v1") + with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as owned: + candidate: Final = owned.gateway + workers: Final = _worker_pids(owned) + burst: Final = await _hold_burst(held, candidate, _GLOBAL_UNSET, calls) + held_by: Final = MappingProxyType({pid: pcb.open_upstream_connections(pid, wire.url) for pid in workers}) + assert sum(held_by.values()) == 20, held_by + victim_pid, survivor_pid = sorted(workers, key=held_by.__getitem__) + victim: Final = psutil.Process(victim_pid) + victim.suspend() + victim.send_signal(signal.SIGKILL) + held.release.set() + served: Final = await burst + assert held_by[survivor_pid] >= 10, held_by + assert len(served) == held_by[survivor_pid], (held_by, len(served)) + for item in served: + _assert_answered_in_its_own_shape(item) + _marked_once(pcb.drained_posts(wire), calls, pcb.EXPLICIT) + follow_up: Final = uuid.uuid4().hex + call_id: Final = _completion(_chat(candidate, _GLOBAL_UNSET, follow_up, "yes"), follow_up) + pcb.assert_marker(_user_block_on_wire(wire, follow_up), None) + spend.landed(_GLOBAL_UNSET, call_id, follow_up) + + +async def test_proxy_restart_mid_burst_never_lands_a_served_call_twice( + gateway: Gateway, tmp_path: Path, record_property: _RecordProperty, spend: pcb.SpendLogs +) -> None: + calls: Final = _calls(20, ("chat",), stream=False) + held: Final = _Held(threading.Event(), SimpleQueue()) + with wire_server(held.respond) as wire: + config: Final = _global_config(tmp_path, f"{wire.url}/v1") + with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as first: + _worker_pids(first) + burst: Final = await _hold_burst(held, first.gateway, _GLOBAL_UNSET, calls) + first.process.terminate() + held.release.set() + served: Final = await burst + for item in served: + _assert_answered_in_its_own_shape(item) + second_directory: Final = tmp_path / "second" + second_directory.mkdir() + with owned_proxy_process(gateway, second_directory, {}, config=config, workers=2) as second: + follow_up: Final = uuid.uuid4().hex + call_id: Final = _completion(_chat(second.gateway, _GLOBAL_UNSET, follow_up, pcb.EXPLICIT), follow_up) + pcb.assert_marker(_user_block_on_wire(wire, follow_up), pcb.EXPLICIT) + spend.landed(_GLOBAL_UNSET, call_id, follow_up) + counts: Final = Counter(string_value(row["litellm_call_id"]) for row in spend.rows_for(_GLOBAL_UNSET)) + assert all(count == 1 for count in counts.values()), counts + landed: Final = sum(1 for item in served if item.call_id in counts) + record_property("served", len(served)) + record_property("landed", landed) + record_property("lost_responses", len(calls) - len(served)) diff --git a/tests/integration/providers/test_responses_bridge_stream_output_items_wire.py b/tests/integration/providers/test_responses_bridge_stream_output_items_wire.py new file mode 100644 index 00000000000..a325686e05a --- /dev/null +++ b/tests/integration/providers/test_responses_bridge_stream_output_items_wire.py @@ -0,0 +1,248 @@ +import json +import uuid +from collections.abc import Callable, Iterator, Sequence +from contextlib import contextmanager +from typing import Final + +import openai +from integration._support.client import Gateway, Scenario, eventually, object_value +from integration._support.database import read_rows +from integration._support.wire import Reply, Request, Wire, wire_server +from pydantic import JsonValue, TypeAdapter + +_BACKEND: Final = "qwen3-bridge-items" +_API_KEY: Final = "synthetic-hosted-vllm-key" +_QUESTION: Final = "What is the weather in Paris?" +_ANSWER: Final = "Paris is 22 degrees Celsius with clear skies." +_PREFACE: Final = "Checking the weather." +_REASONING: Final = "The user asks for the weather, so the weather tool applies." +_CALL_ID: Final = "call_bridge_items_1" +_ARGUMENTS: Final[dict[str, JsonValue]] = {"city": "Paris"} +_TOOL: Final[dict[str, JsonValue]] = { + "type": "function", + "name": "get_weather", + "description": "Weather for a city", + "parameters": {"type": "object", "properties": {"city": {"type": "string"}}, "required": ["city"]}, +} +_NO_CACHE: Final[dict[str, JsonValue]] = {"no-cache": True} +_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) +_DISCOVERY_PROBE: Final = ("GET", "/v1/models") + + +def _frame(identity: str, delta: dict[str, JsonValue], finish_reason: str | None = None) -> bytes: + usage: Final = {"usage": {"prompt_tokens": 30, "completion_tokens": 12, "total_tokens": 42}} if finish_reason else {} + body: Final = { + "id": identity, + "object": "chat.completion.chunk", + "created": 1, + "model": _BACKEND, + "choices": [{"index": 0, "delta": delta, "finish_reason": finish_reason}], + **usage, + } + return b"data: " + json.dumps(body).encode() + b"\n\n" + + +def _sse_reply(identity: str, deltas: Sequence[dict[str, JsonValue]], finish_reason: str) -> Reply: + return Reply( + content_type="text/event-stream", + chunks=( + _frame(identity, {"role": "assistant", "content": ""}), + *(_frame(identity, delta) for delta in deltas), + _frame(identity, {}, finish_reason), + b"data: [DONE]\n\n", + ), + ) + + +def _tool_call_delta() -> dict[str, JsonValue]: + return { + "tool_calls": [ + { + "index": 0, + "id": _CALL_ID, + "type": "function", + "function": {"name": "get_weather", "arguments": json.dumps(_ARGUMENTS)}, + } + ] + } + + +def _is_discovery_probe(request: Request) -> bool: + return (request.method, request.target) == _DISCOVERY_PROBE + + +@contextmanager +def _vllm_server(respond: Callable[[Request], Reply]) -> Iterator[Wire]: + with wire_server( + lambda request: Reply(body=b'{"object":"list","data":[]}') if _is_discovery_probe(request) else respond(request) + ) as wire: + yield wire + + +def _bridged_vllm_model(scenario: Scenario, wire: Wire) -> str: + return scenario.model( + model=f"hosted_vllm/{_BACKEND}", api_base=wire.url + "/v1", api_key=_API_KEY, use_chat_completions_api=True + ) + + +def _only_streamed_chat(wire: Wire, *tool_names: str) -> dict[str, JsonValue]: + received: Final = tuple(request for request in wire.drain() if not _is_discovery_probe(request)) + assert [(request.method, request.target) for request in received] == [("POST", "/v1/chat/completions")] + assert received[0].headers["authorization"] == f"Bearer {_API_KEY}" + body: Final = _JSON_OBJECT.validate_json(received[0].body) + assert body["stream"] is True + tools: Final = body.get("tools", []) + assert isinstance(tools, list) + assert [object_value(object_value(tool)["function"])["name"] for tool in tools] == list(tool_names), tools + messages: Final = body["messages"] + assert isinstance(messages, list) and object_value(messages[-1])["content"] == _QUESTION, messages + return body + + +def _spend_statuses(model: str) -> list[JsonValue]: + rows: Final = eventually( + lambda: read_rows('SELECT status FROM "LiteLLM_SpendLogs" WHERE model_group=%s', (model,)), + lambda found: len(found) >= 1, + seconds=70, + ) + return [row["status"] for row in rows] + + +def _openai_client(gateway: Gateway) -> openai.OpenAI: + return openai.OpenAI(base_url=str(gateway.client.base_url) + "/v1", api_key=gateway.key, max_retries=0) + + +def _async_openai_client(gateway: Gateway) -> openai.AsyncOpenAI: + return openai.AsyncOpenAI(base_url=str(gateway.client.base_url) + "/v1", api_key=gateway.key, max_retries=0) + + +def _raw_events(gateway: Gateway, body: dict[str, JsonValue]) -> list[dict[str, JsonValue]]: + with gateway.client.stream( + "POST", + "/v1/responses", + json={**body, "cache": _NO_CACHE}, + headers={"Authorization": f"Bearer {gateway.key}"}, + ) as response: + assert response.status_code == 200, response.read() + return [ + _JSON_OBJECT.validate_json(line.removeprefix("data: ")) + for line in response.iter_lines() + if line.startswith("data: {") + ] + + +def _raw_added_items(events: Sequence[dict[str, JsonValue]]) -> Iterator[tuple[JsonValue, JsonValue]]: + for event in events: + if event["type"] == "response.output_item.added": + yield object_value(event["item"])["type"], event["output_index"] + + +def _raw_completed_output_types(events: Sequence[dict[str, JsonValue]]) -> list[JsonValue]: + completed: Final = [event for event in events if event["type"] == "response.completed"] + assert len(completed) == 1, [event["type"] for event in events] + output: Final = object_value(completed[0]["response"])["output"] + assert isinstance(output, list) + return [object_value(item)["type"] for item in output] + + +def test_openai_sdk_responses_stream_with_a_tool_call_only_reply_announces_no_message_item(gateway: Gateway) -> None: + identity: Final = f"chatcmpl-bridge-{uuid.uuid4().hex}" + reply: Final = _sse_reply(identity, (_tool_call_delta(),), "tool_calls") + with _vllm_server(lambda _: reply) as wire, gateway.scenario() as scenario: + model: Final = _bridged_vllm_model(scenario, wire) + with _openai_client(gateway).responses.stream( + model=model, + input=_QUESTION, + tools=[_TOOL], # pyright: ignore[reportArgumentType] # plain JSON tool + store=False, + extra_body={"cache": _NO_CACHE}, + ) as stream: + events: Final = list(stream) + final: Final = stream.get_final_response() + added: Final = [ + (event.item.type, event.output_index) for event in events if event.type == "response.output_item.added" + ] + assert added == [("function_call", 0)], [event.type for event in events] + assert [item.type for item in final.output] == ["function_call"], final.output + call: Final = final.output[0] + assert call.type == "function_call" and call.name == "get_weather" and call.call_id == _CALL_ID + assert json.loads(call.arguments) == _ARGUMENTS + assert final.output_text == "" + created: Final = [event for event in events if event.type == "response.created"] + assert len(created) == 1 and created[0].response.id == final.id, [event.type for event in events] + _only_streamed_chat(wire, "get_weather") + assert _spend_statuses(model) == ["success"] + + +async def test_async_openai_sdk_responses_stream_text_reply_has_one_message_item_at_index_zero( + gateway: Gateway, +) -> None: + identity: Final = f"chatcmpl-bridge-{uuid.uuid4().hex}" + reply: Final = _sse_reply(identity, ({"content": _ANSWER[:17]}, {"content": _ANSWER[17:]}), "stop") + with _vllm_server(lambda _: reply) as wire, gateway.scenario() as scenario: + model: Final = _bridged_vllm_model(scenario, wire) + stream: Final = await _async_openai_client(gateway).responses.create( + model=model, input=_QUESTION, store=False, stream=True, extra_body={"cache": _NO_CACHE} + ) + events: Final = [event async for event in stream] + added: Final = [ + (event.item.type, event.output_index) for event in events if event.type == "response.output_item.added" + ] + assert added == [("message", 0)], [event.type for event in events] + done: Final = [(event.item.type, event.output_index) for event in events if event.type == "response.output_item.done"] + assert done == [("message", 0)], [event.type for event in events] + assert "".join(event.delta for event in events if event.type == "response.output_text.delta") == _ANSWER + completed: Final = [event for event in events if event.type == "response.completed"] + assert len(completed) == 1 and completed[0].response.output_text == _ANSWER + assert [item.type for item in completed[0].response.output] == ["message"], completed[0].response.output + created: Final = [event for event in events if event.type == "response.created"] + assert len(created) == 1 and created[0].response.id == completed[0].response.id + _only_streamed_chat(wire) + assert _spend_statuses(model) == ["success"] + + +async def test_async_openai_sdk_responses_stream_reasoning_then_text_gets_contiguous_output_indexes( + gateway: Gateway, +) -> None: + identity: Final = f"chatcmpl-bridge-{uuid.uuid4().hex}" + reply: Final = _sse_reply(identity, ({"reasoning_content": _REASONING}, {"content": _ANSWER}), "stop") + with _vllm_server(lambda _: reply) as wire, gateway.scenario() as scenario: + model: Final = _bridged_vllm_model(scenario, wire) + stream: Final = await _async_openai_client(gateway).responses.create( + model=model, input=_QUESTION, store=False, stream=True, extra_body={"cache": _NO_CACHE} + ) + events: Final = [event async for event in stream] + added: Final = [ + (event.item.type, event.output_index) for event in events if event.type == "response.output_item.added" + ] + assert added == [("reasoning", 0), ("message", 1)], [event.type for event in events] + assert "".join(event.delta for event in events if event.type == "response.output_text.delta") == _ANSWER + completed: Final = [event for event in events if event.type == "response.completed"] + assert len(completed) == 1 and completed[0].response.output_text == _ANSWER + assert [item.type for item in completed[0].response.output] == ["reasoning", "message"], ( + completed[0].response.output + ) + _only_streamed_chat(wire) + assert _spend_statuses(model) == ["success"] + + +def test_raw_responses_stream_text_then_tool_call_keeps_the_message_first(gateway: Gateway) -> None: + identity: Final = f"chatcmpl-bridge-{uuid.uuid4().hex}" + reply: Final = _sse_reply(identity, ({"content": _PREFACE}, _tool_call_delta()), "tool_calls") + with _vllm_server(lambda _: reply) as wire, gateway.scenario() as scenario: + model: Final = _bridged_vllm_model(scenario, wire) + events: Final = _raw_events( + gateway, {"model": model, "input": _QUESTION, "tools": [_TOOL], "store": False, "stream": True} + ) + assert list(_raw_added_items(events)) == [("message", 0), ("function_call", 1)], [e["type"] for e in events] + assert _raw_completed_output_types(events) == ["message", "function_call"] + deltas: Final = [event["delta"] for event in events if event["type"] == "response.output_text.delta"] + assert "".join(str(delta) for delta in deltas) == _PREFACE, deltas + done_calls: Final = [ + object_value(event["item"]) + for event in events + if event["type"] == "response.output_item.done" and object_value(event["item"])["type"] == "function_call" + ] + assert [(item["name"], item["call_id"]) for item in done_calls] == [("get_weather", _CALL_ID)], done_calls + _only_streamed_chat(wire, "get_weather") + assert _spend_statuses(model) == ["success"] diff --git a/tests/integration/providers/test_responses_websocket_session_limit.py b/tests/integration/providers/test_responses_websocket_session_limit.py new file mode 100644 index 00000000000..32b01d8387a --- /dev/null +++ b/tests/integration/providers/test_responses_websocket_session_limit.py @@ -0,0 +1,1479 @@ +from __future__ import annotations + +import asyncio +import itertools +import json +import re +import ssl +import threading +import time +import uuid +from collections.abc import Generator, Iterator, Mapping +from contextlib import contextmanager +from dataclasses import dataclass +from pathlib import Path +from queue import Empty, SimpleQueue +from types import MappingProxyType +from typing import Final, TypeVar +from urllib.parse import parse_qs, urlsplit + +import httpx +import psutil +import pytest +import websockets +from integration._support.client import Gateway, eventually, gateway_from_environment +from integration._support.database import read_rows, scratch_database +from integration._support.process import OwnedProxy, graceful_stop_seconds, owned_proxy, owned_proxy_process +from integration._support.tls import server_context, write_self_signed_cert +from integration._support.upstream import ( + JsonResponse, + RoutedResponse, + delete_scenario, + register_scenario, +) +from pydantic import JsonValue, TypeAdapter, ValidationError +from websockets.asyncio.server import ServerConnection, serve +from websockets.exceptions import ConnectionClosed, InvalidStatus +from websockets.typing import Subprotocol + +OWNED_PROXY_CELL_SECONDS: Final = 2 * graceful_stop_seconds() + 180 +pytestmark: Final = pytest.mark.timeout(max(360.0, OWNED_PROXY_CELL_SECONDS)) + +PROVIDER_MODEL: Final = "ws-peer-model" +STALL_PROVIDER_MODEL: Final = "ws-stall-peer-model" +DEAF_PROVIDER_MODEL: Final = "ws-deaf-peer-model" +DEAF_RESPONSE_DELAY_SECONDS: Final = 10 +DEAF_READ_PAUSE_SECONDS: Final = 20 +PEER_TEXT: Final = "responses websocket peer" +TERMINAL: Final = frozenset({"response.completed", "response.failed", "error"}) +JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) +SESSION_LIMIT_FIELD: Final = "responses_websocket_session_limit_seconds" +LIMIT_CLOSE_REASON: Final = "Session duration limit reached" +RELOAD_INTERVAL_SECONDS: Final = 3 +WORKER_SYNC_SECONDS: Final = RELOAD_INTERVAL_SECONDS + 7 +CAP_SECONDS: Final = 60 +RESTORED_IDLE_SECONDS: Final = 75 +CAPPED_PATHS: Final = ("/v1/responses", "/responses") +CAPPED_SOCKETS_MINIMUM: Final = 8 +CAPPED_SOCKETS_MAXIMUM: Final = 40 +KILLED_POOL_SIZE: Final = 8 +SUBPROTOCOLS: Final = (Subprotocol("litellm-responses-first"), Subprotocol("litellm-responses-second")) +STARTED_WORKER: Final = re.compile(r"Started server process \[(\d+)\]") +WORKER_PIDS: Final = TypeAdapter(tuple[int, ...]) +CLIENT_ADDRESS: Final = TypeAdapter(tuple[str, int]) +T = TypeVar("T") + + +@dataclass(frozen=True, slots=True) +class PeerConnection: + path: str + frames: SimpleQueue[dict[str, JsonValue]] + closed: SimpleQueue[float] + + +@dataclass(frozen=True, slots=True) +class ResponsesPeer: + url: str + connections: SimpleQueue[PeerConnection] + connections_by_model: Mapping[str, SimpleQueue[PeerConnection]] + + +@dataclass(frozen=True, slots=True) +class PeerSnapshot: + path: str + frames: tuple[dict[str, JsonValue], ...] + closed: bool + + +@dataclass(frozen=True, slots=True) +class ReceiveResult: + timed_out: bool + closed: bool + frame: dict[str, JsonValue] | None + close_code: int | None + close_reason: str | None + + +@dataclass(frozen=True, slots=True) +class SessionResult: + text: str + idle: ReceiveResult + events: tuple[dict[str, JsonValue], ...] + error: str | None + + +@dataclass(frozen=True, slots=True) +class HandshakeResult: + status_code: int | None + error: str | None + + +@dataclass(frozen=True, slots=True) +class SubprotocolResult: + text: str + negotiated: str | None + events: tuple[dict[str, JsonValue], ...] + error: str | None + + +@dataclass(frozen=True, slots=True) +class DefaultResults: + pool: tuple[SessionResult, ...] + query: SessionResult + rejected: HandshakeResult + subprotocol: SubprotocolResult + provider: tuple[PeerSnapshot, ...] + + +@dataclass(frozen=True, slots=True) +class AuthResult: + idle: ReceiveResult + frame: dict[str, JsonValue] | None + close_code: int | None + close_reason: str | None + error: str | None + + +@dataclass(frozen=True, slots=True) +class AuthResults: + immediate: AuthResult + delayed: AuthResult + provider_connections: int + + +@dataclass(frozen=True, slots=True) +class CloseResult: + outcome: ReceiveResult + elapsed: float + + +@dataclass(frozen=True, slots=True) +class ActiveResult: + turn: tuple[dict[str, JsonValue], ...] + turn_error: str | None + close: CloseResult + provider_closed: bool + + +@dataclass(frozen=True, slots=True) +class MidResult: + created: bool + turn_error: str | None + close: CloseResult + provider_closed: bool + fresh_completed: bool + fresh_error: str | None + + +@dataclass(frozen=True, slots=True) +class CapResults: + idle: CloseResult + active: ActiveResult + mid: MidResult + deaf: MidResult + provider: tuple[PeerSnapshot, ...] + + +@dataclass(frozen=True, slots=True) +class BurstResult: + path: str + status_code: int | None + body: dict[str, JsonValue] | None + error: str | None + + +@dataclass(frozen=True, slots=True) +class InvalidResult: + session: SessionResult + warning_found: bool + + +@dataclass(frozen=True, slots=True) +class UpdateResult: + status_code: int + body: str + + +@dataclass(frozen=True, slots=True) +class CappedSocket: + path: str + worker_pid: int | None + close: CloseResult + + +@dataclass(frozen=True, slots=True) +class OpenedSocket: + path: str + worker_pid: int | None + connection: websockets.ClientConnection + started: float + + +@dataclass(frozen=True, slots=True) +class HeldResult: + held_seconds: float + events: tuple[dict[str, JsonValue], ...] + error: str | None + + +@dataclass(frozen=True, slots=True) +class OverrideResults: + workers: frozenset[int] + updates: tuple[UpdateResult, ...] + stored_after_updates: JsonValue + capped: tuple[CappedSocket, ...] + earlier: HeldResult + delete: UpdateResult + stored_after_delete: dict[str, JsonValue] | None + restored: SessionResult + + +@dataclass(frozen=True, slots=True) +class TurnOutcome: + events: tuple[dict[str, JsonValue], ...] + error: str | None + + +@dataclass(frozen=True, slots=True) +class KillResults: + held_by: Mapping[int, int] + victim: int + victim_closes: tuple[ReceiveResult, ...] + survivor_turns: tuple[TurnOutcome, ...] + health_status: int + fresh_completed: bool + + +@dataclass(frozen=True, slots=True) +class ChaosResults: + burst: tuple[BurstResult, ...] + health_status: int + websocket_completed: bool + kill: KillResults + + +def _object(value: JsonValue) -> dict[str, JsonValue]: + assert isinstance(value, dict), value + return value + + +def _list(value: JsonValue) -> list[JsonValue]: + assert isinstance(value, list), value + return value + + +def _string(value: JsonValue) -> str: + assert isinstance(value, str), value + return value + + +def _drain(queue: SimpleQueue[T]) -> tuple[T, ...]: + return tuple(queue.get_nowait() for _ in range(queue.qsize())) + + +def _events(response_id: str, model: str, text: str, *, stall: bool = False) -> tuple[dict[str, JsonValue], ...]: + message: Final[dict[str, JsonValue]] = { + "type": "message", + "id": f"msg_{response_id}", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": text, "annotations": []}], + } + response: Final[dict[str, JsonValue]] = { + "id": response_id, + "object": "response", + "created_at": 1700000000, + "model": model, + } + created: Final[dict[str, JsonValue]] = { + "type": "response.created", + "response": {**response, "status": "in_progress", "output": []}, + } + if stall: + return (created,) + return ( + created, + { + "type": "response.output_text.delta", + "item_id": f"msg_{response_id}", + "output_index": 0, + "content_index": 0, + "delta": text, + }, + { + "type": "response.completed", + "response": {**response, "status": "completed", "output": [message]}, + }, + ) + + +def _path_model(path: str) -> str | None: + values: Final = parse_qs(urlsplit(path).query).get("model") + return values[0] if values else None + + +async def _peer_handler(connection: ServerConnection, peer: ResponsesPeer) -> None: + path: Final = connection.request.path if connection.request is not None else "" + record: Final = PeerConnection(path, SimpleQueue(), SimpleQueue()) + peer.connections.put(record) + model: Final = _path_model(path) + if model is not None and model in peer.connections_by_model: + peer.connections_by_model[model].put(record) + turns: Final = itertools.count(1) + try: + async for raw in connection: + await _peer_frame(raw, connection, record, turns) + finally: + record.closed.put(time.monotonic()) + + +async def _peer_frame( + raw: str | bytes, + connection: ServerConnection, + record: PeerConnection, + turns: itertools.count[int], +) -> None: + frame: Final = JSON_OBJECT.validate_json(raw) + record.frames.put(frame) + if frame.get("type") != "response.create": + return + model: Final = _string(frame.get("model", "")) + stall: Final = model in (STALL_PROVIDER_MODEL, DEAF_PROVIDER_MODEL) + if model == DEAF_PROVIDER_MODEL: + await asyncio.sleep(DEAF_RESPONSE_DELAY_SECONDS) + for event in _events(f"resp_peer_{next(turns)}", model, PEER_TEXT, stall=stall): + await connection.send(json.dumps(event)) + if model == DEAF_PROVIDER_MODEL: + transport: Final = connection.transport + assert transport is not None + transport.pause_reading() + try: + await asyncio.sleep(DEAF_READ_PAUSE_SECONDS) + finally: + transport.resume_reading() + + +async def _serve_peer( + tls: ssl.SSLContext, + peer: ResponsesPeer, + ports: SimpleQueue[int], + stop: asyncio.Event, +) -> None: + async with serve(lambda connection: _peer_handler(connection, peer), "127.0.0.1", 0, ssl=tls) as server: + address: object = next(iter(server.sockets)).getsockname() # pyright: ignore[reportAny] # socket.getsockname is typed Any + port: Final = TypeAdapter(tuple[str, int]).validate_python(address)[1] + ports.put(port) + await stop.wait() + + +@contextmanager +def responses_peer(cert: tuple[Path, Path]) -> Generator[ResponsesPeer, None, None]: + loop: Final = asyncio.new_event_loop() + stop: Final = asyncio.Event() + ports: Final = SimpleQueue[int]() + connections_by_model: Final[Mapping[str, SimpleQueue[PeerConnection]]] = MappingProxyType( + {model: SimpleQueue[PeerConnection]() for model in (PROVIDER_MODEL, STALL_PROVIDER_MODEL, DEAF_PROVIDER_MODEL)} + ) + peer: Final = ResponsesPeer("", SimpleQueue(), connections_by_model) + thread: Final = threading.Thread( + target=loop.run_until_complete, + args=(_serve_peer(server_context(*cert), peer, ports, stop),), + daemon=True, + ) + thread.start() + try: + port: Final = ports.get(timeout=10) + yield ResponsesPeer(f"https://127.0.0.1:{port}/v1", peer.connections, peer.connections_by_model) + finally: + loop.call_soon_threadsafe(stop.set) + thread.join(timeout=10) + loop.close() + + +def _proxy_ws_url(candidate: Gateway, path: str) -> str: + base: Final = str(candidate.client.base_url).rstrip("/").replace("http://", "ws://", 1) + return f"{base}{path}" + + +def _create(model: str, text: str, *, include_model: bool = True) -> str: + body: Final[dict[str, JsonValue]] = { + "type": "response.create", + "input": [{"type": "message", "role": "user", "content": [{"type": "input_text", "text": text}]}], + **({"model": model} if include_model else {}), + } + return json.dumps(body) + + +async def _receive(connection: websockets.ClientConnection, timeout: float) -> ReceiveResult: + try: + raw: Final = await asyncio.wait_for(connection.recv(), timeout=timeout) + except asyncio.TimeoutError: + return ReceiveResult(True, False, None, None, None) + except ConnectionClosed as error: + received: Final = error.rcvd + return ReceiveResult( + False, + True, + None, + received.code if received is not None else None, + received.reason if received is not None else None, + ) + return ReceiveResult(False, False, JSON_OBJECT.validate_json(raw), None, None) + + +async def _wait_for_close(connection: websockets.ClientConnection, timeout: float) -> ReceiveResult: + started: Final = time.monotonic() + result: Final = await _receive(connection, timeout) + if result.closed or result.timed_out: + return result + remaining: Final = timeout - (time.monotonic() - started) + if remaining <= 0: + return ReceiveResult(True, False, None, None, None) + return await _wait_for_close(connection, remaining) + + +async def _until_terminal( + connection: websockets.ClientConnection, + received: tuple[dict[str, JsonValue], ...] = (), +) -> tuple[dict[str, JsonValue], ...]: + event: Final = JSON_OBJECT.validate_json(await asyncio.wait_for(connection.recv(), timeout=20)) + collected: Final = (*received, event) + if event.get("type") in TERMINAL or len(collected) >= 50: + return collected + return await _until_terminal(connection, collected) + + +async def _turn(connection: websockets.ClientConnection, frame: str) -> tuple[dict[str, JsonValue], ...]: + await connection.send(frame) + return await _until_terminal(connection) + + +async def _idle_turn( + proxy: str, + key: str, + model: str, + path: str, + text: str, + *, + query_model: bool = False, + idle_seconds: float = 35, +) -> SessionResult: + query: Final = f"?model={model}" if query_model else "" + try: + async with websockets.connect( + f"{proxy}{path}{query}", + additional_headers={"Authorization": f"Bearer {key}"}, + open_timeout=10, + ) as connection: + idle: Final = await _receive(connection, idle_seconds) + try: + events: Final = await _turn(connection, _create(model, text, include_model=not query_model)) + except (ConnectionClosed, asyncio.TimeoutError) as error: + return SessionResult(text, idle, (), f"{type(error).__name__}: {error}") + return SessionResult(text, idle, events, None) + except (ConnectionClosed, asyncio.TimeoutError) as error: + return SessionResult(text, ReceiveResult(False, True, None, None, None), (), f"{type(error).__name__}: {error}") + + +async def _pool_workload(candidate: Gateway, key: str, model: str) -> tuple[SessionResult, ...]: + proxy: Final = _proxy_ws_url(candidate, "") + paths: Final = ("/v1/responses",) * 4 + ("/responses",) * 4 + return tuple( + await asyncio.gather(*tuple(_idle_turn(proxy, key, model, path, f"pool-{uuid.uuid4().hex}") for path in paths)) + ) + + +async def _rejected_handshake(proxy: str) -> HandshakeResult: + try: + async with websockets.connect( + f"{proxy}/v1/responses", + additional_headers={"Authorization": f"Bearer sk-rejected-{uuid.uuid4().hex}"}, + open_timeout=10, + ): + return HandshakeResult(None, "handshake accepted") + except InvalidStatus as error: + return HandshakeResult(error.response.status_code, None) + except (ConnectionClosed, asyncio.TimeoutError, OSError) as error: + return HandshakeResult(None, f"{type(error).__name__}: {error}") + + +async def _subprotocol_turn(proxy: str, key: str, model: str) -> SubprotocolResult: + text: Final = f"subprotocol-{uuid.uuid4().hex}" + try: + async with websockets.connect( + f"{proxy}/v1/responses", + additional_headers={"Authorization": f"Bearer {key}"}, + subprotocols=SUBPROTOCOLS, + open_timeout=10, + ) as connection: + events, error = await _turn_result(connection, _create(model, text)) + return SubprotocolResult(text, connection.subprotocol, events, error) + except (ConnectionClosed, asyncio.TimeoutError, InvalidStatus) as error: + return SubprotocolResult(text, None, (), f"{type(error).__name__}: {error}") + + +async def _default_workload( + candidate: Gateway, key: str, model: str +) -> tuple[tuple[SessionResult, ...], SessionResult, HandshakeResult, SubprotocolResult]: + proxy: Final = _proxy_ws_url(candidate, "") + pool, query, rejected, negotiated = await asyncio.gather( + _pool_workload(candidate, key, model), + _idle_turn(proxy, key, model, "/v1/responses", f"query-{uuid.uuid4().hex}", query_model=True), + _rejected_handshake(proxy), + _subprotocol_turn(proxy, key, model), + ) + return pool, query, rejected, negotiated + + +def _snapshot(record: PeerConnection) -> PeerSnapshot: + return PeerSnapshot(record.path, _drain(record.frames), record.closed.qsize() > 0) + + +def _snapshots(peer: ResponsesPeer, count: int) -> tuple[PeerSnapshot, ...]: + eventually(lambda: peer.connections.qsize(), lambda size: size >= count, seconds=15, return_last_on_timeout=True) + records: Final = _drain(peer.connections) + eventually( + lambda: tuple(record.closed.qsize() for record in records), + lambda counts: all(count > 0 for count in counts), + seconds=10, + return_last_on_timeout=True, + ) + return tuple(_snapshot(record) for record in records) + + +def _available_snapshots(peer: ResponsesPeer) -> tuple[PeerSnapshot, ...]: + eventually(lambda: peer.connections.qsize(), lambda count: count >= 1, seconds=10) + records: Final = _drain(peer.connections) + eventually( + lambda: tuple(record.closed.qsize() for record in records), + lambda counts: all(count > 0 for count in counts), + seconds=10, + ) + return tuple(_snapshot(record) for record in records) + + +def _frame_text(frame: dict[str, JsonValue]) -> str: + input_value: Final = _list(frame["input"]) + message: Final = _object(input_value[0]) + content: Final = _list(message["content"]) + item: Final = _object(content[0]) + return _string(item["text"]) + + +def _completed_text(events: tuple[dict[str, JsonValue], ...]) -> str: + assert events, events + completed: Final = events[-1] + assert completed.get("type") == "response.completed", completed + response: Final = _object(completed["response"]) + output: Final = _list(response["output"]) + message: Final = _object(output[0]) + content: Final = _list(message["content"]) + item: Final = _object(content[0]) + return _string(item["text"]) + + +def _auth_close(error: ConnectionClosed) -> tuple[int | None, str | None]: + received: Final = error.rcvd + return ( + received.code if received is not None else None, + received.reason if received is not None else None, + ) + + +async def _auth_attempt(proxy: str, key: str, model: str, *, delay: bool) -> AuthResult: + try: + async with websockets.connect( + f"{proxy}/v1/responses", + additional_headers={"Authorization": f"Bearer {key}"}, + open_timeout=10, + ) as connection: + idle: Final = await _receive(connection, 35) if delay else ReceiveResult(False, False, None, None, None) + try: + await connection.send(_create(model, f"auth-{uuid.uuid4().hex}")) + frame: Final = JSON_OBJECT.validate_json(await asyncio.wait_for(connection.recv(), timeout=20)) + try: + await asyncio.wait_for(connection.recv(), timeout=20) + except ConnectionClosed as error: + code, reason = _auth_close(error) + return AuthResult(idle, frame, code, reason, None) + return AuthResult(idle, frame, None, None, "connection remained open after rejection") + except (ConnectionClosed, asyncio.TimeoutError) as error: + code, reason = _auth_close(error) if isinstance(error, ConnectionClosed) else (None, None) + return AuthResult(idle, None, code, reason, f"{type(error).__name__}: {error}") + except (ConnectionClosed, asyncio.TimeoutError) as error: + return AuthResult( + ReceiveResult(False, True, None, None, None), + None, + None, + None, + f"{type(error).__name__}: {error}", + ) + + +async def _auth_workload(candidate: Gateway, key: str, model: str, peer: ResponsesPeer) -> AuthResults: + proxy: Final = _proxy_ws_url(candidate, "") + immediate, delayed = await asyncio.gather( + _auth_attempt(proxy, key, model, delay=False), + _auth_attempt(proxy, key, model, delay=True), + ) + return AuthResults(immediate, delayed, peer.connections.qsize()) + + +async def _close_at_limit( + candidate: Gateway, + key: str, + model: str, + *, + delay: float | None = None, +) -> CloseResult: + started: Final = time.monotonic() + outcome: Final = await _close_session(candidate, key, model, delay=delay) + return CloseResult(outcome, time.monotonic() - started) + + +async def _close_session( + candidate: Gateway, + key: str, + model: str, + *, + delay: float | None, +) -> ReceiveResult: + started: Final = time.monotonic() + try: + async with websockets.connect( + _proxy_ws_url(candidate, "/v1/responses"), + additional_headers={"Authorization": f"Bearer {key}"}, + open_timeout=10, + ) as connection: + if delay is not None: + await asyncio.sleep(max(0, started + delay - time.monotonic())) + return await _wait_for_close(connection, 75) + except (ConnectionClosed, asyncio.TimeoutError) as error: + code, reason = _auth_close(error) if isinstance(error, ConnectionClosed) else (None, None) + return ReceiveResult(False, True, None, code, reason) + + +async def _turn_result( + connection: websockets.ClientConnection, + frame: str, +) -> tuple[tuple[dict[str, JsonValue], ...], str | None]: + try: + return await _turn(connection, frame), None + except (ConnectionClosed, asyncio.TimeoutError) as error: + return (), f"{type(error).__name__}: {error}" + + +async def _active_session(candidate: Gateway, key: str, model: str, peer: ResponsesPeer) -> ActiveResult: + started: Final = time.monotonic() + try: + async with websockets.connect( + _proxy_ws_url(candidate, "/v1/responses"), + additional_headers={"Authorization": f"Bearer {key}"}, + open_timeout=10, + ) as connection: + await asyncio.sleep(max(0, started + 5 - time.monotonic())) + turn, turn_error = await _turn_result(connection, _create(model, f"active-{uuid.uuid4().hex}")) + remaining: Final = max(0, 75 - (time.monotonic() - started)) + close: Final = CloseResult(await _wait_for_close(connection, remaining), time.monotonic() - started) + provider_closed: Final = await asyncio.to_thread(_provider_closed_within, peer, PROVIDER_MODEL, 5) + return ActiveResult(turn, turn_error, close, provider_closed) + except (ConnectionClosed, asyncio.TimeoutError) as error: + code, reason = _auth_close(error) if isinstance(error, ConnectionClosed) else (None, None) + return ActiveResult( + (), + f"{type(error).__name__}: {error}", + CloseResult(ReceiveResult(False, True, None, code, reason), time.monotonic() - started), + False, + ) + + +def _provider_closed_within(peer: ResponsesPeer, model: str, seconds: float) -> bool: + deadline: Final = time.monotonic() + seconds + try: + record: Final = peer.connections_by_model[model].get(timeout=seconds) + except Empty: + return False + try: + eventually( + lambda: record.closed.qsize(), + lambda count: count >= 1, + seconds=max(0, deadline - time.monotonic()), + ) + except AssertionError: + return False + return True + + +async def _stall_turn(connection: websockets.ClientConnection, model: str) -> tuple[bool, str | None]: + try: + await connection.send(_create(model, f"stall-{uuid.uuid4().hex}")) + first: Final = JSON_OBJECT.validate_json(await asyncio.wait_for(connection.recv(), timeout=20)) + return first.get("type") == "response.created", None + except (ConnectionClosed, asyncio.TimeoutError) as error: + return False, f"{type(error).__name__}: {error}" + + +async def _mid_session( + candidate: Gateway, + key: str, + stall_model: str, +) -> tuple[bool, str | None, CloseResult]: + started: Final = time.monotonic() + try: + async with websockets.connect( + _proxy_ws_url(candidate, "/v1/responses"), + additional_headers={"Authorization": f"Bearer {key}"}, + open_timeout=10, + ) as connection: + await asyncio.sleep(max(0, started + 40 - time.monotonic())) + created, turn_error = await _stall_turn(connection, stall_model) + remaining: Final = max(0, 75 - (time.monotonic() - started)) + close: Final = CloseResult(await _wait_for_close(connection, remaining), time.monotonic() - started) + return created, turn_error, close + except (ConnectionClosed, asyncio.TimeoutError) as error: + code, reason = _auth_close(error) if isinstance(error, ConnectionClosed) else (None, None) + return ( + False, + f"{type(error).__name__}: {error}", + CloseResult( + ReceiveResult(False, True, None, code, reason), + time.monotonic() - started, + ), + ) + + +async def _fresh_session(candidate: Gateway, key: str, normal_model: str) -> tuple[bool, str | None]: + try: + async with websockets.connect( + _proxy_ws_url(candidate, "/v1/responses"), + additional_headers={"Authorization": f"Bearer {key}"}, + open_timeout=10, + ) as connection: + events: Final = await _turn(connection, _create(normal_model, f"fresh-{uuid.uuid4().hex}")) + return _completed_text(events) == PEER_TEXT, None + except (ConnectionClosed, asyncio.TimeoutError) as error: + return False, f"{type(error).__name__}: {error}" + + +async def _mid_response( + candidate: Gateway, + key: str, + normal_model: str, + stall_model: str, + peer: ResponsesPeer, +) -> MidResult: + created, turn_error, close = await _mid_session(candidate, key, stall_model) + provider_closed: Final = await asyncio.to_thread(_provider_closed_within, peer, STALL_PROVIDER_MODEL, 5) + fresh_completed, fresh_error = await _fresh_session(candidate, key, normal_model) + return MidResult(created, turn_error, close, provider_closed, fresh_completed, fresh_error) + + +async def _deaf_response( + candidate: Gateway, + key: str, + normal_model: str, + deaf_model: str, + peer: ResponsesPeer, +) -> MidResult: + created, turn_error, close = await _mid_session(candidate, key, deaf_model) + provider_closed: Final = await asyncio.to_thread(_provider_closed_within, peer, DEAF_PROVIDER_MODEL, 30) + fresh_completed, fresh_error = await _fresh_session(candidate, key, normal_model) + return MidResult(created, turn_error, close, provider_closed, fresh_completed, fresh_error) + + +async def _cap_workload( + candidate: Gateway, + key: str, + normal_model: str, + stall_model: str, + deaf_model: str, + peer: ResponsesPeer, +) -> tuple[CloseResult, ActiveResult, MidResult, MidResult]: + idle, active, mid, deaf = await asyncio.gather( + _close_at_limit(candidate, key, normal_model), + _active_session(candidate, key, normal_model, peer), + _mid_response(candidate, key, normal_model, stall_model, peer), + _deaf_response(candidate, key, normal_model, deaf_model, peer), + ) + return idle, active, mid, deaf + + +def _session_config(path: Path, seconds: int) -> Path: + source: Final = Path(__file__).resolve().parents[1] / "proxy_config.yaml" + content: Final = source.read_text() + updated: Final = content.replace( + "general_settings:\n", + f"general_settings:\n responses_websocket_session_limit_seconds: {seconds}\n", + 1, + ) + path.write_text(updated) + return path + + +def _worker_pids(log: Path) -> frozenset[int]: + return frozenset( + eventually( + lambda: WORKER_PIDS.validate_python(STARTED_WORKER.findall(log.read_text())), + lambda pids: len(pids) == 2, + seconds=30, + ) + ) + + +def _client_port(connection: websockets.ClientConnection) -> int: + return CLIENT_ADDRESS.validate_python(connection.transport.get_extra_info("sockname"))[1] + + +def _holding_worker(workers: frozenset[int], client_port: int) -> int | None: + def holds(pid: int) -> bool: + return any( + connection.status == psutil.CONN_ESTABLISHED and connection.raddr and connection.raddr.port == client_port + for connection in psutil.Process(pid).net_connections(kind="tcp") + ) + + return next((pid for pid in sorted(workers) if holds(pid)), None) + + +def _settled_on_every_worker(written_at: float) -> None: + eventually( + lambda: time.monotonic() - written_at, + lambda elapsed: elapsed >= WORKER_SYNC_SECONDS, + seconds=WORKER_SYNC_SECONDS + 5, + ) + + +def _config_field(candidate: Gateway, action: str, body: Mapping[str, JsonValue]) -> UpdateResult: + response: Final = candidate.request("POST", f"/config/field/{action}", body) + return UpdateResult(response.status_code, response.text) + + +def _update_limit(candidate: Gateway, value: JsonValue) -> UpdateResult: + return _config_field( + candidate, + "update", + {"field_name": SESSION_LIMIT_FIELD, "field_value": value, "config_type": "general_settings"}, + ) + + +def _delete_limit(candidate: Gateway) -> UpdateResult: + return _config_field(candidate, "delete", {"field_name": SESSION_LIMIT_FIELD, "config_type": "general_settings"}) + + +def _general_settings_row(database_url: str | None) -> dict[str, JsonValue] | None: + rows: Final = read_rows( + 'SELECT param_value FROM "LiteLLM_Config" WHERE param_name = %s', + ("general_settings",), + database_url=database_url, + ) + return _object(rows[0]["param_value"]) if rows else None + + +def _stored_limit(database_url: str) -> JsonValue: + row: Final = _general_settings_row(database_url) + return None if row is None else row.get(SESSION_LIMIT_FIELD) + + +def _health_status(candidate: Gateway) -> int: + try: + return candidate.request("GET", "/health/liveliness").status_code + except httpx.HTTPError: + return 0 + + +async def _update_sequence(candidate: Gateway, values: tuple[int, ...]) -> tuple[UpdateResult, ...]: + if not values: + return () + first: Final = await asyncio.to_thread(_update_limit, candidate, values[0]) + return (first, *await _update_sequence(candidate, values[1:])) + + +async def _open_capped_socket(proxy: str, key: str, path: str, workers: frozenset[int]) -> OpenedSocket: + started: Final = time.monotonic() + connection: Final = await websockets.connect( + f"{proxy}{path}", + additional_headers={"Authorization": f"Bearer {key}"}, + open_timeout=10, + ) + holder: Final = await asyncio.to_thread(_holding_worker, workers, _client_port(connection)) + return OpenedSocket(path, holder, connection, started) + + +async def _sockets_on_every_worker( + proxy: str, key: str, workers: frozenset[int], opened: tuple[OpenedSocket, ...] = () +) -> tuple[OpenedSocket, ...]: + holders: Final = frozenset(socket.worker_pid for socket in opened) + enough: Final = len(opened) >= CAPPED_SOCKETS_MINIMUM + if enough and (holders >= workers or len(opened) >= CAPPED_SOCKETS_MAXIMUM): + return opened + path: Final = CAPPED_PATHS[len(opened) % len(CAPPED_PATHS)] + next_socket: Final = await _open_capped_socket(proxy, key, path, workers) + return await _sockets_on_every_worker(proxy, key, workers, (*opened, next_socket)) + + +async def _capped_close(socket: OpenedSocket) -> CappedSocket: + try: + outcome: Final = await _wait_for_close(socket.connection, CAP_SECONDS + 15) + return CappedSocket(socket.path, socket.worker_pid, CloseResult(outcome, time.monotonic() - socket.started)) + finally: + await socket.connection.close() + + +async def _held_turn(connection: websockets.ClientConnection, model: str, accepted_at: float) -> HeldResult: + events, error = await _turn_result(connection, _create(model, f"earlier-{uuid.uuid4().hex}")) + return HeldResult(time.monotonic() - accepted_at, events, error) + + +async def _override_workload(owned: OwnedProxy, key: str, model: str, database_url: str) -> OverrideResults: + candidate: Final = owned.gateway + proxy: Final = _proxy_ws_url(candidate, "") + workers: Final = await asyncio.to_thread(_worker_pids, owned.log) + async with websockets.connect( + f"{proxy}/v1/responses", + additional_headers={"Authorization": f"Bearer {key}"}, + open_timeout=10, + ) as earlier: + accepted_at: Final = time.monotonic() + updates: Final = await _update_sequence(candidate, (7200, CAP_SECONDS, CAP_SECONDS)) + written_at: Final = time.monotonic() + stored_after_updates: Final = await asyncio.to_thread(_stored_limit, database_url) + await asyncio.to_thread(_settled_on_every_worker, written_at) + opened: Final = await _sockets_on_every_worker(proxy, key, workers) + capped: Final = await asyncio.gather(*tuple(_capped_close(socket) for socket in opened)) + earlier_result: Final = await _held_turn(earlier, model, accepted_at) + delete: Final = await asyncio.to_thread(_delete_limit, candidate) + deleted_at: Final = time.monotonic() + stored_after_delete: Final = await asyncio.to_thread(_general_settings_row, database_url) + await asyncio.to_thread(_settled_on_every_worker, deleted_at) + restored: Final = await _idle_turn( + proxy, + key, + model, + "/v1/responses", + f"restored-{uuid.uuid4().hex}", + idle_seconds=RESTORED_IDLE_SECONDS, + ) + return OverrideResults( + workers, + updates, + stored_after_updates, + tuple(capped), + earlier_result, + delete, + stored_after_delete, + restored, + ) + + +async def _kill_workload(owned: OwnedProxy, key: str, model: str) -> KillResults: + candidate: Final = owned.gateway + workers: Final = await asyncio.to_thread(_worker_pids, owned.log) + connections: Final = await asyncio.gather( + *tuple( + websockets.connect( + _proxy_ws_url(candidate, "/v1/responses"), + additional_headers={"Authorization": f"Bearer {key}"}, + open_timeout=10, + ) + for _ in range(KILLED_POOL_SIZE) + ) + ) + holders: Final = await asyncio.gather( + *tuple(asyncio.to_thread(_holding_worker, workers, _client_port(connection)) for connection in connections) + ) + held_by: Final = MappingProxyType({pid: holders.count(pid) for pid in sorted(workers)}) + victim: Final = max(sorted(workers), key=held_by.__getitem__) + psutil.Process(victim).kill() + victims: Final = tuple( + connection for connection, holder in zip(connections, holders, strict=True) if holder == victim + ) + survivors: Final = tuple( + connection for connection, holder in zip(connections, holders, strict=True) if holder != victim + ) + victim_closes, survivor_turns = await asyncio.gather( + asyncio.gather(*tuple(_wait_for_close(connection, 30) for connection in victims)), + asyncio.gather( + *tuple(_turn_result(connection, _create(model, f"survivor-{uuid.uuid4().hex}")) for connection in survivors) + ), + ) + health: Final = await asyncio.to_thread( + eventually, + lambda: _health_status(candidate), + lambda status_code: status_code == 200, + 30, + ) + fresh: Final = await _chaos_turn(candidate, key, model) + await asyncio.gather(*tuple(connection.close() for connection in survivors)) + return KillResults( + held_by, + victim, + tuple(victim_closes), + tuple(TurnOutcome(events, error) for events, error in survivor_turns), + health, + fresh, + ) + + +def _responses_body(model: str) -> JsonResponse: + message: Final[dict[str, JsonValue]] = { + "type": "message", + "id": "msg_$UNIQUE_ID", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": "responses-$UNIQUE_ID", "annotations": []}], + } + return JsonResponse( + content_type="application/json", + body={ + "id": "$UNIQUE_ID", + "object": "response", + "created_at": 1700000000, + "status": "completed", + "model": model, + "output": [message], + "usage": {"input_tokens": 1, "output_tokens": 1, "total_tokens": 2}, + }, + ) + + +def _chat_body(model: str) -> JsonResponse: + return JsonResponse( + content_type="application/json", + body={ + "id": "$UNIQUE_ID", + "object": "chat.completion", + "created": 1700000000, + "model": model, + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": "chat-$UNIQUE_ID"}, + "finish_reason": "stop", + } + ], + "usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2}, + }, + ) + + +def _responses_spec(model: str) -> tuple[str, dict[str, JsonValue]]: + return "/v1/responses", {"model": model, "input": f"responses-{uuid.uuid4().hex}"} + + +def _chat_spec(model: str) -> tuple[str, dict[str, JsonValue]]: + return "/v1/chat/completions", { + "model": model, + "messages": [{"role": "user", "content": f"chat-{uuid.uuid4().hex}"}], + } + + +async def _burst_request( + client: httpx.AsyncClient, + key: str, + path: str, + body: dict[str, JsonValue], +) -> BurstResult: + try: + response: Final = await client.post(path, json=body, headers={"Authorization": f"Bearer {key}"}) + try: + parsed: Final = JSON_OBJECT.validate_json(response.content) + except ValidationError as error: + return BurstResult(path, response.status_code, None, str(error)) + return BurstResult(path, response.status_code, parsed, None) + except httpx.HTTPError as error: + return BurstResult(path, None, None, str(error)) + + +async def _chaos_workload( + owned: OwnedProxy, + key: str, + scripted_model: str, + peer_model: str, +) -> ChaosResults: + candidate: Final = owned.gateway + base: Final = str(candidate.client.base_url) + specs: Final = tuple( + _responses_spec(scripted_model) if index % 2 == 0 else _chat_spec(scripted_model) for index in range(20) + ) + connections: Final = await asyncio.gather( + *tuple( + websockets.connect( + _proxy_ws_url(candidate, "/v1/responses"), + additional_headers={"Authorization": f"Bearer {key}"}, + open_timeout=10, + ) + for _ in range(30) + ) + ) + async with httpx.AsyncClient(base_url=base, timeout=30, trust_env=False) as client: + results: Final = await asyncio.gather(*tuple(_burst_request(client, key, path, body) for path, body in specs)) + tuple(connection.transport.abort() for connection in connections) + health: Final = candidate.request("GET", "/health/liveliness").status_code + websocket_completed: Final = await _chaos_turn(candidate, key, peer_model) + kill: Final = await _kill_workload(owned, key, peer_model) + return ChaosResults(tuple(results), health, websocket_completed, kill) + + +async def _chaos_turn(candidate: Gateway, key: str, model: str) -> bool: + async with websockets.connect( + _proxy_ws_url(candidate, "/v1/responses"), + additional_headers={"Authorization": f"Bearer {key}"}, + open_timeout=10, + ) as connection: + events: Final = await _turn(connection, _create(model, f"chaos-peer-{uuid.uuid4().hex}")) + return _completed_text(events) == PEER_TEXT + + +@pytest.fixture(scope="module") +def cert(tmp_path_factory: pytest.TempPathFactory) -> tuple[Path, Path]: + return write_self_signed_cert(tmp_path_factory.mktemp("responses-ws-session-limit")) + + +@pytest.fixture(scope="module") +def default_results( + cert: tuple[Path, Path], + tmp_path_factory: pytest.TempPathFactory, +) -> Iterator[DefaultResults]: + with responses_peer(cert) as peer, gateway_from_environment() as base: + with ( + owned_proxy( + base, + tmp_path_factory.mktemp("responses-ws-default"), + {"SSL_CERT_FILE": str(cert[0])}, + workers=2, + ) as candidate, + candidate.scenario() as scenario, + ): + model: Final = scenario.model(model=f"openai/{PROVIDER_MODEL}", api_base=peer.url) + key: Final = scenario.key(models=[model]) + pool, query, rejected, negotiated = asyncio.run(_default_workload(candidate, key, model)) + provider: Final = _snapshots(peer, 10) + yield DefaultResults(pool, query, rejected, negotiated, provider) + + +@pytest.fixture(scope="module") +def cap_results( + cert: tuple[Path, Path], + tmp_path_factory: pytest.TempPathFactory, +) -> Iterator[CapResults]: + with responses_peer(cert) as peer, gateway_from_environment() as base: + config: Final = _session_config(tmp_path_factory.mktemp("responses-ws-cap") / "config.yaml", 60) + with ( + owned_proxy( + base, + tmp_path_factory.mktemp("responses-ws-cap-proxy"), + {"SSL_CERT_FILE": str(cert[0])}, + config=config, + workers=2, + ) as candidate, + candidate.scenario() as scenario, + ): + normal: Final = scenario.model(model=f"openai/{PROVIDER_MODEL}", api_base=peer.url) + stall: Final = scenario.model(model=f"openai/{STALL_PROVIDER_MODEL}", api_base=peer.url) + deaf: Final = scenario.model(model=f"openai/{DEAF_PROVIDER_MODEL}", api_base=peer.url) + key: Final = scenario.key(models=[normal, stall, deaf]) + idle, active, mid, deaf_result = asyncio.run(_cap_workload(candidate, key, normal, stall, deaf, peer)) + provider: Final = _available_snapshots(peer) + yield CapResults(idle, active, mid, deaf_result, provider) + + +@pytest.fixture(scope="module") +def invalid_results( + cert: tuple[Path, Path], + tmp_path_factory: pytest.TempPathFactory, +) -> Iterator[InvalidResult]: + with responses_peer(cert) as peer, gateway_from_environment() as base: + config: Final = _session_config(tmp_path_factory.mktemp("responses-ws-invalid") / "config.yaml", 30) + with owned_proxy_process( + base, + tmp_path_factory.mktemp("responses-ws-invalid-proxy"), + {"SSL_CERT_FILE": str(cert[0])}, + config=config, + workers=2, + ) as owned: + with owned.gateway.scenario() as scenario: + model: Final = scenario.model(model=f"openai/{PROVIDER_MODEL}", api_base=peer.url) + key: Final = scenario.key(models=[model]) + session: Final = asyncio.run( + _idle_turn( + _proxy_ws_url(owned.gateway, ""), + key, + model, + "/v1/responses", + f"invalid-{uuid.uuid4().hex}", + ) + ) + warning: Final = "invalid general_settings.responses_websocket_session_limit_seconds=30" + yield InvalidResult(session, warning in owned.log.read_text()) + + +@pytest.fixture(scope="module") +def chaos_results( + cert: tuple[Path, Path], + tmp_path_factory: pytest.TempPathFactory, +) -> Iterator[ChaosResults]: + with responses_peer(cert) as peer, gateway_from_environment() as base: + with ( + owned_proxy_process( + base, + tmp_path_factory.mktemp("responses-ws-chaos"), + {"SSL_CERT_FILE": str(cert[0])}, + workers=2, + ) as owned, + owned.gateway.scenario() as scenario, + ): + scenario_id: Final = f"responses-chaos-{uuid.uuid4().hex}" + scripted: Final = register_scenario( + scenario_id, + RoutedResponse( + content_type="application/x-routed", + routes={ + "POST /responses": _responses_body("scripted-responses-model"), + "POST /chat/completions": _chat_body("scripted-chat-model"), + }, + ), + control_url=owned.gateway.upstream_url, + ) + scenario.cleanups.callback(delete_scenario, scripted) + scripted_model: Final = scenario.model( + model="openai/scripted-responses-model", + api_base=scripted.api_base(), + ) + peer_model: Final = scenario.model(model=f"openai/{PROVIDER_MODEL}", api_base=peer.url) + key: Final = scenario.key(models=[scripted_model, peer_model]) + yield asyncio.run(_chaos_workload(owned, key, scripted_model, peer_model)) + + +@pytest.fixture(scope="module") +def override_results( + cert: tuple[Path, Path], + tmp_path_factory: pytest.TempPathFactory, +) -> Iterator[OverrideResults]: + with responses_peer(cert) as peer, gateway_from_environment() as base, scratch_database() as database_url: + with owned_proxy_process( + base, + tmp_path_factory.mktemp("responses-ws-override"), + { + "SSL_CERT_FILE": str(cert[0]), + "DATABASE_URL": database_url, + "PROXY_CONFIG_RELOAD_INTERVAL_SECONDS": str(RELOAD_INTERVAL_SECONDS), + }, + remove_environment=("DATABASE_URL_READ_REPLICA",), + workers=2, + ) as owned: + with owned.gateway.scenario() as scenario: + model: Final = scenario.model(model=f"openai/{PROVIDER_MODEL}", api_base=peer.url) + key: Final = scenario.key(models=[model]) + yield asyncio.run(_override_workload(owned, key, model, database_url)) + + +def test_default_config_idle_pool_stays_open_and_completes(default_results: DefaultResults) -> None: + assert all(result.idle.timed_out for result in default_results.pool), default_results.pool + assert all(result.error is None for result in default_results.pool), default_results.pool + assert all(_completed_text(result.events) == PEER_TEXT for result in default_results.pool) + texts: Final = frozenset(result.text for result in default_results.pool) + provider: Final = tuple(snapshot for snapshot in default_results.provider if snapshot.frames) + assert len(provider) == 10, default_results.provider + assert all(snapshot.path == f"/v1/responses?model={PROVIDER_MODEL}" for snapshot in provider) + frames: Final = tuple(snapshot.frames[0] for snapshot in provider) + expected_texts: Final = texts | {default_results.query.text, default_results.subprotocol.text} + assert frozenset(_frame_text(frame) for frame in frames) == expected_texts + assert all(frame.get("model") == PROVIDER_MODEL for frame in frames) + + +def test_default_config_query_model_socket_stays_open_and_completes(default_results: DefaultResults) -> None: + assert default_results.query.idle.timed_out, default_results.query + assert default_results.query.error is None, default_results.query + assert _completed_text(default_results.query.events) == PEER_TEXT + assert any( + default_results.query.text == _frame_text(snapshot.frames[0]) + for snapshot in default_results.provider + if snapshot.frames + ) + + +def test_delayed_first_frame_model_auth_matches_immediate_rejection( + cert: tuple[Path, Path], + tmp_path_factory: pytest.TempPathFactory, +) -> None: + with responses_peer(cert) as peer, gateway_from_environment() as base: + with ( + owned_proxy( + base, + tmp_path_factory.mktemp("responses-ws-auth"), + {"SSL_CERT_FILE": str(cert[0])}, + workers=2, + ) as candidate, + candidate.scenario() as scenario, + ): + model: Final = scenario.model(model=f"openai/{PROVIDER_MODEL}", api_base=peer.url) + key: Final = scenario.key(models=["not-authorized-model"]) + results: Final = asyncio.run(_auth_workload(candidate, key, model, peer)) + assert results.immediate.error is None, results.immediate + assert results.delayed.idle.timed_out, results.delayed + assert results.delayed.error is None, results.delayed + assert results.immediate.frame is not None, results.immediate + assert results.immediate.frame.get("type") == "error", results.immediate + rejection: Final = _object(results.immediate.frame["error"]) + assert rejection.get("type") == "invalid_request_error", results.immediate + assert results.immediate.close_code == 1008, results.immediate + assert results.immediate.close_reason == "Pre-call error", results.immediate + assert results.immediate.frame == results.delayed.frame + assert results.immediate.close_code == results.delayed.close_code + assert results.immediate.close_reason == results.delayed.close_reason + assert results.provider_connections == 0, results + + +def test_session_cap_closes_never_started_socket(cap_results: CapResults) -> None: + assert cap_results.idle.outcome.closed, cap_results.idle + assert cap_results.idle.outcome.close_code == 1000, cap_results.idle + assert cap_results.idle.outcome.close_reason == "Session duration limit reached", cap_results.idle + assert 59 <= cap_results.idle.elapsed <= 75, cap_results.idle + + +def test_session_cap_closes_active_socket_and_provider(cap_results: CapResults) -> None: + assert cap_results.active.turn_error is None, cap_results.active + assert _completed_text(cap_results.active.turn) == PEER_TEXT + assert cap_results.active.close.outcome.closed, cap_results.active + assert cap_results.active.close.outcome.close_code == 1000, cap_results.active + assert cap_results.active.close.outcome.close_reason == "Session duration limit reached", cap_results.active + assert 59 <= cap_results.active.close.elapsed <= 75, cap_results.active + assert cap_results.active.provider_closed, cap_results.active + assert all(snapshot.closed for snapshot in cap_results.provider), cap_results.provider + + +def test_session_cap_closes_mid_response_and_allows_new_session(cap_results: CapResults) -> None: + assert cap_results.mid.created, cap_results.mid + assert cap_results.mid.turn_error is None, cap_results.mid + assert cap_results.mid.close.outcome.closed, cap_results.mid + assert cap_results.mid.close.outcome.close_code == 1000, cap_results.mid + assert cap_results.mid.close.outcome.close_reason == "Session duration limit reached", cap_results.mid + assert 59 <= cap_results.mid.close.elapsed <= 75, cap_results.mid + assert cap_results.mid.provider_closed, cap_results.mid + assert cap_results.mid.fresh_error is None, cap_results.mid + assert cap_results.mid.fresh_completed, cap_results.mid + + +def test_session_cap_closes_client_promptly_when_provider_ignores_close(cap_results: CapResults) -> None: + assert cap_results.deaf.created, cap_results.deaf + assert cap_results.deaf.turn_error is None, cap_results.deaf + assert cap_results.deaf.close.outcome.closed, cap_results.deaf + assert cap_results.deaf.close.outcome.close_code == 1000, cap_results.deaf + assert cap_results.deaf.close.outcome.close_reason == "Session duration limit reached", cap_results.deaf + assert 59 <= cap_results.deaf.close.elapsed <= 63, cap_results.deaf + assert cap_results.deaf.provider_closed, cap_results.deaf + assert cap_results.deaf.fresh_error is None, cap_results.deaf + assert cap_results.deaf.fresh_completed, cap_results.deaf + + +def test_invalid_session_cap_falls_back_to_default(invalid_results: InvalidResult) -> None: + assert invalid_results.session.idle.timed_out, invalid_results.session + assert invalid_results.session.error is None, invalid_results.session + assert _completed_text(invalid_results.session.events) == PEER_TEXT + assert invalid_results.warning_found + + +def _matches_scripted_body(result: BurstResult) -> bool: + if result.body is None: + return False + if result.path == "/v1/responses": + output: Final = _list(result.body["output"]) + responses_message: Final = _object(output[0]) + message_id: Final = _string(responses_message["id"]) + if not message_id.startswith("msg_"): + return False + content: Final = _list(responses_message["content"]) + return _string(_object(content[0])["text"]) == f"responses-{message_id.removeprefix('msg_')}" + identity: Final = _string(result.body["id"]) + choices: Final = _list(result.body["choices"]) + chat_message: Final = _object(_object(choices[0])["message"]) + return _string(chat_message["content"]) == f"chat-{identity}" + + +def test_aborted_idle_pool_preserves_http_health_and_new_websocket(chaos_results: ChaosResults) -> None: + results: Final = chaos_results.burst + assert len(results) == 20 + assert all(result.status_code == 200 and result.error is None for result in results), results + identities: Final = tuple(_string(_object(result.body)["id"]) for result in results if result.body is not None) + assert len(set(identities)) == 20, identities + assert all(_matches_scripted_body(result) for result in results), results + assert chaos_results.health_status == 200 + assert chaos_results.websocket_completed + + +def test_rejected_key_fails_the_handshake_with_403(default_results: DefaultResults) -> None: + assert default_results.rejected == HandshakeResult(403, None), default_results.rejected + + +def test_requested_subprotocol_is_accepted_and_the_turn_completes(default_results: DefaultResults) -> None: + negotiated: Final = default_results.subprotocol + assert negotiated.error is None, negotiated + assert negotiated.negotiated == SUBPROTOCOLS[0], negotiated + assert _completed_text(negotiated.events) == PEER_TEXT + + +def test_db_override_caps_new_sessions_on_every_worker_without_restart(override_results: OverrideResults) -> None: + assert len(override_results.workers) == 2, override_results.workers + capped: Final = override_results.capped + assert CAPPED_SOCKETS_MINIMUM <= len(capped) <= CAPPED_SOCKETS_MAXIMUM, capped + assert tuple(socket.path for socket in capped) == tuple( + CAPPED_PATHS[i % len(CAPPED_PATHS)] for i in range(len(capped)) + ) + assert all(socket.close.outcome.closed and socket.close.outcome.close_code == 1000 for socket in capped), capped + assert all(socket.close.outcome.close_reason == LIMIT_CLOSE_REASON for socket in capped), capped + assert all(CAP_SECONDS - 1 <= socket.close.elapsed <= CAP_SECONDS + 15 for socket in capped), capped + assert frozenset(socket.worker_pid for socket in capped) == override_results.workers, capped + + +def test_db_override_leaves_sessions_accepted_before_it_alone(override_results: OverrideResults) -> None: + earlier: Final = override_results.earlier + assert earlier.error is None, earlier + assert earlier.held_seconds > CAP_SECONDS, earlier + assert _completed_text(earlier.events) == PEER_TEXT + + +def test_db_override_accepts_the_range_bounds_and_repeats(override_results: OverrideResults) -> None: + updates: Final = override_results.updates + assert tuple(update.status_code for update in updates) == (200, 200, 200), updates + assert override_results.stored_after_updates == CAP_SECONDS + + +def test_deleting_the_db_override_restores_the_default(override_results: OverrideResults) -> None: + assert override_results.delete.status_code == 200, override_results.delete + assert override_results.stored_after_delete is not None + assert SESSION_LIMIT_FIELD not in override_results.stored_after_delete + restored: Final = override_results.restored + assert restored.idle.timed_out, restored + assert restored.error is None, restored + assert _completed_text(restored.events) == PEER_TEXT + + +@pytest.mark.parametrize( + ("value", "kind"), + ( + pytest.param(30, "int", id="below-minimum"), + pytest.param(7201, "int", id="above-maximum"), + pytest.param("abc", "str", id="text"), + pytest.param("", "str", id="empty-string"), + pytest.param("x" * 5000, "str", id="five-kilobyte-string"), + pytest.param([], "list", id="list"), + pytest.param({}, "dict", id="object"), + pytest.param(None, "NoneType", id="null"), + ), +) +def test_out_of_range_or_wrong_type_update_is_refused(gateway: Gateway, value: JsonValue, kind: str) -> None: + before: Final = _general_settings_row(None) + refused: Final = _update_limit(gateway, value) + assert refused.status_code == 400, refused + assert json.loads(refused.body) == {"detail": {"error": f"Invalid type of field value= passed in."}} + assert _general_settings_row(None) == before + + +def test_killing_one_worker_drops_only_its_sockets_and_the_proxy_keeps_serving(chaos_results: ChaosResults) -> None: + kill: Final = chaos_results.kill + assert sum(kill.held_by.values()) == KILLED_POOL_SIZE, kill.held_by + assert len(kill.victim_closes) == kill.held_by[kill.victim] >= 1, kill + assert all(close.closed and close.close_code is None for close in kill.victim_closes), kill.victim_closes + assert all(turn.error is None for turn in kill.survivor_turns), kill.survivor_turns + assert all(_completed_text(turn.events) == PEER_TEXT for turn in kill.survivor_turns), kill.survivor_turns + assert kill.health_status == 200 + assert kill.fresh_completed diff --git a/tests/integration/providers/test_vertex_realtime_multi_region_host.py b/tests/integration/providers/test_vertex_realtime_multi_region_host.py new file mode 100644 index 00000000000..df279aa2ced --- /dev/null +++ b/tests/integration/providers/test_vertex_realtime_multi_region_host.py @@ -0,0 +1,468 @@ +"""Vertex AI Live sessions resolve their host through the shared Vertex resolver. + +``VertexAIRealtimeConfig.get_complete_url`` dials ``aiplatform.{us,eu}.rep.googleapis.com`` for the +multi-regions, ``{region}-aiplatform.googleapis.com`` for a region and ``aiplatform.googleapis.com`` for +``global``, and a malformed location fails the shared validator before any socket opens. Every row here +runs against the scripted upstream on 127.0.0.1: an ``api_base`` override pins the path and the host the +upstream sees, the setup frame it receives names the location the proxy resolved, and the malformed rows +pin the validator's answer, which the proxy gives without dialing anything. +""" + +from __future__ import annotations + +import asyncio +import json +import uuid +from collections.abc import AsyncIterator, Callable +from dataclasses import dataclass +from hashlib import sha256 +from itertools import chain, repeat +from pathlib import Path +from typing import Final + +import httpx +import pytest +import websockets +from pydantic import JsonValue +from websockets.asyncio.client import ClientConnection +from websockets.exceptions import ConnectionClosed + +from tests.integration._support.client import JSON_OBJECT, Gateway, Scenario, eventually, object_value, string_value +from tests.integration._support.database import read_rows +from tests.integration._support.process import owned_upstream +from tests.integration._support.upstream import GEMINI_LIVE_PATH, ScenarioHandle, delete_scenario, register_scenario +from tests.integration.cost_calculation.cost_tracking_case import RealtimeResponse + +RecordProperty = Callable[[str, object], None] + +pytestmark: Final = pytest.mark.timeout(180) + +LIVE_MODEL: Final = "gemini-3.8-live" +LOCATIONS: Final = ("us", "eu", "us-central1", "global") +MALFORMED_LOCATIONS: Final = ("US", "us/evil") +DEFAULT_LOCATION: Final = "us-central1" +INVALID_LOCATION: Final = "Invalid vertex_location format" +HANDSHAKE_REFUSED: Final = "Upstream realtime handshake rejected with HTTP 403" +INTERNAL_CLOSE: Final = 1011 +REFUSAL_CLOSE: Final = 1008 +SESSIONS_PER_LOCATION: Final = 3 +OUTAGE_BURST: Final = 20 +INPUT_TOKENS: Final = 7 +OUTPUT_TOKENS: Final = 5 +CONVERGENCE_SECONDS: Final = 30 +SPEND_SQL: Final = 'SELECT call_type, prompt_tokens, completion_tokens FROM "LiteLLM_SpendLogs" WHERE api_key = %s' + + +@dataclass(frozen=True, slots=True) +class Session: + events: tuple[dict[str, JsonValue], ...] + + @property + def types(self) -> tuple[str, ...]: + return tuple(string_value(event["type"]) for event in self.events) + + @property + def close_code(self) -> int | None: + closes: Final = tuple(event for event in self.events if event["type"] == "closed") + return _integer(closes[-1]["code"]) if closes else None + + @property + def session_model(self) -> str: + return string_value(object_value(self.events[0]["session"])["model"]) + + @property + def response_ids(self) -> tuple[str, ...]: + done: Final = tuple(event for event in self.events if event["type"] == "response.done") + return tuple(string_value(object_value(event["response"])["id"]) for event in done) + + @property + def error_messages(self) -> tuple[str, ...]: + errors: Final = tuple(event for event in self.events if event["type"] == "error") + return tuple(string_value(object_value(event["error"])["message"]) for event in errors) + + @property + def text(self) -> str: + deltas: Final = tuple(event for event in self.events if event["type"] == "response.output_text.delta") + return "".join(string_value(event["delta"]) for event in deltas) + + def completed_turn(self, scenario_id: str) -> bool: + return ( + self.types[0] == "session.created" + and self.types[-1] == "response.done" + and self.session_model == LIVE_MODEL + and self.text == f"scripted {scenario_id}" + and len(self.response_ids) == 1 + ) + + +def _integer(value: JsonValue) -> int: + assert isinstance(value, int), value + return value + + +def _ws_base(http_url: str) -> str: + return http_url.replace("https://", "wss://").replace("http://", "ws://") + + +def _proxy_url(gateway: Gateway) -> str: + return str(gateway.client.base_url).rstrip("/") + + +def _upstream_url(gateway: Gateway) -> str: + return gateway.upstream_url.rstrip("/") + + +def _scripted_turn() -> RealtimeResponse: + return RealtimeResponse( + content_type="application/x-realtime", + events=( + {"serverContent": {"modelTurn": {"parts": [{"text": "scripted $REQUEST_ID"}]}}}, + { + "serverContent": {"turnComplete": True}, + "usageMetadata": { + "promptTokenCount": INPUT_TOKENS, + "responseTokenCount": OUTPUT_TOKENS, + "totalTokenCount": INPUT_TOKENS + OUTPUT_TOKENS, + }, + }, + ), + ) + + +def _scripted(scenario: Scenario) -> ScenarioHandle: + handle: Final = register_scenario(f"vertex-live-{uuid.uuid4().hex[:12]}", _scripted_turn()) + scenario.cleanups.callback(delete_scenario, handle) + return handle + + +def _credentials(upstream_url: str) -> str: + return json.dumps( + { + "type": "external_account", + "audience": "synthetic-vertex-audience", + "subject_token_type": "urn:ietf:params:oauth:token-type:jwt", + "token_url": f"{upstream_url}/_oauth/token", + "credential_source": {"url": f"{upstream_url}/health"}, + } + ) + + +def _deployment( + gateway: Gateway, scenario: Scenario, project: str, *, location: str | None, api_base: str | None +) -> str: + name: Final = f"vertex-live-{uuid.uuid4().hex[:12]}" + litellm_params: Final[dict[str, JsonValue]] = { + "model": f"vertex_ai/{LIVE_MODEL}", + "vertex_project": project, + "vertex_credentials": _credentials(_upstream_url(gateway)), + **({} if location is None else {"vertex_location": location}), + **({} if api_base is None else {"api_base": api_base}), + } + created: Final = gateway.post( + "/model/new", {"model_name": name, "litellm_params": litellm_params, "model_info": {"mode": "realtime"}} + ) + scenario.cleanups.callback(scenario.delete_model, string_value(object_value(created["model_info"])["id"])) + return name + + +def _model_path(project: str, location: str) -> str: + return f"projects/{project}/locations/{location}/publishers/google/models/{LIVE_MODEL}" + + +def _user_turn() -> str: + return json.dumps( + { + "type": "conversation.item.create", + "item": { + "type": "message", + "role": "user", + "content": [{"type": "input_text", "text": "say the scripted line"}], + }, + } + ) + + +def _close_code(closed: ConnectionClosed) -> int: + return 1006 if closed.rcvd is None else closed.rcvd.code + + +async def _frames(socket: ClientConnection, turns: int) -> AsyncIterator[dict[str, JsonValue]]: + try: + first: Final = JSON_OBJECT.validate_json(await socket.recv()) + yield first + if first.get("type") != "session.created": + async for message in socket: + yield JSON_OBJECT.validate_json(message) + return + for _ in range(turns): + await socket.send(_user_turn()) + async for message in socket: + event: Final = JSON_OBJECT.validate_json(message) + yield event + if event.get("type") == "response.done": + break + except ConnectionClosed as closed: + yield {"type": "closed", "code": _close_code(closed)} + + +async def _collect(socket: ClientConnection, turns: int) -> tuple[dict[str, JsonValue], ...]: + return tuple([frame async for frame in _frames(socket, turns)]) + + +async def _session(ws_base: str, model: str, key: str, *, turns: int = 1) -> Session: + headers: Final = {"Authorization": f"Bearer {key}"} + async with websockets.connect(f"{ws_base}/v1/realtime?model={model}", additional_headers=headers) as socket: + return Session(await asyncio.wait_for(_collect(socket, turns), 60)) + + +def _run(gateway: Gateway, model: str, key: str, *, turns: int = 1) -> Session: + return asyncio.run(_session(_ws_base(_proxy_url(gateway)), model, key, turns=turns)) + + +def _observations(gateway: Gateway) -> tuple[dict[str, JsonValue], ...]: + with httpx.Client(base_url=gateway.upstream_url, timeout=5, trust_env=False) as upstream: + return tuple(map(object_value, upstream.get("/__observations").json()["requests"])) + + +def _for_project(observed: tuple[dict[str, JsonValue], ...], project: str) -> tuple[dict[str, JsonValue], ...]: + return tuple(request for request in observed if request["api_key"] == project) + + +def _upgrade_hosts(observed: tuple[dict[str, JsonValue], ...]) -> tuple[str, ...]: + upgrades: Final = tuple(request for request in observed if request["method"] == "WEBSOCKET") + assert all(request["path"] == GEMINI_LIVE_PATH for request in upgrades), upgrades + return tuple(string_value(object_value(request["body"])["host"]) for request in upgrades) + + +def _setup_models(observed: tuple[dict[str, JsonValue], ...]) -> tuple[str, ...]: + frames: Final = tuple( + object_value(request["body"]) for request in observed if request["method"] == "WEBSOCKET_FRAME" + ) + return tuple(string_value(object_value(frame["setup"])["model"]) for frame in frames if "setup" in frame) + + +def _authority(gateway: Gateway) -> str: + return _upstream_url(gateway).removeprefix("http://") + + +def _spend_rows(key: str, count: int) -> list[dict[str, JsonValue]]: + return eventually( + lambda: read_rows(SPEND_SQL, (sha256(key.encode()).hexdigest(),)), + lambda rows: len(rows) == count, + seconds=70, + ) + + +def _billed(rows: list[dict[str, JsonValue]]) -> list[tuple[JsonValue, JsonValue, JsonValue]]: + return [(row["call_type"], row["prompt_tokens"], row["completion_tokens"]) for row in rows] + + +def _health(gateway: Gateway, model: str) -> httpx.Response: + return gateway.request("GET", "/health", params={"model": model}) + + +def _probed_count(response: httpx.Response) -> int: + body: Final = JSON_OBJECT.validate_json(response.content) + return _integer(body.get("healthy_count", 0)) + _integer(body.get("unhealthy_count", 0)) + + +def _converged_health(gateway: Gateway, model: str) -> httpx.Response: + return eventually( + lambda: _health(gateway, model), lambda response: _probed_count(response) == 1, seconds=CONVERGENCE_SECONDS + ) + + +def _assert_scripted_turn_reached_upstream(gateway: Gateway, project: str, location: str, session: Session) -> None: + assert session.completed_turn(project), session + observed: Final = _for_project(_observations(gateway), project) + assert _upgrade_hosts(observed) == (_authority(gateway),), observed + assert _setup_models(observed) == (_model_path(project, location),), observed + + +@pytest.mark.parametrize("location", LOCATIONS) +def test_api_base_override_completes_a_turn_for_every_location(gateway: Gateway, location: str) -> None: + with gateway.scenario() as scenario: + handle: Final = _scripted(scenario) + key: Final = scenario.key() + model: Final = _deployment( + gateway, scenario, handle.scenario_id, location=location, api_base=_upstream_url(gateway) + ) + session: Final = _run(gateway, model, key) + _assert_scripted_turn_reached_upstream(gateway, handle.scenario_id, location, session) + assert _billed(_spend_rows(key, 1)) == [("_arealtime", INPUT_TOKENS, OUTPUT_TOKENS)] + + +@pytest.mark.parametrize("location", MALFORMED_LOCATIONS) +def test_malformed_location_is_refused_by_the_validator_before_any_dial(gateway: Gateway, location: str) -> None: + with gateway.scenario() as scenario: + handle: Final = _scripted(scenario) + key: Final = scenario.key() + malformed: Final = _deployment(gateway, scenario, handle.scenario_id, location=location, api_base=None) + refused: Final = _run(gateway, malformed, key) + assert refused.types == ("error", "closed"), refused + assert refused.error_messages == (INVALID_LOCATION,), refused + assert refused.close_code == INTERNAL_CLOSE, refused + model: Final = _deployment( + gateway, scenario, handle.scenario_id, location="us", api_base=_upstream_url(gateway) + ) + session: Final = _run(gateway, model, key) + _assert_scripted_turn_reached_upstream(gateway, handle.scenario_id, "us", session) + + +@pytest.mark.parametrize("location", [None, ""], ids=["omitted", "empty"]) +def test_missing_location_defaults_to_us_central1(gateway: Gateway, location: str | None) -> None: + with gateway.scenario() as scenario: + handle: Final = _scripted(scenario) + key: Final = scenario.key() + model: Final = _deployment( + gateway, scenario, handle.scenario_id, location=location, api_base=_upstream_url(gateway) + ) + session: Final = _run(gateway, model, key) + _assert_scripted_turn_reached_upstream(gateway, handle.scenario_id, DEFAULT_LOCATION, session) + + +@pytest.mark.parametrize("location", ["us", "eu"]) +def test_health_check_handshakes_with_the_override_host(gateway: Gateway, location: str) -> None: + with gateway.scenario() as scenario: + handle: Final = _scripted(scenario) + model: Final = _deployment( + gateway, scenario, handle.scenario_id, location=location, api_base=_upstream_url(gateway) + ) + health: Final = _converged_health(gateway, model) + assert health.status_code == 200, health.text + body: Final = JSON_OBJECT.validate_json(health.content) + assert (body["healthy_count"], body["unhealthy_count"]) == (1, 0), health.text + observed: Final = _for_project(_observations(gateway), handle.scenario_id) + assert _upgrade_hosts(observed) == (_authority(gateway),), observed + assert _setup_models(observed) == (), observed + + +def test_health_check_reports_a_malformed_location_without_dialing(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + handle: Final = _scripted(scenario) + model: Final = _deployment(gateway, scenario, handle.scenario_id, location="US", api_base=None) + health: Final = _converged_health(gateway, model) + assert health.status_code == 503, health.text + body: Final = JSON_OBJECT.validate_json(health.content) + assert (body["healthy_count"], body["unhealthy_count"]) == (0, 1), health.text + unhealthy: Final = body["unhealthy_endpoints"] + assert isinstance(unhealthy, list) and len(unhealthy) == 1, health.text + assert INVALID_LOCATION in string_value(object_value(unhealthy[0])["error"]), health.text + assert _for_project(_observations(gateway), handle.scenario_id) == (), "the upstream saw a dial" + + +def test_upstream_handshake_refusal_reaches_the_client_and_the_next_session_connects(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + key: Final = scenario.key() + unregistered: Final = f"vertex-live-unknown-{uuid.uuid4().hex[:12]}" + refused_model: Final = _deployment( + gateway, scenario, unregistered, location="us", api_base=_upstream_url(gateway) + ) + refused: Final = _run(gateway, refused_model, key) + assert refused.types == ("error", "closed"), refused + assert refused.error_messages == (HANDSHAKE_REFUSED,), refused + assert refused.close_code == REFUSAL_CLOSE, refused + assert _upgrade_hosts(_for_project(_observations(gateway), unregistered)) == (_authority(gateway),) + handle: Final = _scripted(scenario) + model: Final = _deployment( + gateway, scenario, handle.scenario_id, location="us", api_base=_upstream_url(gateway) + ) + session: Final = _run(gateway, model, key) + _assert_scripted_turn_reached_upstream(gateway, handle.scenario_id, "us", session) + chat: Final = gateway.chat(scenario.model(), key=key) + assert string_value(chat["object"]) == "chat.completion", chat + + +async def _burst(ws_base: str, models: tuple[str, ...], key: str) -> tuple[Session, ...]: + return tuple(await asyncio.gather(*(_session(ws_base, model, key) for model in models))) + + +def test_concurrent_sessions_across_locations_each_reach_the_upstream_once(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + key: Final = scenario.key() + handles: Final = {location: _scripted(scenario) for location in LOCATIONS} + models: Final = { + location: _deployment( + gateway, scenario, handles[location].scenario_id, location=location, api_base=_upstream_url(gateway) + ) + for location in LOCATIONS + } + order: Final = tuple(chain.from_iterable(repeat(location, SESSIONS_PER_LOCATION) for location in LOCATIONS)) + sessions: Final = asyncio.run( + _burst(_ws_base(_proxy_url(gateway)), tuple(models[location] for location in order), key) + ) + assert all( + session.completed_turn(handles[location].scenario_id) for location, session in zip(order, sessions) + ), sessions + assert len({session.response_ids[0] for session in sessions}) == len(order), sessions + observed: Final = _observations(gateway) + for location, handle in handles.items(): + mine: Final = _for_project(observed, handle.scenario_id) + assert _upgrade_hosts(mine) == (_authority(gateway),) * SESSIONS_PER_LOCATION, mine + assert _setup_models(mine) == (_model_path(handle.scenario_id, location),) * SESSIONS_PER_LOCATION, mine + assert _billed(_spend_rows(key, len(order))) == [("_arealtime", INPUT_TOKENS, OUTPUT_TOKENS)] * len(order) + + +async def _frames_until_closed(socket: ClientConnection) -> AsyncIterator[dict[str, JsonValue]]: + try: + async for message in socket: + yield JSON_OBJECT.validate_json(message) + except ConnectionClosed as closed: + yield {"type": "closed", "code": _close_code(closed)} + + +async def _hold_until_closed(ws_base: str, model: str, key: str, opened: asyncio.Queue[str]) -> Session: + headers: Final = {"Authorization": f"Bearer {key}"} + async with websockets.connect(f"{ws_base}/v1/realtime?model={model}", additional_headers=headers) as socket: + created: Final = JSON_OBJECT.validate_json(await socket.recv()) + assert created.get("type") == "session.created", created + await opened.put(string_value(object_value(created["session"])["id"])) + return Session(tuple([frame async for frame in _frames_until_closed(socket)])) + + +def _relays_the_upstream_close(session: Session) -> bool: + return f"upstream websocket closed with code {session.close_code}" in session.error_messages[0] + + +async def _drain(opened: asyncio.Queue[str], count: int) -> tuple[str, ...]: + return tuple([await opened.get() for _ in range(count)]) + + +async def _burst_through_outage( + ws_base: str, proxy_url: str, model: str, key: str, stop_upstream: Callable[[], None] +) -> tuple[Session, ...]: + opened: Final[asyncio.Queue[str]] = asyncio.Queue() + holders: Final = tuple( + asyncio.ensure_future(_hold_until_closed(ws_base, model, key, opened)) for _ in range(OUTAGE_BURST) + ) + opened_sessions: Final = await asyncio.wait_for(_drain(opened, OUTAGE_BURST), 60) + assert len(opened_sessions) == OUTAGE_BURST, opened_sessions + await asyncio.to_thread(stop_upstream) + async with httpx.AsyncClient(base_url=proxy_url, timeout=15, trust_env=False) as client: + liveliness: Final = await client.get("/health/liveliness") + assert liveliness.status_code == 200, liveliness.text + return tuple(await asyncio.wait_for(asyncio.gather(*holders), 90)) + + +@pytest.mark.timeout(240) +def test_upstream_outage_closes_every_open_session_and_recovers( + gateway: Gateway, tmp_path: Path, record_property: RecordProperty +) -> None: + with gateway.scenario() as scenario, owned_upstream(tmp_path) as slot: + project: Final = f"vertex-live-outage-{uuid.uuid4().hex[:12]}" + register_scenario(project, _scripted_turn(), control_url=slot.url) + key: Final = scenario.key() + model: Final = _deployment(gateway, scenario, project, location="us", api_base=slot.url) + held: Final = asyncio.run( + _burst_through_outage(_ws_base(_proxy_url(gateway)), _proxy_url(gateway), model, key, slot.stop) + ) + record_property("close_codes_during_upstream_outage", sorted(session.close_code or 0 for session in held)) + assert [session.types for session in held] == [("error", "closed")] * OUTAGE_BURST, held + assert all(_relays_the_upstream_close(session) for session in held), held + assert len({session.close_code for session in held}) == 1, held + slot.start() + register_scenario(project, _scripted_turn(), control_url=slot.url) + recovered: Final = _run(gateway, model, key) + assert recovered.completed_turn(project), recovered + rows: Final = _spend_rows(key, OUTAGE_BURST + 1) + assert {str(row["call_type"]) for row in rows} == {"_arealtime"}, rows diff --git a/tests/integration/proxy_config.yaml b/tests/integration/proxy_config.yaml index 6924b641f13..4c027254838 100644 --- a/tests/integration/proxy_config.yaml +++ b/tests/integration/proxy_config.yaml @@ -322,6 +322,30 @@ model_list: api_base: http://127.0.0.1:8191 api_key: synthetic-bedrock-mantle-key aws_region_name: us-east-1 + - model_name: perplexity/pplx-decider-v1-27b + litellm_params: + model: perplexity/pplx-decider-v1-27b + api_base: http://127.0.0.1:8191 + api_key: synthetic-perplexity-key + - model_name: typesafe/jev-1.13.0 + litellm_params: + model: typesafe/jev-1.13.0 + api_base: http://127.0.0.1:8191 + api_key: synthetic-typesafe-key + - model_name: openrouter/typesafe/jev-1.13 + litellm_params: + model: openrouter/typesafe/jev-1.13 + api_base: http://127.0.0.1:8191 + api_key: synthetic-openrouter-key + - model_name: cloudflare/@cf/cloudflare/clef + litellm_params: + model: cloudflare/@cf/cloudflare/clef + api_base: http://127.0.0.1:8191 + api_key: synthetic-cloudflare-key + - model_name: strands_decider/strands-decider-2B-hobson-v19 + litellm_params: + model: strands_decider/strands-decider-2B-hobson-v19 + api_base: http://127.0.0.1:8191 general_settings: master_key: os.environ/LITELLM_MASTER_KEY database_url: os.environ/DATABASE_URL diff --git a/tests/integration/routing/test_priority_scheduler_queue_cleanup.py b/tests/integration/routing/test_priority_scheduler_queue_cleanup.py new file mode 100644 index 00000000000..a4606554d34 --- /dev/null +++ b/tests/integration/routing/test_priority_scheduler_queue_cleanup.py @@ -0,0 +1,860 @@ +from __future__ import annotations + +import asyncio +import http.client +import json +import re +import signal +import threading +import time +import uuid +from collections.abc import Callable, Iterator, Mapping, Sequence +from concurrent.futures import ThreadPoolExecutor +from contextlib import ExitStack +from dataclasses import dataclass +from pathlib import Path +from queue import SimpleQueue +from typing import Final, Literal, TypeVar +from urllib.parse import urlsplit + +import httpx +import openai +import psutil +import pytest +import yaml +from integration._support.client import Gateway, eventually, gateway_from_environment +from integration._support.database import read_rows +from integration._support.openai_wire import answering_model_discovery, chat_reply, openai_error, responses_reply +from integration._support.process import graceful_stop_seconds, owned_proxy_process +from integration._support.redis_process import OwnedRedis, owned_redis +from integration._support.wire import Reply, Request, Wire, wire_server +from pydantic import JsonValue, TypeAdapter +from redis import Redis + +T = TypeVar("T") +Endpoint = Literal["chat", "chat-stream", "completions", "completions-stream", "queue", "queue-stream"] +Instance = Literal["first", "second"] + +MARKER: Final = re.compile(r"sched-[0-9a-f]{32}") +UNKNOWN_MODEL: Final = "Invalid model name passed in model=" +NO_DEPLOYMENTS: Final = "No deployments available" +QUEUE_TIMEOUT: Final = "Request timed out while polling queue" +ROUTER_TIMEOUT_SECONDS: Final = 8 +PROMPT_SECONDS: Final = 4 +COOLDOWN_SECONDS: Final = 60 +USAGE: Final[dict[str, JsonValue]] = {"prompt_tokens": 5, "completion_tokens": 3, "total_tokens": 8} +STARTED_WORKER: Final = re.compile(r"Started server process \[(\d+)\]") +OWNED_CELL_TIMEOUT: Final = 2 * graceful_stop_seconds() + 120 +PAIR_CELL_TIMEOUT: Final = 3 * graceful_stop_seconds() + 120 +WORKER_HEALTHCHECK_ARGUMENTS: Final = ("--timeout_worker_healthcheck", str(int(graceful_stop_seconds()))) +PINNED_CONNECTION_TIMEOUT_SECONDS: Final = 30 +PINNED_LIMITS: Final = httpx.Limits(max_connections=1, max_keepalive_connections=1, keepalive_expiry=30) +ENDPOINTS: Final[tuple[Endpoint, ...]] = ( + "chat", + "chat-stream", + "completions", + "completions-stream", + "queue", + "queue-stream", +) +PAIR_GROUPS: Final = tuple(f"sched-pair-{row}" for row in ("p1", "p2", "p3", "p4", "p5", "p6", "p7", "p8")) +INMEM_GROUP: Final = "sched-inmem" +OUTAGE_GROUP: Final = "sched-outage" +KILL_GROUP: Final = "sched-kill" +GHOST_ENTRY: Final[list[JsonValue]] = [1, "ghost"] +JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) +JSON_LIST: Final = TypeAdapter(list[JsonValue]) + + +def new_marker() -> str: + return f"sched-{uuid.uuid4().hex}" + + +def marker_in(text: str) -> str: + found: Final = MARKER.search(text) + assert found is not None, text[:300] + return found.group(0) + + +def markers_in(text: str) -> frozenset[str]: + return frozenset(MARKER.findall(text)) + + +def upstream_name(group: str) -> str: + return f"{group}-upstream" + + +def queue_key(group: str) -> str: + return f"scheduler:queue:{group}" + + +def data_frame(frame: Mapping[str, JsonValue]) -> bytes: + return b"data: " + json.dumps(frame).encode() + b"\n\n" + + +def text_completion_reply(marker: str, model: str, *, stream: bool) -> Reply: + choice: Final[dict[str, JsonValue]] = { + "text": f"served {marker}", + "index": 0, + "logprobs": None, + "finish_reason": "stop", + } + body: Final[dict[str, JsonValue]] = { + "id": f"cmpl-{marker}", + "object": "text_completion", + "created": 1, + "model": model, + "choices": [choice], + "usage": USAGE, + } + if not stream: + return Reply(body=json.dumps(body).encode()) + first: Final = data_frame({**body, "choices": [{**choice, "finish_reason": None}]}) + last: Final = data_frame({**body, "choices": [{**choice, "text": ""}]}) + return Reply(content_type="text/event-stream", chunks=(first, last + b"data: [DONE]\n\n")) + + +@dataclass(frozen=True, slots=True) +class Upstream: + refusing: Mapping[str, threading.Event] + held: SimpleQueue[str] + release: threading.Event + hold: frozenset[str] + + def respond(self, request: Request) -> Reply: + body: Final = JSON_OBJECT.validate_json(request.body) + model: Final = str(body["model"]) + refusal: Final = self.refusing.get(model) + if refusal is not None and refusal.is_set(): + return openai_error(401) + marker: Final = marker_in(request.body.decode()) + if model in self.hold: + self.held.put(marker) + assert self.release.wait(timeout=60), "The held burst was never released" + stream: Final = body.get("stream") is True + if request.target.endswith("/responses"): + return responses_reply(f"resp_{marker}", model, f"served {marker}", stream=stream) + if request.target.endswith("/completions") and not request.target.endswith("/chat/completions"): + return text_completion_reply(marker, model, stream=stream) + return chat_reply(f"chatcmpl-{marker}", model, f"served {marker}", stream=stream) + + +def upstream_for(groups: Sequence[str], *, hold: Sequence[str] = ()) -> Upstream: + return Upstream( + {upstream_name(group): threading.Event() for group in groups}, + SimpleQueue(), + threading.Event(), + frozenset(upstream_name(group) for group in hold), + ) + + +def received_markers(wire: Wire) -> tuple[str, ...]: + return tuple(marker_in(request.body.decode()) for request in wire.drain() if request.method == "POST") + + +def assert_once(received: Sequence[str], markers: Sequence[str]) -> None: + counts: Final = {marker: received.count(marker) for marker in markers} + assert all(count == 1 for count in counts.values()), counts + + +def assert_never(received: Sequence[str], markers: Sequence[str]) -> None: + assert not set(markers) & set(received), (markers, received) + + +def deployment_entry(group: str, wire: Wire) -> dict[str, JsonValue]: + return { + "model_name": group, + "litellm_params": { + "model": f"openai/{upstream_name(group)}", + "api_base": wire.url + "/v1", + "api_key": "synthetic-openai-key", + }, + "model_info": {"id": f"{group}-deployment"}, + } + + +def owned_config( + directory: Path, + wire: Wire, + groups: Sequence[str], + settings: Mapping[str, JsonValue], + *, + cancel_on_disconnect: bool = False, +) -> Path: + base: Final = JSON_OBJECT.validate_python( + yaml.safe_load((Path(__file__).parents[1] / "proxy_config.yaml").read_text()) + ) + general_settings: Final = base.get("general_settings", {}) + assert isinstance(general_settings, dict) + general: Final[dict[str, JsonValue]] = { + **general_settings, + **({"cancel_on_disconnect": True} if cancel_on_disconnect else {}), + } + config: Final[dict[str, JsonValue]] = { + **base, + "general_settings": general, + "router_settings": dict(settings), + "model_list": [deployment_entry(group, wire) for group in groups], + } + path: Final = directory / f"scheduler-{uuid.uuid4().hex[:8]}.yaml" + path.write_text(yaml.safe_dump(config)) + return path + + +def cooldown_settings(cache: OwnedRedis | None) -> dict[str, JsonValue]: + redis: Final = {} if cache is None else {"redis_host": cache.host, "redis_port": cache.port} + return {"num_retries": 0, "timeout": ROUTER_TIMEOUT_SECONDS, "cooldown_time": COOLDOWN_SECONDS, **redis} + + +def redis_settings(cache: OwnedRedis) -> dict[str, JsonValue]: + return {"num_retries": 0, "redis_host": cache.host, "redis_port": cache.port} + + +def chat_body(group: str, marker: str, **extra: JsonValue) -> dict[str, JsonValue]: + return {"model": group, "messages": [{"role": "user", "content": marker}], **extra} + + +def completion_body(group: str, marker: str, **extra: JsonValue) -> dict[str, JsonValue]: + return {"model": group, "prompt": marker, **extra} + + +@dataclass(frozen=True, slots=True) +class Answer: + status: int + text: str + headers: Mapping[str, str] + seconds: float + + @property + def identity(self) -> str: + return str(JSON_OBJECT.validate_json(self.text)["id"]) + + +def timed(send: Callable[[], httpx.Response]) -> Answer: + started: Final = time.monotonic() + response: Final = send() + return Answer(response.status_code, response.text, dict(response.headers), time.monotonic() - started) + + +def settled(send: Callable[[], T], text: Callable[[T], str]) -> T: + return eventually(send, lambda observed: UNKNOWN_MODEL not in text(observed), seconds=30) + + +def sdk_settled(call: Callable[[], T]) -> T: + def attempt() -> T | openai.APIStatusError: + try: + return call() + except openai.APIStatusError as error: + if UNKNOWN_MODEL in error.message: + return error + raise + + outcome: Final = eventually(attempt, lambda observed: not isinstance(observed, openai.APIStatusError), seconds=30) + assert not isinstance(outcome, openai.APIStatusError) + return outcome + + +def post(gateway: Gateway, path: str, body: Mapping[str, JsonValue]) -> Answer: + return settled(lambda: timed(lambda: gateway.request("POST", path, body)), lambda answer: answer.text) + + +def stream_text(gateway: Gateway, path: str, body: Mapping[str, JsonValue]) -> Answer: + def send() -> Answer: + started: Final = time.monotonic() + with gateway.client.stream( + "POST", path, json=body, headers={"Authorization": f"Bearer {gateway.key}"} + ) as response: + text: Final = "".join(response.iter_text()) + return Answer(response.status_code, text, dict(response.headers), time.monotonic() - started) + + return settled(send, lambda answer: answer.text) + + +def stream_identities(text: str) -> frozenset[str]: + frames: Final = tuple( + JSON_OBJECT.validate_json(line[len("data: ") :]) + for line in text.splitlines() + if line.startswith("data: ") and line != "data: [DONE]" + ) + assert frames, text + return frozenset(str(frame["id"]) for frame in frames) + + +def call(gateway: Gateway, endpoint: Endpoint, group: str, marker: str, priority: int = 1) -> Answer: + match endpoint: + case "chat": + return post(gateway, "/v1/chat/completions", chat_body(group, marker, priority=priority)) + case "chat-stream": + return stream_text( + gateway, "/v1/chat/completions", chat_body(group, marker, priority=priority, stream=True) + ) + case "completions": + return post(gateway, "/v1/completions", completion_body(group, marker, priority=priority)) + case "completions-stream": + return stream_text( + gateway, "/v1/completions", completion_body(group, marker, priority=priority, stream=True) + ) + case "queue": + return post(gateway, "/queue/chat/completions", chat_body(group, marker, priority=priority)) + case "queue-stream": + return stream_text( + gateway, "/queue/chat/completions", chat_body(group, marker, priority=priority, stream=True) + ) + + +def served_identity(status: int, text: str, marker: str) -> str: + assert status == 200, (status, text) + identities: Final = ( + stream_identities(text) if text.startswith("data:") else frozenset({str(JSON_OBJECT.validate_json(text)["id"])}) + ) + assert identities in ({f"chatcmpl-{marker}"}, {f"cmpl-{marker}"}), (identities, marker) + assert markers_in(text) == {marker}, text + return next(iter(identities)) + + +def assert_served(answer: Answer, marker: str) -> str: + return served_identity(answer.status, answer.text, marker) + + +def spend_row_lands(identity: str) -> None: + rows: Final = eventually( + lambda: read_rows('SELECT spend FROM "LiteLLM_SpendLogs" WHERE request_id=%s', (identity,)), + lambda values: len(values) == 1, + seconds=70, + ) + assert rows[0]["spend"] is not None, rows + + +def sdk(gateway: Gateway) -> openai.OpenAI: + return openai.OpenAI(base_url=str(gateway.client.base_url), api_key=gateway.key, max_retries=0, timeout=30) + + +def queue_entries(cache: OwnedRedis, group: str) -> list[JsonValue] | None: + with Redis(host=cache.host, port=cache.port) as client: + raw: Final = client.get(queue_key(group)) + if raw is None: + return None + assert isinstance(raw, bytes), raw + return JSON_LIST.validate_json(raw) + + +def write_queue(cache: OwnedRedis, group: str, entries: Sequence[JsonValue]) -> None: + with Redis(host=cache.host, port=cache.port) as client: + client.set(queue_key(group), json.dumps(list(entries))) + + +def persist_queue(cache: OwnedRedis, group: str) -> None: + with Redis(host=cache.host, port=cache.port) as client: + client.persist(queue_key(group)) + + +def entry_priority(entry: JsonValue) -> int: + assert isinstance(entry, list), entry + priority: Final = entry[0] + assert isinstance(priority, int), entry + return priority + + +def waiting_entries(cache: OwnedRedis, group: str, priority: int) -> list[JsonValue]: + def read() -> list[JsonValue]: + return queue_entries(cache, group) or [] + + return eventually( + read, lambda found: any(entry_priority(entry) == priority for entry in found), seconds=PROMPT_SECONDS + ) + + +def cooled(gateway: Gateway, upstream: Upstream, group: str) -> None: + upstream.refusing[upstream_name(group)].set() + trip: Final = post(gateway, "/v1/chat/completions", chat_body(group, new_marker())) + assert trip.status == 401, (trip.status, trip.text) + eventually( + lambda: post(gateway, "/v1/chat/completions", chat_body(group, new_marker())), + lambda answer: answer.status == 429 and NO_DEPLOYMENTS in answer.text, + seconds=10, + ) + + +def assert_refused_at_once(answer: Answer) -> None: + assert answer.status == 429, (answer.status, answer.text) + assert NO_DEPLOYMENTS in answer.text, answer.text + assert answer.seconds < PROMPT_SECONDS, answer.seconds + + +def assert_timed_out_polling(answer: Answer) -> None: + assert answer.status == 408, (answer.status, answer.text) + assert QUEUE_TIMEOUT in answer.text, answer.text + assert answer.seconds >= ROUTER_TIMEOUT_SECONDS, answer.seconds + + +@pytest.fixture(scope="module") +def rig_upstream() -> Iterator[tuple[Upstream, Wire]]: + upstream: Final = upstream_for(()) + with wire_server(answering_model_discovery(upstream.respond)) as wire: + yield upstream, wire + + +@pytest.fixture +def rig_model(gateway: Gateway, rig_upstream: tuple[Upstream, Wire]) -> Iterator[tuple[str, Wire]]: + _, wire = rig_upstream + with gateway.scenario() as scenario: + yield scenario.model(api_base=wire.url + "/v1"), wire + + +def test_sdk_chat_with_priority_is_served_once_and_billed(gateway: Gateway, rig_model: tuple[str, Wire]) -> None: + model, wire = rig_model + marker: Final = new_marker() + client: Final = sdk(gateway) + raw: Final = sdk_settled( + lambda: client.chat.completions.with_raw_response.create( + model=model, messages=[{"role": "user", "content": marker}], extra_body={"priority": 1} + ) + ) + assert served_identity(raw.status_code, raw.text, marker) == f"chatcmpl-{marker}" + (upstream_request,) = tuple(request for request in wire.drain() if marker.encode() in request.body) + assert "priority" not in JSON_OBJECT.validate_json(upstream_request.body), upstream_request.body + spend_row_lands(f"chatcmpl-{marker}") + + +def test_sdk_chat_stream_with_priority_is_served_once_and_billed(gateway: Gateway, rig_model: tuple[str, Wire]) -> None: + model, wire = rig_model + marker: Final = new_marker() + client: Final = sdk(gateway) + chunks: Final = sdk_settled( + lambda: tuple( + client.chat.completions.create( + model=model, messages=[{"role": "user", "content": marker}], stream=True, extra_body={"priority": 1} + ) + ) + ) + assert {chunk.id for chunk in chunks} == {f"chatcmpl-{marker}"}, chunks + content: Final = "".join(chunk.choices[0].delta.content or "" for chunk in chunks if chunk.choices) + assert content == f"served {marker}", chunks + assert_once(received_markers(wire), (marker,)) + spend_row_lands(f"chatcmpl-{marker}") + + +def test_async_sdk_chat_with_priority_is_served_once(gateway: Gateway, rig_model: tuple[str, Wire]) -> None: + model, wire = rig_model + marker: Final = new_marker() + + async def send() -> tuple[int, str]: + async with openai.AsyncOpenAI( + base_url=str(gateway.client.base_url), api_key=gateway.key, max_retries=0, timeout=30 + ) as client: + raw: Final = await client.chat.completions.with_raw_response.create( + model=model, messages=[{"role": "user", "content": marker}], extra_body={"priority": 1} + ) + return raw.status_code, raw.text + + status, text = sdk_settled(lambda: asyncio.run(send())) + assert served_identity(status, text, marker) == f"chatcmpl-{marker}" + assert_once(received_markers(wire), (marker,)) + + +def test_raw_chat_with_duplicate_priority_keys_takes_the_last_one( + gateway: Gateway, rig_model: tuple[str, Wire] +) -> None: + model, wire = rig_model + marker: Final = new_marker() + body: Final = ( + '{"model": "%s", "messages": [{"role": "user", "content": "%s"}], "priority": "not-an-int", "priority": 1}' + % (model, marker) + ) + + def send() -> httpx.Response: + return gateway.client.post( + "/v1/chat/completions", + content=body.encode(), + headers={"Authorization": f"Bearer {gateway.key}", "Content-Type": "application/json"}, + ) + + answer: Final = settled(lambda: timed(send), lambda observed: observed.text) + assert_served(answer, marker) + (upstream_request,) = tuple(request for request in wire.drain() if marker.encode() in request.body) + assert "priority" not in JSON_OBJECT.validate_json(upstream_request.body), upstream_request.body + + +@pytest.mark.parametrize("endpoint", ("completions", "completions-stream", "queue", "queue-stream")) +def test_other_prioritized_endpoints_are_served_once_and_billed( + gateway: Gateway, rig_model: tuple[str, Wire], endpoint: Endpoint +) -> None: + model, wire = rig_model + marker: Final = new_marker() + answer: Final = call(gateway, endpoint, model, marker) + identity: Final = assert_served(answer, marker) + if endpoint == "queue": + assert answer.headers.get("x-litellm-priority") == "1", answer.headers + assert_once(received_markers(wire), (marker,)) + spend_row_lands(identity) + + +@pytest.mark.parametrize("priority", ("1", [1], "", "p" * 5120, 0), ids=("string", "list", "empty", "5kb", "zero")) +def test_non_integer_priority_bypasses_the_scheduler_and_is_forwarded( + gateway: Gateway, rig_model: tuple[str, Wire], priority: JsonValue +) -> None: + model, wire = rig_model + marker: Final = new_marker() + answer: Final = post(gateway, "/v1/chat/completions", chat_body(model, marker, priority=priority)) + assert_served(answer, marker) + (upstream_request,) = tuple(request for request in wire.drain() if marker.encode() in request.body) + assert JSON_OBJECT.validate_json(upstream_request.body).get("priority") == priority + + +def test_unauthenticated_prioritized_request_never_reaches_the_upstream( + gateway: Gateway, rig_model: tuple[str, Wire] +) -> None: + model, wire = rig_model + marker: Final = new_marker() + response: Final = gateway.client.post("/v1/chat/completions", json=chat_body(model, marker, priority=1)) + assert response.status_code == 401, response.text + assert_never(received_markers(wire), (marker,)) + + +@pytest.mark.parametrize("path", ("/v1/messages", "/v1/responses"), ids=("messages", "responses")) +def test_priority_on_unscheduled_routes_is_served_once( + gateway: Gateway, rig_model: tuple[str, Wire], path: str +) -> None: + model, wire = rig_model + marker: Final = new_marker() + body: Final[dict[str, JsonValue]] = ( + {"model": model, "max_tokens": 32, "messages": [{"role": "user", "content": marker}], "priority": 1} + if path == "/v1/messages" + else {"model": model, "input": marker, "priority": 1} + ) + answer: Final = post(gateway, path, body) + assert answer.status == 200, (answer.status, answer.text) + assert markers_in(answer.text) == {marker}, answer.text + assert_once(received_markers(wire), (marker,)) + + +def pinned(gateway: Gateway, stack: ExitStack) -> Gateway: + client: Final = stack.enter_context( + httpx.Client(base_url=gateway.client.base_url, timeout=15, trust_env=False, limits=PINNED_LIMITS) + ) + return Gateway(client, gateway.key, gateway.upstream_url) + + +@pytest.mark.timeout(OWNED_CELL_TIMEOUT) +def test_in_memory_queue_forgets_served_requests_before_a_cooldown(gateway: Gateway, tmp_path: Path) -> None: + upstream: Final = upstream_for((INMEM_GROUP,)) + with ExitStack() as stack: + wire: Final = stack.enter_context(wire_server(answering_model_discovery(upstream.respond))) + config: Final = owned_config(tmp_path, wire, (INMEM_GROUP,), cooldown_settings(None)) + owned: Final = stack.enter_context( + owned_proxy_process( + gateway, tmp_path, {}, config=config, workers=2, extra_arguments=WORKER_HEALTHCHECK_ARGUMENTS + ) + ) + worker: Final = pinned(owned.gateway, stack) + served: Final = new_marker() + assert_served(post(worker, "/v1/chat/completions", chat_body(INMEM_GROUP, served, priority=1)), served) + cooled(worker, upstream, INMEM_GROUP) + waiting: Final = (new_marker(), new_marker()) + for marker in waiting: + assert_refused_at_once(post(worker, "/v1/chat/completions", chat_body(INMEM_GROUP, marker, priority=2))) + received: Final = received_markers(wire) + assert_once(received, (served,)) + assert_never(received, waiting) + + +@dataclass(frozen=True, slots=True) +class Pair: + first: Gateway + second: Gateway + cache: OwnedRedis + wire: Wire + upstream: Upstream + + def at(self, instance: Instance) -> Gateway: + return self.first if instance == "first" else self.second + + +@pytest.fixture(scope="module") +def pair(tmp_path_factory: pytest.TempPathFactory) -> Iterator[Pair]: + directory: Final = tmp_path_factory.mktemp("scheduler-pair") + upstream: Final = upstream_for(PAIR_GROUPS) + with ExitStack() as stack: + gateway: Final = stack.enter_context(gateway_from_environment()) + cache: Final = stack.enter_context(owned_redis(directory)) + wire: Final = stack.enter_context(wire_server(answering_model_discovery(upstream.respond))) + config: Final = owned_config(directory, wire, PAIR_GROUPS, cooldown_settings(cache), cancel_on_disconnect=True) + overrides: Final = {"REDIS_HOST": cache.host, "REDIS_PORT": str(cache.port)} + first: Final = stack.enter_context( + owned_proxy_process( + gateway, directory, overrides, config=config, workers=2, extra_arguments=WORKER_HEALTHCHECK_ARGUMENTS + ) + ) + second: Final = stack.enter_context( + owned_proxy_process( + gateway, directory, overrides, config=config, workers=2, extra_arguments=WORKER_HEALTHCHECK_ARGUMENTS + ) + ) + yield Pair(first.gateway, second.gateway, cache, wire, upstream) + + +@pytest.mark.timeout(PAIR_CELL_TIMEOUT) +def test_served_request_leaves_no_queue_entry_for_the_other_instance(pair: Pair) -> None: + group: Final = PAIR_GROUPS[0] + markers: Final = (new_marker(), new_marker(), new_marker()) + assert_served(post(pair.first, "/v1/chat/completions", chat_body(group, markers[0], priority=1)), markers[0]) + assert queue_entries(pair.cache, group) == [] + assert_served(post(pair.second, "/v1/chat/completions", chat_body(group, markers[1], priority=1)), markers[1]) + assert_served(post(pair.first, "/v1/chat/completions", chat_body(group, markers[2], priority=1)), markers[2]) + assert_once(received_markers(pair.wire), markers) + + +@pytest.mark.timeout(PAIR_CELL_TIMEOUT) +def test_concurrent_prioritized_burst_across_instances_and_endpoints_is_served(pair: Pair) -> None: + group: Final = PAIR_GROUPS[1] + primer: Final = new_marker() + assert_served(post(pair.first, "/v1/chat/completions", chat_body(group, primer, priority=1)), primer) + instances: Final[tuple[Instance, ...]] = ("first", "second") + plan: Final = tuple((instance, endpoint, new_marker()) for instance in instances for endpoint in ENDPOINTS) + + def one(item: tuple[Instance, Endpoint, str]) -> Answer: + instance, endpoint, marker = item + return call(pair.at(instance), endpoint, group, marker) + + with ThreadPoolExecutor(max_workers=len(plan)) as pool: + answers: Final = tuple(pool.map(one, plan)) + for (_, _, marker), answer in zip(plan, answers, strict=True): + assert_served(answer, marker) + closing: Final = new_marker() + assert_served(post(pair.second, "/v1/chat/completions", chat_body(group, closing, priority=1)), closing) + assert_once(received_markers(pair.wire), (primer, *(marker for _, _, marker in plan), closing)) + + +@pytest.mark.timeout(PAIR_CELL_TIMEOUT) +def test_fresh_queue_during_a_cooldown_is_refused_at_once(pair: Pair) -> None: + group: Final = PAIR_GROUPS[2] + cooled(pair.second, pair.upstream, group) + client: Final = sdk(pair.second) + started: Final = time.monotonic() + with pytest.raises(openai.APIStatusError) as refused: + client.completions.create(model=group, prompt=new_marker(), extra_body={"priority": 1}) + assert refused.value.status_code == 429, refused.value.message + assert NO_DEPLOYMENTS in refused.value.message + assert time.monotonic() - started < PROMPT_SECONDS + assert queue_entries(pair.cache, group) == [] + + +@pytest.mark.timeout(PAIR_CELL_TIMEOUT) +def test_served_request_on_one_instance_does_not_block_a_waiter_on_the_other(pair: Pair) -> None: + group: Final = PAIR_GROUPS[3] + served: Final = new_marker() + assert_served(post(pair.first, "/v1/chat/completions", chat_body(group, served, priority=1)), served) + cooled(pair.second, pair.upstream, group) + assert_refused_at_once(post(pair.second, "/v1/chat/completions", chat_body(group, new_marker(), priority=2))) + + +def assert_refused_before_the_router_timeout(answer: Answer) -> None: + assert answer.status == 429, (answer.status, answer.text) + assert NO_DEPLOYMENTS in answer.text, answer.text + assert answer.seconds < ROUTER_TIMEOUT_SECONDS, answer.seconds + + +def assert_drained(cache: OwnedRedis, group: str) -> None: + assert queue_entries(cache, group) in ([], None) + + +@pytest.mark.timeout(PAIR_CELL_TIMEOUT) +def test_waiter_behind_an_expired_dead_replica_entry_is_re_enqueued_and_proceeds(pair: Pair) -> None: + group: Final = PAIR_GROUPS[4] + cooled(pair.second, pair.upstream, group) + write_queue(pair.cache, group, [GHOST_ENTRY]) + answer: Final = post(pair.second, "/queue/chat/completions", chat_body(group, new_marker(), priority=2)) + assert_refused_before_the_router_timeout(answer) + assert_drained(pair.cache, group) + + +@pytest.mark.timeout(PAIR_CELL_TIMEOUT) +def test_waiter_behind_a_live_entry_times_out_and_removes_only_itself(pair: Pair) -> None: + group: Final = PAIR_GROUPS[5] + cooled(pair.second, pair.upstream, group) + write_queue(pair.cache, group, [GHOST_ENTRY]) + with ThreadPoolExecutor(max_workers=1) as pool: + waiter: Final = pool.submit( + post, pair.second, "/v1/chat/completions", chat_body(group, new_marker(), priority=2) + ) + waiting_entries(pair.cache, group, 2) + persist_queue(pair.cache, group) + assert_timed_out_polling(waiter.result(timeout=30)) + assert queue_entries(pair.cache, group) == [GHOST_ENTRY] + + +@pytest.mark.timeout(PAIR_CELL_TIMEOUT) +def test_waiter_proceeds_once_the_entry_ahead_is_cleared(pair: Pair) -> None: + group: Final = PAIR_GROUPS[6] + cooled(pair.second, pair.upstream, group) + write_queue(pair.cache, group, [GHOST_ENTRY]) + with ThreadPoolExecutor(max_workers=1) as pool: + waiter: Final = pool.submit( + post, pair.second, "/v1/chat/completions", chat_body(group, new_marker(), priority=2) + ) + entries: Final = waiting_entries(pair.cache, group, 2) + write_queue(pair.cache, group, [entry for entry in entries if entry_priority(entry) == 2]) + assert_refused_before_the_router_timeout(waiter.result(timeout=30)) + assert_drained(pair.cache, group) + + +def pinned_connection(gateway: Gateway) -> http.client.HTTPConnection: + address: Final = urlsplit(str(gateway.client.base_url)) + assert address.hostname is not None and address.port is not None, address + return http.client.HTTPConnection(address.hostname, address.port, timeout=PINNED_CONNECTION_TIMEOUT_SECONDS) + + +def send_pinned(connection: http.client.HTTPConnection, key: str, body: Mapping[str, JsonValue]) -> None: + connection.request( + "POST", + "/v1/chat/completions", + body=json.dumps(body).encode(), + headers={"Authorization": f"Bearer {key}", "Content-Type": "application/json"}, + ) + + +def post_pinned(connection: http.client.HTTPConnection, key: str, body: Mapping[str, JsonValue]) -> tuple[int, str]: + send_pinned(connection, key, body) + response: Final = connection.getresponse() + text: Final = response.read().decode() + assert connection.sock is not None, "the proxy closed the pinned connection" + return response.status, text + + +def cooled_over(connection: http.client.HTTPConnection, key: str, upstream: Upstream, group: str) -> None: + upstream.refusing[upstream_name(group)].set() + trip: Final = post_pinned(connection, key, chat_body(group, new_marker())) + assert trip[0] == 401, trip + eventually( + lambda: post_pinned(connection, key, chat_body(group, new_marker())), + lambda answer: answer[0] == 429 and NO_DEPLOYMENTS in answer[1], + seconds=10, + ) + + +@pytest.mark.timeout(PAIR_CELL_TIMEOUT) +def test_disconnected_waiter_is_removed_from_the_queue(pair: Pair) -> None: + group: Final = PAIR_GROUPS[7] + waiter: Final = pinned_connection(pair.second) + cooled_over(waiter, pair.second.key, pair.upstream, group) + write_queue(pair.cache, group, [GHOST_ENTRY]) + send_pinned(waiter, pair.second.key, chat_body(group, new_marker(), priority=2)) + waiting_entries(pair.cache, group, 2) + persist_queue(pair.cache, group) + waiter.close() + eventually( + lambda: queue_entries(pair.cache, group), lambda entries: entries == [GHOST_ENTRY], seconds=PROMPT_SECONDS + ) + + +Outcome = Answer | httpx.TransportError + + +def burst(gateway: Gateway, group: str, count: int) -> tuple[tuple[Endpoint, str, Outcome], ...]: + plan: Final[tuple[tuple[Endpoint, str], ...]] = tuple( + (ENDPOINTS[index % len(ENDPOINTS)], new_marker()) for index in range(count) + ) + + def one(item: tuple[Endpoint, str]) -> Outcome: + endpoint, marker = item + with httpx.Client(base_url=gateway.client.base_url, timeout=60, trust_env=False) as client: + try: + return call(Gateway(client, gateway.key, gateway.upstream_url), endpoint, group, marker) + except httpx.TransportError as error: + return error + + with ThreadPoolExecutor(max_workers=count) as pool: + outcomes: Final = tuple(pool.map(one, plan)) + return tuple((endpoint, marker, outcome) for (endpoint, marker), outcome in zip(plan, outcomes, strict=True)) + + +def assert_all_served(outcomes: Sequence[tuple[Endpoint, str, Outcome]]) -> tuple[str, ...]: + for _, marker, outcome in outcomes: + assert isinstance(outcome, Answer), repr(outcome) + assert_served(outcome, marker) + return tuple(marker for _, marker, _ in outcomes) + + +def served_eventually(gateway: Gateway, group: str, marker: str) -> None: + def send() -> Outcome: + try: + return post(gateway, "/v1/chat/completions", chat_body(group, marker, priority=1)) + except httpx.TransportError as error: + return error + + answer: Final = eventually(send, lambda outcome: isinstance(outcome, Answer), seconds=30) + assert isinstance(answer, Answer) + assert_served(answer, marker) + + +@pytest.mark.timeout(OWNED_CELL_TIMEOUT) +def test_prioritized_requests_survive_a_redis_outage(gateway: Gateway, tmp_path: Path) -> None: + upstream: Final = upstream_for((OUTAGE_GROUP,)) + with owned_redis(tmp_path) as cache, wire_server(answering_model_discovery(upstream.respond)) as wire: + config: Final = owned_config(tmp_path, wire, (OUTAGE_GROUP,), redis_settings(cache)) + overrides: Final = { + "REDIS_HOST": cache.host, + "REDIS_PORT": str(cache.port), + "REDIS_CIRCUIT_BREAKER_RECOVERY_TIMEOUT": "1", + } + with owned_proxy_process( + gateway, tmp_path, overrides, config=config, workers=2, extra_arguments=WORKER_HEALTHCHECK_ARGUMENTS + ) as owned: + before: Final = assert_all_served(burst(owned.gateway, OUTAGE_GROUP, 12)) + cache.stop() + during: Final = assert_all_served(burst(owned.gateway, OUTAGE_GROUP, 12)) + cache.start() + after: Final = assert_all_served(burst(owned.gateway, OUTAGE_GROUP, 12)) + assert_once(received_markers(wire), (*before, *during, *after)) + eventually(lambda: queue_entries(cache, OUTAGE_GROUP), lambda entries: entries in (None, []), seconds=10) + + +def established_upstream_connections(pid: int, wire: Wire) -> int: + port: Final = urlsplit(wire.url).port + return sum( + 1 + for connection in psutil.Process(pid).net_connections(kind="tcp") + if connection.status == psutil.CONN_ESTABLISHED and connection.raddr and connection.raddr.port == port + ) + + +@pytest.mark.timeout(OWNED_CELL_TIMEOUT) +def test_sibling_worker_keeps_serving_prioritized_requests_after_a_worker_is_killed( + gateway: Gateway, tmp_path: Path +) -> None: + upstream: Final = upstream_for((KILL_GROUP,), hold=(KILL_GROUP,)) + with owned_redis(tmp_path) as cache, wire_server(answering_model_discovery(upstream.respond)) as wire: + config: Final = owned_config(tmp_path, wire, (KILL_GROUP,), redis_settings(cache)) + overrides: Final = {"REDIS_HOST": cache.host, "REDIS_PORT": str(cache.port)} + with owned_proxy_process( + gateway, tmp_path, overrides, config=config, workers=2, extra_arguments=WORKER_HEALTHCHECK_ARGUMENTS + ) as owned: + workers: Final = eventually( + lambda: tuple(int(found.group(1)) for found in STARTED_WORKER.finditer(owned.log.read_text())), + lambda pids: len(pids) == 2, + seconds=graceful_stop_seconds(), + ) + with ThreadPoolExecutor(max_workers=1) as pool: + pending: Final = pool.submit(burst, owned.gateway, KILL_GROUP, 20) + eventually(upstream.held.qsize, lambda size: size == 20, seconds=60) + held_by: Final = {pid: established_upstream_connections(pid, wire) for pid in workers} + assert sum(held_by.values()) == 20, held_by + victim_pid, survivor_pid = sorted(workers, key=held_by.__getitem__) + victim: Final = psutil.Process(victim_pid) + victim.suspend() + victim.send_signal(signal.SIGKILL) + upstream.release.set() + outcomes: Final = pending.result(timeout=90) + answered: Final = tuple(outcome for outcome in outcomes if isinstance(outcome[2], Answer)) + failed: Final = tuple(outcome for outcome in outcomes if not isinstance(outcome[2], Answer)) + assert len(answered) == held_by[survivor_pid], (held_by, len(answered)) + assert len(failed) == held_by[victim_pid], (held_by, len(failed)) + assert_all_served(answered) + follow_up: Final = new_marker() + served_eventually(owned.gateway, KILL_GROUP, follow_up) + eventually( + lambda: len(STARTED_WORKER.findall(owned.log.read_text())), + lambda count: count == 3, + seconds=graceful_stop_seconds(), + ) + assert_once(received_markers(wire), (*(marker for _, marker, _ in outcomes), follow_up)) diff --git a/tests/integration/sdk/test_anthropic_messages_bridge_mid_stream_failure_sdk.py b/tests/integration/sdk/test_anthropic_messages_bridge_mid_stream_failure_sdk.py new file mode 100644 index 00000000000..f719b8fbbb2 --- /dev/null +++ b/tests/integration/sdk/test_anthropic_messages_bridge_mid_stream_failure_sdk.py @@ -0,0 +1,113 @@ +import asyncio +import json +import threading +import uuid +from collections.abc import AsyncIterable, Callable +from typing import Final + +from integration._support.anthropic_sse import ( + SseEvent, + delta_text, + dropping_reply, + event_types, + parse_sse, + stream_reply, + user_prompt, +) +from integration._support.client import eventually, object_value +from integration._support.openai_wire import answering_model_discovery, chat_stream, posted_targets +from integration._support.wire import Reply, Request, Wire, wire_server + +import litellm +from litellm.exceptions import MidStreamFallbackError +from litellm.integrations.custom_logger import CustomLogger + +_BACKEND: Final = "gpt-4o-mini" +_PROVIDER_KEY: Final = "integration-provider-key" +_UPSTREAM_TARGET: Final = "/v1/chat/completions" +_TEXT: Final = "Hello" +_DROPS_AFTER_CONTENT: Final = "drops-after-content" +_SUCCEEDS: Final = "succeeds" + + +class _CallbackProbe(CustomLogger): + def __init__(self) -> None: + super().__init__() + self._lock: Final = threading.Lock() + self._failures: Final[list[object]] = [] + self._successes: Final[list[object]] = [] + + async def async_log_failure_event( + self, kwargs: dict[str, object], response_obj: object, start_time: object, end_time: object + ) -> None: + with self._lock: + self._failures.append(kwargs.get("exception")) + + async def async_log_success_event( + self, kwargs: dict[str, object], response_obj: object, start_time: object, end_time: object + ) -> None: + with self._lock: + self._successes.append(kwargs.get("litellm_call_id")) + + def failures(self) -> tuple[object, ...]: + with self._lock: + return tuple(self._failures) + + def successes(self) -> tuple[object, ...]: + with self._lock: + return tuple(self._successes) + + +def _upstream(marker: str) -> Callable[[Request], Reply]: + def respond(request: Request) -> Reply: + assert (request.method, request.target) == ("POST", _UPSTREAM_TARGET), request + assert request.headers["authorization"] == f"Bearer {_PROVIDER_KEY}", request.headers + body: Final = object_value(json.loads(request.body)) + assert body["model"] == _BACKEND and body["stream"] is True, body + outcome, _, scripted_marker = user_prompt(body).partition(":") + assert scripted_marker == marker, body + chunks: Final = chat_stream(f"chunk-{outcome}-{marker}", _BACKEND, _TEXT) + if outcome == _DROPS_AFTER_CONTENT: + return dropping_reply(chunks, abort_after=2) + return stream_reply(chunks) + + return answering_model_discovery(respond) + + +async def _stream(wire: Wire, prompt: str) -> tuple[SseEvent, ...]: + response: Final = await litellm.anthropic.messages.acreate( + model=f"hosted_vllm/{_BACKEND}", + api_base=wire.url + "/v1", + api_key=_PROVIDER_KEY, + max_tokens=16, + stream=True, + messages=[{"role": "user", "content": prompt}], + ) + assert isinstance(response, AsyncIterable), response + frames: Final = [frame async for frame in response] + assert all(isinstance(frame, bytes) for frame in frames), frames + return parse_sse(b"".join(frame for frame in frames if isinstance(frame, bytes)).decode()) + + +async def _settled(read: Callable[[], tuple[object, ...]]) -> tuple[object, ...]: + return await asyncio.to_thread(eventually, read, lambda seen: len(seen) >= 1) + + +async def test_sdk_bridged_stream_failing_after_content_reports_the_provider_error_once() -> None: + probe: Final = _CallbackProbe() + litellm.callbacks.append(probe) + marker: Final = uuid.uuid4().hex + with wire_server(_upstream(marker)) as wire: + failing: Final = await _stream(wire, f"{_DROPS_AFTER_CONTENT}:{marker}") + assert delta_text(failing) == _TEXT, failing + assert event_types(failing)[-1] == "error" and "message_stop" not in event_types(failing), failing + await _settled(probe.failures) + succeeding: Final = await _stream(wire, f"{_SUCCEEDS}:{marker}") + assert event_types(succeeding)[-1] == "message_stop", succeeding + assert delta_text(succeeding) == _TEXT, succeeding + await _settled(probe.successes) + assert posted_targets(wire) == (_UPSTREAM_TARGET,) * 2 + failures: Final = probe.failures() + assert len(failures) == 1, failures + assert isinstance(failures[0], Exception), failures + assert not isinstance(failures[0], MidStreamFallbackError), failures diff --git a/tests/integration/spend/test_redis_ttl_preserving_token_increment.py b/tests/integration/spend/test_redis_ttl_preserving_token_increment.py index 3f623280c17..8b7a21f73d5 100644 --- a/tests/integration/spend/test_redis_ttl_preserving_token_increment.py +++ b/tests/integration/spend/test_redis_ttl_preserving_token_increment.py @@ -8,7 +8,7 @@ from redis import Redis from litellm.caching.caching import DualCache from litellm.caching.redis_cache import RedisCache from litellm.proxy.hooks.parallel_request_limiter_v3 import ( - _PROXY_MaxParallelRequestsHandler_v3 as _PROXY_MaxParallelRequestsHandler, + PROXY_MaxParallelRequestsHandler_v3 as _PROXY_MaxParallelRequestsHandler, ) from litellm.proxy.utils import InternalUsageCache from litellm.types.caching import RedisPipelineIncrementOperation diff --git a/tests/integration/translation/decisions/__init__.py b/tests/integration/translation/decisions/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/integration/translation/decisions/bases/__init__.py b/tests/integration/translation/decisions/bases/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/integration/translation/decisions/bases/cloudflare.py b/tests/integration/translation/decisions/bases/cloudflare.py new file mode 100644 index 00000000000..1668dc215a8 --- /dev/null +++ b/tests/integration/translation/decisions/bases/cloudflare.py @@ -0,0 +1,84 @@ +from typing import Final + +from integration.translation.case import TranslationTestCase + +"""Provider request and reply shape from https://developers.cloudflare.com/workers-ai (POST /ai/run/@cf/cloudflare/clef, reply wrapped in result). Mock reply captured live on 2026-10-06. +""" +CLEF_TEST_CASE: Final = TranslationTestCase( + scenario="basic", + litellm_endpoint="/v1/decisions", + litellm_request={ + "model": "cloudflare/@cf/cloudflare/clef", + "state": "Ticket (billing): The export job hangs at 99% and never finishes", + "questions": { + "defect": {"type": "noul", "instructions": "Is this a defect?"}, + "severity": { + "type": "choice", + "instructions": "How severe is it?", + "criteria": {"low": "cosmetic", "high": "blocks users"}, + }, + "confidence": {"type": "score", "instructions": "How sure are you?", "criteria": ["unsure", "sure"]}, + }, + "cache": {"no-cache": True}, + }, + expected_provider_endpoint="/ai/run/@cf/cloudflare/clef", + expected_provider_headers={"authorization": "Bearer synthetic-cloudflare-key", "content-type": "application/json"}, + expected_provider_request={ + "model": "clef", + "state": "Ticket (billing): The export job hangs at 99% and never finishes", + "questions": { + "defect": {"type": "noul", "instructions": "Is this a defect?"}, + "severity": { + "type": "choice", + "instructions": "How severe is it?", + "criteria": {"low": "cosmetic", "high": "blocks users"}, + }, + "confidence": {"type": "score", "instructions": "How sure are you?", "criteria": ["unsure", "sure"]}, + }, + }, + mock_provider_response={ + "result": { + "model": "clef", + "answers": { + "defect": {"type": "noul", "noul": 0.9345}, + "severity": { + "type": "choice", + "choice": "high", + "confidence": 0.8067, + "probabilities": {"low": 0.0509, "high": 0.9491}, + }, + "confidence": { + "type": "score", + "score": 0.9036, + "confidence": 0.6515, + "legend": {"0": "unsure", "1": "sure"}, + "probabilities": {"0": 0.0964, "1": 0.9036}, + }, + }, + "usage": {"input_tokens": 290, "output_tokens": 0}, + }, + "success": True, + "errors": [], + "messages": [], + }, + expected_litellm_response={ + "model": "clef", + "answers": { + "defect": {"type": "noul", "noul": 0.9345}, + "severity": { + "type": "choice", + "choice": "high", + "confidence": 0.8067, + "probabilities": {"low": 0.0509, "high": 0.9491}, + }, + "confidence": { + "type": "score", + "score": 0.9036, + "confidence": 0.6515, + "legend": {"0": "unsure", "1": "sure"}, + "probabilities": {"0": 0.0964, "1": 0.9036}, + }, + }, + "usage": {"input_tokens": 290, "output_tokens": 0}, + }, +) diff --git a/tests/integration/translation/decisions/bases/openrouter.py b/tests/integration/translation/decisions/bases/openrouter.py new file mode 100644 index 00000000000..f6217bcf818 --- /dev/null +++ b/tests/integration/translation/decisions/bases/openrouter.py @@ -0,0 +1,83 @@ +from typing import Final + +from integration.translation.case import TranslationTestCase + +"""Provider request and reply shape from https://openrouter.ai/docs (POST /api/alpha/decisions). Mock reply captured live on 2026-10-06. +""" +TYPESAFE_JEV_1_13_TEST_CASE: Final = TranslationTestCase( + scenario="basic", + litellm_endpoint="/v1/decisions", + litellm_request={ + "model": "openrouter/typesafe/jev-1.13", + "state": "Ticket (billing): The export job hangs at 99% and never finishes", + "questions": { + "defect": {"type": "noul", "instructions": "Is this a defect?"}, + "severity": { + "type": "choice", + "instructions": "How severe is it?", + "criteria": {"low": "cosmetic", "high": "blocks users"}, + }, + "confidence": {"type": "score", "instructions": "How sure are you?", "criteria": ["unsure", "sure"]}, + }, + "cache": {"no-cache": True}, + }, + expected_provider_endpoint="/alpha/decisions", + expected_provider_headers={"authorization": "Bearer synthetic-openrouter-key", "content-type": "application/json"}, + expected_provider_request={ + "model": "typesafe/jev-1.13", + "state": "Ticket (billing): The export job hangs at 99% and never finishes", + "questions": { + "defect": {"type": "noul", "instructions": "Is this a defect?"}, + "severity": { + "type": "choice", + "instructions": "How severe is it?", + "criteria": {"low": "cosmetic", "high": "blocks users"}, + }, + "confidence": {"type": "score", "instructions": "How sure are you?", "criteria": ["unsure", "sure"]}, + }, + }, + mock_provider_response={ + "model": "typesafe/jev-1.13-20260917", + "answers": { + "defect": {"type": "noul", "noul": 0.81}, + "severity": { + "type": "choice", + "choice": "high", + "confidence": 0.99, + "probabilities": {"low": 0.01, "high": 0.99}, + }, + "confidence": { + "type": "score", + "score": 0.5, + "confidence": 0, + "legend": {"0": "unsure", "1": "sure"}, + "probabilities": {"0": 0.5, "1": 0.5}, + }, + }, + "usage": {"input_tokens": 377, "output_tokens": 62, "cost": 1.5834e-05}, + "id": "gen-dec-1791323839-GA15kY0nt34oiJ7srfki", + "provider": "TypeSafe", + }, + expected_litellm_response={ + "model": "typesafe/jev-1.13-20260917", + "answers": { + "defect": {"type": "noul", "noul": 0.81}, + "severity": { + "type": "choice", + "choice": "high", + "confidence": 0.99, + "probabilities": {"low": 0.01, "high": 0.99}, + }, + "confidence": { + "type": "score", + "score": 0.5, + "confidence": 0, + "legend": {"0": "unsure", "1": "sure"}, + "probabilities": {"0": 0.5, "1": 0.5}, + }, + }, + "usage": {"input_tokens": 377, "output_tokens": 62, "cost": 1.5834e-05}, + "id": "gen-dec-1791323839-GA15kY0nt34oiJ7srfki", + "provider": "TypeSafe", + }, +) diff --git a/tests/integration/translation/decisions/bases/perplexity.py b/tests/integration/translation/decisions/bases/perplexity.py new file mode 100644 index 00000000000..71ef7a85220 --- /dev/null +++ b/tests/integration/translation/decisions/bases/perplexity.py @@ -0,0 +1,79 @@ +from typing import Final + +from integration.translation.case import TranslationTestCase + +"""Provider request and reply shape from https://docs.perplexity.ai (POST /v1/decisions). Mock reply captured live on 2026-10-06. +""" +PPLX_DECIDER_V1_27B_TEST_CASE: Final = TranslationTestCase( + scenario="basic", + litellm_endpoint="/v1/decisions", + litellm_request={ + "model": "perplexity/pplx-decider-v1-27b", + "state": "Ticket (billing): The export job hangs at 99% and never finishes", + "questions": { + "defect": {"type": "noul", "instructions": "Is this a defect?"}, + "severity": { + "type": "choice", + "instructions": "How severe is it?", + "criteria": {"low": "cosmetic", "high": "blocks users"}, + }, + "confidence": {"type": "score", "instructions": "How sure are you?", "criteria": ["unsure", "sure"]}, + }, + "cache": {"no-cache": True}, + }, + expected_provider_endpoint="/v1/decisions", + expected_provider_headers={"authorization": "Bearer synthetic-perplexity-key", "content-type": "application/json"}, + expected_provider_request={ + "model": "pplx-decider-v1-27b", + "state": "Ticket (billing): The export job hangs at 99% and never finishes", + "questions": { + "defect": {"type": "noul", "instructions": "Is this a defect?"}, + "severity": { + "type": "choice", + "instructions": "How severe is it?", + "criteria": {"low": "cosmetic", "high": "blocks users"}, + }, + "confidence": {"type": "score", "instructions": "How sure are you?", "criteria": ["unsure", "sure"]}, + }, + }, + mock_provider_response={ + "model": "pplx-decider-v1-27b", + "answers": { + "defect": {"type": "noul", "noul": 0.9989100737587077}, + "severity": { + "type": "choice", + "choice": "high", + "confidence": 0.9964631215356778, + "probabilities": {"low": 0.0017684392321610232, "high": 0.9982315607678389}, + }, + "confidence": { + "type": "score", + "score": 0.07367392327139817, + "confidence": 0.8526521534572037, + "legend": {"0": "unsure", "1": "sure"}, + "probabilities": {"0": 0.9263260767286018, "1": 0.07367392327139817}, + }, + }, + "usage": {"input_tokens": 318, "output_tokens": 3}, + }, + expected_litellm_response={ + "model": "pplx-decider-v1-27b", + "answers": { + "defect": {"type": "noul", "noul": 0.9989100737587077}, + "severity": { + "type": "choice", + "choice": "high", + "confidence": 0.9964631215356778, + "probabilities": {"low": 0.0017684392321610232, "high": 0.9982315607678389}, + }, + "confidence": { + "type": "score", + "score": 0.07367392327139817, + "confidence": 0.8526521534572037, + "legend": {"0": "unsure", "1": "sure"}, + "probabilities": {"0": 0.9263260767286018, "1": 0.07367392327139817}, + }, + }, + "usage": {"input_tokens": 318, "output_tokens": 3}, + }, +) diff --git a/tests/integration/translation/decisions/bases/strands_decider.py b/tests/integration/translation/decisions/bases/strands_decider.py new file mode 100644 index 00000000000..83d7ae5e33a --- /dev/null +++ b/tests/integration/translation/decisions/bases/strands_decider.py @@ -0,0 +1,35 @@ +from typing import Final + +from integration.translation.case import TranslationTestCase + +"""Provider request and reply shape from https://github.com/strands-agents/decider (POST /v1/systemone, no auth). Mock reply captured live on 2026-10-06. +""" +STRANDS_DECIDER_2B_HOBSON_V19_TEST_CASE: Final = TranslationTestCase( + scenario="basic", + litellm_endpoint="/v1/decisions", + litellm_request={ + "model": "strands_decider/strands-decider-2B-hobson-v19", + "state": "Help! My payouts have been failing for 3 days!", + "questions": {"is_urgent": {"type": "noul", "instructions": "Does this convey urgency?"}}, + "cache": {"no-cache": True}, + }, + expected_provider_endpoint="/v1/systemone", + expected_provider_headers={"content-type": "application/json"}, + expected_provider_request={ + "model": "strands-decider-2B-hobson-v19", + "state": "Help! My payouts have been failing for 3 days!", + "questions": {"is_urgent": {"type": "noul", "instructions": "Does this convey urgency?"}}, + }, + mock_provider_response={ + "model": "strands-decider-2B-hobson-v19", + "answers": {"is_urgent": {"type": "noul", "noul": 0.8277}}, + "usage": {"input_tokens": 86, "output_tokens": 1}, + "latency_ms": 140.03, + }, + expected_litellm_response={ + "model": "strands-decider-2B-hobson-v19", + "answers": {"is_urgent": {"type": "noul", "noul": 0.8277}}, + "usage": {"input_tokens": 86, "output_tokens": 1}, + "latency_ms": 140.03, + }, +) diff --git a/tests/integration/translation/decisions/bases/typesafe.py b/tests/integration/translation/decisions/bases/typesafe.py new file mode 100644 index 00000000000..2fb384a94ae --- /dev/null +++ b/tests/integration/translation/decisions/bases/typesafe.py @@ -0,0 +1,79 @@ +from typing import Final + +from integration.translation.case import TranslationTestCase + +"""Provider request and reply shape from https://docs.typesafe.ai (POST /v1/systemone). Mock reply captured live on 2026-10-06. +""" +JEV_1_13_0_TEST_CASE: Final = TranslationTestCase( + scenario="basic", + litellm_endpoint="/v1/decisions", + litellm_request={ + "model": "typesafe/jev-1.13.0", + "state": "Ticket (billing): The export job hangs at 99% and never finishes", + "questions": { + "defect": {"type": "noul", "instructions": "Is this a defect?"}, + "severity": { + "type": "choice", + "instructions": "How severe is it?", + "criteria": {"low": "cosmetic", "high": "blocks users"}, + }, + "confidence": {"type": "score", "instructions": "How sure are you?", "criteria": ["unsure", "sure"]}, + }, + "cache": {"no-cache": True}, + }, + expected_provider_endpoint="/v1/systemone", + expected_provider_headers={"authorization": "Bearer synthetic-typesafe-key", "content-type": "application/json"}, + expected_provider_request={ + "model": "jev-1.13.0", + "state": "Ticket (billing): The export job hangs at 99% and never finishes", + "questions": { + "defect": {"type": "noul", "instructions": "Is this a defect?"}, + "severity": { + "type": "choice", + "instructions": "How severe is it?", + "criteria": {"low": "cosmetic", "high": "blocks users"}, + }, + "confidence": {"type": "score", "instructions": "How sure are you?", "criteria": ["unsure", "sure"]}, + }, + }, + mock_provider_response={ + "model": "jev-1.13.0", + "answers": { + "defect": {"type": "noul", "noul": 0.78}, + "severity": { + "type": "choice", + "choice": "high", + "confidence": 0.99, + "probabilities": {"low": 0.01, "high": 0.99}, + }, + "confidence": { + "type": "score", + "score": 0.51, + "confidence": 0.03, + "legend": {"0": "unsure", "1": "sure"}, + "probabilities": {"0": 0.49, "1": 0.51}, + }, + }, + "usage": {"input_tokens": 377, "output_tokens": 62}, + }, + expected_litellm_response={ + "model": "jev-1.13.0", + "answers": { + "defect": {"type": "noul", "noul": 0.78}, + "severity": { + "type": "choice", + "choice": "high", + "confidence": 0.99, + "probabilities": {"low": 0.01, "high": 0.99}, + }, + "confidence": { + "type": "score", + "score": 0.51, + "confidence": 0.03, + "legend": {"0": "unsure", "1": "sure"}, + "probabilities": {"0": 0.49, "1": 0.51}, + }, + }, + "usage": {"input_tokens": 377, "output_tokens": 62}, + }, +) diff --git a/tests/integration/translation/decisions/basic/__init__.py b/tests/integration/translation/decisions/basic/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/integration/translation/decisions/basic/test_decisions_basic_cloudflare.py b/tests/integration/translation/decisions/basic/test_decisions_basic_cloudflare.py new file mode 100644 index 00000000000..79fc0574b3b --- /dev/null +++ b/tests/integration/translation/decisions/basic/test_decisions_basic_cloudflare.py @@ -0,0 +1,11 @@ +import pytest +from integration._support.client import Gateway +from integration._support.provider import SharedProvider +from integration.translation.case import TranslationTestCase +from integration.translation.decisions.bases.cloudflare import CLEF_TEST_CASE +from integration.translation.runner import assert_translation + + +@pytest.mark.parametrize("case", [CLEF_TEST_CASE], ids=lambda case: case.id) +def test_decisions_basic_cloudflare(case: TranslationTestCase, gateway: Gateway, provider: SharedProvider) -> None: + assert_translation(case, gateway, provider) diff --git a/tests/integration/translation/decisions/basic/test_decisions_basic_openrouter.py b/tests/integration/translation/decisions/basic/test_decisions_basic_openrouter.py new file mode 100644 index 00000000000..a74cb23f6b4 --- /dev/null +++ b/tests/integration/translation/decisions/basic/test_decisions_basic_openrouter.py @@ -0,0 +1,11 @@ +import pytest +from integration._support.client import Gateway +from integration._support.provider import SharedProvider +from integration.translation.case import TranslationTestCase +from integration.translation.decisions.bases.openrouter import TYPESAFE_JEV_1_13_TEST_CASE +from integration.translation.runner import assert_translation + + +@pytest.mark.parametrize("case", [TYPESAFE_JEV_1_13_TEST_CASE], ids=lambda case: case.id) +def test_decisions_basic_openrouter(case: TranslationTestCase, gateway: Gateway, provider: SharedProvider) -> None: + assert_translation(case, gateway, provider) diff --git a/tests/integration/translation/decisions/basic/test_decisions_basic_perplexity.py b/tests/integration/translation/decisions/basic/test_decisions_basic_perplexity.py new file mode 100644 index 00000000000..e341d24ba53 --- /dev/null +++ b/tests/integration/translation/decisions/basic/test_decisions_basic_perplexity.py @@ -0,0 +1,11 @@ +import pytest +from integration._support.client import Gateway +from integration._support.provider import SharedProvider +from integration.translation.case import TranslationTestCase +from integration.translation.decisions.bases.perplexity import PPLX_DECIDER_V1_27B_TEST_CASE +from integration.translation.runner import assert_translation + + +@pytest.mark.parametrize("case", [PPLX_DECIDER_V1_27B_TEST_CASE], ids=lambda case: case.id) +def test_decisions_basic_perplexity(case: TranslationTestCase, gateway: Gateway, provider: SharedProvider) -> None: + assert_translation(case, gateway, provider) diff --git a/tests/integration/translation/decisions/basic/test_decisions_basic_strands_decider.py b/tests/integration/translation/decisions/basic/test_decisions_basic_strands_decider.py new file mode 100644 index 00000000000..cae7ca70c78 --- /dev/null +++ b/tests/integration/translation/decisions/basic/test_decisions_basic_strands_decider.py @@ -0,0 +1,11 @@ +import pytest +from integration._support.client import Gateway +from integration._support.provider import SharedProvider +from integration.translation.case import TranslationTestCase +from integration.translation.decisions.bases.strands_decider import STRANDS_DECIDER_2B_HOBSON_V19_TEST_CASE +from integration.translation.runner import assert_translation + + +@pytest.mark.parametrize("case", [STRANDS_DECIDER_2B_HOBSON_V19_TEST_CASE], ids=lambda case: case.id) +def test_decisions_basic_strands_decider(case: TranslationTestCase, gateway: Gateway, provider: SharedProvider) -> None: + assert_translation(case, gateway, provider) diff --git a/tests/integration/translation/decisions/basic/test_decisions_basic_typesafe.py b/tests/integration/translation/decisions/basic/test_decisions_basic_typesafe.py new file mode 100644 index 00000000000..302636e6fb4 --- /dev/null +++ b/tests/integration/translation/decisions/basic/test_decisions_basic_typesafe.py @@ -0,0 +1,11 @@ +import pytest +from integration._support.client import Gateway +from integration._support.provider import SharedProvider +from integration.translation.case import TranslationTestCase +from integration.translation.decisions.bases.typesafe import JEV_1_13_0_TEST_CASE +from integration.translation.runner import assert_translation + + +@pytest.mark.parametrize("case", [JEV_1_13_0_TEST_CASE], ids=lambda case: case.id) +def test_decisions_basic_typesafe(case: TranslationTestCase, gateway: Gateway, provider: SharedProvider) -> None: + assert_translation(case, gateway, provider) diff --git a/tests/llm_translation/base_llm_unit_tests.py b/tests/llm_translation/base_llm_unit_tests.py index 4e2eecd2f31..e744b6275a8 100644 --- a/tests/llm_translation/base_llm_unit_tests.py +++ b/tests/llm_translation/base_llm_unit_tests.py @@ -5,8 +5,6 @@ import sys from typing import Any, Dict, List from unittest.mock import MagicMock, Mock, patch import os -from litellm._uuid import uuid -import time import base64 import inspect @@ -31,36 +29,6 @@ sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "..", -def _usage_format_tests(usage: litellm.Usage): - """ - OpenAI prompt caching - - prompt_tokens = sum of non-cache hit tokens + cache-hit tokens - - total_tokens = prompt_tokens + completion_tokens - - Example - ``` - "usage": { - "prompt_tokens": 2006, - "completion_tokens": 300, - "total_tokens": 2306, - "prompt_tokens_details": { - "cached_tokens": 1920 - }, - "completion_tokens_details": { - "reasoning_tokens": 0 - } - # ANTHROPIC_ONLY # - "cache_creation_input_tokens": 0 - } - ``` - """ - print(f"usage={usage}") - assert usage.total_tokens == usage.prompt_tokens + usage.completion_tokens - - if usage.prompt_tokens_details is not None: - assert usage.prompt_tokens > usage.prompt_tokens_details.cached_tokens - - class BaseLLMChatTest(ABC): """ Abstract base test class that enforces a common test across all test classes. diff --git a/tests/llm_translation/test_bedrock_moonshot.py b/tests/llm_translation/test_bedrock_moonshot.py index ac1a8363845..cf5e554e91b 100644 --- a/tests/llm_translation/test_bedrock_moonshot.py +++ b/tests/llm_translation/test_bedrock_moonshot.py @@ -12,7 +12,6 @@ This test suite verifies: """ from base_llm_unit_tests import BaseLLMChatTest -import json import litellm @@ -37,7 +36,3 @@ class TestBedrockMoonshotInvoke(BaseLLMChatTest): return { "model": "bedrock/invoke/moonshot.kimi-k2-thinking", } - - -class TestBedrockMoonshotToolCalling: - """Unit tests for tool calling functionality.""" diff --git a/tests/llm_translation/test_gemini.py b/tests/llm_translation/test_gemini.py index e30f97df192..d05afe92626 100644 --- a/tests/llm_translation/test_gemini.py +++ b/tests/llm_translation/test_gemini.py @@ -16,7 +16,6 @@ import json class TestGoogleAIStudioGemini(BaseLLMChatTest): - test_tool_call_no_arguments = None test_async_pdf_handling_with_file_id = None test_content_list_handling = None test_developer_role_translation = None diff --git a/tests/local_testing/test_caching.py b/tests/local_testing/test_caching.py index 0c68720473d..b8c7a4c189b 100644 --- a/tests/local_testing/test_caching.py +++ b/tests/local_testing/test_caching.py @@ -1744,7 +1744,7 @@ async def test_redis_proxy_batch_redis_get_cache(): from litellm.caching.caching import Cache, DualCache from litellm.proxy._types import UserAPIKeyAuth - from litellm.proxy.hooks.batch_redis_get import _PROXY_BatchRedisRequests + from litellm.proxy.hooks.batch_redis_get import PROXY_BatchRedisRequests litellm.cache = Cache( type="redis", @@ -1755,7 +1755,7 @@ async def test_redis_proxy_batch_redis_get_cache(): ) batch_redis_get_obj = ( - _PROXY_BatchRedisRequests() + PROXY_BatchRedisRequests() ) # overrides the .async_get_cache method user_api_key_cache = DualCache() diff --git a/tests/local_testing/test_lunary.py b/tests/local_testing/test_lunary.py index c2f67a137c0..4312abcc0a4 100644 --- a/tests/local_testing/test_lunary.py +++ b/tests/local_testing/test_lunary.py @@ -2,7 +2,6 @@ import io import litellm -from litellm import completion litellm.failure_callback = ["lunary"] litellm.success_callback = ["lunary"] diff --git a/tests/local_testing/test_streaming.py b/tests/local_testing/test_streaming.py index 76301f69c3f..c6582754c3a 100644 --- a/tests/local_testing/test_streaming.py +++ b/tests/local_testing/test_streaming.py @@ -22,7 +22,6 @@ load_dotenv() import litellm from litellm import ( AuthenticationError, - BadRequestError, ModelResponse, RateLimitError, acompletion, @@ -913,14 +912,6 @@ def test_openai_stream_options_call_text_completion() -> None: -# # test on together ai completion call - starcoder - - -# # test on together ai completion call - starcoder - - - - #### Test Function calling + streaming #### diff --git a/tests/local_testing/test_text_completion.py b/tests/local_testing/test_text_completion.py index 69064dfd1f4..d62b141d636 100644 --- a/tests/local_testing/test_text_completion.py +++ b/tests/local_testing/test_text_completion.py @@ -1,4 +1,3 @@ -import asyncio import json import os from types import MappingProxyType diff --git a/tests/logging_callback_tests/base_test.py b/tests/logging_callback_tests/base_test.py index cc894498f5e..11d89996474 100644 --- a/tests/logging_callback_tests/base_test.py +++ b/tests/logging_callback_tests/base_test.py @@ -11,7 +11,7 @@ from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler from litellm.utils import CustomStreamWrapper # test_example.py -from abc import ABC, abstractmethod +from abc import ABC class BaseLoggingCallbackTest(ABC): diff --git a/tests/openai_endpoints_tests/test_openai_batches_endpoint.py b/tests/openai_endpoints_tests/test_openai_batches_endpoint.py index bf36ce4d7d1..f488095aa12 100644 --- a/tests/openai_endpoints_tests/test_openai_batches_endpoint.py +++ b/tests/openai_endpoints_tests/test_openai_batches_endpoint.py @@ -9,7 +9,6 @@ from typing import Optional, List, Union from test_openai_files_endpoints import upload_file, delete_file import sys import time -from unittest.mock import patch BASE_URL = "http://localhost:4000" # Replace with your actual base URL diff --git a/tests/pass_through_unit_tests/base_anthropic_unified_messages_test.py b/tests/pass_through_unit_tests/base_anthropic_unified_messages_test.py index 40eda77949d..16055b5a29b 100644 --- a/tests/pass_through_unit_tests/base_anthropic_unified_messages_test.py +++ b/tests/pass_through_unit_tests/base_anthropic_unified_messages_test.py @@ -15,7 +15,6 @@ from litellm.llms.anthropic.pass_through.messages.handler import ( from typing import Optional from litellm.types.utils import StandardLoggingPayload from litellm.integrations.custom_logger import CustomLogger -from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler from litellm.router import Router import importlib from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM diff --git a/tests/pass_through_unit_tests/test_pass_through_unit_tests.py b/tests/pass_through_unit_tests/test_pass_through_unit_tests.py index f5b758843be..7a3109839d8 100644 --- a/tests/pass_through_unit_tests/test_pass_through_unit_tests.py +++ b/tests/pass_through_unit_tests/test_pass_through_unit_tests.py @@ -308,7 +308,7 @@ async def test_pass_through_request_logging_failure_with_stream( # Patch both the logging handler and the httpx client with ( patch( - "litellm.proxy.pass_through_endpoints.streaming_handler.PassThroughStreamingHandler._route_streaming_logging_to_handler", + "litellm.proxy.pass_through_endpoints.streaming_handler.PassThroughStreamingHandler.route_streaming_logging_to_handler", new=mock_logging_failure, ), patch( diff --git a/tests/test_litellm_rust/cache/test_python_cache.py b/tests/test_litellm_rust/cache/test_python_cache.py index 3a2e3b8c143..490643191ea 100644 --- a/tests/test_litellm_rust/cache/test_python_cache.py +++ b/tests/test_litellm_rust/cache/test_python_cache.py @@ -12,10 +12,10 @@ from litellm.caching.caching_handler import ( ) from litellm.proxy._types import Litellm_EntityType, UserAPIKeyAuth from litellm.proxy.hooks.model_max_budget_limiter import ( - _PROXY_VirtualKeyModelMaxBudgetLimiter, + PROXY_VirtualKeyModelMaxBudgetLimiter, model_budget_spend_cache_key, ) -from litellm.proxy.hooks.parallel_request_limiter_v3 import _PROXY_MaxParallelRequestsHandler_v3 +from litellm.proxy.hooks.parallel_request_limiter_v3 import PROXY_MaxParallelRequestsHandler_v3 from litellm.proxy.utils import InternalUsageCache from litellm.router_utils.add_retry_fallback_headers import get_hidden_params_dict from litellm.rust_bridge import runtime @@ -57,8 +57,8 @@ async def test_cache_hit_keeps_model_budget_spend_but_accounts_for_usage( monkeypatch.setenv("LITELLM_RUST", "1" if native else "0") litellm.cache = Cache() if legacy else _v2.Cache.memory() counters: Final = litellm.DualCache() - budget: Final = _PROXY_VirtualKeyModelMaxBudgetLimiter(counters) - limiter: Final = _PROXY_MaxParallelRequestsHandler_v3( + budget: Final = PROXY_VirtualKeyModelMaxBudgetLimiter(counters) + limiter: Final = PROXY_MaxParallelRequestsHandler_v3( InternalUsageCache(counters), model_group_resolver=lambda model: model ) recorder: Final = RecordingLogger() @@ -133,8 +133,8 @@ async def test_response_cache_backend_does_not_control_coordination( ) recording_server.expected_requests = 2 if backend == "disabled" else 1 counters: Final = litellm.DualCache() - budget: Final = _PROXY_VirtualKeyModelMaxBudgetLimiter(counters) - limiter: Final = _PROXY_MaxParallelRequestsHandler_v3( + budget: Final = PROXY_VirtualKeyModelMaxBudgetLimiter(counters) + limiter: Final = PROXY_MaxParallelRequestsHandler_v3( InternalUsageCache(counters), model_group_resolver=lambda model: model ) key_hash: Final = "b" * 64 diff --git a/tests/test_presidio_latency.py b/tests/test_presidio_latency.py index 40a2cc42b25..bb676ca051e 100644 --- a/tests/test_presidio_latency.py +++ b/tests/test_presidio_latency.py @@ -3,7 +3,7 @@ import aiohttp import pytest from unittest.mock import MagicMock, patch from litellm.proxy.guardrails.guardrail_hooks.presidio import ( - _OPTIONAL_PresidioPIIMasking, + OPTIONAL_PresidioPIIMasking, ) @@ -14,7 +14,7 @@ async def test_sanity_presidio_session_reuse_main_thread(): Verify that Presidio guardrail reuses sessions in the main thread. This ensures we don't break existing session pooling functionality. """ - presidio = _OPTIONAL_PresidioPIIMasking( + presidio = OPTIONAL_PresidioPIIMasking( mock_testing=True, presidio_analyzer_api_base="http://mock-analyzer", presidio_anonymizer_api_base="http://mock-anonymizer", @@ -28,9 +28,7 @@ async def test_sanity_presidio_session_reuse_main_thread(): session_creations += 1 original_init(self, *args, **kwargs) - with patch.object( - aiohttp.ClientSession, "__init__", side_effect=mocked_init, autospec=True - ): + with patch.object(aiohttp.ClientSession, "__init__", side_effect=mocked_init, autospec=True): for _ in range(10): async with presidio._get_session_iterator() as session: pass @@ -51,7 +49,7 @@ async def test_bug_presidio_session_explosion_background_thread_causes_latency() """ import threading - presidio = _OPTIONAL_PresidioPIIMasking( + presidio = OPTIONAL_PresidioPIIMasking( mock_testing=True, presidio_analyzer_api_base="http://mock-analyzer", presidio_anonymizer_api_base="http://mock-anonymizer", @@ -68,9 +66,7 @@ async def test_bug_presidio_session_explosion_background_thread_causes_latency() session_creations += 1 original_init(self, *args, **kwargs) - with patch.object( - aiohttp.ClientSession, "__init__", side_effect=mocked_init, autospec=True - ): + with patch.object(aiohttp.ClientSession, "__init__", side_effect=mocked_init, autospec=True): for _ in range(10): async with presidio._get_session_iterator() as session: pass diff --git a/tests/test_spend_logs.py b/tests/test_spend_logs.py index 310900c257d..57109ea01e5 100644 --- a/tests/test_spend_logs.py +++ b/tests/test_spend_logs.py @@ -2,7 +2,7 @@ import os # What this tests? ## Tests /spend endpoints. -import pytest, time, uuid, json +import pytest, uuid, json import asyncio import aiohttp @@ -56,32 +56,6 @@ async def chat_completion(session, key, model="gpt-3.5-turbo"): return await response.json() -async def chat_completion_high_traffic(session, key, model="gpt-3.5-turbo"): - url = "http://0.0.0.0:4000/chat/completions" - headers = { - "Authorization": f"Bearer {key}", - "Content-Type": "application/json", - } - data = { - "model": model, - "messages": [ - {"role": "system", "content": "You are a helpful assistant."}, - {"role": "user", "content": f"Hello! {uuid.uuid4()}"}, - ], - } - try: - async with session.post(url, headers=headers, json=data) as response: - status = response.status - response_text = await response.text() - - if status != 200: - raise Exception(f"Request did not return a 200 status code: {status}") - - return await response.json() - except Exception as e: - return None - - async def get_spend_logs(session, request_id=None, api_key=None): if api_key is not None: url = f"http://0.0.0.0:4000/spend/logs?api_key={api_key}" @@ -142,7 +116,7 @@ async def generate_team(session: aiohttp.ClientSession, org_id: str) -> dict: @pytest.mark.skip( - reason="Flaky in CI: /spend/logs?request_id=... returns 500 even after a 20s wait for the spend log to be written. Same write-then-read race against the spend logs DB as test_spend_logs. Spend-log accuracy is covered by tests/unit/proxy/spend_tracking/ and the proxy_spend_accuracy_tests CircleCI job." + reason="Flaky in CI: /spend/logs?request_id=... returns 500 even after a 20s wait for the spend log to be written. Spend-log accuracy is covered by tests/unit/proxy/spend_tracking/ and the proxy_spend_accuracy_tests CircleCI job." ) @pytest.mark.asyncio async def test_spend_logs_with_org_id(): diff --git a/tests/unit/batches/test_batch_utils.py b/tests/unit/batches/test_batch_utils.py index bd8a546c0ce..293c38e9f20 100644 --- a/tests/unit/batches/test_batch_utils.py +++ b/tests/unit/batches/test_batch_utils.py @@ -1184,7 +1184,7 @@ async def test_output_file_content_fetches_and_parses(monkeypatch): return type("R", (), {"content": b'{"a": 1}\n{"b": 2}'})() monkeypatch.setattr(files_main, "afile_content", fake_afile_content) - monkeypatch.setattr(cu, "_is_base64_encoded_unified_file_id", lambda fid: False) + monkeypatch.setattr(cu, "is_base64_encoded_unified_file_id", lambda fid: False) result = await bu._fetch_batch_output_file_content( _batch("file-out"), @@ -1217,7 +1217,7 @@ async def test_output_file_content_unified_file_id_extraction(monkeypatch): monkeypatch.setattr(files_main, "afile_content", fake_afile_content) monkeypatch.setattr( cu, - "_is_base64_encoded_unified_file_id", + "is_base64_encoded_unified_file_id", lambda fid: "litellm_proxy;llm_output_file_id,real-file-99;rest", ) diff --git a/tests/unit/caching/test_request_redis_batch_post_call.py b/tests/unit/caching/test_request_redis_batch_post_call.py index 8c8c9df5926..5de785e44db 100644 --- a/tests/unit/caching/test_request_redis_batch_post_call.py +++ b/tests/unit/caching/test_request_redis_batch_post_call.py @@ -34,7 +34,7 @@ from litellm.proxy.hooks.parallel_request_limiter_v3 import ( TOKEN_INCREMENT_SCRIPT, ParallelSlotAcquisition, RequestRateLimiterStash, - _PROXY_MaxParallelRequestsHandler_v3, + PROXY_MaxParallelRequestsHandler_v3, ) from litellm.proxy.spend_tracking.spend_counter_batch import PendingSpendIncrement from litellm.proxy.utils import InternalUsageCache @@ -99,10 +99,10 @@ def _names(client: FakeClient, index: int = 0) -> list[str]: return [command[0] for command in client.pipelines[index].commands] -def _limiter(redis_cache: FakeRedisCache) -> _PROXY_MaxParallelRequestsHandler_v3: +def _limiter(redis_cache: FakeRedisCache) -> PROXY_MaxParallelRequestsHandler_v3: dual_cache = DualCache() dual_cache.attach_redis_cache(redis_cache) - return _PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=InternalUsageCache(dual_cache=dual_cache)) + return PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=InternalUsageCache(dual_cache=dual_cache)) def _slot_stash(slot_id: str, *counter_keys: str) -> RequestRateLimiterStash: diff --git a/tests/unit/caching/test_request_redis_batch_pre_call.py b/tests/unit/caching/test_request_redis_batch_pre_call.py index 81a2a1718e9..69ef67e595b 100644 --- a/tests/unit/caching/test_request_redis_batch_pre_call.py +++ b/tests/unit/caching/test_request_redis_batch_pre_call.py @@ -18,14 +18,14 @@ from litellm._internal_context import current_service_target from litellm.caching.dual_cache import DualCache from litellm.caching.redis_batch import active_request_redis_batches, request_redis_batch_scope from litellm.proxy._types import LiteLLM_TeamTableCachedObj, LiteLLM_UserTable -from litellm.proxy.auth.auth_checks import _cache_team_object +from litellm.proxy.auth.auth_checks import cache_team_object from litellm.proxy.auth.auth_object_prefetch import _CacheEntry, _write_back, prefetch_identity_keys from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache from litellm.proxy.hooks.parallel_request_limiter_v3 import ( CHECK_AND_INCREMENT_BY_N_SCRIPT, RateLimitDescriptor, RateLimitUnverifiableError, - _PROXY_MaxParallelRequestsHandler_v3, + PROXY_MaxParallelRequestsHandler_v3, ) from litellm.proxy.utils import InternalUsageCache from litellm.router_utils.cooldown_cache import ROUTER_COOLDOWNS_TARGET, CooldownCache @@ -46,9 +46,9 @@ def sha_of(script: str) -> str: return hashlib.sha1(script.encode()).hexdigest() # noqa: S324 -def _limiter(redis_cache: FakeRedisCache, fail_closed: bool = False) -> _PROXY_MaxParallelRequestsHandler_v3: +def _limiter(redis_cache: FakeRedisCache, fail_closed: bool = False) -> PROXY_MaxParallelRequestsHandler_v3: dual_cache = DualCache() - limiter = _PROXY_MaxParallelRequestsHandler_v3( + limiter = PROXY_MaxParallelRequestsHandler_v3( internal_usage_cache=InternalUsageCache(dual_cache=dual_cache), fail_closed_resolver=lambda: fail_closed, ) @@ -64,7 +64,7 @@ def _descriptor(key: str, value: str, rpm: int) -> RateLimitDescriptor: return {"key": key, "value": value, "rate_limit": {"requests_per_unit": rpm}} -def _refunds(limiter: _PROXY_MaxParallelRequestsHandler_v3) -> list[tuple[str, float]]: +def _refunds(limiter: PROXY_MaxParallelRequestsHandler_v3) -> list[tuple[str, float]]: refund_script = limiter.window_guarded_token_increment_script assert isinstance(refund_script, AsyncMock) return [(call.kwargs["keys"][1], call.kwargs["args"][1]) for call in refund_script.await_args_list] @@ -1029,7 +1029,7 @@ async def test_a_team_refresh_inside_a_request_sends_its_set_and_alias_del_in_on proxy_logging_obj.internal_usage_cache = InternalUsageCache(dual_cache=usage_cache) team = LiteLLM_TeamTableCachedObj(team_id="t1", team_alias="alpha") with request_redis_batch_scope() as request: - await _cache_team_object("t1", team, cache, proxy_logging_obj) + await cache_team_object("t1", team, cache, proxy_logging_obj) assert [c[:2] for c in client.pipelines[0].commands] == [("SET", "team_id:t1"), ("DEL", "team_alias:alpha")], ( "the alias DEL must reach Redis before the refresh returns, or another request can refill memory from it" ) diff --git a/tests/unit/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py b/tests/unit/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py index eedf766ae92..5b2187211df 100644 --- a/tests/unit/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py +++ b/tests/unit/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py @@ -136,9 +136,7 @@ def test_convert_chat_completion_messages_to_responses_api_tool_result_with_imag function_call_output = item break - assert ( - function_call_output is not None - ), "function_call_output not found in response" + assert function_call_output is not None, "function_call_output not found in response" assert function_call_output["call_id"] == "call_abc123" # Check that the output is correctly transformed @@ -148,12 +146,8 @@ def test_convert_chat_completion_messages_to_responses_api_tool_result_with_imag image_item = output[0] # Should be transformed to Responses API format - assert ( - image_item["type"] == "input_image" - ), f"Expected type 'input_image', got '{image_item.get('type')}'" - assert ( - image_item["image_url"] == test_image_base64 - ), "image_url should be a flat string, not a nested object" + assert image_item["type"] == "input_image", f"Expected type 'input_image', got '{image_item.get('type')}'" + assert image_item["image_url"] == test_image_base64, "image_url should be a flat string, not a nested object" assert "detail" in image_item, "detail field should be present" print("✓ Tool result with image correctly transformed to Responses API format") @@ -215,9 +209,7 @@ def test_convert_chat_completion_messages_to_responses_api_tool_result_with_text function_call_output = item break - assert ( - function_call_output is not None - ), "function_call_output not found in response" + assert function_call_output is not None, "function_call_output not found in response" assert function_call_output["call_id"] == "call_abc123" # Check that the output is correctly transformed to use input_text, not output_text @@ -227,16 +219,12 @@ def test_convert_chat_completion_messages_to_responses_api_tool_result_with_text text_item = output[0] # Should be transformed to use input_text for tool results in Responses API format - assert ( - text_item["type"] == "input_text" - ), f"Expected type 'input_text' for tool result, got '{text_item.get('type')}'" - assert ( - text_item["text"] == "15 degrees" - ), f"Expected text '15 degrees', got '{text_item.get('text')}'" - - print( - "✓ Tool result with text correctly transformed to use input_text for Responses API format" + assert text_item["type"] == "input_text", ( + f"Expected type 'input_text' for tool result, got '{text_item.get('type')}'" ) + assert text_item["text"] == "15 degrees", f"Expected text '15 degrees', got '{text_item.get('text')}'" + + print("✓ Tool result with text correctly transformed to use input_text for Responses API format") def test_openai_responses_chunk_parser_reasoning_summary(): @@ -245,9 +233,7 @@ def test_openai_responses_chunk_parser_reasoning_summary(): ) from litellm.types.utils import Delta, ModelResponseStream, StreamingChoices - iterator = OpenAiResponsesToChatCompletionStreamIterator( - streaming_response=None, sync_stream=True - ) + iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True) chunk = { "delta": "**Compar", @@ -279,9 +265,7 @@ def test_chunk_parser_string_output_text_delta_produces_text(): ) from litellm.types.utils import ModelResponseStream - iterator = OpenAiResponsesToChatCompletionStreamIterator( - streaming_response=None, sync_stream=True - ) + iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True) chunk = {"type": "response.output_text.delta", "delta": "literal text"} @@ -302,9 +286,7 @@ def test_chunk_parser_enum_output_text_delta_produces_text(): from litellm.types.llms.openai import ResponsesAPIStreamEvents from litellm.types.utils import ModelResponseStream - iterator = OpenAiResponsesToChatCompletionStreamIterator( - streaming_response=None, sync_stream=True - ) + iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True) chunk = {"type": ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA, "delta": "enum text"} @@ -325,9 +307,7 @@ def test_chunk_parser_function_call_added_produces_tool_use(): from litellm.types.llms.openai import ResponsesAPIStreamEvents from litellm.types.utils import ModelResponseStream - iterator = OpenAiResponsesToChatCompletionStreamIterator( - streaming_response=None, sync_stream=True - ) + iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True) chunk = { "type": ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED, @@ -412,9 +392,7 @@ Tomorrow will bring its petitions and promises, but for now the city breathes slow and wide, and I learn to carry this small calm home.""" - output_text = ResponseOutputText( - annotations=[], text=poem_text, type="output_text", logprobs=[] - ) + output_text = ResponseOutputText(annotations=[], text=poem_text, type="output_text", logprobs=[]) output_message = ResponseOutputMessage( id="msg_04c8021b8b3188a00068e9ae0b92f4819dac64d85b4abb67ec", content=[output_text], @@ -426,9 +404,7 @@ and I learn to carry this small calm home.""" # Create usage information usage = ResponseAPIUsage( input_tokens=16, - input_tokens_details=InputTokensDetails( - audio_tokens=None, cached_tokens=0, text_tokens=None - ), + input_tokens_details=InputTokensDetails(audio_tokens=None, cached_tokens=0, text_tokens=None), output_tokens=195, output_tokens_details=OutputTokensDetails(reasoning_tokens=0, text_tokens=None), total_tokens=211, @@ -777,11 +753,7 @@ def test_recover_output_items_merges_text_only_items_at_distinct_indices(): ] ) - recovered = ( - LiteLLMResponsesTransformationHandler._recover_output_items_from_raw_sse( - raw_sse - ) - ) + recovered = LiteLLMResponsesTransformationHandler._recover_output_items_from_raw_sse(raw_sse) assert len(recovered) == 2 assert recovered[0]["id"] == "msg_item_0" @@ -919,9 +891,7 @@ def test_transform_request_system_only_message_maps_to_system_input_item(): { "type": "message", "role": "system", - "content": [ - {"type": "input_text", "text": "You are a helpful assistant."} - ], + "content": [{"type": "input_text", "text": "You are a helpful assistant."}], } ] # System content lives in input only; not duplicated into instructions. @@ -993,9 +963,7 @@ def test_transform_request_single_char_keys_not_matched(): assert result_correct.get("metadata") == {"user_id": "123"} assert result_correct.get("previous_response_id") == "resp_abc" - print( - "✓ Single-character keys are not incorrectly matched to metadata/previous_response_id" - ) + print("✓ Single-character keys are not incorrectly matched to metadata/previous_response_id") # ============================================================================= @@ -1015,9 +983,7 @@ def test_message_done_does_not_emit_is_finished(): OpenAiResponsesToChatCompletionStreamIterator, ) - iterator = OpenAiResponsesToChatCompletionStreamIterator( - streaming_response=None, sync_stream=True - ) + iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True) chunk = { "type": "response.output_item.done", @@ -1029,9 +995,9 @@ def test_message_done_does_not_emit_is_finished(): # After the fix, message completion should NOT set finish_reason # ModelResponseStream doesn't have is_finished - check finish_reason instead assert len(result.choices) > 0, "result should have choices" - assert ( - result.choices[0].finish_reason is None or result.choices[0].finish_reason == "" - ), "message completion should not emit finish_reason" + assert result.choices[0].finish_reason is None or result.choices[0].finish_reason == "", ( + "message completion should not emit finish_reason" + ) def test_response_completed_emits_is_finished(): @@ -1043,9 +1009,7 @@ def test_response_completed_emits_is_finished(): OpenAiResponsesToChatCompletionStreamIterator, ) - iterator = OpenAiResponsesToChatCompletionStreamIterator( - streaming_response=None, sync_stream=True - ) + iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True) chunk = {"type": "response.completed"} @@ -1053,9 +1017,7 @@ def test_response_completed_emits_is_finished(): # response.completed should emit finish_reason='stop' assert len(result.choices) > 0, "result should have choices" - assert ( - result.choices[0].finish_reason == "stop" - ), "response.completed should emit finish_reason='stop'" + assert result.choices[0].finish_reason == "stop", "response.completed should emit finish_reason='stop'" def test_response_completed_with_function_calls_emits_tool_calls_finish_reason(): @@ -1074,9 +1036,7 @@ def test_response_completed_with_function_calls_emits_tool_calls_finish_reason() OpenAiResponsesToChatCompletionStreamIterator, ) - iterator = OpenAiResponsesToChatCompletionStreamIterator( - streaming_response=None, sync_stream=True - ) + iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True) # Simulate a response.completed event with function_call in output # This matches what Azure/OpenAI sends for gpt-5.1-codex-mini and similar models @@ -1102,9 +1062,9 @@ def test_response_completed_with_function_calls_emits_tool_calls_finish_reason() # response.completed with function_call should emit finish_reason='tool_calls' assert len(result.choices) > 0, "result should have choices" - assert ( - result.choices[0].finish_reason == "tool_calls" - ), "response.completed with function_call output should emit finish_reason='tool_calls'" + assert result.choices[0].finish_reason == "tool_calls", ( + "response.completed with function_call output should emit finish_reason='tool_calls'" + ) def test_response_completed_with_message_only_emits_stop_finish_reason(): @@ -1115,9 +1075,7 @@ def test_response_completed_with_message_only_emits_stop_finish_reason(): OpenAiResponsesToChatCompletionStreamIterator, ) - iterator = OpenAiResponsesToChatCompletionStreamIterator( - streaming_response=None, sync_stream=True - ) + iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True) # Simulate a response.completed event with only message output chunk = { @@ -1141,9 +1099,9 @@ def test_response_completed_with_message_only_emits_stop_finish_reason(): # response.completed with only message should emit finish_reason='stop' assert len(result.choices) > 0, "result should have choices" - assert ( - result.choices[0].finish_reason == "stop" - ), "response.completed with only message output should emit finish_reason='stop'" + assert result.choices[0].finish_reason == "stop", ( + "response.completed with only message output should emit finish_reason='stop'" + ) def test_response_completed_preserves_usage_with_cached_tokens(): @@ -1159,9 +1117,7 @@ def test_response_completed_preserves_usage_with_cached_tokens(): OpenAiResponsesToChatCompletionStreamIterator, ) - iterator = OpenAiResponsesToChatCompletionStreamIterator( - streaming_response=None, sync_stream=True - ) + iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True) chunk = { "type": "response.completed", @@ -1190,18 +1146,12 @@ def test_response_completed_preserves_usage_with_cached_tokens(): result = iterator.chunk_parser(chunk) assert result.usage is not None, "usage should be set on response.completed chunk" - assert ( - result.usage.prompt_tokens == 1226 - ), "prompt_tokens should map from input_tokens" - assert ( - result.usage.completion_tokens == 5 - ), "completion_tokens should map from output_tokens" - assert ( - result.usage.prompt_tokens_details is not None - ), "prompt_tokens_details should be set" - assert ( - result.usage.prompt_tokens_details.cached_tokens == 1024 - ), "cached_tokens should be preserved from input_tokens_details" + assert result.usage.prompt_tokens == 1226, "prompt_tokens should map from input_tokens" + assert result.usage.completion_tokens == 5, "completion_tokens should map from output_tokens" + assert result.usage.prompt_tokens_details is not None, "prompt_tokens_details should be set" + assert result.usage.prompt_tokens_details.cached_tokens == 1024, ( + "cached_tokens should be preserved from input_tokens_details" + ) def test_function_call_done_emits_is_finished(): @@ -1215,9 +1165,7 @@ def test_function_call_done_emits_is_finished(): OpenAiResponsesToChatCompletionStreamIterator, ) - iterator = OpenAiResponsesToChatCompletionStreamIterator( - streaming_response=None, sync_stream=True - ) + iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True) chunk = { "type": "response.output_item.done", @@ -1237,9 +1185,9 @@ def test_function_call_done_emits_is_finished(): "output_item.done for function_call must not emit finish_reason; " "response.completed is responsible for the terminal finish_reason" ) - assert not result.choices[ - 0 - ].delta.tool_calls, "output_item.done for function_call must not include a duplicate tool_calls delta" + assert not result.choices[0].delta.tool_calls, ( + "output_item.done for function_call must not include a duplicate tool_calls delta" + ) def test_text_plus_tool_calls_sequence(): @@ -1254,9 +1202,7 @@ def test_text_plus_tool_calls_sequence(): OpenAiResponsesToChatCompletionStreamIterator, ) - iterator = OpenAiResponsesToChatCompletionStreamIterator( - streaming_response=None, sync_stream=True - ) + iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True) # Simulate the sequence from OpenAI Responses API chunks = [ @@ -1295,28 +1241,23 @@ def test_text_plus_tool_calls_sequence(): # Check message done (index 2) does NOT have finish_reason set message_done_result = results[2] assert len(message_done_result.choices) > 0, "message done should have choices" - assert ( - message_done_result.choices[0].finish_reason is None - or message_done_result.choices[0].finish_reason == "" - ), "message done should not have finish_reason" + assert message_done_result.choices[0].finish_reason is None or message_done_result.choices[0].finish_reason == "", ( + "message done should not have finish_reason" + ) # Check function_call done (index 5) does NOT have finish_reason set # (response.completed is responsible for the terminal finish_reason) function_done_result = results[5] - assert ( - len(function_done_result.choices) > 0 - ), "function_call done should have choices" - assert ( - function_done_result.choices[0].finish_reason is None - ), "output_item.done for function_call must not emit finish_reason" + assert len(function_done_result.choices) > 0, "function_call done should have choices" + assert function_done_result.choices[0].finish_reason is None, ( + "output_item.done for function_call must not emit finish_reason" + ) # Check response.completed (index 6) has finish_reason='stop' # (the mock chunk has no nested 'response' data, so has_function_calls is False → 'stop') completed_result = results[6] assert len(completed_result.choices) > 0, "response.completed should have choices" - assert ( - completed_result.choices[0].finish_reason == "stop" - ), "response.completed should have finish_reason='stop'" + assert completed_result.choices[0].finish_reason == "stop", "response.completed should have finish_reason='stop'" # ============================================================================= @@ -1333,7 +1274,11 @@ def test_developer_message_content_uses_input_text(): assert instructions is None assert input_items == [ - {"type": "message", "role": "developer", "content": [{"type": "input_text", "text": "Always answer in French."}]} + { + "type": "message", + "role": "developer", + "content": [{"type": "input_text", "text": "Always answer in French."}], + } ] @@ -1395,9 +1340,7 @@ def test_tool_message_output_uses_input_text_not_output_text(): output = function_call_output["output"] assert isinstance(output, list), f"output should be a list, got {type(output)}" assert len(output) == 1 - assert ( - output[0]["type"] == "input_text" - ), f"Expected input_text, got {output[0].get('type')}" + assert output[0]["type"] == "input_text", f"Expected input_text, got {output[0].get('type')}" assert output[0]["text"] == '{"temperature": 15, "condition": "sunny"}' print("✓ Tool message output correctly uses input_text type") @@ -1582,13 +1525,9 @@ def test_map_reasoning_effort_adds_summary_detailed(monkeypatch): assert result is not None, f"Result should not be None for effort={effort}" assert result["effort"] == effort, f"Effort should be {effort}" - assert ( - "summary" not in result - ), f"Summary should NOT be present by default for effort={effort}" + assert "summary" not in result, f"Summary should NOT be present by default for effort={effort}" - print( - f"✓ reasoning_effort='{effort}' correctly maps to effort='{effort}' (no summary by default)" - ) + print(f"✓ reasoning_effort='{effort}' correctly maps to effort='{effort}' (no summary by default)") # Test 2: With flag enabled - summary IS added litellm.reasoning_auto_summary = True @@ -1598,9 +1537,9 @@ def test_map_reasoning_effort_adds_summary_detailed(monkeypatch): assert result is not None, f"Result should not be None for effort={effort}" assert result["effort"] == effort, f"Effort should be {effort}" - assert ( - result["summary"] == "detailed" - ), f"Summary should be 'detailed' when flag is enabled for effort={effort}" + assert result["summary"] == "detailed", ( + f"Summary should be 'detailed' when flag is enabled for effort={effort}" + ) print( f"✓ reasoning_effort='{effort}' correctly maps to effort='{effort}', summary='detailed' (flag enabled)" @@ -1611,9 +1550,7 @@ def test_map_reasoning_effort_adds_summary_detailed(monkeypatch): monkeypatch.setenv("LITELLM_REASONING_AUTO_SUMMARY", "true") result = handler.map_reasoning_effort("high") - assert ( - result["summary"] == "detailed" - ), "Summary should be 'detailed' when env var is enabled" + assert result["summary"] == "detailed", "Summary should be 'detailed' when env var is enabled" print("✓ LITELLM_REASONING_AUTO_SUMMARY env var works correctly") # Test 4: Dict input is passed through as-is (no modification) @@ -1627,9 +1564,7 @@ def test_map_reasoning_effort_adds_summary_detailed(monkeypatch): assert result_dict["summary"] == "custom_summary" print("✓ Dict input is passed through without modification") - print( - "✓ All reasoning_effort behaviors work correctly with flag/env var control" - ) + print("✓ All reasoning_effort behaviors work correctly with flag/env var control") finally: # Restore original values @@ -1705,9 +1640,7 @@ def test_transform_response_preserves_annotations(): # Create usage information usage = ResponseAPIUsage( input_tokens=10, - input_tokens_details=InputTokensDetails( - audio_tokens=None, cached_tokens=0, text_tokens=None - ), + input_tokens_details=InputTokensDetails(audio_tokens=None, cached_tokens=0, text_tokens=None), output_tokens=20, output_tokens_details=OutputTokensDetails(reasoning_tokens=0, text_tokens=None), total_tokens=30, @@ -1794,13 +1727,9 @@ def test_transform_response_preserves_annotations(): assert choice.message.content == "Here is some information with citations." # Check that annotations are preserved - assert hasattr( - choice.message, "annotations" - ), "Message should have annotations attribute" + assert hasattr(choice.message, "annotations"), "Message should have annotations attribute" assert choice.message.annotations is not None, "Annotations should not be None" - assert ( - len(choice.message.annotations) == 2 - ), f"Expected 2 annotations, got {len(choice.message.annotations)}" + assert len(choice.message.annotations) == 2, f"Expected 2 annotations, got {len(choice.message.annotations)}" # Verify annotation content annotation1 = choice.message.annotations[0] @@ -1822,9 +1751,7 @@ def test_transform_response_preserves_annotations(): assert result.usage.completion_tokens == 20 assert result.usage.total_tokens == 30 - print( - "✓ Annotations from Responses API are correctly preserved in Chat Completions format" - ) + print("✓ Annotations from Responses API are correctly preserved in Chat Completions format") def test_apply_patch_tool_call_converted_to_chat_completion_tool_call(): @@ -1989,9 +1916,7 @@ def test_multi_tool_call_stream_no_premature_finish(): OpenAiResponsesToChatCompletionStreamIterator, ) - iterator = OpenAiResponsesToChatCompletionStreamIterator( - streaming_response=None, sync_stream=True - ) + iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True) chunks = [ # 0: response created @@ -2067,12 +1992,10 @@ def test_multi_tool_call_stream_no_premature_finish(): r = results[done_idx] assert r is not None, f"{label}: chunk_parser must return a result" assert len(r.choices) > 0, f"{label}: result must have choices" - assert ( - r.choices[0].finish_reason is None - ), f"{label}: output_item.done must not emit finish_reason (stream would terminate prematurely)" - assert not r.choices[ - 0 - ].delta.tool_calls, ( + assert r.choices[0].finish_reason is None, ( + f"{label}: output_item.done must not emit finish_reason (stream would terminate prematurely)" + ) + assert not r.choices[0].delta.tool_calls, ( f"{label}: output_item.done must not include a duplicate tool_calls delta" ) @@ -2084,12 +2007,8 @@ def test_multi_tool_call_stream_no_premature_finish(): r = results[added_idx] if r is not None and r.choices and r.choices[0].delta.tool_calls: tc = r.choices[0].delta.tool_calls[0] - assert ( - tc.function.name == expected_name - ), f"output_item.added for {expected_name}: tool_call name mismatch" - assert ( - tc.id == expected_call_id - ), f"output_item.added for {expected_name}: call_id mismatch" + assert tc.function.name == expected_name, f"output_item.added for {expected_name}: tool_call name mismatch" + assert tc.id == expected_call_id, f"output_item.added for {expected_name}: call_id mismatch" # 3. argument delta events (indices 2 and 5) should carry arguments for delta_idx, expected_args, label in [ @@ -2099,17 +2018,15 @@ def test_multi_tool_call_stream_no_premature_finish(): r = results[delta_idx] if r is not None and r.choices and r.choices[0].delta.tool_calls: tc = r.choices[0].delta.tool_calls[0] - assert ( - tc.function.arguments == expected_args - ), f"{label}: argument delta mismatch" + assert tc.function.arguments == expected_args, f"{label}: argument delta mismatch" # 4. Only response.completed (index 7) emits the terminal finish_reason completed_result = results[7] assert completed_result is not None, "response.completed must return a result" assert len(completed_result.choices) > 0, "response.completed must have choices" - assert ( - completed_result.choices[0].finish_reason == "tool_calls" - ), "response.completed with function_call outputs must emit finish_reason='tool_calls'" + assert completed_result.choices[0].finish_reason == "tool_calls", ( + "response.completed with function_call outputs must emit finish_reason='tool_calls'" + ) # 5. No chunk before the last one should have finish_reason set for idx, r in enumerate(results[:-1]): @@ -2119,9 +2036,7 @@ def test_multi_tool_call_stream_no_premature_finish(): f"— only response.completed should terminate the stream" ) - print( - "✓ Multi-tool-call stream completes without premature finish_reason termination" - ) + print("✓ Multi-tool-call stream completes without premature finish_reason termination") # ============================================================================= @@ -2202,16 +2117,13 @@ def test_streaming_parallel_tool_calls_have_distinct_indices(): ] for chunk in chunks: - result = OpenAiResponsesToChatCompletionStreamIterator.translate_responses_chunk_to_openai_stream( - chunk - ) + result = OpenAiResponsesToChatCompletionStreamIterator.translate_responses_chunk_to_openai_stream(chunk) expected_index = chunk["output_index"] for choice in result.choices: if choice.delta.tool_calls: for tc in choice.delta.tool_calls: assert tc.index == expected_index, ( - f"Event {chunk['type']}: expected tool_call.index={expected_index}, " - f"got {tc.index}" + f"Event {chunk['type']}: expected tool_call.index={expected_index}, got {tc.index}" ) @@ -2339,9 +2251,7 @@ def test_parallel_tool_calls_comprehensive_streaming_integration(): }, ] - iterator = OpenAiResponsesToChatCompletionStreamIterator( - streaming_response=None, sync_stream=True - ) + iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True) results = [iterator.chunk_parser(chunk) for chunk in chunks] # 1. output_item.done events (indices 4 and 8) must NOT emit finish_reason @@ -2353,9 +2263,7 @@ def test_parallel_tool_calls_comprehensive_streaming_integration(): f"{label}: output_item.done must not emit finish_reason " f"(would prematurely terminate stream before subsequent tool calls arrive)" ) - assert not r.choices[ - 0 - ].delta.tool_calls, ( + assert not r.choices[0].delta.tool_calls, ( f"{label}: output_item.done must not emit a duplicate tool_calls delta" ) @@ -2389,19 +2297,15 @@ def test_parallel_tool_calls_comprehensive_streaming_integration(): for tc in tool_calls: if tc.function and tc.function.arguments: idx = tc.index - assembled_args[idx] = ( - assembled_args.get(idx, "") + tc.function.arguments - ) + assembled_args[idx] = assembled_args.get(idx, "") + tc.function.arguments # delta 1 = '{"path":' + delta 2 = '"/etc/foo"}' → '{"path":"/etc/foo"}' assert assembled_args.get(0) == '{"path":"/etc/foo"}', ( - f"Assembled args for index 0 (read_file): " - f"expected '{{\"path\":\"/etc/foo\"}}', got '{assembled_args.get(0)}'" + f"Assembled args for index 0 (read_file): expected '{{\"path\":\"/etc/foo\"}}', got '{assembled_args.get(0)}'" ) # delta 1 = '{"path":' + delta 2 = '"/tmp"}' → '{"path":"/tmp"}' assert assembled_args.get(1) == '{"path":"/tmp"}', ( - f"Assembled args for index 1 (list_dir): " - f"expected '{{\"path\":\"/tmp\"}}', got '{assembled_args.get(1)}'" + f"Assembled args for index 1 (list_dir): expected '{{\"path\":\"/tmp\"}}', got '{assembled_args.get(1)}'" ) # 4. Stream terminates with exactly one finish event, at the final response.completed chunk @@ -2410,16 +2314,13 @@ def test_parallel_tool_calls_comprehensive_streaming_integration(): for i, r in enumerate(results) if r is not None and r.choices and r.choices[0].finish_reason ] - assert ( - len(finish_events) == 1 - ), f"Expected exactly 1 finish event, got {len(finish_events)}: {finish_events}" + assert len(finish_events) == 1, f"Expected exactly 1 finish event, got {len(finish_events)}: {finish_events}" assert finish_events[0][0] == len(chunks) - 1, ( - f"Finish event must be at the last chunk (index {len(chunks) - 1}), " - f"but was at index {finish_events[0][0]}" + f"Finish event must be at the last chunk (index {len(chunks) - 1}), but was at index {finish_events[0][0]}" + ) + assert finish_events[0][1] == "tool_calls", ( + f"Terminal finish_reason must be 'tool_calls', got '{finish_events[0][1]}'" ) - assert ( - finish_events[0][1] == "tool_calls" - ), f"Terminal finish_reason must be 'tool_calls', got '{finish_events[0][1]}'" # 5. Parallel tool calls have distinct indices matching output_index (0 and 1) # Collect indices from output_item.added chunks only (they carry the call id) @@ -2435,9 +2336,7 @@ def test_parallel_tool_calls_comprehensive_streaming_integration(): 1, }, f"Parallel tool calls must have distinct indices {{0, 1}}, got: {set(added_tool_call_indices)}" - print( - "✓ Parallel tool calls with split argument deltas stream correctly end-to-end" - ) + print("✓ Parallel tool calls with split argument deltas stream correctly end-to-end") def test_map_optional_params_preserves_reasoning_summary(): @@ -2461,9 +2360,7 @@ def test_map_optional_params_preserves_reasoning_summary(): } responses_api_request = ResponsesAPIOptionalRequestParams() - handler._map_optional_params_to_responses_api_request( - optional_params, responses_api_request - ) + handler._map_optional_params_to_responses_api_request(optional_params, responses_api_request) # Verify reasoning_effort dict with summary was fully preserved assert "reasoning" in responses_api_request @@ -2736,9 +2633,7 @@ def test_reasoning_items_non_streaming_round_trip(): ) usage = ResponseAPIUsage( input_tokens=10, - input_tokens_details=InputTokensDetails( - audio_tokens=None, cached_tokens=0, text_tokens=None - ), + input_tokens_details=InputTokensDetails(audio_tokens=None, cached_tokens=0, text_tokens=None), output_tokens=20, output_tokens_details=OutputTokensDetails(reasoning_tokens=0, text_tokens=None), total_tokens=30, @@ -2802,9 +2697,7 @@ def test_reasoning_items_non_streaming_round_trip(): assert len(result.choices) == 1 msg = result.choices[0].message - assert ( - msg.reasoning_content == summary_text - ), "reasoning_content should equal summary text" + assert msg.reasoning_content == summary_text, "reasoning_content should equal summary text" assert msg.reasoning_items is not None, "reasoning_items should be set" assert len(msg.reasoning_items) == 1 @@ -2829,13 +2722,9 @@ def test_reasoning_items_non_streaming_round_trip(): # The reasoning input item must appear before the assistant message item types = [item.get("type") for item in input_items] - assert ( - "reasoning" in types - ), "reasoning input item must be emitted for the assistant turn" + assert "reasoning" in types, "reasoning input item must be emitted for the assistant turn" - reasoning_input = next( - item for item in input_items if item.get("type") == "reasoning" - ) + reasoning_input = next(item for item in input_items if item.get("type") == "reasoning") assert reasoning_input["id"] == "rs_test001" assert reasoning_input["encrypted_content"] == encrypted assert reasoning_input["summary"][0]["text"] == summary_text @@ -2843,13 +2732,9 @@ def test_reasoning_items_non_streaming_round_trip(): # reasoning item must come before the assistant message item reasoning_idx = types.index("reasoning") assistant_msg_idx = next( - i - for i, item in enumerate(input_items) - if item.get("type") == "message" and item.get("role") == "assistant" + i for i, item in enumerate(input_items) if item.get("type") == "message" and item.get("role") == "assistant" ) - assert ( - reasoning_idx < assistant_msg_idx - ), "reasoning input item must precede the assistant message item" + assert reasoning_idx < assistant_msg_idx, "reasoning input item must precede the assistant message item" def test_reasoning_items_streaming_emitted_on_response_completed(): @@ -2862,9 +2747,7 @@ def test_reasoning_items_streaming_emitted_on_response_completed(): OpenAiResponsesToChatCompletionStreamIterator, ) - iterator = OpenAiResponsesToChatCompletionStreamIterator( - streaming_response=None, sync_stream=True - ) + iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True) encrypted = "gAAAAABpw5xyz987FAKE==" summary_text = "**Reasoning summary**\n\nModel thought about this carefully." @@ -2908,16 +2791,14 @@ def test_reasoning_items_streaming_emitted_on_response_completed(): assert result.choices[0].finish_reason == "stop" # reasoning_items must be on the delta - assert ( - getattr(delta, "reasoning_items", None) is not None - ), "reasoning_items must be present on the response.completed delta" + assert getattr(delta, "reasoning_items", None) is not None, ( + "reasoning_items must be present on the response.completed delta" + ) assert len(delta.reasoning_items) == 1 ri = delta.reasoning_items[0] assert ri["type"] == "reasoning" assert ri["id"] == "rs_stream001" - assert ( - ri["encrypted_content"] == encrypted - ), "encrypted_content must be preserved in streaming" + assert ri["encrypted_content"] == encrypted, "encrypted_content must be preserved in streaming" assert ri["summary"][0]["text"] == summary_text @@ -2944,9 +2825,7 @@ def test_streaming_function_call_tool_id_for_degenerate_call_id(): "arguments": "", }, } - out = OpenAiResponsesToChatCompletionStreamIterator.translate_responses_chunk_to_openai_stream( - chunk - ) + out = OpenAiResponsesToChatCompletionStreamIterator.translate_responses_chunk_to_openai_stream(chunk) tool_calls = out.model_dump()["choices"][0]["delta"]["tool_calls"] assert tool_calls, "expected a tool_call chunk in the streaming delta" return tool_calls[0]["id"] @@ -2965,9 +2844,7 @@ def test_streaming_chunks_share_one_chat_completion_id(): OpenAiResponsesToChatCompletionStreamIterator, ) - iterator = OpenAiResponsesToChatCompletionStreamIterator( - streaming_response=None, sync_stream=True - ) + iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True) events = [ {"type": "response.created", "response": {"id": "resp_abc", "output": []}}, {"type": "response.output_text.delta", "delta": "Hel"}, @@ -2983,12 +2860,10 @@ def test_streaming_chunks_share_one_chat_completion_id(): assert len(set(ids)) == 1, f"streamed chunks carried different ids: {ids}" assert ids[0], "streamed chunks carried an empty id" - other_stream = OpenAiResponsesToChatCompletionStreamIterator( - streaming_response=None, sync_stream=True + other_stream = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True) + assert other_stream.chunk_parser(events[1]).id != ids[0], ( + "a separate stream must get its own id, not a process-wide one" ) - assert ( - other_stream.chunk_parser(events[1]).id != ids[0] - ), "a separate stream must get its own id, not a process-wide one" @pytest.mark.asyncio @@ -2999,9 +2874,7 @@ def test_streaming_chunks_share_one_chat_completion_id(): ({"include_usage": True}, None), ], ) -async def test_acompletion_bridge_normalizes_stream_options_on_the_wire( - stream_options, expected_wire_stream_options -): +async def test_acompletion_bridge_normalizes_stream_options_on_the_wire(stream_options, expected_wire_stream_options): """include_usage must be stripped from the /v1/responses body; include_obfuscation must survive as a dict.""" from unittest.mock import AsyncMock @@ -3077,9 +2950,7 @@ def test_chunk_parser_custom_tool_call_stream_sequence(): OpenAiResponsesToChatCompletionStreamIterator, ) - iterator = OpenAiResponsesToChatCompletionStreamIterator( - streaming_response=None, sync_stream=True - ) + iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True) added = iterator.chunk_parser( { @@ -3157,9 +3028,7 @@ def test_chunk_parser_remaps_tool_call_indices_sequentially(): OpenAiResponsesToChatCompletionStreamIterator, ) - iterator = OpenAiResponsesToChatCompletionStreamIterator( - streaming_response=None, sync_stream=True - ) + iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True) first = iterator.chunk_parser( { @@ -3736,9 +3605,7 @@ def _make_incomplete_responses_api_response( created_at=1760144904, error=None, incomplete_details=( - {"reason": incomplete_reason} - if incomplete_reason is not None or empty_incomplete_details - else None + {"reason": incomplete_reason} if incomplete_reason is not None or empty_incomplete_details else None ), instructions=None, metadata={}, @@ -3758,13 +3625,9 @@ def _make_incomplete_responses_api_response( truncation="disabled", usage=ResponseAPIUsage( input_tokens=37, - input_tokens_details=InputTokensDetails( - audio_tokens=None, cached_tokens=0, text_tokens=None - ), + input_tokens_details=InputTokensDetails(audio_tokens=None, cached_tokens=0, text_tokens=None), output_tokens=16, - output_tokens_details=OutputTokensDetails( - reasoning_tokens=16, text_tokens=None - ), + output_tokens_details=OutputTokensDetails(reasoning_tokens=16, text_tokens=None), total_tokens=53, cost=None, ), @@ -3814,9 +3677,7 @@ def _call_transform_response( def test_transform_response_incomplete_reasoning_only_returns_empty_length_choice(): handler = LiteLLMResponsesTransformationHandler() - raw_response = _make_incomplete_responses_api_response( - "max_output_tokens", [_make_reasoning_only_output_item()] - ) + raw_response = _make_incomplete_responses_api_response("max_output_tokens", [_make_reasoning_only_output_item()]) result = _call_transform_response(handler, raw_response) @@ -3835,9 +3696,7 @@ def test_transform_response_incomplete_reasoning_only_returns_empty_length_choic def test_transform_response_incomplete_content_filter_maps_finish_reason(): handler = LiteLLMResponsesTransformationHandler() - raw_response = _make_incomplete_responses_api_response( - "content_filter", [_make_reasoning_only_output_item()] - ) + raw_response = _make_incomplete_responses_api_response("content_filter", [_make_reasoning_only_output_item()]) result = _call_transform_response(handler, raw_response) @@ -3860,11 +3719,7 @@ def test_transform_response_completed_with_reasonless_incomplete_details_keeps_s handler = LiteLLMResponsesTransformationHandler() output_message = ResponseOutputMessage( id="msg_complete", - content=[ - ResponseOutputText( - annotations=[], text="full answer", type="output_text", logprobs=[] - ) - ], + content=[ResponseOutputText(annotations=[], text="full answer", type="output_text", logprobs=[])], role="assistant", status="completed", type="message", @@ -3886,11 +3741,7 @@ def test_transform_response_incomplete_partial_text_overrides_finish_reason_to_l handler = LiteLLMResponsesTransformationHandler() output_message = ResponseOutputMessage( id="msg_partial", - content=[ - ResponseOutputText( - annotations=[], text="partial answer", type="output_text", logprobs=[] - ) - ], + content=[ResponseOutputText(annotations=[], text="partial answer", type="output_text", logprobs=[])], role="assistant", status="incomplete", type="message", @@ -3912,9 +3763,7 @@ def test_response_incomplete_stream_event_emits_length_and_usage(): OpenAiResponsesToChatCompletionStreamIterator, ) - iterator = OpenAiResponsesToChatCompletionStreamIterator( - streaming_response=None, sync_stream=True - ) + iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True) chunk = { "type": "response.incomplete", @@ -3955,9 +3804,7 @@ def test_response_incomplete_stream_event_content_filter_maps_finish_reason(): OpenAiResponsesToChatCompletionStreamIterator, ) - iterator = OpenAiResponsesToChatCompletionStreamIterator( - streaming_response=None, sync_stream=True - ) + iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True) chunk = { "type": "response.incomplete", @@ -3979,9 +3826,7 @@ def test_response_incomplete_stream_event_without_details_defaults_to_length(): OpenAiResponsesToChatCompletionStreamIterator, ) - iterator = OpenAiResponsesToChatCompletionStreamIterator( - streaming_response=None, sync_stream=True - ) + iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True) chunk = { "type": "response.incomplete", @@ -4061,9 +3906,7 @@ def test_thinking_only_assistant_turn_still_sends_its_reasoning(): { "role": "assistant", "content": None, - "thinking_blocks": [ - {"type": "thinking", "thinking": "August in Denver is dry.", "signature": "sig1"} - ], + "thinking_blocks": [{"type": "thinking", "thinking": "August in Denver is dry.", "signature": "sig1"}], }, {"role": "user", "content": "Why?"}, ] @@ -4089,9 +3932,7 @@ def test_stored_reasoning_items_win_over_thinking_blocks(): "summary": [{"type": "summary_text", "text": "August in Denver is dry."}], } ], - "thinking_blocks": [ - {"type": "thinking", "thinking": "August in Denver is dry.", "signature": "rs_real"} - ], + "thinking_blocks": [{"type": "thinking", "thinking": "August in Denver is dry.", "signature": "rs_real"}], }, ] @@ -4581,9 +4422,7 @@ def test_convert_chat_completion_messages_to_responses_api_drops_prompt_cache_br { "role": "tool", "tool_call_id": "call_1", - "content": [ - {"type": "text", "text": "Tool result", "prompt_cache_breakpoint": cache_breakpoint} - ], + "content": [{"type": "text", "text": "Tool result", "prompt_cache_breakpoint": cache_breakpoint}], }, ], ) @@ -4897,3 +4736,104 @@ def test_every_bridged_chunk_after_response_created_carries_the_served_service_t relayed = [iterator.chunk_parser(event).model_dump().get("service_tier") for event in events] assert relayed == ["default"] * len(events), relayed + + +def test_convert_chat_completion_messages_to_responses_api_keeps_prompt_cache_breakpoint_on_unknown_block(): + """The hook marks the last block of its target message, so a message ending in a block the bridge + cannot map reaches the stringify path and has to keep the marker there.""" + from litellm.completion_extras.litellm_responses_transformation.transformation import ( + LiteLLMResponsesTransformationHandler, + ) + + handler = LiteLLMResponsesTransformationHandler() + breakpoint_marker = {"mode": "explicit"} + + messages = [ + { + "role": "user", + "content": [ + {"type": "text", "text": "describe this"}, + { + "type": "input_audio", + "input_audio": {"data": "Zm9v", "format": "wav"}, + "prompt_cache_breakpoint": breakpoint_marker, + }, + ], + }, + ] + + response, _ = handler.convert_chat_completion_messages_to_responses_api( + messages, keep_prompt_cache_breakpoints=True + ) + + content = response[0]["content"] + assert [block["type"] for block in content] == ["input_text", "input_text"] + assert content[1]["prompt_cache_breakpoint"] == breakpoint_marker + + +_HAND_WRITTEN_PROMPT_CACHE_BREAKPOINT_BLOCKS: Final = ( + {"type": "text", "text": "a string marker", "prompt_cache_breakpoint": "explicit"}, + {"type": "text", "text": "an unknown mode", "prompt_cache_breakpoint": {"mode": "bogus"}}, + { + "type": "input_audio", + "input_audio": {"data": "Zm9v", "format": "wav"}, + "prompt_cache_breakpoint": ["explicit"], + }, + {"type": "text", "text": "unsupported ttl", "prompt_cache_breakpoint": {"mode": "explicit", "ttl": "1h"}}, + {"type": "text", "text": "unknown key", "prompt_cache_breakpoint": {"mode": "explicit", "scope": "all"}}, + {"type": "text", "text": "supported ttl", "prompt_cache_breakpoint": {"mode": "explicit", "ttl": "30m"}}, + {"type": "text", "text": "well formed", "prompt_cache_breakpoint": {"mode": "explicit"}}, +) + + +def test_convert_chat_completion_messages_to_responses_api_drops_malformed_prompt_cache_breakpoint_under_drop_params(): + """OpenAI's Responses API answered "Supported values are: '30m'" for a 1h breakpoint ttl on 2026-10-07, + so an unsupported ttl drops the marker as a unit while an unknown key is dropped from a valid one.""" + handler = LiteLLMResponsesTransformationHandler() + messages = [{"role": "user", "content": list(_HAND_WRITTEN_PROMPT_CACHE_BREAKPOINT_BLOCKS)}] + + response, _ = handler.convert_chat_completion_messages_to_responses_api( + messages, drop_params=True, keep_prompt_cache_breakpoints=True + ) + + content = response[0]["content"] + assert [block.get("prompt_cache_breakpoint") for block in content] == [ + None, + None, + None, + None, + {"mode": "explicit"}, + {"mode": "explicit", "ttl": "30m"}, + {"mode": "explicit"}, + ] + assert all("prompt_cache_breakpoint" not in block for block in content[:4]) + + +def test_convert_chat_completion_messages_to_responses_api_keeps_malformed_prompt_cache_breakpoint_by_default(): + handler = LiteLLMResponsesTransformationHandler() + messages = [{"role": "user", "content": list(_HAND_WRITTEN_PROMPT_CACHE_BREAKPOINT_BLOCKS)}] + + response, _ = handler.convert_chat_completion_messages_to_responses_api( + messages, keep_prompt_cache_breakpoints=True + ) + + content = response[0]["content"] + assert [block["prompt_cache_breakpoint"] for block in content] == [ + block["prompt_cache_breakpoint"] for block in _HAND_WRITTEN_PROMPT_CACHE_BREAKPOINT_BLOCKS + ] + + +def test_transform_request_drop_params_in_litellm_params_gates_the_prompt_cache_breakpoint_carry(): + handler = LiteLLMResponsesTransformationHandler() + messages = [{"role": "user", "content": [{"type": "text", "text": "hi", "prompt_cache_breakpoint": "explicit"}]}] + + result = handler.transform_request( + model="gpt-6.1-sol", + messages=messages, + optional_params={}, + litellm_params={"drop_params": True}, + headers={}, + litellm_logging_obj=Mock(), + ) + + assert "prompt_cache_breakpoint" not in result["input"][0]["content"][0] diff --git a/tests/unit/containers/test_container_proxy_ownership.py b/tests/unit/containers/test_container_proxy_ownership.py index dd21faa1059..780ae9bf6d1 100644 --- a/tests/unit/containers/test_container_proxy_ownership.py +++ b/tests/unit/containers/test_container_proxy_ownership.py @@ -425,7 +425,7 @@ async def test_should_validate_owner_and_forward_decoded_id_for_multipart_upload async def base_process_llm_request(self, **kwargs): return captured["data"] - async def _handle_llm_api_exception(self, **kwargs): + async def handle_llm_api_exception(self, **kwargs): raise kwargs["e"] monkeypatch.setattr( @@ -499,7 +499,7 @@ async def test_should_forward_decoded_container_id_for_proxy_retrieve(monkeypatc async def base_process_llm_request(self, **kwargs): return captured["data"] - async def _handle_llm_api_exception(self, **kwargs): + async def handle_llm_api_exception(self, **kwargs): raise kwargs["e"] monkeypatch.setattr(endpoints, "ProxyBaseLLMRequestProcessing", FakeProcessor) @@ -554,7 +554,7 @@ async def test_should_record_container_owner_inside_create_endpoint(monkeypatch) async def base_process_llm_request(self, **kwargs): return response - async def _handle_llm_api_exception(self, **kwargs): + async def handle_llm_api_exception(self, **kwargs): raise kwargs["e"] record_owner = AsyncMock(return_value=response) @@ -608,7 +608,7 @@ async def test_should_not_route_owner_record_errors_through_llm_error_handler( async def base_process_llm_request(self, **kwargs): return _container("cntr_provider") - async def _handle_llm_api_exception(self, **kwargs): + async def handle_llm_api_exception(self, **kwargs): raise AssertionError("ownership errors should not use LLM error handler") monkeypatch.setattr(endpoints, "ProxyBaseLLMRequestProcessing", FakeProcessor) @@ -668,7 +668,7 @@ async def test_should_return_response_when_owner_recording_raises_unexpected( async def base_process_llm_request(self, **kwargs): return created - async def _handle_llm_api_exception(self, **kwargs): + async def handle_llm_api_exception(self, **kwargs): raise AssertionError("upstream-create errors only") monkeypatch.setattr(endpoints, "ProxyBaseLLMRequestProcessing", FakeProcessor) @@ -772,7 +772,7 @@ async def test_should_forward_decoded_container_id_for_proxy_delete(monkeypatch) async def base_process_llm_request(self, **kwargs): return captured["data"] - async def _handle_llm_api_exception(self, **kwargs): + async def handle_llm_api_exception(self, **kwargs): raise kwargs["e"] monkeypatch.setattr(endpoints, "ProxyBaseLLMRequestProcessing", FakeProcessor) diff --git a/tests/unit/decisions/test_main.py b/tests/unit/decisions/test_main.py index 684248b3561..820451ed943 100644 --- a/tests/unit/decisions/test_main.py +++ b/tests/unit/decisions/test_main.py @@ -6,6 +6,7 @@ from collections.abc import Mapping from types import MappingProxyType from typing import Final +import httpx import pytest import respx @@ -17,6 +18,7 @@ from litellm.types.decisions import ( DecisionsResponse, DecisionsUsage, NoulAnswer, + OpenAIDecisionResponse, ScoreAnswer, ) @@ -29,6 +31,8 @@ _QUESTIONS: Final[Mapping[str, object]] = MappingProxyType( ) _INPUT_TOKENS: Final[int] = 367 _OUTPUT_TOKENS: Final[int] = 3 +_CACHED_TOKENS: Final[int] = 256 +_CACHE_WRITE_TOKENS: Final[int] = 64 _RESPONSE: Final[Mapping[str, object]] = { "model": "jev-1.13", "answers": { @@ -79,6 +83,21 @@ _PROVIDERS: Final[tuple[tuple[str, str, str, str], ...]] = ( ), ) +_OPENAI_RESPONSE: Final[Mapping[str, object]] = { + "model": "gpt-6-luna", + "answers": [ + {"type": "predicate", "name": "is_defect", "probability": 0.9}, + {"type": "refusal", "name": "sentiment"}, + ], + "usage": { + "input_tokens": _INPUT_TOKENS, + "input_tokens_details": {"cached_tokens": _CACHED_TOKENS, "cache_write_tokens": _CACHE_WRITE_TOKENS}, + "output_tokens": _OUTPUT_TOKENS, + "output_tokens_details": {"reasoning_tokens": 0}, + "total_tokens": _INPUT_TOKENS + _OUTPUT_TOKENS, + }, +} + class _RecordingLogger(CustomLogger): def __init__(self) -> None: @@ -252,10 +271,7 @@ def test_decisions_cost_uses_litellm_token_pricing() -> None: answers={}, usage=DecisionsUsage(input_tokens=_INPUT_TOKENS, output_tokens=_OUTPUT_TOKENS), ) - response._hidden_params = { - "model": "perplexity/pplx-decider-v1-27b", - "custom_llm_provider": "perplexity", - } + response.set_hidden_params({"model": "perplexity/pplx-decider-v1-27b", "custom_llm_provider": "perplexity"}) cost: Final = litellm.completion_cost(completion_response=response) perplexity_cost: Final = litellm.model_cost["perplexity/pplx-decider-v1-27b"] @@ -267,6 +283,46 @@ def test_decisions_cost_uses_litellm_token_pricing() -> None: assert cost == pytest.approx(expected_cost) +@pytest.mark.parametrize( + "response", + ( + DecisionsResponse( + model="jev-latest", + answers={}, + usage=DecisionsUsage( + input_tokens=_INPUT_TOKENS, + output_tokens=_OUTPUT_TOKENS, + cached_tokens=_CACHED_TOKENS, + cache_write_tokens=_CACHE_WRITE_TOKENS, + ), + ), + OpenAIDecisionResponse.model_validate(_OPENAI_RESPONSE), + ), + ids=("systemone", "openai"), +) +def test_custom_token_pricing_bills_cached_decisions_input_tokens_once( + response: DecisionsResponse | OpenAIDecisionResponse, +) -> None: + cost: Final = litellm.completion_cost( + completion_response=response, + model="gpt-6-luna", + custom_llm_provider="openai", + custom_cost_per_token={ + "input_cost_per_token": 1.0, + "output_cost_per_token": 2.0, + "cache_read_input_token_cost": 0.1, + "cache_creation_input_token_cost": 1.25, + }, + ) + + assert cost == pytest.approx( + (_INPUT_TOKENS - _CACHED_TOKENS - _CACHE_WRITE_TOKENS) * 1.0 + + _CACHED_TOKENS * 0.1 + + _CACHE_WRITE_TOKENS * 1.25 + + _OUTPUT_TOKENS * 2.0 + ) + + def test_decisions_response_hidden_params_getter_preserves_mutable_identity() -> None: response: Final = DecisionsResponse(model="decider", answers={}, usage=None) @@ -308,7 +364,7 @@ async def test_decisions_cost_is_in_standard_logging_object(respx_mock: respx.Mo @pytest.mark.asyncio async def test_unknown_provider_is_rejected_before_http(respx_mock: respx.MockRouter) -> None: - with pytest.raises(litellm.BadRequestError, match="Supported providers"): + with pytest.raises(litellm.BadRequestError, match="LLM Provider NOT provided"): await litellm.adecisions( model="unknown/jev-1.13", state="review", @@ -320,17 +376,79 @@ async def test_unknown_provider_is_rejected_before_http(respx_mock: respx.MockRo @pytest.mark.asyncio -async def test_empty_custom_provider_is_rejected_before_http(respx_mock: respx.MockRouter) -> None: - with pytest.raises(litellm.BadRequestError, match="Supported providers"): +async def test_provider_without_decisions_support_is_rejected_before_http(respx_mock: respx.MockRouter) -> None: + with pytest.raises(litellm.BadRequestError, match=r"Unknown Decisions provider 'anthropic'\. Supported providers"): + await litellm.adecisions( + model="anthropic/claude-sonnet-4-5", + state="review", + questions={"is_defect": {"type": "noul", "instructions": "Is this a defect?"}}, + api_key="caller-key", + ) + + assert len(respx_mock.calls) == 0 + + +@pytest.mark.asyncio +async def test_empty_custom_provider_falls_back_to_the_model_prefix(respx_mock: respx.MockRouter) -> None: + route: Final = respx_mock.post("https://api.perplexity.ai/v1/decisions").respond(json=_RESPONSE) + + response: Final = await litellm.adecisions( + model="perplexity/pplx-decider-v1-27b", + state="review", + questions={"is_defect": {"type": "noul", "instructions": "Is this a defect?"}}, + api_key="caller-key", + custom_llm_provider="", + ) + + assert route.called + assert json.loads(respx_mock.calls[0].request.content)["model"] == "pplx-decider-v1-27b" + assert response._hidden_params["custom_llm_provider"] == "perplexity" + + +@pytest.mark.asyncio +async def test_upstream_reply_without_answers_is_a_server_error(respx_mock: respx.MockRouter) -> None: + respx_mock.post("https://api.perplexity.ai/v1/decisions").respond( + json={"model": "pplx-decider-v1-27b", "usage": {"input_tokens": 10, "output_tokens": 0}} + ) + + with pytest.raises(litellm.InternalServerError, match="unexpected response"): await litellm.adecisions( model="perplexity/pplx-decider-v1-27b", state="review", questions={"is_defect": {"type": "noul", "instructions": "Is this a defect?"}}, api_key="caller-key", - custom_llm_provider="", ) - assert len(respx_mock.calls) == 0 + +@pytest.mark.asyncio +async def test_router_sends_the_openrouter_deployment_key_when_the_provider_is_already_resolved( + monkeypatch: pytest.MonkeyPatch, + respx_mock: respx.MockRouter, +) -> None: + monkeypatch.delenv("OPENROUTER_API_KEY", raising=False) + router: Final = litellm.Router( + model_list=[ + { + "model_name": "jev", + "litellm_params": { + "model": "openrouter/typesafe/jev-1.13", + "api_key": "deployment-key", + "api_base": "https://egress.example/openrouter", + }, + } + ] + ) + upstream: Final = respx_mock.post("https://egress.example/openrouter/alpha/decisions").respond(json=_RESPONSE) + + await router.adecisions( + model="jev", + state="review", + questions={"is_defect": {"type": "noul", "instructions": "Is this a defect?"}}, + ) + + assert upstream.called + assert respx_mock.calls[0].request.headers["authorization"] == "Bearer deployment-key" + assert json.loads(respx_mock.calls[0].request.content)["model"] == "typesafe/jev-1.13" @pytest.mark.asyncio @@ -361,6 +479,20 @@ def test_upstream_bad_request_maps_to_litellm_error(respx_mock: respx.MockRouter ) +def test_unreachable_upstream_maps_to_a_connection_error(respx_mock: respx.MockRouter) -> None: + respx_mock.post("https://api.perplexity.ai/v1/decisions").mock( + side_effect=httpx.ConnectError("Cannot connect to host api.perplexity.ai:443") + ) + + with pytest.raises(litellm.APIConnectionError, match="PerplexityException - Cannot connect to host"): + litellm.decisions( + model="perplexity/pplx-decider-v1-27b", + state="review", + questions={"is_defect": {"type": "noul", "instructions": "Is this a defect?"}}, + api_key="caller-key", + ) + + def test_server_key_is_sent_to_an_explicit_api_base( monkeypatch: pytest.MonkeyPatch, respx_mock: respx.MockRouter, @@ -413,7 +545,7 @@ async def test_cloudflare_clef_resolves_model_and_response_envelope( "state": "review", "questions": {"is_defect": {"type": "noul", "instructions": "Is this a defect?"}}, } - assert response.answers == DecisionsResponse.model_validate(_RESPONSE).answers + assert response.answers == {"is_defect": NoulAnswer(type="noul", noul=0.9)} assert response._hidden_params["model"] == "cloudflare/@cf/cloudflare/clef" @@ -599,3 +731,194 @@ async def test_strands_decider_provider_resolution_and_router_dispatch( assert provider_resolution[:2] == ("strands-decider-2B-hobson-v19", "strands_decider") assert route.called assert response.model == _STRANDS_RESPONSE["model"] + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("api_base", "url"), + ( + (None, "https://api.openai.com/v1/decisions"), + ("https://gateway.example/v1", "https://gateway.example/v1/decisions"), + ), +) +async def test_openai_decisions_translate_systemone_to_the_openai_wire_contract_and_back( + api_base: str | None, + url: str, + monkeypatch: pytest.MonkeyPatch, + respx_mock: respx.MockRouter, +) -> None: + monkeypatch.setattr(litellm, "api_base", None) + monkeypatch.delenv("OPENAI_BASE_URL", raising=False) + monkeypatch.delenv("OPENAI_API_BASE", raising=False) + route: Final = respx_mock.post(url).respond(json=_OPENAI_RESPONSE) + + response: Final = await litellm.adecisions( + model="openai/gpt-6-luna", + state="The package arrived broken.", + questions={ + "is_defect": {"type": "noul", "instructions": "Is this a defect?"}, + "sentiment": { + "type": "choice", + "instructions": "How does the customer feel?", + "criteria": {"positive": None, "negative": "unhappy"}, + }, + }, + api_key="caller-key", + api_base=api_base, + ) + + assert route.called + request: Final = respx_mock.calls[0].request + assert request.headers["authorization"] == "Bearer caller-key" + assert json.loads(request.content) == { + "model": "gpt-6-luna", + "input": "The package arrived broken.", + "questions": [ + {"type": "predicate", "name": "is_defect", "instructions": "Is this a defect?"}, + { + "type": "choice", + "name": "sentiment", + "instructions": "How does the customer feel?", + "choices": [{"value": "positive"}, {"value": "negative", "description": "unhappy"}], + }, + ], + } + assert response.answers == {"is_defect": NoulAnswer(type="noul", noul=0.9)} + assert response.hidden_params["custom_llm_provider"] == "openai" + luna_cost: Final = litellm.model_cost["gpt-6-luna"] + expected_cost: Final = ( + (_INPUT_TOKENS - _CACHED_TOKENS - _CACHE_WRITE_TOKENS) * float(luna_cost["input_cost_per_token"]) + + _CACHED_TOKENS * float(luna_cost["cache_read_input_token_cost"]) + + _CACHE_WRITE_TOKENS * float(luna_cost["cache_creation_input_token_cost"]) + + _OUTPUT_TOKENS * float(luna_cost["output_cost_per_token"]) + ) + assert expected_cost > 0 + assert litellm.completion_cost(completion_response=response) == pytest.approx(expected_cost) + + +@pytest.mark.parametrize( + ("settings", "env", "url", "authorization"), + ( + ({"openai_key": "sdk-key"}, {}, "https://api.openai.com/v1/decisions", "Bearer sdk-key"), + ( + {"api_key": "global-key", "openai_key": "sdk-key"}, + {"OPENAI_API_KEY": "env-key"}, + "https://api.openai.com/v1/decisions", + "Bearer global-key", + ), + ( + {}, + {"OPENAI_API_KEY": "env-key", "OPENAI_API_BASE": "https://legacy.example/v1"}, + "https://legacy.example/v1/decisions", + "Bearer env-key", + ), + ( + {"api_base": "https://sdk.example"}, + {"OPENAI_API_KEY": "env-key", "OPENAI_BASE_URL": "https://env.example"}, + "https://sdk.example/v1/decisions", + "Bearer env-key", + ), + ), + ids=("openai_key", "api_key_before_env", "openai_api_base_env", "api_base_before_env"), +) +def test_openai_decisions_use_the_same_settings_as_other_openai_calls( + settings: Mapping[str, str], + env: Mapping[str, str], + url: str, + authorization: str, + monkeypatch: pytest.MonkeyPatch, + respx_mock: respx.MockRouter, +) -> None: + for name in ("api_key", "openai_key", "api_base"): + monkeypatch.setattr(litellm, name, settings.get(name)) + for name in ("OPENAI_API_KEY", "OPENAI_BASE_URL", "OPENAI_API_BASE"): + monkeypatch.delenv(name, raising=False) + for name, value in env.items(): + monkeypatch.setenv(name, value) + route: Final = respx_mock.post(url).respond(json=_OPENAI_RESPONSE) + + litellm.decisions( + model="openai/gpt-6-luna", + state="review", + questions={"is_defect": {"type": "noul", "instructions": "Is this a defect?"}}, + ) + + assert route.call_count == 1 + assert route.calls[0].request.headers["authorization"] == authorization + + +@pytest.mark.asyncio +async def test_openai_format_calls_to_a_systemone_provider_get_openai_format_answers( + respx_mock: respx.MockRouter, +) -> None: + route: Final = respx_mock.post("https://api.typesafe.ai/v1/systemone").respond(json=_RESPONSE) + + response: Final = await litellm.adecisions( + model="typesafe/jev-1.13", + input="review", + questions=[ + {"type": "predicate", "name": "is_defect", "instructions": "Is this a defect?"}, + { + "type": "choice", + "name": "sentiment", + "instructions": "Tone?", + "choices": [{"value": "positive"}, {"value": "negative"}], + }, + ], + api_key="caller-key", + ) + + assert route.called + assert tuple(json.loads(respx_mock.calls[0].request.content)["questions"]) == ("is_defect", "sentiment") + assert isinstance(response, OpenAIDecisionResponse) + assert [answer.model_dump(mode="json") for answer in response.answers] == [ + {"type": "predicate", "name": "is_defect", "probability": 0.9}, + { + "type": "choice", + "name": "sentiment", + "choice": "positive", + "probabilities": [{"value": "positive", "probability": 0.8}, {"value": "negative", "probability": 0.2}], + "confidence": 0.8, + }, + ] + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("model", "request_kwargs", "message"), + ( + ( + "typesafe/jev-1.13", + { + "input": [ + { + "role": "user", + "content": [{"type": "input_image", "image_url": "data:image/png;base64,AA=="}], + } + ], + "questions": [{"type": "predicate", "instructions": "Is this a defect?"}], + }, + "cannot serve this request", + ), + ( + "openai/gpt-6-luna", + { + "state": "review", + "input": "review", + "questions": [{"type": "predicate", "instructions": "Is this a defect?"}], + }, + "not both", + ), + ), + ids=("image_to_systemone_provider", "state_and_input"), +) +async def test_requests_a_provider_cannot_serve_are_rejected_before_http( + respx_mock: respx.MockRouter, + model: str, + request_kwargs: Mapping[str, object], + message: str, +) -> None: + with pytest.raises(litellm.BadRequestError, match=message): + await litellm.adecisions(model=model, api_key="caller-key", **request_kwargs) + + assert len(respx_mock.calls) == 0 diff --git a/tests/unit/enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py b/tests/unit/enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py index 78545d3fb62..de6a5601d5c 100644 --- a/tests/unit/enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py +++ b/tests/unit/enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py @@ -1711,7 +1711,7 @@ async def test_initialize_remaining_budget_metrics_exception_handling( "litellm.proxy.management_endpoints.team_endpoints.get_paginated_teams" ) as mock_get_teams, patch( - "litellm.proxy.management_endpoints.key_management_endpoints._list_key_helper" + "litellm.proxy.management_endpoints.key_management_endpoints.list_key_helper" ) as mock_list_keys, ): # Make get_paginated_teams raise an exception @@ -1786,7 +1786,7 @@ async def test_initialize_api_key_budget_metrics(prometheus_logger): with ( patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch( - "litellm.proxy.management_endpoints.key_management_endpoints._list_key_helper" + "litellm.proxy.management_endpoints.key_management_endpoints.list_key_helper" ) as mock_list_keys, ): # Create mock key data with proper datetime objects for budget_reset_at diff --git a/tests/unit/enterprise/proxy/auth/test_user_api_key_auth.py b/tests/unit/enterprise/proxy/auth/test_user_api_key_auth.py index d21bb046549..a93e1368c49 100644 --- a/tests/unit/enterprise/proxy/auth/test_user_api_key_auth.py +++ b/tests/unit/enterprise/proxy/auth/test_user_api_key_auth.py @@ -22,9 +22,7 @@ async def test_enterprise_custom_auth_mode_on(): mock_user_auth = AsyncMock(return_value={"user_id": "test-user"}) request = MagicMock(spec=Request) - with patch( - "litellm_enterprise.proxy.proxy_server.custom_auth_settings", {"mode": "on"} - ): + with patch("litellm_enterprise.proxy.proxy_server.custom_auth_settings", {"mode": "on"}): result = await enterprise_custom_auth(request, "test-api-key", mock_user_auth) assert result == {"user_id": "test-user"} mock_user_auth.assert_called_once_with(request, "test-api-key") @@ -36,9 +34,7 @@ async def test_enterprise_custom_auth_mode_auto_with_error(): mock_user_auth = AsyncMock(side_effect=Exception("Auth failed")) request = MagicMock(spec=Request) - with patch( - "litellm_enterprise.proxy.proxy_server.custom_auth_settings", {"mode": "auto"} - ): + with patch("litellm_enterprise.proxy.proxy_server.custom_auth_settings", {"mode": "auto"}): result = await enterprise_custom_auth(request, "test-api-key", mock_user_auth) assert result is None mock_user_auth.assert_called_once_with(request, "test-api-key") @@ -61,9 +57,7 @@ async def test_enterprise_custom_auth_returns_string(): patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), ): # Verify the key is correctly handled in _user_api_key_auth_builder - with patch( - "litellm.proxy.auth.resolvers.store.IdentityStore._resolve_key" - ) as mock_get_key_object: + with patch("litellm.proxy.auth.resolvers.store.IdentityStore._resolve_key") as mock_get_key_object: mock_get_key_object.return_value = UserAPIKeyAuth( token="sk-test-key", user_role="internal_user", @@ -72,10 +66,10 @@ async def test_enterprise_custom_auth_returns_string(): ) # Call _user_api_key_auth_builder with the returned key - from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth_builder try: - auth_obj = await _user_api_key_auth_builder( + auth_obj = await user_api_key_auth_builder( request=request, api_key="my-custom-key", azure_api_key_header="", diff --git a/tests/unit/enterprise/proxy/hooks/test_managed_files.py b/tests/unit/enterprise/proxy/hooks/test_managed_files.py index 7d8c8b5c425..11d9dbec49c 100644 --- a/tests/unit/enterprise/proxy/hooks/test_managed_files.py +++ b/tests/unit/enterprise/proxy/hooks/test_managed_files.py @@ -11,7 +11,7 @@ from litellm.caching import DualCache from litellm.proxy._types import CallTypes from litellm.proxy.openai_files_endpoints.common_utils import ( BATCH_CREATE_HIDDEN_PARAM, - _is_base64_encoded_unified_file_id, + is_base64_encoded_unified_file_id, encode_file_id_with_model, ) @@ -365,7 +365,7 @@ async def test_async_post_call_success_hook_for_unified_finetuning_job(): ) assert isinstance(response, LiteLLMFineTuningJob) - assert _is_base64_encoded_unified_file_id(response.id) + assert is_base64_encoded_unified_file_id(response.id) @pytest.mark.asyncio @@ -603,7 +603,7 @@ async def test_output_file_id_preserves_target_model_names_when_model_name_missi response=batch, ) - decoded_output_file_id = _is_base64_encoded_unified_file_id( + decoded_output_file_id = is_base64_encoded_unified_file_id( cast(LiteLLMBatch, response).output_file_id ) assert decoded_output_file_id @@ -691,7 +691,7 @@ async def test_error_file_id_for_failed_batch(): assert cast(LiteLLMBatch, response).error_file_id is not None assert not cast(LiteLLMBatch, response).error_file_id.startswith("error-") # Verify it's a base64 encoded managed file ID - assert _is_base64_encoded_unified_file_id( + assert is_base64_encoded_unified_file_id( cast(LiteLLMBatch, response).error_file_id ) @@ -754,7 +754,7 @@ async def test_async_post_call_success_hook_twice_assert_no_unique_violation(): assert task.exception() is None, f"Error: {task.exception()}" assert isinstance(response, LiteLLMBatch) - assert _is_base64_encoded_unified_file_id(response.id) + assert is_base64_encoded_unified_file_id(response.id) # second retrieve batch tasks = [] @@ -2762,7 +2762,7 @@ async def test_return_unified_file_id_includes_expires_at(): assert result.filename == "test.jsonl" assert result.bytes == 1234 assert result.created_at == 1234567890 - assert _is_base64_encoded_unified_file_id(result.id) + assert is_base64_encoded_unified_file_id(result.id) # ============================================================================ diff --git a/tests/unit/enterprise/proxy/test_batch_retrieve_input_file_id.py b/tests/unit/enterprise/proxy/test_batch_retrieve_input_file_id.py index 6e9c3c0354b..93ed7420b4c 100644 --- a/tests/unit/enterprise/proxy/test_batch_retrieve_input_file_id.py +++ b/tests/unit/enterprise/proxy/test_batch_retrieve_input_file_id.py @@ -13,7 +13,7 @@ import pytest from unittest.mock import AsyncMock, MagicMock, patch from litellm.proxy.openai_files_endpoints.common_utils import ( - _is_base64_encoded_unified_file_id, + is_base64_encoded_unified_file_id, ) @@ -32,7 +32,7 @@ async def test_should_resolve_raw_input_file_id_to_unified(): contains a record for that raw ID, the retrieve endpoint should resolve it to the unified file ID. """ - unified_batch_id = _is_base64_encoded_unified_file_id(B64_UNIFIED_BATCH_ID) + unified_batch_id = is_base64_encoded_unified_file_id(B64_UNIFIED_BATCH_ID) assert unified_batch_id, "Test setup: batch_id should decode as unified" from litellm.types.utils import LiteLLMBatch diff --git a/tests/unit/integrations/open_telemetry/test_otel_admin_endpoints.py b/tests/unit/integrations/open_telemetry/test_otel_admin_endpoints.py index 34103449dad..ebeb514ffff 100644 --- a/tests/unit/integrations/open_telemetry/test_otel_admin_endpoints.py +++ b/tests/unit/integrations/open_telemetry/test_otel_admin_endpoints.py @@ -105,9 +105,7 @@ def test_key_generate_failure_stamps_server_span( ) -def test_key_generate_success_stamps_server_span( - server_span_factory, otel_with_exporter -): +def test_key_generate_success_stamps_server_span(server_span_factory, otel_with_exporter): otel, exporter = otel_with_exporter server_span = server_span_factory(KEY_GENERATE_PATH) @@ -230,9 +228,7 @@ def test_management_wrapper_success_ends_server_span_without_http_request( ) -def test_management_wrapper_failure_ends_server_span( - server_span_factory, otel_with_exporter, monkeypatch -): +def test_management_wrapper_failure_ends_server_span(server_span_factory, otel_with_exporter, monkeypatch): """When the handler raises, the wrapper must route through the failure hook and stamp + end the parent SERVER span with the error status — even for an ``http_request``-less handler (route falls back to ``func.__name__``).""" @@ -249,9 +245,7 @@ def test_management_wrapper_failure_ends_server_span( raise HttpStatusException(500, "boom") with pytest.raises(HttpStatusException): - asyncio.run( - failing_fn(data={}, user_api_key_dict=_real_user_api_key_dict(server_span)) - ) + asyncio.run(failing_fn(data={}, user_api_key_dict=_real_user_api_key_dict(server_span))) assert_server_span_attrs( exporter, @@ -261,9 +255,7 @@ def test_management_wrapper_failure_ends_server_span( ) -def test_management_wrapper_success_with_http_request( - server_span_factory, otel_with_exporter, monkeypatch -): +def test_management_wrapper_success_with_http_request(server_span_factory, otel_with_exporter, monkeypatch): """Cover the branch where the handler DOES declare ``http_request``: the route comes from ``http_request.url.path`` and the body is read from it.""" import litellm.proxy.proxy_server as proxy_server @@ -276,7 +268,7 @@ def test_management_wrapper_success_with_http_request( async def _fake_body(request=None): return {"team_alias": "t"} - monkeypatch.setattr(mgmt_utils, "_read_request_body", _fake_body) + monkeypatch.setattr(mgmt_utils, "read_request_body", _fake_body) server_span = server_span_factory("/team/new") http_request = MagicMock() @@ -302,9 +294,7 @@ def test_management_wrapper_success_with_http_request( ) -def test_management_wrapper_noop_when_otel_logger_absent( - server_span_factory, otel_with_exporter, monkeypatch -): +def test_management_wrapper_noop_when_otel_logger_absent(server_span_factory, otel_with_exporter, monkeypatch): """When no OTEL logger is registered, the helper early-returns and no SERVER span is exported — and the handler result is still returned unchanged.""" import litellm.proxy.proxy_server as proxy_server @@ -320,17 +310,13 @@ def test_management_wrapper_noop_when_otel_logger_absent( async def fake_fn(data=None, user_api_key_dict=None): return {"ok": True} - result = asyncio.run( - fake_fn(data={}, user_api_key_dict=_real_user_api_key_dict(server_span)) - ) + result = asyncio.run(fake_fn(data={}, user_api_key_dict=_real_user_api_key_dict(server_span))) assert result == {"ok": True} assert get_server_span(exporter) is None -def test_management_wrapper_swallows_post_success_errors( - server_span_factory, otel_with_exporter, monkeypatch -): +def test_management_wrapper_swallows_post_success_errors(server_span_factory, otel_with_exporter, monkeypatch): """A failure in post-success bookkeeping (cache invalidation, alerting) must not propagate — the handler result is returned regardless (non-blocking).""" import litellm.proxy.proxy_server as proxy_server @@ -351,8 +337,6 @@ def test_management_wrapper_swallows_post_success_errors( async def fake_fn(data=None, user_api_key_dict=None): return {"ok": True} - result = asyncio.run( - fake_fn(data={}, user_api_key_dict=_real_user_api_key_dict(server_span)) - ) + result = asyncio.run(fake_fn(data={}, user_api_key_dict=_real_user_api_key_dict(server_span))) assert result == {"ok": True} diff --git a/tests/unit/integrations/test_shadow_eval_logger.py b/tests/unit/integrations/test_shadow_eval_logger.py index e3f059a7941..55796f7bbcf 100644 --- a/tests/unit/integrations/test_shadow_eval_logger.py +++ b/tests/unit/integrations/test_shadow_eval_logger.py @@ -1522,7 +1522,7 @@ class TestShadowPipeline: monkeypatch.setattr( auth_checks, - "_virtual_key_max_budget_check", + "virtual_key_max_budget_check", AsyncMock(side_effect=BudgetExceededError(current_cost=11.0, max_budget=10.0)), ) router = _router() diff --git a/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py b/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py index 7962d5f6d73..56e97173eb3 100644 --- a/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py +++ b/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py @@ -207,6 +207,84 @@ def test_ollama_pt_consecutive_user_messages(): assert result["prompt"] == expected_prompt +def _ollama_tool_turn(*results: object) -> list[dict]: + call: Final = {"id": "call_1", "type": "function", "function": {"name": "get_weather", "arguments": '{"city": "Paris"}'}} + return [ + {"role": "user", "content": "Weather in Paris?"}, + {"role": "assistant", "content": None, "tool_calls": [call]}, + *({"role": "tool", "tool_call_id": "call_1", "content": result} for result in results), + ] + + +@pytest.mark.parametrize( + ("results", "forwarded"), + [ + pytest.param(("Paris: 22 degrees", "Sky: clear"), "Paris: 22 degrees\nSky: clear", id="two-tool-messages"), + pytest.param( + ([{"type": "text", "text": "Paris: 22 degrees"}, {"type": "text", "text": "clear skies"}],), + "Paris: 22 degrees\nclear skies", + id="text-parts-of-one-tool-message", + ), + pytest.param( + ([{"type": "text", "text": "Paris: 22 degrees"}], "Sky: clear"), + "Paris: 22 degrees\nSky: clear", + id="text-part-then-string", + ), + pytest.param( + ([{"type": "text", "text": ""}, {"type": "text", "text": "clear skies"}],), + "clear skies", + id="empty-text-part-adds-no-blank-line", + ), + ], +) +def test_ollama_pt_separates_merged_tool_results_with_a_newline(results: tuple[object, ...], forwarded: str): + result: Final = ollama_pt(model="llama2", messages=_ollama_tool_turn(*results)) + + assert isinstance(result, dict) + assert result["prompt"].endswith(f"### User:\n{forwarded}\n\n"), result["prompt"] + + +@pytest.mark.parametrize("content", [22, 22.5, True, {"temperature": 22}], ids=type) +def test_ollama_pt_rejects_non_text_tool_content_as_a_bad_request(content: object): + with pytest.raises(litellm.BadRequestError) as excinfo: + ollama_pt(model="llama2", messages=_ollama_tool_turn(content)) + + assert excinfo.value.status_code == 400 + assert "content" in excinfo.value.message + assert "tool message at index 2" in excinfo.value.message + assert type(content).__name__ in excinfo.value.message + + +@pytest.mark.parametrize( + ("part", "expected_detail"), + ( + ({"type": "image_url", "image_url": None}, "NoneType image_url"), + ({"type": "text", "text": 22}, "int text part"), + ({"type": "text"}, "text part with no text"), + ({"type": "image_url"}, "image_url part with no image_url"), + ({"type": "image_url", "image_url": {"detail": "high"}}, "image_url object without a url string"), + ("hello", "str content part"), + ), + ids=( + "none-image-url", + "int-text", + "text-without-text", + "image-url-without-image-url", + "image-url-object-without-url", + "str-part", + ), +) +def test_ollama_pt_rejects_a_malformed_content_part_as_a_bad_request(part: object, expected_detail: str): + messages: Final = [{"role": "user", "content": [part]}] + + with pytest.raises(litellm.BadRequestError) as excinfo: + ollama_pt(model="llava", messages=messages) + + assert excinfo.value.status_code == 400 + assert "user message at index 0" in excinfo.value.message + assert expected_detail in excinfo.value.message + + @pytest.mark.asyncio async def test_anthropic_bedrock_thinking_blocks_with_none_content(): """ diff --git a/tests/unit/litellm_core_utils/test_exception_mapping_utils.py b/tests/unit/litellm_core_utils/test_exception_mapping_utils.py index 486070781ac..7e7f1b536f8 100644 --- a/tests/unit/litellm_core_utils/test_exception_mapping_utils.py +++ b/tests/unit/litellm_core_utils/test_exception_mapping_utils.py @@ -1231,6 +1231,34 @@ def test_branchless_provider_transport_error_maps_to_api_connection_error(): ) +def test_openrouter_transport_error_maps_to_api_connection_error(): + from litellm.llms.base_llm.chat.transformation import BaseLLMException + + original_exception = BaseLLMException(status_code=500, message="[Errno 111] Connection refused") + original_exception.status_code_is_synthesized = True + + with pytest.raises(litellm.APIConnectionError): + exception_type( + model="typesafe/jev-1.13", + original_exception=original_exception, + custom_llm_provider="openrouter", + ) + + +def test_openrouter_upstream_500_still_maps_to_api_error(): + from litellm.llms.base_llm.chat.transformation import BaseLLMException + + original_exception = BaseLLMException(status_code=500, message="upstream exploded") + + with pytest.raises(litellm.APIError) as excinfo: + exception_type( + model="typesafe/jev-1.13", + original_exception=original_exception, + custom_llm_provider="openrouter", + ) + assert excinfo.value.status_code == 500 + + def test_branchless_provider_upstream_500_still_maps_to_internal_server_error(): from litellm.llms.base_llm.chat.transformation import BaseLLMException diff --git a/tests/unit/litellm_core_utils/test_health_check_helpers.py b/tests/unit/litellm_core_utils/test_health_check_helpers.py index 9a686ba9ab3..ffe50c137d2 100644 --- a/tests/unit/litellm_core_utils/test_health_check_helpers.py +++ b/tests/unit/litellm_core_utils/test_health_check_helpers.py @@ -840,12 +840,12 @@ async def test_ahealth_check_without_mode_reports_the_real_failure( def test_update_litellm_params_for_health_check(): """ - Test if _update_litellm_params_for_health_check correctly: + Test if update_litellm_params_for_health_check correctly: 1. Updates messages with a random message 2. Updates model name when health_check_model is provided 3. Updates voice when health_check_voice is provided for audio_speech mode """ - from litellm.proxy.health_check import _update_litellm_params_for_health_check + from litellm.proxy.health_check import update_litellm_params_for_health_check model_info = {"health_check_model": "gpt-5-mini"} litellm_params = { @@ -853,7 +853,7 @@ def test_update_litellm_params_for_health_check(): "api_key": "fake_key", } - updated_params = _update_litellm_params_for_health_check(model_info, litellm_params) + updated_params = update_litellm_params_for_health_check(model_info, litellm_params) assert "messages" in updated_params assert isinstance(updated_params["messages"], list) @@ -865,7 +865,7 @@ def test_update_litellm_params_for_health_check(): "api_key": "fake_key", } - updated_params = _update_litellm_params_for_health_check(model_info, litellm_params) + updated_params = update_litellm_params_for_health_check(model_info, litellm_params) assert "messages" in updated_params assert isinstance(updated_params["messages"], list) @@ -876,7 +876,7 @@ def test_update_litellm_params_for_health_check(): "model": "gpt-5.5", "api_key": "fake_key", } - updated_params = _update_litellm_params_for_health_check(model_info, litellm_params) + updated_params = update_litellm_params_for_health_check(model_info, litellm_params) assert "voice" in updated_params assert updated_params["voice"] == "en-US-JennyNeural" @@ -885,7 +885,7 @@ def test_update_litellm_params_for_health_check(): "model": "gpt-5.5", "api_key": "fake_key", } - updated_params = _update_litellm_params_for_health_check(model_info, litellm_params) + updated_params = update_litellm_params_for_health_check(model_info, litellm_params) assert "voice" in updated_params assert updated_params["voice"] == "alloy" @@ -894,7 +894,7 @@ def test_update_litellm_params_for_health_check(): "model": "gpt-5.5", "api_key": "fake_key", } - updated_params = _update_litellm_params_for_health_check(model_info, litellm_params) + updated_params = update_litellm_params_for_health_check(model_info, litellm_params) assert "voice" not in updated_params model_info = {} @@ -902,28 +902,28 @@ def test_update_litellm_params_for_health_check(): "model": "bedrock/us-gov-west-1/anthropic.claude-sonnet-4-5-20250929-v1:0", "api_key": "fake_key", } - updated_params = _update_litellm_params_for_health_check(model_info, litellm_params) + updated_params = update_litellm_params_for_health_check(model_info, litellm_params) assert updated_params["model"] == "anthropic.claude-sonnet-4-5-20250929-v1:0" litellm_params = { "model": "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0", "api_key": "fake_key", } - updated_params = _update_litellm_params_for_health_check(model_info, litellm_params) + updated_params = update_litellm_params_for_health_check(model_info, litellm_params) assert updated_params["model"] == "us.anthropic.claude-haiku-4-5-20251001-v1:0" litellm_params = { "model": "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0", "api_key": "fake_key", } - updated_params = _update_litellm_params_for_health_check(model_info, litellm_params) + updated_params = update_litellm_params_for_health_check(model_info, litellm_params) assert updated_params["model"] == "us.anthropic.claude-haiku-4-5-20251001-v1:0" litellm_params = { "model": "openai/gpt-5.5", "api_key": "fake_key", } - updated_params = _update_litellm_params_for_health_check(model_info, litellm_params) + updated_params = update_litellm_params_for_health_check(model_info, litellm_params) assert updated_params["model"] == "openai/gpt-5.5" cris_prefixes = ["us.", "eu.", "apac.", "jp.", "au.", "us-gov.", "global."] @@ -932,7 +932,7 @@ def test_update_litellm_params_for_health_check(): "model": f"bedrock/{prefix}anthropic.claude-3-haiku-20240307-v1:0", "api_key": "fake_key", } - updated_params = _update_litellm_params_for_health_check( + updated_params = update_litellm_params_for_health_check( model_info, litellm_params ) assert ( @@ -943,21 +943,21 @@ def test_update_litellm_params_for_health_check(): "model": "bedrock/us-east-2/us.anthropic.claude-3-haiku-20240307-v1:0", "api_key": "fake_key", } - updated_params = _update_litellm_params_for_health_check(model_info, litellm_params) + updated_params = update_litellm_params_for_health_check(model_info, litellm_params) assert updated_params["model"] == "us.anthropic.claude-3-haiku-20240307-v1:0" litellm_params = { "model": "bedrock/us-gov-east-1/anthropic.claude-instant-v1", "api_key": "fake_key", } - updated_params = _update_litellm_params_for_health_check(model_info, litellm_params) + updated_params = update_litellm_params_for_health_check(model_info, litellm_params) assert updated_params["model"] == "anthropic.claude-instant-v1" litellm_params = { "model": "bedrock/llama/arn:aws:bedrock:us-east-1:123:imported-model/abc", "api_key": "fake_key", } - updated_params = _update_litellm_params_for_health_check(model_info, litellm_params) + updated_params = update_litellm_params_for_health_check(model_info, litellm_params) assert ( updated_params["model"] == "llama/arn:aws:bedrock:us-east-1:123:imported-model/abc" @@ -967,7 +967,7 @@ def test_update_litellm_params_for_health_check(): "model": "bedrock/deepseek_r1/arn:aws:bedrock:us-west-2:456:imported-model/xyz", "api_key": "fake_key", } - updated_params = _update_litellm_params_for_health_check(model_info, litellm_params) + updated_params = update_litellm_params_for_health_check(model_info, litellm_params) assert ( updated_params["model"] == "deepseek_r1/arn:aws:bedrock:us-west-2:456:imported-model/xyz" @@ -977,7 +977,7 @@ def test_update_litellm_params_for_health_check(): "model": "bedrock/converse/us.anthropic.claude-haiku-4-5-20251001-v1:0", "api_key": "fake_key", } - updated_params = _update_litellm_params_for_health_check(model_info, litellm_params) + updated_params = update_litellm_params_for_health_check(model_info, litellm_params) assert ( updated_params["model"] == "converse/us.anthropic.claude-haiku-4-5-20251001-v1:0" @@ -987,14 +987,14 @@ def test_update_litellm_params_for_health_check(): "model": "bedrock/invoke/us-west-2/anthropic.claude-instant-v1", "api_key": "fake_key", } - updated_params = _update_litellm_params_for_health_check(model_info, litellm_params) + updated_params = update_litellm_params_for_health_check(model_info, litellm_params) assert updated_params["model"] == "invoke/anthropic.claude-instant-v1" litellm_params = { "model": "bedrock/arn:aws:bedrock:eu-central-1:000:application-inference-profile/abc", "api_key": "fake_key", } - updated_params = _update_litellm_params_for_health_check(model_info, litellm_params) + updated_params = update_litellm_params_for_health_check(model_info, litellm_params) assert ( updated_params["model"] == "arn:aws:bedrock:eu-central-1:000:application-inference-profile/abc" @@ -1004,7 +1004,7 @@ def test_update_litellm_params_for_health_check(): "model": "bedrock/us-west-2/llama/arn:aws:bedrock:us-east-1:123:imported-model/abc", "api_key": "fake_key", } - updated_params = _update_litellm_params_for_health_check(model_info, litellm_params) + updated_params = update_litellm_params_for_health_check(model_info, litellm_params) assert ( updated_params["model"] == "llama/arn:aws:bedrock:us-east-1:123:imported-model/abc" @@ -1014,7 +1014,7 @@ def test_update_litellm_params_for_health_check(): "model": "bedrock/converse/us-west-2/eu.anthropic.claude-3-sonnet-20240229-v1:0", "api_key": "fake_key", } - updated_params = _update_litellm_params_for_health_check(model_info, litellm_params) + updated_params = update_litellm_params_for_health_check(model_info, litellm_params) assert ( updated_params["model"] == "converse/eu.anthropic.claude-3-sonnet-20240229-v1:0" ) diff --git a/tests/unit/litellm_core_utils/test_litellm_logging.py b/tests/unit/litellm_core_utils/test_litellm_logging.py index 3fbc425da7f..caa8708ee9a 100644 --- a/tests/unit/litellm_core_utils/test_litellm_logging.py +++ b/tests/unit/litellm_core_utils/test_litellm_logging.py @@ -28,6 +28,7 @@ import litellm from litellm._internal_context import in_post_response_phase from litellm._logging import session_id_var, trace_id_var, verbose_logger from litellm._service_logger import ServiceLogging +from litellm.caching.caching import DualCache from litellm.constants import LOGGING_WORKER_MAX_TIME_PER_COROUTINE, REDACTED_BY_LITELLM, SENTRY_PII_DENYLIST from litellm.cost_calculator import ocr_batch_cost from litellm.integrations.custom_logger import CustomLogger @@ -42,8 +43,10 @@ from litellm.litellm_core_utils.litellm_logging import ( from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER from litellm.llms.base_llm.ocr.transformation import OCRUsageInfo from litellm.proxy._types import UserAPIKeyAuth -from litellm.proxy.hooks.cache_control_check import _PROXY_CacheControlCheck -from litellm.proxy.hooks.max_iterations_limiter import _PROXY_MaxIterationsHandler +from litellm.proxy.hooks.cache_control_check import PROXY_CacheControlCheck +from litellm.proxy.hooks.max_iterations_limiter import PROXY_MaxIterationsHandler +from litellm.proxy.hooks.parallel_request_limiter_v3 import PROXY_MaxParallelRequestsHandler_v3 +from litellm.proxy.utils import InternalUsageCache from litellm.types.llms.openai import ResponseAPIUsage, ResponseCompletedEvent, ResponsesAPIResponse from litellm.types.utils import ( CallTypes, @@ -11074,6 +11077,24 @@ def test_litellm_logging_no_log_param(monkeypatch, disable_no_log_param): else: assert should_run is False + proxy_callback = PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=InternalUsageCache(DualCache())) + should_run_proxy_callback = litellm_logging_obj.should_run_callback( + callback=proxy_callback, + litellm_params={"no-log": True}, + event_hook="success_handler", + ) + assert should_run_proxy_callback is True + + from litellm_enterprise.proxy.hooks.managed_files import PROXY_LiteLLMManagedFiles + + managed_files_callback = PROXY_LiteLLMManagedFiles(DualCache(), prisma_client=MagicMock()) + should_run_managed_files_callback = litellm_logging_obj.should_run_callback( + callback=managed_files_callback, + litellm_params={"no-log": True}, + event_hook="success_handler", + ) + assert should_run_managed_files_callback is True + @pytest.mark.usefixtures("_vcr_outcome_gate", "drain_logging_worker", "isolate_litellm_state", "setup_and_teardown") def test_get_callback_name(): @@ -11108,7 +11129,7 @@ def test_is_internal_litellm_proxy_callback(): """ logging = setup_logging() - assert logging._is_internal_litellm_proxy_callback(_PROXY_MaxIterationsHandler) == True + assert logging._is_internal_litellm_proxy_callback(PROXY_MaxIterationsHandler) == True # Test non-internal callbacks def regular_callback(): @@ -11141,7 +11162,7 @@ def test_should_run_sync_callbacks_for_async_calls(): assert logging._should_run_sync_callbacks_for_async_calls() == True # Test with internal callback only - litellm.success_callback = [_PROXY_MaxIterationsHandler] + litellm.success_callback = [PROXY_MaxIterationsHandler] assert logging._should_run_sync_callbacks_for_async_calls() == False @pytest.mark.usefixtures("_vcr_outcome_gate", "drain_logging_worker", "isolate_litellm_state", "setup_and_teardown") @@ -11153,8 +11174,8 @@ def test_remove_internal_litellm_callbacks(): callbacks = [ regular_callback, - _PROXY_MaxIterationsHandler, - _PROXY_CacheControlCheck, + PROXY_MaxIterationsHandler, + PROXY_CacheControlCheck, "string_callback", ] @@ -11162,8 +11183,8 @@ def test_remove_internal_litellm_callbacks(): assert len(filtered) == 2 # Should only keep regular_callback and string_callback assert regular_callback in filtered assert "string_callback" in filtered - assert _PROXY_MaxIterationsHandler not in filtered - assert _PROXY_CacheControlCheck not in filtered + assert PROXY_MaxIterationsHandler not in filtered + assert PROXY_CacheControlCheck not in filtered @pytest.mark.asyncio async def test_background_interaction_completion_logs_while_in_progress_handler_is_parked(monkeypatch): diff --git a/tests/unit/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py b/tests/unit/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py index 85e4c055f07..9c0fb5c6a98 100644 --- a/tests/unit/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py +++ b/tests/unit/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py @@ -182,7 +182,7 @@ class TestAnthropicMessagesHandlerStreamingRequestData: with ( patch.object(handler, "_check_streaming_has_ended", return_value=True), patch( - "litellm.llms.anthropic.chat.guardrail_translation.handler.AnthropicPassthroughLoggingHandler._build_complete_streaming_response", + "litellm.llms.anthropic.chat.guardrail_translation.handler.AnthropicPassthroughLoggingHandler.build_complete_streaming_response", return_value=mock_response, ), ): @@ -241,7 +241,7 @@ class TestAnthropicMessagesHandlerStreamingOutputProcessing: with ( patch.object(handler, "_check_streaming_has_ended", return_value=True), patch( - "litellm.llms.anthropic.chat.guardrail_translation.handler.AnthropicPassthroughLoggingHandler._build_complete_streaming_response", + "litellm.llms.anthropic.chat.guardrail_translation.handler.AnthropicPassthroughLoggingHandler.build_complete_streaming_response", return_value=None, ), ): @@ -1353,7 +1353,7 @@ class TestAnthropicMessagesHandlerInputProcessing: with ( patch.object(handler, "_check_streaming_has_ended", return_value=True), patch( - "litellm.llms.anthropic.chat.guardrail_translation.handler.AnthropicPassthroughLoggingHandler._build_complete_streaming_response", + "litellm.llms.anthropic.chat.guardrail_translation.handler.AnthropicPassthroughLoggingHandler.build_complete_streaming_response", return_value=mock_response, ), ): @@ -1400,7 +1400,7 @@ class TestAnthropicMessagesHandlerInputProcessing: with ( patch.object(handler, "_check_streaming_has_ended", return_value=True), patch( - "litellm.llms.anthropic.chat.guardrail_translation.handler.AnthropicPassthroughLoggingHandler._build_complete_streaming_response", + "litellm.llms.anthropic.chat.guardrail_translation.handler.AnthropicPassthroughLoggingHandler.build_complete_streaming_response", return_value=mock_response, ), ): diff --git a/tests/unit/llms/anthropic/pass_through/adapters/test_streaming_iterator_mid_stream_error.py b/tests/unit/llms/anthropic/pass_through/adapters/test_streaming_iterator_mid_stream_error.py index 4798d522182..2f1152174c2 100644 --- a/tests/unit/llms/anthropic/pass_through/adapters/test_streaming_iterator_mid_stream_error.py +++ b/tests/unit/llms/anthropic/pass_through/adapters/test_streaming_iterator_mid_stream_error.py @@ -14,9 +14,13 @@ The async SSE wrapper must instead surface the failure as a well-formed Anthropic ``error`` event so the stream stays valid and the client can retry. """ +import asyncio import json import os import sys +import threading +from datetime import datetime +from collections.abc import Callable from typing import List, Optional from unittest.mock import MagicMock @@ -25,6 +29,8 @@ import pytest sys.path.insert(0, os.path.abspath("../../../../..")) from litellm.exceptions import MidStreamFallbackError +from litellm.integrations.custom_logger import CustomLogger +from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.llms.anthropic.pass_through.adapters.streaming_iterator import ( AnthropicStreamWrapper, _mid_stream_error_sse_event, @@ -140,3 +146,145 @@ def test_error_event_preserves_midstream_fallback_error(): assert name == "error" assert payload["error"]["type"] == "api_error" assert "internalServerException" in payload["error"]["message"] + + +class _AsyncFailureRecorder(CustomLogger): + def __init__(self): + super().__init__() + self.exceptions: list[BaseException] = [] + + async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time): + self.exceptions.append(kwargs["exception"]) + + +class _SyncFailureRecorder: + def __init__(self): + self.exceptions: list[BaseException] = [] + self.called = threading.Event() + + def __call__(self, kwargs, completion_response, start_time, end_time): + self.exceptions.append(kwargs["exception"]) + self.called.set() + + +def _make_logging_obj( + test_name: str, + async_recorder: _AsyncFailureRecorder, + sync_recorder: _SyncFailureRecorder, +) -> LiteLLMLoggingObj: + return LiteLLMLoggingObj( + model="bedrock-converse-sonnet-4-6", + messages=[{"role": "user", "content": "hi"}], + stream=True, + call_type="anthropic_messages", + start_time=datetime.now(), + litellm_call_id=test_name, + function_id=test_name, + dynamic_failure_callbacks=[sync_recorder], + dynamic_async_failure_callbacks=[async_recorder], + ) + + +async def _proxy_boundary_hook(exc: Exception) -> None: + return None + + +async def _wait_for_sync_failure(sync_recorder: _SyncFailureRecorder) -> None: + for _ in range(500): + if sync_recorder.called.is_set(): + return + await asyncio.sleep(0.01) + + +def _bedrock_drop() -> BedrockError: + return BedrockError(status_code=500, message="ConverseStream ended without messageStop") + + +def _chat_wrapper_envelope() -> MidStreamFallbackError: + provider_error = _bedrock_drop() + return MidStreamFallbackError( + message=str(provider_error), + model="bedrock-converse-sonnet-4-6", + llm_provider="bedrock", + original_exception=provider_error, + is_pre_first_chunk=False, + ) + + +_RAISED_ERRORS = pytest.mark.parametrize( + "raised", + [_bedrock_drop, _chat_wrapper_envelope], + ids=["provider_error", "chat_wrapper_envelope"], +) + + +def _failing_wrapper( + logging_obj: LiteLLMLoggingObj | None, + raised: Callable[[], Exception] = _bedrock_drop, +) -> AnthropicStreamWrapper: + return AnthropicStreamWrapper( + completion_stream=_AsyncStreamThenRaise([_make_chunk(Delta(content="partial"))], raised()), + model="bedrock-converse-sonnet-4-6", + litellm_logging_obj=logging_obj, + ) + + +@_RAISED_ERRORS +@pytest.mark.asyncio +async def test_mid_stream_error_reraises_the_provider_error_for_proxy_managed_stream(raised): + async_recorder = _AsyncFailureRecorder() + sync_recorder = _SyncFailureRecorder() + logging_obj = _make_logging_obj("proxy-managed", async_recorder, sync_recorder) + logging_obj.on_detached_stream_failure = _proxy_boundary_hook + + with pytest.raises(BedrockError) as raised_info: + await _drain_sse(_failing_wrapper(logging_obj, raised)) + + assert str(raised_info.value) == "ConverseStream ended without messageStop" + await asyncio.sleep(0.1) + assert async_recorder.exceptions == [] + assert sync_recorder.exceptions == [] + + +@_RAISED_ERRORS +@pytest.mark.asyncio +async def test_mid_stream_error_dispatches_the_provider_error_to_failure_handlers_for_standalone_stream(raised): + async_recorder = _AsyncFailureRecorder() + sync_recorder = _SyncFailureRecorder() + wrapper = _failing_wrapper(_make_logging_obj("standalone", async_recorder, sync_recorder), raised) + + events = await _drain_sse(wrapper) + await _wait_for_sync_failure(sync_recorder) + + assert [str(exc) for exc in async_recorder.exceptions] == ["ConverseStream ended without messageStop"] + assert [str(exc) for exc in sync_recorder.exceptions] == ["ConverseStream ended without messageStop"] + assert _parse_sse(events[-1])[0] == "error" + + +class _BrokenFailureDispatchLogging(LiteLLMLoggingObj): + async def dispatch_failure_handlers(self, *args, **kwargs): + raise RuntimeError("failure sink is down") + + +@pytest.mark.asyncio +async def test_mid_stream_error_frame_survives_a_raising_failure_dispatch(): + logging_obj = _BrokenFailureDispatchLogging( + model="bedrock-converse-sonnet-4-6", + messages=[{"role": "user", "content": "hi"}], + stream=True, + call_type="anthropic_messages", + start_time=datetime.now(), + litellm_call_id="broken-dispatch", + function_id="broken-dispatch", + ) + + events = await _drain_sse(_failing_wrapper(logging_obj)) + + assert _parse_sse(events[-1])[0] == "error" + + +@pytest.mark.asyncio +async def test_mid_stream_error_emits_error_event_without_logging_obj(): + events = await _drain_sse(_failing_wrapper(None)) + + assert _parse_sse(events[-1])[0] == "error" diff --git a/tests/unit/llms/anthropic/pass_through/context_management/test_compact.py b/tests/unit/llms/anthropic/pass_through/context_management/test_compact.py index a94c47d73f1..ced53794cc2 100644 --- a/tests/unit/llms/anthropic/pass_through/context_management/test_compact.py +++ b/tests/unit/llms/anthropic/pass_through/context_management/test_compact.py @@ -1651,10 +1651,10 @@ async def test_summary_model_denied_when_user_over_model_budget(): import inspect from litellm.proxy.hooks.model_max_budget_limiter import ( - _PROXY_VirtualKeyModelMaxBudgetLimiter, + PROXY_VirtualKeyModelMaxBudgetLimiter, ) - real_params = inspect.signature(_PROXY_VirtualKeyModelMaxBudgetLimiter.is_user_within_model_budget).parameters + real_params = inspect.signature(PROXY_VirtualKeyModelMaxBudgetLimiter.is_user_within_model_budget).parameters for kwarg in ("user_id", "user_model_max_budget", "model"): assert kwarg in real_params, f"compact.py passes {kwarg}=, which the limiter no longer accepts" @@ -1767,7 +1767,7 @@ class _FakeRateLimiter: self._raises = raises self.read_only_checked = False - def _create_rate_limit_descriptors(self, **kwargs): + def create_rate_limit_descriptors(self, **kwargs): return [ { "key": "api_key", @@ -1921,12 +1921,12 @@ async def test_summary_model_allowed_while_the_caller_holds_the_keys_only_parall caller's own in-flight slot must not trip a ``max_parallel_requests`` gauge.""" from litellm.caching.caching import DualCache from litellm.proxy._types import UserAPIKeyAuth - from litellm.proxy.hooks.parallel_request_limiter_v3 import _PROXY_MaxParallelRequestsHandler_v3 + from litellm.proxy.hooks.parallel_request_limiter_v3 import PROXY_MaxParallelRequestsHandler_v3 from litellm.proxy.utils import InternalUsageCache, hash_token messages = _simple_messages() mock_call = AsyncMock(return_value=_make_mock_response("

ok")) - limiter = _PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=InternalUsageCache(DualCache())) + limiter = PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=InternalUsageCache(DualCache())) auth = UserAPIKeyAuth( api_key=hash_token("sk-compact-parallel-slot"), max_parallel_requests=1, models=["all-proxy-models"] ) @@ -2061,11 +2061,11 @@ async def test_summary_model_denied_when_team_over_model_budget(): import inspect from litellm.proxy.hooks.model_max_budget_limiter import ( - _PROXY_VirtualKeyModelMaxBudgetLimiter, + PROXY_VirtualKeyModelMaxBudgetLimiter, ) real_params = inspect.signature( - _PROXY_VirtualKeyModelMaxBudgetLimiter.is_team_within_model_budget + PROXY_VirtualKeyModelMaxBudgetLimiter.is_team_within_model_budget ).parameters for kwarg in ("team_id", "team_model_max_budget", "key_model_max_budget", "model"): assert kwarg in real_params, f"compact.py passes {kwarg}=, which the limiter does not accept" diff --git a/tests/unit/llms/anthropic/pass_through/messages/test_response_cache.py b/tests/unit/llms/anthropic/pass_through/messages/test_response_cache.py index 8690d2ee68f..b88a57f440d 100644 --- a/tests/unit/llms/anthropic/pass_through/messages/test_response_cache.py +++ b/tests/unit/llms/anthropic/pass_through/messages/test_response_cache.py @@ -289,7 +289,7 @@ async def test_cached_stream_replay_logs_once_when_polled_after_exhaustion(): with patch.object( PassThroughStreamingHandler, - "_route_streaming_logging_to_handler", + "route_streaming_logging_to_handler", new=AsyncMock(), ) as mock_route: assert await _collect(iterator) == STREAM_EVENTS diff --git a/tests/unit/llms/anthropic/test_anthropic_common_utils.py b/tests/unit/llms/anthropic/test_anthropic_common_utils.py index 337108c80f8..0b0630b5e56 100644 --- a/tests/unit/llms/anthropic/test_anthropic_common_utils.py +++ b/tests/unit/llms/anthropic/test_anthropic_common_utils.py @@ -14,6 +14,7 @@ import json import os import sys import threading +from pathlib import Path from types import SimpleNamespace from typing import Final from unittest.mock import patch @@ -2535,6 +2536,29 @@ class TestWifZeroBehaviorChange: assert "ANTHROPIC_SERVICE_ACCOUNT_ID" in exc_info.value.message assert "ANTHROPIC_IDENTITY_TOKEN_FILE" in exc_info.value.message + def test_a_token_file_credential_missing_an_id_names_the_id_not_the_key(self, clean_anthropic_env: None, tmp_path: Path): + import litellm + from litellm.llms.anthropic.common_utils import AnthropicModelInfo + + with pytest.raises(litellm.AuthenticationError) as exc_info: + AnthropicModelInfo().validate_environment( + headers={}, + model="claude-haiku-5-5", + messages=[{"role": "user", "content": "Hello"}], + optional_params={}, + litellm_params={ + "anthropic_federation_rule_id": "fdrl_1", + "anthropic_identity_token_file": str(tmp_path / "token"), + }, + api_key=None, + api_base=None, + ) + + assert ( + "anthropic_identity_token_file is set, but anthropic_organization_id is not set" in exc_info.value.message + ) + assert "Missing Anthropic API Key" not in exc_info.value.message + class TestWifHeaderContract: def test_minted_token_headers(self, monkeypatch, wif_engine): diff --git a/tests/unit/llms/anthropic/test_anthropic_wif.py b/tests/unit/llms/anthropic/test_anthropic_wif.py index a054b4d130c..733e1472a8a 100644 --- a/tests/unit/llms/anthropic/test_anthropic_wif.py +++ b/tests/unit/llms/anthropic/test_anthropic_wif.py @@ -277,12 +277,18 @@ class TestExchangeHostTrust: def test_a_gateway_listed_with_its_port_is_trusted(self, monkeypatch: pytest.MonkeyPatch): monkeypatch.setenv("LITELLM_ANTHROPIC_WIF_ALLOWED_HOSTS", "gateway.internal:8443") - assert self._mint("https://gateway.internal:8443", monkeypatch) == "https://gateway.internal:8443/v1/oauth/token" + assert ( + self._mint("https://gateway.internal:8443", monkeypatch) == "https://gateway.internal:8443/v1/oauth/token" + ) def test_allowlist_matching_ignores_hostname_case(self, monkeypatch: pytest.MonkeyPatch): monkeypatch.setenv("LITELLM_ANTHROPIC_WIF_ALLOWED_HOSTS", "Gateway.Internal:8443") - assert self._mint("https://gateway.internal:8443", monkeypatch) == "https://gateway.internal:8443/v1/oauth/token" - assert self._mint("https://GATEWAY.internal:8443", monkeypatch) == "https://GATEWAY.internal:8443/v1/oauth/token" + assert ( + self._mint("https://gateway.internal:8443", monkeypatch) == "https://gateway.internal:8443/v1/oauth/token" + ) + assert ( + self._mint("https://GATEWAY.internal:8443", monkeypatch) == "https://GATEWAY.internal:8443/v1/oauth/token" + ) def test_a_gateway_listed_with_a_port_is_not_trusted_on_another_port(self, monkeypatch: pytest.MonkeyPatch): monkeypatch.setenv("LITELLM_ANTHROPIC_WIF_ALLOWED_HOSTS", "gateway.internal:8443") @@ -302,7 +308,9 @@ class TestExchangeHostTrust: def test_a_gateway_listed_without_a_port_is_trusted_on_every_port(self, monkeypatch: pytest.MonkeyPatch): monkeypatch.setenv("LITELLM_ANTHROPIC_WIF_ALLOWED_HOSTS", "gateway.internal") - assert self._mint("https://gateway.internal:9443", monkeypatch) == "https://gateway.internal:9443/v1/oauth/token" + assert ( + self._mint("https://gateway.internal:9443", monkeypatch) == "https://gateway.internal:9443/v1/oauth/token" + ) def test_an_entry_spelling_the_scheme_default_port_matches_a_base_that_omits_it( self, monkeypatch: pytest.MonkeyPatch @@ -311,7 +319,6 @@ class TestExchangeHostTrust: assert self._mint("https://gateway.internal", monkeypatch) == "https://gateway.internal/v1/oauth/token" - class TestBaseUrlDerivation: def _mint(self, api_base: str | None, monkeypatch: pytest.MonkeyPatch) -> str: monkeypatch.setenv("ANTHROPIC_IDENTITY_TOKEN", "inline-jwt") @@ -548,8 +555,6 @@ class TestResolutionMatrix: {"anthropic_federation_rule_id": "fdrl_1"}, {"anthropic_organization_id": "org-1"}, {"anthropic_federation_rule_id": "fdrl_1", "anthropic_organization_id": "org-1"}, - {"anthropic_organization_id": "org-1", "anthropic_identity_token": "oidc/env/TOK"}, - {"anthropic_federation_rule_id": "fdrl_1", "anthropic_identity_token": "oidc/env/TOK"}, ], ) def test_gate_unmet_returns_none(self, litellm_params: dict): @@ -834,7 +839,11 @@ class TestDenialHints: def test_only_console_pointer_when_both_set(self, monkeypatch: pytest.MonkeyPatch): message = self._raise( - {**self.BASE_PARAMS, "anthropic_federation_workspace_id": "wrkspc_1", "anthropic_service_account_id": "svac_1"}, + { + **self.BASE_PARAMS, + "anthropic_federation_workspace_id": "wrkspc_1", + "anthropic_service_account_id": "svac_1", + }, 401, monkeypatch, ) @@ -1064,7 +1073,6 @@ class TestKeycloakIdentitySourceDispatch: assert first.assertion_ref != second.assertion_ref - @pytest.mark.parametrize( "sparse_params", [TestInternalIssuerIdentitySourceDispatch.LITELLM_PARAMS, TestKeycloakIdentitySourceDispatch.LITELLM_PARAMS], @@ -1229,9 +1237,82 @@ class TestMissingIdsFailClosedWhenIdentitySourceConfigured: with pytest.raises(litellm.AuthenticationError, match="must be one of internal_issuer, keycloak"): resolve_anthropic_wif_params({"anthropic_identity_source": "bogus"}) - def test_legacy_token_params_without_ids_still_return_none(self, monkeypatch: pytest.MonkeyPatch): + +class TestLegacyRefsFailClosedWithoutIds: + """A token file or inline token on the deployment asks to federate as explicitly as a named + identity source does, so a missing rule or organization id is reported by name instead of + letting the request die later as a missing API key.""" + + def test_token_file_without_organization_id_names_the_file_param(self, tmp_path: Path): + with pytest.raises(litellm.AuthenticationError) as exc_info: + resolve_anthropic_wif_params( + {"anthropic_federation_rule_id": "fdrl_1", "anthropic_identity_token_file": str(tmp_path / "token")} + ) + + message: Final = exc_info.value.message + assert "anthropic_identity_token_file is set, but anthropic_organization_id is not set. Copy" in message + assert "Settings > Workload identity" in message + assert "ANTHROPIC_FEDERATION_RULE_ID" in message + assert not message.endswith(".") + + def test_inline_token_without_rule_id_names_the_token_param(self): + with pytest.raises(litellm.AuthenticationError) as exc_info: + resolve_anthropic_wif_params( + {"anthropic_organization_id": "org-1", "anthropic_identity_token": "oidc/env/TOK"} + ) + + assert ( + "anthropic_identity_token is set, but anthropic_federation_rule_id is not set. Copy" + in exc_info.value.message + ) + + def test_token_file_with_both_ids_missing_names_both(self, tmp_path: Path): + with pytest.raises( + litellm.AuthenticationError, match="anthropic_federation_rule_id and anthropic_organization_id are not set" + ): + resolve_anthropic_wif_params({"anthropic_identity_token_file": str(tmp_path / "token")}) + + def test_fleet_wide_env_source_does_not_relabel_a_legacy_param(self, monkeypatch: pytest.MonkeyPatch): monkeypatch.setenv("ANTHROPIC_IDENTITY_SOURCE", "internal_issuer") - assert resolve_anthropic_wif_params({"anthropic_identity_token": "oidc/env/TOK"}) is None + with pytest.raises(litellm.AuthenticationError) as exc_info: + resolve_anthropic_wif_params({"anthropic_identity_token": "oidc/env/TOK"}) + + assert "anthropic_identity_token is set, but" in exc_info.value.message + assert "internal_issuer" not in exc_info.value.message + + def test_env_token_file_without_ids_still_returns_none(self, monkeypatch: pytest.MonkeyPatch, tmp_path: Path): + monkeypatch.setenv("ANTHROPIC_IDENTITY_TOKEN_FILE", str(tmp_path / "token")) + assert resolve_anthropic_wif_params({"anthropic_federation_rule_id": "fdrl_1"}) is None + + def test_environment_ids_complete_a_token_file_param(self, monkeypatch: pytest.MonkeyPatch, tmp_path: Path): + monkeypatch.setenv("ANTHROPIC_FEDERATION_RULE_ID", "fdrl_env") + monkeypatch.setenv("ANTHROPIC_ORGANIZATION_ID", "org-env") + token_file: Final = tmp_path / "token" + + params: Final = resolve_anthropic_wif_params({"anthropic_identity_token_file": str(token_file)}) + + assert params is not None + assert (params.federation_rule_id, params.organization_id) == ("fdrl_env", "org-env") + assert params.assertion_ref == f"oidc/file/{token_file}" + + def test_disabling_federation_wins_over_the_gate(self, tmp_path: Path): + litellm_params: Final = { + "anthropic_disable_workload_identity_federation": True, + "anthropic_identity_token_file": str(tmp_path / "token"), + } + assert resolve_anthropic_wif_params(litellm_params) is None + + def test_facade_raises_without_an_engine_call(self): + poster: Final = ScriptedPoster([token_response()]) + engine: Final = make_engine(poster) + with pytest.raises(litellm.AuthenticationError, match="anthropic_identity_token is set, but"): + get_anthropic_wif_token( + {"anthropic_organization_id": "org-1", "anthropic_identity_token": "oidc/env/TOK"}, + None, + "claude-haiku-5-5", + engine, + ) + assert poster.requests == [] class TestConfigYamlShapedIdentitySources: diff --git a/tests/unit/llms/base_llm/decisions/__init__.py b/tests/unit/llms/base_llm/decisions/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/base_llm/decisions/test_base_decisions_transformation.py b/tests/unit/llms/base_llm/decisions/test_base_decisions_transformation.py new file mode 100644 index 00000000000..a452a621d2e --- /dev/null +++ b/tests/unit/llms/base_llm/decisions/test_base_decisions_transformation.py @@ -0,0 +1,180 @@ +from collections.abc import Mapping +from typing import Final + +import pytest +from pydantic import TypeAdapter + +from litellm.llms.base_llm.decisions.transformation import ( + ir_to_systemone_request, + ir_to_systemone_response, + parse_systemone_response, + systemone_request_to_ir, +) +from litellm.llms.openai.decisions.transformation import openai_request_to_ir +from litellm.types.decisions import ( + MAX_DECISION_QUESTIONS, + DecisionsIRRefusal, + DecisionsIRRequest, + DecisionsRequestBody, + OpenAIDecisionRequestBody, + UnsupportedDecisionsRequest, +) + +_SYSTEMONE_BODY: Final[TypeAdapter[DecisionsRequestBody]] = TypeAdapter(DecisionsRequestBody) +_OPENAI_BODY: Final[TypeAdapter[OpenAIDecisionRequestBody]] = TypeAdapter(OpenAIDecisionRequestBody) + +_SYSTEMONE_REQUEST: Final[Mapping[str, object]] = { + "state": {"ticket": 1234, "text": "Screen cracked"}, + "questions": { + "damaged": { + "type": "noul", + "instructions": "Is the item damaged?", + "criteria": {"true": "Visible damage", "false": None}, + "provider_field": "dropped", + }, + "rubric_only": {"type": "noul", "criteria": {"true": {"signal": "refund"}}}, + "action": { + "type": "choice", + "instructions": {"policy": "refund-v2"}, + "criteria": {"refund": "Within 30 days", "escalate": None}, + }, + "severity": {"type": "score", "criteria": ["minor", {"label": "major"}]}, + }, +} + + +def _openai_ir(raw: Mapping[str, object]) -> DecisionsIRRequest: + return openai_request_to_ir(_OPENAI_BODY.validate_python(raw)) + + +def test_a_systemone_request_reaches_a_systemone_provider_unchanged() -> None: + ir: Final = systemone_request_to_ir(_SYSTEMONE_BODY.validate_python(_SYSTEMONE_REQUEST)) + + assert ir_to_systemone_request("jev-latest", ir) == {"model": "jev-latest", **_SYSTEMONE_REQUEST} + + +@pytest.mark.parametrize( + ("names", "keys"), + ( + (("damaged", "action"), ("damaged", "action")), + (("damaged", None), ("0", "1")), + (("damaged", "damaged"), ("0", "1")), + ), + ids=("all_named", "one_unnamed", "duplicate_names"), +) +def test_openai_questions_are_keyed_by_name_only_when_every_name_is_unique( + names: tuple[str | None, str | None], keys: tuple[str, str] +) -> None: + questions: Final = [ + {"type": "predicate", "instructions": f"Question {index}?", **({} if name is None else {"name": name})} + for index, name in enumerate(names) + ] + + body: Final = ir_to_systemone_request("jev-latest", _openai_ir({"input": "review", "questions": questions})) + + assert tuple(_SYSTEMONE_BODY.validate_python(body).questions) == keys + + +def test_openai_messages_become_systemone_state_text() -> None: + ir: Final = _openai_ir( + { + "input": [ + { + "role": "user", + "content": [ + {"type": "input_text", "text": "The package arrived broken."}, + {"type": "input_text", "text": "I want a refund."}, + ], + }, + {"role": "user", "content": "Order 1234."}, + ], + "questions": [{"type": "predicate", "instructions": "Is this a defect?"}], + } + ) + + body: Final = ir_to_systemone_request("jev-latest", ir) + + assert isinstance(body, Mapping) + assert body["state"] == "The package arrived broken.\n\nI want a refund.\n\nOrder 1234." + + +def test_image_input_is_unsupported_by_systemone_providers() -> None: + ir: Final = _openai_ir( + { + "input": [ + { + "role": "user", + "content": [ + {"type": "input_text", "text": "Is the screen cracked?"}, + {"type": "input_image", "image_url": "data:image/png;base64,AA=="}, + ], + } + ], + "questions": [{"type": "predicate", "instructions": "Is this a defect?"}], + } + ) + + assert isinstance(ir_to_systemone_request("jev-latest", ir), UnsupportedDecisionsRequest) + + +def test_the_largest_openai_request_accepted_translates_to_a_valid_systemone_request() -> None: + ir: Final = _openai_ir( + { + "input": "review", + "questions": [ + {"type": "predicate", "instructions": f"Question {index}?"} for index in range(MAX_DECISION_QUESTIONS) + ], + } + ) + + translated: Final = _SYSTEMONE_BODY.validate_python(ir_to_systemone_request("jev-latest", ir)) + + assert len(translated.questions) == MAX_DECISION_QUESTIONS + + +def test_a_systemone_response_reaches_the_caller_unchanged_with_provider_extras() -> None: + payload: Final = { + "model": "jev-latest", + "answers": { + "damaged": {"type": "noul", "noul": 0.95, "rationale": "crack visible"}, + "action": { + "type": "choice", + "choice": "refund", + "confidence": 0.8, + "probabilities": {"refund": 0.9, "escalate": 0.1}, + "calibrated": True, + }, + "severity": { + "type": "score", + "score": 0.4, + "confidence": 0.6, + "legend": {"0": "minor", "1": {"label": "major"}}, + "probabilities": {"0": 0.6, "1": 0.4}, + "raw_logits": [0.1, 0.2], + }, + }, + "usage": {"input_tokens": 383, "output_tokens": 2, "cost": 0.25}, + "latency_ms": 3722.17, + } + ir: Final = systemone_request_to_ir(_SYSTEMONE_BODY.validate_python(_SYSTEMONE_REQUEST)) + + response: Final = ir_to_systemone_response(parse_systemone_response(payload, ir), ir) + + assert response.model_dump(mode="json") == payload + + +def test_missing_or_mismatched_systemone_answers_are_refusals_left_out_of_the_response() -> None: + ir: Final = systemone_request_to_ir(_SYSTEMONE_BODY.validate_python(_SYSTEMONE_REQUEST)) + payload: Final = { + "model": "jev-latest", + "answers": { + "damaged": {"type": "noul", "noul": 0.95}, + "action": {"type": "noul", "noul": 0.5}, + }, + "usage": {"input_tokens": 10, "output_tokens": 1}, + } + + parsed: Final = parse_systemone_response(payload, ir) + + assert parsed.answers[1:] == (DecisionsIRRefusal(), DecisionsIRRefusal(), DecisionsIRRefusal()) + assert tuple(ir_to_systemone_response(parsed, ir).answers) == ("damaged",) diff --git a/tests/unit/llms/base_llm/decisions/test_systemone.py b/tests/unit/llms/base_llm/decisions/test_systemone.py new file mode 100644 index 00000000000..42a1c41b3d3 --- /dev/null +++ b/tests/unit/llms/base_llm/decisions/test_systemone.py @@ -0,0 +1,262 @@ +from __future__ import annotations + +from collections.abc import Mapping, Sequence +from typing import Final + +import pytest +from pydantic import TypeAdapter + +from litellm.llms.base_llm.chat.transformation import BaseLLMException +from litellm.llms.base_llm.decisions.systemone import ( + SYSTEM_ONE_RESPONSE_ADAPTER, + question_keys, + to_decisions_response, + to_system_one_request, +) +from litellm.types.openai_decisions import ( + ChoiceAnswer, + DecisionsRequest, + DecisionsRequestBody, + PredicateAnswer, + ScoreAnswer, +) + +_INPUT: Final = "The export job hangs at 99% and never finishes" +_QUESTIONS: Final[Sequence[Mapping[str, object]]] = ( + {"type": "predicate", "name": "is_defect", "instructions": "Is this a defect?"}, + { + "type": "choice", + "name": "sentiment", + "instructions": "How does the customer feel?", + "choices": [{"value": "positive"}, {"value": "negative", "description": "unhappy"}], + }, + { + "type": "score", + "name": "severity", + "instructions": "How severe is it?", + "levels": [{"label": "none"}, {"label": "low"}, {"label": "high", "description": "blocks users"}], + }, +) +_SYSTEM_ONE_QUESTIONS: Final[Mapping[str, object]] = { + "is_defect": {"type": "noul", "instructions": "Is this a defect?"}, + "sentiment": { + "type": "choice", + "instructions": "How does the customer feel?", + "criteria": {"positive": None, "negative": "unhappy"}, + }, + "severity": {"type": "score", "instructions": "How severe is it?", "criteria": ["none", "low", "blocks users"]}, +} +_SYSTEM_ONE_RESPONSE: Final[Mapping[str, object]] = { + "model": "jev-1.13", + "answers": { + "is_defect": {"type": "noul", "noul": 0.9}, + "sentiment": { + "type": "choice", + "choice": "positive", + "confidence": 0.8, + "probabilities": {"positive": 0.8, "negative": 0.2}, + }, + "severity": { + "type": "score", + "score": 1, + "confidence": 0.7, + "legend": {"0": "none", "1": "low", "2": "high"}, + "probabilities": {"0": 0.1, "1": 0.8, "2": 0.1}, + }, + }, + "usage": {"input_tokens": 367, "output_tokens": 3}, +} +_EXPECTED_ANSWERS: Final[Sequence[Mapping[str, object]]] = ( + {"type": "predicate", "name": "is_defect", "probability": 0.9}, + { + "type": "choice", + "name": "sentiment", + "choice": "positive", + "probabilities": [{"value": "positive", "probability": 0.8}, {"value": "negative", "probability": 0.2}], + "confidence": 0.8, + }, + { + "type": "score", + "name": "severity", + "score": 1.0, + "probabilities": [ + {"value": 0, "label": "none", "probability": 0.1}, + {"value": 1, "label": "low", "probability": 0.8}, + {"value": 2, "label": "high", "probability": 0.1}, + ], + "confidence": 0.7, + }, +) +_EXPECTED_USAGE: Final[Mapping[str, object]] = { + "input_tokens": 367, + "input_tokens_details": {"cached_tokens": 0, "cache_write_tokens": 0}, + "output_tokens": 3, + "output_tokens_details": {"reasoning_tokens": 0}, + "total_tokens": 370, +} +_BODY_ADAPTER: Final[TypeAdapter[DecisionsRequestBody]] = TypeAdapter(DecisionsRequestBody) + + +def _body( + input_value: object = _INPUT, + questions: Sequence[Mapping[str, object]] = _QUESTIONS, +) -> DecisionsRequestBody: + return _BODY_ADAPTER.validate_python({"input": input_value, "questions": questions}) + + +def _request( + input_value: object = _INPUT, + questions: Sequence[Mapping[str, object]] = _QUESTIONS, + model: str = "jev-1.13", +) -> DecisionsRequest: + return DecisionsRequest(model=model, body=_body(input_value, questions)) + + +def _predicate(name: str | None = "is_defect") -> tuple[Mapping[str, object]]: + return ({"type": "predicate", "name": name, "instructions": "Is this a defect?"},) + + +def test_openai_request_becomes_the_system_one_body() -> None: + assert to_system_one_request("jev-1.13", _body(), "typesafe") == { + "model": "jev-1.13", + "state": _INPUT, + "questions": _SYSTEM_ONE_QUESTIONS, + } + + +def test_system_one_answers_become_openai_answers_in_question_order() -> None: + response: Final = to_decisions_response( + SYSTEM_ONE_RESPONSE_ADAPTER.validate_python(_SYSTEM_ONE_RESPONSE), _request(), "typesafe" + ) + + assert response.model_dump(mode="json") == { + "model": "jev-1.13", + "answers": list(_EXPECTED_ANSWERS), + "usage": _EXPECTED_USAGE, + } + assert isinstance(response.answers[0], PredicateAnswer) + assert isinstance(response.answers[1], ChoiceAnswer) + assert isinstance(response.answers[2], ScoreAnswer) + + +def test_user_messages_are_joined_into_one_system_one_state() -> None: + messages: Final = [ + {"role": "user", "content": "first"}, + { + "role": "user", + "content": [{"type": "input_text", "text": "second"}, {"type": "input_text", "text": "third"}], + }, + ] + + body: Final = to_system_one_request("jev-1.13", _body(messages, _predicate()), "typesafe") + + assert body["state"] == "first\nsecond\nthird" + + +@pytest.mark.parametrize( + ("label", "input_value", "questions"), + ( + ( + "input_image", + [{"role": "user", "content": [{"type": "input_image", "image_url": "data:image/png;base64,AA=="}]}], + _predicate(), + ), + ("boolean choice", _INPUT, [{"type": "choice", "instructions": "Refund?", "choices": [{"value": True}]}]), + ("unique name", _INPUT, [*_predicate(), *_predicate()]), + ( + "repeated choice", + _INPUT, + [{"type": "choice", "instructions": "Refund?", "choices": [{"value": "yes"}, {"value": "yes"}]}], + ), + ), +) +def test_what_system_one_cannot_express_is_a_400( + label: str, + input_value: object, + questions: Sequence[Mapping[str, object]], +) -> None: + with pytest.raises(BaseLLMException, match=label) as error: + to_system_one_request("jev-1.13", _body(input_value, questions), "perplexity") + + assert error.value.status_code == 400 + assert "perplexity" in error.value.message + + +def test_unnamed_questions_get_positional_keys_that_never_shadow_a_supplied_name() -> None: + body: Final = _body(questions=[*_predicate(None), *_predicate("q0"), *_predicate(None)]) + + assert question_keys(body.questions, "typesafe") == ("_q0", "q0", "q2") + noul: Final = {"type": "noul", "instructions": "Is this a defect?"} + assert to_system_one_request("jev-1.13", body, "typesafe")["questions"] == {"_q0": noul, "q0": noul, "q2": noul} + + +def test_positional_answers_come_back_in_question_order_without_a_name() -> None: + request: Final = _request(questions=[*_predicate(None), *_predicate("q0")]) + system_one: Final = SYSTEM_ONE_RESPONSE_ADAPTER.validate_python( + {"answers": {"_q0": {"type": "noul", "noul": 0.25}, "q0": {"type": "noul", "noul": 0.75}}} + ) + + response: Final = to_decisions_response(system_one, request, "typesafe") + + assert [answer.model_dump(mode="json") for answer in response.answers] == [ + {"type": "predicate", "name": None, "probability": 0.25}, + {"type": "predicate", "name": "q0", "probability": 0.75}, + ] + assert response.usage.model_dump(mode="json") == { + **_EXPECTED_USAGE, + "input_tokens": 0, + "output_tokens": 0, + "total_tokens": 0, + } + + +def test_a_reply_without_a_model_reports_the_requested_model() -> None: + system_one: Final = SYSTEM_ONE_RESPONSE_ADAPTER.validate_python( + {k: v for k, v in _SYSTEM_ONE_RESPONSE.items() if k != "model"} + ) + + response: Final = to_decisions_response(system_one, _request(model="typesafe/jev-1.13.0"), "typesafe") + + assert response.model == "typesafe/jev-1.13.0" + + +def test_a_choice_the_provider_left_out_of_probabilities_is_reported_at_zero() -> None: + system_one: Final = SYSTEM_ONE_RESPONSE_ADAPTER.validate_python( + { + "answers": { + "sentiment": { + "type": "choice", + "choice": "positive", + "confidence": 1.0, + "probabilities": {"positive": 1.0}, + } + } + } + ) + request: Final = _request(questions=_QUESTIONS[1:2]) + + response: Final = to_decisions_response(system_one, request, "typesafe") + + assert response.answers[0].model_dump(mode="json") == { + "type": "choice", + "name": "sentiment", + "choice": "positive", + "probabilities": [{"value": "positive", "probability": 1.0}, {"value": "negative", "probability": 0.0}], + "confidence": 1.0, + } + + +@pytest.mark.parametrize( + "answers", + ( + {}, + {"is_defect": {"type": "choice", "choice": "yes", "confidence": 1.0, "probabilities": {"yes": 1.0}}}, + ), +) +def test_a_reply_without_a_matching_answer_is_a_server_error(answers: Mapping[str, object]) -> None: + system_one: Final = SYSTEM_ONE_RESPONSE_ADAPTER.validate_python({"answers": answers}) + + with pytest.raises(BaseLLMException, match="no predicate answer for question 'is_defect'") as error: + to_decisions_response(system_one, _request(questions=_predicate()), "typesafe") + + assert error.value.status_code == 500 diff --git a/tests/unit/llms/bedrock/batches/test_handler.py b/tests/unit/llms/bedrock/batches/test_handler.py index 4d7f77a6590..d920eabcdf6 100644 --- a/tests/unit/llms/bedrock/batches/test_handler.py +++ b/tests/unit/llms/bedrock/batches/test_handler.py @@ -222,7 +222,7 @@ def test_missing_error_count_maps_to_zero_failed(patched_boto3): [ ("Submitted", "validating"), ("Validating", "validating"), - ("Scheduled", "validating"), + ("Scheduled", "in_progress"), ("InProgress", "in_progress"), ("Stopping", "cancelling"), ("Stopped", "cancelled"), diff --git a/tests/unit/llms/custom_httpx/test_http_handler.py b/tests/unit/llms/custom_httpx/test_http_handler.py index f5ae65e9898..b6b0db69162 100644 --- a/tests/unit/llms/custom_httpx/test_http_handler.py +++ b/tests/unit/llms/custom_httpx/test_http_handler.py @@ -24,7 +24,9 @@ from litellm.llms.custom_httpx.http_handler import ( HTTPHandler, MaskedHTTPStatusError, get_httpx_client, + get_shared_realtime_ssl_context, get_ssl_configuration, + realtime_ssl_for_url, ) from litellm.types.llms.custom_http import VerifyTypes from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer @@ -1493,7 +1495,6 @@ async def test_connection_error_retry_forwards_content(method: str): await handler.close() - @pytest.fixture def forward_proxy_server(): """Plain HTTP forward proxy that records the absolute URIs it is asked to fetch.""" @@ -1624,9 +1625,7 @@ def private_ca_tls_upstream(tmp_path: pathlib.Path): ca_pem.write_bytes(cert.public_bytes(serialization.Encoding.PEM)) key_pem = tmp_path / "key.pem" key_pem.write_bytes( - key.private_bytes( - serialization.Encoding.PEM, serialization.PrivateFormat.PKCS8, serialization.NoEncryption() - ) + key.private_bytes(serialization.Encoding.PEM, serialization.PrivateFormat.PKCS8, serialization.NoEncryption()) ) class OkTlsHandler(BaseHTTPRequestHandler): @@ -1801,8 +1800,11 @@ async def test_bounded_get_preserves_sdk_redirect_auth_and_query_handling(respx_ handler = AsyncHTTPHandler() try: response = await handler.get( - "https://example.com/spec.json?original=1", max_response_bytes=100, follow_redirects=True, - headers={"Authorization": "Bearer sentinel", "Accept-Encoding": "gzip"}, timeout=2.0, + "https://example.com/spec.json?original=1", + max_response_bytes=100, + follow_redirects=True, + headers={"Authorization": "Bearer sentinel", "Accept-Encoding": "gzip"}, + timeout=2.0, ) finally: await handler.close() @@ -1888,6 +1890,7 @@ def _vcr_outcome_gate(request, vcr): yield record_vcr_outcome(request, vcr) + @pytest.fixture(scope="function") def isolate_litellm_state(): """ @@ -1940,6 +1943,7 @@ def isolate_litellm_state(): setattr(litellm, attr, original_value) _invalidate_model_cost_lowercase_map() + _SCALAR_DEFAULTS = { "num_retries": getattr(litellm, "num_retries", None), "num_retries_per_request": getattr(litellm, "num_retries_per_request", None), @@ -1960,6 +1964,7 @@ _SCALAR_DEFAULTS = { "api_key": getattr(litellm, "api_key", None), } + @pytest.fixture(scope="module") def setup_and_teardown(): """ @@ -1982,12 +1987,14 @@ def setup_and_teardown(): litellm.in_memory_llm_clients_cache.flush_cache() yield + _SERVER_DELAY_S = 5 _PER_REQUEST_TIMEOUT_S = 1.0 _CLIENT_DEFAULT_TIMEOUT_S = 60.0 + class _SlowHandler(BaseHTTPRequestHandler): def do_POST(self): time.sleep(_SERVER_DELAY_S) @@ -2001,6 +2008,7 @@ class _SlowHandler(BaseHTTPRequestHandler): def log_message(self, *args): pass + @pytest.mark.usefixtures("_vcr_outcome_gate", "isolate_litellm_state", "setup_and_teardown") def test_post_delay_exceeds_per_request_timeout_raises(): server = ThreadingHTTPServer(("127.0.0.1", 0), _SlowHandler) @@ -2022,3 +2030,24 @@ def test_post_delay_exceeds_per_request_timeout_raises(): handler.close() server.shutdown() server.server_close() + + +def test_realtime_ssl_for_url_sends_no_tls_argument_for_a_plain_ws_endpoint() -> None: + assert ( + realtime_ssl_for_url("ws://127.0.0.1:8080/ws/google.cloud.aiplatform.v1.LlmBidiService/BidiGenerateContent") + is None + ) + + +def test_realtime_ssl_for_url_keeps_the_shared_context_for_wss_endpoints() -> None: + shared: Final = get_shared_realtime_ssl_context() + assert isinstance(shared, ssl.SSLContext) + assert realtime_ssl_for_url("wss://aiplatform.us.rep.googleapis.com/ws") is shared + + +def test_realtime_ssl_for_url_turns_ssl_verify_false_into_an_unverified_tls_context(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr("litellm.llms.custom_httpx.http_handler._shared_realtime_ssl_context", False) + selected: Final = realtime_ssl_for_url("wss://aiplatform.us.rep.googleapis.com/ws") + assert isinstance(selected, ssl.SSLContext) + assert selected.verify_mode == ssl.CERT_NONE + assert selected.check_hostname is False diff --git a/tests/unit/llms/ollama/test_ollama_completion_transformation.py b/tests/unit/llms/ollama/test_ollama_completion_transformation.py index 8558bf50bb9..2979cc7d572 100644 --- a/tests/unit/llms/ollama/test_ollama_completion_transformation.py +++ b/tests/unit/llms/ollama/test_ollama_completion_transformation.py @@ -596,6 +596,28 @@ class TestOllamaConfig: class TestOllamaTextCompletionResponseIterator: + def test_every_chunk_of_one_stream_carries_the_same_response_id(self): + iterator: Final = OllamaTextCompletionResponseIterator( + streaming_response=iter([]), sync_stream=True, json_mode=False + ) + ollama_chunks: Final = ( + {"model": "qwen3:0.6b", "created_at": "2026-10-07T00:00:00Z", "response": "", "done": False}, + {"model": "qwen3:0.6b", "created_at": "2026-10-07T00:00:00Z", "response": "", "thinking": "Hm", "done": False}, + {"model": "qwen3:0.6b", "created_at": "2026-10-07T00:00:00Z", "response": "Hel", "done": False}, + {"model": "qwen3:0.6b", "created_at": "2026-10-07T00:00:00Z", "response": "lo", "done": False}, + ) + + results: Final = tuple(iterator.chunk_parser(chunk) for chunk in ollama_chunks) + + ids: Final = {result.id for result in results if isinstance(result, ModelResponseStream)} + assert len(results) == len(ollama_chunks) and len(ids) == 1, ids + assert next(iter(ids)).startswith("chatcmpl-") + other: Final = OllamaTextCompletionResponseIterator( + streaming_response=iter([]), sync_stream=True, json_mode=False + ) + other_result: Final = other.chunk_parser(ollama_chunks[2]) + assert isinstance(other_result, ModelResponseStream) and other_result.id not in ids + def test_chunk_parser_with_thinking_field(self): """Test that chunks with 'thinking' field and empty 'response' are handled correctly.""" iterator = OllamaTextCompletionResponseIterator( diff --git a/tests/unit/llms/openai/decisions/__init__.py b/tests/unit/llms/openai/decisions/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/openai/decisions/test_openai_decisions_transformation.py b/tests/unit/llms/openai/decisions/test_openai_decisions_transformation.py new file mode 100644 index 00000000000..e73c5b92ab6 --- /dev/null +++ b/tests/unit/llms/openai/decisions/test_openai_decisions_transformation.py @@ -0,0 +1,256 @@ +from collections.abc import Mapping +from typing import Final + +import pytest +from pydantic import TypeAdapter + +from litellm.llms.base_llm.decisions.transformation import ir_to_systemone_response, systemone_request_to_ir +from litellm.llms.openai.decisions.transformation import ( + OpenAIDecisionsConfig, + ir_to_openai_request, + ir_to_openai_response, + openai_request_to_ir, +) +from litellm.types.decisions import ( + DecisionsIRRequest, + DecisionsRequestBody, + OpenAIDecisionRequestBody, + UnsupportedDecisionsRequest, +) + +_SYSTEMONE_BODY: Final[TypeAdapter[DecisionsRequestBody]] = TypeAdapter(DecisionsRequestBody) +_OPENAI_BODY: Final[TypeAdapter[OpenAIDecisionRequestBody]] = TypeAdapter(OpenAIDecisionRequestBody) + +_OPENAI_REQUEST: Final[Mapping[str, object]] = { + "input": [ + { + "type": "message", + "role": "user", + "content": [ + {"type": "input_text", "text": "The screen is cracked."}, + {"type": "input_image", "image_url": "data:image/png;base64,AA==", "detail": "high"}, + ], + }, + {"type": "message", "role": "user", "content": "Order 1234."}, + ], + "questions": [ + {"type": "predicate", "name": "damaged", "instructions": "Is the item damaged?"}, + { + "type": "choice", + "instructions": "Should we refund?", + "choices": [{"value": True, "description": "Refund now"}, {"value": "escalate"}], + }, + { + "type": "score", + "name": "severity", + "instructions": "How severe is it?", + "levels": [{"label": "minor"}, {"label": "major", "description": "Product unusable"}], + }, + {"type": "predicate", "name": "fraud", "instructions": "Is this fraud?"}, + ], + "safety_identifier": "end-user-1", +} +_PREDICATE_ANSWER: Final[Mapping[str, object]] = {"type": "predicate", "name": "damaged", "probability": 0.95} +_CHOICE_ANSWER: Final[Mapping[str, object]] = { + "type": "choice", + "name": None, + "choice": True, + "probabilities": [{"value": True, "probability": 0.9}, {"value": "escalate", "probability": 0.1}], + "confidence": 0.8, +} +_REFUSAL_ANSWER: Final[Mapping[str, object]] = {"type": "refusal", "name": "fraud"} +_USAGE: Final[Mapping[str, object]] = { + "input_tokens": 383, + "input_tokens_details": {"cached_tokens": 256, "cache_write_tokens": 64}, + "output_tokens": 2, + "output_tokens_details": {"reasoning_tokens": 1}, + "total_tokens": 385, +} +_OPENAI_RESPONSE: Final[Mapping[str, object]] = { + "model": "gpt-6-luna", + "answers": [ + _PREDICATE_ANSWER, + _CHOICE_ANSWER, + { + "type": "score", + "name": "severity", + "score": 0.7, + "probabilities": [ + {"value": 0, "label": "minor", "probability": 0.3}, + {"value": 1, "label": "major", "probability": 0.7}, + ], + "confidence": 0.6, + }, + _REFUSAL_ANSWER, + ], + "usage": _USAGE, +} + + +def _openai_ir(raw: Mapping[str, object]) -> DecisionsIRRequest: + return openai_request_to_ir(_OPENAI_BODY.validate_python(raw)) + + +def test_an_openai_request_reaches_openai_unchanged() -> None: + assert ir_to_openai_request("gpt-6-luna", _openai_ir(_OPENAI_REQUEST)) == { + "model": "gpt-6-luna", + **_OPENAI_REQUEST, + } + + +def test_an_openai_response_reaches_the_caller_unchanged() -> None: + ir: Final = _openai_ir(_OPENAI_REQUEST) + + parsed: Final = OpenAIDecisionsConfig().parse_response(_OPENAI_RESPONSE, ir) + + assert ir_to_openai_response(parsed, ir, "requested").model_dump(mode="json") == _OPENAI_RESPONSE + + +def test_answers_openai_did_not_return_are_refusals() -> None: + ir: Final = _openai_ir(_OPENAI_REQUEST) + payload: Final = {**_OPENAI_RESPONSE, "answers": [_PREDICATE_ANSWER]} + + response: Final = ir_to_openai_response(OpenAIDecisionsConfig().parse_response(payload, ir), ir, "requested") + + assert [answer.type for answer in response.answers] == ["predicate", "refusal", "refusal", "refusal"] + assert [answer.name for answer in response.answers] == ["damaged", None, "severity", "fraud"] + + +def test_a_systemone_request_becomes_an_openai_request_with_questions_named_by_their_keys() -> None: + request: Final = _SYSTEMONE_BODY.validate_python( + { + "state": {"ticket": 1234, "text": "Screen cracked"}, + "questions": { + "damaged": { + "type": "noul", + "instructions": "Is the item damaged?", + "criteria": {"true": "Visible damage", "false": None}, + "provider_field": "dropped", + }, + "rubric_only": {"type": "noul", "criteria": {"true": {"signal": "refund"}}}, + "action": { + "type": "choice", + "instructions": {"policy": "refund-v2"}, + "criteria": {"refund": "Within 30 days", "escalate": None}, + }, + "severity": {"type": "score", "criteria": ["minor", {"label": "major"}]}, + }, + } + ) + + assert ir_to_openai_request("gpt-6-luna", systemone_request_to_ir(request)) == { + "model": "gpt-6-luna", + "input": '{"ticket": 1234, "text": "Screen cracked"}', + "questions": [ + { + "type": "predicate", + "name": "damaged", + "instructions": "Is the item damaged?\n\nAnswer true when: Visible damage", + }, + {"type": "predicate", "name": "rubric_only", "instructions": 'Answer true when: {"signal": "refund"}'}, + { + "type": "choice", + "name": "action", + "instructions": '{"policy": "refund-v2"}', + "choices": [{"value": "refund", "description": "Within 30 days"}, {"value": "escalate"}], + }, + { + "type": "score", + "name": "severity", + "instructions": "Which level best fits the input?", + "levels": [{"label": "minor"}, {"label": '{"label": "major"}'}], + }, + ], + } + + +def test_systemone_questions_without_instructions_become_valid_openai_questions() -> None: + request: Final = _SYSTEMONE_BODY.validate_python( + { + "state": "Screen cracked", + "questions": { + "damaged": {"type": "noul", "criteria": {"false": None}}, + "action": {"type": "choice", "criteria": {"refund": None, "escalate": None}}, + "severity": {"type": "score", "criteria": ["minor", "major"]}, + }, + } + ) + + body: Final = ir_to_openai_request("gpt-6-luna", systemone_request_to_ir(request)) + + assert all(question.instructions for question in _OPENAI_BODY.validate_python(body).questions) + + +@pytest.mark.parametrize( + "question", + ({"type": "choice", "criteria": {"refund": None}}, {"type": "score", "criteria": ["minor"]}), + ids=("choice", "score"), +) +def test_a_systemone_question_with_one_option_is_unsupported_by_openai(question: Mapping[str, object]) -> None: + request: Final = _SYSTEMONE_BODY.validate_python( + { + "state": "Screen cracked", + "questions": {"damaged": {"type": "noul", "instructions": "Damaged?"}, "q": question}, + } + ) + + assert isinstance(ir_to_openai_request("gpt-6-luna", systemone_request_to_ir(request)), UnsupportedDecisionsRequest) + + +def test_an_openai_response_becomes_systemone_answers_with_the_callers_score_labels() -> None: + ir: Final = systemone_request_to_ir( + _SYSTEMONE_BODY.validate_python( + { + "state": "Screen cracked", + "questions": { + "damaged": {"type": "noul", "instructions": "Damaged?"}, + "action": {"type": "choice", "criteria": {"true": None, "escalate": None}}, + "severity": {"type": "score", "criteria": ["minor", {"label": "major"}]}, + "fraud": {"type": "noul", "instructions": "Fraud?"}, + }, + } + ) + ) + payload: Final = { + **_OPENAI_RESPONSE, + "answers": [ + _PREDICATE_ANSWER, + _CHOICE_ANSWER, + { + "type": "score", + "name": "severity", + "score": 0.7, + "probabilities": [ + {"value": 0, "label": "minor", "probability": 0.3}, + {"value": 1, "label": '{"label": "major"}', "probability": 0.7}, + ], + "confidence": 0.6, + }, + _REFUSAL_ANSWER, + ], + } + + response: Final = ir_to_systemone_response(OpenAIDecisionsConfig().parse_response(payload, ir), ir) + + assert response.model_dump(mode="json") == { + "model": "gpt-6-luna", + "answers": { + "damaged": {"type": "noul", "noul": 0.95}, + "action": { + "type": "choice", + "choice": "true", + "confidence": 0.8, + "probabilities": {"true": 0.9, "escalate": 0.1}, + }, + "severity": { + "type": "score", + "score": 0.7, + "confidence": 0.6, + "legend": {"0": "minor", "1": {"label": "major"}}, + "probabilities": {"0": 0.3, "1": 0.7}, + }, + }, + "usage": {"input_tokens": 383, "output_tokens": 2}, + } + assert response.usage is not None + assert (response.usage.cached_tokens, response.usage.cache_write_tokens) == (256, 64) diff --git a/tests/unit/llms/vertex_ai/realtime/test_vertex_ai_realtime_transformation.py b/tests/unit/llms/vertex_ai/realtime/test_vertex_ai_realtime_transformation.py index 14b3bdb48a1..d6986a2697a 100644 --- a/tests/unit/llms/vertex_ai/realtime/test_vertex_ai_realtime_transformation.py +++ b/tests/unit/llms/vertex_ai/realtime/test_vertex_ai_realtime_transformation.py @@ -2,7 +2,7 @@ Unit tests for VertexAIRealtimeConfig. Validates: -- URL construction (regional and global) +- URL construction (regional, multi-region and global) - Auth headers (Bearer token + project header) - Session setup message format - Full text-in / text-out round-trip via RealTimeStreaming with a mocked @@ -10,12 +10,14 @@ Validates: """ import json +from typing import Final from unittest.mock import AsyncMock, MagicMock import pytest import websockets.exceptions # registers websockets.exceptions on the websockets namespace import litellm +from litellm.llms.vertex_ai.common_utils import get_vertex_base_url from litellm.llms.vertex_ai.realtime.transformation import VertexAIRealtimeConfig # --------------------------------------------------------------------------- @@ -45,6 +47,26 @@ def test_get_complete_url_global(): ) +@pytest.mark.parametrize("location", ["us", "eu"]) +def test_get_complete_url_multi_region_uses_rep_host(location: str): + cfg: Final = VertexAIRealtimeConfig(access_token="tok", project="my-proj", location=location) + url: Final = cfg.get_complete_url(api_base=None, model="gemini-3.8-live") + # Google documents the multi-region Vertex endpoints as aiplatform.{us,eu}.rep.googleapis.com + # (https://docs.cloud.google.com/vertex-ai/generative-ai/docs/learn/locations, read 2026-10-07) + assert url == ( + f"wss://aiplatform.{location}.rep.googleapis.com" + "/ws/google.cloud.aiplatform.v1.LlmBidiService/BidiGenerateContent" + ) + + +@pytest.mark.parametrize("location", ["us", "eu", "global", "us-central1", "europe-west4"]) +def test_get_complete_url_host_matches_shared_vertex_host(location: str): + cfg: Final = VertexAIRealtimeConfig(access_token="tok", project="my-proj", location=location) + url: Final = cfg.get_complete_url(api_base=None, model="gemini-3.8-live") + shared_host: Final = get_vertex_base_url(location).removeprefix("https://") + assert url == f"wss://{shared_host}/ws/google.cloud.aiplatform.v1.LlmBidiService/BidiGenerateContent" + + def test_get_complete_url_custom_api_base(): cfg = VertexAIRealtimeConfig( access_token="tok", project="my-proj", location="us-central1" diff --git a/tests/unit/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py b/tests/unit/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py index 8e439cc822f..281bb8f413b 100644 --- a/tests/unit/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py +++ b/tests/unit/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py @@ -14,7 +14,7 @@ from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( MCPRequestHandler, UnloadableEntitlementError, _agent_capped_servers, - _is_mcp_admitted_user_subject, + is_mcp_admitted_user_subject, ) from litellm.proxy._types import ( LiteLLM_ObjectPermissionTable, @@ -364,7 +364,7 @@ class TestMCPRequestHandler: "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager", mock_manager, ), - patch.object(MCPRequestHandler, "_get_mcp_servers_from_access_groups", AsyncMock(return_value=[])), + patch.object(MCPRequestHandler, "get_mcp_servers_from_access_groups", AsyncMock(return_value=[])), ): result = await MCPRequestHandler._get_allowed_mcp_servers_for_key(user_api_key_auth) @@ -385,7 +385,7 @@ class TestMCPRequestHandler: "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager", mock_manager, ), - patch.object(MCPRequestHandler, "_get_mcp_servers_from_access_groups", AsyncMock(return_value=[])), + patch.object(MCPRequestHandler, "get_mcp_servers_from_access_groups", AsyncMock(return_value=[])), ): result = await MCPRequestHandler._get_allowed_mcp_servers_for_key(user_api_key_auth) @@ -405,7 +405,7 @@ class TestMCPRequestHandler: "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager", mock_manager, ), - patch.object(MCPRequestHandler, "_get_mcp_servers_from_access_groups", AsyncMock(return_value=[])), + patch.object(MCPRequestHandler, "get_mcp_servers_from_access_groups", AsyncMock(return_value=[])), patch.object(MCPRequestHandler, "_get_allowed_mcp_servers_for_team", AsyncMock(return_value=[])), patch.object(MCPRequestHandler, "_get_key_access_group_mcp_server_extras", AsyncMock(return_value=[])), ): @@ -430,7 +430,7 @@ class TestMCPRequestHandler: "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager", mock_manager, ), - patch.object(MCPRequestHandler, "_get_mcp_servers_from_access_groups", AsyncMock(return_value=[])), + patch.object(MCPRequestHandler, "get_mcp_servers_from_access_groups", AsyncMock(return_value=[])), patch.object( MCPRequestHandler, "_get_allowed_mcp_servers_for_team", @@ -733,7 +733,7 @@ class TestMCPRequestHandler: mock_manager, ), patch.object( # test-quality-ok: access-group lookup hits the DB, not under test here - MCPRequestHandler, "_get_mcp_servers_from_access_groups", AsyncMock(return_value=[]) + MCPRequestHandler, "get_mcp_servers_from_access_groups", AsyncMock(return_value=[]) ), ): servers = await MCPRequestHandler._team_granted_servers(team_obj, []) @@ -763,7 +763,7 @@ class TestMCPRequestHandler: "litellm.proxy.auth.auth_checks.get_team_object", AsyncMock(return_value=team_obj) ), patch( # test-quality-ok: access-group lookup hits the DB, not under test here - "litellm.proxy.auth.auth_checks._get_mcp_server_ids_from_access_groups", + "litellm.proxy.auth.auth_checks.get_mcp_server_ids_from_access_groups", AsyncMock(return_value=[]), ), patch( # test-quality-ok: isolate the MCP registry, same seam as the sibling tests @@ -801,7 +801,7 @@ class TestMCPRequestHandler: MCPRequestHandler, "_get_key_access_group_mcp_server_extras", AsyncMock(return_value=[]) ), patch.object( # test-quality-ok: access-group lookup hits the DB, not under test here - MCPRequestHandler, "_get_mcp_servers_from_access_groups", AsyncMock(return_value=[]) + MCPRequestHandler, "get_mcp_servers_from_access_groups", AsyncMock(return_value=[]) ), patch( # test-quality-ok: isolate the MCP registry, same seam as the sibling tests "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager", @@ -838,7 +838,7 @@ class TestMCPRequestHandler: MCPRequestHandler, "_get_key_access_group_mcp_server_extras", AsyncMock(return_value=[]) ), patch.object( # test-quality-ok: access-group lookup hits the DB, not under test here - MCPRequestHandler, "_get_mcp_servers_from_access_groups", AsyncMock(return_value=[]) + MCPRequestHandler, "get_mcp_servers_from_access_groups", AsyncMock(return_value=[]) ), patch( # test-quality-ok: isolate the MCP registry, same seam as the sibling tests "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager", @@ -880,7 +880,7 @@ class TestMCPRequestHandler: "litellm.proxy.auth.auth_checks.get_team_object", AsyncMock(return_value=team_obj) ), patch( # test-quality-ok: access-group lookup hits the DB, not under test here - "litellm.proxy.auth.auth_checks._get_mcp_server_ids_from_access_groups", + "litellm.proxy.auth.auth_checks.get_mcp_server_ids_from_access_groups", AsyncMock(return_value=[]), ), patch( # test-quality-ok: isolate the MCP registry, same seam as the sibling tests @@ -940,7 +940,7 @@ class TestMCPRequestHandler: MCPRequestHandler, "_get_org_object_permission", AsyncMock(return_value=org_object_permission) ), patch.object( # test-quality-ok: access-group lookup hits the DB, not under test here - MCPRequestHandler, "_get_mcp_servers_from_access_groups", AsyncMock(return_value=[]) + MCPRequestHandler, "get_mcp_servers_from_access_groups", AsyncMock(return_value=[]) ), patch( # test-quality-ok: isolate the MCP registry, same seam as the sibling tests "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager", @@ -992,7 +992,7 @@ class TestMCPRequestHandler: MCPRequestHandler, "_get_user_object_permission", AsyncMock(return_value=user_object_permission) ), patch.object( # test-quality-ok: access-group lookup hits the DB, not under test here - MCPRequestHandler, "_get_mcp_servers_from_access_groups", AsyncMock(return_value=[]) + MCPRequestHandler, "get_mcp_servers_from_access_groups", AsyncMock(return_value=[]) ), patch( # test-quality-ok: isolate the MCP registry, same seam as the sibling tests "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager", @@ -1045,7 +1045,7 @@ class TestMCPRequestHandler: MCPRequestHandler, "_get_org_object_permission", AsyncMock(return_value=org_object_permission) ), patch.object( # test-quality-ok: access-group lookup hits the DB, not under test here - MCPRequestHandler, "_get_mcp_servers_from_access_groups", AsyncMock(return_value=[]) + MCPRequestHandler, "get_mcp_servers_from_access_groups", AsyncMock(return_value=[]) ), patch( # test-quality-ok: isolate the MCP registry, same seam as the sibling tests "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager", @@ -1067,7 +1067,7 @@ class TestMCPRequestHandler: MCPRequestHandler, "_get_user_object_permission", AsyncMock(return_value=user_object_permission) ), patch.object( # test-quality-ok: access-group lookup hits the DB, not under test here - MCPRequestHandler, "_get_mcp_servers_from_access_groups", AsyncMock(return_value=[]) + MCPRequestHandler, "get_mcp_servers_from_access_groups", AsyncMock(return_value=[]) ), patch( # test-quality-ok: isolate the MCP registry, same seam as the sibling tests "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager", @@ -1220,7 +1220,7 @@ class TestMCPRequestHandler: auth = UserAPIKeyAuth(api_key="k", access_group_ids=[]) with ( patch( - "litellm.proxy.auth.auth_checks._get_mcp_server_ids_from_access_groups", + "litellm.proxy.auth.auth_checks.get_mcp_server_ids_from_access_groups", new=AsyncMock(return_value=[]), ), patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as mock_mgr, @@ -1235,7 +1235,7 @@ class TestMCPRequestHandler: auth = UserAPIKeyAuth(api_key="k", access_group_ids=["grp-mcp"]) with ( patch( - "litellm.proxy.auth.auth_checks._get_mcp_server_ids_from_access_groups", + "litellm.proxy.auth.auth_checks.get_mcp_server_ids_from_access_groups", new=AsyncMock(return_value=["alias-a", "srv-b"]), ), patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as mock_mgr, @@ -1249,7 +1249,7 @@ class TestMCPRequestHandler: """Resolution failures degrade to no grants rather than raising.""" auth = UserAPIKeyAuth(api_key="k", access_group_ids=["grp-mcp"]) with patch( - "litellm.proxy.auth.auth_checks._get_mcp_server_ids_from_access_groups", + "litellm.proxy.auth.auth_checks.get_mcp_server_ids_from_access_groups", new=AsyncMock(side_effect=Exception("db down")), ): result = await MCPRequestHandler._get_key_access_group_mcp_server_extras(auth) @@ -1720,7 +1720,7 @@ class TestMCPOAuth2AuthFlow: [b"sk-litellm-valid-key", b"Bearer sk-litellm-valid-key", b"bearer sk-litellm-valid-key"], ) async def test_x_litellm_api_key_survives_bearer_only_strip(self, header_value): - from litellm.proxy.auth.user_api_key_auth import _get_bearer_token + from litellm.proxy.auth.user_api_key_auth import get_bearer_token scope = { "type": "http", @@ -1741,7 +1741,7 @@ class TestMCPOAuth2AuthFlow: auth_result, *_rest = await MCPRequestHandler.process_mcp_request(scope) mock_auth.assert_called_once() - assert _get_bearer_token(api_key=mock_auth.call_args.kwargs["api_key"]) == "sk-litellm-valid-key" + assert get_bearer_token(api_key=mock_auth.call_args.kwargs["api_key"]) == "sk-litellm-valid-key" assert auth_result.user_id == "test-user" async def test_litellm_key_in_authorization_backward_compat(self): @@ -4135,7 +4135,7 @@ async def test_get_allowed_mcp_servers_for_team_uses_helper(): ), patch.object( MCPRequestHandler, - "_get_mcp_servers_from_access_groups", + "get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=["group-server1", "group-server2"], ) as mock_get_access_group_servers, @@ -4311,7 +4311,7 @@ async def test_get_allowed_mcp_servers_for_key_prefers_in_memory_permission(): "litellm.proxy.auth.auth_checks.get_object_permission", new_callable=AsyncMock, ) as mock_get_perm: - with patch.object(MCPRequestHandler, "_get_mcp_servers_from_access_groups") as mock_access_groups: + with patch.object(MCPRequestHandler, "get_mcp_servers_from_access_groups") as mock_access_groups: mock_access_groups.return_value = ["group-server"] result = await MCPRequestHandler._get_allowed_mcp_servers_for_key(user_api_key_auth) @@ -4706,7 +4706,7 @@ class TestAgentMCPPermissions: mock_manager, ), patch.object( # test-quality-ok: access-group lookup hits the DB, not under test here - MCPRequestHandler, "_get_mcp_servers_from_access_groups", AsyncMock(return_value=[]) + MCPRequestHandler, "get_mcp_servers_from_access_groups", AsyncMock(return_value=[]) ), ) @@ -4944,7 +4944,7 @@ async def test_tool_permission_servers_included_in_allowed_servers(): patch.object(MCPRequestHandler, "_get_key_object_permission", return_value=perm), patch.object( MCPRequestHandler, - "_get_mcp_servers_from_access_groups", + "get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ), @@ -5095,7 +5095,7 @@ class TestOrgMCPPermissions: ), patch.object( MCPRequestHandler, - "_get_mcp_servers_from_access_groups", + "get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ), @@ -5121,7 +5121,7 @@ class TestOrgMCPPermissions: ), patch.object( MCPRequestHandler, - "_get_mcp_servers_from_access_groups", + "get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=["group_server_1"], ), @@ -5147,7 +5147,7 @@ class TestOrgMCPPermissions: ), patch.object( MCPRequestHandler, - "_get_mcp_servers_from_access_groups", + "get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ), @@ -5552,7 +5552,7 @@ async def test_team_access_group_ids_resolve_to_mcp_servers(): return_value=mock_team, ), patch( - "litellm.proxy.auth.auth_checks._get_mcp_server_ids_from_access_groups", + "litellm.proxy.auth.auth_checks.get_mcp_server_ids_from_access_groups", new_callable=AsyncMock, return_value=["srv-stripe"], ) as mock_resolver, @@ -5611,7 +5611,7 @@ async def test_team_access_group_ids_union_with_object_permission(): return_value=mock_team, ), patch( - "litellm.proxy.auth.auth_checks._get_mcp_server_ids_from_access_groups", + "litellm.proxy.auth.auth_checks.get_mcp_server_ids_from_access_groups", new_callable=AsyncMock, return_value=["srv-stripe"], ), @@ -5649,7 +5649,7 @@ async def test_team_access_group_ids_empty_returns_no_extras(): return_value=mock_team, ), patch( - "litellm.proxy.auth.auth_checks._get_mcp_server_ids_from_access_groups", + "litellm.proxy.auth.auth_checks.get_mcp_server_ids_from_access_groups", new_callable=AsyncMock, return_value=[], ) as mock_resolver, @@ -5712,7 +5712,7 @@ async def test_allowed_mcp_servers_for_key_excludes_access_group_ids(): with ( patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), patch( - "litellm.proxy.auth.auth_checks._get_mcp_server_ids_from_access_groups", + "litellm.proxy.auth.auth_checks.get_mcp_server_ids_from_access_groups", new_callable=AsyncMock, return_value=["srv-stripe"], ) as mock_resolver, @@ -5760,7 +5760,7 @@ async def test_allowed_mcp_servers_for_key_uses_object_permission_not_access_gro with ( patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), patch( - "litellm.proxy.auth.auth_checks._get_mcp_server_ids_from_access_groups", + "litellm.proxy.auth.auth_checks.get_mcp_server_ids_from_access_groups", new_callable=AsyncMock, return_value=["srv-stripe"], ) as mock_resolver, @@ -5787,7 +5787,7 @@ async def test_get_allowed_mcp_servers_surfaces_ungated_key_access_group_grant_e patches = _patch_proxy_server_globals_for_mcp() + [ patch( - "litellm.proxy.auth.auth_checks._get_mcp_server_ids_from_access_groups", + "litellm.proxy.auth.auth_checks.get_mcp_server_ids_from_access_groups", new_callable=AsyncMock, return_value=["srv-deepwiki"], ), @@ -5887,7 +5887,7 @@ async def test_get_allowed_mcp_servers_for_team_expands_all_proxy_sentinel_dynam return_value=team_obj, ), patch( - "litellm.proxy.auth.auth_checks._get_mcp_server_ids_from_access_groups", + "litellm.proxy.auth.auth_checks.get_mcp_server_ids_from_access_groups", new_callable=AsyncMock, return_value=[], ), @@ -6022,13 +6022,13 @@ async def test_get_allowed_mcp_servers_team_all_proxy_key_scoped_to_one_end_to_e return_value=team_obj, ), patch( - "litellm.proxy.auth.auth_checks._get_mcp_server_ids_from_access_groups", + "litellm.proxy.auth.auth_checks.get_mcp_server_ids_from_access_groups", new_callable=AsyncMock, return_value=[], ), patch.object( MCPRequestHandler, - "_get_mcp_servers_from_access_groups", + "get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ), @@ -6737,7 +6737,7 @@ class TestMCPDcrBridgeDelegateAdmission: assert exc_info.value.status_code == 401 - _POLICY_GATE = "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp._run_centralized_common_checks" + _POLICY_GATE = "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.run_centralized_common_checks" async def _enforce_with_gate_error(self, error): """Drive _enforce_admitted_live_policy with the centralized gate raising ``error`` and return @@ -8273,7 +8273,7 @@ class TestUserSubjectTeamUnion: patch("litellm.proxy.auth.auth_checks.get_team_object", _get_team_object), patch("litellm.proxy.auth.auth_checks.get_user_object", _get_user_object), patch("litellm.proxy.auth.auth_checks.get_org_object", _get_org_object), - patch("litellm.proxy.auth.auth_checks._get_mcp_server_ids_from_access_groups", AsyncMock(return_value=[])), + patch("litellm.proxy.auth.auth_checks.get_mcp_server_ids_from_access_groups", AsyncMock(return_value=[])), patch("litellm.proxy.proxy_server.get_current_spend", _spend_from_fallback), patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()), @@ -8460,7 +8460,7 @@ class TestUserSubjectTeamUnion: one cross-team user drain several teams' buckets on a single call, blocking their other members for access those teams did not provide. Exactly one source is charged, and it is the SAME source billing picks — one owner for both, so they cannot disagree.""" - from litellm.proxy.hooks.parallel_request_limiter_v3 import _PROXY_MaxParallelRequestsHandler_v3 + from litellm.proxy.hooks.parallel_request_limiter_v3 import PROXY_MaxParallelRequestsHandler_v3 t1 = _make_team("t1", ["srv1"]) t1.metadata = {"mcp_rpm_limit": {"srv1": 5}} @@ -8474,7 +8474,7 @@ class TestUserSubjectTeamUnion: assert billed is not None and billed.team_id == "t1", "throttling and billing pick the same source" descriptors: list = [] - limiter = _PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=MagicMock()) + limiter = PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=MagicMock()) limiter._add_mcp_per_team_rate_limit_descriptor(auth, "srv1", descriptors) charged = {d["value"]: d["rate_limit"]["requests_per_unit"] for d in descriptors} assert charged == {"t1:srv1": 5}, "only the attributing team's bucket is charged" @@ -8536,7 +8536,7 @@ class TestUserSubjectTeamUnion: server = MagicMock(server_id="srv1") with self._patch(teams_by_id={"t-grant": t_grant}, user_teams=["t-grant"]): with patch( - "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager._get_mcp_server_from_tool_name", + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager.get_mcp_server_from_tool_name", MagicMock(return_value=server), ): billed = await MCPRequestHandler.billing_auth_for_tool_call(auth, tool_name="t-grant/tool_a") @@ -8985,7 +8985,7 @@ class TestUserSubjectTeamUnion: api_key="sk-real-key", metadata={"mcp_admitted_user_subject": True}, # caller-forged marker in key metadata ) - assert _is_mcp_admitted_user_subject(forged) is False + assert is_mcp_admitted_user_subject(forged) is False with self._patch(teams_by_id=teams, user_teams=["team-a", "team-b"]): assert await MCPRequestHandler._team_ids_for_mcp_grant(forged) == [] assert await MCPRequestHandler._get_allowed_mcp_servers_for_team(forged) == [] @@ -9065,8 +9065,8 @@ class TestUserSubjectTeamUnion: via_validate = UserAPIKeyAuth.model_validate({"user_id": "u", "mcp_admitted_user_subject": True}) assert via_kwarg.mcp_admitted_user_subject is False assert via_validate.mcp_admitted_user_subject is False - assert _is_mcp_admitted_user_subject(via_kwarg) is False - assert _is_mcp_admitted_user_subject(via_validate) is False + assert is_mcp_admitted_user_subject(via_kwarg) is False + assert is_mcp_admitted_user_subject(via_validate) is False @pytest.mark.asyncio @@ -9126,7 +9126,7 @@ class TestAdmittedSubjectPerTeamOrgCap: patch("litellm.proxy.auth.auth_checks.get_org_object", _get_org_object), patch("litellm.proxy.auth.auth_checks.get_object_permission", _get_object_permission), patch( - "litellm.proxy.auth.auth_checks._get_mcp_server_ids_from_access_groups", + "litellm.proxy.auth.auth_checks.get_mcp_server_ids_from_access_groups", AsyncMock(return_value=[]), ), patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), @@ -9574,7 +9574,7 @@ class TestUserMCPEntitlement: ), patch.object( MCPRequestHandler, - "_get_mcp_servers_from_access_groups", + "get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ), @@ -9731,7 +9731,7 @@ class TestUserMCPEntitlement: with self._entitled(self._perm(tool_permissions={"srv-a": ["read"]})): with patch.object( MCPRequestHandler, - "_get_mcp_servers_from_access_groups", + "get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ): diff --git a/tests/unit/proxy/_experimental/mcp_server/test_db_credentials.py b/tests/unit/proxy/_experimental/mcp_server/test_db_credentials.py index f57b5121fa5..de86c1b62f4 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_db_credentials.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_db_credentials.py @@ -1541,7 +1541,7 @@ async def test_master_key_rotation_reencrypts_oauth_client_store(monkeypatch): ) ) - monkeypatch.setattr(enc, "_get_salt_key", lambda: key_old) + monkeypatch.setattr(enc, "get_salt_key", lambda: key_old) prisma = MagicMock() prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[]) @@ -1558,7 +1558,7 @@ async def test_master_key_rotation_reencrypts_oauth_client_store(monkeypatch): assert store_update.await_args.kwargs["where"] == {"server_id": "config_faros"} rotated_blob = store_update.await_args.kwargs["data"]["credentials"] - monkeypatch.setattr(enc, "_get_salt_key", lambda: key_new) + monkeypatch.setattr(enc, "get_salt_key", lambda: key_new) recovered = decrypt_credentials(credentials=json.loads(rotated_blob)) assert recovered["client_id"] == "cid-123" assert recovered["client_secret"] == "sec-456" diff --git a/tests/unit/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/unit/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index bf21e3434ca..23bdfa3eadf 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -251,7 +251,7 @@ async def test_register_resolves_cold_oauth_metadata(): "_discover_oauth_metadata_for_server", new=AsyncMock(return_value=_resolved_oauth_metadata()), ) as discovery, - patch.object(discoverable_endpoints, "_read_request_body", new=AsyncMock(return_value={})), + patch.object(discoverable_endpoints, "read_request_body", new=AsyncMock(return_value={})), patch.object( discoverable_endpoints, "get_async_httpx_client", @@ -307,7 +307,7 @@ async def test_register_route_bridge_missing_registration_url_joins_discovery(): ) as discovery, patch.object( # test-quality-ok: the MagicMock Request carries no body; this seam feeds the RFC 7591 redirect_uris discoverable_endpoints, - "_read_request_body", + "read_request_body", new=AsyncMock(return_value={"redirect_uris": ["https://client.example.com/cb"]}), ), patch.object( # test-quality-ok: keeps the DCR POST off the network so its target URL can be asserted @@ -766,7 +766,7 @@ async def test_register_client_without_mcp_server_name_returns_dummy(server_name mock_request.base_url = "https://proxy.litellm.example/" mock_request.headers = {} with patch( - "litellm.proxy._experimental.mcp_server.discoverable_endpoints._read_request_body", + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.read_request_body", new=AsyncMock(return_value={}), ): result = await register_client(request=mock_request, mcp_server_name=server_name) @@ -817,7 +817,7 @@ async def test_register_client_returns_existing_server_credentials(use_root): try: with patch( - "litellm.proxy._experimental.mcp_server.discoverable_endpoints._read_request_body", + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.read_request_body", new=AsyncMock(return_value={}), ): result = await register_client( @@ -889,7 +889,7 @@ async def test_register_client_remote_registration_success(): try: with ( patch( - "litellm.proxy._experimental.mcp_server.discoverable_endpoints._read_request_body", + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.read_request_body", new=AsyncMock(return_value=request_payload), ), patch( @@ -975,7 +975,7 @@ async def test_register_client_non_bridge_returns_client_redirect_not_gateway_ca try: with ( patch( - "litellm.proxy._experimental.mcp_server.discoverable_endpoints._read_request_body", + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.read_request_body", new=AsyncMock(return_value=request_payload), ), patch( @@ -1031,7 +1031,7 @@ async def test_register_client_admin_client_id_echoes_client_redirect_uris(): try: with patch( - "litellm.proxy._experimental.mcp_server.discoverable_endpoints._read_request_body", + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.read_request_body", new=AsyncMock(return_value={"redirect_uris": [client_redirect]}), ): result = await register_client(request=mock_request, mcp_server_name=oauth2_server.server_name) @@ -1107,7 +1107,7 @@ async def test_dcr_full_loop_lands_on_client_redirect_not_gateway_callback(monke try: with ( patch( - "litellm.proxy._experimental.mcp_server.discoverable_endpoints._read_request_body", + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.read_request_body", new=AsyncMock( return_value={ "client_name": "Open WebUI", @@ -1268,7 +1268,7 @@ async def test_register_client_malformed_redirect_uris_falls_back_to_gateway_cal try: with patch( - "litellm.proxy._experimental.mcp_server.discoverable_endpoints._read_request_body", + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.read_request_body", new=AsyncMock(return_value={"redirect_uris": malformed_redirect_uris}), ): result = await register_client(request=mock_request, mcp_server_name=oauth2_server.server_name) @@ -1312,7 +1312,7 @@ async def test_register_client_valid_multi_redirect_uris_all_echoed(): client_redirects = ["https://app.example/cb", "http://127.0.0.1:6274/callback"] try: with patch( - "litellm.proxy._experimental.mcp_server.discoverable_endpoints._read_request_body", + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.read_request_body", new=AsyncMock(return_value={"redirect_uris": client_redirects}), ): result = await register_client(request=mock_request, mcp_server_name=oauth2_server.server_name) @@ -2258,7 +2258,7 @@ async def test_register_client_reuses_existing_client_id_without_re_dcr(): try: with ( patch( - "litellm.proxy._experimental.mcp_server.discoverable_endpoints._read_request_body", + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.read_request_body", new=AsyncMock(return_value=request_payload), ), patch( @@ -2337,7 +2337,7 @@ async def test_public_register_route_does_not_persist_client_credentials(): try: with ( patch( - "litellm.proxy._experimental.mcp_server.discoverable_endpoints._read_request_body", + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.read_request_body", new=AsyncMock(return_value=request_payload), ), patch( @@ -2631,7 +2631,7 @@ async def test_register_client_respects_x_forwarded_proto(): mock_request.headers = {"X-Forwarded-Proto": "https"} # Behind HTTPS proxy with patch( - "litellm.proxy._experimental.mcp_server.discoverable_endpoints._read_request_body", + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.read_request_body", new=AsyncMock(return_value={}), ): result = await register_client(request=mock_request) @@ -3737,7 +3737,7 @@ async def test_register_client_resolves_server_by_id_when_name_lookup_fails(): with ( patch.object(global_mcp_server_manager, "get_mcp_server_by_name", return_value=None) as by_name, # test-quality-ok: resolver seam patch.object(global_mcp_server_manager, "get_mcp_server_by_id", return_value=server) as by_id, # test-quality-ok: resolver seam - patch.object(discoverable_endpoints, "_read_request_body", new=AsyncMock(return_value={})), # test-quality-ok: request seam + patch.object(discoverable_endpoints, "read_request_body", new=AsyncMock(return_value={})), # test-quality-ok: request seam patch.object(discoverable_endpoints, "get_async_httpx_client", return_value=client), # test-quality-ok: HTTP seam ): result = await discoverable_endpoints.register_client(request=request, mcp_server_name=server.server_id) @@ -4065,7 +4065,7 @@ async def test_register_root_does_aggregate_dcr_not_single_server_resolution(): try: with ( patch( - "litellm.proxy._experimental.mcp_server.discoverable_endpoints._read_request_body", + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.read_request_body", new=AsyncMock(return_value={"redirect_uris": ["https://claude.ai/cb"]}), ), patch("litellm.proxy.proxy_server.master_key", "sk-test-salt-for-lit3637"), @@ -4106,7 +4106,7 @@ async def test_register_root_does_not_leak_a_private_server(): try: with ( patch( - "litellm.proxy._experimental.mcp_server.discoverable_endpoints._read_request_body", + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.read_request_body", new=AsyncMock(return_value={"redirect_uris": ["https://claude.ai/cb"]}), ), patch( @@ -5879,7 +5879,7 @@ async def test_interactive_bridge_authorize_seals_sso_user_into_state(): with ( patch( - "litellm.proxy._experimental.mcp_server.byok_oauth_endpoints._user_id_from_session_cookie", + "litellm.proxy._experimental.mcp_server.byok_oauth_endpoints.user_id_from_session_cookie", return_value="sso-user-42", ), patch( @@ -5944,7 +5944,7 @@ async def test_bridge_authorize_gates_on_the_egress_server_access_resolver(user_ try: with ( patch( - "litellm.proxy._experimental.mcp_server.byok_oauth_endpoints._user_id_from_session_cookie", + "litellm.proxy._experimental.mcp_server.byok_oauth_endpoints.user_id_from_session_cookie", return_value="bridge-user-1", ), patch( @@ -6015,7 +6015,7 @@ async def test_bridge_authorize_reload_failure_denies_or_stays_retryable(reload_ try: with ( patch( - "litellm.proxy._experimental.mcp_server.byok_oauth_endpoints._user_id_from_session_cookie", + "litellm.proxy._experimental.mcp_server.byok_oauth_endpoints.user_id_from_session_cookie", return_value="bridge-user-1", ), patch( @@ -6061,7 +6061,7 @@ async def test_interactive_bridge_authorize_without_session_redirects_to_login() server = _bridge_server(auth_type=MCPAuth.oauth_delegate, client_id="admin-client", registration_url=None) with patch( - "litellm.proxy._experimental.mcp_server.byok_oauth_endpoints._user_id_from_session_cookie", + "litellm.proxy._experimental.mcp_server.byok_oauth_endpoints.user_id_from_session_cookie", return_value=None, ): response = await authorize_with_server( @@ -6600,7 +6600,7 @@ async def test_revalidate_active_subject_dispatches_on_subject_type(): new=AsyncMock(return_value=_ResolvedKey(key_hash="kh", key=MagicMock())), ) as key_reload, patch( - "litellm.proxy._experimental.mcp_server.bridge_token_flow._reload_active_user_by_id", + "litellm.proxy._experimental.mcp_server.bridge_token_flow.reload_active_user_by_id", new=AsyncMock(return_value=None), ) as user_reload, ): @@ -6614,7 +6614,7 @@ async def test_revalidate_active_subject_dispatches_on_subject_type(): new=AsyncMock(), ) as key_reload2, patch( - "litellm.proxy._experimental.mcp_server.bridge_token_flow._reload_active_user_by_id", + "litellm.proxy._experimental.mcp_server.bridge_token_flow.reload_active_user_by_id", new=AsyncMock(return_value="no_active_key"), ) as user_reload2, ): @@ -7262,7 +7262,7 @@ def test_bridge_reported_expires_in_can_be_zero_at_jwt_exp_boundary(): from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( _BridgeMintReady, - _finish_bridge_mint, + finish_bridge_mint, ) from litellm.proxy._experimental.mcp_server.outbound_credentials.bridge_credentials import ( envelope_keys_from_master_key, @@ -7274,7 +7274,7 @@ def test_bridge_reported_expires_in_can_be_zero_at_jwt_exp_boundary(): identity=key_hash_identity(server_id="bridge_srv", key_hash="hashed-litellm-key-77"), keys=envelope_keys_from_master_key(_BRIDGE_MASTER_KEY), ) - response = _finish_bridge_mint( + response = finish_bridge_mint( ready=ready, mcp_server=_bridge_server(auth_type=MCPAuth.oauth_delegate), token_response={"access_token": "UP", "expires_in": 1}, @@ -7439,7 +7439,7 @@ async def _exchange_persistence_attempted_for_auth_type(auth_type) -> bool: return_value=fake_http_client, ), patch( - "litellm.proxy._experimental.mcp_server.discoverable_endpoints._extract_user_id_from_request", + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.extract_user_id_from_request", new_callable=AsyncMock, return_value="admin-user", ), @@ -7623,7 +7623,7 @@ async def test_extract_user_id_reads_x_litellm_api_key_header(proxy_globals): Authorization. Reading only Authorization dropped the identity, so the per-user token was never stored and the egress 401'd forever. Resolution must honor x-litellm-api-key.""" from litellm.proxy._experimental.mcp_server.bridge_token_flow import ( - _extract_user_id_from_request, + extract_user_id_from_request, ) from litellm.proxy._types import UserAPIKeyAuth, hash_token from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache @@ -7639,7 +7639,7 @@ async def test_extract_user_id_reads_x_litellm_api_key_header(proxy_globals): proxy_globals.prisma_client = object() request = _token_request({"x-litellm-api-key": f"Bearer {key}"}) - assert await _extract_user_id_from_request(request) == "alice" + assert await extract_user_id_from_request(request) == "alice" @pytest.mark.asyncio @@ -7648,7 +7648,7 @@ async def test_extract_user_id_rehydrates_cross_replica_dict_cache(proxy_globals Resolution must rehydrate it; the old getattr(cached, "user_id") returned None on a dict, which is exactly why a multi-replica gateway never found the stored token.""" from litellm.proxy._experimental.mcp_server.bridge_token_flow import ( - _extract_user_id_from_request, + extract_user_id_from_request, ) from litellm.proxy._types import hash_token from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache @@ -7660,7 +7660,7 @@ async def test_extract_user_id_rehydrates_cross_replica_dict_cache(proxy_globals proxy_globals.prisma_client = object() request = _token_request({"Authorization": f"Bearer {key}"}) - assert await _extract_user_id_from_request(request) == "alice" + assert await extract_user_id_from_request(request) == "alice" @pytest.mark.asyncio @@ -7669,7 +7669,7 @@ async def test_extract_user_id_falls_back_to_db_on_cache_miss(proxy_globals): cache-only peek and skipped the DB, so any replica that hadn't just authenticated the key failed to store the token.""" from litellm.proxy._experimental.mcp_server.bridge_token_flow import ( - _extract_user_id_from_request, + extract_user_id_from_request, ) from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache @@ -7684,14 +7684,14 @@ async def test_extract_user_id_falls_back_to_db_on_cache_miss(proxy_globals): proxy_globals.prisma_client = _FakePrisma() request = _token_request({"x-litellm-api-key": key}) - assert await _extract_user_id_from_request(request) == "db-bob" + assert await extract_user_id_from_request(request) == "db-bob" @pytest.mark.asyncio async def test_extract_user_id_none_without_litellm_key(proxy_globals): """No LiteLLM key on the request resolves to None without consulting the resolver.""" from litellm.proxy._experimental.mcp_server.bridge_token_flow import ( - _extract_user_id_from_request, + extract_user_id_from_request, ) from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache @@ -7699,7 +7699,7 @@ async def test_extract_user_id_none_without_litellm_key(proxy_globals): proxy_globals.prisma_client = object() request = _token_request({"content-type": "application/json"}) - assert await _extract_user_id_from_request(request) is None + assert await extract_user_id_from_request(request) is None @pytest.mark.asyncio @@ -7708,7 +7708,7 @@ async def test_extract_user_id_rejects_blocked_key(proxy_globals): checking blocked/expiry (the main auth pipeline does, and the public token endpoint bypasses it), so a revoked key could otherwise overwrite the stored per-user OAuth token for its user.""" from litellm.proxy._experimental.mcp_server.bridge_token_flow import ( - _extract_user_id_from_request, + extract_user_id_from_request, ) from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache @@ -7721,7 +7721,7 @@ async def test_extract_user_id_rejects_blocked_key(proxy_globals): proxy_globals.prisma_client = _FakePrisma() request = _token_request({"x-litellm-api-key": "sk-blocked-key"}) - assert await _extract_user_id_from_request(request) is None + assert await extract_user_id_from_request(request) is None @pytest.mark.asyncio @@ -7730,7 +7730,7 @@ async def test_extract_user_id_rejects_expired_key(proxy_globals): from datetime import datetime, timedelta, timezone from litellm.proxy._experimental.mcp_server.bridge_token_flow import ( - _extract_user_id_from_request, + extract_user_id_from_request, ) from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache @@ -7745,7 +7745,7 @@ async def test_extract_user_id_rejects_expired_key(proxy_globals): proxy_globals.prisma_client = _FakePrisma() request = _token_request({"x-litellm-api-key": "sk-expired-key"}) - assert await _extract_user_id_from_request(request) is None + assert await extract_user_id_from_request(request) is None @pytest.mark.asyncio @@ -7785,7 +7785,7 @@ async def test_resolve_active_litellm_key_resolves_key_without_user_id(proxy_glo blocked and expiry, and the key hash (not the user) is what the mint seals. The per-user token store still gets no user for such a key, since there is none to key a stored credential by.""" from litellm.proxy._experimental.mcp_server.bridge_token_flow import ( - _extract_user_id_from_request, + extract_user_id_from_request, _ResolvedKey, _resolve_active_litellm_key, ) @@ -7806,7 +7806,7 @@ async def test_resolve_active_litellm_key_resolves_key_without_user_id(proxy_glo resolved = await _resolve_active_litellm_key(request) assert isinstance(resolved, _ResolvedKey) assert resolved.key_hash == hash_token(key) - assert await _extract_user_id_from_request(request) is None + assert await extract_user_id_from_request(request) is None @pytest.mark.asyncio @@ -7979,7 +7979,7 @@ async def test_reload_active_user_by_id_missing_user_is_no_active_key(proxy_glob refresh path maps it to invalid_grant), not unresolvable/500. get_user_object catches the missing row and re-raises a bare ValueError, so a missing user must not be misclassified as a DB outage or an opaque gateway fault.""" - from litellm.proxy._experimental.mcp_server.bridge_token_flow import _reload_active_user_by_id + from litellm.proxy._experimental.mcp_server.bridge_token_flow import reload_active_user_by_id from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache proxy_globals.user_api_key_cache = UserApiKeyCache() @@ -7989,7 +7989,7 @@ async def test_reload_active_user_by_id_missing_user_is_no_active_key(proxy_glob "litellm.proxy.auth.auth_checks.get_user_object", new=AsyncMock(side_effect=_wrapped_user_lookup_error(Exception())), ): - assert await _reload_active_user_by_id("gone-user") == "no_active_key" + assert await reload_active_user_by_id("gone-user") == "no_active_key" @pytest.mark.asyncio @@ -7998,7 +7998,7 @@ async def test_reload_active_user_by_id_db_outage_is_unavailable(proxy_globals): a missing user, so the refresh path surfaces "unavailable" (a 503) rather than blaming the caller. get_user_object wraps the outage in a bare ValueError, so this exercises the chain-aware classifier; a raw ConnectionError would falsely pass even a chain-blind check because it is an OSError.""" - from litellm.proxy._experimental.mcp_server.bridge_token_flow import _reload_active_user_by_id + from litellm.proxy._experimental.mcp_server.bridge_token_flow import reload_active_user_by_id from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache proxy_globals.user_api_key_cache = UserApiKeyCache() @@ -8008,7 +8008,7 @@ async def test_reload_active_user_by_id_db_outage_is_unavailable(proxy_globals): "litellm.proxy.auth.auth_checks.get_user_object", new=AsyncMock(side_effect=_wrapped_user_lookup_error(ConnectionError("user database unreachable"))), ): - assert await _reload_active_user_by_id("sso-user-7") == "unavailable" + assert await reload_active_user_by_id("sso-user-7") == "unavailable" @pytest.mark.asyncio @@ -8018,7 +8018,7 @@ async def test_reload_active_user_by_id_permanent_engine_fault_is_faulted(proxy_ wraps the fault in a bare ValueError, so the classification has to read the wrapped cause.""" from prisma.engine.errors import MismatchedVersionsError - from litellm.proxy._experimental.mcp_server.bridge_token_flow import _reload_active_user_by_id + from litellm.proxy._experimental.mcp_server.bridge_token_flow import reload_active_user_by_id from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache proxy_globals.user_api_key_cache = UserApiKeyCache() @@ -8028,7 +8028,7 @@ async def test_reload_active_user_by_id_permanent_engine_fault_is_faulted(proxy_ "litellm.proxy.auth.auth_checks.get_user_object", new=AsyncMock(side_effect=_wrapped_user_lookup_error(MismatchedVersionsError(expected="1", got="2"))), ): - assert await _reload_active_user_by_id("sso-user-7") == "faulted" + assert await reload_active_user_by_id("sso-user-7") == "faulted" @pytest.mark.asyncio @@ -8067,7 +8067,7 @@ async def test_load_active_user_by_id_serves_a_cached_row_without_a_database_rea cached row answers without a database read, and only a caller that asks for the database row pays for one.""" from litellm.proxy._experimental.mcp_server.bridge_token_flow import ( - _reload_active_user_by_id, + reload_active_user_by_id, load_active_user_by_id, ) from litellm.proxy._types import LiteLLM_UserTable @@ -8090,7 +8090,7 @@ async def test_load_active_user_by_id_serves_a_cached_row_without_a_database_rea assert not isinstance(loaded, str) assert loaded.teams == ["team-a"] - assert await _reload_active_user_by_id("cached-jwt-user") is None + assert await reload_active_user_by_id("cached-jwt-user") is None prisma.db.litellm_usertable.find_unique.assert_not_awaited() @@ -8343,7 +8343,7 @@ async def test_register_client_rejects_non_oauth2_server(): try: with pytest.raises(HTTPException) as exc_info: with patch( - "litellm.proxy._experimental.mcp_server.discoverable_endpoints._read_request_body", + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.read_request_body", new=AsyncMock(return_value={}), ): await register_client(request=mock_request, mcp_server_name="access_group_server") @@ -8731,7 +8731,7 @@ async def test_token_exchange_refresh_passes_presented_refresh_ownership(): return_value=client, ), patch( # test-quality-ok: no injection seam exists for request identity extraction - "litellm.proxy._experimental.mcp_server.discoverable_endpoints._extract_user_id_from_request", + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.extract_user_id_from_request", new=AsyncMock(return_value="user-a"), ), patch( # test-quality-ok: captures the ownership value at the exchange boundary @@ -8797,7 +8797,7 @@ async def test_token_exchange_authorization_code_passes_no_refresh_ownership(mon return_value=client, ), patch( # test-quality-ok: no injection seam exists for request identity extraction - "litellm.proxy._experimental.mcp_server.discoverable_endpoints._extract_user_id_from_request", + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.extract_user_id_from_request", new=AsyncMock(return_value="user-a"), ), patch( # test-quality-ok: captures the ownership value at the exchange boundary @@ -9294,7 +9294,7 @@ async def test_hydrate_config_server_applies_stored_dcr_client(monkeypatch, lega issuer="https://idp.example", ) - monkeypatch.setattr(enc, "_get_salt_key", lambda: "salt-hydrate-key") + monkeypatch.setattr(enc, "get_salt_key", lambda: "salt-hydrate-key") stored_blob = safe_dumps( encrypt_credentials( credentials={**({} if legacy else {"dcr_issuer": "https://idp.example", "dcr_server_url": "https://resource.example/mcp"}), @@ -9351,7 +9351,7 @@ async def test_reuse_config_server_reads_store_with_real_crypto(monkeypatch): issuer="https://idp.example", ) - monkeypatch.setattr(enc, "_get_salt_key", lambda: "salt-reuse-key") + monkeypatch.setattr(enc, "get_salt_key", lambda: "salt-reuse-key") blob = safe_dumps( encrypt_credentials( credentials={"dcr_issuer": "https://idp.example", "dcr_server_url": "https://resource.example/mcp","client_id": "stored-client", "client_secret": "sec", "redirect_uris": ["https://x/callback"]}, @@ -11697,7 +11697,7 @@ def test_introspect_route_answers_for_authenticated_caller(monkeypatch): return None monkeypatch.setattr( - "litellm.proxy._experimental.mcp_server.discoverable_endpoints._reload_active_user_by_id", fake_reload + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.reload_active_user_by_id", fake_reload ) app = FastAPI() app.include_router(router) @@ -12153,7 +12153,7 @@ async def test_oauth_jwt_identity_rejects_untrusted_or_inactive_owner( from litellm.proxy import proxy_server from litellm.proxy._experimental.mcp_server import mcp_server_manager from litellm.proxy._experimental.mcp_server.bridge_token_flow import ( - _extract_user_id_from_request, authorize_oauth_credential_request, + extract_user_id_from_request, authorize_oauth_credential_request, ) allowed_servers: Final = AsyncMock(return_value=["server-a"]) @@ -12184,7 +12184,7 @@ async def test_oauth_jwt_identity_rejects_untrusted_or_inactive_owner( request: Final = _token_request({"Authorization": f"Bearer {bearer}"}) result: Final = ( await authorize_oauth_credential_request(request, "server-a") - if credential_write else await _extract_user_id_from_request(request) + if credential_write else await extract_user_id_from_request(request) ) assert result is None allowed_servers.assert_not_awaited() @@ -12196,7 +12196,7 @@ async def test_oauth_jwt_cannot_override_explicit_litellm_key( jwt_oauth_identity: tuple["JWTHandler", "RSAPrivateKey"], blocked: bool, ) -> None: - from litellm.proxy._experimental.mcp_server.bridge_token_flow import _extract_user_id_from_request + from litellm.proxy._experimental.mcp_server.bridge_token_flow import extract_user_id_from_request from litellm.proxy._types import UserAPIKeyAuth, hash_token handler, signing_key = jwt_oauth_identity @@ -12208,7 +12208,7 @@ async def test_oauth_jwt_cannot_override_explicit_litellm_key( "x-litellm-api-key": key, } ) - assert await _extract_user_id_from_request(request) == (None if blocked else "key-owner") + assert await extract_user_id_from_request(request) == (None if blocked else "key-owner") @pytest.mark.asyncio @@ -12220,7 +12220,7 @@ async def test_oauth_jwt_uses_configured_virtual_key_owner( mapping: str, ) -> None: from litellm.models.user import LiteLLM_UserTable - from litellm.proxy._experimental.mcp_server.bridge_token_flow import _extract_user_id_from_request + from litellm.proxy._experimental.mcp_server.bridge_token_flow import extract_user_id_from_request from litellm.proxy._types import UserAPIKeyAuth, UnregisteredJWTClientBehavior, hash_token from litellm.proxy.auth.auth_checks import jwt_key_mapping_cache_key @@ -12248,7 +12248,7 @@ async def test_oauth_jwt_uses_configured_virtual_key_owner( ) request: Final = _token_request({"Authorization": f"Bearer {_oauth_identity_jwt(signing_key)}"}) expected: Final = "jwt-owner" if mapping == "fallback" else "mapped-owner" if mapping == "active" else None - assert await _extract_user_id_from_request(request) == expected + assert await extract_user_id_from_request(request) == expected @pytest.mark.asyncio @@ -12257,13 +12257,13 @@ async def test_oauth_jwt_respects_custom_validation_and_email_policy( jwt_oauth_identity: tuple["JWTHandler", "RSAPrivateKey"], allowed_domain: str | None, ) -> None: - from litellm.proxy._experimental.mcp_server.bridge_token_flow import _extract_user_id_from_request + from litellm.proxy._experimental.mcp_server.bridge_token_flow import extract_user_id_from_request handler, signing_key = jwt_oauth_identity handler.litellm_jwtauth.custom_validate = lambda claims: True handler.litellm_jwtauth.user_allowed_email_domain = allowed_domain request: Final = _token_request({"Authorization": f"Bearer {_oauth_identity_jwt(signing_key)}"}) - assert await _extract_user_id_from_request(request) == (None if allowed_domain else "jwt-owner") + assert await extract_user_id_from_request(request) == (None if allowed_domain else "jwt-owner") @pytest.mark.asyncio @@ -12274,7 +12274,7 @@ async def test_oauth_jwt_identity_preserves_separate_mcp_route_authorization( route_allowed: bool, ) -> None: from litellm.proxy import proxy_server - from litellm.proxy._experimental.mcp_server.bridge_token_flow import _extract_user_id_from_request + from litellm.proxy._experimental.mcp_server.bridge_token_flow import extract_user_id_from_request from litellm.proxy._types import LitellmUserRoles, RoleBasedPermissions, RoleMapping from litellm.proxy.auth.handle_jwt import JWTAuthManager @@ -12301,7 +12301,7 @@ async def test_oauth_jwt_identity_preserves_separate_mcp_route_authorization( ) bearer: Final = _oauth_identity_jwt(signing_key) request: Final = _token_request({"Authorization": f"Bearer {bearer}"}, path="/example/token") - assert await _extract_user_id_from_request(request) == "jwt-owner" + assert await extract_user_id_from_request(request) == "jwt-owner" admission: Final = JWTAuthManager.auth_builder( api_key=bearer, jwt_handler=handler, @@ -12335,7 +12335,7 @@ async def test_oauth_jwt_resolves_canonical_owner_without_cached_identity( ) -> None: from litellm.models.user import LiteLLM_UserTable from litellm.proxy import proxy_server - from litellm.proxy._experimental.mcp_server.bridge_token_flow import _extract_user_id_from_request + from litellm.proxy._experimental.mcp_server.bridge_token_flow import extract_user_id_from_request from litellm.proxy.auth.handle_jwt import JWTAuthManager handler, signing_key = jwt_oauth_identity @@ -12356,7 +12356,7 @@ async def test_oauth_jwt_resolves_canonical_owner_without_cached_identity( monkeypatch.setattr(proxy_server, "prisma_client", database) bearer: Final = _oauth_identity_jwt(signing_key, owner=external_id, scope="litellm_proxy_admin" if admin else "") request: Final = _token_request({"Authorization": f"Bearer {bearer}"}) - stored_owner: Final = await _extract_user_id_from_request(request) + stored_owner: Final = await extract_user_id_from_request(request) assert stored_owner == (None if inactive else external_id if admin else "canonical-oauth-owner") assert table.find_unique.await_count == 2 if identity == "email": @@ -12382,7 +12382,7 @@ async def test_oauth_jwt_identity_does_not_provision_or_synchronize_teams( ) -> None: from litellm.models.user import LiteLLM_UserTable from litellm.proxy import proxy_server - from litellm.proxy._experimental.mcp_server.bridge_token_flow import _extract_user_id_from_request + from litellm.proxy._experimental.mcp_server.bridge_token_flow import extract_user_id_from_request handler, signing_key = jwt_oauth_identity handler.litellm_jwtauth.enforce_team_based_model_access = True @@ -12394,7 +12394,7 @@ async def test_oauth_jwt_identity_does_not_provision_or_synchronize_teams( request: Final = _token_request( {"Authorization": f"Bearer {_oauth_identity_jwt(signing_key)}"}, path="/example/token" ) - assert await _extract_user_id_from_request(request) == "jwt-owner" + assert await extract_user_id_from_request(request) == "jwt-owner" assert owner.teams == ["existing-team"] proxy_server.prisma_client.db.litellm_teamtable.find_unique.assert_not_called() proxy_server.prisma_client.db.litellm_teamtable.upsert.assert_not_called() @@ -12410,7 +12410,7 @@ async def test_oauth_refresh_revalidates_the_same_active_user_rule( ) -> None: from litellm.models.user import LiteLLM_UserTable from litellm.proxy import proxy_server - from litellm.proxy._experimental.mcp_server.bridge_token_flow import _reload_active_user_by_id + from litellm.proxy._experimental.mcp_server.bridge_token_flow import reload_active_user_by_id handler, _ = jwt_oauth_identity user_id: Final = f"jwt-owner-{state}" @@ -12419,7 +12419,7 @@ async def test_oauth_refresh_revalidates_the_same_active_user_rule( if state == "missing_database": monkeypatch.setattr(proxy_server, "prisma_client", None) expected: Final = None if state == "active" else "no_active_key" if state == "inactive" else "unresolvable" - assert await _reload_active_user_by_id(user_id) == expected + assert await reload_active_user_by_id(user_id) == expected if state != "missing_database": cached: Final = handler.user_api_key_cache.get_cache(user_id, model_type=LiteLLM_UserTable) assert cached is not None diff --git a/tests/unit/proxy/_experimental/mcp_server/test_jwt_mcp_enforcement.py b/tests/unit/proxy/_experimental/mcp_server/test_jwt_mcp_enforcement.py index d8e4a342e52..eddfdc4f884 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_jwt_mcp_enforcement.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_jwt_mcp_enforcement.py @@ -308,7 +308,7 @@ async def test_e2e_jwt_team_mcp_permissions_enforced(monkeypatch): # Mock _get_mcp_servers_from_access_groups to return empty with patch.object( - MCPRequestHandler, "_get_mcp_servers_from_access_groups" + MCPRequestHandler, "get_mcp_servers_from_access_groups" ) as mock_access_groups: mock_access_groups.return_value = [] @@ -506,7 +506,7 @@ async def test_e2e_jwt_team_mcp_key_intersection(monkeypatch): ), patch.object( MCPRequestHandler, - "_get_mcp_servers_from_access_groups", + "get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ), diff --git a/tests/unit/proxy/_experimental/mcp_server/test_jwt_mcp_simple.py b/tests/unit/proxy/_experimental/mcp_server/test_jwt_mcp_simple.py index 052231b562a..7ee340a69f4 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_jwt_mcp_simple.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_jwt_mcp_simple.py @@ -64,7 +64,7 @@ async def test_simple_jwt_mcp_permissions_enforced(): ), patch.object( MCPRequestHandler, - "_get_mcp_servers_from_access_groups", + "get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ), @@ -143,7 +143,7 @@ async def test_simple_jwt_team_id_required_for_mcp_permissions(): ) as mock_get_team, patch.object( MCPRequestHandler, - "_get_mcp_servers_from_access_groups", + "get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ), diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_env_vars.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_env_vars.py index 6a08e8eeab6..6970ffb6b08 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_env_vars.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_env_vars.py @@ -631,7 +631,7 @@ async def test_health_check_reaches_servers_without_forwarding_per_user_env_vars manager: Final = MCPServerManager() manager.registry[mock_server.server_id] = mock_server create_client: Final = AsyncMock() - monkeypatch.setattr(manager, "_create_mcp_client", create_client) + monkeypatch.setattr(manager, "create_mcp_client", create_client) monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") route: Final = respx_mock.get(mock_server.url).respond(401) diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_guardrail_usage_monitor.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_guardrail_usage_monitor.py index 1578ea8e601..8d5e36e9795 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_guardrail_usage_monitor.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_guardrail_usage_monitor.py @@ -68,9 +68,9 @@ def _fake_proxy_logging(capture: dict, *, guardrail_effect=None): """ plo = mock.MagicMock() plo.enforce_mcp_server_rate_limits = mock.AsyncMock() - plo._create_mcp_request_object_from_kwargs.return_value = mock.MagicMock() + plo.create_mcp_request_object_from_kwargs.return_value = mock.MagicMock() # Mirror the real conversion's metadata bucket so a test can prove it survives. - plo._convert_mcp_to_llm_format.side_effect = lambda *_a, **_k: { + plo.convert_mcp_to_llm_format.side_effect = lambda *_a, **_k: { "metadata": {"headers": {"x-forwarded-for": "1.2.3.4"}} } diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py index 333984bde68..c477a8d34cb 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py @@ -99,10 +99,10 @@ class TestPreCallToolCheckReturnsHeaders: server = self._make_server() proxy_logging = MagicMock(spec=ProxyLogging) - proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(return_value=MagicMock()) - proxy_logging._convert_mcp_to_llm_format = MagicMock(return_value={"model": "fake"}) + proxy_logging.create_mcp_request_object_from_kwargs = MagicMock(return_value=MagicMock()) + proxy_logging.convert_mcp_to_llm_format = MagicMock(return_value={"model": "fake"}) proxy_logging.pre_call_hook = AsyncMock(return_value={"modified_arguments": {"key": "val"}}) - proxy_logging._convert_mcp_hook_response_to_kwargs = MagicMock(return_value={"arguments": {"key": "val"}}) + proxy_logging.convert_mcp_hook_response_to_kwargs = MagicMock(return_value={"arguments": {"key": "val"}}) with patch.object(manager, "check_allowed_or_banned_tools", return_value=True): with patch.object( @@ -130,10 +130,10 @@ class TestPreCallToolCheckReturnsHeaders: hook_headers = {"Authorization": "Bearer signed-jwt", "X-Trace-Id": "abc123"} proxy_logging = MagicMock(spec=ProxyLogging) - proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(return_value=MagicMock()) - proxy_logging._convert_mcp_to_llm_format = MagicMock(return_value={"model": "fake"}) + proxy_logging.create_mcp_request_object_from_kwargs = MagicMock(return_value=MagicMock()) + proxy_logging.convert_mcp_to_llm_format = MagicMock(return_value={"model": "fake"}) proxy_logging.pre_call_hook = AsyncMock(return_value={"extra_headers": hook_headers}) - proxy_logging._convert_mcp_hook_response_to_kwargs = MagicMock( + proxy_logging.convert_mcp_hook_response_to_kwargs = MagicMock( return_value={"arguments": {"key": "val"}, "extra_headers": hook_headers} ) @@ -161,8 +161,8 @@ class TestPreCallToolCheckReturnsHeaders: server = self._make_server() proxy_logging = MagicMock(spec=ProxyLogging) - proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(return_value=MagicMock()) - proxy_logging._convert_mcp_to_llm_format = MagicMock(return_value={"model": "fake"}) + proxy_logging.create_mcp_request_object_from_kwargs = MagicMock(return_value=MagicMock()) + proxy_logging.convert_mcp_to_llm_format = MagicMock(return_value={"model": "fake"}) proxy_logging.pre_call_hook = AsyncMock(return_value=None) with patch.object(manager, "check_allowed_or_banned_tools", return_value=True): @@ -193,10 +193,10 @@ class TestPreCallToolCheckReturnsHeaders: modified_args = {"key": "modified", "extra": "added"} proxy_logging = MagicMock(spec=ProxyLogging) - proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(return_value=MagicMock()) - proxy_logging._convert_mcp_to_llm_format = MagicMock(return_value={"model": "fake"}) + proxy_logging.create_mcp_request_object_from_kwargs = MagicMock(return_value=MagicMock()) + proxy_logging.convert_mcp_to_llm_format = MagicMock(return_value={"model": "fake"}) proxy_logging.pre_call_hook = AsyncMock(return_value={"modified_arguments": modified_args}) - proxy_logging._convert_mcp_hook_response_to_kwargs = MagicMock(return_value={"arguments": modified_args}) + proxy_logging.convert_mcp_hook_response_to_kwargs = MagicMock(return_value={"arguments": modified_args}) with patch.object(manager, "check_allowed_or_banned_tools", return_value=True): with patch.object( @@ -226,10 +226,10 @@ class TestPreCallToolCheckReturnsHeaders: hook_headers = {"Authorization": "Bearer jwt"} proxy_logging = MagicMock(spec=ProxyLogging) - proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(return_value=MagicMock()) - proxy_logging._convert_mcp_to_llm_format = MagicMock(return_value={"model": "fake"}) + proxy_logging.create_mcp_request_object_from_kwargs = MagicMock(return_value=MagicMock()) + proxy_logging.convert_mcp_to_llm_format = MagicMock(return_value={"model": "fake"}) proxy_logging.pre_call_hook = AsyncMock(return_value={"dummy": True}) - proxy_logging._convert_mcp_hook_response_to_kwargs = MagicMock( + proxy_logging.convert_mcp_hook_response_to_kwargs = MagicMock( return_value={"arguments": modified_args, "extra_headers": hook_headers} ) @@ -276,7 +276,7 @@ class TestCallToolFlowsHookHeaders: with patch.object( manager, - "_get_mcp_server_from_tool_name", + "get_mcp_server_from_tool_name", return_value=server, ): with patch.object( @@ -317,7 +317,7 @@ class TestCallToolFlowsHookHeaders: with patch.object( manager, - "_get_mcp_server_from_tool_name", + "get_mcp_server_from_tool_name", return_value=server, ): with patch.object( @@ -347,7 +347,7 @@ class TestCallToolFlowsHookHeaders: with patch.object( manager, - "_get_mcp_server_from_tool_name", + "get_mcp_server_from_tool_name", return_value=server, ): with patch.object( @@ -394,7 +394,7 @@ class TestCallToolFlowsHookHeaders: spec_path="/path/to/spec.yaml", ) - with patch.object(manager, "_get_mcp_server_from_tool_name", return_value=server): + with patch.object(manager, "get_mcp_server_from_tool_name", return_value=server): with patch.object( manager, "pre_call_tool_check", @@ -441,7 +441,7 @@ class TestCallToolFlowsHookHeaders: spec_path="/path/to/spec.yaml", ) - with patch.object(manager, "_get_mcp_server_from_tool_name", return_value=server): + with patch.object(manager, "get_mcp_server_from_tool_name", return_value=server): with patch.object( manager, "pre_call_tool_check", @@ -504,8 +504,8 @@ class TestHookHeaderMergePriority: mock_client.call_tool = AsyncMock(return_value=MagicMock()) return mock_client - with patch.object(manager, "_create_mcp_client", side_effect=fake_create_mcp_client): - with patch.object(manager, "_build_stdio_env", return_value=None): + with patch.object(manager, "create_mcp_client", side_effect=fake_create_mcp_client): + with patch.object(manager, "build_stdio_env", return_value=None): try: await manager._call_regular_mcp_tool( mcp_server=server, @@ -540,8 +540,8 @@ class TestHookHeaderMergePriority: mock_client.call_tool = AsyncMock(return_value=MagicMock()) return mock_client - with patch.object(manager, "_create_mcp_client", side_effect=fake_create_mcp_client): - with patch.object(manager, "_build_stdio_env", return_value=None): + with patch.object(manager, "create_mcp_client", side_effect=fake_create_mcp_client): + with patch.object(manager, "build_stdio_env", return_value=None): try: await manager._call_regular_mcp_tool( mcp_server=server, @@ -584,8 +584,8 @@ class TestHookHeaderMergePriority: mock_client.call_tool = AsyncMock(return_value=MagicMock()) return mock_client - with patch.object(manager, "_create_mcp_client", side_effect=fake_create_mcp_client): - with patch.object(manager, "_build_stdio_env", return_value=None): + with patch.object(manager, "create_mcp_client", side_effect=fake_create_mcp_client): + with patch.object(manager, "build_stdio_env", return_value=None): try: await manager._call_regular_mcp_tool( mcp_server=server, @@ -631,8 +631,8 @@ class TestHookHeaderMergePriority: mock_client.call_tool = AsyncMock(return_value=MagicMock()) return mock_client - with patch.object(manager, "_create_mcp_client", side_effect=fake_create_mcp_client): - with patch.object(manager, "_build_stdio_env", return_value=None): + with patch.object(manager, "create_mcp_client", side_effect=fake_create_mcp_client): + with patch.object(manager, "build_stdio_env", return_value=None): try: await manager._call_regular_mcp_tool( mcp_server=server, @@ -667,8 +667,8 @@ class TestHookHeaderMergePriority: mock_client.call_tool = AsyncMock(return_value=MagicMock()) return mock_client - with patch.object(manager, "_create_mcp_client", side_effect=fake_create_mcp_client): - with patch.object(manager, "_build_stdio_env", return_value=None): + with patch.object(manager, "create_mcp_client", side_effect=fake_create_mcp_client): + with patch.object(manager, "build_stdio_env", return_value=None): try: await manager._call_regular_mcp_tool( mcp_server=server, @@ -708,8 +708,8 @@ class TestHookHeaderMergePriority: mock_client.call_tool = AsyncMock(return_value=MagicMock()) return mock_client - with patch.object(manager, "_create_mcp_client", side_effect=fake_create_mcp_client): - with patch.object(manager, "_build_stdio_env", return_value=None): + with patch.object(manager, "create_mcp_client", side_effect=fake_create_mcp_client): + with patch.object(manager, "build_stdio_env", return_value=None): try: await manager._call_regular_mcp_tool( mcp_server=server, @@ -752,8 +752,8 @@ class TestHookHeaderMergePriority: mock_client.call_tool = AsyncMock(return_value=MagicMock()) return mock_client - with patch.object(manager, "_create_mcp_client", side_effect=fake_create_mcp_client): - with patch.object(manager, "_build_stdio_env", return_value=None): + with patch.object(manager, "create_mcp_client", side_effect=fake_create_mcp_client): + with patch.object(manager, "build_stdio_env", return_value=None): try: await manager._call_regular_mcp_tool( mcp_server=server, @@ -790,8 +790,8 @@ class TestHookHeaderMergePriority: mock_client.call_tool = AsyncMock(return_value=MagicMock()) return mock_client - with patch.object(manager, "_create_mcp_client", side_effect=fake_create_mcp_client): - with patch.object(manager, "_build_stdio_env", return_value=None): + with patch.object(manager, "create_mcp_client", side_effect=fake_create_mcp_client): + with patch.object(manager, "build_stdio_env", return_value=None): try: await manager._call_regular_mcp_tool( mcp_server=server, @@ -828,8 +828,8 @@ class TestHookHeaderMergePriority: mock_client.call_tool = AsyncMock(return_value=MagicMock()) return mock_client - with patch.object(manager, "_create_mcp_client", side_effect=fake_create_mcp_client): - with patch.object(manager, "_build_stdio_env", return_value=None): + with patch.object(manager, "create_mcp_client", side_effect=fake_create_mcp_client): + with patch.object(manager, "build_stdio_env", return_value=None): try: await manager._call_regular_mcp_tool( mcp_server=server, @@ -875,8 +875,8 @@ class TestHookHeaderMergePriority: mock_client.call_tool = AsyncMock(return_value=MagicMock()) return mock_client - with patch.object(manager, "_create_mcp_client", side_effect=fake_create_mcp_client): - with patch.object(manager, "_build_stdio_env", return_value=None): + with patch.object(manager, "create_mcp_client", side_effect=fake_create_mcp_client): + with patch.object(manager, "build_stdio_env", return_value=None): try: await manager._call_regular_mcp_tool( mcp_server=server, @@ -919,8 +919,8 @@ class TestHookHeaderMergePriority: mock_client.call_tool = AsyncMock(return_value=MagicMock()) return mock_client - with patch.object(manager, "_create_mcp_client", side_effect=fake_create_mcp_client): - with patch.object(manager, "_build_stdio_env", return_value=None): + with patch.object(manager, "create_mcp_client", side_effect=fake_create_mcp_client): + with patch.object(manager, "build_stdio_env", return_value=None): try: await manager._call_regular_mcp_tool( mcp_server=server, @@ -1031,10 +1031,10 @@ class TestMcpRateLimitServerNameSurfacing: return {"model": "fake"} proxy_logging = MagicMock(spec=ProxyLogging) - proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(return_value=MagicMock()) - proxy_logging._convert_mcp_to_llm_format = MagicMock(side_effect=capture_convert) + proxy_logging.create_mcp_request_object_from_kwargs = MagicMock(return_value=MagicMock()) + proxy_logging.convert_mcp_to_llm_format = MagicMock(side_effect=capture_convert) proxy_logging.pre_call_hook = AsyncMock(return_value=None) - proxy_logging._convert_mcp_hook_response_to_kwargs = MagicMock(return_value={"arguments": {}}) + proxy_logging.convert_mcp_hook_response_to_kwargs = MagicMock(return_value={"arguments": {}}) with patch.object(manager, "check_allowed_or_banned_tools", return_value=True): with patch.object( @@ -1060,7 +1060,7 @@ class TestOpenApiByokCallTool: async def test_call_tool_openapi_byok_injects_request_auth_contextvar(self): """Playground/responses call call_tool directly; BYOK must reach OpenAPI handlers.""" from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( - _request_auth_header, + request_auth_header, ) manager = MCPServerManager() @@ -1078,7 +1078,7 @@ class TestOpenApiByokCallTool: captured_auth: dict[str, Optional[str]] = {} async def fake_openapi_handler(_server, _name, _arguments, _wire_compat): - captured_auth["value"] = _request_auth_header.get() + captured_auth["value"] = request_auth_header.get() return CallToolResult(content=[], isError=False) with patch.object(manager, "_resolve_mcp_server_for_tool_call", return_value=server): @@ -1305,7 +1305,7 @@ class TestOpenApiResolvedUpstreamAuth: """The managed spec_path arm resolves the v2 credential and sets the ContextVar; kills the mutant that drops the resolve_openapi_upstream_auth call in call_tool.""" from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( - _request_resolved_auth_headers, + request_resolved_auth_headers, ) from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import ( StaticHeaderAuth, @@ -1318,7 +1318,7 @@ class TestOpenApiResolvedUpstreamAuth: captured: Dict[str, Any] = {} async def fake_openapi_handler(_server, _name, _arguments, _wire_compat): - captured["resolved"] = _request_resolved_auth_headers.get() + captured["resolved"] = request_resolved_auth_headers.get() return CallToolResult(content=[], isError=False) with patch.object(manager, "_resolve_mcp_server_for_tool_call", return_value=server): @@ -1336,7 +1336,7 @@ class TestOpenApiResolvedUpstreamAuth: ) assert captured["resolved"] == {"Authorization": "Bearer stored-user-token"} - assert _request_resolved_auth_headers.get() is None + assert request_resolved_auth_headers.get() is None @pytest.mark.asyncio async def test_call_tool_openapi_m2m_missing_token_url_fails_closed(self): @@ -1467,8 +1467,8 @@ class TestPreCallToolCheckExposesClientHeaders: return {"model": "fake"} proxy_logging = MagicMock(spec=ProxyLogging) - proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(return_value=MagicMock()) - proxy_logging._convert_mcp_to_llm_format = MagicMock(side_effect=capture) + proxy_logging.create_mcp_request_object_from_kwargs = MagicMock(return_value=MagicMock()) + proxy_logging.convert_mcp_to_llm_format = MagicMock(side_effect=capture) proxy_logging.pre_call_hook = AsyncMock(return_value=None) with patch.object(manager, "check_allowed_or_banned_tools", return_value=True): diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_max_concurrent_requests.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_max_concurrent_requests.py index 0dc7ac5ecd9..d046871c8e2 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_max_concurrent_requests.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_max_concurrent_requests.py @@ -57,7 +57,7 @@ def _patch_client_with_tracker(manager: MCPServerManager, tracker: _ConcurrencyT return _ProbeClient() - return patch.object(manager, "_create_mcp_client", side_effect=fake_create_mcp_client) + return patch.object(manager, "create_mcp_client", side_effect=fake_create_mcp_client) async def _fire(manager: MCPServerManager, server: MCPServer, n: int) -> None: diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_tools.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_tools.py index 1909e3306a2..7290f0de86d 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_tools.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_tools.py @@ -347,7 +347,7 @@ async def test_aggregate_list_tools_absorbs_one_unauthenticated_server(): ), patch.object( mcp_operations, "filter_tools_by_key_team_permissions", AsyncMock(side_effect=lambda tools, **k: tools) ), patch.object( - mcp_operations.global_mcp_server_manager, "_get_tools_from_server", AsyncMock(side_effect=fake_get_tools) + mcp_operations.global_mcp_server_manager, "get_tools_from_server", AsyncMock(side_effect=fake_get_tools) ): listing = await mcp_operations._get_tools_from_mcp_servers( user_api_key_auth=UserAPIKeyAuth(token="h", user_id="u1"), @@ -370,7 +370,7 @@ async def test_single_server_route_also_absorbs_upstream_auth_error(): from unittest.mock import patch from litellm.proxy._experimental.mcp_server import server as mcp_server - from litellm.proxy._experimental.mcp_server.mcp_context import _mcp_gateway_server_name + from litellm.proxy._experimental.mcp_server.mcp_context import mcp_gateway_server_name from litellm.proxy._types import UserAPIKeyAuth delegate = _http_server( @@ -381,14 +381,14 @@ async def test_single_server_route_also_absorbs_upstream_auth_error(): raise MCPUpstreamAuthError(status_code=401, www_authenticate=None, server_name=server.name) # //mcp sets the path-derived single-server scope; absorption must hold even then. - token = _mcp_gateway_server_name.set("delegate_docs") + token = mcp_gateway_server_name.set("delegate_docs") try: with patch.object(mcp_operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[delegate])), patch.object( mcp_operations, "_prefetch_oauth_creds_for_user", AsyncMock(return_value={}) ), patch.object(mcp_operations, "_prepare_mcp_server_headers", MagicMock(return_value=(None, None))), patch.object( mcp_operations, "_get_user_oauth_extra_headers_from_db", AsyncMock(return_value=None) ), patch.object( - mcp_operations.global_mcp_server_manager, "_get_tools_from_server", AsyncMock(side_effect=fake_get_tools) + mcp_operations.global_mcp_server_manager, "get_tools_from_server", AsyncMock(side_effect=fake_get_tools) ): listing = await mcp_operations._get_tools_from_mcp_servers( user_api_key_auth=UserAPIKeyAuth(token="h", user_id="u1"), @@ -398,7 +398,7 @@ async def test_single_server_route_also_absorbs_upstream_auth_error(): assert listing.tools == [] assert listing.outcomes["delegate_docs"].tag == "auth_required" finally: - _mcp_gateway_server_name.reset(token) + mcp_gateway_server_name.reset(token) @pytest.mark.asyncio @@ -425,7 +425,7 @@ async def test_aggregate_with_single_accessible_server_still_absorbs(): ), patch.object(mcp_operations, "_prepare_mcp_server_headers", MagicMock(return_value=(None, None))), patch.object( mcp_operations, "_get_user_oauth_extra_headers_from_db", AsyncMock(return_value=None) ), patch.object( - mcp_operations.global_mcp_server_manager, "_get_tools_from_server", AsyncMock(side_effect=fake_get_tools) + mcp_operations.global_mcp_server_manager, "get_tools_from_server", AsyncMock(side_effect=fake_get_tools) ): # Aggregate route: no explicit server filter, even though only one server is accessible. listing = await mcp_operations._get_tools_from_mcp_servers( @@ -470,7 +470,7 @@ async def test_client_creation_failure_logs_sanitized_exchange(monkeypatch, capl request = httpx.Request("POST", "https://upstream/mcp?credential=query-secret") response = httpx.Response(500, request=request, json={"error":"missing_scope"}) error = httpx.HTTPStatusError("query-secret", request=request, response=response) - monkeypatch.setattr(manager, "_create_mcp_client", AsyncMock(side_effect=error)) + monkeypatch.setattr(manager, "create_mcp_client", AsyncMock(side_effect=error)) with caplog.at_level(logging.WARNING, logger="LiteLLM"): with pytest.raises(MCPServerListError): await manager._get_tools_from_server(server) diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_proxy_mode.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_proxy_mode.py index e52a86d76af..6f1bbad5040 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_proxy_mode.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_proxy_mode.py @@ -164,7 +164,7 @@ async def test_proxy_call_tool_on_a_never_listed_tool_hands_the_pre_hook_no_list patch.dict(manager.registry, {server.server_id: server}), patch.dict(manager.tool_name_to_mcp_server_name_mapping), patch.object(mcp_operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[server])), - patch.object(manager, "_create_mcp_client", AsyncMock(return_value=object())), + patch.object(manager, "create_mcp_client", AsyncMock(return_value=object())), patch.object(manager, "_fetch_tools_with_timeout", AsyncMock(return_value=upstream)), patch.object(manager, "pre_call_tool_check", pre_call_tool_check), patch.object(manager, "_call_regular_mcp_tool", call_regular_mcp_tool), diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py index 941d67cb587..23defa6f513 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py @@ -925,7 +925,7 @@ async def test_get_tools_from_mcp_servers(): mock_manager.get_mcp_server_by_id = lambda server_id: ( mock_server_1 if server_id == "server1_id" else mock_server_2 ) - mock_manager._get_tools_from_server = AsyncMock(return_value=[mock_tool_1]) + mock_manager.get_tools_from_server = AsyncMock(return_value=[mock_tool_1]) # Mock filter_server_ids_by_ip_with_info to return input unchanged (no IP filtering in test) mock_manager.filter_server_ids_by_ip_with_info = MagicMock( side_effect=lambda server_ids, client_ip: (server_ids, 0) @@ -972,7 +972,7 @@ async def test_get_tools_from_mcp_servers(): return [mock_tool_1] return [mock_tool_2] - mock_manager_2._get_tools_from_server = AsyncMock( + mock_manager_2.get_tools_from_server = AsyncMock( side_effect=mock_get_tools_side_effect ) # Mock filter_server_ids_by_ip_with_info to return input unchanged (no IP filtering in test) @@ -1007,7 +1007,7 @@ async def test_get_tools_from_mcp_servers(): if server_id == "server1_id" else (mock_server_2 if server_id == "server2_id" else mock_server_3) ) - mock_manager._get_tools_from_server = AsyncMock(return_value=[mock_tool_1]) + mock_manager.get_tools_from_server = AsyncMock(return_value=[mock_tool_1]) # Mock filter_server_ids_by_ip_with_info to return input unchanged (no IP filtering in test) mock_manager.filter_server_ids_by_ip_with_info = MagicMock( side_effect=lambda server_ids, client_ip: (server_ids, 0) @@ -1018,7 +1018,7 @@ async def test_get_tools_from_mcp_servers(): mock_manager, ): with patch( - "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups", + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_mcp_servers_from_access_groups", AsyncMock(return_value=["server3_id"]), ): # Test with specific servers @@ -1399,7 +1399,7 @@ async def test_mcp_server_manager_alias_tool_prefixing(): mock_client_constructor, ): # Get tools from server - tools = await test_manager._get_tools_from_server(mock_server) + tools = await test_manager.get_tools_from_server(mock_server) # Verify tool is prefixed with alias assert len(tools) == 1 @@ -1459,7 +1459,7 @@ async def test_mcp_server_manager_server_name_tool_prefixing(): mock_client_constructor, ): # Get tools from server - tools = await test_manager._get_tools_from_server(mock_server) + tools = await test_manager.get_tools_from_server(mock_server) # Verify tool is prefixed with server_name (normalized) assert len(tools) == 1 @@ -1519,7 +1519,7 @@ async def test_mcp_server_manager_server_id_tool_prefixing(): mock_client_constructor, ): # Get tools from server - tools = await test_manager._get_tools_from_server(mock_server) + tools = await test_manager.get_tools_from_server(mock_server) # Verify tool is prefixed with server_id assert len(tools) == 1 @@ -1989,12 +1989,12 @@ async def test_get_tools_for_single_server(): with patch( "litellm.proxy._experimental.mcp_server.rest_endpoints.global_mcp_server_manager" ) as mock_manager: - mock_manager._get_tools_from_server = AsyncMock(return_value=mock_tools) + mock_manager.get_tools_from_server = AsyncMock(return_value=mock_tools) result = await _get_tools_for_single_server(mock_server, "Bearer test_token") # Verify the manager was called with correct parameters - mock_manager._get_tools_from_server.assert_called_once_with( + mock_manager.get_tools_from_server.assert_called_once_with( server=mock_server, mcp_auth_header="Bearer test_token", extra_headers=None, @@ -2050,7 +2050,7 @@ async def test_get_tools_for_single_server_applies_disallowed_tools_without_allo with patch( "litellm.proxy._experimental.mcp_server.rest_endpoints.global_mcp_server_manager" ) as mock_manager: - mock_manager._get_tools_from_server = AsyncMock(return_value=mock_tools) + mock_manager.get_tools_from_server = AsyncMock(return_value=mock_tools) result = await _get_tools_for_single_server(mock_server, "Bearer test_token") @@ -2104,7 +2104,7 @@ async def test_rest_listing_hides_key_grants_dispatch_would_refuse(): "get_allowed_tools_for_server", AsyncMock(return_value=[f"{server_id}-read_wiki_contents"]), ): - mock_manager._get_tools_from_server = AsyncMock(return_value=mock_tools) + mock_manager.get_tools_from_server = AsyncMock(return_value=mock_tools) mock_server_manager.get_mcp_server_by_id.return_value = mock_server result = await _get_tools_for_single_server( @@ -2135,10 +2135,10 @@ async def test_list_tool_rest_api_with_server_specific_auth(): # Mock the MCPRequestHandler methods with patch.object( - MCPRequestHandler, "_get_mcp_auth_header_from_headers" + MCPRequestHandler, "get_mcp_auth_header_from_headers" ) as mock_get_auth: with patch.object( - MCPRequestHandler, "_get_mcp_server_auth_headers_from_headers" + MCPRequestHandler, "get_mcp_server_auth_headers_from_headers" ) as mock_get_server_auth: mock_get_auth.return_value = "Bearer default_token" mock_get_server_auth.return_value = { @@ -2232,10 +2232,10 @@ async def test_list_tool_rest_api_with_default_auth(): # Mock the MCPRequestHandler methods with patch.object( - MCPRequestHandler, "_get_mcp_auth_header_from_headers" + MCPRequestHandler, "get_mcp_auth_header_from_headers" ) as mock_get_auth: with patch.object( - MCPRequestHandler, "_get_mcp_server_auth_headers_from_headers" + MCPRequestHandler, "get_mcp_server_auth_headers_from_headers" ) as mock_get_server_auth: mock_get_auth.return_value = "Bearer default_token" mock_get_server_auth.return_value = {} # No server-specific headers @@ -2327,10 +2327,10 @@ async def test_list_tool_rest_api_all_servers_with_auth(): # Mock the MCPRequestHandler methods with patch.object( - MCPRequestHandler, "_get_mcp_auth_header_from_headers" + MCPRequestHandler, "get_mcp_auth_header_from_headers" ) as mock_get_auth: with patch.object( - MCPRequestHandler, "_get_mcp_server_auth_headers_from_headers" + MCPRequestHandler, "get_mcp_server_auth_headers_from_headers" ) as mock_get_server_auth: mock_get_auth.return_value = "Bearer default_token" mock_get_server_auth.return_value = { @@ -2508,7 +2508,7 @@ async def test_filter_tools_by_allowed_tools_integration(): ) # Mock the _get_tools_from_server method to return all tools - mock_manager._get_tools_from_server = AsyncMock(return_value=mock_tools) + mock_manager.get_tools_from_server = AsyncMock(return_value=mock_tools) # Mock the MCPClient constructor with patch( @@ -2547,7 +2547,7 @@ async def test_filter_tools_by_allowed_tools_integration(): # Note: get_mcp_server_by_id is now called for each server ID instead of batch # Verify it was called with the correct server ID assert mock_manager.get_mcp_server_by_id.call_count > 0 - mock_manager._get_tools_from_server.assert_called_once() + mock_manager.get_tools_from_server.assert_called_once() @pytest.mark.asyncio @@ -2622,7 +2622,7 @@ async def test_filter_tools_by_disallowed_tools_integration(): side_effect=lambda server_ids, client_ip: (server_ids, 0) ) # Mock the _get_tools_from_server method to return all tools - mock_manager._get_tools_from_server = AsyncMock(return_value=mock_tools) + mock_manager.get_tools_from_server = AsyncMock(return_value=mock_tools) # Mock the MCPClient constructor with patch( @@ -2661,7 +2661,7 @@ async def test_filter_tools_by_disallowed_tools_integration(): # Note: get_mcp_server_by_id is now called for each server ID instead of batch # Verify it was called with the correct server ID assert mock_manager.get_mcp_server_by_id.call_count > 0 - mock_manager._get_tools_from_server.assert_called_once() + mock_manager.get_tools_from_server.assert_called_once() @pytest.mark.asyncio @@ -2724,7 +2724,7 @@ async def test_filter_tools_no_restrictions_integration(): ) # Mock the _get_tools_from_server method to return all tools - mock_manager._get_tools_from_server = AsyncMock(return_value=mock_tools) + mock_manager.get_tools_from_server = AsyncMock(return_value=mock_tools) # Mock the MCPClient constructor with patch( @@ -2988,7 +2988,7 @@ async def test_call_mcp_tool_uses_manager_permission_lookup(): ), patch.object( global_mcp_server_manager, - "_get_mcp_server_from_tool_name", + "get_mcp_server_from_tool_name", return_value=mock_server, ) as mock_get_server, patch( @@ -3064,7 +3064,7 @@ async def test_call_mcp_tool_resolves_unprefixed_tool_name_and_checks_permission ), patch.object( global_mcp_server_manager, - "_get_mcp_server_from_tool_name", + "get_mcp_server_from_tool_name", return_value=mock_server, ) as mock_get_server, patch( diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 9aaba0e9356..5bd4ff5d894 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -54,8 +54,8 @@ from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( _deserialize_json_list, _normalize_mcp_server_cost_info, _obo_retry_applies, - _resolve_openapi_tool_auth, - _should_strip_caller_authorization, + resolve_openapi_tool_auth, + should_strip_caller_authorization, listed_tools_caller_for, ) from litellm.proxy._types import ( @@ -1366,7 +1366,7 @@ class TestMCPServerManager: manager.registry[server.server_id] = server manager._set_oauth_discovery_deferred(server.server_id, True) - with patch.object(manager, "_get_tools_from_server", new=AsyncMock()) as get_tools: + with patch.object(manager, "get_tools_from_server", new=AsyncMock()) as get_tools: await manager._initialize_tool_name_to_mcp_server_name_mapping() get_tools.assert_not_awaited() @@ -1883,7 +1883,7 @@ class TestMCPServerManager: tool1.name = "zapier_tool_1" return [tool1] - manager._get_tools_from_server = mock_get_tools_from_server + manager.get_tools_from_server = mock_get_tools_from_server # Test with server-specific auth headers mcp_server_auth_headers = { @@ -1928,7 +1928,7 @@ class TestMCPServerManager: tool.name = "github_tool_1" return [tool] - manager._get_tools_from_server = mock_get_tools_from_server + manager.get_tools_from_server = mock_get_tools_from_server # Test with only legacy auth header (no server-specific headers) result = await manager.list_tools( @@ -1965,7 +1965,7 @@ class TestMCPServerManager: tool.name = "github_tool_1" return [tool] - manager._get_tools_from_server = mock_get_tools_from_server + manager.get_tools_from_server = mock_get_tools_from_server # Test with both legacy and server-specific headers result = await manager.list_tools( @@ -2001,7 +2001,7 @@ class TestMCPServerManager: captured_extra_headers = extra_headers return mock_client - manager._create_mcp_client = AsyncMock(side_effect=capture_create_mcp_client) + manager.create_mcp_client = AsyncMock(side_effect=capture_create_mcp_client) result = await manager._call_regular_mcp_tool( mcp_server=server, @@ -2029,7 +2029,7 @@ class TestMCPServerManager: captured["subject_token"] = subject_token return AsyncMock() - manager._create_mcp_client = AsyncMock(side_effect=capture_create_mcp_client) + manager.create_mcp_client = AsyncMock(side_effect=capture_create_mcp_client) manager._fetch_tools_with_timeout = AsyncMock(return_value=[]) await manager._get_tools_from_server(server=server, oauth2_headers=oauth2_headers, raw_headers=raw_headers) return captured["subject_token"] @@ -2102,7 +2102,7 @@ class TestMCPServerManager: challenge = ( 'Bearer resource_metadata="/.well-known/oauth-protected-resource/mcp/te-401-server", error="invalid_token"' ) - manager._create_mcp_client = AsyncMock( + manager.create_mcp_client = AsyncMock( side_effect=HTTPException(status_code=401, detail="Unauthorized", headers={"WWW-Authenticate": challenge}) ) with pytest.raises(MCPUpstreamAuthError) as exc_info: @@ -2128,7 +2128,7 @@ class TestMCPServerManager: client_secret="csec", ) manager = MCPServerManager() - manager._create_mcp_client = AsyncMock( + manager.create_mcp_client = AsyncMock( side_effect=HTTPException(status_code=412, detail="token exchange endpoint is not configured") ) with pytest.raises(MCPServerListError) as exc_info: @@ -2178,7 +2178,7 @@ class TestMCPServerManager: manager = MCPServerManager() mock_client = AsyncMock() mock_client.call_tool = AsyncMock(side_effect=self._upstream_status_error(401, challenge)) - manager._create_mcp_client = AsyncMock(return_value=mock_client) + manager.create_mcp_client = AsyncMock(return_value=mock_client) with pytest.raises(MCPUpstreamAuthError) as exc_info: await self._run_call_regular(manager, server) @@ -2198,7 +2198,7 @@ class TestMCPServerManager: expected = CallToolResult(content=[], isError=is_error) mock_client = AsyncMock() mock_client.call_tool = AsyncMock(return_value=expected) - manager._create_mcp_client = AsyncMock(return_value=mock_client) + manager.create_mcp_client = AsyncMock(return_value=mock_client) result = await self._run_call_regular(manager, server) @@ -2219,7 +2219,7 @@ class TestMCPServerManager: mock_client = AsyncMock() mock_client.call_tool = AsyncMock(side_effect=self._upstream_status_error(status_code)) mock_client.error_tool_result = MCPClient.error_tool_result - manager._create_mcp_client = AsyncMock(return_value=mock_client) + manager.create_mcp_client = AsyncMock(return_value=mock_client) import litellm.proxy._experimental.mcp_server.mcp_server_manager as _mgr_mod @@ -2246,7 +2246,7 @@ class TestMCPServerManager: manager = MCPServerManager() mock_client = AsyncMock() mock_client.call_tool = AsyncMock(return_value=CallToolResult(content=[], isError=False)) - manager._create_mcp_client = AsyncMock(return_value=mock_client) + manager.create_mcp_client = AsyncMock(return_value=mock_client) result = await manager._call_regular_mcp_tool( mcp_server=server, @@ -2818,7 +2818,7 @@ class TestMCPServerManager: captured["subject_token"] = subject_token return AsyncMock() - manager._create_mcp_client = AsyncMock(side_effect=capture_create_mcp_client) + manager.create_mcp_client = AsyncMock(side_effect=capture_create_mcp_client) await call(manager) return captured.get("subject_token") @@ -3431,7 +3431,7 @@ class TestMCPServerManager: captured_extra_headers = extra_headers return mock_client - manager._create_mcp_client = AsyncMock(side_effect=capture_create_mcp_client) + manager.create_mcp_client = AsyncMock(side_effect=capture_create_mcp_client) await manager._call_regular_mcp_tool( mcp_server=server, @@ -3472,7 +3472,7 @@ class TestMCPServerManager: ) # Migrated authorization_code => the centralized strip decision says drop the # caller's Authorization (the v2 resolver injects the stored token). - assert _should_strip_caller_authorization(mcp_server=server, raw_headers=None, user_api_key_auth=None) is True + assert should_strip_caller_authorization(mcp_server=server, raw_headers=None, user_api_key_auth=None) is True mock_client = AsyncMock() mock_client.call_tool = AsyncMock(return_value=CallToolResult(content=[], isError=False)) @@ -3490,7 +3490,7 @@ class TestMCPServerManager: captured_extra_headers = extra_headers return mock_client - manager._create_mcp_client = AsyncMock(side_effect=capture_create_mcp_client) + manager.create_mcp_client = AsyncMock(side_effect=capture_create_mcp_client) await manager._call_regular_mcp_tool( mcp_server=server, @@ -3558,7 +3558,7 @@ class TestMCPServerManager: captured_extra_headers = extra_headers return mock_client - manager._create_mcp_client = AsyncMock(side_effect=capture_create_mcp_client) + manager.create_mcp_client = AsyncMock(side_effect=capture_create_mcp_client) await manager._call_regular_mcp_tool( mcp_server=server, @@ -3615,7 +3615,7 @@ class TestMCPServerManager: captured_extra_headers = extra_headers return mock_client - manager._create_mcp_client = AsyncMock(side_effect=capture_create_mcp_client) + manager.create_mcp_client = AsyncMock(side_effect=capture_create_mcp_client) await manager._call_regular_mcp_tool( mcp_server=server, @@ -3644,7 +3644,7 @@ class TestMCPServerManager: captured["extra_headers"] = extra_headers return mock_client - manager._create_mcp_client = AsyncMock(side_effect=capture_create_mcp_client) + manager.create_mcp_client = AsyncMock(side_effect=capture_create_mcp_client) await manager._call_regular_mcp_tool( mcp_server=server, original_tool_name="tool", @@ -3796,7 +3796,7 @@ class TestMCPServerManager: auth_type=MCPAuth.true_passthrough, ) assert ( - _should_strip_caller_authorization( + should_strip_caller_authorization( mcp_server=true_passthrough, raw_headers={"authorization": "Bearer upstream"}, user_api_key_auth=UserAPIKeyAuth(api_key=None), @@ -3812,7 +3812,7 @@ class TestMCPServerManager: auth_type=MCPAuth.oauth_delegate, ) assert ( - _should_strip_caller_authorization( + should_strip_caller_authorization( mcp_server=oauth_delegate, raw_headers={ "x-litellm-api-key": "Bearer sk-litellm-key", @@ -3823,7 +3823,7 @@ class TestMCPServerManager: is False ) assert ( - _should_strip_caller_authorization( + should_strip_caller_authorization( mcp_server=oauth_delegate, raw_headers={"authorization": "Bearer sk-litellm-key"}, user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-key"), @@ -3845,7 +3845,7 @@ class TestMCPServerManager: auth_type=MCPAuth.oauth_delegate, ) assert ( - _should_strip_caller_authorization( + should_strip_caller_authorization( mcp_server=oauth_delegate, raw_headers={"authorization": "Bearer eyJ-idp-jwt"}, user_api_key_auth=UserAPIKeyAuth(user_id="alice", api_key=None), @@ -3853,7 +3853,7 @@ class TestMCPServerManager: is True ) assert ( - _should_strip_caller_authorization( + should_strip_caller_authorization( mcp_server=oauth_delegate, raw_headers={ "x-litellm-api-key": "Bearer sk-9876", @@ -4055,7 +4055,7 @@ class TestMCPServerManager: def test_caller_authorization_fans_out_only_with_second_consumer(self): from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( - _caller_authorization_fans_out, + caller_authorization_fans_out, ) delegate = MCPServer( @@ -4081,10 +4081,10 @@ class TestMCPServerManager: authentication_token="x", ) - assert _caller_authorization_fans_out(delegate, None) is False - assert _caller_authorization_fans_out(delegate, [delegate]) is False - assert _caller_authorization_fans_out(delegate, [delegate, static_server]) is False - assert _caller_authorization_fans_out(delegate, [delegate, second]) is True + assert caller_authorization_fans_out(delegate, None) is False + assert caller_authorization_fans_out(delegate, [delegate]) is False + assert caller_authorization_fans_out(delegate, [delegate, static_server]) is False + assert caller_authorization_fans_out(delegate, [delegate, second]) is True @pytest.mark.asyncio async def test_get_prompts_from_server_success(self): @@ -4107,7 +4107,7 @@ class TestMCPServerManager: with patch.object( manager, - "_create_mcp_client", + "create_mcp_client", new_callable=AsyncMock, return_value=mock_client, ): @@ -4140,7 +4140,7 @@ class TestMCPServerManager: with patch.object( manager, - "_create_mcp_client", + "create_mcp_client", new_callable=AsyncMock, return_value=mock_client, ): @@ -4181,7 +4181,7 @@ class TestMCPServerManager: with ( patch.object( manager, - "_create_mcp_client", + "create_mcp_client", new_callable=AsyncMock, return_value=mock_client, ) as mock_create_client, @@ -4236,7 +4236,7 @@ class TestMCPServerManager: with ( patch.object( manager, - "_create_mcp_client", + "create_mcp_client", new_callable=AsyncMock, return_value=mock_client, ) as mock_create_client, @@ -4291,7 +4291,7 @@ class TestMCPServerManager: with patch.object( manager, - "_create_mcp_client", + "create_mcp_client", new_callable=AsyncMock, return_value=mock_client, ) as mock_create_client: @@ -4854,7 +4854,7 @@ class TestMCPServerManager: tool.name = "github_tool_1" return [tool] - manager._get_tools_from_server = mock_get_tools_from_server + manager.get_tools_from_server = mock_get_tools_from_server # Test with server-specific headers that match server_name (even without alias) result = await manager.list_tools( @@ -5011,7 +5011,7 @@ class TestMCPServerManager: # Mock successful client.run_with_session mock_client = AsyncMock() mock_client.run_with_session = AsyncMock(return_value="ok") - manager._create_mcp_client = AsyncMock(return_value=mock_client) + manager.create_mcp_client = AsyncMock(return_value=mock_client) # Perform health check result = await manager.health_check_server("test-server") @@ -5043,7 +5043,7 @@ class TestMCPServerManager: # Mock failed client.run_with_session mock_client = AsyncMock() mock_client.run_with_session = AsyncMock(side_effect=Exception("Connection timeout")) - manager._create_mcp_client = AsyncMock(return_value=mock_client) + manager.create_mcp_client = AsyncMock(return_value=mock_client) # Perform health check result = await manager.health_check_server("test-server") @@ -5069,7 +5069,7 @@ class TestMCPServerManager: ) manager.get_mcp_server_by_id = MagicMock(return_value=server) manager._resolve_static_headers_with_env_vars = AsyncMock(return_value=None) - manager._create_mcp_client = AsyncMock( + manager.create_mcp_client = AsyncMock( side_effect=HTTPException(status_code=503, detail="OAuth discovery unavailable") ) @@ -5119,12 +5119,12 @@ class TestMCPServerManager: static_headers={"Authorization": "Bearer static-secret", "X-API-Key": "key-secret", "Cookie": "secret"}, ) manager.registry[server.server_id] = server - manager._create_mcp_client = AsyncMock() + manager.create_mcp_client = AsyncMock() route: Final = respx_mock.get(server.url).respond(401) result: Final = await manager.health_check_server(server.server_id, mcp_auth_header="caller-secret") - manager._create_mcp_client.assert_not_called() + manager.create_mcp_client.assert_not_called() assert result.status == "reachable" assert result.health_check_error is None assert result.last_health_check is not None @@ -5167,12 +5167,12 @@ class TestMCPServerManager: url="http://no-token-server.com", ) manager.registry[server.server_id] = server - manager._create_mcp_client = AsyncMock() + manager.create_mcp_client = AsyncMock() route: Final = respx_mock.get(server.url).respond(response_code) result: Final = await manager.health_check_server(server.server_id) - manager._create_mcp_client.assert_not_called() + manager.create_mcp_client.assert_not_called() assert route.call_count == 1 assert result.status == "reachable" assert result.health_check_error is None @@ -5435,7 +5435,7 @@ class TestMCPServerManager: captured_extra_headers = extra_headers return mock_client - manager._create_mcp_client = AsyncMock(side_effect=capture_create_mcp_client) + manager.create_mcp_client = AsyncMock(side_effect=capture_create_mcp_client) # Perform health check result = await manager.health_check_server("test-server") @@ -5465,12 +5465,12 @@ class TestMCPServerManager: extra_headers=["Authorization"], ) manager.registry[server.server_id] = server - manager._create_mcp_client = AsyncMock() + manager.create_mcp_client = AsyncMock() route: Final = respx_mock.get(server.url).respond(401) result: Final = await manager.health_check_server(server.server_id) - manager._create_mcp_client.assert_not_called() + manager.create_mcp_client.assert_not_called() assert route.call_count == 1 assert "authorization" not in route.calls[0].request.headers assert result.status == "reachable" @@ -5493,12 +5493,12 @@ class TestMCPServerManager: extra_headers=["x-api-key"], ) manager.registry[server.server_id] = server - manager._create_mcp_client = AsyncMock() + manager.create_mcp_client = AsyncMock() route: Final = respx_mock.get(server.url).respond(403) result: Final = await manager.health_check_server(server.server_id) - manager._create_mcp_client.assert_not_called() + manager.create_mcp_client.assert_not_called() assert route.call_count == 1 assert "x-api-key" not in route.calls[0].request.headers assert result.status == "reachable" @@ -5526,13 +5526,13 @@ class TestMCPServerManager: # Mock successful client mock_client = AsyncMock() mock_client.run_with_session = AsyncMock(return_value="ok") - manager._create_mcp_client = AsyncMock(return_value=mock_client) + manager.create_mcp_client = AsyncMock(return_value=mock_client) # Perform health check result = await manager.health_check_server("public-server") # Verify that client WAS created (health check should run) - manager._create_mcp_client.assert_called_once() + manager.create_mcp_client.assert_called_once() # Verify results assert isinstance(result, LiteLLM_MCPServerTable) @@ -5562,13 +5562,13 @@ class TestMCPServerManager: # Mock successful client mock_client = AsyncMock() mock_client.run_with_session = AsyncMock(return_value="ok") - manager._create_mcp_client = AsyncMock(return_value=mock_client) + manager.create_mcp_client = AsyncMock(return_value=mock_client) # Perform health check result = await manager.health_check_server("custom-server") # Verify that client WAS created (health check should run) - manager._create_mcp_client.assert_called_once() + manager.create_mcp_client.assert_called_once() # Verify results assert isinstance(result, LiteLLM_MCPServerTable) @@ -5858,8 +5858,8 @@ class TestMCPServerManager: proxy_logging_obj = _mock_proxy_logging() # Mock the async methods that pre_call_tool_check calls - proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(return_value={}) - proxy_logging_obj._convert_mcp_to_llm_format = MagicMock(return_value={}) + proxy_logging_obj.create_mcp_request_object_from_kwargs = MagicMock(return_value={}) + proxy_logging_obj.convert_mcp_to_llm_format = MagicMock(return_value={}) proxy_logging_obj.pre_call_hook = AsyncMock(return_value={}) # This should not raise an exception @@ -5925,8 +5925,8 @@ class TestMCPServerManager: proxy_logging_obj = _mock_proxy_logging() # Mock the async methods that pre_call_tool_check calls - proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(return_value={}) - proxy_logging_obj._convert_mcp_to_llm_format = MagicMock(return_value={}) + proxy_logging_obj.create_mcp_request_object_from_kwargs = MagicMock(return_value={}) + proxy_logging_obj.convert_mcp_to_llm_format = MagicMock(return_value={}) proxy_logging_obj.pre_call_hook = AsyncMock(return_value={}) # This should not raise an exception @@ -5992,8 +5992,8 @@ class TestMCPServerManager: proxy_logging_obj = _mock_proxy_logging() # Mock the async methods that pre_call_tool_check calls - proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(return_value={}) - proxy_logging_obj._convert_mcp_to_llm_format = MagicMock(return_value={}) + proxy_logging_obj.create_mcp_request_object_from_kwargs = MagicMock(return_value={}) + proxy_logging_obj.convert_mcp_to_llm_format = MagicMock(return_value={}) proxy_logging_obj.pre_call_hook = AsyncMock(return_value={}) # This should not raise an exception @@ -6027,8 +6027,8 @@ class TestMCPServerManager: proxy_logging_obj = _mock_proxy_logging() # Mock the async methods that pre_call_tool_check calls - proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(return_value={}) - proxy_logging_obj._convert_mcp_to_llm_format = MagicMock(return_value={}) + proxy_logging_obj.create_mcp_request_object_from_kwargs = MagicMock(return_value={}) + proxy_logging_obj.convert_mcp_to_llm_format = MagicMock(return_value={}) proxy_logging_obj.pre_call_hook = AsyncMock(return_value={}) # tool2 should be allowed since it's in allowed_tools (takes precedence) @@ -6067,7 +6067,7 @@ class TestMCPServerManager: ) # Mock client creation and fetching tools - manager._create_mcp_client = AsyncMock(return_value=object()) + manager.create_mcp_client = AsyncMock(return_value=object()) # Tools returned upstream (unprefixed from provider) upstream_tool = MCPTool( @@ -6105,7 +6105,7 @@ class TestMCPServerManager: transport=MCPTransport.http, ) - manager._create_mcp_client = AsyncMock(return_value=object()) + manager.create_mcp_client = AsyncMock(return_value=object()) manager._fetch_tools_with_timeout = AsyncMock(return_value=[]) user_auth = UserAPIKeyAuth(api_key="sk-test", user_id="alice") @@ -6739,7 +6739,7 @@ class TestMCPServerManager: with patch.object( rest_endpoints.global_mcp_server_manager, - "_get_tools_from_server", + "get_tools_from_server", new=AsyncMock(return_value=[tool1, tool2, tool3]), ): # Call the REST endpoint helper @@ -6780,7 +6780,7 @@ class TestMCPServerManager: with patch.object( rest_endpoints.global_mcp_server_manager, - "_get_tools_from_server", + "get_tools_from_server", new=AsyncMock(return_value=[tool1, tool2, tool3]), ): # Call the REST endpoint helper @@ -6819,7 +6819,7 @@ class TestMCPServerManager: with patch.object( rest_endpoints.global_mcp_server_manager, - "_get_tools_from_server", + "get_tools_from_server", new=AsyncMock(return_value=[tool1, tool2]), ): # Call the REST endpoint helper @@ -6893,8 +6893,8 @@ class TestMCPServerManager: ) proxy_logging = _mock_proxy_logging() - proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(return_value={}) - proxy_logging._convert_mcp_to_llm_format = MagicMock(return_value={}) + proxy_logging.create_mcp_request_object_from_kwargs = MagicMock(return_value={}) + proxy_logging.convert_mcp_to_llm_format = MagicMock(return_value={}) proxy_logging.pre_call_hook = AsyncMock(return_value=None) # Should succeed @@ -6936,8 +6936,8 @@ class TestMCPServerManager: ) proxy_logging = _mock_proxy_logging() - proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(return_value={}) - proxy_logging._convert_mcp_to_llm_format = MagicMock(return_value={}) + proxy_logging.create_mcp_request_object_from_kwargs = MagicMock(return_value={}) + proxy_logging.convert_mcp_to_llm_format = MagicMock(return_value={}) proxy_logging.pre_call_hook = AsyncMock(return_value=None) # Should fail with 403 @@ -7079,8 +7079,8 @@ class TestMCPServerManager: proxy_logging_obj = _mock_proxy_logging() # Mock the async methods that pre_call_tool_check calls - proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(return_value={}) - proxy_logging_obj._convert_mcp_to_llm_format = MagicMock(return_value={}) + proxy_logging_obj.create_mcp_request_object_from_kwargs = MagicMock(return_value={}) + proxy_logging_obj.convert_mcp_to_llm_format = MagicMock(return_value={}) proxy_logging_obj.pre_call_hook = AsyncMock(return_value={}) # Test 1: Call getpetbyid (unprefixed in allowed_tools) - should succeed @@ -7156,14 +7156,14 @@ class TestMCPServerManager: mock_client.call_tool.side_effect = mock_call_tool # Mock _create_mcp_client to return our mock client - manager._create_mcp_client = AsyncMock(return_value=mock_client) + manager.create_mcp_client = AsyncMock(return_value=mock_client) user_api_key_auth: Final = UserAPIKeyAuth(api_key="sk-test") # Mock proxy logging proxy_logging_obj = _mock_proxy_logging() - proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(return_value={}) - proxy_logging_obj._convert_mcp_to_llm_format = MagicMock(return_value={}) + proxy_logging_obj.create_mcp_request_object_from_kwargs = MagicMock(return_value={}) + proxy_logging_obj.convert_mcp_to_llm_format = MagicMock(return_value={}) proxy_logging_obj.pre_call_hook = AsyncMock(return_value={}) proxy_logging_obj.during_call_hook = AsyncMock(return_value=None) @@ -7205,11 +7205,11 @@ class TestMCPServerManager: mock_client = AsyncMock() mock_client.call_tool.return_value = MagicMock(spec=CallToolResult, content=[], isError=False) - manager._create_mcp_client = AsyncMock(return_value=mock_client) + manager.create_mcp_client = AsyncMock(return_value=mock_client) proxy_logging_obj = _mock_proxy_logging() - proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(return_value={}) - proxy_logging_obj._convert_mcp_to_llm_format = MagicMock(return_value={}) + proxy_logging_obj.create_mcp_request_object_from_kwargs = MagicMock(return_value={}) + proxy_logging_obj.convert_mcp_to_llm_format = MagicMock(return_value={}) proxy_logging_obj.pre_call_hook = AsyncMock(return_value={}) proxy_logging_obj.during_call_hook = AsyncMock(return_value=None) return manager, proxy_logging_obj @@ -7235,7 +7235,7 @@ class TestMCPServerManager: proxy_logging_obj=proxy_logging_obj, ) - hook_kwargs = proxy_logging_obj._create_mcp_request_object_from_kwargs.call_args.args[0] + hook_kwargs = proxy_logging_obj.create_mcp_request_object_from_kwargs.call_args.args[0] assert (hook_kwargs["tool_description"], hook_kwargs["tool_input_schema"]) == ("Runs the test tool", schema) @pytest.mark.asyncio @@ -7279,7 +7279,7 @@ class TestMCPServerManager: proxy_logging_obj=proxy_logging_obj, ) - hook_kwargs = proxy_logging_obj._create_mcp_request_object_from_kwargs.call_args.args[0] + hook_kwargs = proxy_logging_obj.create_mcp_request_object_from_kwargs.call_args.args[0] assert (hook_kwargs["tool_description"], hook_kwargs["tool_input_schema"]) == (None, None) def test_get_listed_tool_resolves_the_bare_name_from_the_latest_listing(self): @@ -7384,7 +7384,7 @@ class TestMCPServerManager: await release_fetch.wait() return [MCPTool(name="turn", description="before save", inputSchema={})] - manager._create_mcp_client = AsyncMock(return_value=AsyncMock()) + manager.create_mcp_client = AsyncMock(return_value=AsyncMock()) manager._fetch_tools_with_timeout = fetch caller = ListedToolsCaller(user_api_key_auth=user) @@ -7636,7 +7636,7 @@ class TestMCPServerManager: auth_type=MCPAuth.api_key, ) user = UserAPIKeyAuth(api_key="sk-litellm", user_id="byok-cold-user") - manager._create_mcp_client = AsyncMock(return_value=AsyncMock()) + manager.create_mcp_client = AsyncMock(return_value=AsyncMock()) manager._fetch_tools_with_timeout = AsyncMock( return_value=[MCPTool(name="turn", description="listed while db down", inputSchema={})] ) @@ -7688,13 +7688,13 @@ class TestMCPServerManager: user = UserAPIKeyAuth(api_key="sk-litellm", user_id="byok-user") mock_client = AsyncMock() mock_client.call_tool.return_value = MagicMock(spec=CallToolResult, content=[], isError=False) - manager._create_mcp_client = AsyncMock(return_value=mock_client) + manager.create_mcp_client = AsyncMock(return_value=mock_client) manager._fetch_tools_with_timeout = AsyncMock( return_value=[MCPTool(name="turn", description="stored cred catalog", inputSchema={})] ) proxy_logging_obj = _mock_proxy_logging() - proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(return_value={}) - proxy_logging_obj._convert_mcp_to_llm_format = MagicMock(return_value={}) + proxy_logging_obj.create_mcp_request_object_from_kwargs = MagicMock(return_value={}) + proxy_logging_obj.convert_mcp_to_llm_format = MagicMock(return_value={}) proxy_logging_obj.pre_call_hook = AsyncMock(return_value={}) proxy_logging_obj.during_call_hook = AsyncMock(return_value=None) cache_byok_credential("byok-user", "byok-catalog", "stored-secret") @@ -7718,7 +7718,7 @@ class TestMCPServerManager: finally: byok_credential_cache.delete_cache(byok_credential_cache_key("byok-user", "byok-catalog")) - hook_kwargs = proxy_logging_obj._create_mcp_request_object_from_kwargs.call_args.args[0] + hook_kwargs = proxy_logging_obj.create_mcp_request_object_from_kwargs.call_args.args[0] assert hook_kwargs["tool_description"] == "stored cred catalog" @pytest.mark.asyncio @@ -7731,7 +7731,7 @@ class TestMCPServerManager: url="http://byok-catalog", is_byok=True, ) - manager._create_mcp_client = AsyncMock(return_value=AsyncMock()) + manager.create_mcp_client = AsyncMock(return_value=AsyncMock()) manager._fetch_tools_with_timeout = AsyncMock( return_value=[MCPTool(name="turn", description="t", inputSchema={})] ) @@ -7749,7 +7749,7 @@ class TestMCPServerManager: ) listed = manager.get_listed_tool(server, "turn", caller) assert listed is not None and listed.description == "t" - assert manager._create_mcp_client.await_args.kwargs["mcp_auth_header"] == "Bearer hdr" + assert manager.create_mcp_client.await_args.kwargs["mcp_auth_header"] == "Bearer hdr" @pytest.mark.parametrize( "server_auth", @@ -7794,7 +7794,7 @@ class TestMCPServerManager: **server_auth, ) alice = UserAPIKeyAuth(api_key="sk-alice", user_id="alice") - manager._create_mcp_client = AsyncMock(return_value=AsyncMock()) + manager.create_mcp_client = AsyncMock(return_value=AsyncMock()) manager._fetch_tools_with_timeout = AsyncMock( return_value=[MCPTool(name="echo", description="listed catalog", inputSchema={})] ) @@ -7815,7 +7815,7 @@ class TestMCPServerManager: finally: byok_credential_cache.delete_cache(byok_credential_cache_key("alice", "cc1")) - client_kwargs = manager._create_mcp_client.await_args.kwargs + client_kwargs = manager.create_mcp_client.await_args.kwargs assert client_kwargs["mcp_auth_header"] is None, client_kwargs assert client_kwargs["extra_headers"] == {"Authorization": "Bearer signed-jwt"} signer_headers.assert_awaited_once() @@ -8014,10 +8014,10 @@ class TestMCPServerManager: } mock_client = AsyncMock() mock_client.call_tool.return_value = MagicMock(spec=CallToolResult, content=[], isError=False) - manager._create_mcp_client = AsyncMock(return_value=mock_client) + manager.create_mcp_client = AsyncMock(return_value=mock_client) manager._fetch_tools_with_timeout = AsyncMock(side_effect=lambda client, name: catalogs[client.workspace]) for workspace in ("A", "B"): - manager._create_mcp_client.return_value.workspace = workspace + manager.create_mcp_client.return_value.workspace = workspace await manager._get_tools_from_server( server=server, extra_headers={"X-Workspace": workspace}, @@ -8027,8 +8027,8 @@ class TestMCPServerManager: ) proxy_logging_obj = _mock_proxy_logging() - proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(return_value={}) - proxy_logging_obj._convert_mcp_to_llm_format = MagicMock(return_value={}) + proxy_logging_obj.create_mcp_request_object_from_kwargs = MagicMock(return_value={}) + proxy_logging_obj.convert_mcp_to_llm_format = MagicMock(return_value={}) proxy_logging_obj.pre_call_hook = AsyncMock(return_value={}) proxy_logging_obj.during_call_hook = AsyncMock(return_value=None) await manager.call_tool( @@ -8040,7 +8040,7 @@ class TestMCPServerManager: raw_headers={"x-workspace": "A", "authorization": "Bearer sk-litellm"}, ) - hook_kwargs = proxy_logging_obj._create_mcp_request_object_from_kwargs.call_args.args[0] + hook_kwargs = proxy_logging_obj.create_mcp_request_object_from_kwargs.call_args.args[0] assert (hook_kwargs["tool_description"], hook_kwargs["tool_input_schema"]) == ( "Catalog A", {"properties": {"turn": {"description": "A"}}}, @@ -8093,7 +8093,7 @@ class TestMCPServerManager: spec_path="/spec.yaml", ) manager = MCPServerManager() - manager._create_mcp_client = AsyncMock(return_value=AsyncMock()) + manager.create_mcp_client = AsyncMock(return_value=AsyncMock()) async def _handler(**kwargs): return None @@ -8128,7 +8128,7 @@ class TestMCPServerManager: spec_path="/spec.yaml", ) manager = MCPServerManager() - manager._create_mcp_client = AsyncMock(return_value=AsyncMock()) + manager.create_mcp_client = AsyncMock(return_value=AsyncMock()) async def _handler(**kwargs): return None @@ -8170,7 +8170,7 @@ class TestMCPServerManager: server_id="srv", name="srv", alias="srv", transport=MCPTransport.http, url=None, spec_path="/spec.yaml" ) manager = MCPServerManager() - manager._create_mcp_client = AsyncMock(return_value=AsyncMock()) + manager.create_mcp_client = AsyncMock(return_value=AsyncMock()) global_mcp_tool_registry.unregister_tools_with_prefix("srv-") global_mcp_tool_registry.register_tool( name="srv-echo", description="Echoes", input_schema={"type": "object"}, handler=lambda **kwargs: None @@ -8433,7 +8433,7 @@ class TestMCPServerManager: MCPServerAccess, ) from litellm.proxy._experimental.mcp_server.mcp_context import ( - _mcp_active_toolset_id, + mcp_active_toolset_id, ) from litellm.proxy._types import UserAPIKeyAuth @@ -8456,7 +8456,7 @@ class TestMCPServerManager: user_api_key_auth = UserAPIKeyAuth(api_key="sk-test", user_id="user-123") user_api_key_auth.mcp_toolset_id = "toolset-abc" - token = _mcp_active_toolset_id.set("unrelated-ambient-toolset") + token = mcp_active_toolset_id.set("unrelated-ambient-toolset") try: with ( patch.object(proxy_server_module, "user_api_key_cache", cache), @@ -8474,7 +8474,7 @@ class TestMCPServerManager: ): result = await manager.get_allowed_mcp_servers(user_api_key_auth) finally: - _mcp_active_toolset_id.reset(token) + mcp_active_toolset_id.reset(token) assert result == ["toolset-server"] @@ -9599,7 +9599,7 @@ class TestMCPServerTimestamps: timeout=0.01, ) - with patch.object(manager, "_create_mcp_client", return_value=mock_client): + with patch.object(manager, "create_mcp_client", return_value=mock_client): with pytest.raises(HTTPException) as exc_info: await manager._call_regular_mcp_tool( mcp_server=server, @@ -10770,7 +10770,7 @@ class TestHealthCheckInterpolatesGlobalEnvVars: captured["extra_headers"] = extra_headers return mock_client - manager._create_mcp_client = AsyncMock(side_effect=_create) + manager.create_mcp_client = AsyncMock(side_effect=_create) return captured @pytest.mark.asyncio @@ -11331,7 +11331,7 @@ class TestMCPToolsListAuthSurfacing: manager = MCPServerManager() server = MCPServer(server_id="oauth-srv", name="oauth-srv", transport=MCPTransport.http) challenge = 'Bearer resource_metadata="/.well-known/oauth-protected-resource/mcp/oauth-srv"' - manager._create_mcp_client = AsyncMock( + manager.create_mcp_client = AsyncMock( side_effect=HTTPException( status_code=401, detail="Unauthorized", @@ -11355,7 +11355,7 @@ class TestMCPToolsListAuthSurfacing: 401/403 remain the challenge-class statuses routed to MCPUpstreamAuthError.""" manager = MCPServerManager() server = MCPServer(server_id="stdio-srv", name="stdio-srv", transport=MCPTransport.http) - manager._create_mcp_client = AsyncMock( + manager.create_mcp_client = AsyncMock( side_effect=HTTPException( status_code=500, detail="MCP stdio command 'foo' is not in the allowlist", @@ -11393,7 +11393,7 @@ class TestMCPToolsListAuthSurfacing: ) wrapper = RuntimeError("client build failed") wrapper.__cause__ = causal - manager._create_mcp_client = AsyncMock(side_effect=wrapper) + manager.create_mcp_client = AsyncMock(side_effect=wrapper) with pytest.raises(MCPUpstreamAuthError) as exc_info: await manager._get_tools_from_server(server) @@ -11432,7 +11432,7 @@ class TestMCPToolsListAuthSurfacing: ) wrapper = RuntimeError("client build failed") wrapper.__cause__ = causal - manager._create_mcp_client = AsyncMock(side_effect=wrapper) + manager.create_mcp_client = AsyncMock(side_effect=wrapper) with pytest.raises(MCPUpstreamAuthError) as exc_info: await manager._get_tools_from_server(bridge_server) @@ -11462,7 +11462,7 @@ class TestMCPToolsListAuthSurfacing: upstream_challenge = 'Bearer resource_metadata="https://upstream.example/.well-known/oauth-protected-resource"' client = MagicMock() client.list_tools = AsyncMock(side_effect=_upstream_status_error(401, upstream_challenge)) - manager._create_mcp_client = AsyncMock(return_value=client) + manager.create_mcp_client = AsyncMock(return_value=client) with pytest.raises(MCPUpstreamAuthError) as exc_info: await manager._get_tools_from_server(bridge_server) @@ -11487,7 +11487,7 @@ class TestMCPToolsListAuthSurfacing: auth_type=MCPAuth.oauth_delegate, dcr_bridge=True, ) - manager._create_mcp_client = AsyncMock( + manager.create_mcp_client = AsyncMock( side_effect=HTTPException( status_code=401, detail="Unauthorized", @@ -11525,7 +11525,7 @@ class TestMCPToolsListAuthSurfacing: ) return [good_tool] - manager._get_tools_from_server = fake_get_tools + manager.get_tools_from_server = fake_get_tools result = await manager.list_tools() @@ -11544,7 +11544,7 @@ def test_should_strip_caller_authorization_for_token_exchange(): client_id="cid", client_secret="csec", ) - assert _should_strip_caller_authorization(mcp_server=server, raw_headers=None, user_api_key_auth=None) is True + assert should_strip_caller_authorization(mcp_server=server, raw_headers=None, user_api_key_auth=None) is True def _retry_gate_server(auth_type: MCPAuthType) -> MCPServer: @@ -11632,7 +11632,7 @@ class TestOBOCallToolRetry: success = CallToolResult(content=[], isError=False) first = _RetryFakeClient(raises=_UpstreamAuthError(401)) retry = _RetryFakeClient(result=success) - manager._create_mcp_client = AsyncMock(return_value=retry) + manager.create_mcp_client = AsyncMock(return_value=retry) result = await manager._obo_call_tool_with_retry( client=first, @@ -11648,7 +11648,7 @@ class TestOBOCallToolRetry: assert result is success manager._cred_provider.invalidate_credentials.assert_awaited_once() - manager._create_mcp_client.assert_awaited_once() + manager.create_mcp_client.assert_awaited_once() assert first.attempts == 1 and retry.attempts == 1 @pytest.mark.asyncio @@ -11663,7 +11663,7 @@ class TestOBOCallToolRetry: success = CallToolResult(content=[], isError=False) first = _RetryFakeClient(raises=_UpstreamAuthError(401)) retry = _RetryFakeClient(result=success) - manager._create_mcp_client = AsyncMock(return_value=retry) + manager.create_mcp_client = AsyncMock(return_value=retry) server = MCPServer( server_id="id-jag-srv", name="id-jag", @@ -11702,7 +11702,7 @@ class TestOBOCallToolRetry: success = CallToolResult(content=[], isError=False) first = _RetryFakeClient(raises=_UpstreamAuthError(401)) retry = _RetryFakeClient(result=success) - manager._create_mcp_client = AsyncMock(side_effect=[first, retry]) + manager.create_mcp_client = AsyncMock(side_effect=[first, retry]) server = MCPServer( server_id="id-jag-srv", name="id-jag", @@ -11735,7 +11735,7 @@ class TestOBOCallToolRetry: async def test_non_auth_error_does_not_retry(self): manager = self._manager() first = _RetryFakeClient(raises=ValueError("tool blew up")) - manager._create_mcp_client = AsyncMock() + manager.create_mcp_client = AsyncMock() result = await manager._obo_call_tool_with_retry( client=first, @@ -11751,7 +11751,7 @@ class TestOBOCallToolRetry: assert result.is_error is True manager._cred_provider.invalidate_credentials.assert_not_awaited() - manager._create_mcp_client.assert_not_awaited() + manager.create_mcp_client.assert_not_awaited() assert first.attempts == 1 @pytest.mark.asyncio @@ -11760,7 +11760,7 @@ class TestOBOCallToolRetry: first = _RetryFakeClient(raises=_UpstreamAuthError(401)) # The retry client still fails; with raise_on_error defaulting False it returns isError. retry = _RetryFakeClient(raises=_UpstreamAuthError(401)) - manager._create_mcp_client = AsyncMock(return_value=retry) + manager.create_mcp_client = AsyncMock(return_value=retry) result = await manager._obo_call_tool_with_retry( client=first, @@ -11775,7 +11775,7 @@ class TestOBOCallToolRetry: ) assert result.is_error is True - manager._create_mcp_client.assert_awaited_once() + manager.create_mcp_client.assert_awaited_once() assert first.attempts == 1 and retry.attempts == 1 @@ -11819,7 +11819,7 @@ class TestOBOConcurrencyLimit: return CallToolResult(content=[], isError=False) manager = MCPServerManager() - manager._create_mcp_client = AsyncMock(return_value=_ConcurrencyRecordingClient()) + manager.create_mcp_client = AsyncMock(return_value=_ConcurrencyRecordingClient()) async def _dispatch(): return await manager._call_regular_mcp_tool( @@ -12041,7 +12041,7 @@ async def test_aggregate_list_still_absorbs_step_up_challenged_server(): ) return [good_tool] - manager._get_tools_from_server = fake_get_tools + manager.get_tools_from_server = fake_get_tools result = await manager.list_tools() @@ -12869,8 +12869,8 @@ def _mock_proxy_logging() -> MagicMock: def _permissive_proxy_logging() -> MagicMock: proxy_logging_obj: Final = _mock_proxy_logging() - proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(return_value={}) - proxy_logging_obj._convert_mcp_to_llm_format = MagicMock(return_value={}) + proxy_logging_obj.create_mcp_request_object_from_kwargs = MagicMock(return_value={}) + proxy_logging_obj.convert_mcp_to_llm_format = MagicMock(return_value={}) proxy_logging_obj.pre_call_hook = AsyncMock(return_value={}) return proxy_logging_obj @@ -13733,7 +13733,7 @@ class TestResolveOpenapiToolAuth: expected_extra_keys: set, expected_credential: object, ): - auth_value, forwarded, credential = _resolve_openapi_tool_auth( + auth_value, forwarded, credential = resolve_openapi_tool_auth( mcp_server=self._server(), mcp_auth_header=byok, mcp_server_auth_headers=per_server, @@ -13747,7 +13747,7 @@ class TestResolveOpenapiToolAuth: def test_per_server_value_is_never_re_prefixed(self): """The regression that a naive wiring produces: the caller already sent ``Bearer ``.""" - auth_value, _, credential = _resolve_openapi_tool_auth( + auth_value, _, credential = resolve_openapi_tool_auth( mcp_server=self._server(auth_type=MCPAuth.api_key), mcp_auth_header="byok-secret", mcp_server_auth_headers={"report_api": "Bearer caller-token"}, @@ -13762,7 +13762,7 @@ class TestResolveOpenapiToolAuth: def test_per_server_authorization_is_not_also_left_in_forwarded_headers(self): """``resolve_openapi_upstream_auth`` pops Authorization out of the forwarded headers, so a second copy there would give the passthrough arm two sources to reconcile.""" - _, forwarded, _ = _resolve_openapi_tool_auth( + _, forwarded, _ = resolve_openapi_tool_auth( mcp_server=self._server(), mcp_auth_header=None, mcp_server_auth_headers={"report_api": {"Authorization": "Bearer caller-token"}}, @@ -14518,12 +14518,12 @@ class TestLitellmAdmissionKeyIsNeverTheSubjectToken: client.call_tool = AsyncMock(return_value=CallToolResult(content=[], isError=False)) client.list_prompts_result = AsyncMock(return_value=ListPromptsResult(prompts=[])) client.read_resource = AsyncMock(return_value=ReadResourceResult(contents=[])) - manager._create_mcp_client = AsyncMock(return_value=client) + manager.create_mcp_client = AsyncMock(return_value=client) return manager @staticmethod def _subject_token_given_to_client(manager: MCPServerManager) -> str | None: - return manager._create_mcp_client.call_args.kwargs["subject_token"] + return manager.create_mcp_client.call_args.kwargs["subject_token"] async def _call_tool_subject(self, server: MCPServer, oauth2_headers, raw_headers, user_api_key_auth): manager: Final = self._manager_with_recording_client() @@ -15549,7 +15549,7 @@ async def test_openapi_listing_ignores_overlapping_server_prefix() -> None: from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry manager: Final = MCPServerManager() - manager._create_mcp_client = AsyncMock(return_value=AsyncMock()) + manager.create_mcp_client = AsyncMock(return_value=AsyncMock()) for prefix in ("pet-", "petstore-"): global_mcp_tool_registry.unregister_tools_with_prefix(prefix) _register_local_tool("pet-list", "Local pet tool") @@ -15570,7 +15570,7 @@ async def test_openapi_listing_finds_tools_registered_under_the_normalized_prefi from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry manager: Final = MCPServerManager() - manager._create_mcp_client = AsyncMock(return_value=AsyncMock()) + manager.create_mcp_client = AsyncMock(return_value=AsyncMock()) global_mcp_tool_registry.unregister_tools_with_prefix("pet_store-") _register_local_tool("pet_store-list", "Pet store tool") try: @@ -16151,9 +16151,9 @@ class TestProtectedCredentialPreparation: caller: str | None, ) -> None: from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( - _request_auth_header, - _request_extra_headers, create_tool_function, + request_auth_header, + request_extra_headers, ) tool: Final = create_tool_function( @@ -16166,8 +16166,8 @@ class TestProtectedCredentialPreparation: ) monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") destination: Final = respx_mock.get("https://upstream.example/echo").respond(200, text="authenticated") - caller_token: Final = _request_auth_header.set(caller) - extra_token: Final = _request_extra_headers.set(forwarded) + caller_token: Final = request_auth_header.set(caller) + extra_token: Final = request_extra_headers.set(forwarded) try: assert await tool() == TextResult("authenticated") sent: Final = destination.calls.last.request.headers @@ -16176,8 +16176,8 @@ class TestProtectedCredentialPreparation: assert sent["authorization"] == caller assert destination.call_count == 1 finally: - _request_auth_header.reset(caller_token) - _request_extra_headers.reset(extra_token) + request_auth_header.reset(caller_token) + request_extra_headers.reset(extra_token) @pytest.mark.asyncio async def test_static_resolution_cancellation_closes_flow(self) -> None: @@ -17925,7 +17925,7 @@ def catalog_guardrail(monkeypatch): def _catalog_manager(*upstream_tools: MCPTool) -> MCPServerManager: manager = MCPServerManager() - manager._create_mcp_client = AsyncMock(return_value=object()) + manager.create_mcp_client = AsyncMock(return_value=object()) manager._fetch_tools_with_timeout = AsyncMock(return_value=list(upstream_tools)) return manager @@ -18436,8 +18436,8 @@ class TestToolCatalogGuard: server = _notes_server({"list_notes": _pin(LIST_NOTES)}) user_api_key_auth = MagicMock(object_permission=None, object_permission_id=None) proxy_logging_obj = _mock_proxy_logging() - proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(return_value={}) - proxy_logging_obj._convert_mcp_to_llm_format = MagicMock(return_value={}) + proxy_logging_obj.create_mcp_request_object_from_kwargs = MagicMock(return_value={}) + proxy_logging_obj.convert_mcp_to_llm_format = MagicMock(return_value={}) proxy_logging_obj.pre_call_hook = AsyncMock(return_value={}) with pytest.raises(HTTPException) as exc_info: @@ -18910,7 +18910,7 @@ async def test_catalog_page_registers_bare_routes_only_for_complete_initial_disc other = MCPServer(server_id="other", name="other", transport=MCPTransport.http) manager.registry = {server.server_id: server, other.server_id: other} manager._create_prefixed_tools([LIST_NOTES], other) - manager._create_mcp_client.return_value = SimpleNamespace( + manager.create_mcp_client.return_value = SimpleNamespace( list_tools_page=AsyncMock(return_value=ListToolsResult(tools=[LIST_NOTES], next_cursor=next_cursor)) ) @@ -18943,7 +18943,7 @@ async def test_paginated_listing_keeps_earlier_tool_metadata_and_caller_isolatio ListToolsResult(tools=[second]), ListToolsResult(tools=[second]), ] - manager._create_mcp_client = AsyncMock(return_value=client) + manager.create_mcp_client = AsyncMock(return_value=client) monkeypatch.setattr(operations, "global_mcp_server_manager", manager) caller = UserAPIKeyAuth(api_key="owned-caller", user_id="alice") context = operations.prepare_context(caller) @@ -19025,7 +19025,7 @@ async def test_failed_aggregate_continuation_preserves_only_delivered_tool_metad ] async def create_client(server, **kwargs): return clients[server.server_id] - manager._create_mcp_client = create_client + manager.create_mcp_client = create_client monkeypatch.setattr(operations, "global_mcp_server_manager", manager) caller = UserAPIKeyAuth(api_key="owned-caller", user_id="alice") context = operations.prepare_context(caller) @@ -19066,7 +19066,7 @@ async def test_aggregate_publishes_complete_bare_routes_only_after_delivering_a_ async def create_client(server, **kwargs): return clients[server.server_id] - manager._create_mcp_client = create_client + manager.create_mcp_client = create_client monkeypatch.setattr(operations, "global_mcp_server_manager", manager) context = operations.prepare_context(UserAPIKeyAuth(api_key="owned-caller", user_id="alice")) listing = catalog.aggregate_gateway_tools(context, PaginatedRequestParams(), servers, {}, record_listing=True) diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py index df3f1d78a7e..97c6e2703ad 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py @@ -826,7 +826,7 @@ async def test_call_tool_m2m_skips_authorization_headers(): mock_client = MagicMock() mock_client.call_tool = AsyncMock(return_value=MagicMock()) - with patch.object(manager, "_create_mcp_client", new=AsyncMock(return_value=mock_client)) as create_client_mock: + with patch.object(manager, "create_mcp_client", new=AsyncMock(return_value=mock_client)) as create_client_mock: await manager._call_regular_mcp_tool( mcp_server=server, original_tool_name="echo", @@ -1302,7 +1302,7 @@ async def test_get_tools_from_mcp_servers_continues_when_one_server_fails(): # Failing server raises an exception raise Exception("Server connection failed") - mock_manager._get_tools_from_server = mock_get_tools_from_server + mock_manager.get_tools_from_server = mock_get_tools_from_server with patch( "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager", @@ -1398,7 +1398,7 @@ async def test_get_tools_from_mcp_servers_handles_all_servers_failing(): # All servers fail raise Exception(f"Server {server.name} connection failed") - mock_manager._get_tools_from_server = mock_get_tools_from_server + mock_manager.get_tools_from_server = mock_get_tools_from_server with patch( "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager", @@ -4237,11 +4237,11 @@ async def test_mcp_routing_with_conflicting_alias_and_group_name(): mock_get_allowed, ), patch( - "litellm.proxy._experimental.mcp_server.server.MCPRequestHandler._get_mcp_servers_from_access_groups", + "litellm.proxy._experimental.mcp_server.server.MCPRequestHandler.get_mcp_servers_from_access_groups", mock_db_lookup, ), patch( - "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager._get_tools_from_server", + "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_tools_from_server", mock_get_tools_spy, ), ): @@ -4341,7 +4341,7 @@ async def test_oauth2_caller_headers_not_forwarded_for_migrated_server(): with ( patch.object( global_mcp_server_manager, - "_create_mcp_client", + "create_mcp_client", side_effect=mock_create_mcp_client, ) as mock_create_client, patch.object( @@ -4438,7 +4438,7 @@ async def test_list_tools_single_server_unprefixed_names(): tool.input_schema = {} return [tool] - mock_manager._get_tools_from_server = mock_get_tools_from_server + mock_manager.get_tools_from_server = mock_get_tools_from_server with patch( "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager", @@ -4517,7 +4517,7 @@ async def test_list_tools_multiple_servers_prefixed_names(): tool.input_schema = {} return [tool] - mock_manager._get_tools_from_server = mock_get_tools_from_server + mock_manager.get_tools_from_server = mock_get_tools_from_server with patch( "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager", @@ -4557,7 +4557,7 @@ async def test_mcp_manager_allows_public_servers_without_permissions(): with ( patch( - "litellm.proxy.management_endpoints.common_utils._user_has_admin_view", + "litellm.proxy.management_endpoints.common_utils.user_api_key_has_admin_view", return_value=False, ), patch( @@ -4592,7 +4592,7 @@ async def test_mcp_manager_returns_public_when_permission_lookup_fails(): with ( patch( - "litellm.proxy.management_endpoints.common_utils._user_has_admin_view", + "litellm.proxy.management_endpoints.common_utils.user_api_key_has_admin_view", return_value=False, ), patch( @@ -4636,7 +4636,7 @@ async def test_mcp_manager_merges_public_and_restricted_servers(): with ( patch( - "litellm.proxy.management_endpoints.common_utils._user_has_admin_view", + "litellm.proxy.management_endpoints.common_utils.user_api_key_has_admin_view", return_value=False, ), patch( @@ -4946,7 +4946,7 @@ async def test_list_tools_filters_by_key_team_permissions(): return [tool1, tool2, tool3, tool4] - mock_manager._get_tools_from_server = mock_get_tools_from_server + mock_manager.get_tools_from_server = mock_get_tools_from_server with patch( "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager", @@ -5057,7 +5057,7 @@ async def test_list_tools_with_team_tool_permissions_inheritance(): return [tool1, tool2, tool3, tool4] - mock_manager._get_tools_from_server = mock_get_tools_from_server + mock_manager.get_tools_from_server = mock_get_tools_from_server with patch( "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager", @@ -5149,7 +5149,7 @@ async def test_list_tools_with_no_tool_permissions_shows_all(): return [tool1, tool2, tool3] - mock_manager._get_tools_from_server = mock_get_tools_from_server + mock_manager.get_tools_from_server = mock_get_tools_from_server with patch( "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager", @@ -5255,7 +5255,7 @@ async def test_list_tools_strips_prefix_when_matching_permissions(): return [tool1, tool2, tool3, tool4] - mock_manager._get_tools_from_server = mock_get_tools_from_server + mock_manager.get_tools_from_server = mock_get_tools_from_server with patch( "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager", @@ -5796,7 +5796,7 @@ async def test_get_tools_from_mcp_servers_logs_list_tools_to_spendlogs_when_enab side_effect=_capture_function_setup, ), ): - mock_manager._get_tools_from_server = AsyncMock(return_value=[tool_1]) + mock_manager.get_tools_from_server = AsyncMock(return_value=[tool_1]) listing = await _get_tools_from_mcp_servers( user_api_key_auth=user_auth, @@ -5878,7 +5878,7 @@ async def test_get_tools_from_mcp_servers_returns_tools_when_success_logging_fai return_value=(dummy_logging_obj, None), ), ): - mock_manager._get_tools_from_server = AsyncMock(return_value=[tool_1]) + mock_manager.get_tools_from_server = AsyncMock(return_value=[tool_1]) listing = await _get_tools_from_mcp_servers( user_api_key_auth=user_auth, @@ -6195,7 +6195,7 @@ async def test_get_tools_from_mcp_servers_injects_stored_oauth2_token(): new=AsyncMock(side_effect=lambda tools, **_: tools), ), ): - mock_manager._get_tools_from_server = AsyncMock(return_value=[tool_1]) + mock_manager.get_tools_from_server = AsyncMock(return_value=[tool_1]) listing = await _get_tools_from_mcp_servers( user_api_key_auth=user_auth, @@ -6209,8 +6209,8 @@ async def test_get_tools_from_mcp_servers_injects_stored_oauth2_token(): mock_prefetch.assert_awaited_once_with(user_auth) # The stored token was forwarded to the MCP transport layer as extra_headers - mock_manager._get_tools_from_server.assert_awaited_once() - call_kwargs = mock_manager._get_tools_from_server.await_args.kwargs + mock_manager.get_tools_from_server.assert_awaited_once() + call_kwargs = mock_manager.get_tools_from_server.await_args.kwargs assert call_kwargs["extra_headers"] == {"Authorization": f"Bearer {STORED_TOKEN}"} assert listing.tools == [tool_1] @@ -6351,7 +6351,7 @@ class TestEnsureUpstreamInitializeInstructionsCached: ) server = _make_instruction_server(server_id="yaml-only", instructions="from yaml") - with patch.object(global_mcp_server_manager, "_create_mcp_client", AsyncMock()) as mock_create: + with patch.object(global_mcp_server_manager, "create_mcp_client", AsyncMock()) as mock_create: await global_mcp_server_manager._ensure_upstream_initialize_instructions_cached(server) mock_create.assert_not_awaited() @@ -6366,7 +6366,7 @@ class TestEnsureUpstreamInitializeInstructionsCached: server = _make_instruction_server(server_id="cached-only", instructions=None) global_mcp_server_manager._upstream_initialize_instructions_by_server_id["cached-only"] = "warm" try: - with patch.object(global_mcp_server_manager, "_create_mcp_client", AsyncMock()) as mock_create: + with patch.object(global_mcp_server_manager, "create_mcp_client", AsyncMock()) as mock_create: await global_mcp_server_manager._ensure_upstream_initialize_instructions_cached(server) mock_create.assert_not_awaited() finally: @@ -6381,7 +6381,7 @@ class TestEnsureUpstreamInitializeInstructionsCached: ) server = _make_instruction_server(server_id="openapi-spec", spec_path="/openapi.json", url=None) - with patch.object(global_mcp_server_manager, "_create_mcp_client", AsyncMock()) as mock_create: + with patch.object(global_mcp_server_manager, "create_mcp_client", AsyncMock()) as mock_create: await global_mcp_server_manager._ensure_upstream_initialize_instructions_cached(server) mock_create.assert_not_awaited() @@ -6400,7 +6400,7 @@ class TestEnsureUpstreamInitializeInstructionsCached: with patch.object( global_mcp_server_manager, - "_create_mcp_client", + "create_mcp_client", AsyncMock(return_value=fake_client), ): try: @@ -6428,7 +6428,7 @@ class TestEnsureUpstreamInitializeInstructionsCached: fake_client._last_initialize_instructions = None # upstream sent nothing create = AsyncMock(return_value=fake_client) - with patch.object(global_mcp_server_manager, "_create_mcp_client", create): + with patch.object(global_mcp_server_manager, "create_mcp_client", create): try: await global_mcp_server_manager._ensure_upstream_initialize_instructions_cached(server) await global_mcp_server_manager._ensure_upstream_initialize_instructions_cached(server) @@ -6453,7 +6453,7 @@ class TestEnsureUpstreamInitializeInstructionsCached: fake_client._last_initialize_instructions = None create = AsyncMock(return_value=fake_client) - with patch.object(global_mcp_server_manager, "_create_mcp_client", create): + with patch.object(global_mcp_server_manager, "create_mcp_client", create): try: await global_mcp_server_manager._ensure_upstream_initialize_instructions_cached(server) await global_mcp_server_manager._ensure_upstream_initialize_instructions_cached(server) @@ -6485,22 +6485,22 @@ class TestGatewayCreateInitializationOptions: """When ContextVar is None, instructions are absent.""" try: from litellm.proxy._experimental.mcp_server.mcp_context import ( - _mcp_gateway_initialize_instructions, - _mcp_gateway_server_name, + mcp_gateway_initialize_instructions, + mcp_gateway_server_name, ) from litellm.proxy._experimental.mcp_server.server import server except ImportError: pytest.skip("MCP server not available") - instructions_token = _mcp_gateway_initialize_instructions.set(None) - server_name_token = _mcp_gateway_server_name.set(None) + instructions_token = mcp_gateway_initialize_instructions.set(None) + server_name_token = mcp_gateway_server_name.set(None) try: opts = server.create_initialization_options() assert getattr(opts, "instructions", None) is None assert opts.server_name == "litellm-mcp-server" finally: - _mcp_gateway_initialize_instructions.reset(instructions_token) - _mcp_gateway_server_name.reset(server_name_token) + mcp_gateway_initialize_instructions.reset(instructions_token) + mcp_gateway_server_name.reset(server_name_token) @pytest.mark.asyncio async def test_scoped_request_uses_configured_server_alias(self): @@ -6529,7 +6529,7 @@ class TestGatewayCreateInitializationOptions: ), patch.object( global_mcp_server_manager, - "_ensure_upstream_initialize_instructions_cached", + "ensure_upstream_initialize_instructions_cached", new_callable=AsyncMock, ), ): @@ -6599,7 +6599,7 @@ class TestGatewayCreateInitializationOptions: async def test_non_initialize_request_with_no_granted_servers_is_not_rejected_here(self): from litellm.proxy._experimental.mcp_server.server import ( _gateway_initialize_instructions_request_scope, - _mcp_gateway_initialize_instructions, + mcp_gateway_initialize_instructions, ) from litellm.proxy._types import UserAPIKeyAuth @@ -6613,7 +6613,7 @@ class TestGatewayCreateInitializationOptions: mcp_servers=None, client_ip=None, ): - assert _mcp_gateway_initialize_instructions.get() is None + assert mcp_gateway_initialize_instructions.get() is None @pytest.mark.asyncio async def test_sse_handler_scopes_server_name_from_single_server_path(self): @@ -6678,7 +6678,7 @@ class TestGatewayCreateInitializationOptions: ), patch.object( global_mcp_server_manager, - "_ensure_upstream_initialize_instructions_cached", + "ensure_upstream_initialize_instructions_cached", new_callable=AsyncMock, ), patch( @@ -6708,31 +6708,31 @@ class TestGatewayCreateInitializationOptions: """When ContextVar has a value, it appears in InitializationOptions.""" try: from litellm.proxy._experimental.mcp_server.mcp_context import ( - _mcp_gateway_initialize_instructions, + mcp_gateway_initialize_instructions, ) from litellm.proxy._experimental.mcp_server.server import server except ImportError: pytest.skip("MCP server not available") - tok = _mcp_gateway_initialize_instructions.set("hello from merge") + tok = mcp_gateway_initialize_instructions.set("hello from merge") try: opts = server.create_initialization_options() assert opts.instructions == "hello from merge" finally: - _mcp_gateway_initialize_instructions.reset(tok) + mcp_gateway_initialize_instructions.reset(tok) def test_contextvar_reset_removes_instructions(self): """After resetting the ContextVar, instructions disappear.""" try: from litellm.proxy._experimental.mcp_server.mcp_context import ( - _mcp_gateway_initialize_instructions, + mcp_gateway_initialize_instructions, ) from litellm.proxy._experimental.mcp_server.server import server except ImportError: pytest.skip("MCP server not available") - tok = _mcp_gateway_initialize_instructions.set("temporary") - _mcp_gateway_initialize_instructions.reset(tok) + tok = mcp_gateway_initialize_instructions.set("temporary") + mcp_gateway_initialize_instructions.reset(tok) opts = server.create_initialization_options() assert getattr(opts, "instructions", None) is None @@ -6821,7 +6821,7 @@ async def test_list_tools_with_legacy_db_m2m_server_resolves_oauth2_flow(): mock_manager.get_allowed_mcp_servers = AsyncMock(return_value=["legacy-m2m-id"]) mock_manager.get_mcp_server_by_id = MagicMock(return_value=legacy_server) mock_manager.filter_server_ids_by_ip_with_info = MagicMock(return_value=(["legacy-m2m-id"], 0)) - mock_manager._get_tools_from_server = AsyncMock(side_effect=capture_extra_headers) + mock_manager.get_tools_from_server = AsyncMock(side_effect=capture_extra_headers) listing = await _get_tools_from_mcp_servers( user_api_key_auth=user_auth, @@ -6891,7 +6891,7 @@ async def test_call_tool_empty_extra_headers_returns_none(): with ( patch.object( manager, - "_create_mcp_client", + "create_mcp_client", side_effect=capture_create_mcp_client, ), patch.object( @@ -7488,7 +7488,7 @@ async def test_execute_mcp_tool_rest_server_id_authoritative_for_unprefixed_tool ), patch.object( mcp_operations.global_mcp_server_manager, - "_get_mcp_server_from_tool_name", + "get_mcp_server_from_tool_name", return_value=oauth_server, ), patch.object( @@ -7553,7 +7553,7 @@ def _worker_that_never_listed(server: MCPServer, upstream_tools: tuple[str, ...] with ( patch.object( # test-quality-ok: the upstream MCP session is the boundary; a real one needs an initialize handshake over a live server mcp_operations.global_mcp_server_manager, - "_create_mcp_client", + "create_mcp_client", new=AsyncMock(return_value=MagicMock()), ) as create_client, patch.object( # test-quality-ok: same boundary, this is the tools/list answer the upstream would give @@ -7735,7 +7735,7 @@ async def test_execute_mcp_tool_strips_a_prefix_that_contains_the_separator(): with ( patch.object( mcp_operations.global_mcp_server_manager, - "_get_mcp_server_from_tool_name", + "get_mcp_server_from_tool_name", return_value=alias_less_server, ), patch.object( @@ -7821,7 +7821,7 @@ async def test_execute_mcp_tool_rest_server_id_injects_requested_server_credenti ), patch.object( mcp_operations.global_mcp_server_manager, - "_create_mcp_client", + "create_mcp_client", new=fake_create_mcp_client, ), patch.object( @@ -7888,7 +7888,7 @@ async def test_execute_mcp_tool_rest_prefixed_tool_still_validates_server_id(): ), patch.object( mcp_operations.global_mcp_server_manager, - "_get_mcp_server_from_tool_name", + "get_mcp_server_from_tool_name", return_value=oauth_server, ), patch.object( @@ -7950,7 +7950,7 @@ async def test_execute_mcp_tool_rest_unauthorized_prefix_still_mismatches(): ), patch.object( mcp_operations.global_mcp_server_manager, - "_get_mcp_server_from_tool_name", + "get_mcp_server_from_tool_name", return_value=restricted_server, ), patch.object( @@ -8011,7 +8011,7 @@ async def test_execute_mcp_tool_rest_hyphenated_upstream_tool_name_routes_to_req ), patch.object( mcp_operations.global_mcp_server_manager, - "_get_mcp_server_from_tool_name", + "get_mcp_server_from_tool_name", return_value=None, ), patch.object( @@ -8094,7 +8094,7 @@ async def test_execute_mcp_tool_sets_model_in_model_call_details(): with ( patch.object( mcp_operations.global_mcp_server_manager, - "_get_mcp_server_from_tool_name", + "get_mcp_server_from_tool_name", return_value=fake_server, ), patch.object( @@ -8170,7 +8170,7 @@ async def test_execute_mcp_tool_hands_openapi_hooks_the_listed_entry_and_nothing try: with ( - patch.object(manager, "_get_mcp_server_from_tool_name", return_value=petstore), + patch.object(manager, "get_mcp_server_from_tool_name", return_value=petstore), patch.object(manager, "pre_call_tool_check", new=pre_call_tool_check), patch("litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging), ): @@ -8225,7 +8225,7 @@ async def test_execute_mcp_tool_hands_openapi_hooks_the_guarded_catalog_entry_cl try: with ( - patch.object(manager, "_get_mcp_server_from_tool_name", return_value=petstore), + patch.object(manager, "get_mcp_server_from_tool_name", return_value=petstore), patch.object(manager, "pre_call_tool_check", new=pre_call_tool_check), ): await mcp_module.execute_mcp_tool( @@ -8280,7 +8280,7 @@ async def test_execute_mcp_tool_hands_openapi_hooks_each_callers_own_listed_entr try: with ( - patch.object(manager, "_get_mcp_server_from_tool_name", return_value=petstore), + patch.object(manager, "get_mcp_server_from_tool_name", return_value=petstore), patch.object(manager, "pre_call_tool_check", new=pre_call_tool_check), ): for caller in (guarded, opted_out): @@ -8328,7 +8328,7 @@ async def test_execute_mcp_tool_runs_the_longer_colliding_operation_and_hands_ho try: with ( - patch.object(manager, "_get_mcp_server_from_tool_name", return_value=petstore), + patch.object(manager, "get_mcp_server_from_tool_name", return_value=petstore), patch.object(manager, "pre_call_tool_check", new=pre_call_tool_check), ): result = await mcp_module.execute_mcp_tool( @@ -8374,7 +8374,7 @@ async def test_execute_mcp_tool_hands_hooks_nothing_for_a_never_listed_operation try: with ( - patch.object(manager, "_get_mcp_server_from_tool_name", return_value=petstore), + patch.object(manager, "get_mcp_server_from_tool_name", return_value=petstore), patch.object(manager, "pre_call_tool_check", new=pre_call_tool_check), ): result = await mcp_module.execute_mcp_tool( @@ -8404,7 +8404,7 @@ async def test_execute_mcp_tool_implicit_listing_before_the_first_call_hands_hoo upstream = AsyncMock() upstream.call_tool.return_value = CallToolResult(content=[TextContent(type="text", text="ok")], isError=False) proxy_logging = _mock_mcp_proxy_logging() - proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(return_value={}) + proxy_logging.create_mcp_request_object_from_kwargs = MagicMock(return_value={}) proxy_logging._convert_mcp_to_llm_format = MagicMock(return_value={}) proxy_logging.pre_call_hook = AsyncMock(return_value={}) proxy_logging.during_call_hook = AsyncMock(return_value=None) @@ -8413,7 +8413,7 @@ async def test_execute_mcp_tool_implicit_listing_before_the_first_call_hands_hoo ) with ( - patch.object(manager, "_create_mcp_client", new=AsyncMock(return_value=upstream)), + patch.object(manager, "create_mcp_client", new=AsyncMock(return_value=upstream)), patch.object(manager, "_fetch_tools_with_timeout", new=fetch_tools), patch("litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging), ): @@ -8429,7 +8429,7 @@ async def test_execute_mcp_tool_implicit_listing_before_the_first_call_hands_hoo assert fetch_tools.await_count == 1 assert upstream.call_tool.await_count == 1 assert result.content[0].text == "ok" - hook_kwargs = proxy_logging._create_mcp_request_object_from_kwargs.call_args.args[0] + hook_kwargs = proxy_logging.create_mcp_request_object_from_kwargs.call_args.args[0] assert (hook_kwargs["tool_description"], hook_kwargs["tool_input_schema"]) == (None, None) assert server.server_id not in manager._listed_tools_by_server_id @@ -8460,7 +8460,7 @@ async def test_fetch_pinnable_tool_catalog_records_no_listed_catalog_for_the_adm ) with ( - patch.object(manager, "_create_mcp_client", new=AsyncMock(return_value=MagicMock())), + patch.object(manager, "create_mcp_client", new=AsyncMock(return_value=MagicMock())), patch.object(manager, "_fetch_tools_with_timeout", new=fetch_tools), patch("litellm.proxy.proxy_server.proxy_logging_obj", ProxyLogging(user_api_key_cache=DualCache())), ): @@ -8523,7 +8523,7 @@ async def test_execute_mcp_tool_rest_unresolved_prefixed_name_routes_to_requeste ), patch.object( mcp_operations.global_mcp_server_manager, - "_get_mcp_server_from_tool_name", + "get_mcp_server_from_tool_name", return_value=None, ), patch.object( @@ -8603,7 +8603,7 @@ async def test_execute_mcp_tool_rest_prefix_retry_resolution_still_enforces_serv ), patch.object( mcp_operations.global_mcp_server_manager, - "_get_mcp_server_from_tool_name", + "get_mcp_server_from_tool_name", side_effect=resolve_only_when_requested_prefix_added, ), patch.object( @@ -8671,7 +8671,7 @@ async def test_get_allowed_mcp_servers_from_mcp_server_names_unknown_name_fails_ with patch( "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp." - "MCPRequestHandler._get_mcp_servers_from_access_groups", + "MCPRequestHandler.get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ): @@ -8730,7 +8730,7 @@ async def test_get_allowed_mcp_servers_from_mcp_server_names_known_alias_returns with patch( "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp." - "MCPRequestHandler._get_mcp_servers_from_access_groups", + "MCPRequestHandler.get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ): @@ -8763,7 +8763,7 @@ async def test_get_allowed_mcp_servers_from_mcp_server_names_mixed_known_and_unk with patch( "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp." - "MCPRequestHandler._get_mcp_servers_from_access_groups", + "MCPRequestHandler.get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ): @@ -8796,7 +8796,7 @@ async def test_get_allowed_mcp_servers_from_mcp_server_names_access_group_resolv with patch( "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp." - "MCPRequestHandler._get_mcp_servers_from_access_groups", + "MCPRequestHandler.get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=["id-b"], ): @@ -9630,9 +9630,9 @@ def test_redact_mcp_resource_url_strips_credentials(url, expected): """The MCP tool-call log records the upstream resource, so the URL must be redacted to scheme+host+path: userinfo, query string, and fragment (which can carry embedded tokens or secret parameters) must never reach spend-log metadata or logging callbacks.""" - from litellm.proxy._experimental.mcp_server.server import _redact_mcp_resource_url + from litellm.proxy._experimental.mcp_server.server import redact_mcp_resource_url - assert _redact_mcp_resource_url(url) == expected + assert redact_mcp_resource_url(url) == expected @pytest.mark.asyncio @@ -9721,7 +9721,7 @@ def _managed_tool_returning(server, upstream_result, proxy_logging_mock): return_value=[server.server_id], ), patch.object(global_mcp_server_manager, "get_mcp_server_by_id", return_value=server), - patch.object(global_mcp_server_manager, "_get_mcp_server_from_tool_name", return_value=server), + patch.object(global_mcp_server_manager, "get_mcp_server_from_tool_name", return_value=server), patch.object(global_mcp_server_manager, "server_owning_tool_name_prefix", return_value=server), patch( "litellm.proxy._experimental.mcp_server.operations._get_allowed_mcp_servers_from_mcp_server_names", @@ -9898,7 +9898,7 @@ async def test_aggregate_listing_reports_per_server_outcomes(): return [tool1] raise MCPServerListError(ServerListFault(tag="upstream_error", status_code=500), server.name) - mock_manager._get_tools_from_server = mock_get_tools_from_server + mock_manager.get_tools_from_server = mock_get_tools_from_server with patch( "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager", @@ -10747,7 +10747,7 @@ async def test_list_tools_injects_byok_credential_for_non_oauth2_auth_types(auth mock_manager.get_allowed_mcp_servers = AsyncMock(return_value=[server.server_id]) mock_manager.get_mcp_server_by_id = MagicMock(return_value=server) mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: (server_ids, 0) - mock_manager._get_tools_from_server = mock_get_tools_from_server + mock_manager.get_tools_from_server = mock_get_tools_from_server with ( patch( @@ -10949,7 +10949,7 @@ async def test_tool_listing_preserves_permission_denial_when_failure_logging_fai patch.object(operations, "_get_allowed_mcp_servers", AsyncMock(side_effect=denial)), patch.object(operations, "function_setup", return_value=(None, None)), patch.object(proxy_server, "proxy_logging_obj", logger), - patch.object(operations.global_mcp_server_manager, "_get_tools_from_server", upstream), + patch.object(operations.global_mcp_server_manager, "get_tools_from_server", upstream), ): with pytest.raises(HTTPException) as rejected: await operations._get_tools_from_mcp_servers( diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_sigv4_auth.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_sigv4_auth.py index 6430d5f9259..dda8f0b1044 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_sigv4_auth.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_sigv4_auth.py @@ -615,7 +615,7 @@ class TestCredentialMergeOnUpdate: with ( patch( - "litellm.proxy._experimental.mcp_server.db._get_salt_key", + "litellm.proxy._experimental.mcp_server.db.get_salt_key", return_value=None, ), patch( @@ -652,7 +652,7 @@ class TestCredentialMergeOnUpdate: ) with patch( - "litellm.proxy._experimental.mcp_server.db._get_salt_key", + "litellm.proxy._experimental.mcp_server.db.get_salt_key", return_value=None, ): await update_mcp_server(mock_prisma, data, "test-user") @@ -682,7 +682,7 @@ class TestCredentialMergeOnUpdate: with ( patch( - "litellm.proxy._experimental.mcp_server.db._get_salt_key", + "litellm.proxy._experimental.mcp_server.db.get_salt_key", return_value=None, ), patch( @@ -724,7 +724,7 @@ class TestCredentialMergeOnUpdate: with ( patch( - "litellm.proxy._experimental.mcp_server.db._get_salt_key", + "litellm.proxy._experimental.mcp_server.db.get_salt_key", return_value=None, ), patch( @@ -767,7 +767,7 @@ class TestCredentialMergeOnUpdate: with ( patch( - "litellm.proxy._experimental.mcp_server.db._get_salt_key", + "litellm.proxy._experimental.mcp_server.db.get_salt_key", return_value=None, ), patch( @@ -1001,7 +1001,7 @@ class TestRotateCredentials: with ( patch( - "litellm.proxy._experimental.mcp_server.db._get_salt_key", + "litellm.proxy._experimental.mcp_server.db.get_salt_key", return_value="old-key", ), patch( @@ -1050,7 +1050,7 @@ class TestRotateCredentials: with ( patch( - "litellm.proxy._experimental.mcp_server.db._get_salt_key", + "litellm.proxy._experimental.mcp_server.db.get_salt_key", return_value="old-key", ), patch( @@ -1101,7 +1101,7 @@ class TestAuthTypeSwitchClearsCredentials: ) with patch( - "litellm.proxy._experimental.mcp_server.db._get_salt_key", + "litellm.proxy._experimental.mcp_server.db.get_salt_key", return_value=None, ): await update_mcp_server(mock_prisma, data, "test-user") @@ -1121,7 +1121,7 @@ class TestInheritCredentials: def test_inherits_sigv4_credentials(self): """SigV4 fields are copied from existing server to inherited credentials.""" from litellm.proxy.management_endpoints.mcp_management_endpoints import ( - _inherit_credentials_from_existing_server, + inherit_credentials_from_existing_server, ) from litellm.proxy._types import NewMCPServerRequest from litellm.types.mcp_server.mcp_server_manager import MCPServer @@ -1152,7 +1152,7 @@ class TestInheritCredentials: "litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager" ) as mock_manager: mock_manager.get_mcp_server_by_id.return_value = existing - result = _inherit_credentials_from_existing_server(payload) + result = inherit_credentials_from_existing_server(payload) assert result.credentials is not None assert result.credentials["aws_access_key_id"] == "AKIAEXAMPLE" diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_tool_search.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_tool_search.py index 1c987778da6..c3a94e04b9f 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_tool_search.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_tool_search.py @@ -1373,7 +1373,7 @@ async def test_mcp_tool_search_leaves_the_listed_tools_slot_empty(monkeypatch: p with ( patch.dict(manager.tool_name_to_mcp_server_name_mapping), patch.object(mcp_operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[server])), - patch.object(manager, "_create_mcp_client", AsyncMock(return_value=object())), + patch.object(manager, "create_mcp_client", AsyncMock(return_value=object())), patch.object(manager, "_fetch_tools_with_timeout", AsyncMock(return_value=upstream)), ): try: diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_toolset_scope.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_toolset_scope.py index 95be2b8b12b..097b0924b8d 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_toolset_scope.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_toolset_scope.py @@ -50,7 +50,7 @@ class TestApplyToolsetScope: @pytest.mark.asyncio async def test_restricts_to_toolset_servers_and_tools(self): - from litellm.proxy._experimental.mcp_server.server import _apply_toolset_scope + from litellm.proxy._experimental.mcp_server.server import apply_toolset_scope toolset_perms = { "server-a": ["tool1", "tool2"], @@ -66,7 +66,7 @@ class TestApplyToolsetScope: mcp_servers=["server-a", "server-b", "server-c"], mcp_toolsets=["toolset-123"], ) - result = await _apply_toolset_scope(auth, "toolset-123") + result = await apply_toolset_scope(auth, "toolset-123") op = result.object_permission assert op is not None @@ -92,7 +92,7 @@ class TestApplyToolsetScope: @pytest.mark.asyncio async def test_admin_creates_object_permission_when_none(self): """Admin key with object_permission=None can access any toolset.""" - from litellm.proxy._experimental.mcp_server.server import _apply_toolset_scope + from litellm.proxy._experimental.mcp_server.server import apply_toolset_scope toolset_perms = {"server-a": ["tool1"]} with patch( @@ -105,7 +105,7 @@ class TestApplyToolsetScope: user_role=LitellmUserRoles.PROXY_ADMIN, object_permission=None, ) - result = await _apply_toolset_scope(auth, "toolset-123") + result = await apply_toolset_scope(auth, "toolset-123") op = result.object_permission assert op is not None @@ -116,7 +116,7 @@ class TestApplyToolsetScope: async def test_team_granted_toolset_is_served_to_a_key_without_its_own_grant(self): """A team key whose own row carries no toolset grant is admitted to the toolset its team holds (LIT-6029), scoped to that toolset's servers and tools.""" - from litellm.proxy._experimental.mcp_server.server import _apply_toolset_scope + from litellm.proxy._experimental.mcp_server.server import apply_toolset_scope toolset_perms = {"server-a": ["tool1"]} auth = UserAPIKeyAuth(api_key="sk-test", team_id="team-a", object_permission=None) @@ -125,7 +125,7 @@ class TestApplyToolsetScope: "global_mcp_server_manager.resolve_toolset_tool_permissions", new=AsyncMock(return_value=toolset_perms), ): - result = await _apply_toolset_scope(auth, "toolset-123", granted=_granted_through_team("toolset-123")) + result = await apply_toolset_scope(auth, "toolset-123", granted=_granted_through_team("toolset-123")) assert result.mcp_toolset_id == "toolset-123" assert result.object_permission is not None @@ -138,7 +138,7 @@ class TestApplyToolsetScope: team-granted toolset is not capped by the user's own row: the row stays intact and the toolset rides along as mcp_toolset_id (LIT-6029).""" from litellm.constants import UI_SESSION_TOKEN_TEAM_ID - from litellm.proxy._experimental.mcp_server.server import _apply_toolset_scope + from litellm.proxy._experimental.mcp_server.server import apply_toolset_scope session = UserAPIKeyAuth(team_id=UI_SESSION_TOKEN_TEAM_ID, user_id="user-1") own_row = LiteLLM_ObjectPermissionTable(object_permission_id="user-op", mcp_servers=["server-own"]) @@ -151,7 +151,7 @@ class TestApplyToolsetScope: "global_mcp_server_manager.resolve_toolset_tool_permissions", new=resolve, ): - result = await _apply_toolset_scope( + result = await apply_toolset_scope( session, "toolset-123", acting_user=AsyncMock(return_value=admitted), granted=granted ) @@ -163,13 +163,13 @@ class TestApplyToolsetScope: @pytest.mark.asyncio async def test_a_gateway_admitted_user_without_the_toolset_in_any_source_is_denied(self): - from litellm.proxy._experimental.mcp_server.server import _apply_toolset_scope + from litellm.proxy._experimental.mcp_server.server import apply_toolset_scope admitted = UserAPIKeyAuth(user_id="user-1", object_permission=None) admitted.mcp_admitted_user_subject = True granted = AsyncMock(return_value=frozenset({"toolset-other"})) with pytest.raises(HTTPException) as exc_info: - await _apply_toolset_scope(admitted, "toolset-123", granted=granted) + await apply_toolset_scope(admitted, "toolset-123", granted=granted) assert exc_info.value.status_code == 403 granted.assert_awaited_once_with(admitted) @@ -178,7 +178,7 @@ class TestApplyToolsetScope: async def test_a_resource_scoped_admitted_user_is_denied_a_team_toolset_on_another_server(self): """A gateway bearer scoped to server-own (RFC 8707 resource) cannot open a team toolset whose servers lie outside that resource, even though the team grants it (Devin Review 4150024267).""" - from litellm.proxy._experimental.mcp_server.server import _apply_toolset_scope + from litellm.proxy._experimental.mcp_server.server import apply_toolset_scope admitted = UserAPIKeyAuth(user_id="user-1", object_permission=None) admitted.mcp_admitted_user_subject = True @@ -192,14 +192,14 @@ class TestApplyToolsetScope: new=resolve, ): with pytest.raises(HTTPException) as exc_info: - await _apply_toolset_scope(admitted, "toolset-123", granted=granted) + await apply_toolset_scope(admitted, "toolset-123", granted=granted) assert exc_info.value.status_code == 403 resolve.assert_awaited_once_with(toolset_ids=["toolset-123"], requires_fresh_policy=True) @pytest.mark.asyncio async def test_a_resource_scoped_admitted_user_opens_a_toolset_inside_its_resource(self): - from litellm.proxy._experimental.mcp_server.server import _apply_toolset_scope + from litellm.proxy._experimental.mcp_server.server import apply_toolset_scope admitted = UserAPIKeyAuth(user_id="user-1", object_permission=None) admitted.mcp_admitted_user_subject = True @@ -211,7 +211,7 @@ class TestApplyToolsetScope: "global_mcp_server_manager.resolve_toolset_tool_permissions", new=resolve, ): - result = await _apply_toolset_scope(admitted, "toolset-123", granted=granted) + result = await apply_toolset_scope(admitted, "toolset-123", granted=granted) assert result.mcp_toolset_id == "toolset-123" assert result.mcp_session_resource_server_id == "server-team" @@ -219,12 +219,12 @@ class TestApplyToolsetScope: @pytest.mark.asyncio async def test_team_grant_for_another_toolset_does_not_admit_a_key_to_this_one(self): - from litellm.proxy._experimental.mcp_server.server import _apply_toolset_scope + from litellm.proxy._experimental.mcp_server.server import apply_toolset_scope auth = _make_auth(mcp_toolsets=[]) auth.team_id = "team-a" with pytest.raises(HTTPException) as exc_info: - await _apply_toolset_scope(auth, "toolset-123", granted=_granted_through_team("toolset-other")) + await apply_toolset_scope(auth, "toolset-123", granted=_granted_through_team("toolset-other")) assert exc_info.value.status_code == 403 @@ -233,11 +233,11 @@ class TestApplyToolsetScope: """Non-admin key with object_permission=None is denied (no grants configured).""" from starlette.exceptions import HTTPException - from litellm.proxy._experimental.mcp_server.server import _apply_toolset_scope + from litellm.proxy._experimental.mcp_server.server import apply_toolset_scope auth = UserAPIKeyAuth(api_key="sk-test", object_permission=None) with pytest.raises(HTTPException) as exc_info: - await _apply_toolset_scope(auth, "toolset-123") + await apply_toolset_scope(auth, "toolset-123") assert exc_info.value.status_code == 403 @pytest.mark.asyncio @@ -248,7 +248,7 @@ class TestApplyToolsetScope: toolset path, which replaces mcp_servers and would drop the sentinel.""" from starlette.exceptions import HTTPException - from litellm.proxy._experimental.mcp_server.server import _apply_toolset_scope + from litellm.proxy._experimental.mcp_server.server import apply_toolset_scope op = LiteLLM_ObjectPermissionTable( object_permission_id="test", @@ -266,7 +266,7 @@ class TestApplyToolsetScope: new=resolve, ): with pytest.raises(HTTPException) as exc_info: - await _apply_toolset_scope(auth, "toolset-123") + await apply_toolset_scope(auth, "toolset-123") assert exc_info.value.status_code == 403 resolve.assert_not_awaited() @@ -786,17 +786,17 @@ class TestMCPActiveToolsetContextVar: """Tests for _mcp_active_toolset_id ContextVar — clients cannot inject it.""" def test_contextvar_default_is_none(self): - from litellm.proxy._experimental.mcp_server.server import _mcp_active_toolset_id + from litellm.proxy._experimental.mcp_server.server import mcp_active_toolset_id - assert _mcp_active_toolset_id.get() is None + assert mcp_active_toolset_id.get() is None def test_contextvar_set_and_reset(self): - from litellm.proxy._experimental.mcp_server.server import _mcp_active_toolset_id + from litellm.proxy._experimental.mcp_server.server import mcp_active_toolset_id - token = _mcp_active_toolset_id.set("toolset-abc") - assert _mcp_active_toolset_id.get() == "toolset-abc" - _mcp_active_toolset_id.reset(token) - assert _mcp_active_toolset_id.get() is None + token = mcp_active_toolset_id.set("toolset-abc") + assert mcp_active_toolset_id.get() == "toolset-abc" + mcp_active_toolset_id.reset(token) + assert mcp_active_toolset_id.get() is None @pytest.mark.asyncio async def test_client_header_is_stripped_in_scope(self): diff --git a/tests/unit/proxy/_experimental/mcp_server/test_oauth2_token_cache.py b/tests/unit/proxy/_experimental/mcp_server/test_oauth2_token_cache.py index 30d0f17a099..6eb93b090e9 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_oauth2_token_cache.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_oauth2_token_cache.py @@ -237,44 +237,44 @@ def test_storage_ttl_capped_at_token_lifetime(): while the stored refresh_token sat unused because refresh only runs on the DB read-through.""" from litellm.constants import MCP_PER_USER_TOKEN_EXPIRY_BUFFER_SECONDS from litellm.proxy._experimental.mcp_server.oauth2_token_cache import ( - _compute_per_user_token_ttl, + compute_per_user_token_ttl, ) server = _server(oauth2_flow=None, token_storage_ttl_seconds=604800) - assert _compute_per_user_token_ttl(server, expires_in=86400) == 86400 - MCP_PER_USER_TOKEN_EXPIRY_BUFFER_SECONDS + assert compute_per_user_token_ttl(server, expires_in=86400) == 86400 - MCP_PER_USER_TOKEN_EXPIRY_BUFFER_SECONDS def test_storage_ttl_shorter_than_token_lifetime_wins(): """A configured TTL below the token lifetime is the operative value: the knob's purpose is to force earlier DB re-checks (staleness backstop), so the shorter side must win the min().""" from litellm.proxy._experimental.mcp_server.oauth2_token_cache import ( - _compute_per_user_token_ttl, + compute_per_user_token_ttl, ) server = _server(oauth2_flow=None, token_storage_ttl_seconds=3600) - assert _compute_per_user_token_ttl(server, expires_in=86400) == 3600 + assert compute_per_user_token_ttl(server, expires_in=86400) == 3600 def test_storage_ttl_verbatim_when_token_lifetime_unknown(): """With no expires_in from the upstream there is nothing to cap against, so the configured TTL applies as-is (matching the pre-cap behavior for lifetime-less tokens).""" from litellm.proxy._experimental.mcp_server.oauth2_token_cache import ( - _compute_per_user_token_ttl, + compute_per_user_token_ttl, ) server = _server(oauth2_flow=None, token_storage_ttl_seconds=604800) - assert _compute_per_user_token_ttl(server, expires_in=None) == 604800 + assert compute_per_user_token_ttl(server, expires_in=None) == 604800 def test_storage_ttl_floors_at_one_second_for_nearly_dead_token(): """A token already inside the expiry buffer yields the 1-second floor, not zero or a negative TTL, mirroring the floor the default (unconfigured) path has always had.""" from litellm.proxy._experimental.mcp_server.oauth2_token_cache import ( - _compute_per_user_token_ttl, + compute_per_user_token_ttl, ) server = _server(oauth2_flow=None, token_storage_ttl_seconds=3600) - assert _compute_per_user_token_ttl(server, expires_in=30) == 1 + assert compute_per_user_token_ttl(server, expires_in=30) == 1 def test_default_ttl_paths_unchanged_without_storage_ttl(): @@ -285,12 +285,12 @@ def test_default_ttl_paths_unchanged_without_storage_ttl(): MCP_PER_USER_TOKEN_EXPIRY_BUFFER_SECONDS, ) from litellm.proxy._experimental.mcp_server.oauth2_token_cache import ( - _compute_per_user_token_ttl, + compute_per_user_token_ttl, ) server = _server(oauth2_flow=None) - assert _compute_per_user_token_ttl(server, expires_in=86400) == 86400 - MCP_PER_USER_TOKEN_EXPIRY_BUFFER_SECONDS - assert _compute_per_user_token_ttl(server, expires_in=None) == MCP_PER_USER_TOKEN_DEFAULT_TTL + assert compute_per_user_token_ttl(server, expires_in=86400) == 86400 - MCP_PER_USER_TOKEN_EXPIRY_BUFFER_SECONDS + assert compute_per_user_token_ttl(server, expires_in=None) == MCP_PER_USER_TOKEN_DEFAULT_TTL @pytest.mark.asyncio diff --git a/tests/unit/proxy/_experimental/mcp_server/test_openapi_to_mcp_generator.py b/tests/unit/proxy/_experimental/mcp_server/test_openapi_to_mcp_generator.py index 6b0211c3866..9f3db632db1 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_openapi_to_mcp_generator.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_openapi_to_mcp_generator.py @@ -20,9 +20,9 @@ from respx import MockRouter from litellm.types.mcp import MCPAuth, MCPAuthType from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( - _request_auth_header, - _request_extra_headers, - _request_resolved_auth_headers, + request_auth_header, + request_extra_headers, + request_resolved_auth_headers, _request_upstream_url, _resolve_param_list, _resolve_ref, @@ -99,7 +99,7 @@ async def test_authorization_validates_credentials_before_http( "/echo", "get", {}, "https://upstream.example", auth_type=auth_type, ) destination: Final = respx_mock.get("https://upstream.example/echo").respond(200, text="authenticated") - caller_token: Final = _request_auth_header.set(value) + caller_token: Final = request_auth_header.set(value) try: if accepted: assert await tool() == TextResult("authenticated") @@ -111,7 +111,7 @@ async def test_authorization_validates_credentials_before_http( assert exc.value.status_code == 500 assert destination.call_count == 0 finally: - _request_auth_header.reset(caller_token) + request_auth_header.reset(caller_token) @pytest.mark.asyncio @@ -132,9 +132,9 @@ async def test_static_auth_validates_headers_after_existing_precedence( ) monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") destination: Final = respx_mock.get("https://upstream.example/echo").respond(200, text="authenticated") - caller_token: Final = _request_auth_header.set(caller) - extra_token: Final = _request_extra_headers.set(forwarded) - resolved_token: Final = _request_resolved_auth_headers.set(resolved) + caller_token: Final = request_auth_header.set(caller) + extra_token: Final = request_extra_headers.set(forwarded) + resolved_token: Final = request_resolved_auth_headers.set(resolved) try: if expected is None: with pytest.raises(HTTPException, match="requires a usable upstream credential") as exc: @@ -146,9 +146,9 @@ async def test_static_auth_validates_headers_after_existing_precedence( assert destination.call_count == 1 assert destination.calls.last.request.headers["authorization"] == expected finally: - _request_auth_header.reset(caller_token) - _request_extra_headers.reset(extra_token) - _request_resolved_auth_headers.reset(resolved_token) + request_auth_header.reset(caller_token) + request_extra_headers.reset(extra_token) + request_resolved_auth_headers.reset(resolved_token) @pytest.mark.asyncio @@ -204,13 +204,13 @@ async def test_static_validation_preserves_no_auth_and_resolved_oauth( monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") tool: Final = create_tool_function("/echo", "get", {}, "https://upstream.example", auth_type=auth_type) destination: Final = respx_mock.get("https://upstream.example/echo").respond(200, text="echo") - token: Final = _request_resolved_auth_headers.set(resolved) + token: Final = request_resolved_auth_headers.set(resolved) try: assert await tool() == TextResult("echo") assert destination.call_count == 1 assert destination.calls.last.request.headers.get("authorization") == (resolved or {}).get("Authorization") finally: - _request_resolved_auth_headers.reset(token) + request_resolved_auth_headers.reset(token) def _create_mock_client(method: str, response_text: str, status_code: int = 200) -> AsyncMock: @@ -1214,11 +1214,11 @@ class TestRequestExtraHeaders: async_client = _create_mock_client("get", "ok") mock_client.return_value = async_client - token = _request_extra_headers.set({"X-TOKEN": "secret-value"}) + token = request_extra_headers.set({"X-TOKEN": "secret-value"}) try: result = await func() finally: - _request_extra_headers.reset(token) + request_extra_headers.reset(token) assert result == TextResult("ok") call_args = async_client.get.call_args @@ -1265,11 +1265,11 @@ class TestRequestExtraHeaders: async_client = _create_mock_client("post", "created") mock_client.return_value = async_client - token = _request_extra_headers.set({"X-TOKEN": "dynamic-value"}) + token = request_extra_headers.set({"X-TOKEN": "dynamic-value"}) try: result = await func() finally: - _request_extra_headers.reset(token) + request_extra_headers.reset(token) assert result == TextResult("created") call_args = async_client.post.call_args @@ -1293,11 +1293,11 @@ class TestRequestExtraHeaders: async_client = _create_mock_client("get", "ok") mock_client.return_value = async_client - token = _request_extra_headers.set({"X-Tenant": "caller-spoofed"}) + token = request_extra_headers.set({"X-Tenant": "caller-spoofed"}) try: result = await func() finally: - _request_extra_headers.reset(token) + request_extra_headers.reset(token) assert result == TextResult("ok") call_args = async_client.get.call_args @@ -1321,11 +1321,11 @@ class TestRequestExtraHeaders: async_client = _create_mock_client("get", "ok") mock_client.return_value = async_client - token = _request_extra_headers.set({"x-tenant": "caller-spoofed"}) + token = request_extra_headers.set({"x-tenant": "caller-spoofed"}) try: result = await func() finally: - _request_extra_headers.reset(token) + request_extra_headers.reset(token) assert result == TextResult("ok") call_args = async_client.get.call_args @@ -1349,15 +1349,15 @@ class TestRequestExtraHeaders: async_client = _create_mock_client("get", "secure-data") mock_client.return_value = async_client - extra_token = _request_extra_headers.set( + extra_token = request_extra_headers.set( {"Authorization": "Bearer extra", "X-TOKEN": "token-value"} ) - auth_token = _request_auth_header.set("Bearer byok-credential") + auth_token = request_auth_header.set("Bearer byok-credential") try: result = await func() finally: - _request_auth_header.reset(auth_token) - _request_extra_headers.reset(extra_token) + request_auth_header.reset(auth_token) + request_extra_headers.reset(extra_token) assert result == TextResult("secure-data") call_args = async_client.get.call_args @@ -1380,8 +1380,8 @@ class TestRequestExtraHeaders: async_client = _create_mock_client("get", "ok") mock_client.return_value = async_client - token = _request_extra_headers.set({"X-TOKEN": "first-call"}) - _request_extra_headers.reset(token) + token = request_extra_headers.set({"X-TOKEN": "first-call"}) + request_extra_headers.reset(token) await func() @@ -1409,15 +1409,15 @@ class TestRequestExtraHeaders: async_client = _create_mock_client("get", "secure-data") mock_client.return_value = async_client - extra_token = _request_extra_headers.set({"Authorization": "Bearer caller-forwarded"}) - auth_token = _request_auth_header.set("Bearer byok-credential") - resolved_token = _request_resolved_auth_headers.set({"Authorization": "Bearer resolved-oauth"}) + extra_token = request_extra_headers.set({"Authorization": "Bearer caller-forwarded"}) + auth_token = request_auth_header.set("Bearer byok-credential") + resolved_token = request_resolved_auth_headers.set({"Authorization": "Bearer resolved-oauth"}) try: result = await func() finally: - _request_auth_header.reset(auth_token) - _request_extra_headers.reset(extra_token) - _request_resolved_auth_headers.reset(resolved_token) + request_auth_header.reset(auth_token) + request_extra_headers.reset(extra_token) + request_resolved_auth_headers.reset(resolved_token) assert result == TextResult("secure-data") headers_sent = async_client.get.call_args[1]["headers"] @@ -1439,8 +1439,8 @@ class TestRequestExtraHeaders: async_client = _create_mock_client("get", "ok") mock_client.return_value = async_client - token = _request_resolved_auth_headers.set({"Authorization": "Bearer resolved-oauth"}) - _request_resolved_auth_headers.reset(token) + token = request_resolved_auth_headers.set({"Authorization": "Bearer resolved-oauth"}) + request_resolved_auth_headers.reset(token) await func() diff --git a/tests/unit/proxy/_experimental/mcp_server/test_openapi_tool_auth.py b/tests/unit/proxy/_experimental/mcp_server/test_openapi_tool_auth.py index c570d498f44..ffa25d63391 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_openapi_tool_auth.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_openapi_tool_auth.py @@ -56,7 +56,7 @@ async def test_openapi_local_tool_runs_pre_call_tool_check(): with ( patch.object( mcp_operations.global_mcp_server_manager, - "_get_mcp_server_from_tool_name", + "get_mcp_server_from_tool_name", return_value=fake_server, ), patch.object( @@ -145,7 +145,7 @@ async def test_openapi_local_tool_blocked_when_pre_call_check_raises(): with ( patch.object( mcp_operations.global_mcp_server_manager, - "_get_mcp_server_from_tool_name", + "get_mcp_server_from_tool_name", return_value=fake_server, ), patch.object( @@ -207,10 +207,10 @@ async def test_openapi_local_tool_denied_when_server_not_resolvable(): # `_get_mcp_server_from_tool_name` returns None — no server context. with ( - patch.object(mcp_operations, "_resolve_openapi_tool_auth", new=resolve_auth), + patch.object(mcp_operations, "resolve_openapi_tool_auth", new=resolve_auth), patch.object( mcp_operations.global_mcp_server_manager, - "_get_mcp_server_from_tool_name", + "get_mcp_server_from_tool_name", return_value=None, ), patch.object( @@ -258,7 +258,7 @@ async def test_openapi_local_tool_injects_resolved_oauth_token(): reads. Kills the mutant that deletes the resolve_openapi_upstream_auth call in server.py.""" from litellm.proxy._experimental.mcp_server import server as mcp_module from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( - _request_resolved_auth_headers, + request_resolved_auth_headers, ) from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import ( StaticHeaderAuth, @@ -290,13 +290,13 @@ async def test_openapi_local_tool_injects_resolved_oauth_token(): captured: dict = {} async def handle_local(_name, _arguments, _wire_compat): - captured["resolved"] = _request_resolved_auth_headers.get() + captured["resolved"] = request_resolved_auth_headers.get() return CallToolResult(content=[], is_error=False) with ( patch.object( mcp_operations.global_mcp_server_manager, - "_get_mcp_server_from_tool_name", + "get_mcp_server_from_tool_name", return_value=oauth_server, ), patch.object( @@ -332,7 +332,7 @@ async def test_openapi_local_tool_injects_resolved_oauth_token(): ) assert captured["resolved"] == {"Authorization": "Bearer stored-user-token"} - assert _request_resolved_auth_headers.get() is None + assert request_resolved_auth_headers.get() is None @@ -605,7 +605,7 @@ async def test_per_server_auth_header_reaches_both_openapi_dispatch_arms(dispatc """ from litellm.proxy._experimental.mcp_server import server as mcp_module from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( - _request_auth_header, + request_auth_header, ) server = _spec_path_server() @@ -618,11 +618,11 @@ async def test_per_server_auth_header_reaches_both_openapi_dispatch_arms(dispatc return None, kwargs["forwarded_headers"] async def capture_local(_name, _arguments, _wire_compat): - captured["injected"] = _request_auth_header.get() + captured["injected"] = request_auth_header.get() return CallToolResult(content=[], is_error=False) async def capture_openapi_handler(_server, _name, _arguments, _wire_compat): - captured["injected"] = _request_auth_header.get() + captured["injected"] = request_auth_header.get() return CallToolResult(content=[], is_error=False) manager = mcp_operations.global_mcp_server_manager @@ -637,7 +637,7 @@ async def test_per_server_auth_header_reaches_both_openapi_dispatch_arms(dispatc fake_tool.input_schema = {"type": "object"} fake_tool.server_id = server.server_id with ( - patch.object(manager, "_get_mcp_server_from_tool_name", return_value=server), + patch.object(manager, "get_mcp_server_from_tool_name", return_value=server), patch.object(mcp_operations.global_mcp_tool_registry, "get_tool", return_value=fake_tool), patch( "litellm.proxy._experimental.mcp_server.operations._handle_local_mcp_tool", @@ -671,7 +671,7 @@ async def test_per_server_auth_header_reaches_both_openapi_dispatch_arms(dispatc assert captured["resolver_credential"] == {"Authorization": OPENAPI_PER_SERVER_TOKEN} assert captured["injected"] == OPENAPI_PER_SERVER_TOKEN - assert _request_auth_header.get() is None + assert request_auth_header.get() is None @pytest.mark.parametrize("failure", ["auth", "other"]) @@ -723,7 +723,7 @@ async def test_local_dispatch_reports_the_outcome_instead_of_success(failure: st user = UserAPIKeyAuth(api_key="sk-user", user_id="alice", user_role=LitellmUserRoles.INTERNAL_USER.value) with ( - patch.object(mcp_operations.global_mcp_server_manager, "_get_mcp_server_from_tool_name", return_value=server), + patch.object(mcp_operations.global_mcp_server_manager, "get_mcp_server_from_tool_name", return_value=server), patch.object(mcp_operations.global_mcp_server_manager, "pre_call_tool_check", new=AsyncMock(return_value={})), patch.object(mcp_operations.global_mcp_tool_registry, "get_tool", return_value=fake_tool), patch.object( @@ -798,16 +798,16 @@ def test_the_openapi_arm_installs_the_guard_when_a_credential_rides_a_custom_slo hook alone passes even if this arm never installs it. """ from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( - _request_resolved_auth_headers, + request_resolved_auth_headers, _upstream_client, ) - token = _request_resolved_auth_headers.set({"esb-oauth": "Bearer minted"}) + token = request_resolved_auth_headers.set({"esb-oauth": "Bearer minted"}) try: client = _upstream_client() assert client.client.event_hooks["request"], "custom slot must install a redirect guard" finally: - _request_resolved_auth_headers.reset(token) + request_resolved_auth_headers.reset(token) def test_the_guarded_client_is_reused_rather_than_built_per_call(): @@ -816,15 +816,15 @@ def test_the_guarded_client_is_reused_rather_than_built_per_call(): have to come from the shared cache. """ from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( - _request_resolved_auth_headers, + request_resolved_auth_headers, _upstream_client, ) - token = _request_resolved_auth_headers.set({"esb-oauth": "Bearer minted"}) + token = request_resolved_auth_headers.set({"esb-oauth": "Bearer minted"}) try: assert _upstream_client() is _upstream_client() finally: - _request_resolved_auth_headers.reset(token) + request_resolved_auth_headers.reset(token) @pytest.mark.asyncio @@ -836,11 +836,11 @@ async def test_the_shared_guard_reads_the_url_from_the_request_context(): from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( _drop_credential_across_origin, - _request_resolved_auth_headers, + request_resolved_auth_headers, _request_upstream_url, ) - creds = _request_resolved_auth_headers.set({"esb-oauth": "Bearer minted"}) + creds = request_resolved_auth_headers.set({"esb-oauth": "Bearer minted"}) url = _request_upstream_url.set("https://api.example.com/v1/things") try: same = httpx.Request("POST", "https://api.example.com/v1/other", headers={"esb-oauth": "Bearer m"}) @@ -852,7 +852,7 @@ async def test_the_shared_guard_reads_the_url_from_the_request_context(): assert "esb-oauth" not in foreign.headers finally: _request_upstream_url.reset(url) - _request_resolved_auth_headers.reset(creds) + request_resolved_auth_headers.reset(creds) @pytest.mark.parametrize("resolved", [{"Authorization": "Bearer minted"}, {}, None]) @@ -860,16 +860,16 @@ def test_the_openapi_arm_keeps_the_shared_client_when_no_guard_is_needed(resolve # Authorization is already stripped across origins by the HTTP client, so taking the guarded # path for it would give up the shared connection pool for nothing. from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( - _request_resolved_auth_headers, + request_resolved_auth_headers, _upstream_client, ) - token = _request_resolved_auth_headers.set(resolved) + token = request_resolved_auth_headers.set(resolved) try: client = _upstream_client() assert not client.client.event_hooks.get("request") finally: - _request_resolved_auth_headers.reset(token) + request_resolved_auth_headers.reset(token) @pytest.mark.asyncio diff --git a/tests/unit/proxy/_experimental/mcp_server/test_operations.py b/tests/unit/proxy/_experimental/mcp_server/test_operations.py index 089426eb5cb..d9997ec0fce 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_operations.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_operations.py @@ -22,7 +22,7 @@ from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError from litellm.proxy.hooks.parallel_request_limiter_v3 import ( - _PROXY_MaxParallelRequestsHandler_v3, + PROXY_MaxParallelRequestsHandler_v3, ) from litellm.proxy.utils import InternalUsageCache, ProxyLogging, hash_token from litellm.types.mcp import MCPAuth, MCPTransport @@ -294,7 +294,7 @@ def _catalog_case(method): def _mcp_rate_limited_proxy_logging() -> ProxyLogging: proxy_logging: Final = ProxyLogging(user_api_key_cache=UserApiKeyCache()) - proxy_logging.proxy_hook_mapping["parallel_request_limiter"] = _PROXY_MaxParallelRequestsHandler_v3( + proxy_logging.proxy_hook_mapping["parallel_request_limiter"] = PROXY_MaxParallelRequestsHandler_v3( internal_usage_cache=InternalUsageCache(DualCache()) ) return proxy_logging @@ -324,7 +324,7 @@ async def test_mcp_server_rpm_limits_every_catalog_operation(operation: str) -> ) caller: Final = UserAPIKeyAuth(api_key=hash_token("sk-catalog-rpm")) operation_to_manager_method: Final = { - "tools/list": "_get_tools_from_server", + "tools/list": "get_tools_from_server", "prompts/list": "get_prompts_from_server", "resources/list": "get_resources_from_server", "resources/templates/list": "get_resource_templates_from_server", @@ -433,7 +433,7 @@ async def test_tools_call_warmup_does_not_consume_mcp_server_rpm() -> None: patch.object(operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[server])), patch.object(operations.global_mcp_server_manager, "server_exposes_tool", return_value=False), patch("litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging), - patch.object(operations.global_mcp_server_manager, "_get_tools_from_server", upstream), + patch.object(operations.global_mcp_server_manager, "get_tools_from_server", upstream), ): await operations._list_tools_before_first_call( server=server, @@ -502,8 +502,8 @@ async def test_tools_call_pre_call_hook_rejection_does_not_enforce_mcp_server_rp ) rate_limit_error: Final = ProxyRateLimitError(detail="ordinary key rate limit") proxy_logging: Final = MagicMock() - proxy_logging._create_mcp_request_object_from_kwargs.return_value = {} - proxy_logging._convert_mcp_to_llm_format.return_value = {} + proxy_logging.create_mcp_request_object_from_kwargs.return_value = {} + proxy_logging.convert_mcp_to_llm_format.return_value = {} proxy_logging.pre_call_hook = AsyncMock(side_effect=rate_limit_error) proxy_logging.enforce_mcp_server_rate_limits = AsyncMock() @@ -714,7 +714,7 @@ async def test_tool_listing_returns_empty_result_without_dispatch_for_unavailabl upstream = AsyncMock() with ( patch.object(operations, "_get_allowed_mcp_servers", allowed), - patch.object(operations.global_mcp_server_manager, "_get_tools_from_server", upstream), + patch.object(operations.global_mcp_server_manager, "get_tools_from_server", upstream), ): result = await GatewayOperations().execute(ListToolsRequest(), prepare_context()) assert result.tools == [] @@ -1142,7 +1142,7 @@ async def test_list_mcp_tools_records_the_catalog_only_when_asked( upstream = [MCPTool(name="echo", description="Echo text back", inputSchema={"type": "object"})] with ( patch.object(operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[server])), - patch.object(manager, "_create_mcp_client", AsyncMock(return_value=object())), + patch.object(manager, "create_mcp_client", AsyncMock(return_value=object())), patch.object(manager, "_fetch_tools_with_timeout", AsyncMock(return_value=upstream)), patch.dict(manager.tool_name_to_mcp_server_name_mapping), ): diff --git a/tests/unit/proxy/_experimental/mcp_server/test_per_user_oauth_cache.py b/tests/unit/proxy/_experimental/mcp_server/test_per_user_oauth_cache.py index ac453df8fa5..efddd1beef3 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_per_user_oauth_cache.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_per_user_oauth_cache.py @@ -21,7 +21,7 @@ for _mod in ("orjson",): from litellm.proxy._experimental.mcp_server.oauth2_token_cache import ( # noqa: E402 MCPPerUserTokenCache, - _compute_per_user_token_ttl, + compute_per_user_token_ttl, mcp_per_user_token_cache, ) from litellm.types.mcp import MCPAuth, MCPTransport # noqa: E402 @@ -215,25 +215,25 @@ class TestValidateTokenResponse: class TestComputePerUserTokenTtl: def test_uses_server_override_when_set(self): server = _make_server(token_storage_ttl_seconds=7200) - assert _compute_per_user_token_ttl(server, expires_in=99999) == 7200 + assert compute_per_user_token_ttl(server, expires_in=99999) == 7200 def test_uses_expires_in_minus_buffer(self): server = _make_server() # Default buffer is 60s - ttl = _compute_per_user_token_ttl(server, expires_in=3600) + ttl = compute_per_user_token_ttl(server, expires_in=3600) assert ttl == 3600 - 60 def test_minimum_ttl_is_1(self): server = _make_server() # expires_in smaller than buffer → clamp to 1 - ttl = _compute_per_user_token_ttl(server, expires_in=30) + ttl = compute_per_user_token_ttl(server, expires_in=30) assert ttl == 1 def test_default_ttl_when_expires_in_none(self): from litellm.constants import MCP_PER_USER_TOKEN_DEFAULT_TTL server = _make_server() - ttl = _compute_per_user_token_ttl(server, expires_in=None) + ttl = compute_per_user_token_ttl(server, expires_in=None) assert ttl == MCP_PER_USER_TOKEN_DEFAULT_TTL diff --git a/tests/unit/proxy/_experimental/mcp_server/test_rest_endpoints.py b/tests/unit/proxy/_experimental/mcp_server/test_rest_endpoints.py index c748bdc6a6b..791e6d18ea8 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_rest_endpoints.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_rest_endpoints.py @@ -218,7 +218,7 @@ class TestExecuteWithMcpClient: monkeypatch.setattr( rest_endpoints.global_mcp_server_manager, - "_create_mcp_client", + "create_mcp_client", fake_create_client, ) @@ -243,7 +243,7 @@ class TestExecuteWithMcpClient: monkeypatch.setattr( rest_endpoints.global_mcp_server_manager, - "_create_mcp_client", + "create_mcp_client", fake_create_client, ) @@ -271,7 +271,7 @@ class TestExecuteWithMcpClient: monkeypatch.setattr( rest_endpoints.global_mcp_server_manager, - "_create_mcp_client", + "create_mcp_client", hanging_create_client, ) @@ -350,13 +350,13 @@ class TestExecuteWithMcpClient: monkeypatch.setattr( rest_endpoints.global_mcp_server_manager, - "_build_stdio_env", + "build_stdio_env", fake_build_stdio_env, raising=False, ) monkeypatch.setattr( rest_endpoints.global_mcp_server_manager, - "_create_mcp_client", + "create_mcp_client", fake_create_client, raising=False, ) @@ -400,13 +400,13 @@ class TestExecuteWithMcpClient: monkeypatch.setattr( rest_endpoints.global_mcp_server_manager, - "_build_stdio_env", + "build_stdio_env", fake_build_stdio_env, raising=False, ) monkeypatch.setattr( rest_endpoints.global_mcp_server_manager, - "_create_mcp_client", + "create_mcp_client", fake_create_client, raising=False, ) @@ -451,13 +451,13 @@ class TestExecuteWithMcpClient: monkeypatch.setattr( rest_endpoints.global_mcp_server_manager, - "_build_stdio_env", + "build_stdio_env", fake_build_stdio_env, raising=False, ) monkeypatch.setattr( rest_endpoints.global_mcp_server_manager, - "_create_mcp_client", + "create_mcp_client", fake_create_client, raising=False, ) @@ -492,13 +492,13 @@ class TestExecuteWithMcpClient: monkeypatch.setattr( rest_endpoints.global_mcp_server_manager, - "_build_stdio_env", + "build_stdio_env", fake_build_stdio_env, raising=False, ) monkeypatch.setattr( rest_endpoints.global_mcp_server_manager, - "_create_mcp_client", + "create_mcp_client", fake_create_client, raising=False, ) @@ -546,13 +546,13 @@ class TestExecuteWithMcpClient: monkeypatch.setattr( rest_endpoints.global_mcp_server_manager, - "_build_stdio_env", + "build_stdio_env", fake_build_stdio_env, raising=False, ) monkeypatch.setattr( rest_endpoints.global_mcp_server_manager, - "_create_mcp_client", + "create_mcp_client", fake_create_client, raising=False, ) @@ -600,13 +600,13 @@ class TestExecuteWithMcpClient: monkeypatch.setattr( rest_endpoints.global_mcp_server_manager, - "_build_stdio_env", + "build_stdio_env", fake_build_stdio_env, raising=False, ) monkeypatch.setattr( rest_endpoints.global_mcp_server_manager, - "_create_mcp_client", + "create_mcp_client", fake_create_client, raising=False, ) @@ -648,13 +648,13 @@ class TestExecuteWithMcpClient: monkeypatch.setattr( rest_endpoints.global_mcp_server_manager, - "_build_stdio_env", + "build_stdio_env", fake_build_stdio_env, raising=False, ) monkeypatch.setattr( rest_endpoints.global_mcp_server_manager, - "_create_mcp_client", + "create_mcp_client", fake_create_client, raising=False, ) @@ -692,13 +692,13 @@ class TestExecuteWithMcpClient: monkeypatch.setattr( rest_endpoints.global_mcp_server_manager, - "_build_stdio_env", + "build_stdio_env", fake_build_stdio_env, raising=False, ) monkeypatch.setattr( rest_endpoints.global_mcp_server_manager, - "_create_mcp_client", + "create_mcp_client", fake_create_client, raising=False, ) @@ -900,7 +900,7 @@ class TestTestToolsList: monkeypatch.setattr( auth_mcp.MCPRequestHandler, - "_get_oauth2_headers_from_headers", + "get_oauth2_headers_from_headers", staticmethod(fake_oauth), raising=False, ) @@ -1082,7 +1082,7 @@ class TestTestToolsList: monkeypatch.setattr( auth_mcp.MCPRequestHandler, - "_get_oauth2_headers_from_headers", + "get_oauth2_headers_from_headers", staticmethod(fake_oauth), raising=False, ) @@ -1136,7 +1136,7 @@ class TestTestToolsList: monkeypatch.setattr( auth_mcp.MCPRequestHandler, - "_get_oauth2_headers_from_headers", + "get_oauth2_headers_from_headers", staticmethod(lambda headers: oauth_headers), raising=False, ) @@ -1233,7 +1233,7 @@ class TestListToolsRestAPI: monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging) monkeypatch.setattr(manager, "get_allowed_mcp_servers", allowed_servers) monkeypatch.setattr(manager, "filter_server_ids_by_ip_with_info", lambda ids, _ip: (ids, 0)) - monkeypatch.setattr(manager, "_get_tools_from_server", upstream) + monkeypatch.setattr(manager, "get_tools_from_server", upstream) with pytest.raises(HTTPException) as error: await rest_endpoints.list_tool_rest_api( @@ -1274,7 +1274,7 @@ class TestListToolsRestAPI: monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging) monkeypatch.setattr(manager, "get_allowed_mcp_servers", allowed_servers) monkeypatch.setattr(manager, "filter_server_ids_by_ip_with_info", lambda ids, _ip: (ids, 0)) - monkeypatch.setattr(manager, "_get_tools_from_server", upstream) + monkeypatch.setattr(manager, "get_tools_from_server", upstream) result: Final = await rest_endpoints.list_tool_rest_api( _build_request(path="/mcp-rest/tools/list", method="GET"), @@ -1532,7 +1532,7 @@ class TestListToolsRestAPI: fake_get_toolset_by_name_cached, raising=False, ) - monkeypatch.setattr(rest_endpoints, "_apply_toolset_scope", fake_apply_toolset_scope, raising=False) + monkeypatch.setattr(rest_endpoints, "apply_toolset_scope", fake_apply_toolset_scope, raising=False) monkeypatch.setattr( rest_endpoints.global_mcp_server_manager, "get_allowed_mcp_servers", @@ -2319,7 +2319,7 @@ class TestListToolsRestAPI: ) monkeypatch.setattr( rest_endpoints, - "_apply_toolset_scope", + "apply_toolset_scope", fake_apply_toolset_scope, raising=False, ) @@ -3497,7 +3497,7 @@ class TestGetToolsForSingleServer: monkeypatch.setattr( rest_endpoints.global_mcp_server_manager, - "_get_tools_from_server", + "get_tools_from_server", fake_get_tools_from_server, raising=False, ) @@ -3550,7 +3550,7 @@ class TestGetToolsForSingleServer: monkeypatch.setattr( rest_endpoints.global_mcp_server_manager, - "_get_tools_from_server", + "get_tools_from_server", fake_get_tools_from_server, raising=False, ) @@ -3592,7 +3592,7 @@ class TestGetToolsForSingleServer: monkeypatch.setattr( rest_endpoints.global_mcp_server_manager, - "_get_tools_from_server", + "get_tools_from_server", fake_get_tools_from_server, raising=False, ) @@ -3639,7 +3639,7 @@ class TestGetToolsForSingleServer: monkeypatch.setattr( rest_endpoints.global_mcp_server_manager, - "_get_tools_from_server", + "get_tools_from_server", fake_get_tools_from_server, raising=False, ) @@ -3688,7 +3688,7 @@ class TestGetToolsForSingleServer: monkeypatch.setattr( rest_endpoints.global_mcp_server_manager, - "_get_tools_from_server", + "get_tools_from_server", fake_get_tools_from_server, raising=False, ) @@ -3739,7 +3739,7 @@ class TestGetToolsForSingleServer: monkeypatch.setattr( rest_endpoints.global_mcp_server_manager, - "_get_tools_from_server", + "get_tools_from_server", fake_get_tools_from_server, raising=False, ) @@ -4518,7 +4518,7 @@ class TestRestListToolsetFiltering: monkeypatch.setattr( rest_endpoints.global_mcp_server_manager, - "_get_tools_from_server", + "get_tools_from_server", AsyncMock(return_value=upstream_tools), ) diff --git a/tests/unit/proxy/_experimental/mcp_server/test_server_resolution.py b/tests/unit/proxy/_experimental/mcp_server/test_server_resolution.py index 51b97f7c8c5..4b9214916a8 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_server_resolution.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_server_resolution.py @@ -43,11 +43,11 @@ class FakeMCPServerManager: self.name_lookup_spy(server_name, client_ip) return self.servers_by_name.get(server_name) - def _is_server_accessible_from_ip(self, server: MCPServer, client_ip: str | None) -> bool: + def is_server_accessible_from_ip(self, server: MCPServer, client_ip: str | None) -> bool: self.ip_filter_spy(server, client_ip) return self.ip_accessible - def _build_mcp_server_table(self, server: MCPServer) -> LiteLLM_MCPServerTable: + def build_mcp_server_table(self, server: MCPServer) -> LiteLLM_MCPServerTable: return LiteLLM_MCPServerTable( server_id=server.server_id, alias=server.alias, @@ -132,7 +132,7 @@ async def test_temp_resolution_precedes_db_and_registry() -> None: ) assert resolved == ResolvedMCPServer( - table=manager._build_mcp_server_table(temporary_server), + table=manager.build_mcp_server_table(temporary_server), runtime=temporary_server, source="temp", ) @@ -179,7 +179,7 @@ async def test_registry_id_resolution_precedes_name() -> None: ) assert resolved == ResolvedMCPServer( - table=manager._build_mcp_server_table(server), + table=manager.build_mcp_server_table(server), runtime=server, source="registry", ) @@ -244,7 +244,7 @@ async def test_db_lookup_none_skips_db_and_returns_registry_source() -> None: resolved: Final = await resolve_mcp_server(server.server_id, manager=manager, db_lookup=None) assert resolved == ResolvedMCPServer( - table=manager._build_mcp_server_table(server), + table=manager.build_mcp_server_table(server), runtime=server, source="registry", ) @@ -286,7 +286,7 @@ async def test_non_admin_temp_resolution_is_denied_before_allowed_lookup() -> No server: Final = _runtime_server() manager: Final = _manager(allowed_server_ids=(server.server_id,)) resolved: Final = ResolvedMCPServer( - table=manager._build_mcp_server_table(server), + table=manager.build_mcp_server_table(server), runtime=server, source="temp", ) @@ -384,7 +384,7 @@ async def test_catalog_visibility_never_opens_temporary_setup_to_non_admins( ) -> None: server: Final = _runtime_server() manager: Final = _manager() - resolved: Final = ResolvedMCPServer(manager._build_mcp_server_table(server), server, source) + resolved: Final = ResolvedMCPServer(manager.build_mcp_server_table(server), server, source) operation: Final = authorize_mcp_server( resolved, _auth(), diff --git a/tests/unit/proxy/agent_endpoints/auth/test_agent_permission_handler.py b/tests/unit/proxy/agent_endpoints/auth/test_agent_permission_handler.py index a1a022fdd35..b5fa1cbc53d 100644 --- a/tests/unit/proxy/agent_endpoints/auth/test_agent_permission_handler.py +++ b/tests/unit/proxy/agent_endpoints/auth/test_agent_permission_handler.py @@ -526,7 +526,7 @@ class TestAgentRequestHandler: AgentRequestHandler, "_get_key_object_permission", return_value=None ): with patch( - "litellm.proxy.auth.auth_checks._get_agent_ids_from_access_groups", + "litellm.proxy.auth.auth_checks.get_agent_ids_from_access_groups", new_callable=AsyncMock, return_value=["agent-from-ag-1", "agent-from-ag-2"], ): @@ -558,7 +558,7 @@ class TestAgentRequestHandler: mock_user_auth.object_permission = mock_permission with patch( - "litellm.proxy.auth.auth_checks._get_agent_ids_from_access_groups", + "litellm.proxy.auth.auth_checks.get_agent_ids_from_access_groups", new_callable=AsyncMock, return_value=["agent-from-ag"], ): @@ -966,7 +966,7 @@ async def test_managed_target_rechecks_authoritative_key_after_peer_revocation( warm.object_permission = None warm.access_group_ids = ["old-group"] from litellm.proxy.auth import auth_checks - monkeypatch.setattr(auth_checks, "_get_agent_ids_from_access_groups", AsyncMock(return_value=["target"])) + monkeypatch.setattr(auth_checks, "get_agent_ids_from_access_groups", AsyncMock(return_value=["target"])) if change == "deleted": client.get_data.return_value = None if change == "outage": diff --git a/tests/unit/proxy/anthropic_endpoints/test_endpoints.py b/tests/unit/proxy/anthropic_endpoints/test_endpoints.py index 801f61aa498..c0e7a4fdb47 100644 --- a/tests/unit/proxy/anthropic_endpoints/test_endpoints.py +++ b/tests/unit/proxy/anthropic_endpoints/test_endpoints.py @@ -107,7 +107,7 @@ class TestBlockedResponseUsage: ) with ( - patch.object(ep, "_read_request_body", new=AsyncMock(return_value={})), + patch.object(ep, "read_request_body", new=AsyncMock(return_value={})), patch.object( ep.ProxyBaseLLMRequestProcessing, "base_process_llm_request", @@ -147,7 +147,7 @@ class TestProxyExceptionAnthropicEnvelope: request.headers = {"x-request-id": "req_test_6468"} with ( - patch.object(ep, "_read_request_body", new=AsyncMock(return_value={})), + patch.object(ep, "read_request_body", new=AsyncMock(return_value={})), patch.object( ep.ProxyBaseLLMRequestProcessing, "base_process_llm_request", @@ -278,7 +278,7 @@ class TestHttpExceptionDictDetail: request.headers = {} with ( - patch.object(ep, "_read_request_body", new=AsyncMock(return_value={})), # test-quality-ok: endpoint reads the body via a module function; no injection seam + patch.object(ep, "read_request_body", new=AsyncMock(return_value={})), # test-quality-ok: endpoint reads the body via a module function; no injection seam patch.object( # test-quality-ok: the guardrail raise happens deep inside this call; the test targets the endpoint's except block ep.ProxyBaseLLMRequestProcessing, "base_process_llm_request", @@ -324,7 +324,7 @@ class TestFailureHookRequestData: request.headers = {} with ( - patch.object(ep, "_read_request_body", new=AsyncMock(return_value={"model": "claude-sonnet"})), + patch.object(ep, "read_request_body", new=AsyncMock(return_value={"model": "claude-sonnet"})), patch.object(ep.ProxyBaseLLMRequestProcessing, "base_process_llm_request", new=fake_process), patch.object(proxy_server, "proxy_logging_obj") as mock_logging, ): @@ -375,7 +375,7 @@ class TestErrorLogCarriesCallId: request.headers = {} with ( - patch.object(ep, "_read_request_body", new=AsyncMock(return_value={"model": "claude-sonnet"})), # test-quality-ok: endpoint reads the body via a module function; no injection seam + patch.object(ep, "read_request_body", new=AsyncMock(return_value={"model": "claude-sonnet"})), # test-quality-ok: endpoint reads the body via a module function; no injection seam patch.object(ep.ProxyBaseLLMRequestProcessing, "base_process_llm_request", new=fake_process), # test-quality-ok: the provider failure happens inside this call; the test targets the endpoint's except block patch.object(proxy_server, "proxy_logging_obj") as mock_logging, # test-quality-ok: module global imported at call time; no injection seam caplog.at_level(logging.ERROR, logger="LiteLLM Proxy"), @@ -408,7 +408,7 @@ class TestErrorLogCarriesCallId: request.headers = {} with ( - patch.object(ep, "_read_request_body", new=AsyncMock(return_value={"model": "claude-sonnet"})), # test-quality-ok: endpoint reads the body via a module function; no injection seam + patch.object(ep, "read_request_body", new=AsyncMock(return_value={"model": "claude-sonnet"})), # test-quality-ok: endpoint reads the body via a module function; no injection seam patch.object(ep.ProxyBaseLLMRequestProcessing, "base_process_llm_request", new=fake_process), # test-quality-ok: the proxy shaped failure happens inside this call; the test targets the endpoint's except block patch.object(proxy_server, "proxy_logging_obj") as mock_logging, # test-quality-ok: module global imported at call time; no injection seam ): @@ -437,7 +437,7 @@ class TestErrorLogCarriesCallId: with ( patch.object( # test-quality-ok: endpoint reads the body via a module function; no injection seam ep, - "_read_request_body", + "read_request_body", new=AsyncMock(return_value={"model": "claude-sonnet", "messages": [{"role": "user", "content": "hi"}]}), ), patch.object(proxy_server, "token_counter", new=AsyncMock(side_effect=RuntimeError("tokenizer down"))), # test-quality-ok: module global imported at call time; the test targets the endpoint's except block diff --git a/tests/unit/proxy/auth/test_auth_checks.py b/tests/unit/proxy/auth/test_auth_checks.py index b698a3d02fa..dfd00a7458f 100644 --- a/tests/unit/proxy/auth/test_auth_checks.py +++ b/tests/unit/proxy/auth/test_auth_checks.py @@ -32,8 +32,8 @@ from litellm.proxy._types import ( from litellm.proxy.utils import PrismaClient from litellm.proxy.auth.auth_checks import ( can_team_access_model, - _is_model_cost_zero, - _virtual_key_soft_budget_check, + is_model_cost_zero, + virtual_key_soft_budget_check, _team_soft_budget_check, ) from litellm.proxy.utils import ProxyLogging @@ -91,7 +91,7 @@ async def test_check_end_user_budget(customer_spend, customer_budget): Note: Budget enforcement for end users happens in common_checks() via _check_end_user_budget(), not in get_end_user_object(). """ - from litellm.proxy.auth.auth_checks import _check_end_user_budget + from litellm.proxy.auth.auth_checks import check_end_user_budget _budget = LiteLLM_BudgetTable(max_budget=customer_budget) end_user_obj = LiteLLM_EndUserTable( @@ -104,14 +104,14 @@ async def test_check_end_user_budget(customer_spend, customer_budget): should_exceed = customer_spend > customer_budget if not should_exceed: - await _check_end_user_budget( + await check_end_user_budget( end_user_obj=end_user_obj, route="/v1/chat/completions", ) return with pytest.raises(litellm.BudgetExceededError) as exc_info: - await _check_end_user_budget( + await check_end_user_budget( end_user_obj=end_user_obj, route="/v1/chat/completions", ) @@ -476,7 +476,7 @@ async def test_virtual_key_max_budget_check( 1. Triggers budget alert for all cases 2. Raises BudgetExceededError when spend >= max_budget """ - from litellm.proxy.auth.auth_checks import _virtual_key_max_budget_check + from litellm.proxy.auth.auth_checks import virtual_key_max_budget_check # Setup test data valid_token = UserAPIKeyAuth( @@ -508,7 +508,7 @@ async def test_virtual_key_max_budget_check( if expect_budget_error: with pytest.raises(litellm.BudgetExceededError) as exc_info: - await _virtual_key_max_budget_check( + await virtual_key_max_budget_check( valid_token=valid_token, proxy_logging_obj=proxy_logging_obj, user_obj=user_obj, @@ -516,7 +516,7 @@ async def test_virtual_key_max_budget_check( assert exc_info.value.current_cost == token_spend assert exc_info.value.max_budget == max_budget else: - await _virtual_key_max_budget_check( + await virtual_key_max_budget_check( valid_token=valid_token, proxy_logging_obj=proxy_logging_obj, user_obj=user_obj, @@ -633,7 +633,7 @@ async def test_virtual_key_soft_budget_check(spend, soft_budget, expect_alert): proxy_logging_obj = MockProxyLogging() - await _virtual_key_soft_budget_check( + await virtual_key_soft_budget_check( valid_token=valid_token, proxy_logging_obj=proxy_logging_obj, ) @@ -970,7 +970,7 @@ async def test_can_key_call_model_with_aliases(model, alias_map, expect_to_work) @pytest.mark.asyncio async def test_cache_access_object(): """Test _cache_access_object stores access group in cache with correct key.""" - from litellm.proxy.auth.auth_checks import _cache_access_object + from litellm.proxy.auth.auth_checks import cache_access_object from litellm.proxy._types import LiteLLM_AccessGroupTable cache = DualCache() @@ -980,7 +980,7 @@ async def test_cache_access_object(): access_group_name="test-group", access_model_names=["gpt-4"], ) - await _cache_access_object( + await cache_access_object( access_group_id=ag_id, access_group_table=ag_table, user_api_key_cache=cache, @@ -998,7 +998,7 @@ async def test_cache_access_object(): @pytest.mark.asyncio async def test_delete_cache_access_object(): """Test _delete_cache_access_object removes access group from in-memory cache.""" - from litellm.proxy.auth.auth_checks import _delete_cache_access_object + from litellm.proxy.auth.auth_checks import delete_cache_access_object from litellm.proxy._types import LiteLLM_AccessGroupTable cache = DualCache() @@ -1008,7 +1008,7 @@ async def test_delete_cache_access_object(): access_group_name="to-delete", ) await cache.async_set_cache(key=f"access_group_id:{ag_id}", value=ag_table, ttl=60) - await _delete_cache_access_object(access_group_id=ag_id, user_api_key_cache=cache) + await delete_cache_access_object(access_group_id=ag_id, user_api_key_cache=cache) cached = await cache.async_get_cache(key=f"access_group_id:{ag_id}") assert cached is None @@ -1047,8 +1047,8 @@ async def test_get_resources_from_access_groups( from litellm.proxy._types import LiteLLM_AccessGroupTable from litellm.proxy.auth.auth_checks import ( - _get_agent_ids_from_access_groups, - _get_models_from_access_groups, + get_agent_ids_from_access_groups, + get_models_from_access_groups, ) ag_table = LiteLLM_AccessGroupTable( @@ -1064,13 +1064,13 @@ async def test_get_resources_from_access_groups( return_value=ag_table, ): if resource_field == "access_model_names": - result = await _get_models_from_access_groups( + result = await get_models_from_access_groups( access_group_ids=[access_group_data["access_group_id"]], prisma_client=MagicMock(), user_api_key_cache=DualCache(), ) else: - result = await _get_agent_ids_from_access_groups( + result = await get_agent_ids_from_access_groups( access_group_ids=[access_group_data["access_group_id"]], prisma_client=MagicMock(), user_api_key_cache=DualCache(), @@ -1081,9 +1081,9 @@ async def test_get_resources_from_access_groups( @pytest.mark.asyncio async def test_get_models_from_access_groups_empty_ids(): """Test _get_models_from_access_groups returns empty list when access_group_ids is empty.""" - from litellm.proxy.auth.auth_checks import _get_models_from_access_groups + from litellm.proxy.auth.auth_checks import get_models_from_access_groups - result = await _get_models_from_access_groups(access_group_ids=[]) + result = await get_models_from_access_groups(access_group_ids=[]) assert result == [] @@ -1106,7 +1106,7 @@ async def test_can_team_access_model_via_access_group_ids(): ) with patch( - "litellm.proxy.auth.auth_checks._get_models_from_access_groups", + "litellm.proxy.auth.auth_checks.get_models_from_access_groups", new_callable=AsyncMock, return_value=["gpt-4"], ): @@ -1134,7 +1134,7 @@ async def test_can_team_access_model_access_group_ids_denied(): ) with patch( - "litellm.proxy.auth.auth_checks._get_models_from_access_groups", + "litellm.proxy.auth.auth_checks.get_models_from_access_groups", new_callable=AsyncMock, return_value=["claude-3"], ): @@ -1174,7 +1174,7 @@ async def test_can_key_call_model_via_access_group_ids(): ) with patch( - "litellm.proxy.auth.auth_checks._get_models_from_access_groups", + "litellm.proxy.auth.auth_checks.get_models_from_access_groups", new_callable=AsyncMock, return_value=["gpt-4"], ): @@ -1231,7 +1231,7 @@ async def test_key_access_group_grants_model_when_team_authorized(): """ from unittest.mock import AsyncMock, patch - from litellm.proxy.auth.auth_checks import _key_access_group_grants_model + from litellm.proxy.auth.auth_checks import key_access_group_grants_model valid_token = UserAPIKeyAuth( token="test-token", @@ -1262,7 +1262,7 @@ async def test_key_access_group_grants_model_when_team_authorized(): p.start() try: assert ( - await _key_access_group_grants_model( + await key_access_group_grants_model( model="claude-haiku-4-5", valid_token=valid_token, team_object=team_object, @@ -1284,7 +1284,7 @@ async def test_key_access_group_grants_model_when_key_directly_authorized(): """ from unittest.mock import AsyncMock, patch - from litellm.proxy.auth.auth_checks import _key_access_group_grants_model + from litellm.proxy.auth.auth_checks import key_access_group_grants_model valid_token = UserAPIKeyAuth( token="test-token-hashed", @@ -1316,7 +1316,7 @@ async def test_key_access_group_grants_model_when_key_directly_authorized(): p.start() try: assert ( - await _key_access_group_grants_model( + await key_access_group_grants_model( model="claude-haiku-4-5", valid_token=valid_token, team_object=team_object, @@ -1332,7 +1332,7 @@ async def test_key_access_group_grants_model_when_key_directly_authorized(): @pytest.mark.asyncio async def test_key_access_group_grants_model_when_key_has_no_groups(): """Key with no access_group_ids → False (early return, no DB read).""" - from litellm.proxy.auth.auth_checks import _key_access_group_grants_model + from litellm.proxy.auth.auth_checks import key_access_group_grants_model valid_token = UserAPIKeyAuth( token="test-token", @@ -1346,7 +1346,7 @@ async def test_key_access_group_grants_model_when_key_has_no_groups(): access_group_ids=["any-group"], ) assert ( - await _key_access_group_grants_model( + await key_access_group_grants_model( model="claude-haiku-4-5", valid_token=valid_token, team_object=team_object, @@ -1361,7 +1361,7 @@ async def test_key_access_group_grants_model_when_group_does_not_cover_model(): """Group authorizes the team but does not grant the requested model → False.""" from unittest.mock import AsyncMock, patch - from litellm.proxy.auth.auth_checks import _key_access_group_grants_model + from litellm.proxy.auth.auth_checks import key_access_group_grants_model valid_token = UserAPIKeyAuth( token="test-token", @@ -1392,7 +1392,7 @@ async def test_key_access_group_grants_model_when_group_does_not_cover_model(): p.start() try: assert ( - await _key_access_group_grants_model( + await key_access_group_grants_model( model="claude-haiku-4-5", valid_token=valid_token, team_object=team_object, @@ -1415,7 +1415,7 @@ async def test_key_access_group_grants_model_when_group_authorizes_neither(): """ from unittest.mock import AsyncMock, patch - from litellm.proxy.auth.auth_checks import _key_access_group_grants_model + from litellm.proxy.auth.auth_checks import key_access_group_grants_model valid_token = UserAPIKeyAuth( token="team-a-token", @@ -1447,7 +1447,7 @@ async def test_key_access_group_grants_model_when_group_authorizes_neither(): p.start() try: assert ( - await _key_access_group_grants_model( + await key_access_group_grants_model( model="claude-opus-4-5", valid_token=valid_token, team_object=team_object, @@ -1465,7 +1465,7 @@ async def test_key_access_group_grants_model_when_get_access_object_raises(): """Group lookup failure (404, network, etc.) is treated as no authorization.""" from unittest.mock import AsyncMock, patch - from litellm.proxy.auth.auth_checks import _key_access_group_grants_model + from litellm.proxy.auth.auth_checks import key_access_group_grants_model valid_token = UserAPIKeyAuth( token="test-token", @@ -1490,7 +1490,7 @@ async def test_key_access_group_grants_model_when_get_access_object_raises(): p.start() try: assert ( - await _key_access_group_grants_model( + await key_access_group_grants_model( model="claude-haiku-4-5", valid_token=valid_token, team_object=team_object, @@ -1633,7 +1633,7 @@ def test_is_model_cost_zero_judges_an_alias_chain_by_the_deployment_its_entry_ro expected: Final = {"chain-entry": True, "local-free": False, "paid-gpt": False} order: Final = ("chain-entry", "local-free", "paid-gpt") if entry_first else ("local-free", "paid-gpt", "chain-entry") - verdicts: Final = {name: _is_model_cost_zero(model=name, llm_router=router) for name in order} + verdicts: Final = {name: is_model_cost_zero(model=name, llm_router=router) for name in order} assert verdicts == expected - assert {name: _is_model_cost_zero(model=name, llm_router=router) for name in order} == expected + assert {name: is_model_cost_zero(model=name, llm_router=router) for name in order} == expected diff --git a/tests/unit/proxy/auth/test_auth_checks_object_access_and_lookup.py b/tests/unit/proxy/auth/test_auth_checks_object_access_and_lookup.py index e855d8cd346..e26845e5b2a 100644 --- a/tests/unit/proxy/auth/test_auth_checks_object_access_and_lookup.py +++ b/tests/unit/proxy/auth/test_auth_checks_object_access_and_lookup.py @@ -43,24 +43,24 @@ from litellm.proxy.auth.auth_checks import ( LITELLM_SESSION_TOKEN_PREFIX, ExperimentalUIJWTToken, _cache_management_object, - _can_object_call_model, + can_object_call_model, _can_object_call_vector_stores, _check_agent_access_group_model_access, - _check_end_user_budget, + check_end_user_budget, _check_team_member_budget, - _fetch_key_object_from_db_with_reconnect, + fetch_key_object_from_db_with_reconnect, _get_fuzzy_user_object, CallerTeamLoader, CallerUserLoader, _get_team_db_check, _log_budget_lookup_failure, _tag_max_budget_check, - _team_max_budget_check, - _team_member_max_budget_alert_check, - _virtual_key_max_budget_alert_check, + team_max_budget_check, + team_member_max_budget_alert_check, + virtual_key_max_budget_alert_check, _check_agent_caller_model_access, - _virtual_key_max_budget_check, - _virtual_key_soft_budget_check, + virtual_key_max_budget_check, + virtual_key_soft_budget_check, common_checks, get_key_object, get_user_object, @@ -343,7 +343,7 @@ def test_get_key_object_from_ui_hash_key_invalid(): ) def test_can_object_call_model_denials_return_forbidden(object_type, expected_error_type): with pytest.raises(ProxyException) as exc_info: - _can_object_call_model( + can_object_call_model( model="restricted-model", llm_router=None, models=["allowed-model"], @@ -481,7 +481,7 @@ async def test_enforce_key_access_teamless_all_team_models_passes(): the sentinel is present, regardless of team_id. Fails if someone adds a team_id guard to the pass branch.""" from litellm.proxy._types import SpecialModelNames - from litellm.proxy.auth.user_api_key_auth import _enforce_key_and_fallback_model_access + from litellm.proxy.auth.user_api_key_auth import enforce_key_and_fallback_model_access valid_token = UserAPIKeyAuth( api_key="sk-orphan", @@ -489,7 +489,7 @@ async def test_enforce_key_access_teamless_all_team_models_passes(): team_models=[], ) - await _enforce_key_and_fallback_model_access( + await enforce_key_and_fallback_model_access( valid_token=valid_token, request_data={"model": "gpt-4o"}, route="/chat/completions", @@ -576,7 +576,7 @@ async def test_can_team_access_model_error_lists_direct_and_access_group_models( ) with patch( # test-quality-ok: access-group lookup has no dependency-injection seam - "litellm.proxy.auth.auth_checks._get_models_from_access_groups", + "litellm.proxy.auth.auth_checks.get_models_from_access_groups", new=AsyncMock(return_value=["group-model"]), ): assert await can_team_access_model("direct-model", team_object, None) is True @@ -669,7 +669,7 @@ async def test_fetch_key_object_from_db_bounds_in_flight_prisma_requests(): results: Final = await asyncio.gather( *( - _fetch_key_object_from_db_with_reconnect( + fetch_key_object_from_db_with_reconnect( hashed_token=f"hashed-token-{i}", prisma_client=prisma, # pyright: ignore[reportArgumentType] # fake stands in for PrismaClient parent_otel_span=None, @@ -730,7 +730,7 @@ async def test_fetch_key_object_from_db_fails_a_stalled_burst_within_the_deadlin results: Final = await asyncio.gather( *( - _fetch_key_object_from_db_with_reconnect( + fetch_key_object_from_db_with_reconnect( hashed_token=f"hashed-token-{i}", prisma_client=prisma, # pyright: ignore[reportArgumentType] # fake stands in for PrismaClient parent_otel_span=None, @@ -753,7 +753,7 @@ async def test_fetch_key_object_from_db_fails_a_stalled_burst_within_the_deadlin after: Final = await asyncio.wait_for( asyncio.gather( *( - _fetch_key_object_from_db_with_reconnect( + fetch_key_object_from_db_with_reconnect( hashed_token=f"after-{i}", prisma_client=recovered, # pyright: ignore[reportArgumentType] # fake stands in for PrismaClient parent_otel_span=None, @@ -2379,7 +2379,7 @@ async def test_key_and_team_grants_are_read_through_the_object_permission_cache( def test_can_object_call_model_with_alias(): """Test that can_object_call_model works with model aliases""" from litellm import Router - from litellm.proxy.auth.auth_checks import _can_object_call_model + from litellm.proxy.auth.auth_checks import can_object_call_model model = "[ip-approved] gpt-4o" llm_router = Router( @@ -2400,7 +2400,7 @@ def test_can_object_call_model_with_alias(): }, ) - result = _can_object_call_model( + result = can_object_call_model( model=model, llm_router=llm_router, models=["gpt-3.5-turbo"], @@ -2423,7 +2423,7 @@ def test_can_object_call_model_access_via_alias_only(): - The call should succeed because access is granted via the alias """ from litellm import Router - from litellm.proxy.auth.auth_checks import _can_object_call_model + from litellm.proxy.auth.auth_checks import can_object_call_model model = "my-fake-gpt" llm_router = Router( @@ -2445,7 +2445,7 @@ def test_can_object_call_model_access_via_alias_only(): ) # Key has access to the alias but NOT the underlying model - result = _can_object_call_model( + result = can_object_call_model( model=model, llm_router=llm_router, models=["my-fake-gpt"], # Only has access to alias, not "gpt-4" @@ -2460,9 +2460,9 @@ def test_can_object_call_model_access_via_alias_only(): def test_can_object_call_model_key_alias_to_allowed_target_is_allowed(): """A key alias whose target is on the key allowlist resolves like a team alias.""" - from litellm.proxy.auth.auth_checks import _can_object_call_model + from litellm.proxy.auth.auth_checks import can_object_call_model - result = _can_object_call_model( + result = can_object_call_model( model="mistral-7b", llm_router=None, models=["gpt-4o-mini"], @@ -2477,10 +2477,10 @@ def test_can_object_call_model_key_alias_to_allowed_target_is_allowed(): def test_can_object_call_model_key_alias_to_disallowed_target_is_denied(): """A key alias whose target is outside the key allowlist stays denied.""" from litellm.proxy._types import ProxyErrorTypes, ProxyException - from litellm.proxy.auth.auth_checks import _can_object_call_model + from litellm.proxy.auth.auth_checks import can_object_call_model with pytest.raises(ProxyException) as exc_info: - _can_object_call_model( + can_object_call_model( model="mistral-7b", llm_router=None, models=["gpt-4o-mini"], @@ -2563,12 +2563,12 @@ async def test_can_key_call_model_honors_key_alias(): def test_can_object_call_model_key_alias_applies_before_global_alias(monkeypatch): """The key alias rewrite precedes the global one at dispatch, so the key target is authorized.""" - from litellm.proxy.auth.auth_checks import _can_object_call_model + from litellm.proxy.auth.auth_checks import can_object_call_model monkeypatch.setattr(litellm, "model_alias_map", {"foo": "bar"}) assert ( - _can_object_call_model( + can_object_call_model( model="foo", llm_router=None, models=["baz"], @@ -2580,7 +2580,7 @@ def test_can_object_call_model_key_alias_applies_before_global_alias(monkeypatch ) with pytest.raises(ProxyException) as exc_info: - _can_object_call_model( + can_object_call_model( model="foo", llm_router=None, models=["bar"], @@ -2594,12 +2594,12 @@ def test_can_object_call_model_key_alias_applies_before_global_alias(monkeypatch def test_can_object_call_model_key_alias_matches_global_rewritten_name(monkeypatch): """A key alias on the globally rewritten name resolves the same way the request chain does.""" - from litellm.proxy.auth.auth_checks import _can_object_call_model + from litellm.proxy.auth.auth_checks import can_object_call_model monkeypatch.setattr(litellm, "model_alias_map", {"foo": "bar"}) assert ( - _can_object_call_model( + can_object_call_model( model="foo", llm_router=None, models=["baz"], @@ -2613,12 +2613,12 @@ def test_can_object_call_model_key_alias_matches_global_rewritten_name(monkeypat def test_can_object_call_model_chained_alias_requires_final_target(monkeypatch): """When a key alias fires on the globally rewritten name, only the final target is dispatched.""" - from litellm.proxy.auth.auth_checks import _can_object_call_model + from litellm.proxy.auth.auth_checks import can_object_call_model monkeypatch.setattr(litellm, "model_alias_map", {"foo": "bar"}) with pytest.raises(ProxyException) as exc_info: - _can_object_call_model( + can_object_call_model( model="foo", llm_router=None, models=["bar"], @@ -2630,7 +2630,7 @@ def test_can_object_call_model_chained_alias_requires_final_target(monkeypatch): assert exc_info.value.type == ProxyErrorTypes.key_model_access_denied assert ( - _can_object_call_model( + can_object_call_model( model="foo", llm_router=None, models=["baz"], @@ -2644,10 +2644,10 @@ def test_can_object_call_model_chained_alias_requires_final_target(monkeypatch): def test_can_object_call_model_key_alias_name_alone_is_not_enough(): """A key that may call the alias name but not its target cannot call the alias.""" - from litellm.proxy.auth.auth_checks import _can_object_call_model + from litellm.proxy.auth.auth_checks import can_object_call_model with pytest.raises(ProxyException) as exc_info: - _can_object_call_model( + can_object_call_model( model="bar", llm_router=None, models=["bar"], @@ -2659,7 +2659,7 @@ def test_can_object_call_model_key_alias_name_alone_is_not_enough(): assert exc_info.value.type == ProxyErrorTypes.key_model_access_denied assert ( - _can_object_call_model( + can_object_call_model( model="bar", llm_router=None, models=["baz"], @@ -2673,10 +2673,10 @@ def test_can_object_call_model_key_alias_name_alone_is_not_enough(): def test_can_object_call_model_team_alias_applies_before_key_alias(): """A key alias on the raw name loses to the team alias that rewrites it first at dispatch.""" - from litellm.proxy.auth.auth_checks import _can_object_call_model + from litellm.proxy.auth.auth_checks import can_object_call_model assert ( - _can_object_call_model( + can_object_call_model( model="foo", llm_router=None, models=["bar"], @@ -2691,10 +2691,10 @@ def test_can_object_call_model_team_alias_applies_before_key_alias(): def test_can_object_call_model_key_alias_on_team_alias_target(): """A key alias on the team-rewritten name resolves like the dispatch chain does.""" - from litellm.proxy.auth.auth_checks import _can_object_call_model + from litellm.proxy.auth.auth_checks import can_object_call_model assert ( - _can_object_call_model( + can_object_call_model( model="foo", llm_router=None, models=["baz"], @@ -2707,7 +2707,7 @@ def test_can_object_call_model_key_alias_on_team_alias_target(): ) with pytest.raises(ProxyException) as exc_info: - _can_object_call_model( + can_object_call_model( model="foo", llm_router=None, models=["bar"], @@ -2751,7 +2751,7 @@ async def test_can_user_call_model_honors_key_alias(): async def test_check_team_member_model_access_honors_key_alias(): """A key alias resolves against the member allowlist, not just the raw alias name.""" from litellm.proxy._types import LiteLLM_TeamMembership - from litellm.proxy.auth.auth_checks import _check_team_member_model_access + from litellm.proxy.auth.auth_checks import check_team_member_model_access membership = LiteLLM_TeamMembership( user_id="alice", @@ -2759,7 +2759,7 @@ async def test_check_team_member_model_access_honors_key_alias(): litellm_budget_table=LiteLLM_BudgetTable(allowed_models=["gpt-4o-mini"]), ) - await _check_team_member_model_access( + await check_team_member_model_access( model="mistral-7b", team_object=LiteLLM_TeamTable(team_id="team-a"), valid_token=UserAPIKeyAuth(token="sk-test", user_id="alice", team_id="team-a"), @@ -2773,7 +2773,7 @@ async def test_check_team_member_model_access_honors_key_alias(): ) with pytest.raises(ProxyException) as exc_info: - await _check_team_member_model_access( + await check_team_member_model_access( model="mistral-7b", team_object=LiteLLM_TeamTable(team_id="team-a"), valid_token=UserAPIKeyAuth(token="sk-test", user_id="alice", team_id="team-a"), @@ -2799,7 +2799,7 @@ def test_can_object_call_model_access_via_underlying_model_only(): - The call should succeed because access is granted via the underlying model """ from litellm import Router - from litellm.proxy.auth.auth_checks import _can_object_call_model + from litellm.proxy.auth.auth_checks import can_object_call_model model = "my-fake-gpt" llm_router = Router( @@ -2821,7 +2821,7 @@ def test_can_object_call_model_access_via_underlying_model_only(): ) # Key has access to the underlying model but NOT the alias - result = _can_object_call_model( + result = can_object_call_model( model=model, llm_router=llm_router, models=["gpt-4"], # Only has access to underlying model, not "my-fake-gpt" @@ -2840,7 +2840,7 @@ def test_can_object_call_model_no_access_to_alias_or_underlying(): """ from litellm import Router from litellm.proxy._types import ProxyErrorTypes, ProxyException - from litellm.proxy.auth.auth_checks import _can_object_call_model + from litellm.proxy.auth.auth_checks import can_object_call_model model = "my-fake-gpt" llm_router = Router( @@ -2863,7 +2863,7 @@ def test_can_object_call_model_no_access_to_alias_or_underlying(): # Key has access to neither the alias nor the underlying model with pytest.raises(ProxyException) as exc_info: - _can_object_call_model( + can_object_call_model( model=model, llm_router=llm_router, models=["gpt-3.5-turbo"], # Has access to different model entirely @@ -2887,7 +2887,7 @@ _DENIED_MESSAGE_TEMPLATE: Final = ( def test_can_object_call_model_denial_hides_allowlist_and_keeps_detail_on_exception(caplog): with caplog.at_level("DEBUG", logger="LiteLLM Proxy"): with pytest.raises(ModelAccessDeniedProxyException) as exc_info: - _can_object_call_model( + can_object_call_model( model="anthropic-sonnet-4-5", llm_router=None, models=["internal-models"], @@ -2914,7 +2914,7 @@ async def test_access_group_fallback_grant_does_not_log_a_denial(caplog): with ( patch( # test-quality-ok: access-group lookup has no dependency-injection seam - "litellm.proxy.auth.auth_checks._get_models_from_access_groups", + "litellm.proxy.auth.auth_checks.get_models_from_access_groups", new=AsyncMock(return_value=["group-model"]), ), caplog.at_level("DEBUG", logger="LiteLLM Proxy"), @@ -2934,7 +2934,7 @@ async def test_access_group_fallback_grant_does_not_log_a_denial(caplog): ) def test_can_object_call_model_denial_same_client_message_for_every_object_type(object_type, expected_type): with pytest.raises(ModelAccessDeniedProxyException) as exc_info: - _can_object_call_model( + can_object_call_model( model="anthropic-sonnet-4-5", llm_router=None, models=["internal-models"], @@ -2964,7 +2964,7 @@ async def test_can_user_call_model_no_default_models_hides_policy_detail(): @pytest.mark.asyncio async def test_check_team_member_model_access_denied_hides_member_allowlist(): from litellm.proxy._types import LiteLLM_TeamMembership - from litellm.proxy.auth.auth_checks import _check_team_member_model_access + from litellm.proxy.auth.auth_checks import check_team_member_model_access from litellm.proxy.common_utils.user_api_key_cache import team_membership_reservation_cache_key membership = LiteLLM_TeamMembership( @@ -2980,7 +2980,7 @@ async def test_check_team_member_model_access_denied_hides_member_allowlist(): ) with pytest.raises(ModelAccessDeniedProxyException) as exc_info: - await _check_team_member_model_access( + await check_team_member_model_access( model="mock-vision", team_object=LiteLLM_TeamTable(team_id="team-a"), valid_token=UserAPIKeyAuth(token="sk-test", user_id="alice", team_id="team-a"), @@ -3058,11 +3058,11 @@ def test_can_object_call_model_access_group_with_team_id(): model_info.access_groups for team-scoped DB models and allow access via group name. """ - from litellm.proxy.auth.auth_checks import _can_object_call_model + from litellm.proxy.auth.auth_checks import can_object_call_model router = _make_team_scoped_router() - result = _can_object_call_model( + result = can_object_call_model( model="mock-fast-1", llm_router=router, models=["fast-models", "mock-power"], @@ -3079,12 +3079,12 @@ def test_can_object_call_model_access_group_without_team_id_fails(): This is the pre-fix behavior. """ from litellm.proxy._types import ProxyException - from litellm.proxy.auth.auth_checks import _can_object_call_model + from litellm.proxy.auth.auth_checks import can_object_call_model router = _make_team_scoped_router() with pytest.raises(ProxyException): - _can_object_call_model( + can_object_call_model( model="mock-fast-1", llm_router=router, models=["fast-models", "mock-power"], @@ -3098,11 +3098,11 @@ def test_can_object_call_model_literal_name_with_team_id(): Literal model name matching should still work when team_id is passed — no regression from adding team_id. """ - from litellm.proxy.auth.auth_checks import _can_object_call_model + from litellm.proxy.auth.auth_checks import can_object_call_model router = _make_team_scoped_router() - result = _can_object_call_model( + result = can_object_call_model( model="mock-power", llm_router=router, models=["fast-models", "mock-power"], @@ -3118,12 +3118,12 @@ def test_can_object_call_model_denied_model_with_team_id(): still be denied even when team_id is passed. """ from litellm.proxy._types import ProxyException - from litellm.proxy.auth.auth_checks import _can_object_call_model + from litellm.proxy.auth.auth_checks import can_object_call_model router = _make_team_scoped_router() with pytest.raises(ProxyException): - _can_object_call_model( + can_object_call_model( model="mock-vision", llm_router=router, models=["fast-models", "mock-power"], @@ -3137,11 +3137,11 @@ def test_can_object_call_model_second_group_member_with_team_id(): Both models in the access group should be reachable, not just the first one. """ - from litellm.proxy.auth.auth_checks import _can_object_call_model + from litellm.proxy.auth.auth_checks import can_object_call_model router = _make_team_scoped_router() - result = _can_object_call_model( + result = can_object_call_model( model="mock-fast-2", llm_router=router, models=["fast-models"], @@ -3164,7 +3164,7 @@ async def test_check_team_member_model_access_with_access_group(): LiteLLM_TeamTable, UserAPIKeyAuth, ) - from litellm.proxy.auth.auth_checks import _check_team_member_model_access + from litellm.proxy.auth.auth_checks import check_team_member_model_access router = _make_team_scoped_router() team = LiteLLM_TeamTable(team_id="team-a") @@ -3182,7 +3182,7 @@ async def test_check_team_member_model_access_with_access_group(): return_value=membership, ): # Should not raise — mock-fast-1 is in the fast-models group - await _check_team_member_model_access( + await check_team_member_model_access( model="mock-fast-1", team_object=team, valid_token=token, @@ -3206,7 +3206,7 @@ async def test_check_team_member_model_access_denied_model(): ProxyException, UserAPIKeyAuth, ) - from litellm.proxy.auth.auth_checks import _check_team_member_model_access + from litellm.proxy.auth.auth_checks import check_team_member_model_access router = _make_team_scoped_router() team = LiteLLM_TeamTable(team_id="team-a") @@ -3224,7 +3224,7 @@ async def test_check_team_member_model_access_denied_model(): return_value=membership, ): with pytest.raises(ProxyException) as exc_info: - await _check_team_member_model_access( + await check_team_member_model_access( model="mock-vision", team_object=team, valid_token=token, @@ -3249,7 +3249,7 @@ async def test_check_team_member_model_access_no_override_inherits_team(): LiteLLM_TeamTable, UserAPIKeyAuth, ) - from litellm.proxy.auth.auth_checks import _check_team_member_model_access + from litellm.proxy.auth.auth_checks import check_team_member_model_access router = _make_team_scoped_router() team = LiteLLM_TeamTable(team_id="team-a") @@ -3265,7 +3265,7 @@ async def test_check_team_member_model_access_no_override_inherits_team(): return_value=membership, ): # Should return without raising — no per-member restriction - await _check_team_member_model_access( + await check_team_member_model_access( model="mock-vision", team_object=team, valid_token=token, @@ -4278,7 +4278,7 @@ async def test_virtual_key_soft_budget_check_with_user_obj(): proxy_logging_obj = MockProxyLogging() - await _virtual_key_soft_budget_check( + await virtual_key_soft_budget_check( valid_token=valid_token, proxy_logging_obj=proxy_logging_obj, user_obj=user_obj, @@ -4326,7 +4326,7 @@ async def test_virtual_key_soft_budget_check_without_user_obj(): proxy_logging_obj = MockProxyLogging() - await _virtual_key_soft_budget_check( + await virtual_key_soft_budget_check( valid_token=valid_token, proxy_logging_obj=proxy_logging_obj, user_obj=None, @@ -4370,7 +4370,7 @@ async def test_virtual_key_soft_budget_check_scenarios(spend, soft_budget, expec proxy_logging_obj = MockProxyLogging() - await _virtual_key_soft_budget_check( + await virtual_key_soft_budget_check( valid_token=valid_token, proxy_logging_obj=proxy_logging_obj, user_obj=None, @@ -4417,7 +4417,7 @@ async def test_virtual_key_max_budget_alert_check_with_user_obj(): proxy_logging_obj = MockProxyLogging() - await _virtual_key_max_budget_alert_check( + await virtual_key_max_budget_alert_check( valid_token=valid_token, proxy_logging_obj=proxy_logging_obj, user_obj=user_obj, @@ -4465,7 +4465,7 @@ async def test_virtual_key_max_budget_alert_check_without_user_obj(): proxy_logging_obj = MockProxyLogging() - await _virtual_key_max_budget_alert_check( + await virtual_key_max_budget_alert_check( valid_token=valid_token, proxy_logging_obj=proxy_logging_obj, user_obj=None, @@ -4503,7 +4503,7 @@ async def test_team_member_max_budget_alert_check_dispatches_only_at_configured_ async def budget_alerts(self, type, user_info): captured.append((type, user_info)) - _team_member_max_budget_alert_check( + team_member_max_budget_alert_check( team_id="team-1", team_alias="platform", team_metadata=team_metadata, @@ -4537,7 +4537,7 @@ async def test_team_member_max_budget_alert_check_drops_thresholds_outside_1_to_ async def budget_alerts(self, type, user_info): captured.append(user_info) - _team_member_max_budget_alert_check( + team_member_max_budget_alert_check( team_id="team-1", team_alias="platform", team_metadata={ @@ -4647,7 +4647,7 @@ async def test_virtual_key_max_budget_alert_check_scenarios(spend, max_budget, e proxy_logging_obj = MockProxyLogging() - await _virtual_key_max_budget_alert_check( + await virtual_key_max_budget_alert_check( valid_token=valid_token, proxy_logging_obj=proxy_logging_obj, user_obj=None, @@ -4690,7 +4690,7 @@ async def test_virtual_key_max_budget_alert_check_with_multi_threshold_map(): max_budget=None, ) - await _virtual_key_max_budget_alert_check( + await virtual_key_max_budget_alert_check( valid_token=valid_token, proxy_logging_obj=MockProxyLogging(), user_obj=user_obj, @@ -4726,7 +4726,7 @@ async def test_virtual_key_max_budget_alert_check_old_path_no_map(): metadata={}, ) - await _virtual_key_max_budget_alert_check( + await virtual_key_max_budget_alert_check( valid_token=valid_token, proxy_logging_obj=MockProxyLogging(), user_obj=None, @@ -4758,7 +4758,7 @@ async def test_virtual_key_max_budget_alert_check_old_path_below_threshold_no_al metadata={}, ) - await _virtual_key_max_budget_alert_check( + await virtual_key_max_budget_alert_check( valid_token=valid_token, proxy_logging_obj=MockProxyLogging(), user_obj=None, @@ -4798,7 +4798,7 @@ async def test_virtual_key_max_budget_alert_check_global_fallback(): original = litellm.default_key_max_budget_alert_emails try: litellm.default_key_max_budget_alert_emails = global_config - await _virtual_key_max_budget_alert_check( + await virtual_key_max_budget_alert_check( valid_token=valid_token, proxy_logging_obj=MockProxyLogging(), user_obj=None, @@ -4838,7 +4838,7 @@ async def test_virtual_key_max_budget_alert_check_per_key_merges_with_global(): original = litellm.default_key_max_budget_alert_emails try: litellm.default_key_max_budget_alert_emails = global_config - await _virtual_key_max_budget_alert_check( + await virtual_key_max_budget_alert_check( valid_token=valid_token, proxy_logging_obj=MockProxyLogging(), user_obj=None, @@ -4906,7 +4906,7 @@ async def test_custom_auth_common_checks_opt_in(): the pre-existing RPS guarantee for custom-auth hot paths. """ import litellm.proxy.proxy_server as _proxy_server_mod - from litellm.proxy.auth.user_api_key_auth import _run_centralized_common_checks + from litellm.proxy.auth.user_api_key_auth import run_centralized_common_checks valid_token = UserAPIKeyAuth(token="test-token", user_id="u1") mock_request = MagicMock() @@ -4934,7 +4934,7 @@ async def test_custom_auth_common_checks_opt_in(): "litellm.proxy.auth.user_api_key_auth.common_checks", new_callable=AsyncMock, ) as mock_common: - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=valid_token, request=mock_request, request_data={}, @@ -4955,7 +4955,7 @@ async def test_custom_auth_common_checks_opt_in(): "litellm.proxy.auth.user_api_key_auth.common_checks", new_callable=AsyncMock, ) as mock_common: - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=valid_token, request=mock_request, request_data={}, @@ -4995,7 +4995,7 @@ async def test_virtual_key_budget_check_reads_from_spend_counter(): with patch("litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend): with pytest.raises(litellm.BudgetExceededError) as exc_info: - await _virtual_key_max_budget_check( + await virtual_key_max_budget_check( valid_token=valid_token, proxy_logging_obj=proxy_logging_obj, ) @@ -5027,7 +5027,7 @@ async def test_virtual_key_budget_check_fallback_no_counter(): with patch("litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend): with pytest.raises(litellm.BudgetExceededError) as exc_info: - await _virtual_key_max_budget_check( + await virtual_key_max_budget_check( valid_token=valid_token, proxy_logging_obj=proxy_logging_obj, ) @@ -5093,7 +5093,7 @@ async def test_budget_exceeded_throttles_instead_of_blocking(monkeypatch): ) with _patched_spend(20.0): - await _virtual_key_max_budget_check( + await virtual_key_max_budget_check( valid_token=valid_token, proxy_logging_obj=_budget_logging_obj(), ) @@ -5111,12 +5111,12 @@ async def test_budget_exceeded_throttles_instead_of_blocking(monkeypatch): async def test_budget_throttle_decision_cleared_before_caching(): """The request-scoped throttle decision must not persist into the key cache, otherwise it would re-apply (and compound) on every subsequent request.""" - from litellm.proxy.auth.auth_checks import _copy_user_api_key_auth_for_cache + from litellm.proxy.auth.auth_checks import copy_user_api_key_auth_for_cache valid_token = _over_budget_token(tpm_limit=1000, rpm_limit=100, metadata={"throttle_on_budget_exceeded": True}) valid_token.budget_throttle_pct = 0.1 - cached = _copy_user_api_key_auth_for_cache(user_api_key_obj=valid_token) + cached = copy_user_api_key_auth_for_cache(user_api_key_obj=valid_token) assert cached.budget_throttle_pct is None assert cached.tpm_limit == 1000 @@ -5132,7 +5132,7 @@ async def test_budget_exceeded_throttle_no_configured_limits(monkeypatch): with _patched_spend(20.0): with pytest.raises(litellm.BudgetExceededError): - await _virtual_key_max_budget_check( + await virtual_key_max_budget_check( valid_token=valid_token, proxy_logging_obj=_budget_logging_obj(), ) @@ -5147,7 +5147,7 @@ async def test_budget_exceeded_not_opted_in_still_blocks(monkeypatch): with _patched_spend(20.0): with pytest.raises(litellm.BudgetExceededError): - await _virtual_key_max_budget_check( + await virtual_key_max_budget_check( valid_token=valid_token, proxy_logging_obj=_budget_logging_obj(), ) @@ -5167,7 +5167,7 @@ async def test_budget_exceeded_invalid_percentage_blocks(monkeypatch, pct): with _patched_spend(20.0): with pytest.raises(litellm.BudgetExceededError): - await _virtual_key_max_budget_check( + await virtual_key_max_budget_check( valid_token=valid_token, proxy_logging_obj=_budget_logging_obj(), ) @@ -5186,7 +5186,7 @@ async def test_under_budget_does_not_throttle(monkeypatch): ) with _patched_spend(5.0): - await _virtual_key_max_budget_check( + await virtual_key_max_budget_check( valid_token=valid_token, proxy_logging_obj=_budget_logging_obj(), ) @@ -5216,7 +5216,7 @@ async def test_team_budget_check_reads_from_spend_counter(): with patch("litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend): with pytest.raises(litellm.BudgetExceededError) as exc_info: - await _team_max_budget_check( + await team_max_budget_check( team_object=team_object, valid_token=valid_token, proxy_logging_obj=proxy_logging_obj, @@ -5243,7 +5243,7 @@ async def test_end_user_budget_check_reads_from_spend_counter(): with patch("litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend): with pytest.raises(litellm.BudgetExceededError) as exc_info: - await _check_end_user_budget( + await check_end_user_budget( end_user_obj=end_user_object, route="/chat/completions", ) @@ -6311,7 +6311,7 @@ async def test_cache_team_object_writes_team_id_and_invalidates_team_alias(): from unittest.mock import AsyncMock, MagicMock from litellm.proxy._types import LiteLLM_TeamTableCachedObj - from litellm.proxy.auth.auth_checks import _cache_team_object + from litellm.proxy.auth.auth_checks import cache_team_object base_team_row = { "team_id": "team-1234", @@ -6328,7 +6328,7 @@ async def test_cache_team_object_writes_team_id_and_invalidates_team_alias(): logging_obj = MagicMock() logging_obj.internal_usage_cache.dual_cache.async_delete_cache = AsyncMock() - await _cache_team_object( + await cache_team_object( team_id="team-1234", team_table=team_table, user_api_key_cache=cache, @@ -6365,7 +6365,7 @@ async def test_cache_team_object_writes_team_id_and_invalidates_team_alias(): logging_obj2 = MagicMock() logging_obj2.internal_usage_cache.dual_cache.async_delete_cache = AsyncMock() - await _cache_team_object( + await cache_team_object( team_id="team-no-alias", team_table=aliasless, user_api_key_cache=cache2, @@ -6427,7 +6427,7 @@ async def test_team_update_not_shadowed_by_internal_usage_cache_lit_4391(): """ from litellm.caching.dual_cache import DualCache from litellm.proxy._types import LiteLLM_TeamTableCachedObj - from litellm.proxy.auth.auth_checks import _cache_team_object, get_team_object + from litellm.proxy.auth.auth_checks import cache_team_object, get_team_object team_id = "team-lit-4391" shared_redis = _SharedFakeRedis() @@ -6439,7 +6439,7 @@ async def test_team_update_not_shadowed_by_internal_usage_cache_lit_4391(): ) prisma_client = MagicMock() - await _cache_team_object( + await cache_team_object( team_id=team_id, team_table=LiteLLM_TeamTableCachedObj(team_id=team_id, models=["model-a"]), user_api_key_cache=user_api_key_cache, @@ -6454,7 +6454,7 @@ async def test_team_update_not_shadowed_by_internal_usage_cache_lit_4391(): ) assert primed is not None and primed.models == ["model-a"] - await _cache_team_object( + await cache_team_object( team_id=team_id, team_table=LiteLLM_TeamTableCachedObj(team_id=team_id, models=["model-a", "model-b"]), user_api_key_cache=user_api_key_cache, @@ -6514,7 +6514,7 @@ async def test_warm_team_object_reads_issue_no_redis_ops_lit_5944(): """ from litellm.caching.dual_cache import DualCache from litellm.proxy._types import LiteLLM_TeamTableCachedObj - from litellm.proxy.auth.auth_checks import _cache_team_object, get_team_object + from litellm.proxy.auth.auth_checks import cache_team_object, get_team_object team_id = "team-lit-5944" counting_redis = _CountingFakeRedis() @@ -6526,7 +6526,7 @@ async def test_warm_team_object_reads_issue_no_redis_ops_lit_5944(): ) prisma_client = MagicMock() - await _cache_team_object( + await cache_team_object( team_id=team_id, team_table=LiteLLM_TeamTableCachedObj(team_id=team_id, models=["model-a"]), user_api_key_cache=user_api_key_cache, @@ -6560,7 +6560,7 @@ async def test_cache_team_object_tolerates_cache_invalidation_failures(): a 500. The authoritative team_id-keyed write must still happen. """ from litellm.proxy._types import LiteLLM_TeamTableCachedObj - from litellm.proxy.auth.auth_checks import _cache_team_object + from litellm.proxy.auth.auth_checks import cache_team_object cache = MagicMock() cache.async_set_cache = AsyncMock() @@ -6568,7 +6568,7 @@ async def test_cache_team_object_tolerates_cache_invalidation_failures(): logging_obj = MagicMock() logging_obj.internal_usage_cache.dual_cache.async_delete_cache = AsyncMock(side_effect=Exception("redis down")) - await _cache_team_object( + await cache_team_object( team_id="team-cache-outage", team_table=LiteLLM_TeamTableCachedObj( team_id="team-cache-outage", @@ -6711,7 +6711,7 @@ async def test_virtual_key_max_budget_error_names_the_key(): new=AsyncMock(return_value=25.0), ): with pytest.raises(litellm.BudgetExceededError) as exc_info: - await _virtual_key_max_budget_check( + await virtual_key_max_budget_check( valid_token=valid_token, proxy_logging_obj=proxy_logging_obj, ) @@ -6737,7 +6737,7 @@ async def test_virtual_key_max_budget_not_exceeded_does_not_raise(): "litellm.proxy.proxy_server.get_current_spend", new=AsyncMock(return_value=1.0), ): - await _virtual_key_max_budget_check( + await virtual_key_max_budget_check( valid_token=valid_token, proxy_logging_obj=proxy_logging_obj, ) @@ -6990,7 +6990,7 @@ async def test_common_checks_personal_user_budget_blocks_in_gather(): async def _common_checks_for_over_budget_personal_key(*, model: str) -> bool: from litellm import Router - from litellm.proxy.auth.auth_checks import _is_model_cost_zero, common_checks + from litellm.proxy.auth.auth_checks import is_model_cost_zero, common_checks llm_router: Final = Router( model_list=[ @@ -7030,7 +7030,7 @@ async def _common_checks_for_over_budget_personal_key(*, model: str) -> bool: proxy_logging_obj=proxy_logging_obj, valid_token=token, request=MagicMock(spec=Request), - skip_budget_checks=_is_model_cost_zero(model=model, llm_router=llm_router), + skip_budget_checks=is_model_cost_zero(model=model, llm_router=llm_router), ) await asyncio.sleep(0) return result @@ -7388,7 +7388,7 @@ async def test_organization_budget_check_carries_org_state_on_the_token(): (Prometheus org budget gauges) reads it from request metadata instead of calling get_org_object again.""" from litellm.proxy._types import LiteLLM_OrganizationTable - from litellm.proxy.auth.auth_checks import _organization_max_budget_check + from litellm.proxy.auth.auth_checks import organization_max_budget_check from litellm.types.proxy.carried_budget_state import OrgBudgetSnapshot org_table = LiteLLM_OrganizationTable( @@ -7406,7 +7406,7 @@ async def test_organization_budget_check_carries_org_state_on_the_token(): key="org_id:o1:with_budget", value=org_table, model_type=LiteLLM_OrganizationTable ) - await _organization_max_budget_check( + await organization_max_budget_check( valid_token=token, team_object=None, prisma_client=MagicMock(), @@ -7540,7 +7540,7 @@ async def test_organization_zero_max_budget_is_enforced(max_budget, spend, expec spend without limit. """ from litellm.proxy._types import LiteLLM_OrganizationTable - from litellm.proxy.auth.auth_checks import _organization_max_budget_check + from litellm.proxy.auth.auth_checks import organization_max_budget_check org_table = LiteLLM_OrganizationTable( organization_id="o1", @@ -7568,7 +7568,7 @@ async def test_organization_zero_max_budget_is_enforced(max_budget, spend, expec ): if expect_blocked: with pytest.raises(litellm.BudgetExceededError) as exc_info: - await _organization_max_budget_check( + await organization_max_budget_check( valid_token=token, team_object=None, prisma_client=MagicMock(), @@ -7577,7 +7577,7 @@ async def test_organization_zero_max_budget_is_enforced(max_budget, spend, expec ) assert exc_info.value.max_budget == max_budget else: - await _organization_max_budget_check( + await organization_max_budget_check( valid_token=token, team_object=None, prisma_client=MagicMock(), @@ -8692,11 +8692,11 @@ def _restricted_member_check_deps() -> dict[str, object]: @pytest.mark.asyncio async def test_check_team_member_model_access_fails_closed_when_the_membership_read_hits_a_db_outage(): - from litellm.proxy.auth.auth_checks import _check_team_member_model_access + from litellm.proxy.auth.auth_checks import check_team_member_model_access from litellm.proxy.auth.auth_exception_handler import _as_proxy_exception with pytest.raises(httpx.ConnectError) as raised: - await _check_team_member_model_access( + await check_team_member_model_access( model="claude-sonnet-5", llm_router=None, **_restricted_member_check_deps() ) @@ -9105,7 +9105,7 @@ def test_is_user_proxy_admin_rejects_view_only_admin(): """This predicate skips `non_proxy_admin_allowed_routes_check` entirely, so an Admin Viewer answering True here would gain every write route. Read parity for that role belongs in the route checks, never here.""" - from litellm.proxy.auth.auth_checks import _is_user_proxy_admin + from litellm.proxy.auth.auth_checks import is_user_proxy_admin viewer = LiteLLM_UserTable( user_id="viewer_user", @@ -9118,9 +9118,9 @@ def test_is_user_proxy_admin_rejects_view_only_admin(): user_role=LitellmUserRoles.PROXY_ADMIN.value, ) - assert _is_user_proxy_admin(user_obj=viewer) is False - assert _is_user_proxy_admin(user_obj=admin) is True - assert _is_user_proxy_admin(user_obj=None) is False + assert is_user_proxy_admin(user_obj=viewer) is False + assert is_user_proxy_admin(user_obj=admin) is True + assert is_user_proxy_admin(user_obj=None) is False def _make_wildcard_access_group_router(): @@ -9156,12 +9156,12 @@ def test_can_object_call_model_access_group_wildcard_accepts_bare_model_name(): pattern router's raw regex and skipped the `{provider}/{model}` retry that both routing and the direct-wildcard grant already perform. """ - from litellm.proxy.auth.auth_checks import _can_object_call_model + from litellm.proxy.auth.auth_checks import can_object_call_model router = _make_wildcard_access_group_router() assert ( - _can_object_call_model( + can_object_call_model( model="gpt-4o", llm_router=router, models=["default-models"], @@ -9172,12 +9172,12 @@ def test_can_object_call_model_access_group_wildcard_accepts_bare_model_name(): def test_can_object_call_model_access_group_wildcard_accepts_prefixed_model_name(): - from litellm.proxy.auth.auth_checks import _can_object_call_model + from litellm.proxy.auth.auth_checks import can_object_call_model router = _make_wildcard_access_group_router() assert ( - _can_object_call_model( + can_object_call_model( model="openai/gpt-4o", llm_router=router, models=["default-models"], @@ -9197,12 +9197,12 @@ def test_can_object_call_model_access_group_wildcard_accepts_prefixed_model_name def test_can_object_call_model_access_group_wildcard_does_not_over_grant(model): """The bare-name retry must not turn an access group into a blanket grant.""" from litellm.proxy._types import ProxyException - from litellm.proxy.auth.auth_checks import _can_object_call_model + from litellm.proxy.auth.auth_checks import can_object_call_model router = _make_wildcard_access_group_router() with pytest.raises(ProxyException): - _can_object_call_model( + can_object_call_model( model=model, llm_router=router, models=["default-models"], @@ -9217,7 +9217,7 @@ def test_can_object_call_model_access_group_rejects_unconsumed_namespace(): """ from litellm import Router from litellm.proxy._types import ProxyException - from litellm.proxy.auth.auth_checks import _can_object_call_model + from litellm.proxy.auth.auth_checks import can_object_call_model router = Router( model_list=[ @@ -9233,7 +9233,7 @@ def test_can_object_call_model_access_group_rejects_unconsumed_namespace(): ) assert ( - _can_object_call_model( + can_object_call_model( model="anthropic.claude-3-5-sonnet-20240620-v1:0", llm_router=router, models=["bedrock-models"], @@ -9243,7 +9243,7 @@ def test_can_object_call_model_access_group_rejects_unconsumed_namespace(): ) with pytest.raises(ProxyException): - _can_object_call_model( + can_object_call_model( model="bedrockz/anthropic.claude-3-5-sonnet-20240620-v1:0", llm_router=router, models=["bedrock-models"], @@ -9258,7 +9258,7 @@ def test_can_object_call_model_team_scoped_wildcard_accepts_bare_model_name(): index that needed the same `{provider}/{model}` retry. """ from litellm import Router - from litellm.proxy.auth.auth_checks import _can_object_call_model + from litellm.proxy.auth.auth_checks import can_object_call_model router = Router( model_list=[ @@ -9277,7 +9277,7 @@ def test_can_object_call_model_team_scoped_wildcard_accepts_bare_model_name(): for model in ("gpt-4o", "openai/gpt-4o"): assert ( - _can_object_call_model( + can_object_call_model( model=model, llm_router=router, models=["team-models"], @@ -10211,7 +10211,7 @@ async def test_delete_cache_key_object_is_best_effort_when_the_cache_backend_fai import logging from unittest.mock import AsyncMock, MagicMock - from litellm.proxy.auth.auth_checks import _delete_cache_key_object + from litellm.proxy.auth.auth_checks import delete_cache_key_object hashed_token = "a" * 64 caplog.set_level(logging.WARNING, logger="LiteLLM Proxy") @@ -10223,7 +10223,7 @@ async def test_delete_cache_key_object_is_best_effort_when_the_cache_backend_fai side_effect=Exception("No permissions to access a key") ) - await _delete_cache_key_object( + await delete_cache_key_object( hashed_token=hashed_token, user_api_key_cache=failing_cache, proxy_logging_obj=failing_logging_obj, @@ -10241,7 +10241,7 @@ async def test_delete_cache_key_object_is_best_effort_when_the_cache_backend_fai healthy_logging_obj = MagicMock() healthy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache = AsyncMock() - await _delete_cache_key_object( + await delete_cache_key_object( hashed_token=hashed_token, user_api_key_cache=healthy_cache, proxy_logging_obj=healthy_logging_obj, @@ -10272,7 +10272,7 @@ async def _run_key_budget_check(key_name: str) -> str: max_budget=1.0, ) with pytest.raises(litellm.BudgetExceededError, match="Budget has been exceeded") as exc_info: - await _virtual_key_max_budget_check( + await virtual_key_max_budget_check( valid_token=valid_token, proxy_logging_obj=_BudgetAlertRecorder(), ) @@ -10920,7 +10920,7 @@ async def test_agent_key_without_an_echoed_caller_keeps_its_own_models(): def test_can_object_call_model_allows_listed_model_for_key(): - result: Final = _can_object_call_model( + result: Final = can_object_call_model( model="allowed-model", llm_router=None, models=["allowed-model"], @@ -11113,7 +11113,7 @@ async def test_authoritative_group_grants_propagate_policy_outages( from fastapi import HTTPException from litellm.proxy import proxy_server - from litellm.proxy.auth.auth_checks import _get_agent_ids_from_access_groups + from litellm.proxy.auth.auth_checks import get_agent_ids_from_access_groups from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache database: Final = MagicMock() @@ -11125,6 +11125,6 @@ async def test_authoritative_group_grants_propagate_policy_outages( monkeypatch.setattr(proxy_server, "user_api_key_cache", UserApiKeyCache()) if strict: with pytest.raises(HTTPException): - await _get_agent_ids_from_access_groups(["group"], check_db_only=True) + await get_agent_ids_from_access_groups(["group"], check_db_only=True) else: - assert await _get_agent_ids_from_access_groups(["group"]) == [] + assert await get_agent_ids_from_access_groups(["group"]) == [] diff --git a/tests/unit/proxy/auth/test_auth_exception_handler.py b/tests/unit/proxy/auth/test_auth_exception_handler.py index 521cbd8daad..dd6e1f57fbc 100644 --- a/tests/unit/proxy/auth/test_auth_exception_handler.py +++ b/tests/unit/proxy/auth/test_auth_exception_handler.py @@ -68,7 +68,7 @@ async def test_handle_authentication_error_db_unavailable_connectivity(db_error) "litellm.proxy.proxy_server.general_settings", {"allow_requests_on_db_unavailable": True}, ): - result = await handler._handle_authentication_error( + result = await handler.handle_authentication_error( db_error, mock_request, {}, @@ -113,7 +113,7 @@ async def test_handle_authentication_error_permanent_fault_gets_no_fallback_iden {"allow_requests_on_db_unavailable": True}, ): with pytest.raises(ProxyException) as exc_info: - await handler._handle_authentication_error( + await handler.handle_authentication_error( prisma_error, mock_request, {}, @@ -148,7 +148,7 @@ async def test_handle_authentication_error_permanent_fault_503_is_not_worded_as_ "litellm.proxy.proxy_server.general_settings", {"allow_requests_on_db_unavailable": False} ): with pytest.raises(ProxyException) as exc_info: - await handler._handle_authentication_error(prisma_error, MagicMock(), {}, "/test", None, "test-key") + await handler.handle_authentication_error(prisma_error, MagicMock(), {}, "/test", None, "test-key") assert exc_info.value.code == str(status.HTTP_503_SERVICE_UNAVAILABLE) assert exc_info.value.type == ProxyErrorTypes.no_db_connection @@ -176,7 +176,7 @@ async def test_handle_authentication_error_transport_error_raised_over_a_permane "litellm.proxy.proxy_server.general_settings", {"allow_requests_on_db_unavailable": False} ): with pytest.raises(ProxyException) as exc_info: - await handler._handle_authentication_error(transport_over_fault, MagicMock(), {}, "/test", None, "k") + await handler.handle_authentication_error(transport_over_fault, MagicMock(), {}, "/test", None, "k") assert exc_info.value.code == str(status.HTTP_503_SERVICE_UNAVAILABLE) assert "temporarily unreachable" not in exc_info.value.message @@ -202,7 +202,7 @@ async def test_handle_authentication_error_transient_outage_503_keeps_retry_word "litellm.proxy.proxy_server.general_settings", {"allow_requests_on_db_unavailable": False} ): with pytest.raises(ProxyException) as exc_info: - await handler._handle_authentication_error(db_error, MagicMock(), {}, "/test", None, "test-key") + await handler.handle_authentication_error(db_error, MagicMock(), {}, "/test", None, "test-key") assert exc_info.value.code == str(status.HTTP_503_SERVICE_UNAVAILABLE) assert exc_info.value.message == ( @@ -249,7 +249,7 @@ async def test_handle_authentication_error_data_layer_errors_do_not_fall_back( {"allow_requests_on_db_unavailable": True}, ): with pytest.raises(ProxyException): - await handler._handle_authentication_error( + await handler.handle_authentication_error( prisma_error, mock_request, {}, @@ -300,7 +300,7 @@ async def test_handle_authentication_error_db_infra_error_returns_503(db_error): ), ): with pytest.raises(ProxyException) as exc_info: - await handler._handle_authentication_error( + await handler.handle_authentication_error( db_error, MagicMock(), {}, @@ -357,7 +357,7 @@ async def test_handle_authentication_error_prisma_engine_teardown_returns_503(): ), ): with pytest.raises(ProxyException) as exc_info: - await handler._handle_authentication_error( + await handler.handle_authentication_error( teardown_error, MagicMock(), {}, @@ -407,7 +407,7 @@ async def test_handle_authentication_error_genuine_auth_failure_stays_401(auth_e ), ): with pytest.raises(ProxyException) as exc_info: - await handler._handle_authentication_error( + await handler.handle_authentication_error( auth_error, MagicMock(), {}, @@ -438,7 +438,7 @@ async def test_handle_authentication_error_budget_exceeded(): ) with pytest.raises(ProxyException) as exc_info: - await handler._handle_authentication_error( + await handler.handle_authentication_error( budget_error, mock_request, mock_request_data, @@ -476,7 +476,7 @@ async def test_route_passed_to_post_call_failure_hook(): {"allow_requests_on_db_unavailable": False}, ): try: - await handler._handle_authentication_error( + await handler.handle_authentication_error( PrismaError(), mock_request, mock_request_data, @@ -507,7 +507,7 @@ async def test_dynamic_route_normalized_on_auth_failure(): ), pytest.raises(ProxyException), ): - await handler._handle_authentication_error( + await handler.handle_authentication_error( HTTPException(status_code=401, detail="Authentication Error, Invalid proxy server token passed"), MagicMock(), {}, @@ -567,7 +567,7 @@ async def test_resolved_identity_exported_on_auth_failure(): ), ): with pytest.raises(ProxyException): - await handler._handle_authentication_error( + await handler.handle_authentication_error( expired_key_error, MagicMock(), {"model": "gpt-4o"}, @@ -666,7 +666,7 @@ async def test_expired_key_error_log_names_the_key_owner( verbose_proxy_logger.propagate = True try: with caplog.at_level("ERROR", logger="LiteLLM Proxy"), pytest.raises(ProxyException): - await handler._handle_authentication_error( + await handler.handle_authentication_error( expired_key_error, MagicMock(), {"model": "gpt-4o"}, @@ -709,7 +709,7 @@ async def test_auth_failure_without_resolved_identity_still_logs(): ), ): with pytest.raises(ProxyException): - await handler._handle_authentication_error( + await handler.handle_authentication_error( ProxyException( message="Invalid API key", type=ProxyErrorTypes.auth_error, @@ -1066,7 +1066,7 @@ async def test_handle_authentication_error_traceback_only_for_unexpected_errors( raise auth_error except (ProxyException, ValueError, HTTPException) as caught: with caplog.at_level(expect_level, logger="LiteLLM Proxy"), pytest.raises((ProxyException, HTTPException)): - await handler._handle_authentication_error( + await handler.handle_authentication_error( caught, MagicMock(), {}, @@ -1139,7 +1139,7 @@ async def test_handle_authentication_error_keeps_internal_message_on_model_acces caplog.at_level("WARNING", logger="LiteLLM Proxy"), pytest.raises(ModelAccessDeniedProxyException) as exc_info, ): - await handler._handle_authentication_error(denial, MagicMock(), {}, "/v1/chat/completions", None, "sk-bad-key") + await handler.handle_authentication_error(denial, MagicMock(), {}, "/v1/chat/completions", None, "sk-bad-key") assert exc_info.value.code == str(status.HTTP_403_FORBIDDEN) assert "internal-models" not in str(exc_info.value.message) diff --git a/tests/unit/proxy/auth/test_auth_utils.py b/tests/unit/proxy/auth/test_auth_utils.py index 08757534059..130d29a9560 100644 --- a/tests/unit/proxy/auth/test_auth_utils.py +++ b/tests/unit/proxy/auth/test_auth_utils.py @@ -706,7 +706,7 @@ def _cache_prediction_auth_app( from litellm.caching.dual_cache import DualCache from litellm.proxy._types import LiteLLM_TeamTableCachedObj, LiteLLM_UserTable, LitellmUserRoles, ProxyException from litellm.proxy.auth import auth_checks - from litellm.proxy.hooks.parallel_request_limiter_v3 import _PROXY_MaxParallelRequestsHandler_v3 + from litellm.proxy.hooks.parallel_request_limiter_v3 import PROXY_MaxParallelRequestsHandler_v3 from litellm.proxy.management_endpoints import prompt_cache_prediction as endpoint from litellm.proxy.utils import InternalUsageCache, ProxyLogging @@ -729,7 +729,7 @@ def _cache_prediction_auth_app( ) return token - monkeypatch.setattr(auth, "_user_api_key_auth_builder", authenticate) + monkeypatch.setattr(auth, "user_api_key_auth_builder", authenticate) monkeypatch.setattr(auth, "get_user_object", AsyncMock(return_value=user)) team = LiteLLM_TeamTableCachedObj(team_id=team_id, models=token.team_models) if team_id else None monkeypatch.setattr(auth, "get_team_object", AsyncMock(return_value=team)) @@ -744,7 +744,7 @@ def _cache_prediction_auth_app( monkeypatch.setattr(proxy_server, "prisma_client", None) monkeypatch.setattr(proxy_server, "user_api_key_cache", DualCache()) logging = ProxyLogging(user_api_key_cache=DualCache()) - logging.proxy_hook_mapping["parallel_request_limiter"] = _PROXY_MaxParallelRequestsHandler_v3( + logging.proxy_hook_mapping["parallel_request_limiter"] = PROXY_MaxParallelRequestsHandler_v3( InternalUsageCache(dual_cache=DualCache()) ) monkeypatch.setattr(proxy_server, "proxy_logging_obj", logging) diff --git a/tests/unit/proxy/auth/test_custom_auth_end_user_budget.py b/tests/unit/proxy/auth/test_custom_auth_end_user_budget.py index e83c5cf8419..0c3487255f6 100644 --- a/tests/unit/proxy/auth/test_custom_auth_end_user_budget.py +++ b/tests/unit/proxy/auth/test_custom_auth_end_user_budget.py @@ -151,11 +151,11 @@ async def test_custom_auth_defers_end_user_budget_to_common_checks_when_enabled( return_value=end_user_obj, ), patch( - "litellm.proxy.auth.user_api_key_auth._check_end_user_budget", + "litellm.proxy.auth.user_api_key_auth.check_end_user_budget", new_callable=AsyncMock, ) as mock_check, patch( - "litellm.proxy.auth.user_api_key_auth._enforce_key_and_fallback_model_access", + "litellm.proxy.auth.user_api_key_auth.enforce_key_and_fallback_model_access", new_callable=AsyncMock, ), patch( diff --git a/tests/unit/proxy/auth/test_default_end_user_budget_simple.py b/tests/unit/proxy/auth/test_default_end_user_budget_simple.py index edd0409343a..c58ea4c19c3 100644 --- a/tests/unit/proxy/auth/test_default_end_user_budget_simple.py +++ b/tests/unit/proxy/auth/test_default_end_user_budget_simple.py @@ -137,7 +137,7 @@ async def test_budget_enforcement_blocks_over_budget_users(): Note: Budget enforcement happens in common_checks() via _check_end_user_budget(), not in get_end_user_object(). get_end_user_object only fetches the user data. """ - from litellm.proxy.auth.auth_checks import _check_end_user_budget + from litellm.proxy.auth.auth_checks import check_end_user_budget end_user_id = f"test_user_{uuid.uuid4().hex}" default_budget_id = str(uuid.uuid4()) @@ -187,7 +187,7 @@ async def test_budget_enforcement_blocks_over_budget_users(): # Now test budget enforcement separately via _check_end_user_budget with pytest.raises(litellm.BudgetExceededError) as exc_info: - await _check_end_user_budget( + await check_end_user_budget( end_user_obj=result, route="/chat/completions", ) diff --git a/tests/unit/proxy/auth/test_login_utils.py b/tests/unit/proxy/auth/test_login_utils.py index e095af20b9c..1fb648fe5e6 100644 --- a/tests/unit/proxy/auth/test_login_utils.py +++ b/tests/unit/proxy/auth/test_login_utils.py @@ -639,7 +639,7 @@ class TestEncodeUiSessionJwt: from unittest.mock import MagicMock from litellm.proxy._experimental.mcp_server.byok_oauth_endpoints import ( - _user_id_from_session_cookie, + user_id_from_session_cookie, ) from litellm.proxy.auth.login_utils import encode_ui_session_jwt @@ -649,7 +649,7 @@ class TestEncodeUiSessionJwt: request = MagicMock() request.cookies = {"token": token} with patch("litellm.proxy.proxy_server.master_key", "sk-master-for-tests"): - assert _user_id_from_session_cookie(request) == "cornell-user" + assert user_id_from_session_cookie(request) == "cornell-user" def _throttle( diff --git a/tests/unit/proxy/auth/test_model_checks.py b/tests/unit/proxy/auth/test_model_checks.py index 13171a42cda..88c2d9c684a 100644 --- a/tests/unit/proxy/auth/test_model_checks.py +++ b/tests/unit/proxy/auth/test_model_checks.py @@ -967,6 +967,81 @@ def test_get_complete_model_list_sentinel_only_grants_nothing(): assert result == [] +def test_get_provider_models_admits_providers_without_a_static_catalog(): + """Providers without a static model list are no longer rejected up front. + + With endpoint discovery off, get_valid_models falls back to the (empty) + static list, so the result is [] rather than None. Before the fix this + returned None and the wildcard was never expanded. + """ + import litellm + from litellm.proxy.auth.model_checks import get_provider_models + from litellm.types.router import LiteLLM_Params + + assert "litellm_proxy" not in litellm.models_by_provider + assert "hosted_vllm" not in litellm.models_by_provider + + result = get_provider_models( + "litellm_proxy", + litellm_params=LiteLLM_Params( + model="litellm_proxy/*", + api_base="http://upstream:4000", + api_key="sk-upstream", + ), + ) + + assert result == [] + + +def test_get_complete_model_list_discovers_litellm_proxy_wildcard_models(monkeypatch): + """A litellm_proxy/* deployment lists the upstream proxy's models when endpoint discovery is on.""" + import litellm + from litellm import Router + from litellm.llms.openai.chat.gpt_transformation import OpenAIGPTConfig + from litellm.proxy.auth.model_checks import get_complete_model_list + + monkeypatch.setattr(litellm, "check_provider_endpoint", True) + captured = {} + + def fake_get_models(self, api_key=None, api_base=None): + captured["api_key"] = api_key + captured["api_base"] = api_base + return ["gpt-4o", "claude-sonnet"] + + monkeypatch.setattr(OpenAIGPTConfig, "get_models", fake_get_models) + + router = Router( + model_list=[ + { + "model_name": "litellm_proxy/*", + "litellm_params": { + "model": "litellm_proxy/*", + "api_base": "http://upstream:4000", + "api_key": "sk-upstream", + }, + } + ] + ) + result = get_complete_model_list( + key_models=[], + team_models=[], + proxy_model_list=["litellm_proxy/*"], + user_model=None, + infer_model_from_keys=False, + llm_router=router, + ) + + assert captured == {"api_key": "sk-upstream", "api_base": "http://upstream:4000"} + assert "litellm_proxy/gpt-4o" in result + assert "litellm_proxy/claude-sonnet" in result + + +def test_get_provider_models_returns_none_for_an_unknown_provider(): + from litellm.proxy.auth.model_checks import get_provider_models + + assert get_provider_models("not-a-real-provider") is None + + def test_transcribe_is_a_known_provider_for_wildcard_expansion(): import litellm from litellm.proxy.auth.model_checks import ( diff --git a/tests/unit/proxy/auth/test_object_permission_loading.py b/tests/unit/proxy/auth/test_object_permission_loading.py index 8db4e210107..867ce9c998f 100644 --- a/tests/unit/proxy/auth/test_object_permission_loading.py +++ b/tests/unit/proxy/auth/test_object_permission_loading.py @@ -53,7 +53,7 @@ async def test_get_key_object_loads_object_permission(): "litellm.proxy.auth.auth_checks.get_object_permission", AsyncMock(return_value=mock_object_permission), ), - patch("litellm.proxy.auth.auth_checks._cache_key_object", AsyncMock()), + patch("litellm.proxy.auth.auth_checks.cache_key_object", AsyncMock()), patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj), ): result = await get_key_object( @@ -94,7 +94,7 @@ async def test_get_key_object_no_permission_id(): mock_proxy_logging_obj.service_logging_obj.async_service_failure_hook = AsyncMock() with ( - patch("litellm.proxy.auth.auth_checks._cache_key_object", AsyncMock()), + patch("litellm.proxy.auth.auth_checks.cache_key_object", AsyncMock()), patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj), ): result = await get_key_object( @@ -147,7 +147,7 @@ async def test_get_team_object_loads_object_permission(): "litellm.proxy.auth.auth_checks.get_object_permission", AsyncMock(return_value=mock_object_permission), ), - patch("litellm.proxy.auth.auth_checks._cache_team_object", AsyncMock()), + patch("litellm.proxy.auth.auth_checks.cache_team_object", AsyncMock()), patch("litellm.proxy.auth.auth_checks._should_check_db", return_value=True), patch("litellm.proxy.auth.auth_checks._update_last_db_access_time"), patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj), diff --git a/tests/unit/proxy/auth/test_proxy_routes.py b/tests/unit/proxy/auth/test_proxy_routes.py index 129a93ea08d..24ef2a0e150 100644 --- a/tests/unit/proxy/auth/test_proxy_routes.py +++ b/tests/unit/proxy/auth/test_proxy_routes.py @@ -246,9 +246,9 @@ def _is_assistants(req): def _metadata_var_name(req): - from litellm.proxy.litellm_pre_call_utils import _get_metadata_variable_name + from litellm.proxy.litellm_pre_call_utils import get_metadata_variable_name - return _get_metadata_variable_name(req) + return get_metadata_variable_name(req) def _vector_store_id_in_path(req): diff --git a/tests/unit/proxy/auth/test_route_checks.py b/tests/unit/proxy/auth/test_route_checks.py index 36799b73876..540b66ddaba 100644 --- a/tests/unit/proxy/auth/test_route_checks.py +++ b/tests/unit/proxy/auth/test_route_checks.py @@ -14,7 +14,7 @@ from litellm.proxy._types import ( UserAPIKeyAuth, ) from litellm.proxy.auth.auth_checks import _is_api_route_allowed -from litellm.proxy.auth.auth_checks_organization import _user_is_org_admin +from litellm.proxy.auth.auth_checks_organization import user_is_org_admin from litellm.proxy.auth.route_checks import RouteChecks from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import router as llm_passthrough_router @@ -2946,7 +2946,7 @@ def test_available_roles_accessible_to_non_admin_users(user_role): ) -# ── _user_is_org_admin tests ────────────────────────────────────────────────── +# ── user_is_org_admin tests ────────────────────────────────────────────────── def _make_org_admin_user(org_id: str) -> LiteLLM_UserTable: @@ -2967,25 +2967,25 @@ def _make_org_admin_user(org_id: str) -> LiteLLM_UserTable: def test_user_is_org_admin_with_organizations_list(): """Org admin can be identified via the `organizations` list field (used by /user/new).""" user_obj = _make_org_admin_user("org-1") - assert _user_is_org_admin({"organizations": ["org-1"]}, user_obj) is True + assert user_is_org_admin({"organizations": ["org-1"]}, user_obj) is True def test_user_is_org_admin_with_singular_organization_id(): """Backward-compat: org admin can still be identified via singular `organization_id`.""" user_obj = _make_org_admin_user("org-1") - assert _user_is_org_admin({"organization_id": "org-1"}, user_obj) is True + assert user_is_org_admin({"organization_id": "org-1"}, user_obj) is True def test_user_is_org_admin_organizations_list_wrong_org(): """Non-member of the requested org is not considered an org admin for it.""" user_obj = _make_org_admin_user("org-2") - assert _user_is_org_admin({"organizations": ["org-1"]}, user_obj) is False + assert user_is_org_admin({"organizations": ["org-1"]}, user_obj) is False def test_user_is_org_admin_no_org_fields(): """Returns False when neither `organization_id` nor `organizations` is in the request.""" user_obj = _make_org_admin_user("org-1") - assert _user_is_org_admin({}, user_obj) is False + assert user_is_org_admin({}, user_obj) is False def test_non_org_admin_with_organizations_list(): @@ -3002,13 +3002,13 @@ def test_non_org_admin_with_organizations_list(): user_role=LitellmUserRoles.INTERNAL_USER.value, organization_memberships=[membership], ) - assert _user_is_org_admin({"organizations": ["org-1"]}, user_obj) is False + assert user_is_org_admin({"organizations": ["org-1"]}, user_obj) is False def test_org_admin_cannot_escalate_to_other_org(): """Regression: admin of org-A requesting [org-A, org-B] must be rejected.""" user_obj = _make_org_admin_user("org-A") - assert _user_is_org_admin({"organizations": ["org-A", "org-B"]}, user_obj) is False + assert user_is_org_admin({"organizations": ["org-A", "org-B"]}, user_obj) is False def test_org_admin_of_multiple_orgs_can_operate_on_both(): @@ -3034,7 +3034,7 @@ def test_org_admin_of_multiple_orgs_can_operate_on_both(): user_role=LitellmUserRoles.INTERNAL_USER.value, organization_memberships=memberships, ) - assert _user_is_org_admin({"organizations": ["org-A", "org-B"]}, user_obj) is True + assert user_is_org_admin({"organizations": ["org-A", "org-B"]}, user_obj) is True # ── LIT-4221: /team/update org-context resolution from team_id ──────────────── @@ -3742,7 +3742,7 @@ def test_organization_daily_activity_not_granted_by_org_admin_request_data_branc self_managed_routes entry is load-bearing rather than redundant. Query params do reach request_data, so the reason is not body-vs-query: it - is the key name. _user_is_org_admin reads ``organization_id`` (singular) and + is the key name. user_is_org_admin reads ``organization_id`` (singular) and ``organizations``, while this endpoint's filter is ``organization_ids`` (plural), and the dashboard's first page load sends no organization filter at all. Both shapes are pinned below because renaming the query param would @@ -3764,11 +3764,11 @@ def test_organization_daily_activity_not_granted_by_org_admin_request_data_branc ) # The dashboard's default page load: no organization filter at all. - assert not _user_is_org_admin(request_data={}, user_object=user_obj) + assert not user_is_org_admin(request_data={}, user_object=user_obj) # The filtered load, naming an org this user really does administer. - assert not _user_is_org_admin(request_data={"organization_ids": "org-a"}, user_object=user_obj) + assert not user_is_org_admin(request_data={"organization_ids": "org-a"}, user_object=user_obj) # The key name the helper would have had to see to grant it. - assert _user_is_org_admin(request_data={"organization_id": "org-a"}, user_object=user_obj) + assert user_is_org_admin(request_data={"organization_id": "org-a"}, user_object=user_obj) assert not RouteChecks.check_route_access( route="/organization/daily/activity", allowed_routes=LiteLLMRoutes.org_admin_only_routes.value, diff --git a/tests/unit/proxy/auth/test_router_override_fallback_auth.py b/tests/unit/proxy/auth/test_router_override_fallback_auth.py index c34614a93ed..d35a5851659 100644 --- a/tests/unit/proxy/auth/test_router_override_fallback_auth.py +++ b/tests/unit/proxy/auth/test_router_override_fallback_auth.py @@ -12,7 +12,7 @@ import pytest from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth from litellm.proxy.auth.auth_utils import fallback_target_model_name, iter_request_fallback_targets -from litellm.proxy.auth.user_api_key_auth import _enforce_key_and_fallback_model_access +from litellm.proxy.auth.user_api_key_auth import enforce_key_and_fallback_model_access def _fallback_model_names(fallbacks): @@ -100,7 +100,7 @@ async def test_router_override_fallbacks_validated_against_key_allowlist(): new=AsyncMock(), ), ): - await _enforce_key_and_fallback_model_access( + await enforce_key_and_fallback_model_access( valid_token=valid_token, request_data=request_data, route="/v1/chat/completions", @@ -151,7 +151,7 @@ async def test_router_override_all_fallback_fields_validated(fallback_field): new=AsyncMock(), ), ): - await _enforce_key_and_fallback_model_access( + await enforce_key_and_fallback_model_access( valid_token=valid_token, request_data=request_data, route="/v1/chat/completions", @@ -198,7 +198,7 @@ async def test_top_level_fallback_fields_validated(fallback_field): new=AsyncMock(), ), ): - await _enforce_key_and_fallback_model_access( + await enforce_key_and_fallback_model_access( valid_token=valid_token, request_data=request_data, route="/v1/chat/completions", @@ -244,7 +244,7 @@ async def test_nested_deployment_fallback_inner_model_validated(): new=AsyncMock(), ), ): - await _enforce_key_and_fallback_model_access( + await enforce_key_and_fallback_model_access( valid_token=valid_token, request_data=request_data, route="/v1/chat/completions", @@ -289,7 +289,7 @@ async def test_model_less_fallback_dict_is_skipped_never_passed_as_none(): new=AsyncMock(), ), ): - await _enforce_key_and_fallback_model_access( + await enforce_key_and_fallback_model_access( valid_token=valid_token, request_data=request_data, route="/v1/chat/completions", @@ -327,7 +327,7 @@ async def test_router_override_without_fallbacks_does_not_break_auth(): new=AsyncMock(), ), ): - await _enforce_key_and_fallback_model_access( + await enforce_key_and_fallback_model_access( valid_token=valid_token, request_data=request_data, route="/v1/chat/completions", diff --git a/tests/unit/proxy/auth/test_unmapped_model_budget_enforcement.py b/tests/unit/proxy/auth/test_unmapped_model_budget_enforcement.py index be7b438a442..c2851f608f0 100644 --- a/tests/unit/proxy/auth/test_unmapped_model_budget_enforcement.py +++ b/tests/unit/proxy/auth/test_unmapped_model_budget_enforcement.py @@ -12,7 +12,7 @@ import copy import pytest import litellm -from litellm.proxy.auth.auth_checks import _is_model_cost_zero +from litellm.proxy.auth.auth_checks import is_model_cost_zero from litellm.router import Router @@ -88,7 +88,7 @@ class TestUnmappedModelBudgetEnforcement: }, ] ) - result = _is_model_cost_zero(model="custom-model", llm_router=router) + result = is_model_cost_zero(model="custom-model", llm_router=router) assert result is False, "Unmapped model should enforce budget (return False), not bypass it (return True)" def test_explicitly_free_model_bypasses_budget(self): @@ -111,7 +111,7 @@ class TestUnmappedModelBudgetEnforcement: }, ] ) - result = _is_model_cost_zero(model="free-model", llm_router=router) + result = is_model_cost_zero(model="free-model", llm_router=router) assert result is True, "Explicitly free model should bypass budget (return True)" def test_known_paid_model_enforces_budget(self): @@ -127,7 +127,7 @@ class TestUnmappedModelBudgetEnforcement: }, ] ) - result = _is_model_cost_zero(model="paid-model", llm_router=router) + result = is_model_cost_zero(model="paid-model", llm_router=router) assert result is False, "Known paid model should enforce budget (return False)" def test_unmapped_model_with_litellm_params_pricing(self): @@ -145,7 +145,7 @@ class TestUnmappedModelBudgetEnforcement: }, ] ) - result = _is_model_cost_zero(model="free-via-params", llm_router=router) + result = is_model_cost_zero(model="free-via-params", llm_router=router) assert result is True, "Model with explicit cost=0 in litellm_params should bypass budget" def test_cache_invalidates_on_in_place_pricing_update(self): @@ -177,7 +177,7 @@ class TestUnmappedModelBudgetEnforcement: ] ) # Warm the cache as zero-cost. - assert _is_model_cost_zero(model="ramping-model", llm_router=router) is True + assert is_model_cost_zero(model="ramping-model", llm_router=router) is True assert router._zero_cost_cache.get("ramping-model") is True # In-place pricing update: same deployment count, same router id, @@ -203,7 +203,7 @@ class TestUnmappedModelBudgetEnforcement: # Cache must have been cleared by ``_invalidate_model_group_info_cache``. assert router._zero_cost_cache == {} # Subsequent call sees the new pricing and enforces budget. - assert _is_model_cost_zero(model="ramping-model", llm_router=router) is False + assert is_model_cost_zero(model="ramping-model", llm_router=router) is False def test_strategy_router_alias_with_zero_pricing_enforces_budget(self): """An auto-router alias is never the deployment that gets called or @@ -231,7 +231,7 @@ class TestUnmappedModelBudgetEnforcement: ) assert "input_cost_per_token" not in litellm.model_cost.get("alias-id", {}) - assert _is_model_cost_zero(model="smart-router", llm_router=router) is False + assert is_model_cost_zero(model="smart-router", llm_router=router) is False def test_model_group_alias_to_free_model_bypasses_budget(self): """A zero-cost group reached through model_group_alias bypasses budget, like its own name. @@ -255,8 +255,8 @@ class TestUnmappedModelBudgetEnforcement: model_group_alias={"free-model-alias": "free-model"}, ) - assert _is_model_cost_zero(model="free-model", llm_router=router) is True - assert _is_model_cost_zero(model="free-model-alias", llm_router=router) is True, ( + assert is_model_cost_zero(model="free-model", llm_router=router) is True + assert is_model_cost_zero(model="free-model-alias", llm_router=router) is True, ( "An alias pointing at an explicitly-zero-cost group must be read as free, like its own name" ) @@ -278,7 +278,7 @@ class TestUnmappedModelBudgetEnforcement: model_group_alias={"free-model-alias": {"model": "free-model", "hidden": False}}, ) - assert _is_model_cost_zero(model="free-model-alias", llm_router=router) is True + assert is_model_cost_zero(model="free-model-alias", llm_router=router) is True def test_model_group_alias_to_paid_model_enforces_budget(self): """An alias does not turn a priced group into a free one.""" @@ -293,7 +293,7 @@ class TestUnmappedModelBudgetEnforcement: model_group_alias={"paid-model-alias": "paid-model"}, ) - assert _is_model_cost_zero(model="paid-model-alias", llm_router=router) is False + assert is_model_cost_zero(model="paid-model-alias", llm_router=router) is False def test_model_group_alias_to_ptu_flat_cost_enforces_budget(self): """A PTU group keeps budget enforced through an alias. @@ -323,8 +323,8 @@ class TestUnmappedModelBudgetEnforcement: model_group_alias={"ptu-model-alias": "ptu-model"}, ) - assert _is_model_cost_zero(model="ptu-model", llm_router=router) is False - assert _is_model_cost_zero(model="ptu-model-alias", llm_router=router) is False, ( + assert is_model_cost_zero(model="ptu-model", llm_router=router) is False + assert is_model_cost_zero(model="ptu-model-alias", llm_router=router) is False, ( "An aliased PTU group must not be read as free" ) @@ -350,7 +350,7 @@ class TestUnmappedModelBudgetEnforcement: model_group_alias={"hidden-alias": {"model": "free-model", "hidden": True}}, ) - assert _is_model_cost_zero(model="hidden-alias", llm_router=router) is True + assert is_model_cost_zero(model="hidden-alias", llm_router=router) is True def test_hidden_model_group_alias_to_paid_model_enforces_budget(self): """A hidden alias to a priced group keeps budget enforced.""" @@ -365,7 +365,7 @@ class TestUnmappedModelBudgetEnforcement: model_group_alias={"hidden-paid-alias": {"model": "paid-model", "hidden": True}}, ) - assert _is_model_cost_zero(model="hidden-paid-alias", llm_router=router) is False + assert is_model_cost_zero(model="hidden-paid-alias", llm_router=router) is False def test_dangling_model_group_alias_enforces_budget(self): """An alias pointing at a group that does not exist keeps budget enforced.""" @@ -385,7 +385,7 @@ class TestUnmappedModelBudgetEnforcement: model_group_alias={"dangling-alias": "model-that-does-not-exist"}, ) - assert _is_model_cost_zero(model="dangling-alias", llm_router=router) is False + assert is_model_cost_zero(model="dangling-alias", llm_router=router) is False def test_repointed_hidden_alias_does_not_reuse_cached_free_result(self): """Repointing a hidden alias from a free group to a paid group re-evaluates the cost. @@ -415,9 +415,9 @@ class TestUnmappedModelBudgetEnforcement: model_group_alias={"hidden-alias": {"model": "free-model", "hidden": True}}, ) - assert _is_model_cost_zero(model="hidden-alias", llm_router=router) is True + assert is_model_cost_zero(model="hidden-alias", llm_router=router) is True router.update_settings(model_group_alias={"hidden-alias": {"model": "paid-model", "hidden": True}}) - assert _is_model_cost_zero(model="hidden-alias", llm_router=router) is False + assert is_model_cost_zero(model="hidden-alias", llm_router=router) is False @pytest.mark.parametrize("alias_name_first", [True, False]) def test_alias_shadowing_a_real_group_answers_for_its_target_in_either_order(self, alias_name_first: bool): @@ -456,10 +456,10 @@ class TestUnmappedModelBudgetEnforcement: order = ("ptu-model", "free-model") if alias_name_first else ("free-model", "ptu-model") expected = {"ptu-model": True, "free-model": True} - assert [_is_model_cost_zero(model=name, llm_router=router) for name in order] == [ + assert [is_model_cost_zero(model=name, llm_router=router) for name in order] == [ expected[name] for name in order ] - assert [_is_model_cost_zero(model=name, llm_router=router) for name in order] == [ + assert [is_model_cost_zero(model=name, llm_router=router) for name in order] == [ expected[name] for name in order ], "the cached verdicts must match the first evaluation" @@ -508,8 +508,8 @@ class TestUnmappedModelBudgetEnforcement: model_group_alias={"free-model": alias_entry}, ) - assert _is_model_cost_zero(model="unpriced-target", llm_router=router) is False - assert _is_model_cost_zero(model="free-model", llm_router=router) is False, ( + assert is_model_cost_zero(model="unpriced-target", llm_router=router) is False + assert is_model_cost_zero(model="free-model", llm_router=router) is False, ( "the alias routes to the unpriced target, so it must be refused like the target by name" ) @@ -546,8 +546,8 @@ class TestUnmappedModelBudgetEnforcement: model_group_alias={"ptu-model": {"model": "free-model", "hidden": True}}, ) - assert _is_model_cost_zero(model="free-model", llm_router=router) is True - assert _is_model_cost_zero(model="ptu-model", llm_router=router) is True, ( + assert is_model_cost_zero(model="free-model", llm_router=router) is True + assert is_model_cost_zero(model="ptu-model", llm_router=router) is True, ( "the alias routes to the free target, so it must bypass budget like the target by name" ) @@ -579,7 +579,7 @@ class TestUnmappedModelBudgetEnforcement: model_group_alias={"free-model": {"model": "paid-model", "hidden": True}}, ) - assert _is_model_cost_zero(model="free-model", llm_router=router) is False + assert is_model_cost_zero(model="free-model", llm_router=router) is False def test_alias_chain_through_a_priced_group_enforces_budget(self): """An alias to a group that is itself an alias key resolves one hop, like the router does. @@ -613,7 +613,7 @@ class TestUnmappedModelBudgetEnforcement: model_group_alias={"chain-smart": "chain-legacy", "chain-legacy": "free-model"}, ) - assert _is_model_cost_zero(model="chain-smart", llm_router=router) is False + assert is_model_cost_zero(model="chain-smart", llm_router=router) is False @pytest.mark.parametrize("hidden", [False, True], ids=["plain_alias", "hidden_alias"]) def test_alias_to_an_unpriced_group_that_is_also_an_alias_enforces_budget(self, hidden: bool): @@ -634,8 +634,8 @@ class TestUnmappedModelBudgetEnforcement: ) assert _served_model(router, "chain-entry") == UNPRICED_ZERO_COST_MODEL - assert _is_model_cost_zero(model="chain-entry", llm_router=router) is False - assert _is_model_cost_zero(model="chain-middle", llm_router=router) is True, ( + assert is_model_cost_zero(model="chain-entry", llm_router=router) is False + assert is_model_cost_zero(model="chain-middle", llm_router=router) is True, ( "asked by its own name, chain-middle routes to the explicitly free group" ) @@ -655,8 +655,8 @@ class TestUnmappedModelBudgetEnforcement: ) assert _served_model(router, "chain-entry") == "gpt-3.5-turbo" - assert _is_model_cost_zero(model="chain-entry", llm_router=router) is True - assert _is_model_cost_zero(model="chain-middle", llm_router=router) is False, ( + assert is_model_cost_zero(model="chain-entry", llm_router=router) is True + assert is_model_cost_zero(model="chain-middle", llm_router=router) is False, ( "asked by its own name, chain-middle routes to the PTU group" ) @@ -676,7 +676,7 @@ class TestUnmappedModelBudgetEnforcement: ) assert _served_model(router, "chain-entry") == "openai/gpt-4o-mini" - assert _is_model_cost_zero(model="chain-entry", llm_router=router) is False + assert is_model_cost_zero(model="chain-entry", llm_router=router) is False @pytest.mark.parametrize("alias_name", ["openai/smart", "smart"], ids=["alias_on_pattern", "alias_off_pattern"]) def test_alias_chain_served_by_an_explicitly_priced_wildcard_route_enforces_budget(self, alias_name: str): @@ -704,7 +704,7 @@ class TestUnmappedModelBudgetEnforcement: ) assert _served_model(router, alias_name) == "openai/gpt-4o-mini" - assert _is_model_cost_zero(model=alias_name, llm_router=router) is False + assert is_model_cost_zero(model=alias_name, llm_router=router) is False def test_alias_shadowing_a_free_group_is_judged_by_its_unpriced_target_through_an_alias_chain(self): """A shadowing alias stays enforced when its unpriced target is itself an alias key to a free group.""" @@ -718,7 +718,7 @@ class TestUnmappedModelBudgetEnforcement: ) assert _served_model(router, "shadowed-free") == UNPRICED_ZERO_COST_MODEL - assert _is_model_cost_zero(model="shadowed-free", llm_router=router) is False + assert is_model_cost_zero(model="shadowed-free", llm_router=router) is False @pytest.mark.parametrize("alias_name", ["ollama/fast", "fast"], ids=["alias_on_pattern", "alias_off_pattern"]) def test_alias_to_a_name_served_by_an_explicitly_free_wildcard_route_bypasses_budget(self, alias_name: str): @@ -729,8 +729,8 @@ class TestUnmappedModelBudgetEnforcement: ) assert _served_model(router, alias_name) == "ollama/llama3" - assert _is_model_cost_zero(model="ollama/llama3", llm_router=router) is True - assert _is_model_cost_zero(model=alias_name, llm_router=router) is True + assert is_model_cost_zero(model="ollama/llama3", llm_router=router) is True + assert is_model_cost_zero(model=alias_name, llm_router=router) is True def test_alias_chain_to_a_name_served_by_an_explicitly_free_wildcard_route_bypasses_budget(self): """An alias to an alias key no deployment is named after reads the wildcard route serving it. @@ -745,8 +745,8 @@ class TestUnmappedModelBudgetEnforcement: ) assert _served_model(router, "ollama/fast") == "ollama/llama3" - assert _is_model_cost_zero(model="ollama/fast", llm_router=router) is True - assert _is_model_cost_zero(model="ollama/llama3", llm_router=router) is False, ( + assert is_model_cost_zero(model="ollama/fast", llm_router=router) is True + assert is_model_cost_zero(model="ollama/llama3", llm_router=router) is False, ( "asked by its own name, ollama/llama3 routes to the unpriced group" ) @@ -769,5 +769,5 @@ class TestUnmappedModelBudgetEnforcement: # Strip the attribute so the helper falls back to the no-cache path. del mock_router._zero_cost_cache - result = _is_model_cost_zero(model="paid-model", llm_router=mock_router) + result = is_model_cost_zero(model="paid-model", llm_router=mock_router) assert result is False diff --git a/tests/unit/proxy/auth/test_user_api_key_auth.py b/tests/unit/proxy/auth/test_user_api_key_auth.py index e84ecea2acd..a58599d28e9 100644 --- a/tests/unit/proxy/auth/test_user_api_key_auth.py +++ b/tests/unit/proxy/auth/test_user_api_key_auth.py @@ -572,7 +572,7 @@ def test_allowed_route_inside_route(user_role, auth_user_id, requested_user_id, def test_read_request_body(): - from litellm.proxy.common_utils.http_parsing_utils import _read_request_body + from litellm.proxy.common_utils.http_parsing_utils import read_request_body from fastapi import Request payload = "()" * 1000000 @@ -582,7 +582,7 @@ def test_read_request_body(): return payload request.body = return_body - result = _read_request_body(request) + result = read_request_body(request) assert result is not None @@ -814,9 +814,9 @@ def test_is_allowed_route(): ], ) def test_is_user_proxy_admin(user_obj, expected_result): - from litellm.proxy.auth.auth_checks import _is_user_proxy_admin + from litellm.proxy.auth.auth_checks import is_user_proxy_admin - assert _is_user_proxy_admin(user_obj) == expected_result + assert is_user_proxy_admin(user_obj) == expected_result @pytest.mark.parametrize( @@ -846,9 +846,9 @@ def test_is_user_proxy_admin(user_obj, expected_result): ], ) def test_get_user_role(user_obj, expected_role): - from litellm.proxy.auth.user_api_key_auth import _get_user_role + from litellm.proxy.auth.auth_checks import get_user_role - assert _get_user_role(user_obj) == expected_role + assert get_user_role(user_obj) == expected_role @pytest.mark.asyncio diff --git a/tests/unit/proxy/auth/test_user_api_key_auth_request_flow.py b/tests/unit/proxy/auth/test_user_api_key_auth_request_flow.py index 5de1d9300fe..4ba2c2c56f2 100644 --- a/tests/unit/proxy/auth/test_user_api_key_auth_request_flow.py +++ b/tests/unit/proxy/auth/test_user_api_key_auth_request_flow.py @@ -45,7 +45,7 @@ from litellm.proxy.auth.auth_checks import ( TeamNotFoundError, UserNotFoundError, get_key_object, - _cache_key_object, + cache_key_object, jwt_key_mapping_cache_key, ) from litellm.proxy.auth.route_checks import RouteChecks @@ -59,9 +59,9 @@ from litellm.proxy.auth.user_api_key_auth import ( _reserve_budget_after_common_checks, _route_requires_auth_despite_public, _routing_selector_matches_claim, - _run_centralized_common_checks, + run_centralized_common_checks, _run_post_custom_auth_checks, - _user_api_key_auth_builder, + user_api_key_auth_builder, get_api_key, user_api_key_auth, user_api_key_auth_websocket_for_model, @@ -311,7 +311,7 @@ async def test_should_not_reuse_cached_key_object_for_request_state(): }, ) - await _cache_key_object( + await cache_key_object( hashed_token="cached-token", user_api_key_obj=cached_key, user_api_key_cache=key_cache, @@ -528,7 +528,7 @@ async def test_user_custom_auth_skips_post_custom_auth_checks_by_default(): import litellm import litellm.proxy.proxy_server as _proxy_server_mod from litellm.proxy._types import LitellmUserRoles - from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth_builder trusted_token = UserAPIKeyAuth( api_key="sk-custom-auth-trusted", @@ -553,7 +553,7 @@ async def test_user_custom_auth_skips_post_custom_auth_checks_by_default(): request = Request(scope={"type": "http"}) request._url = URL(url="/chat/completions") - result = await _user_api_key_auth_builder( + result = await user_api_key_auth_builder( request=request, api_key="Bearer sk-custom-auth-trusted", azure_api_key_header="", @@ -586,7 +586,7 @@ async def test_user_custom_auth_runs_post_custom_auth_checks_when_opt_in(): import litellm import litellm.proxy.proxy_server as _proxy_server_mod from litellm.proxy._types import LitellmUserRoles - from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth_builder trusted_token = UserAPIKeyAuth( api_key="sk-custom-auth-trusted", @@ -612,7 +612,7 @@ async def test_user_custom_auth_runs_post_custom_auth_checks_when_opt_in(): request = Request(scope={"type": "http"}) request._url = URL(url="/chat/completions") - await _user_api_key_auth_builder( + await user_api_key_auth_builder( request=request, api_key="Bearer sk-custom-auth-trusted", azure_api_key_header="", @@ -643,7 +643,7 @@ async def test_enterprise_custom_auth_skips_post_custom_auth_checks_by_default() import litellm import litellm.proxy.proxy_server as _proxy_server_mod from litellm.proxy._types import LitellmUserRoles - from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth_builder trusted_token = UserAPIKeyAuth( api_key="sk-enterprise-custom-auth-trusted", @@ -674,7 +674,7 @@ async def test_enterprise_custom_auth_skips_post_custom_auth_checks_by_default() request = Request(scope={"type": "http"}) request._url = URL(url="/chat/completions") - result = await _user_api_key_auth_builder( + result = await user_api_key_auth_builder( request=request, api_key="Bearer sk-enterprise-custom-auth-trusted", azure_api_key_header="", @@ -706,7 +706,7 @@ async def test_enterprise_custom_auth_runs_post_custom_auth_checks_when_opt_in() import litellm import litellm.proxy.proxy_server as _proxy_server_mod from litellm.proxy._types import LitellmUserRoles - from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth_builder trusted_token = UserAPIKeyAuth( api_key="sk-enterprise-custom-auth-trusted", @@ -738,7 +738,7 @@ async def test_enterprise_custom_auth_runs_post_custom_auth_checks_when_opt_in() request = Request(scope={"type": "http"}) request._url = URL(url="/chat/completions") - await _user_api_key_auth_builder( + await user_api_key_auth_builder( request=request, api_key="Bearer sk-enterprise-custom-auth-trusted", azure_api_key_header="", @@ -1158,7 +1158,7 @@ async def test_proxy_admin_expired_key_from_cache(): ProxyException, UserAPIKeyAuth, ) - from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth_builder from litellm.proxy.proxy_server import hash_token # Create an expired PROXY_ADMIN key @@ -1196,7 +1196,7 @@ async def test_proxy_admin_expired_key_from_cache(): new_callable=AsyncMock, ) as mock_get_key_object, patch( - "litellm.proxy.auth.user_api_key_auth._delete_cache_key_object", + "litellm.proxy.auth.user_api_key_auth.delete_cache_key_object", new_callable=AsyncMock, ) as mock_delete_cache, ): @@ -1232,7 +1232,7 @@ async def test_proxy_admin_expired_key_from_cache(): # Call the auth builder - should raise ProxyException for expired key # Note: api_key needs "Bearer " prefix for get_api_key() to process it correctly with pytest.raises(ProxyException) as exc_info: - await _user_api_key_auth_builder( + await user_api_key_auth_builder( request=request, api_key=f"Bearer {api_key}", # Add Bearer prefix azure_api_key_header="", @@ -1281,7 +1281,7 @@ async def test_scim_deactivated_user_key_is_rejected(): from fastapi import Request from starlette.datastructures import URL - from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth_builder from litellm.proxy.proxy_server import hash_token api_key = "sk-scim-deactivated-user-key" @@ -1346,7 +1346,7 @@ async def test_scim_deactivated_user_key_is_rejected(): ), ): with pytest.raises(ProxyException) as exc_info: - await _user_api_key_auth_builder( + await user_api_key_auth_builder( request=request, api_key=f"Bearer {api_key}", azure_api_key_header="", @@ -1371,7 +1371,7 @@ async def test_cached_proxy_admin_key_sets_via_virtual_key_marker(): from fastapi import Request from starlette.datastructures import URL - from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth_builder from litellm.proxy.proxy_server import hash_token api_key = "sk-cached-admin-marker-test" @@ -1424,7 +1424,7 @@ async def test_cached_proxy_admin_key_sets_via_virtual_key_marker(): new_callable=AsyncMock, return_value=cached_token, ): - result = await _user_api_key_auth_builder( + result = await user_api_key_auth_builder( request=request, api_key=f"Bearer {api_key}", azure_api_key_header="", @@ -1452,7 +1452,7 @@ async def test_master_key_auth_sets_via_virtual_key_marker(): from starlette.datastructures import URL from litellm.constants import LITELLM_PROXY_MASTER_KEY_ALIAS - from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth_builder master_key = "sk-master-key" @@ -1490,7 +1490,7 @@ async def test_master_key_auth_sets_via_virtual_key_marker(): request = Request(scope={"type": "http"}) request._url = URL(url="/chat/completions") - result = await _user_api_key_auth_builder( + result = await user_api_key_auth_builder( request=request, api_key=f"Bearer {master_key}", azure_api_key_header="", @@ -1516,7 +1516,7 @@ async def test_db_virtual_key_auth_sets_via_virtual_key_marker(): from fastapi import Request from starlette.datastructures import URL - from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth_builder from litellm.proxy.proxy_server import hash_token api_key = "sk-via-virtual-key-marker-test" @@ -1576,7 +1576,7 @@ async def test_db_virtual_key_auth_sets_via_virtual_key_marker(): return_value=None, ), ): - result = await _user_api_key_auth_builder( + result = await user_api_key_auth_builder( request=request, api_key=f"Bearer {api_key}", azure_api_key_header="", @@ -1601,7 +1601,7 @@ async def test_auth_prefetches_referenced_objects_only_after_the_key_may_call_th from fastapi import Request from starlette.datastructures import URL - from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth_builder from litellm.proxy.proxy_server import hash_token api_key = "sk-prefetch-order-test" @@ -1649,7 +1649,7 @@ async def test_auth_prefetches_referenced_objects_only_after_the_key_may_call_th return_value=valid_token, ), patch( # test-quality-ok: the observable is whether the prefetch runs before or after this check - "litellm.proxy.auth.user_api_key_auth._enforce_key_and_fallback_model_access", + "litellm.proxy.auth.user_api_key_auth.enforce_key_and_fallback_model_access", new_callable=AsyncMock, side_effect=None if model_allowed else denied, ), @@ -1660,7 +1660,7 @@ async def test_auth_prefetches_referenced_objects_only_after_the_key_may_call_th "litellm.proxy.auth.user_api_key_auth.get_user_object", new_callable=AsyncMock, return_value=None ), ): - call = _user_api_key_auth_builder( + call = user_api_key_auth_builder( request=request, api_key=f"Bearer {api_key}", azure_api_key_header="", @@ -1868,7 +1868,7 @@ async def test_standard_jwt_auth_propagates_user_email(): return_value=mock_jwt_result, ), ): - result = await _user_api_key_auth_builder( + result = await user_api_key_auth_builder( request=mock_request, api_key=jwt_token, azure_api_key_header="", @@ -1938,7 +1938,7 @@ async def test_jwt_auth_propagates_agent_id_to_user_api_key_auth(is_proxy_admin: return_value=mock_jwt_result, ), ): - result = await _user_api_key_auth_builder( + result = await user_api_key_auth_builder( request=mock_request, api_key=jwt_token, azure_api_key_header="", @@ -2452,7 +2452,7 @@ async def test_auto_register_map_existing_key_first_request_runs_key_checks( return_value=reused_key, ), ): - call = _user_api_key_auth_builder( + call = user_api_key_auth_builder( request=mock_request, api_key=jwt_token, azure_api_key_header="", @@ -2562,7 +2562,7 @@ async def test_auto_register_first_request_propagates_user_email(active: bool) - ): if not active: with pytest.raises(ProxyException, match="deactivated via SCIM") as exc: - await _user_api_key_auth_builder( + await user_api_key_auth_builder( request=mock_request, api_key=jwt_token, azure_api_key_header="", @@ -2574,7 +2574,7 @@ async def test_auto_register_first_request_propagates_user_email(active: bool) - assert int(exc.value.code) == 401 auto_register.assert_not_awaited() return - result = await _user_api_key_auth_builder( + result = await user_api_key_auth_builder( request=mock_request, api_key=jwt_token, azure_api_key_header="", @@ -2791,7 +2791,7 @@ async def test_jwt_auto_register_forwards_bound_agent_id(): auto_register, ), ): - result = await _user_api_key_auth_builder( + result = await user_api_key_auth_builder( request=mock_request, api_key=jwt_token, azure_api_key_header="", @@ -3107,7 +3107,7 @@ class TestJWTOAuth2Coexistence: return_value=auto_registered_key, ) as mock_auto_register, ): - result = await _user_api_key_auth_builder( + result = await user_api_key_auth_builder( request=mock_request, api_key=jwt_token, azure_api_key_header="", @@ -3187,7 +3187,7 @@ class TestJWTOAuth2Coexistence: return_value=backfilled_user, ) as mock_get_user_object, ): - result = await _user_api_key_auth_builder( + result = await user_api_key_auth_builder( request=mock_request, api_key=jwt_token, azure_api_key_header="", @@ -3262,7 +3262,7 @@ class TestJWTOAuth2Coexistence: return_value=other_owner, ) as mock_get_user_object, ): - result = await _user_api_key_auth_builder( + result = await user_api_key_auth_builder( request=mock_request, api_key=jwt_token, azure_api_key_header="", @@ -3333,7 +3333,7 @@ class TestJWTOAuth2Coexistence: side_effect=Exception("can't reach database server"), ), ): - result = await _user_api_key_auth_builder( + result = await user_api_key_auth_builder( request=mock_request, api_key=jwt_token, azure_api_key_header="", @@ -4050,7 +4050,7 @@ async def test_user_api_key_auth_builder_no_blocking_calls(): from starlette.requests import Request from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth - from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth_builder _blocking_methods = [ "set_cache", @@ -4142,7 +4142,7 @@ async def test_user_api_key_auth_builder_no_blocking_calls(): return_value=None, ) ) - await _user_api_key_auth_builder( + await user_api_key_auth_builder( request=request, api_key=f"Bearer {api_key}", azure_api_key_header="", @@ -4184,7 +4184,7 @@ async def test_team_metadata_refreshed_from_team_object_during_auth(): LitellmUserRoles, UserAPIKeyAuth, ) - from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth_builder api_key = "sk-test-team-metadata-refresh" @@ -4252,7 +4252,7 @@ async def test_team_metadata_refreshed_from_team_object_during_auth(): return_value=fresh_team_obj, ), ): - result = await _user_api_key_auth_builder( + result = await user_api_key_auth_builder( request=request, api_key=f"Bearer {api_key}", azure_api_key_header="", @@ -4297,7 +4297,7 @@ async def test_auth_flow_never_persists_fallback_team_object_lit_4391(): from fastapi import HTTPException from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth - from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth_builder api_key = "sk-test-lit-4391-no-team-writeback" valid_token = UserAPIKeyAuth( @@ -4358,7 +4358,7 @@ async def test_auth_flow_never_persists_fallback_team_object_lit_4391(): ), ), ): - result = await _user_api_key_auth_builder( + result = await user_api_key_auth_builder( request=request, api_key=f"Bearer {api_key}", azure_api_key_header="", @@ -4400,7 +4400,7 @@ async def test_auth_flow_fallback_team_resolves_object_permission_by_id(): LitellmUserRoles, UserAPIKeyAuth, ) - from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth_builder api_key = "sk-test-fallback-team-object-permission" valid_token = UserAPIKeyAuth( @@ -4471,7 +4471,7 @@ async def test_auth_flow_fallback_team_resolves_object_permission_by_id(): return_value=restricted_object_permission, ) as mock_get_object_permission, ): - result = await _user_api_key_auth_builder( + result = await user_api_key_auth_builder( request=request, api_key=f"Bearer {api_key}", azure_api_key_header="", @@ -4501,7 +4501,7 @@ async def test_auth_flow_fallback_team_object_permission_none_when_unreadable(): from fastapi import HTTPException from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth - from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth_builder api_key = "sk-test-fallback-team-object-permission-unreadable" valid_token = UserAPIKeyAuth( @@ -4566,7 +4566,7 @@ async def test_auth_flow_fallback_team_object_permission_none_when_unreadable(): return_value=None, ), ): - result = await _user_api_key_auth_builder( + result = await user_api_key_auth_builder( request=request, api_key=f"Bearer {api_key}", azure_api_key_header="", @@ -4631,7 +4631,7 @@ async def test_centralized_common_checks_runs_for_standard_auth(): "litellm.proxy.auth.user_api_key_auth.common_checks", new_callable=AsyncMock, ) as mock_checks: - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={"model": "gpt-4o"}, @@ -4690,7 +4690,7 @@ async def test_centralized_common_checks_routes_header_tags_to_litellm_metadata( new_callable=AsyncMock, ), ): - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data=request_data, @@ -4744,7 +4744,7 @@ async def test_centralized_common_checks_carries_team_and_user_budget_state_on_t "litellm.proxy.auth.user_api_key_auth.common_checks", new_callable=AsyncMock, ): - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={"model": "gpt-5.4-mini"}, @@ -4817,7 +4817,7 @@ async def _run_centralized_checks_with_key_end_user_budget( new_callable=AsyncMock, ) as mock_reserve, ): - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={"model": "gpt-5.4-mini", "user": request_user or token.end_user_id}, @@ -4998,7 +4998,7 @@ async def test_centralized_common_checks_enforces_team_model_max_budget_from_the new_callable=AsyncMock, ), ): - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={"model": "gpt-4o"}, @@ -5034,7 +5034,7 @@ async def test_centralized_common_checks_skipped_for_custom_auth_without_flag(): "litellm.proxy.auth.user_api_key_auth.common_checks", new_callable=AsyncMock, ) as mock_checks: - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={"model": "gpt-4o"}, @@ -5131,7 +5131,7 @@ async def test_centralized_checks_enforce_token_end_user_budget_against_row_spen ), pytest.raises(litellm.BudgetExceededError) as exc_info, ): - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=_chat_request(), request_data={ @@ -5167,7 +5167,7 @@ async def test_centralized_checks_skip_end_user_lookup_without_a_token_budget(): return_value=0.6, ), ): - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=_chat_request(), request_data={ @@ -5201,7 +5201,7 @@ async def test_centralized_common_checks_runs_for_custom_auth_with_flag(): "litellm.proxy.auth.user_api_key_auth.common_checks", new_callable=AsyncMock, ) as mock_checks: - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={"model": "gpt-4o"}, @@ -5242,7 +5242,7 @@ async def test_centralized_common_checks_runs_for_oauth2_fallback_token(): ), ): with pytest.raises(ProxyException) as exc: - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={"model": "gpt-4"}, @@ -5292,7 +5292,7 @@ async def test_centralized_common_checks_tolerates_db_errors_when_fetching_conte new_callable=AsyncMock, ) as mock_checks, ): - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={"model": "gpt-4o"}, @@ -5350,7 +5350,7 @@ async def test_team_key_vector_store_access_when_team_cannot_be_resolved( "litellm.proxy.auth.user_api_key_auth.get_team_object", AsyncMock(side_effect=team_lookup_error) ) - checks: Final = _run_centralized_common_checks( + checks: Final = run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={ @@ -5434,7 +5434,7 @@ async def test_keyless_team_member_vector_store_access_uses_only_the_resolved_te ), ) - checks: Final = _run_centralized_common_checks( + checks: Final = run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={ @@ -5501,7 +5501,7 @@ async def test_keyless_proxy_admin_keeps_personal_vector_store_grants_under_deny ), ) - checks: Final = _run_centralized_common_checks( + checks: Final = run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={ @@ -5561,7 +5561,7 @@ async def test_centralized_common_checks_propagates_end_user_budget_error(): ) as mock_checks, ): with pytest.raises(litellm.BudgetExceededError): - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={"user": "alice", "model": "gpt-4o"}, @@ -5624,7 +5624,7 @@ async def test_centralized_common_checks_reserves_request_end_user_budget(): ): assert token.end_user_id is None - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data=request_data, @@ -5674,7 +5674,7 @@ async def test_centralized_common_checks_short_circuits_when_master_key_unset(): "litellm.proxy.auth.user_api_key_auth.common_checks", new_callable=AsyncMock, ) as mock_checks: - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={}, @@ -5710,7 +5710,7 @@ async def test_centralized_common_checks_skips_public_routes(): "litellm.proxy.auth.user_api_key_auth.common_checks", new_callable=AsyncMock, ) as mock_checks: - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={}, @@ -5756,7 +5756,7 @@ async def test_centralized_common_checks_skips_passthrough_endpoint_with_auth_fa "litellm.proxy.auth.user_api_key_auth.common_checks", new_callable=AsyncMock, ) as mock_checks: - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={}, @@ -5800,7 +5800,7 @@ async def test_centralized_common_checks_runs_for_passthrough_endpoint_with_auth "litellm.proxy.auth.user_api_key_auth.common_checks", new_callable=AsyncMock, ) as mock_checks: - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={}, @@ -5859,7 +5859,7 @@ async def test_centralized_common_checks_master_key_admin_overrides_db_user_role new_callable=AsyncMock, ) as mock_checks, ): - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={"team_id": "t1", "max_budget": 10}, @@ -5910,7 +5910,7 @@ async def test_centralized_common_checks_http_exception_without_team_id(): ): # Should NOT raise AssertionError from _team_obj_from_token; # should proceed with team_object=None. - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={"model": "gpt-4o"}, @@ -5999,7 +5999,7 @@ async def test_centralized_common_checks_team_404_does_not_zero_other_contexts() new_callable=AsyncMock, ) as mock_checks, ): - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={"user": "alice", "model": "gpt-4o"}, @@ -6057,7 +6057,7 @@ async def test_centralized_common_checks_unresolvable_team_without_grant_is_refu side_effect=team_read_failure, ): with pytest.raises(HTTPException) as exc_info: - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={"model": "gpt-4.1"}, @@ -6111,7 +6111,7 @@ async def test_centralized_common_checks_absent_team_refused_despite_db_unavaila side_effect=team_absent, ): with pytest.raises(HTTPException) as exc_info: - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={"model": "gpt-4.1"}, @@ -6164,7 +6164,7 @@ async def test_centralized_common_checks_unreadable_team_keeps_db_unavailable_op _capturing_common_checks, ), ): - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={"model": "gpt-4.1"}, @@ -6215,7 +6215,7 @@ async def test_centralized_common_checks_unresolvable_team_with_grant_enforces_i side_effect=HTTPException(status_code=404, detail={"error": "team unreadable"}), ): if is_granted: - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={"model": requested_model}, @@ -6223,7 +6223,7 @@ async def test_centralized_common_checks_unresolvable_team_with_grant_enforces_i ) else: with pytest.raises(ProxyException) as exc_info: - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={"model": requested_model}, @@ -6283,7 +6283,7 @@ async def test_centralized_common_checks_ui_sentinel_team_vouches_despite_absent _capturing_common_checks, ), ): - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={}, @@ -6344,7 +6344,7 @@ async def test_centralized_common_checks_ui_sentinel_team_skips_db_lookup(): _capturing_common_checks, ), ): - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={}, @@ -6371,7 +6371,7 @@ async def test_builder_ui_sentinel_team_never_hits_get_team_object(): # test-qu from starlette.datastructures import URL from litellm.proxy._types import UI_TEAM_ID - from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth_builder from litellm.proxy.proxy_server import hash_token api_key = "sk-test-ui-session-key" @@ -6419,7 +6419,7 @@ async def test_builder_ui_sentinel_team_never_hits_get_team_object(): # test-qu new_callable=AsyncMock, ) as mock_get_team_object, ): - result = await _user_api_key_auth_builder( + result = await user_api_key_auth_builder( request=request, api_key=f"Bearer {api_key}", azure_api_key_header="", @@ -6503,7 +6503,7 @@ async def test_centralized_common_checks_user_http_exception_isolates_to_user_on new_callable=AsyncMock, ) as mock_checks, ): - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={"user": "alice", "model": "gpt-4o"}, @@ -6567,7 +6567,7 @@ async def test_centralized_common_checks_backfills_org_id_from_team(key_org_id, side_effect=lambda **kw: org_id_seen_by_common_checks.append(kw["valid_token"].org_id), ) as mock_checks, ): - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={"model": "gpt-4o"}, @@ -6713,14 +6713,14 @@ async def test_centralized_common_checks_inherits_org_identity( if expect_lookup_error: with pytest.raises(ConnectionRefusedError, match="db unavailable"): - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={"model": "gpt-4o"}, route="/chat/completions", ) else: - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={"model": "gpt-4o"}, @@ -6809,7 +6809,7 @@ async def test_cli_session_token_org_backfilled_from_team(monkeypatch): new_callable=AsyncMock, ), ): - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={"model": "gpt-4o"}, @@ -6851,7 +6851,7 @@ async def test_centralized_common_checks_org_backfill_survives_team_fetch_failur new_callable=AsyncMock, ) as mock_checks, ): - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={"model": "gpt-4o"}, @@ -6878,7 +6878,7 @@ async def test_master_key_auth_substitutes_alias_for_api_key(): from starlette.datastructures import URL from litellm.constants import LITELLM_PROXY_MASTER_KEY_ALIAS - from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth_builder from litellm.proxy.utils import hash_token import litellm.proxy.proxy_server as _proxy_server_mod @@ -6893,7 +6893,7 @@ async def test_master_key_auth_substitutes_alias_for_api_key(): request = Request(scope={"type": "http"}) request._url = URL(url="/chat/completions") - result = await _user_api_key_auth_builder( + result = await user_api_key_auth_builder( request=request, api_key=f"Bearer {master_key}", azure_api_key_header="", @@ -6951,12 +6951,12 @@ async def test_user_api_key_auth_sets_end_user_id_when_builder_skips_it(): # auth state machine; we only care about the wrapper's safety net. with ( patch( - "litellm.proxy.auth.user_api_key_auth._user_api_key_auth_builder", + "litellm.proxy.auth.user_api_key_auth.user_api_key_auth_builder", new_callable=AsyncMock, return_value=builder_token, ), patch( - "litellm.proxy.auth.user_api_key_auth._run_centralized_common_checks", + "litellm.proxy.auth.user_api_key_auth.run_centralized_common_checks", new_callable=AsyncMock, ), patch( @@ -7004,12 +7004,12 @@ async def test_user_api_key_auth_does_not_overwrite_end_user_id_set_by_builder() setattr(_proxy_server_mod, k, v) with ( patch( - "litellm.proxy.auth.user_api_key_auth._user_api_key_auth_builder", + "litellm.proxy.auth.user_api_key_auth.user_api_key_auth_builder", new_callable=AsyncMock, return_value=builder_token, ), patch( - "litellm.proxy.auth.user_api_key_auth._run_centralized_common_checks", + "litellm.proxy.auth.user_api_key_auth.run_centralized_common_checks", new_callable=AsyncMock, ), patch( @@ -7060,12 +7060,12 @@ async def test_user_api_key_auth_authenticates_before_raising_malformed_body_err setattr(_proxy_server_mod, k, v) with ( patch( - "litellm.proxy.auth.user_api_key_auth._user_api_key_auth_builder", + "litellm.proxy.auth.user_api_key_auth.user_api_key_auth_builder", new_callable=AsyncMock, return_value=builder_token, ) as mock_builder, patch( - "litellm.proxy.auth.user_api_key_auth._run_centralized_common_checks", + "litellm.proxy.auth.user_api_key_auth.run_centralized_common_checks", new_callable=AsyncMock, ) as mock_common_checks, patch( @@ -7121,7 +7121,7 @@ async def _run_auth_with_malformed_body(post_call_failure_hook): setattr(_proxy_server_mod, k, v) with ( patch( - "litellm.proxy.auth.user_api_key_auth._user_api_key_auth_builder", + "litellm.proxy.auth.user_api_key_auth.user_api_key_auth_builder", new_callable=AsyncMock, return_value=builder_token, ), @@ -7193,7 +7193,7 @@ async def test_user_api_key_auth_malformed_body_with_rejected_key_still_returns_ setattr(_proxy_server_mod, k, v) with ( patch( - "litellm.proxy.auth.user_api_key_auth._user_api_key_auth_builder", + "litellm.proxy.auth.user_api_key_auth.user_api_key_auth_builder", new_callable=AsyncMock, side_effect=ProxyException( message="Authentication Error, invalid key", @@ -7246,7 +7246,7 @@ async def test_user_api_key_auth_does_not_double_log_a_malformed_body_from_a_rej setattr(_proxy_server_mod, k, v) with ( patch( - "litellm.proxy.auth.user_api_key_auth._user_api_key_auth_builder", + "litellm.proxy.auth.user_api_key_auth.user_api_key_auth_builder", new_callable=AsyncMock, side_effect=ProxyException( message="Authentication Error, invalid key", @@ -7298,7 +7298,7 @@ async def _run_builder_with_key_lookup(get_key_object_mock): 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.auth.user_api_key_auth import user_api_key_auth_builder attrs = _proxy_attrs_for_db_lookup() originals = {a: getattr(_proxy_server_mod, a, None) for a in attrs} @@ -7316,7 +7316,7 @@ async def _run_builder_with_key_lookup(get_key_object_mock): "litellm.proxy.auth.auth_exception_handler.seed_request_identity", ), ): - return await _user_api_key_auth_builder( + return await user_api_key_auth_builder( request=request, api_key="Bearer sk-db-lookup-test", azure_api_key_header="", @@ -7628,7 +7628,7 @@ async def test_non_admin_cli_session_token_reaches_production_auth_path(monkeypa side_effect=__import__("fastapi").HTTPException(status_code=404), ), ): - result = await _user_api_key_auth_builder( + result = await user_api_key_auth_builder( request=request, api_key=f"Bearer {cli_token}", azure_api_key_header="", @@ -7715,7 +7715,7 @@ async def _authenticate_session_token_against_db( AsyncMock(return_value=membership_row), ), ): - await _user_api_key_auth_builder( + await user_api_key_auth_builder( request=request, api_key=f"Bearer {cli_token}", azure_api_key_header="", @@ -7925,7 +7925,7 @@ async def test_auth_does_not_rewrite_cached_key_object_back_into_cache(): 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.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 @@ -7969,10 +7969,10 @@ async def test_auth_does_not_rewrite_cached_key_object_back_into_cache(): 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", + "litellm.proxy.auth.resolvers.store.fetch_key_object_from_db_with_reconnect", fetch_from_db, ): - result = await _user_api_key_auth_builder( + result = await user_api_key_auth_builder( request=request, api_key=f"Bearer {api_key}", azure_api_key_header="", @@ -8388,7 +8388,7 @@ async def test_global_proxy_spend_reads_resettable_proxy_budget_row(): loaded from the MonthlyGlobalSpend view, whose window is hardcoded to a trailing 30 days and never resets on the configured duration.""" from litellm.proxy.auth.user_api_key_auth import ( - _fetch_global_spend_with_event_coordination, + fetch_global_spend_with_event_coordination, ) from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache @@ -8400,7 +8400,7 @@ async def test_global_proxy_spend_reads_resettable_proxy_budget_row(): side_effect=AssertionError("global spend must not be loaded from the fixed-30d MonthlyGlobalSpend view") ) - result = await _fetch_global_spend_with_event_coordination( + result = await fetch_global_spend_with_event_coordination( cache_key="default_user_id:spend", user_api_key_cache=UserApiKeyCache(), prisma_client=prisma_client, @@ -8415,14 +8415,14 @@ async def test_global_proxy_spend_none_when_proxy_budget_row_missing(): """Before the startup upsert creates the aggregate row, enforcement must see None (no cap applied) rather than raising.""" from litellm.proxy.auth.user_api_key_auth import ( - _fetch_global_spend_with_event_coordination, + fetch_global_spend_with_event_coordination, ) from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache prisma_client = MagicMock() prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=None) - result = await _fetch_global_spend_with_event_coordination( + result = await fetch_global_spend_with_event_coordination( cache_key="default_user_id:spend", user_api_key_cache=UserApiKeyCache(), prisma_client=prisma_client, @@ -8463,7 +8463,7 @@ async def test_temp_budget_increase_applied_for_cached_key(): ) user_api_key_cache = DualCache() - await _cache_key_object( + await cache_key_object( hashed_token=hashed_token, user_api_key_obj=cached_key, user_api_key_cache=user_api_key_cache, @@ -8487,13 +8487,13 @@ async def test_temp_budget_increase_applied_for_cached_key(): patch("litellm.proxy.proxy_server.user_api_key_cache", user_api_key_cache), patch("litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging_obj), patch( - "litellm.proxy.auth.user_api_key_auth._virtual_key_max_budget_alert_check", + "litellm.proxy.auth.user_api_key_auth.virtual_key_max_budget_alert_check", new_callable=AsyncMock, ), ): results = tuple( [ - await _user_api_key_auth_builder( + await user_api_key_auth_builder( request=mock_request, api_key=f"Bearer {api_key}", azure_api_key_header="", @@ -8535,7 +8535,7 @@ async def test_cached_key_team_member_budget_blocks_at_exact_cap(team_member_spe max_budget = 2.4 user_api_key_cache = DualCache() - await _cache_key_object( + await cache_key_object( hashed_token=hashed_token, user_api_key_obj=UserAPIKeyAuth( token=hashed_token, @@ -8577,7 +8577,7 @@ async def test_cached_key_team_member_budget_blocks_at_exact_cap(team_member_spe proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None) async def _auth(): - return await _user_api_key_auth_builder( + return await user_api_key_auth_builder( request=mock_request, api_key=f"Bearer {api_key}", azure_api_key_header="", @@ -8637,7 +8637,7 @@ async def test_cached_key_team_member_budget_honours_temp_increase(expiry_offset team_member_spend = 2.5 user_api_key_cache = DualCache() - await _cache_key_object( + await cache_key_object( hashed_token=hashed_token, user_api_key_obj=UserAPIKeyAuth( token=hashed_token, @@ -8683,7 +8683,7 @@ async def test_cached_key_team_member_budget_honours_temp_increase(expiry_offset proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None) async def _auth(): - return await _user_api_key_auth_builder( + return await user_api_key_auth_builder( request=mock_request, api_key=f"Bearer {api_key}", azure_api_key_header="", @@ -8726,7 +8726,7 @@ async def _authenticate_and_authorize(mock_request, api_key): from litellm.proxy.auth.user_api_key_auth import _authorize_authenticated_request request_data = {"model": "claude-sonnet-5", "messages": [{"role": "user", "content": "hi"}]} - auth_obj = await _user_api_key_auth_builder( + auth_obj = await user_api_key_auth_builder( request=mock_request, api_key=f"Bearer {api_key}", azure_api_key_header="", @@ -8773,7 +8773,7 @@ async def test_cached_key_team_member_budget_emails_configured_thresholds( alert_emails = {"50": [], "100": ["finance@example.com"]} user_api_key_cache = DualCache() - await _cache_key_object( + await cache_key_object( hashed_token=hashed_token, user_api_key_obj=UserAPIKeyAuth( token=hashed_token, @@ -8900,7 +8900,7 @@ async def _proxy_exception_for_key( patch("litellm.proxy.proxy_server.jwt_handler", jwt_handler), ): with pytest.raises(ProxyException) as exc_info: - await _user_api_key_auth_builder( + await user_api_key_auth_builder( request=mock_request, api_key=f"Bearer {api_key}", azure_api_key_header="", @@ -9208,7 +9208,7 @@ async def test_jwt_builder_returns_every_team_grant_the_key_path_gets(is_proxy_a new_callable=AsyncMock, return_value=builder_result, ): - token = await _user_api_key_auth_builder( + token = await user_api_key_auth_builder( request=request, api_key="Bearer header.payload.signature", azure_api_key_header="", @@ -9243,7 +9243,7 @@ async def test_jwt_builder_returns_every_team_grant_the_key_path_gets(is_proxy_a ) async def test_claude_view_normalizes_before_model_access(monkeypatch, route): from starlette.requests import Request - from litellm.proxy.auth.user_api_key_auth import _enforce_key_and_fallback_model_access + from litellm.proxy.auth.user_api_key_auth import enforce_key_and_fallback_model_access source = "foo[1m]" encoded = "claude-router-" + source.encode().hex() + "[1m]" @@ -9254,7 +9254,7 @@ async def test_claude_view_normalizes_before_model_access(monkeypatch, route): data = {"model": encoded, "messages": [{"role": "user", "content": "hi"}]} request = Request({"type": "http", "method": "POST", "path": route, "headers": [], "query_string": b""}) token = UserAPIKeyAuth(models=[source]) - await _enforce_key_and_fallback_model_access( + await enforce_key_and_fallback_model_access( valid_token=token, request_data=data, route=route, @@ -9267,7 +9267,7 @@ async def test_claude_view_normalizes_before_model_access(monkeypatch, route): assert json.loads(await request.body())["model"] == source assert request.scope["parsed_body"][1]["model"] == source with pytest.raises(ProxyException): - await _enforce_key_and_fallback_model_access( + await enforce_key_and_fallback_model_access( valid_token=UserAPIKeyAuth(models=["other"]), request_data=data, route=route, @@ -9668,7 +9668,7 @@ async def test_auth_flow_enters_virtual_key_mapping_when_only_an_issuer_configur side_effect=AssertionError("standard JWT auth must not run for a mapped virtual key"), ), ): - result = await _user_api_key_auth_builder( + result = await user_api_key_auth_builder( request=mock_request, api_key=jwt_token, azure_api_key_header="", @@ -9712,9 +9712,9 @@ def _alias_request(route: str, data: dict, content_type: str = "application/json async def _enforce_alias_access(token: UserAPIKeyAuth, data: dict, route: str, request, router: litellm.Router): - from litellm.proxy.auth.user_api_key_auth import _enforce_key_and_fallback_model_access + from litellm.proxy.auth.user_api_key_auth import enforce_key_and_fallback_model_access - await _enforce_key_and_fallback_model_access( + await enforce_key_and_fallback_model_access( valid_token=token, request_data=data, route=route, @@ -9783,18 +9783,18 @@ async def test_router_settings_model_group_alias_leaves_form_bodies_alone(monkey @pytest.mark.asyncio async def test_router_settings_model_group_alias_rewrite_keeps_query_params_out_of_body(monkeypatch): """LIT-3054: auth merges query params into its own copy of the body; the rewrite must not forward them.""" - from litellm.proxy.common_utils.http_parsing_utils import _read_request_body, populate_request_with_path_params + from litellm.proxy.common_utils.http_parsing_utils import read_request_body, populate_request_with_path_params router = _alias_router() monkeypatch.setattr(litellm.proxy.proxy_server, "llm_router", router) body = {"model": "AgentX-LLM", "messages": [{"role": "user", "content": "hi"}]} request = _alias_request("/v1/chat/completions", body) request.scope["query_string"] = b"api-version=2024-10-21&stream=true" - data = populate_request_with_path_params(request_data=await _read_request_body(request), request=request) + data = populate_request_with_path_params(request_data=await read_request_body(request), request=request) assert data["api-version"] == "2024-10-21" token = _alias_token(monkeypatch, "key", {"AgentX-LLM": "claude-haiku"}, ["claude-haiku"]) await _enforce_alias_access(token, data, "/v1/chat/completions", request, router) - downstream = await _read_request_body(request) + downstream = await read_request_body(request) assert downstream == {**body, "model": "claude-haiku"} assert json.loads(await request.body()) == downstream assert await request.json() == downstream @@ -10003,7 +10003,7 @@ async def test_admission_and_budget_reservation_read_the_key_spend_counter_with_ ), spend_counter_batch_scope(redis), ): - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={"model": "gpt-5.4-mini", "messages": [{"role": "user", "content": "hi"}]}, @@ -10242,7 +10242,7 @@ async def test_managed_jwt_cannot_be_downgraded_into_virtual_key_mapping(monkeyp monkeypatch.setattr(proxy_server, name, value) for _ in range(2): with pytest.raises(ProxyException) as failure: - await _user_api_key_auth_builder( + await user_api_key_auth_builder( request=_alias_request("/v1/chat/completions", {}), api_key="Bearer verified.jwt.token", azure_api_key_header="", anthropic_api_key_header=None, google_ai_studio_api_key_header=None, azure_apim_header=None, request_data={}, @@ -10271,7 +10271,7 @@ async def test_virtual_key_cannot_enter_checks_as_an_identity_managed_actor(monk client.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=target) monkeypatch.setattr(proxy_server, "prisma_client", client) checks: Final = AsyncMock() - monkeypatch.setattr(auth_module, "_run_centralized_common_checks", checks) + monkeypatch.setattr(auth_module, "run_centralized_common_checks", checks) monkeypatch.setattr(proxy_server, "proxy_logging_obj", MagicMock(post_call_failure_hook=AsyncMock(return_value=None))) data: Final = {"model": "allowed", "messages": [{"role": "user", "content": "hello"}]} request: Final = _alias_request("/v1/chat/completions", data) @@ -10324,7 +10324,7 @@ async def test_custom_auth_grants_reach_managed_targets_without_a_virtual_key_ro monkeypatch.setattr(module, "enterprise_custom_auth", custom if enterprise else None) monkeypatch.setattr(agent_registry, "global_agent_registry", registry) monkeypatch.setattr(litellm, "enable_post_custom_auth_checks", False, raising=False) - admitted: Final = await _user_api_key_auth_builder( + admitted: Final = await user_api_key_auth_builder( request=_alias_request("/a2a/target/message/send", {}), api_key=f"Bearer {credential}", azure_api_key_header="", anthropic_api_key_header=None, google_ai_studio_api_key_header=None, azure_apim_header=None, request_data={}, @@ -10346,7 +10346,7 @@ async def test_enterprise_custom_auth_key_return_stays_a_proxy_validated_key(mon monkeypatch.setattr(proxy_server, name, value) module: Final = importlib.import_module("litellm.proxy.auth.user_api_key_auth") monkeypatch.setattr(module, "enterprise_custom_auth", custom) - admitted: Final = await _user_api_key_auth_builder( + admitted: Final = await user_api_key_auth_builder( request=_alias_request("/v1/chat/completions", {}), api_key="Bearer external-credential", azure_api_key_header="", anthropic_api_key_header=None, google_ai_studio_api_key_header=None, azure_apim_header=None, request_data={}, diff --git a/tests/unit/proxy/batches_endpoints/test_endpoints.py b/tests/unit/proxy/batches_endpoints/test_endpoints.py index e2ef8e1c789..3469dbee748 100644 --- a/tests/unit/proxy/batches_endpoints/test_endpoints.py +++ b/tests/unit/proxy/batches_endpoints/test_endpoints.py @@ -206,7 +206,7 @@ def harness(monkeypatch: pytest.MonkeyPatch): monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) with ExitStack() as stack: - stack.enter_context(patch.object(endpoints, "_read_request_body", read_body)) + stack.enter_context(patch.object(endpoints, "read_request_body", read_body)) stack.enter_context( patch.object( ProxyBaseLLMRequestProcessing, @@ -593,7 +593,7 @@ async def test_create__unified_file_id_single_model_disables_cross_model_fallbac }, ) with ( - patch.object(endpoints, "_is_base64_encoded_unified_file_id", return_value="unified-xyz"), + patch.object(endpoints, "is_base64_encoded_unified_file_id", return_value="unified-xyz"), patch.object(endpoints, "get_models_from_unified_file_id", return_value=["gpt-4o-mini"]), ): resp = await call_create(harness) @@ -620,7 +620,7 @@ async def test_create__unified_file_id_not_exactly_one_model_400(harness, models }, ) with ( - patch.object(endpoints, "_is_base64_encoded_unified_file_id", return_value="unified-xyz"), + patch.object(endpoints, "is_base64_encoded_unified_file_id", return_value="unified-xyz"), patch.object(endpoints, "get_models_from_unified_file_id", return_value=models), ): with pytest.raises(ProxyException) as exc: @@ -662,7 +662,7 @@ async def test_create__unified_file_id_resolves_real_storage_url(harness): fake_repo_cls = MagicMock(return_value=fake_repo_instance) with ( - patch.object(endpoints, "_is_base64_encoded_unified_file_id", return_value="unified-xyz"), + patch.object(endpoints, "is_base64_encoded_unified_file_id", return_value="unified-xyz"), patch.object(endpoints, "get_models_from_unified_file_id", return_value=["gemini-2.0"]), patch.object(proxy_server, "prisma_client", MagicMock()), patch.object(endpoints, "ManagedFileRepository", fake_repo_cls), @@ -696,7 +696,7 @@ async def test_create__unified_file_id_db_error_falls_back_to_raw_id(harness): fake_repo_cls = MagicMock(return_value=fake_repo_instance) with ( - patch.object(endpoints, "_is_base64_encoded_unified_file_id", return_value="unified-xyz"), + patch.object(endpoints, "is_base64_encoded_unified_file_id", return_value="unified-xyz"), patch.object(endpoints, "get_models_from_unified_file_id", return_value=["gemini-2.0"]), patch.object(proxy_server, "prisma_client", MagicMock()), patch.object(endpoints, "ManagedFileRepository", fake_repo_cls), @@ -728,7 +728,7 @@ async def test_create__multi_model_unified_file_with_loadbalancing_keeps_router_ with ( patch.object(litellm, "enable_loadbalancing_on_batch_endpoints", True), - patch.object(endpoints, "_is_base64_encoded_unified_file_id", return_value="unified-xyz"), + patch.object(endpoints, "is_base64_encoded_unified_file_id", return_value="unified-xyz"), patch.object(endpoints, "get_models_from_unified_file_id", return_value=["model-a", "model-b"]), ): await call_create(harness) @@ -759,7 +759,7 @@ async def test_create__unified_file_id_missing_row_falls_back_to_raw_id(harness) fake_repo_cls = MagicMock(return_value=fake_repo_instance) with ( - patch.object(endpoints, "_is_base64_encoded_unified_file_id", return_value="unified-xyz"), + patch.object(endpoints, "is_base64_encoded_unified_file_id", return_value="unified-xyz"), patch.object(endpoints, "get_models_from_unified_file_id", return_value=["gemini-2.0"]), patch.object(proxy_server, "prisma_client", MagicMock()), patch.object(endpoints, "ManagedFileRepository", fake_repo_cls), @@ -792,7 +792,7 @@ async def test_create__unified_file_id_legacy_row_without_storage_url_dispatches fake_repo_cls = MagicMock(return_value=fake_repo_instance) with ( - patch.object(endpoints, "_is_base64_encoded_unified_file_id", return_value="unified-xyz"), + patch.object(endpoints, "is_base64_encoded_unified_file_id", return_value="unified-xyz"), patch.object(endpoints, "get_models_from_unified_file_id", return_value=["gemini-2.0"]), patch.object(proxy_server, "prisma_client", MagicMock()), patch.object(endpoints, "ManagedFileRepository", fake_repo_cls), @@ -946,7 +946,7 @@ async def test_create__model_encoded_beats_unified(harness): }, ) with ( - patch.object(endpoints, "_is_base64_encoded_unified_file_id", return_value="unified-xyz"), + patch.object(endpoints, "is_base64_encoded_unified_file_id", return_value="unified-xyz"), patch.object(endpoints, "get_models_from_unified_file_id", return_value=["something-else"]), ): await call_create(harness) @@ -1596,7 +1596,7 @@ async def test_retrieve__model_encoded_beats_loadbalancing(retrieve_harness): @pytest.mark.asyncio async def test_retrieve__unified_batch_id_routes_to_router(retrieve_harness): - with patch.object(endpoints, "_is_base64_encoded_unified_file_id", return_value=UNIFIED_BATCH_ID): + with patch.object(endpoints, "is_base64_encoded_unified_file_id", return_value=UNIFIED_BATCH_ID): resp = await call_retrieve(retrieve_harness, "batch-unified-blob") # DISPATCH - router fired, direct litellm did not. @@ -1762,7 +1762,7 @@ async def test_retrieve__db_terminal_unified_resolves_file_ids(retrieve_harness) db_batch_object = MagicMock() retrieve_harness.get_batch_from_db.return_value = (db_batch_object, db_response) - with patch.object(endpoints, "_is_base64_encoded_unified_file_id", return_value=UNIFIED_BATCH_ID): + with patch.object(endpoints, "is_base64_encoded_unified_file_id", return_value=UNIFIED_BATCH_ID): await call_retrieve(retrieve_harness, "batch-unified-blob") # Terminal short-circuit still registers/normalizes raw provider file ids. @@ -1960,7 +1960,7 @@ def list_harness(): litellm_alist = AsyncMock(return_value=FakeListPage([])) with ExitStack() as stack: - stack.enter_context(patch.object(endpoints, "_read_request_body", read_body)) + stack.enter_context(patch.object(endpoints, "read_request_body", read_body)) stack.enter_context( patch.object( ProxyBaseLLMRequestProcessing, @@ -2495,7 +2495,7 @@ async def test_cancel__model_encoded_id_forwards_deployment_model(cancel_harness @pytest.mark.asyncio async def test_cancel__model_encoded_beats_unified(cancel_harness): - with patch.object(endpoints, "_is_base64_encoded_unified_file_id", return_value=UNIFIED_BATCH_ID): + with patch.object(endpoints, "is_base64_encoded_unified_file_id", return_value=UNIFIED_BATCH_ID): await call_cancel(cancel_harness, AZURE_BATCH_ID) assert cancel_harness.litellm_acancel.call_count == 1 @@ -2511,7 +2511,7 @@ async def test_cancel__model_encoded_beats_unified(cancel_harness): @pytest.mark.asyncio async def test_cancel__unified_batch_id_routes_to_router(cancel_harness): - with patch.object(endpoints, "_is_base64_encoded_unified_file_id", return_value=UNIFIED_BATCH_ID): + with patch.object(endpoints, "is_base64_encoded_unified_file_id", return_value=UNIFIED_BATCH_ID): resp = await call_cancel(cancel_harness, "batch-unified-blob") # DISPATCH - router fired, litellm did not, no creds lookup. @@ -2537,7 +2537,7 @@ async def test_cancel__db_write_receives_caller_auth(cancel_harness): """update_batch_in_database can only mint managed IDs for a cancelled batch's output files when it has an auth context, so cancel must forward the caller's.""" caller = UserAPIKeyAuth(api_key="sk-test", user_id="user-cancel-1") - with patch.object(endpoints, "_is_base64_encoded_unified_file_id", return_value=UNIFIED_BATCH_ID): + with patch.object(endpoints, "is_base64_encoded_unified_file_id", return_value=UNIFIED_BATCH_ID): await call_cancel(cancel_harness, "batch-unified-blob", user=caller) assert cancel_harness.update_batch_in_db.call_args.kwargs["user_api_key_dict"] is caller @@ -2548,7 +2548,7 @@ async def test_cancel__unified_missing_model_id_400(cancel_harness): # unified id with no model_id segment -> get_model_id returns None -> 400. with patch.object( endpoints, - "_is_base64_encoded_unified_file_id", + "is_base64_encoded_unified_file_id", return_value="litellm_proxy;llm_batch_id:batch-xyz", ): with pytest.raises(ProxyException) as exc: @@ -2563,7 +2563,7 @@ async def test_cancel__unified_missing_model_id_400(cancel_harness): async def test_cancel__unified_no_router_500(cancel_harness): with ( patch.object(proxy_server, "llm_router", None), - patch.object(endpoints, "_is_base64_encoded_unified_file_id", return_value=UNIFIED_BATCH_ID), + patch.object(endpoints, "is_base64_encoded_unified_file_id", return_value=UNIFIED_BATCH_ID), ): with pytest.raises(ProxyException) as exc: await call_cancel(cancel_harness, "batch-unified-blob") @@ -2769,7 +2769,7 @@ async def test_create__unified_no_router_500(harness): }, ) with ( - patch.object(endpoints, "_is_base64_encoded_unified_file_id", return_value="unified-xyz"), + patch.object(endpoints, "is_base64_encoded_unified_file_id", return_value="unified-xyz"), patch.object(endpoints, "get_models_from_unified_file_id", return_value=["gpt-4o-mini"]), patch.object(proxy_server, "llm_router", None), ): @@ -2782,7 +2782,7 @@ async def test_create__unified_no_router_500(harness): @pytest.mark.asyncio async def test_retrieve__unified_no_router_500(retrieve_harness): with ( - patch.object(endpoints, "_is_base64_encoded_unified_file_id", return_value=UNIFIED_BATCH_ID), + patch.object(endpoints, "is_base64_encoded_unified_file_id", return_value=UNIFIED_BATCH_ID), patch.object(proxy_server, "llm_router", None), ): with pytest.raises(ProxyException) as exc: diff --git a/tests/unit/proxy/batches_endpoints/test_litellm_executed_batches.py b/tests/unit/proxy/batches_endpoints/test_litellm_executed_batches.py index ad99cebc342..00f84e75664 100644 --- a/tests/unit/proxy/batches_endpoints/test_litellm_executed_batches.py +++ b/tests/unit/proxy/batches_endpoints/test_litellm_executed_batches.py @@ -30,7 +30,7 @@ from litellm.proxy.batches_endpoints.litellm_executed_batches import ( upstream_lacks_files_api, ) from litellm.proxy.openai_files_endpoints.common_utils import ( - _is_base64_encoded_unified_file_id, + is_base64_encoded_unified_file_id, get_batch_id_from_unified_batch_id, is_litellm_executed_batch, ) @@ -1129,7 +1129,7 @@ async def test_run_completes_under_the_real_hooks_base64_batch_id() -> None: harness = make_runner(store_factory=RealIdManagedBatchStore) created, finished = await harness.create_and_finish() - assert _is_base64_encoded_unified_file_id(created.id) + assert is_base64_encoded_unified_file_id(created.id) assert finished.status == "completed" assert [call.model_object_id.startswith("litellm_batch_") for call in harness.store.calls] == [True] assert [write.unified_object_id for write in harness.table.writes] == [created.id] * 3 diff --git a/tests/unit/proxy/common_utils/test_cache_aware_routing.py b/tests/unit/proxy/common_utils/test_cache_aware_routing.py index 000d6773d04..24fc1f22a6a 100644 --- a/tests/unit/proxy/common_utils/test_cache_aware_routing.py +++ b/tests/unit/proxy/common_utils/test_cache_aware_routing.py @@ -363,7 +363,7 @@ async def test_router_applies_opt_in_and_preserves_failure_semantics( from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler from litellm.proxy import proxy_server from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache - from litellm.proxy.hooks.parallel_request_limiter_v3 import _PROXY_MaxParallelRequestsHandler_v3 + from litellm.proxy.hooks.parallel_request_limiter_v3 import PROXY_MaxParallelRequestsHandler_v3 from litellm.proxy.utils import ProxyLogging from litellm.router_strategy.complexity_router.context_compaction import initialize_compaction_state @@ -390,7 +390,7 @@ async def test_router_applies_opt_in_and_preserves_failure_semantics( ] ) logging: Final = ProxyLogging(UserApiKeyCache()) - logging.proxy_hook_mapping["parallel_request_limiter"] = _PROXY_MaxParallelRequestsHandler_v3( + logging.proxy_hook_mapping["parallel_request_limiter"] = PROXY_MaxParallelRequestsHandler_v3( logging.internal_usage_cache ) await _observed(logging.internal_usage_cache.dual_cache, expires_at=1e100) diff --git a/tests/unit/proxy/common_utils/test_check_batch_cost.py b/tests/unit/proxy/common_utils/test_check_batch_cost.py index d970978704e..87df82c01d0 100644 --- a/tests/unit/proxy/common_utils/test_check_batch_cost.py +++ b/tests/unit/proxy/common_utils/test_check_batch_cost.py @@ -729,7 +729,7 @@ class TestCheckBatchCost: let the original bug ship undetected. """ import litellm - from litellm.proxy.hooks.proxy_track_cost_callback import _ProxyDBLogger + from litellm.proxy.hooks.proxy_track_cost_callback import ProxyDBLogger mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock(return_value=1) mock_prisma_client.db.litellm_managedobjecttable.update = AsyncMock() @@ -773,7 +773,7 @@ class TestCheckBatchCost: decoded_id = "llm_model_id,model-123;llm_batch_id,batch-456;" - db_logger = _ProxyDBLogger() + db_logger = ProxyDBLogger() mock_update_database = AsyncMock() # Unlike the other tests in this file, this one runs the real @@ -1253,6 +1253,7 @@ class TestCheckBatchCost: mock_response = MagicMock() mock_response.status = terminal_status mock_response.output_file_id = "file-output-123" + mock_response.error_file_id = None mock_response.model_dump_json.return_value = f'{{"id":"batch-1","status":"{terminal_status}"}}' mock_llm_router.aretrieve_batch = AsyncMock(return_value=mock_response) @@ -2272,13 +2273,13 @@ class TestManagedOutputFileIdEncodesPublicModelGroup: @pytest.mark.asyncio async def test_target_model_names_comes_from_input_file_not_provider_model(self): from litellm.proxy.openai_files_endpoints.common_utils import ( - _is_base64_encoded_unified_file_id, + is_base64_encoded_unified_file_id, get_models_from_unified_file_id, ) output_file_id = await self._run(self._job(self._managed_input_file_id(self._PUBLIC_MODEL_GROUP))) - decoded = _is_base64_encoded_unified_file_id(output_file_id) + decoded = is_base64_encoded_unified_file_id(output_file_id) assert get_models_from_unified_file_id(decoded) == [self._PUBLIC_MODEL_GROUP] assert "gpt-5.5" not in decoded @@ -2307,13 +2308,13 @@ class TestManagedOutputFileIdEncodesPublicModelGroup: @pytest.mark.asyncio async def test_falls_back_to_deployment_model_group_without_managed_input_file(self): from litellm.proxy.openai_files_endpoints.common_utils import ( - _is_base64_encoded_unified_file_id, + is_base64_encoded_unified_file_id, get_models_from_unified_file_id, ) output_file_id = await self._run(self._job("file-raw-provider-input")) - decoded = _is_base64_encoded_unified_file_id(output_file_id) + decoded = is_base64_encoded_unified_file_id(output_file_id) assert get_models_from_unified_file_id(decoded) == [self._PUBLIC_MODEL_GROUP] diff --git a/tests/unit/proxy/common_utils/test_encrypt_decrypt_utils.py b/tests/unit/proxy/common_utils/test_encrypt_decrypt_utils.py index 7c07f19d72f..1163090c510 100644 --- a/tests/unit/proxy/common_utils/test_encrypt_decrypt_utils.py +++ b/tests/unit/proxy/common_utils/test_encrypt_decrypt_utils.py @@ -13,7 +13,7 @@ import pytest from litellm.proxy import proxy_server from litellm.proxy.common_utils.encrypt_decrypt_utils import ( - _V2_GCM_PREFIX, + V2_GCM_PREFIX, decrypt_bearer_token, decrypt_if_encrypted_with, decrypt_value_helper, @@ -43,7 +43,7 @@ def test_aes_gcm_round_trip(monkeypatch): ct = encrypt_value_helper("super-secret") - assert ct.startswith(_V2_GCM_PREFIX) + assert ct.startswith(V2_GCM_PREFIX) assert decrypt_value_helper(ct, key="t") == "super-secret" @@ -51,7 +51,7 @@ def test_default_is_legacy_algorithm(monkeypatch): """With no config, writes stay on the legacy algorithm (no v2: marker).""" ct = encrypt_value_helper("legacy-secret") - assert not ct.startswith(_V2_GCM_PREFIX) + assert not ct.startswith(V2_GCM_PREFIX) assert decrypt_value_helper(ct, key="t") == "legacy-secret" @@ -62,12 +62,12 @@ def test_legacy_nacl_value_still_decrypts_after_flag_flip(monkeypatch): flipping the flag forward never strands previously-written data. """ legacy = encrypt_value_helper("legacy-secret") # default = xsalsa20 - assert not legacy.startswith(_V2_GCM_PREFIX) + assert not legacy.startswith(V2_GCM_PREFIX) _use_aes(monkeypatch) # New writes are now AES, but the old value must still come back. assert decrypt_value_helper(legacy, key="t") == "legacy-secret" - assert encrypt_value_helper("fresh").startswith(_V2_GCM_PREFIX) + assert encrypt_value_helper("fresh").startswith(V2_GCM_PREFIX) def test_v2_prefix_is_idempotent_marker(monkeypatch): @@ -79,11 +79,11 @@ def test_v2_prefix_is_idempotent_marker(monkeypatch): _use_aes(monkeypatch) ct = encrypt_value_helper("secret") - assert ct.startswith(_V2_GCM_PREFIX) + assert ct.startswith(V2_GCM_PREFIX) # Round-tripping does not change the plaintext, and the marker is stable. again = encrypt_value_helper(decrypt_value_helper(ct, key="t")) - assert again.startswith(_V2_GCM_PREFIX) + assert again.startswith(V2_GCM_PREFIX) assert decrypt_value_helper(again, key="t") == "secret" @@ -91,7 +91,7 @@ def test_aes_decrypt_failure_returns_none_not_raise(monkeypatch): """Decrypt contract preserved: a garbled v2 value returns None, never raises.""" _use_aes(monkeypatch) - garbled = _V2_GCM_PREFIX + "not-valid-base64-or-ciphertext!!!" + garbled = V2_GCM_PREFIX + "not-valid-base64-or-ciphertext!!!" # exception_type="debug" exercises the swallow path; must not raise. assert decrypt_value_helper(garbled, key="t", exception_type="debug") is None @@ -100,7 +100,7 @@ def test_aes_decrypt_failure_returns_original_when_requested(monkeypatch): """With return_original_value=True a bad v2 value comes back as-is, not None.""" _use_aes(monkeypatch) - garbled = _V2_GCM_PREFIX + "###" + garbled = V2_GCM_PREFIX + "###" assert decrypt_value_helper(garbled, key="t", exception_type="debug", return_original_value=True) == garbled @@ -109,7 +109,7 @@ def test_empty_string_round_trips_under_aes(monkeypatch): _use_aes(monkeypatch) ct = encrypt_value_helper("") - assert ct.startswith(_V2_GCM_PREFIX) + assert ct.startswith(V2_GCM_PREFIX) assert decrypt_value_helper(ct, key="t") == "" @@ -121,7 +121,7 @@ def test_callback_prefix_composes_with_v2(monkeypatch): helper is ``v2:gcm:...``. Ordering must work end to end. """ from litellm.proxy.common_utils.callback_utils import ( - _CALLBACK_VAR_ENCRYPTED_PREFIX, + CALLBACK_VAR_ENCRYPTED_PREFIX, _decrypt_or_passthrough, _encrypt_if_plaintext, ) @@ -131,9 +131,9 @@ def test_callback_prefix_composes_with_v2(monkeypatch): # "gcs_path_service_account" is a known-sensitive callback key. stored = _encrypt_if_plaintext("gcs_path_service_account", "my-sa-secret") - assert stored.startswith(_CALLBACK_VAR_ENCRYPTED_PREFIX) - inner = stored[len(_CALLBACK_VAR_ENCRYPTED_PREFIX) :] - assert inner.startswith(_V2_GCM_PREFIX) + assert stored.startswith(CALLBACK_VAR_ENCRYPTED_PREFIX) + inner = stored[len(CALLBACK_VAR_ENCRYPTED_PREFIX) :] + assert inner.startswith(V2_GCM_PREFIX) assert _decrypt_or_passthrough("gcs_path_service_account", stored) == "my-sa-secret" @@ -142,7 +142,7 @@ def test_unknown_algorithm_falls_back_to_legacy(monkeypatch): monkeypatch.setattr(proxy_server, "general_settings", {"encryption_algorithm": "rot13"}) ct = encrypt_value_helper("secret") - assert not ct.startswith(_V2_GCM_PREFIX) + assert not ct.startswith(V2_GCM_PREFIX) assert decrypt_value_helper(ct, key="t") == "secret" @@ -256,7 +256,7 @@ def test_stored_value_is_not_a_bearer_token_even_when_reshaped(monkeypatch, use_ _use_aes(monkeypatch) stored = encrypt_value_helper("stored-secret") - for candidate in (stored, "kind_a_" + stored.removeprefix(_V2_GCM_PREFIX).rstrip("=")): + for candidate in (stored, "kind_a_" + stored.removeprefix(V2_GCM_PREFIX).rstrip("=")): assert decrypt_bearer_token(candidate, prefix="kind_a_") is None diff --git a/tests/unit/proxy/common_utils/test_http_parsing_utils.py b/tests/unit/proxy/common_utils/test_http_parsing_utils.py index 222a2cb329b..7677837aa0f 100644 --- a/tests/unit/proxy/common_utils/test_http_parsing_utils.py +++ b/tests/unit/proxy/common_utils/test_http_parsing_utils.py @@ -17,11 +17,11 @@ import litellm.proxy.common_utils.http_parsing_utils as http_parsing_utils from litellm.proxy._types import ProxyException from litellm.proxy.common_utils.http_parsing_utils import ( _is_form_content_type, - _read_request_body, - _safe_get_request_headers, + read_request_body, + safe_get_request_headers, _safe_get_request_parsed_body, - _safe_get_request_query_params, - _safe_set_request_parsed_body, + safe_get_request_query_params, + safe_set_request_parsed_body, coerce_numeric_form_fields, get_form_data, get_request_body, @@ -74,8 +74,8 @@ async def test_read_request_body_marks_body_received_once_with_its_size(monkeypa body: Final = orjson.dumps({"model": "claude-sonnet-4-5", "messages": [{"role": "user", "content": "x" * 4096}]}) request: Final = _starlette_request(body, "application/json") - assert await _read_request_body(request) == orjson.loads(body) - assert await _read_request_body(request) == orjson.loads(body) + assert await read_request_body(request) == orjson.loads(body) + assert await read_request_body(request) == orjson.loads(body) assert events == [("litellm.request.body_received", {"litellm.request.body_bytes": len(body)})] @@ -92,9 +92,9 @@ async def test_read_request_body_marks_body_received_for_binary_and_form_bodies( form: Final = b"model=whisper-1&language=en" form_type: Final = "application/x-www-form-urlencoded" - await _read_request_body(_starlette_request(protobuf, "application/x-protobuf")) - await _read_request_body(_starlette_request(form, form_type, content_length=str(len(form)))) - await _read_request_body(_starlette_request(form, form_type)) + await read_request_body(_starlette_request(protobuf, "application/x-protobuf")) + await read_request_body(_starlette_request(form, form_type, content_length=str(len(form)))) + await read_request_body(_starlette_request(form, form_type)) assert events == [ ("litellm.request.body_received", {"litellm.request.body_bytes": len(protobuf)}), @@ -108,7 +108,7 @@ async def test_read_raw_json_body_returns_the_bytes_the_parsed_body_came_from(): body = b'{"model": "claude-sonnet-4-5", "messages": [{"role": "user", "content": "hi"}]}' request = _starlette_request(body, "application/json") - assert await _read_request_body(request) == orjson.loads(body) + assert await read_request_body(request) == orjson.loads(body) assert await read_raw_json_body(request) == body @@ -124,7 +124,7 @@ async def test_read_raw_json_body_is_none_until_the_body_has_been_parsed(): async def test_read_raw_json_body_is_none_for_form_bodies(): request = _starlette_request(b"model=claude-sonnet-4-5", "application/x-www-form-urlencoded") - assert await _read_request_body(request) == {"model": "claude-sonnet-4-5"} + assert await read_request_body(request) == {"model": "claude-sonnet-4-5"} assert await read_raw_json_body(request) is None @@ -136,7 +136,7 @@ async def test_protobuf_body_is_not_parsed_as_json(content_type): body = b"\n\xa2\x01\n\x1c\n\x0cservice.name\x12\x0c\n\nswarm\xed\xa0\x80\xff" request = _starlette_request(body, content_type) - assert await _read_request_body(request) == {} + assert await read_request_body(request) == {} assert await request.body() == body # body is still readable by the endpoint @@ -144,7 +144,7 @@ async def test_protobuf_body_is_not_parsed_as_json(content_type): async def test_gzipped_json_trace_body_survives_auth_pre_read(): body = gzip.compress(b'{"resourceSpans": []}') request = _starlette_request(body, "application/json", "/v1/traces", "gzip") - assert await _read_request_body(request) == {} + assert await read_request_body(request) == {} assert await request.body() == body @@ -170,7 +170,7 @@ async def test_request_body_caching(): mock_request.scope = {} # First call should parse the body - result1 = await _read_request_body(mock_request) + result1 = await read_request_body(mock_request) assert result1 == test_data assert "parsed_body" in mock_request.scope assert mock_request.scope["parsed_body"] == (("key",), {"key": "value"}) @@ -182,7 +182,7 @@ async def test_request_body_caching(): mock_request.body.reset_mock() # Second call should use the cached body - result2 = await _read_request_body(mock_request) + result2 = await read_request_body(mock_request) assert result2 == {"key": "value"} # Verify the body was not read again @@ -205,7 +205,7 @@ async def test_form_data_parsing(): mock_request.state._cached_headers = None # Parse the form data - result = await _read_request_body(mock_request) + result = await read_request_body(mock_request) # Verify the form data was correctly parsed assert result == test_data @@ -256,7 +256,7 @@ async def test_form_data_with_json_metadata(): mock_request.state._cached_headers = None # Parse the form data - result = await _read_request_body(mock_request) + result = await read_request_body(mock_request) # Verify the metadata was parsed from JSON string to dict assert "metadata" in result @@ -298,7 +298,7 @@ async def test_form_data_with_invalid_json_metadata(): # Should raise JSONDecodeError when trying to parse invalid JSON metadata with pytest.raises(json.JSONDecodeError): - await _read_request_body(mock_request) + await read_request_body(mock_request) @pytest.mark.asyncio @@ -320,7 +320,7 @@ async def test_form_data_without_metadata(): mock_request.state._cached_headers = None # Parse the form data - result = await _read_request_body(mock_request) + result = await read_request_body(mock_request) # Verify all fields are preserved as-is assert result == test_data @@ -351,7 +351,7 @@ async def test_form_data_with_empty_metadata(): mock_request.state._cached_headers = None # Parse the form data - result = await _read_request_body(mock_request) + result = await read_request_body(mock_request) # Verify the metadata was parsed to an empty dict assert "metadata" in result @@ -386,7 +386,7 @@ async def test_form_data_with_dict_metadata(): mock_request.state._cached_headers = None # Parse the form data - result = await _read_request_body(mock_request) + result = await read_request_body(mock_request) # Verify the metadata remains as a dict and is not parsed assert "metadata" in result @@ -417,7 +417,7 @@ async def test_form_data_with_none_metadata(): mock_request.state._cached_headers = None # Parse the form data - result = await _read_request_body(mock_request) + result = await read_request_body(mock_request) # Verify the metadata remains None (not parsed) assert "metadata" in result @@ -437,7 +437,7 @@ async def test_empty_request_body(): mock_request.scope = {} # Parse the empty body - result = await _read_request_body(mock_request) + result = await read_request_body(mock_request) # Verify an empty dict is returned assert result == {} @@ -466,7 +466,7 @@ async def test_circular_reference_handling(): mock_request.scope = {} # First parse - result = await _read_request_body(mock_request) + result = await read_request_body(mock_request) # Verify initial parse assert result["model"] == "gpt-4" @@ -481,7 +481,7 @@ async def test_circular_reference_handling(): } # Second parse using the same request - will use the modified cached value - result2 = await _read_request_body(mock_request) + result2 = await read_request_body(mock_request) assert "proxy_server_request" not in result2 # This will pass, showing the cache pollution @@ -513,7 +513,7 @@ async def test_json_parsing_error_handling(): # Should raise ProxyException for trailing comma with pytest.raises(ProxyException) as exc_info: - await _read_request_body(mock_request) + await read_request_body(mock_request) assert exc_info.value.code == "400" assert "Invalid JSON payload" in exc_info.value.message @@ -538,7 +538,7 @@ async def test_json_parsing_error_handling(): # Should raise ProxyException for unquoted property with pytest.raises(ProxyException) as exc_info2: - await _read_request_body(mock_request2) + await read_request_body(mock_request2) assert exc_info2.value.code == "400" assert "Invalid JSON payload" in exc_info2.value.message @@ -564,7 +564,7 @@ async def test_json_parsing_error_handling(): mock_request3.scope = {} # Should parse successfully - result = await _read_request_body(mock_request3) + result = await read_request_body(mock_request3) assert result["model"] == "gpt-4o" assert result["input"] == "Run available tools" assert len(result["tools"]) == 1 @@ -597,21 +597,21 @@ async def test_surrogate_repair_skipped_above_size_limit(monkeypatch): small_body = b'{"model":"gpt-4o","x":NaN}' assert len(small_body) <= 100 - repaired = await _read_request_body(_make_json_request(small_body)) + repaired = await read_request_body(_make_json_request(small_body)) assert repaired["model"] == "gpt-4o" padding = "a" * 200 large_body = b'{"model":"gpt-4o","pad":"' + padding.encode() + b'","x":NaN}' assert len(large_body) > 100 with pytest.raises(ProxyException) as exc_info: - await _read_request_body(_make_json_request(large_body)) + await read_request_body(_make_json_request(large_body)) assert exc_info.value.code == "400" assert "Invalid JSON payload" in exc_info.value.message # Disabling the cap (0) restores repair for the same large body, proving the cap # — not the malformed content — is what short-circuits the repair. monkeypatch.setattr(http_parsing_utils, "MAX_REQUEST_BODY_SIZE_TO_REPAIR_MB", 0) - repaired_large = await _read_request_body(_make_json_request(large_body)) + repaired_large = await read_request_body(_make_json_request(large_body)) assert repaired_large["model"] == "gpt-4o" @@ -632,13 +632,13 @@ async def test_lone_surrogate_escape_is_rejected_with_400(content: bytes): """ body = b'{"model":"gpt-4o","messages":[{"role":"user","content":"' + content + b'"}]}' with pytest.raises(ProxyException) as exc_info: - await _read_request_body(_make_json_request(body)) + await read_request_body(_make_json_request(body)) assert exc_info.value.code == "400" assert exc_info.value.type == "invalid_request_error" assert "Invalid JSON payload" in exc_info.value.message paired = body.replace(content, b"say ok \\ud83d\\ude00") - parsed = await _read_request_body(_make_json_request(paired)) + parsed = await read_request_body(_make_json_request(paired)) assert parsed["messages"][0]["content"] == "say ok \U0001f600" @@ -646,7 +646,7 @@ async def test_lone_surrogate_escape_is_rejected_with_400(content: bytes): @pytest.mark.parametrize("media_type", ["application/x-protobuf", "application/protobuf", "application/octet-stream"]) async def test_json_body_under_a_binary_content_type_is_still_parsed(media_type: str): request = _starlette_request(b'{"model": "claude-sonnet-5"}', media_type) - assert await _read_request_body(request) == {"model": "claude-sonnet-5"} + assert await read_request_body(request) == {"model": "claude-sonnet-5"} @pytest.mark.asyncio @@ -920,7 +920,7 @@ async def test_request_body_with_html_script_tags(): mock_request.headers = {"content-type": "application/json"} mock_request.scope = {} - result = await _read_request_body(mock_request) + result = await read_request_body(mock_request) assert result["model"] == "gpt-4o" assert len(result["messages"]) == 3 @@ -943,7 +943,7 @@ def test_safe_get_request_headers_caches_on_request_state(): mock_request.state = MagicMock(spec=[]) # empty spec so getattr returns default # First call — should create and cache - result1 = _safe_get_request_headers(mock_request) + result1 = safe_get_request_headers(mock_request) assert result1 == { "content-type": "application/json", "authorization": "Bearer sk-123", @@ -951,7 +951,7 @@ def test_safe_get_request_headers_caches_on_request_state(): assert mock_request.state._cached_headers is result1 # Second call — should return the cached object (same identity) - result2 = _safe_get_request_headers(mock_request) + result2 = safe_get_request_headers(mock_request) assert result2 is result1 @@ -959,7 +959,7 @@ def test_safe_get_request_headers_none_request(): """ Test that _safe_get_request_headers returns empty dict for None request. """ - result = _safe_get_request_headers(None) + result = safe_get_request_headers(None) assert result == {} @@ -971,15 +971,15 @@ def test_safe_get_request_headers_copy_protects_cache(): mock_request.headers = {"authorization": "Bearer sk-123", "host": "localhost"} mock_request.state = MagicMock(spec=[]) - original = _safe_get_request_headers(mock_request) + original = safe_get_request_headers(mock_request) # Simulate what mutation call sites do: copy then pop - mutable = _safe_get_request_headers(mock_request).copy() + mutable = safe_get_request_headers(mock_request).copy() mutable.pop("authorization", None) # Cache must be unaffected - assert "authorization" in _safe_get_request_headers(mock_request) - assert _safe_get_request_headers(mock_request) is original + assert "authorization" in safe_get_request_headers(mock_request) + assert safe_get_request_headers(mock_request) is original def test_safe_get_request_headers_state_unavailable(): @@ -1001,7 +1001,7 @@ def test_safe_get_request_headers_state_unavailable(): mock_request.headers = {"content-type": "application/json"} mock_request.state = ReadOnlyState() - result = _safe_get_request_headers(mock_request) + result = safe_get_request_headers(mock_request) assert result == {"content-type": "application/json"} @@ -1095,7 +1095,7 @@ class TestReadRequestBodyNonCanonicalContentType: mock_request.headers = {"content-type": content_type} mock_request.scope = {} - result = await _read_request_body(mock_request) + result = await read_request_body(mock_request) assert result == payload mock_request.form.assert_not_called() @@ -1107,7 +1107,7 @@ class TestReadRequestBodyNonCanonicalContentType: mock_request.headers = {"content-type": "application/x-www-form-urlencoded"} mock_request.scope = {} - result = await _read_request_body(mock_request) + result = await read_request_body(mock_request) assert result == {"k": "v"} mock_request.form.assert_awaited_once() @@ -1136,7 +1136,7 @@ class TestReadRequestBodyFormParseFailure: mock_request.scope = {} with pytest.raises(ProxyException) as exc_info: - await _read_request_body(mock_request) + await read_request_body(mock_request) assert str(exc_info.value.code) == "400" @@ -1341,7 +1341,7 @@ async def test_only_trace_ingest_skips_json_body(method: str, path: str, skip_pa receive, ) - parsed: Final = await _read_request_body(request) + parsed: Final = await read_request_body(request) if skip_parse: assert parsed == {} receive.assert_not_awaited() @@ -1383,7 +1383,7 @@ async def test_otlp_auth_does_not_consume_chunked_bodies_before_the_receiver_lim }, receive, ) - assert await _read_request_body(request) == {} + assert await read_request_body(request) == {} assert received == [] storage = MagicMock() storage.ingest = AsyncMock() @@ -1535,27 +1535,27 @@ def _request_with_body(body: bytes) -> Request_http_parsing: @pytest.mark.usefixtures("_vcr_outcome_gate", "isolate_litellm_state", "setup_and_teardown") @pytest.mark.asyncio async def test_read_request_body_valid_json(): - result = await _read_request_body(_request_with_body(b'{"key": "value"}')) + result = await read_request_body(_request_with_body(b'{"key": "value"}')) assert result == {"key": "value"} @pytest.mark.usefixtures("_vcr_outcome_gate", "isolate_litellm_state", "setup_and_teardown") @pytest.mark.asyncio async def test_read_request_body_empty_body(): - result = await _read_request_body(_request_with_body(b"")) + result = await read_request_body(_request_with_body(b"")) assert result == {} @pytest.mark.usefixtures("_vcr_outcome_gate", "isolate_litellm_state", "setup_and_teardown") @pytest.mark.asyncio async def test_read_request_body_invalid_json(): with pytest.raises(ProxyException): - await _read_request_body(_request_with_body(b'{"key": value}')) + await read_request_body(_request_with_body(b'{"key": value}')) @pytest.mark.usefixtures("_vcr_outcome_gate", "isolate_litellm_state", "setup_and_teardown") @pytest.mark.asyncio async def test_read_request_body_large_payload(): large_payload = '{"key":' + '"a"' * 10**6 + "}" with pytest.raises(ProxyException): - await _read_request_body(_request_with_body(large_payload.encode())) + await read_request_body(_request_with_body(large_payload.encode())) @pytest.mark.usefixtures("_vcr_outcome_gate", "isolate_litellm_state", "setup_and_teardown") @pytest.mark.asyncio @@ -1563,5 +1563,5 @@ async def test_read_request_body_unexpected_error(): async def receive() -> Message: raise ValueError("Unexpected error") - result = await _read_request_body(_request(receive)) + result = await read_request_body(_request(receive)) assert result == {} diff --git a/tests/unit/proxy/common_utils/test_rbac_utils.py b/tests/unit/proxy/common_utils/test_rbac_utils.py index 54b465a94ba..9c2af565efb 100644 --- a/tests/unit/proxy/common_utils/test_rbac_utils.py +++ b/tests/unit/proxy/common_utils/test_rbac_utils.py @@ -60,9 +60,7 @@ async def test_feature_not_disabled_allows_internal_user(): @pytest.mark.asyncio async def test_feature_not_disabled_allows_vector_stores(): user = _make_user(LitellmUserRoles.INTERNAL_USER.value) - with patch.dict( - _GS_PATH, {"disable_vector_stores_for_internal_users": False}, clear=True - ): + with patch.dict(_GS_PATH, {"disable_vector_stores_for_internal_users": False}, clear=True): await check_feature_access_for_user(user, "vector_stores") @@ -120,7 +118,7 @@ async def test_agents_disabled_team_admin_allowed(): clear=True, ): with patch( - "litellm.proxy.management_endpoints.common_utils._user_has_admin_privileges", + "litellm.proxy.management_endpoints.common_utils.user_has_admin_privileges", new=AsyncMock(return_value=True), ): await check_feature_access_for_user(user, "agents") @@ -138,7 +136,7 @@ async def test_agents_disabled_non_team_admin_blocked(): clear=True, ): with patch( - "litellm.proxy.management_endpoints.common_utils._user_has_admin_privileges", + "litellm.proxy.management_endpoints.common_utils.user_has_admin_privileges", new=AsyncMock(return_value=False), ): with pytest.raises(HTTPException) as exc_info: @@ -158,7 +156,7 @@ async def test_vector_stores_disabled_team_admin_allowed(): clear=True, ): with patch( - "litellm.proxy.management_endpoints.common_utils._user_has_admin_privileges", + "litellm.proxy.management_endpoints.common_utils.user_has_admin_privileges", new=AsyncMock(return_value=True), ): await check_feature_access_for_user(user, "vector_stores") @@ -176,7 +174,7 @@ async def test_vector_stores_disabled_non_team_admin_blocked(): clear=True, ): with patch( - "litellm.proxy.management_endpoints.common_utils._user_has_admin_privileges", + "litellm.proxy.management_endpoints.common_utils.user_has_admin_privileges", new=AsyncMock(return_value=False), ): with pytest.raises(HTTPException) as exc_info: @@ -238,8 +236,6 @@ async def test_org_admin_role_enum_and_string_both_blocked(): with pytest.raises(HTTPException): await check_org_admin_can_generate_keys(user_str) - user_enum = UserAPIKeyAuth( - user_role=LitellmUserRoles.ORG_ADMIN, user_id="user-1" - ) + user_enum = UserAPIKeyAuth(user_role=LitellmUserRoles.ORG_ADMIN, user_id="user-1") with pytest.raises(HTTPException): await check_org_admin_can_generate_keys(user_enum) diff --git a/tests/unit/proxy/common_utils/test_realtime_cache.py b/tests/unit/proxy/common_utils/test_realtime_cache.py index 8316ed1d29a..6f5751fa923 100644 --- a/tests/unit/proxy/common_utils/test_realtime_cache.py +++ b/tests/unit/proxy/common_utils/test_realtime_cache.py @@ -2,21 +2,21 @@ from typing import Any, cast import pytest -from litellm.proxy.common_utils.realtime_utils import _realtime_request_body +from litellm.proxy.common_utils.realtime_utils import realtime_request_body from litellm.proxy.proxy_server import _realtime_query_params_template @pytest.fixture(autouse=True) def clear_realtime_caches(): - _realtime_request_body.cache_clear() + realtime_request_body.cache_clear() _realtime_query_params_template.cache_clear() yield - _realtime_request_body.cache_clear() + realtime_request_body.cache_clear() _realtime_query_params_template.cache_clear() def test_realtime_request_body_returns_immutable_bytes(): - cached_body = _realtime_request_body("gpt-4o") + cached_body = realtime_request_body("gpt-4o") with pytest.raises(TypeError): cast(Any, cached_body)[0] = ord("x") @@ -30,9 +30,9 @@ def test_realtime_query_params_template_returns_immutable_tuples(): def test_realtime_request_body_caches_each_model_separately(): - gpt4o_body_first = _realtime_request_body("gpt-4o") - gpt4o_body_second = _realtime_request_body("gpt-4o") - gpt4o_mini_body = _realtime_request_body("gpt-4o-mini") + gpt4o_body_first = realtime_request_body("gpt-4o") + gpt4o_body_second = realtime_request_body("gpt-4o") + gpt4o_mini_body = realtime_request_body("gpt-4o-mini") assert gpt4o_body_first is gpt4o_body_second assert gpt4o_body_first == b'{"model": "gpt-4o"}' @@ -44,9 +44,7 @@ def test_realtime_query_params_template_caches_each_pair_separately(): params_with_intent_first = _realtime_query_params_template("gpt-4o", "intent-a") params_with_intent_second = _realtime_query_params_template("gpt-4o", "intent-a") params_without_intent = _realtime_query_params_template("gpt-4o", None) - params_transcription_without_model = _realtime_query_params_template( - None, "transcription" - ) + params_transcription_without_model = _realtime_query_params_template(None, "transcription") assert params_with_intent_first is params_with_intent_second assert params_with_intent_first == (("model", "gpt-4o"), ("intent", "intent-a")) diff --git a/tests/unit/proxy/common_utils/test_upsert_budget_membership.py b/tests/unit/proxy/common_utils/test_upsert_budget_membership.py index 7a75ec395f1..74d1748d4de 100644 --- a/tests/unit/proxy/common_utils/test_upsert_budget_membership.py +++ b/tests/unit/proxy/common_utils/test_upsert_budget_membership.py @@ -6,7 +6,7 @@ from unittest.mock import AsyncMock, MagicMock import pytest from litellm.proxy.management_endpoints.common_utils import ( - _upsert_budget_and_membership, + upsert_budget_and_membership, ) # --------------------------------------------------------------------------- @@ -68,7 +68,7 @@ def stored_budget_row(mock_tx): # role must not silently wipe their budget. @pytest.mark.asyncio async def test_empty_patch_is_noop(mock_tx, fake_user): - await _upsert_budget_and_membership( + await upsert_budget_and_membership( mock_tx, team_id="team-1", user_id="user-1", @@ -89,7 +89,7 @@ async def test_empty_patch_is_noop(mock_tx, fake_user): async def test_clearing_all_limits_disconnects(mock_tx, fake_user): mock_tx.litellm_budgettable.find_unique = AsyncMock(return_value=budget_row(max_budget=100.0)) - await _upsert_budget_and_membership( + await upsert_budget_and_membership( mock_tx, team_id="team-1", user_id="user-1", @@ -114,7 +114,7 @@ async def test_clear_one_field_keeps_others(mock_tx, fake_user): return_value=budget_row(max_budget=100.0, budget_duration="24h") ) - await _upsert_budget_and_membership( + await upsert_budget_and_membership( mock_tx, team_id="team-1", user_id="user-1", @@ -140,7 +140,7 @@ async def test_clear_one_field_keeps_others(mock_tx, fake_user): async def test_update_in_place_seeds_reset_at(mock_tx, fake_user): mock_tx.litellm_budgettable.find_unique = AsyncMock(return_value=budget_row(max_budget=20.0)) - await _upsert_budget_and_membership( + await upsert_budget_and_membership( mock_tx, team_id="team-dur", user_id="user-dur", @@ -165,7 +165,7 @@ async def test_update_in_place_seeds_reset_at(mock_tx, fake_user): async def test_update_in_place_single_field_leaves_reset_at_alone(mock_tx, fake_user): mock_tx.litellm_budgettable.find_unique = AsyncMock(return_value=budget_row(max_budget=50.0)) - await _upsert_budget_and_membership( + await upsert_budget_and_membership( mock_tx, team_id="team-rpm", user_id="user-rpm", @@ -185,7 +185,7 @@ async def test_update_in_place_single_field_leaves_reset_at_alone(mock_tx, fake_ # the duration and a future reset time, then links the membership. @pytest.mark.asyncio async def test_create_seeds_reset_at_and_links(mock_tx, fake_user): - await _upsert_budget_and_membership( + await upsert_budget_and_membership( mock_tx, team_id="team-new", user_id="user-new", @@ -220,7 +220,7 @@ async def test_create_seeds_reset_at_and_links(mock_tx, fake_user): @pytest.mark.asyncio async def test_create_from_temp_budget_pair_only(mock_tx, fake_user): expiry = datetime(2100, 1, 1, tzinfo=timezone.utc) - await _upsert_budget_and_membership( + await upsert_budget_and_membership( mock_tx, team_id="team-new", user_id="user-new", @@ -241,7 +241,7 @@ async def test_create_from_temp_pair_never_snapshots_team_default(mock_tx, fake_ mock_tx.litellm_budgettable.find_unique = AsyncMock( return_value=budget_row(budget_id="team-default-budget-1", max_budget=0.4, rpm_limit=10) ) - await _upsert_budget_and_membership( + await upsert_budget_and_membership( mock_tx, team_id="team-default", user_id="user-unlinked", @@ -262,7 +262,7 @@ async def test_temp_pair_on_shared_default_member_creates_bare_row(mock_tx, fake mock_tx.litellm_budgettable.find_unique = AsyncMock( return_value=budget_row(budget_id="team-default-budget-1", max_budget=0.4, rpm_limit=10) ) - await _upsert_budget_and_membership( + await upsert_budget_and_membership( mock_tx, team_id="team-default", user_id="user-on-default", @@ -280,7 +280,7 @@ async def test_temp_pair_on_shared_default_member_creates_bare_row(mock_tx, fake @pytest.mark.asyncio async def test_clearing_temp_pair_on_shared_default_member_is_noop(mock_tx, fake_user): - await _upsert_budget_and_membership( + await upsert_budget_and_membership( mock_tx, team_id="team-default", user_id="user-on-default", @@ -302,7 +302,7 @@ async def test_temp_pair_with_permanent_field_still_clones_shared_default(mock_t mock_tx.litellm_budgettable.find_unique = AsyncMock( return_value=budget_row(budget_id="team-default-budget-1", max_budget=0.4, rpm_limit=10) ) - await _upsert_budget_and_membership( + await upsert_budget_and_membership( mock_tx, team_id="team-default", user_id="user-on-default", @@ -326,7 +326,7 @@ async def test_create_from_plain_patch_does_not_snapshot_team_default(mock_tx, f mock_tx.litellm_budgettable.find_unique = AsyncMock( return_value=budget_row(budget_id="team-default-budget-1", max_budget=0.4, rpm_limit=10) ) - await _upsert_budget_and_membership( + await upsert_budget_and_membership( mock_tx, team_id="team-default", user_id="user-unlinked", @@ -363,7 +363,7 @@ async def test_clone_on_write_from_shared_default(mock_tx, fake_user): ) ) - await _upsert_budget_and_membership( + await upsert_budget_and_membership( mock_tx, team_id="team-shared", user_id="user-shared", @@ -416,7 +416,7 @@ async def test_clone_on_write_clears_duration(mock_tx, fake_user): ) ) - await _upsert_budget_and_membership( + await upsert_budget_and_membership( mock_tx, team_id="team-shared", user_id="user-shared", @@ -444,7 +444,7 @@ async def test_clone_on_write_clears_duration(mock_tx, fake_user): async def test_private_budget_updates_in_place(mock_tx, fake_user): mock_tx.litellm_budgettable.find_unique = AsyncMock(return_value=budget_row(max_budget=10.0)) - await _upsert_budget_and_membership( + await upsert_budget_and_membership( mock_tx, team_id="team-mixed", user_id="user-private", diff --git a/tests/unit/proxy/db/db_transaction_queue/test_redis_update_buffer.py b/tests/unit/proxy/db/db_transaction_queue/test_redis_update_buffer.py index e04e2402e1b..5ecc8784384 100644 --- a/tests/unit/proxy/db/db_transaction_queue/test_redis_update_buffer.py +++ b/tests/unit/proxy/db/db_transaction_queue/test_redis_update_buffer.py @@ -675,7 +675,7 @@ class _ListRedis: async def test_store_spend_logs_in_redis_drops_oldest_rows_past_the_cap(): redis = _ListRedis() buffer = RedisUpdateBuffer(redis_cache=redis) - buffer._should_commit_spend_updates_to_redis = MagicMock(return_value=True) + buffer.should_commit_spend_updates_to_redis = MagicMock(return_value=True) assert await buffer.store_spend_logs_in_redis([{"request_id": "old"}, {"request_id": "mid"}], max_rows=2) is True assert await buffer.store_spend_logs_in_redis([{"request_id": "new"}], max_rows=2) is True @@ -697,7 +697,7 @@ async def test_store_spend_logs_in_redis_reports_failure_without_redis(): async def test_store_spend_logs_in_redis_is_off_unless_transaction_buffering_is_enabled(): redis = _ListRedis() buffer = RedisUpdateBuffer(redis_cache=redis) - buffer._should_commit_spend_updates_to_redis = MagicMock(return_value=False) + buffer.should_commit_spend_updates_to_redis = MagicMock(return_value=False) assert await buffer.store_spend_logs_in_redis([{"request_id": "a"}]) is False assert redis.rows == [] diff --git a/tests/unit/proxy/db/test_db_spend_update_writer.py b/tests/unit/proxy/db/test_db_spend_update_writer.py index 8822ecd5bd0..dc9b8fd395b 100644 --- a/tests/unit/proxy/db/test_db_spend_update_writer.py +++ b/tests/unit/proxy/db/test_db_spend_update_writer.py @@ -997,7 +997,7 @@ async def test_commit_spend_updates_to_db_increments_agent_spend(): "agent_list_transactions": {agent_id: response_cost}, } - with patch("litellm.proxy.utils._raise_failed_update_spend_exception"): + with patch("litellm.proxy.utils.raise_failed_update_spend_exception"): await db_writer._commit_spend_updates_to_db( prisma_client=mock_prisma_client, n_retry_times=0, @@ -2029,7 +2029,7 @@ async def test_commit_key_spend_updates_includes_last_active(): before_call = datetime.now(timezone.utc) - with patch("litellm.proxy.utils._raise_failed_update_spend_exception"): + with patch("litellm.proxy.utils.raise_failed_update_spend_exception"): await db_writer._commit_spend_updates_to_db( prisma_client=mock_prisma_client, n_retry_times=0, @@ -3608,7 +3608,7 @@ async def test_commit_spend_updates_to_db_does_not_stamp_key_settings_updated_at "agent_list_transactions": {}, } - with patch("litellm.proxy.utils._raise_failed_update_spend_exception"): + with patch("litellm.proxy.utils.raise_failed_update_spend_exception"): await db_writer._commit_spend_updates_to_db( prisma_client=mock_prisma_client, n_retry_times=0, diff --git a/tests/unit/proxy/db/test_update_daily_tag_spend.py b/tests/unit/proxy/db/test_update_daily_tag_spend.py index dda41ab4543..d720f607942 100644 --- a/tests/unit/proxy/db/test_update_daily_tag_spend.py +++ b/tests/unit/proxy/db/test_update_daily_tag_spend.py @@ -14,23 +14,23 @@ async def test_update_daily_tag_spend_delegates_to_tag_commit_writer(): prisma_client = MagicMock() proxy_logging_obj = MagicMock() redis_update_buffer = MagicMock() - redis_update_buffer._should_commit_spend_updates_to_redis.return_value = False + redis_update_buffer.should_commit_spend_updates_to_redis.return_value = False proxy_logging_obj.db_spend_update_writer = MagicMock() proxy_logging_obj.db_spend_update_writer.redis_update_buffer = redis_update_buffer - proxy_logging_obj.db_spend_update_writer._commit_daily_tag_spend_to_db = AsyncMock() - proxy_logging_obj.db_spend_update_writer._commit_daily_tag_spend_to_db_with_redis = AsyncMock() + proxy_logging_obj.db_spend_update_writer.commit_daily_tag_spend_to_db = AsyncMock() + proxy_logging_obj.db_spend_update_writer.commit_daily_tag_spend_to_db_with_redis = AsyncMock() await update_daily_tag_spend( prisma_client, proxy_logging_obj, ) - proxy_logging_obj.db_spend_update_writer._commit_daily_tag_spend_to_db.assert_awaited_once_with( + proxy_logging_obj.db_spend_update_writer.commit_daily_tag_spend_to_db.assert_awaited_once_with( prisma_client=prisma_client, n_retry_times=3, proxy_logging_obj=proxy_logging_obj, ) - proxy_logging_obj.db_spend_update_writer._commit_daily_tag_spend_to_db_with_redis.assert_not_awaited() + proxy_logging_obj.db_spend_update_writer.commit_daily_tag_spend_to_db_with_redis.assert_not_awaited() @pytest.mark.asyncio @@ -38,11 +38,11 @@ async def test_update_daily_tag_spend_logs_error_and_does_not_raise(): prisma_client = MagicMock() proxy_logging_obj = MagicMock() redis_update_buffer = MagicMock() - redis_update_buffer._should_commit_spend_updates_to_redis.return_value = False + redis_update_buffer.should_commit_spend_updates_to_redis.return_value = False proxy_logging_obj.db_spend_update_writer = MagicMock() proxy_logging_obj.db_spend_update_writer.redis_update_buffer = redis_update_buffer - proxy_logging_obj.db_spend_update_writer._commit_daily_tag_spend_to_db = AsyncMock(side_effect=ValueError("boom")) - proxy_logging_obj.db_spend_update_writer._commit_daily_tag_spend_to_db_with_redis = AsyncMock() + proxy_logging_obj.db_spend_update_writer.commit_daily_tag_spend_to_db = AsyncMock(side_effect=ValueError("boom")) + proxy_logging_obj.db_spend_update_writer.commit_daily_tag_spend_to_db_with_redis = AsyncMock() with patch("litellm.proxy.utils.verbose_proxy_logger.error") as error_logger: await update_daily_tag_spend( @@ -50,7 +50,7 @@ async def test_update_daily_tag_spend_logs_error_and_does_not_raise(): proxy_logging_obj, ) - proxy_logging_obj.db_spend_update_writer._commit_daily_tag_spend_to_db.assert_awaited_once() + proxy_logging_obj.db_spend_update_writer.commit_daily_tag_spend_to_db.assert_awaited_once() error_logger.assert_called_once() @@ -59,23 +59,23 @@ async def test_update_daily_tag_spend_uses_redis_writer_when_enabled(): prisma_client = MagicMock() proxy_logging_obj = MagicMock() redis_update_buffer = MagicMock() - redis_update_buffer._should_commit_spend_updates_to_redis.return_value = True + redis_update_buffer.should_commit_spend_updates_to_redis.return_value = True proxy_logging_obj.db_spend_update_writer = MagicMock() - proxy_logging_obj.db_spend_update_writer._commit_daily_tag_spend_to_db = AsyncMock() + proxy_logging_obj.db_spend_update_writer.commit_daily_tag_spend_to_db = AsyncMock() proxy_logging_obj.db_spend_update_writer.redis_update_buffer = redis_update_buffer - proxy_logging_obj.db_spend_update_writer._commit_daily_tag_spend_to_db_with_redis = AsyncMock() + proxy_logging_obj.db_spend_update_writer.commit_daily_tag_spend_to_db_with_redis = AsyncMock() await update_daily_tag_spend( prisma_client, proxy_logging_obj, ) - proxy_logging_obj.db_spend_update_writer._commit_daily_tag_spend_to_db_with_redis.assert_awaited_once_with( + proxy_logging_obj.db_spend_update_writer.commit_daily_tag_spend_to_db_with_redis.assert_awaited_once_with( prisma_client=prisma_client, n_retry_times=3, proxy_logging_obj=proxy_logging_obj, ) - proxy_logging_obj.db_spend_update_writer._commit_daily_tag_spend_to_db.assert_not_awaited() + proxy_logging_obj.db_spend_update_writer.commit_daily_tag_spend_to_db.assert_not_awaited() @pytest.mark.asyncio diff --git a/tests/unit/proxy/decisions_endpoints/test_endpoints.py b/tests/unit/proxy/decisions_endpoints/test_endpoints.py index 6b1ac9e3404..ccfa12e10d6 100644 --- a/tests/unit/proxy/decisions_endpoints/test_endpoints.py +++ b/tests/unit/proxy/decisions_endpoints/test_endpoints.py @@ -16,7 +16,7 @@ from starlette.routing import Match import litellm from litellm.proxy._lazy_features import LAZY_FEATURES, LazyFeature, attach_lazy_features -from litellm.proxy.decisions_endpoints.endpoints import decisions +from litellm.proxy.decisions_endpoints.endpoints import decisions, systemone from litellm.proxy.pass_through_endpoints.pass_through_endpoints import SafeRouteAdder from litellm.proxy.proxy_server import ( app, @@ -77,7 +77,7 @@ def client(monkeypatch: pytest.MonkeyPatch) -> Iterator[TestClient]: litellm.in_memory_llm_clients_cache.flush_cache() -@pytest.mark.parametrize("endpoint", ("/v1/decisions", "/decisions")) +@pytest.mark.parametrize("endpoint", ("/v1/systemone", "/systemone")) def test_proxy_decisions_route_returns_answers_and_cost( client: TestClient, respx_mock: respx.MockRouter, @@ -126,7 +126,7 @@ def test_proxy_decisions_dispatches_typesafe_deployment( upstream: Final = respx_mock.post("https://api.typesafe.ai/v1/systemone").respond(json=_RESPONSE) response: Final = client.post( - "/v1/decisions", + "/v1/systemone", json={ "model": "jev", "state": {"source": "proxy-test"}, @@ -165,7 +165,7 @@ def test_proxy_decisions_sends_the_env_key_to_the_deployment_api_base( monkeypatch.setattr(litellm.proxy.proxy_server, "llm_router", router) upstream: Final = respx_mock.post("https://egress.example/perplexity/v1/decisions").respond(json=_RESPONSE) - response: Final = client.post("/v1/decisions", json=_REQUEST) + response: Final = client.post("/v1/systemone", json=_REQUEST) assert response.status_code == 200, response.text assert upstream.call_count == 1 @@ -177,7 +177,7 @@ def test_proxy_decisions_unknown_model_is_a_client_error( respx_mock: respx.MockRouter, ) -> None: response: Final = client.post( - "/v1/decisions", + "/v1/systemone", json={ "model": "missing-model", "state": "review", @@ -210,7 +210,7 @@ def test_proxy_decisions_missing_required_field_is_a_client_error( ) -> None: upstream: Final = respx_mock.post("https://api.perplexity.ai/v1/decisions").respond(json=_RESPONSE) - response: Final = client.post("/v1/decisions", json=request_body) + response: Final = client.post("/v1/systemone", json=request_body) assert response.status_code == 400, response.text assert not upstream.called @@ -238,7 +238,7 @@ def test_proxy_decisions_dispatches_strands_decider( upstream: Final = respx_mock.post("https://strands.example/v1/systemone").respond(json=_STRANDS_RESPONSE) response: Final = client.post( - "/v1/decisions", + "/v1/systemone", json={ "model": "strands", "state": {"source": "proxy-test"}, @@ -266,7 +266,7 @@ def test_proxy_decisions_without_model_uses_the_proxy_default_model( upstream: Final = respx_mock.post("https://api.perplexity.ai/v1/decisions").respond(json=_RESPONSE) response: Final = client.post( - "/v1/decisions", json={key: value for key, value in _REQUEST.items() if key != "model"} + "/v1/systemone", json={key: value for key, value in _REQUEST.items() if key != "model"} ) assert response.status_code == 200, response.text @@ -275,6 +275,261 @@ def test_proxy_decisions_without_model_uses_the_proxy_default_model( assert json.loads(upstream.calls[0].request.content)["model"] == "pplx-decider-v1-27b" +_OPENAI_FORMAT_REQUEST: Final[Mapping[str, object]] = { + "model": "decider", + "input": [ + { + "role": "user", + "content": [ + {"type": "input_text", "text": "The package arrived with a broken screen."}, + {"type": "input_text", "text": "I want a refund."}, + ], + }, + {"role": "user", "content": "Order 1234."}, + ], + "questions": [ + {"type": "predicate", "name": "damaged", "instructions": "Does the customer report a damaged item?"}, + { + "type": "choice", + "instructions": "Should we refund?", + "choices": [{"value": True, "description": "Refund now"}, {"value": "escalate"}], + }, + { + "type": "score", + "name": "severity", + "instructions": "How severe is the issue?", + "levels": [{"label": "minor"}, {"label": "major", "description": "Product unusable"}], + }, + {"type": "predicate", "name": "fraud", "instructions": "Is this fraud?"}, + ], + "safety_identifier": "end-user-1", +} +_SYSTEMONE_ANSWERS_FOR_OPENAI_REQUEST: Final[Mapping[str, object]] = { + "model": "pplx-decider-v1-27b", + "answers": { + "0": {"type": "noul", "noul": 0.95}, + "1": {"type": "choice", "choice": "true", "confidence": 0.8, "probabilities": {"true": 0.9, "escalate": 0.1}}, + "2": { + "type": "score", + "score": 0.7, + "confidence": 0.6, + "legend": {"0": "minor", "1": "major: Product unusable"}, + "probabilities": {"0": 0.3, "1": 0.7}, + }, + }, + "usage": {"input_tokens": _INPUT_TOKENS, "output_tokens": _OUTPUT_TOKENS}, +} + +_OPENAI_FORMAT_ANSWERS: Final = [ + {"type": "predicate", "name": "damaged", "probability": 0.95}, + { + "type": "choice", + "name": None, + "choice": True, + "probabilities": [{"value": True, "probability": 0.9}, {"value": "escalate", "probability": 0.1}], + "confidence": 0.8, + }, + { + "type": "score", + "name": "severity", + "score": 0.7, + "probabilities": [ + {"value": 0, "label": "minor", "probability": 0.3}, + {"value": 1, "label": "major", "probability": 0.7}, + ], + "confidence": 0.6, + }, + {"type": "refusal", "name": "fraud"}, +] + + +@pytest.mark.parametrize("endpoint", ("/v1/decisions", "/decisions")) +def test_openai_format_decisions_translate_through_systemone( + client: TestClient, + respx_mock: respx.MockRouter, + endpoint: str, +) -> None: + upstream: Final = respx_mock.post("https://api.perplexity.ai/v1/decisions").respond( + json=_SYSTEMONE_ANSWERS_FOR_OPENAI_REQUEST + ) + + response: Final = client.post(endpoint, json=_OPENAI_FORMAT_REQUEST) + + assert response.status_code == 200, response.text + assert json.loads(upstream.calls[0].request.content) == { + "model": "pplx-decider-v1-27b", + "state": "The package arrived with a broken screen.\n\nI want a refund.\n\nOrder 1234.", + "questions": { + "0": {"type": "noul", "instructions": "Does the customer report a damaged item?"}, + "1": { + "type": "choice", + "instructions": "Should we refund?", + "criteria": {"true": "Refund now", "escalate": None}, + }, + "2": { + "type": "score", + "instructions": "How severe is the issue?", + "criteria": ["minor", "major: Product unusable"], + }, + "3": {"type": "noul", "instructions": "Is this fraud?"}, + }, + } + body: Final = response.json() + assert body["model"] == _SYSTEMONE_ANSWERS_FOR_OPENAI_REQUEST["model"] + assert body["answers"] == _OPENAI_FORMAT_ANSWERS + assert body["usage"] == { + "input_tokens": _INPUT_TOKENS, + "input_tokens_details": {"cached_tokens": 0, "cache_write_tokens": 0}, + "output_tokens": _OUTPUT_TOKENS, + "output_tokens_details": {"reasoning_tokens": 0}, + "total_tokens": _INPUT_TOKENS + _OUTPUT_TOKENS, + } + assert float(response.headers["x-litellm-response-cost"]) > 0 + + +@pytest.mark.parametrize( + ("endpoint", "request_body", "upstream_response"), + ( + ("/v1/systemone", _REQUEST, _RESPONSE), + ("/v1/decisions", _OPENAI_FORMAT_REQUEST, _SYSTEMONE_ANSWERS_FOR_OPENAI_REQUEST), + ), + ids=("systemone", "openai_format"), +) +def test_decisions_return_guardrail_information_when_requested( + client: TestClient, + respx_mock: respx.MockRouter, + endpoint: str, + request_body: Mapping[str, object], + upstream_response: Mapping[str, object], +) -> None: + respx_mock.post("https://api.perplexity.ai/v1/decisions").respond(json=upstream_response) + + response: Final = client.post(endpoint, json={**request_body, "include_guardrail_response": True}) + + assert response.status_code == 200, response.text + assert response.json()["guardrail_information"] == [] + + +@pytest.mark.parametrize( + "request_body", + ( + _REQUEST, + { + "model": "decider", + "input": "review", + "questions": [ + { + "type": "choice", + "instructions": "Pick one", + "choices": [{"value": True}, {"value": "true"}], + } + ], + }, + { + "model": "decider", + "input": [ + {"role": "user", "content": [{"type": "input_image", "image_url": "data:image/png;base64,AA=="}]} + ], + "questions": [{"type": "predicate", "instructions": "Is this a defect?"}], + }, + ), + ids=("systemone_body", "colliding_choice_values", "image_input"), +) +def test_openai_format_decisions_rejects_bodies_it_cannot_translate( + client: TestClient, + respx_mock: respx.MockRouter, + request_body: Mapping[str, object], +) -> None: + upstream: Final = respx_mock.post("https://api.perplexity.ai/v1/decisions").respond(json=_RESPONSE) + + response: Final = client.post("/v1/decisions", json=request_body) + + assert response.status_code == 400, response.text + assert not upstream.called + + +@pytest.mark.parametrize("endpoint", ("/v1/systemone", "/v1/decisions")) +@pytest.mark.parametrize("raw_body", (b"", b"{not json"), ids=("empty", "malformed")) +def test_a_body_that_is_not_json_is_a_client_error( + client: TestClient, + respx_mock: respx.MockRouter, + endpoint: str, + raw_body: bytes, +) -> None: + upstream: Final = respx_mock.post("https://api.perplexity.ai/v1/decisions").respond(json=_RESPONSE) + + response: Final = client.post(endpoint, content=raw_body, headers={"Content-Type": "application/json"}) + + assert response.status_code == 400, response.text + assert not upstream.called + + +def test_openai_format_decisions_reach_an_openai_deployment_unchanged_including_images( + client: TestClient, + monkeypatch: pytest.MonkeyPatch, + respx_mock: respx.MockRouter, +) -> None: + monkeypatch.setattr(litellm, "api_base", None) + monkeypatch.delenv("OPENAI_BASE_URL", raising=False) + monkeypatch.delenv("OPENAI_API_BASE", raising=False) + monkeypatch.setattr( + litellm.proxy.proxy_server, + "llm_router", + litellm.Router( + model_list=[{"model_name": "decider", "litellm_params": {"model": "openai/gpt-6-luna", "api_key": "k"}}] + ), + ) + image_message: Final = { + "type": "message", + "role": "user", + "content": [{"type": "input_image", "image_url": "data:image/png;base64,AA==", "detail": "low"}], + } + request_body: Final = {**_OPENAI_FORMAT_REQUEST, "input": [*_OPENAI_FORMAT_REQUEST["input"], image_message]} + cached_tokens: Final = 128 + upstream_usage: Final = { + "input_tokens": _INPUT_TOKENS, + "input_tokens_details": {"cached_tokens": cached_tokens, "cache_write_tokens": 0}, + "output_tokens": _OUTPUT_TOKENS, + "output_tokens_details": {"reasoning_tokens": 0}, + "total_tokens": _INPUT_TOKENS + _OUTPUT_TOKENS, + } + upstream: Final = respx_mock.post("https://api.openai.com/v1/decisions").respond( + json={"model": "gpt-6-luna", "answers": _OPENAI_FORMAT_ANSWERS, "usage": upstream_usage} + ) + + response: Final = client.post("/v1/decisions", json=request_body) + + assert response.status_code == 200, response.text + assert json.loads(upstream.calls[0].request.content) == { + "model": "gpt-6-luna", + "input": [ + { + "type": "message", + "role": "user", + "content": [ + {"type": "input_text", "text": "The package arrived with a broken screen."}, + {"type": "input_text", "text": "I want a refund."}, + ], + }, + {"type": "message", "role": "user", "content": "Order 1234."}, + image_message, + ], + "questions": _OPENAI_FORMAT_REQUEST["questions"], + "safety_identifier": "end-user-1", + } + body: Final = response.json() + assert body["answers"] == _OPENAI_FORMAT_ANSWERS + assert body["usage"] == upstream_usage + luna_cost: Final = litellm.model_cost["gpt-6-luna"] + expected_cost: Final = ( + (_INPUT_TOKENS - cached_tokens) * float(luna_cost["input_cost_per_token"]) + + cached_tokens * float(luna_cost["cache_read_input_token_cost"]) + + _OUTPUT_TOKENS * float(luna_cost["output_cost_per_token"]) + ) + assert expected_cost > 0 + assert float(response.headers["x-litellm-response-cost"]) == pytest.approx(expected_cost) + + def _decisions_feature() -> LazyFeature: return next(feature for feature in LAZY_FEATURES if feature.name == "decisions") @@ -301,6 +556,8 @@ def test_a_config_pass_through_at_v1_decisions_keeps_its_route_and_the_native_ap assert client.post("/v1/decisions", json={"model": "gpt-6-luna"}).json() == {"served_by": "pass-through"} assert _serving_endpoint(bare, "/v1/decisions") is pass_through assert _serving_endpoint(bare, "/decisions") is decisions + assert _serving_endpoint(bare, "/v1/systemone") is systemone + assert _serving_endpoint(bare, "/systemone") is systemone def test_with_lazy_routes_disabled_a_config_pass_through_at_v1_decisions_still_wins( diff --git a/tests/unit/proxy/google_endpoints/test_gemini_agents_endpoints.py b/tests/unit/proxy/google_endpoints/test_gemini_agents_endpoints.py index cddb0e526b4..79062f15eca 100644 --- a/tests/unit/proxy/google_endpoints/test_gemini_agents_endpoints.py +++ b/tests/unit/proxy/google_endpoints/test_gemini_agents_endpoints.py @@ -63,9 +63,7 @@ class TestMergeQueryParamsIntoData: assert "api_key" not in data def test_litellm_params_template_json_is_expanded(self): - template = json.dumps( - {"api_key": "AIzaFromTemplate", "api_base": "https://example.com"} - ) + template = json.dumps({"api_key": "AIzaFromTemplate", "api_base": "https://example.com"}) from urllib.parse import quote request = _make_request(f"litellm_params_template={quote(template)}") @@ -77,9 +75,7 @@ class TestMergeQueryParamsIntoData: assert "litellm_params_template" not in data def test_litellm_params_template_does_not_overwrite_existing(self): - template = json.dumps( - {"api_key": "FromTemplate", "custom_llm_provider": "openai"} - ) + template = json.dumps({"api_key": "FromTemplate", "custom_llm_provider": "openai"}) from urllib.parse import quote request = _make_request(f"litellm_params_template={quote(template)}") @@ -160,18 +156,14 @@ def _make_endpoint_request(query_string: str = "") -> MagicMock: @pytest.mark.asyncio -async def test_list_gemini_agents_passes_api_key_to_processor( - mock_srv, user_api_key_dict -): +async def test_list_gemini_agents_passes_api_key_to_processor(mock_srv, user_api_key_dict): from urllib.parse import quote from litellm.proxy.google_endpoints.agents_endpoints import list_gemini_agents template = json.dumps({"api_key": "AIzaListTest"}) - with patch( - "litellm.proxy.google_endpoints.agents_endpoints.ProxyBaseLLMRequestProcessing" - ) as MockProcessor: + with patch("litellm.proxy.google_endpoints.agents_endpoints.ProxyBaseLLMRequestProcessing") as MockProcessor: instance = MockProcessor.return_value instance.base_process_llm_request = AsyncMock(return_value=MagicMock()) @@ -188,18 +180,14 @@ async def test_list_gemini_agents_passes_api_key_to_processor( @pytest.mark.asyncio -async def test_get_gemini_agent_passes_api_key_to_processor( - mock_srv, user_api_key_dict -): +async def test_get_gemini_agent_passes_api_key_to_processor(mock_srv, user_api_key_dict): from urllib.parse import quote from litellm.proxy.google_endpoints.agents_endpoints import get_gemini_agent template = json.dumps({"api_key": "AIzaGetTest"}) - with patch( - "litellm.proxy.google_endpoints.agents_endpoints.ProxyBaseLLMRequestProcessing" - ) as MockProcessor: + with patch("litellm.proxy.google_endpoints.agents_endpoints.ProxyBaseLLMRequestProcessing") as MockProcessor: instance = MockProcessor.return_value instance.base_process_llm_request = AsyncMock(return_value=MagicMock()) @@ -218,18 +206,14 @@ async def test_get_gemini_agent_passes_api_key_to_processor( @pytest.mark.asyncio -async def test_delete_gemini_agent_passes_api_key_to_processor( - mock_srv, user_api_key_dict -): +async def test_delete_gemini_agent_passes_api_key_to_processor(mock_srv, user_api_key_dict): from urllib.parse import quote from litellm.proxy.google_endpoints.agents_endpoints import delete_gemini_agent template = json.dumps({"api_key": "AIzaDeleteTest"}) - with patch( - "litellm.proxy.google_endpoints.agents_endpoints.ProxyBaseLLMRequestProcessing" - ) as MockProcessor: + with patch("litellm.proxy.google_endpoints.agents_endpoints.ProxyBaseLLMRequestProcessing") as MockProcessor: instance = MockProcessor.return_value instance.base_process_llm_request = AsyncMock(return_value=MagicMock()) @@ -248,9 +232,7 @@ async def test_delete_gemini_agent_passes_api_key_to_processor( @pytest.mark.asyncio -async def test_list_gemini_agent_versions_passes_api_key_to_processor( - mock_srv, user_api_key_dict -): +async def test_list_gemini_agent_versions_passes_api_key_to_processor(mock_srv, user_api_key_dict): from urllib.parse import quote from litellm.proxy.google_endpoints.agents_endpoints import ( @@ -259,9 +241,7 @@ async def test_list_gemini_agent_versions_passes_api_key_to_processor( template = json.dumps({"api_key": "AIzaVersionsTest"}) - with patch( - "litellm.proxy.google_endpoints.agents_endpoints.ProxyBaseLLMRequestProcessing" - ) as MockProcessor: + with patch("litellm.proxy.google_endpoints.agents_endpoints.ProxyBaseLLMRequestProcessing") as MockProcessor: instance = MockProcessor.return_value instance.base_process_llm_request = AsyncMock(return_value=MagicMock()) @@ -280,17 +260,13 @@ async def test_list_gemini_agent_versions_passes_api_key_to_processor( @pytest.mark.asyncio -async def test_get_gemini_agent_name_not_overwritten_by_query_param( - mock_srv, user_api_key_dict -): +async def test_get_gemini_agent_name_not_overwritten_by_query_param(mock_srv, user_api_key_dict): """Path-param ``name`` must not be replaced by an attacker-controlled query param.""" from urllib.parse import quote from litellm.proxy.google_endpoints.agents_endpoints import get_gemini_agent - with patch( - "litellm.proxy.google_endpoints.agents_endpoints.ProxyBaseLLMRequestProcessing" - ) as MockProcessor: + with patch("litellm.proxy.google_endpoints.agents_endpoints.ProxyBaseLLMRequestProcessing") as MockProcessor: instance = MockProcessor.return_value instance.base_process_llm_request = AsyncMock(return_value=MagicMock()) @@ -299,9 +275,7 @@ async def test_get_gemini_agent_name_not_overwritten_by_query_param( # ``api_key`` is supplied via the JSON template (required for non-admin # callers — see test_*_non_admin_without_api_key_is_rejected below). template = json.dumps({"api_key": "AIzaTest"}) - request = _make_endpoint_request( - f"name=INJECTED&litellm_params_template={quote(template)}" - ) + request = _make_endpoint_request(f"name=INJECTED&litellm_params_template={quote(template)}") await get_gemini_agent( request=request, name="real-agent", @@ -321,9 +295,7 @@ async def test_list_agents_template_via_query_param(mock_srv, user_api_key_dict) template = json.dumps({"api_key": "TemplateKey", "vertex_project": "proj-x"}) - with patch( - "litellm.proxy.google_endpoints.agents_endpoints.ProxyBaseLLMRequestProcessing" - ) as MockProcessor: + with patch("litellm.proxy.google_endpoints.agents_endpoints.ProxyBaseLLMRequestProcessing") as MockProcessor: instance = MockProcessor.return_value instance.base_process_llm_request = AsyncMock(return_value=MagicMock()) @@ -356,9 +328,7 @@ def proxy_admin_user_api_key_dict(): @pytest.mark.asyncio -async def test_list_agents_non_admin_without_api_key_is_rejected( - mock_srv, user_api_key_dict -): +async def test_list_agents_non_admin_without_api_key_is_rejected(mock_srv, user_api_key_dict): """Non-admin callers must supply an explicit api_key — the proxy must not silently fall back to the operator's shared GOOGLE_API_KEY/GEMINI_API_KEY. """ @@ -366,9 +336,7 @@ async def test_list_agents_non_admin_without_api_key_is_rejected( from litellm.proxy.google_endpoints.agents_endpoints import list_gemini_agents - with patch( - "litellm.proxy.google_endpoints.agents_endpoints.ProxyBaseLLMRequestProcessing" - ) as MockProcessor: + with patch("litellm.proxy.google_endpoints.agents_endpoints.ProxyBaseLLMRequestProcessing") as MockProcessor: instance = MockProcessor.return_value instance.base_process_llm_request = AsyncMock(return_value=MagicMock()) @@ -385,16 +353,12 @@ async def test_list_agents_non_admin_without_api_key_is_rejected( @pytest.mark.asyncio -async def test_delete_agent_non_admin_without_api_key_is_rejected( - mock_srv, user_api_key_dict -): +async def test_delete_agent_non_admin_without_api_key_is_rejected(mock_srv, user_api_key_dict): from fastapi import HTTPException from litellm.proxy.google_endpoints.agents_endpoints import delete_gemini_agent - with patch( - "litellm.proxy.google_endpoints.agents_endpoints.ProxyBaseLLMRequestProcessing" - ) as MockProcessor: + with patch("litellm.proxy.google_endpoints.agents_endpoints.ProxyBaseLLMRequestProcessing") as MockProcessor: instance = MockProcessor.return_value instance.base_process_llm_request = AsyncMock(return_value=MagicMock()) @@ -411,19 +375,15 @@ async def test_delete_agent_non_admin_without_api_key_is_rejected( @pytest.mark.asyncio -async def test_create_agent_non_admin_without_api_key_is_rejected( - mock_srv, user_api_key_dict -): +async def test_create_agent_non_admin_without_api_key_is_rejected(mock_srv, user_api_key_dict): from fastapi import HTTPException from litellm.proxy.google_endpoints.agents_endpoints import create_gemini_agent with ( + patch("litellm.proxy.google_endpoints.agents_endpoints.ProxyBaseLLMRequestProcessing") as MockProcessor, patch( - "litellm.proxy.google_endpoints.agents_endpoints.ProxyBaseLLMRequestProcessing" - ) as MockProcessor, - patch( - "litellm.proxy.google_endpoints.agents_endpoints._read_request_body", + "litellm.proxy.google_endpoints.agents_endpoints.read_request_body", new=AsyncMock(return_value={"name": "agent-1", "base_agent": "waverunner"}), ), ): @@ -442,15 +402,11 @@ async def test_create_agent_non_admin_without_api_key_is_rejected( @pytest.mark.asyncio -async def test_list_agents_proxy_admin_may_use_env_fallback( - mock_srv, proxy_admin_user_api_key_dict -): +async def test_list_agents_proxy_admin_may_use_env_fallback(mock_srv, proxy_admin_user_api_key_dict): """Proxy admins (master key) keep the env-fallback convenience.""" from litellm.proxy.google_endpoints.agents_endpoints import list_gemini_agents - with patch( - "litellm.proxy.google_endpoints.agents_endpoints.ProxyBaseLLMRequestProcessing" - ) as MockProcessor: + with patch("litellm.proxy.google_endpoints.agents_endpoints.ProxyBaseLLMRequestProcessing") as MockProcessor: instance = MockProcessor.return_value instance.base_process_llm_request = AsyncMock(return_value=MagicMock()) diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/noma/test_noma.py b/tests/unit/proxy/guardrails/guardrail_hooks/noma/test_noma.py index e23796705b3..f06ba3ee68e 100644 --- a/tests/unit/proxy/guardrails/guardrail_hooks/noma/test_noma.py +++ b/tests/unit/proxy/guardrails/guardrail_hooks/noma/test_noma.py @@ -11,7 +11,7 @@ import litellm from litellm.caching.caching import DualCache from litellm.integrations.custom_logger import CustomLogger from litellm.proxy._types import UserAPIKeyAuth -from litellm.proxy.guardrails.guardrail_hooks.presidio import _OPTIONAL_PresidioPIIMasking +from litellm.proxy.guardrails.guardrail_hooks.presidio import OPTIONAL_PresidioPIIMasking from litellm.types.guardrails import GuardrailEventHooks from litellm.types.utils import StandardLoggingPayload from tests._vcr_conftest_common import install_live_call_probe, record_vcr_outcome @@ -34,7 +34,7 @@ async def test_standard_logging_payload_includes_guardrail_information(): """ test_custom_logger = CustomLoggerForTesting() litellm.callbacks = [test_custom_logger] - presidio_guard = _OPTIONAL_PresidioPIIMasking( + presidio_guard = OPTIONAL_PresidioPIIMasking( guardrail_name="presidio_guard", event_hook=GuardrailEventHooks.pre_call, presidio_analyzer_api_base="https://mock-presidio-analyzer.com/", diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py b/tests/unit/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py index d0b0799817b..01be28470d1 100644 --- a/tests/unit/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py +++ b/tests/unit/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py @@ -3520,7 +3520,7 @@ async def test_chat_completion_modify_response_exception_streaming_logging_obj_n raise exc with ( - patch("litellm.proxy.proxy_server._read_request_body", AsyncMock(return_value=request_data)), + patch("litellm.proxy.proxy_server.read_request_body", AsyncMock(return_value=request_data)), patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging), patch( "litellm.proxy.proxy_server.ProxyBaseLLMRequestProcessing.base_process_llm_request", diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/test_presidio.py b/tests/unit/proxy/guardrails/guardrail_hooks/test_presidio.py index ff035a3741e..0cac6228085 100644 --- a/tests/unit/proxy/guardrails/guardrail_hooks/test_presidio.py +++ b/tests/unit/proxy/guardrails/guardrail_hooks/test_presidio.py @@ -23,7 +23,7 @@ from litellm.caching.caching import DualCache from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.guardrails.guardrail_hooks.presidio import ( PresidioPerRequestConfig, - _OPTIONAL_PresidioPIIMasking, + OPTIONAL_PresidioPIIMasking, ) from litellm.exceptions import GuardrailRaisedException from litellm.types.guardrails import LitellmParams, PiiAction, PiiEntityType @@ -80,7 +80,7 @@ def _make_mock_session_iterator(json_response, status=200, content_type="applica @pytest.fixture def presidio_guardrail(): """Create a Presidio guardrail instance for testing""" - return _OPTIONAL_PresidioPIIMasking( + return OPTIONAL_PresidioPIIMasking( mock_testing=True, output_parse_pii=False, pii_entities_config={ @@ -638,7 +638,7 @@ async def test_logging_only_does_not_mask_pre_call_request(mock_user_api_key, mo causing the model's response to contain anonymization tokens (e.g. ) instead of the real output. """ - presidio_guardrail = _OPTIONAL_PresidioPIIMasking( + presidio_guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, logging_only=True, pii_entities_config={PiiEntityType.PHONE_NUMBER: PiiAction.MASK}, @@ -677,7 +677,7 @@ async def test_presidio_sets_guardrail_information_in_request_data(): This validates that add_standard_logging_guardrail_information_to_request_data correctly sets the guardrail information that will be used for logging. """ - presidio = _OPTIONAL_PresidioPIIMasking( + presidio = OPTIONAL_PresidioPIIMasking( guardrail_name="test_presidio", output_parse_pii=True, mock_testing=True, @@ -735,7 +735,7 @@ async def test_request_data_flows_to_apply_guardrail(): This validates the fix where guardrail translation handler passes data as request_data to apply_guardrail so guardrails can store metadata for logging. """ - presidio = _OPTIONAL_PresidioPIIMasking( + presidio = OPTIONAL_PresidioPIIMasking( guardrail_name="test_presidio", output_parse_pii=True, mock_testing=True, @@ -775,7 +775,7 @@ async def test_output_masking_apply_to_output_only(mock_user_api_key): Ensure output masking runs when apply_to_output is enabled. """ - presidio = _OPTIONAL_PresidioPIIMasking( + presidio = OPTIONAL_PresidioPIIMasking( mock_testing=True, apply_to_output=True, pii_entities_config={PiiEntityType.CREDIT_CARD: PiiAction.MASK}, @@ -842,8 +842,8 @@ async def test_presidio_filter_scope_initializer(monkeypatch): import litellm.proxy.guardrails.guardrail_hooks.presidio as presidio_mod import litellm.proxy.guardrails.guardrail_initializers as gi - monkeypatch.setattr(presidio_mod, "_OPTIONAL_PresidioPIIMasking", DummyGuardrail, raising=False) - monkeypatch.setattr(gi, "_OPTIONAL_PresidioPIIMasking", DummyGuardrail, raising=False) + monkeypatch.setattr(presidio_mod, "OPTIONAL_PresidioPIIMasking", DummyGuardrail, raising=False) + monkeypatch.setattr(gi, "OPTIONAL_PresidioPIIMasking", DummyGuardrail, raising=False) # input-only created.clear() @@ -982,7 +982,7 @@ async def test_analyze_text_with_empty_string(): Should return empty list without making API call to Presidio. """ - presidio = _OPTIONAL_PresidioPIIMasking( + presidio = OPTIONAL_PresidioPIIMasking( presidio_analyzer_api_base="http://test:5002/", presidio_anonymizer_api_base="http://test:5001/", output_parse_pii=False, @@ -1015,7 +1015,7 @@ async def test_analyze_text_error_dict_handling(): When Presidio returns {'error': 'No text provided'}, should handle gracefully instead of crashing with TypeError. """ - presidio = _OPTIONAL_PresidioPIIMasking( + presidio = OPTIONAL_PresidioPIIMasking( presidio_analyzer_api_base="http://mock-presidio:5002/", presidio_anonymizer_api_base="http://mock-presidio:5001/", output_parse_pii=False, @@ -1044,7 +1044,7 @@ async def test_analyze_text_string_response_handling(): When Presidio returns a string (e.g. error message from websearch/hosted models), should handle gracefully instead of crashing with TypeError about mapping vs str. """ - presidio = _OPTIONAL_PresidioPIIMasking( + presidio = OPTIONAL_PresidioPIIMasking( presidio_analyzer_api_base="http://mock-presidio:5002/", presidio_anonymizer_api_base="http://mock-presidio:5001/", output_parse_pii=False, @@ -1069,7 +1069,7 @@ async def test_analyze_text_invalid_response_raises_when_block_configured(): When pii_entities_config has BLOCK and Presidio returns invalid response, should raise GuardrailRaisedException (fail-closed) rather than silently allowing content. """ - presidio = _OPTIONAL_PresidioPIIMasking( + presidio = OPTIONAL_PresidioPIIMasking( presidio_analyzer_api_base="http://mock-presidio:5002/", presidio_anonymizer_api_base="http://mock-presidio:5001/", output_parse_pii=False, @@ -1096,7 +1096,7 @@ async def test_analyze_text_invalid_response_raises_when_mask_configured(): When pii_entities_config has MASK and Presidio returns invalid response, should raise GuardrailRaisedException (fail-closed) because PII masking is expected. """ - presidio = _OPTIONAL_PresidioPIIMasking( + presidio = OPTIONAL_PresidioPIIMasking( presidio_analyzer_api_base="http://mock-presidio:5002/", presidio_anonymizer_api_base="http://mock-presidio:5001/", output_parse_pii=False, @@ -1125,7 +1125,7 @@ async def test_analyze_text_list_with_non_dict_items(): When Presidio returns a list containing strings (malformed response), should skip invalid items and return parsed valid ones. """ - presidio = _OPTIONAL_PresidioPIIMasking( + presidio = OPTIONAL_PresidioPIIMasking( presidio_analyzer_api_base="http://mock-presidio:5002/", presidio_anonymizer_api_base="http://mock-presidio:5001/", output_parse_pii=False, @@ -1210,7 +1210,7 @@ def test_filter_drops_low_score_detection(): """ Detections below the configured score threshold should be removed. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, presidio_score_thresholds={PiiEntityType.CREDIT_CARD: 0.8}, ) @@ -1224,7 +1224,7 @@ def test_filter_preserves_high_score_detection(): """ Detections meeting the score threshold should be preserved. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, presidio_score_thresholds={PiiEntityType.CREDIT_CARD: 0.8}, ) @@ -1239,7 +1239,7 @@ def test_no_thresholds_returns_all(): """ With no thresholds configured, all detections are kept. """ - guardrail = _OPTIONAL_PresidioPIIMasking(mock_testing=True) + guardrail = OPTIONAL_PresidioPIIMasking(mock_testing=True) analyze_results = [ {"entity_type": PiiEntityType.CREDIT_CARD, "score": 0.1, "start": 0, "end": 4}, { @@ -1258,7 +1258,7 @@ def test_entity_specific_threshold_only_applies_to_that_entity(): """ Entity-specific thresholds do not affect other entity types. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, presidio_score_thresholds={PiiEntityType.CREDIT_CARD: 0.8}, ) @@ -1282,7 +1282,7 @@ def test_filter_uses_default_all_threshold(): """ Default ALL threshold applies to any entity without a specific override. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, presidio_score_thresholds={"ALL": 0.75}, ) @@ -1305,7 +1305,7 @@ def test_entity_specific_overrides_default_threshold(): """ Entity-specific threshold should override the ALL default. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, presidio_score_thresholds={ "ALL": 0.8, @@ -1333,7 +1333,7 @@ async def test_anonymize_skips_when_no_detections_after_filter(): """ When all detections are filtered out, anonymize_text should return the original text. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, presidio_score_thresholds={PiiEntityType.CREDIT_CARD: 0.8}, ) @@ -1359,7 +1359,7 @@ def test_blocking_respects_threshold_filter(): """ Entities filtered out by score should not trigger blocking, but high-score detections should. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, pii_entities_config={PiiEntityType.CREDIT_CARD: PiiAction.BLOCK}, presidio_score_thresholds={PiiEntityType.CREDIT_CARD: 0.9}, @@ -1379,7 +1379,7 @@ def test_update_in_memory_applies_score_thresholds(): """ update_in_memory_litellm_params should refresh score thresholds. """ - guardrail = _OPTIONAL_PresidioPIIMasking(mock_testing=True) + guardrail = OPTIONAL_PresidioPIIMasking(mock_testing=True) assert guardrail.presidio_score_thresholds == {} params = LitellmParams( @@ -1458,7 +1458,7 @@ async def test_streaming_with_bytes_chunks_does_not_crash(mock_user_api_key): gracefully handle raw bytes in the stream instead of crashing with 'bytes' object has no attribute 'id'. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, apply_to_output=True, mock_redacted_text={"text": "redacted"}, @@ -1491,7 +1491,7 @@ def test_entity_deny_list_filters_detections(): """ Verify presidio_entities_deny_list removes matching entity types. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, presidio_entities_deny_list=["US_DRIVER_LICENSE"], ) @@ -1511,7 +1511,7 @@ def test_deny_list_and_score_threshold_combined(): """ Verify deny list + score threshold work together correctly. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, presidio_entities_deny_list=["US_DRIVER_LICENSE"], presidio_score_thresholds={"ALL": 0.8}, @@ -1538,7 +1538,7 @@ async def test_analyze_text_non_json_content_type_fail_closed(): Test that analyze_text raises GuardrailRaisedException when Presidio health endpoint returns text/html and fail-closed is enabled. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( presidio_analyzer_api_base="http://test-analyzer/", presidio_anonymizer_api_base="http://test-anonymizer/", pii_entities_config={"PERSON": PiiAction.BLOCK}, @@ -1569,7 +1569,7 @@ async def test_analyze_text_non_json_content_type_fail_open(): Test that analyze_text returns empty list when Presidio returns text/html and fail-closed is NOT enabled. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( presidio_analyzer_api_base="http://test-analyzer/", presidio_anonymizer_api_base="http://test-anonymizer/", mock_testing=False, @@ -1596,7 +1596,7 @@ async def test_analyze_text_http_error_status(): """ Test that analyze_text handles 5xx HTTP errors properly. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( presidio_analyzer_api_base="http://test-analyzer/", presidio_anonymizer_api_base="http://test-anonymizer/", pii_entities_config={"PERSON": PiiAction.BLOCK}, @@ -1625,7 +1625,7 @@ async def test_anonymize_text_non_json_content_type(): """ Test that anonymize_text raises Exception for non-JSON responses. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( presidio_analyzer_api_base="http://test-analyzer/", presidio_anonymizer_api_base="http://test-anonymizer/", mock_testing=False, @@ -1653,7 +1653,7 @@ async def test_anonymize_text_http_error_status(): """ Test that anonymize_text raises Exception on HTTP error. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( presidio_analyzer_api_base="http://test-analyzer/", presidio_anonymizer_api_base="http://test-anonymizer/", mock_testing=False, @@ -1684,7 +1684,7 @@ async def test_pii_tokens_stored_in_metadata_not_top_level(presidio_guardrail): providers like Anthropic, which reject unknown fields with 'pii_tokens: Extra inputs are not permitted'. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, output_parse_pii=True, pii_entities_config={ @@ -1743,7 +1743,7 @@ async def test_pii_tokens_in_metadata_used_for_unmasking(): Regression test: _process_response_for_pii must read pii_tokens from data['metadata']['pii_tokens'] and correctly unmask the response. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, output_parse_pii=True, ) @@ -1786,7 +1786,7 @@ def test_event_hook_auto_expansion_for_all_string_hooks(initial_hook): 'post_call' to event_hook regardless of the initial string hook value, not just when it's 'pre_call'. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, output_parse_pii=True, event_hook=initial_hook, @@ -1798,7 +1798,7 @@ def test_event_hook_auto_expansion_for_all_string_hooks(initial_hook): def test_event_hook_no_expansion_when_already_post_call(): """post_call alone should stay as-is — no expansion needed.""" - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, output_parse_pii=True, event_hook="post_call", @@ -1813,7 +1813,7 @@ async def test_metadata_none_does_not_crash(): Regression test: if metadata is explicitly None in request_data, the guardrail must not crash with TypeError on the write or read path. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, output_parse_pii=True, ) @@ -1859,7 +1859,7 @@ def test_unmask_exact_match_with_sequential_tokens(): Normal unmasking: LLM echoes numbered tokens verbatim → original PII restored. """ from litellm.proxy.guardrails.guardrail_hooks.presidio import ( - _OPTIONAL_PresidioPIIMasking, + OPTIONAL_PresidioPIIMasking, ) pii_tokens = { @@ -1867,7 +1867,7 @@ def test_unmask_exact_match_with_sequential_tokens(): "": "555-123-4567", } text = "Hello , your number is ." - result = _OPTIONAL_PresidioPIIMasking._unmask_pii_text(text, pii_tokens) + result = OPTIONAL_PresidioPIIMasking._unmask_pii_text(text, pii_tokens) assert result == "Hello John Smith, your number is 555-123-4567." @@ -1876,7 +1876,7 @@ def test_unmask_multiple_same_entity_type(): Two phone numbers get distinct numbered tokens and unmask correctly. """ from litellm.proxy.guardrails.guardrail_hooks.presidio import ( - _OPTIONAL_PresidioPIIMasking, + OPTIONAL_PresidioPIIMasking, ) pii_tokens = { @@ -1884,7 +1884,7 @@ def test_unmask_multiple_same_entity_type(): "": "555-222-0000", } text = "Call or ." - result = _OPTIONAL_PresidioPIIMasking._unmask_pii_text(text, pii_tokens) + result = OPTIONAL_PresidioPIIMasking._unmask_pii_text(text, pii_tokens) assert result == "Call 555-111-0000 or 555-222-0000." @@ -1894,7 +1894,7 @@ def test_unmask_graceful_degradation(): in the output — clean and readable, not garbage hex. """ from litellm.proxy.guardrails.guardrail_hooks.presidio import ( - _OPTIONAL_PresidioPIIMasking, + OPTIONAL_PresidioPIIMasking, ) pii_tokens = { @@ -1902,7 +1902,7 @@ def test_unmask_graceful_degradation(): } # LLM paraphrased instead of echoing the token text = "I see you provided a name." - result = _OPTIONAL_PresidioPIIMasking._unmask_pii_text(text, pii_tokens) + result = OPTIONAL_PresidioPIIMasking._unmask_pii_text(text, pii_tokens) # No change — no garbage, just clean text assert result == text @@ -1918,7 +1918,7 @@ async def test_anonymize_text_multiple_items_position_correctness(): Regression test: when multiple PII items exist, coordinates reference the ORIGINAL text. Processing in reverse order prevents coordinate drift. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( presidio_analyzer_api_base="http://test-analyzer/", presidio_anonymizer_api_base="http://test-anonymizer/", mock_testing=False, @@ -1985,7 +1985,7 @@ async def test_anthropic_native_response_unmasking(): Anthropic native dict responses (type='message') should be unmasked when output_parse_pii is enabled. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, output_parse_pii=True, ) @@ -2032,7 +2032,7 @@ async def test_anthropic_native_response_masking(): Anthropic native dict responses should be masked when apply_to_output is enabled. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, apply_to_output=True, ) @@ -2071,7 +2071,7 @@ async def test_anthropic_native_response_non_text_blocks_untouched(): Non-text blocks (tool_use, thinking) in Anthropic responses should be left untouched during unmasking. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, output_parse_pii=True, ) @@ -2123,7 +2123,7 @@ async def test_streaming_bytes_chunks_are_yielded_not_discarded(): through the streaming hook, not silently discarded. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, apply_to_output=True, ) @@ -2151,7 +2151,7 @@ async def test_streaming_unmask_path_bytes_passthrough(): """ Bytes chunks in the unmasking path should also pass through. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, output_parse_pii=True, ) @@ -2183,7 +2183,7 @@ async def test_apply_to_output_streaming_unknown_events_passthrough(): Regression test: /v1/responses-style event objects (neither bytes nor ModelResponseStream) must be preserved in order and not dropped. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, apply_to_output=True, ) @@ -2227,7 +2227,7 @@ async def test_apply_to_output_streaming_mixed_chunks_flushes_and_warns(): a buffered ModelResponseStream chunk followed by unknown responses-style events should be preserved, and masking skip should be visible via warnings. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, apply_to_output=True, ) @@ -2284,7 +2284,7 @@ async def test_apply_guardrail_unmask_on_response(output_parse_pii: bool) -> Non When input_type is 'response' and pii_tokens exist, apply_guardrail should unmask text instead of masking it. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( guardrail_name="test_presidio", output_parse_pii=output_parse_pii, mock_testing=True, @@ -2322,7 +2322,7 @@ async def test_standalone_scans_without_restoration_tokens(input_type: Literal[" """ Standalone callbacks retain scanning without tokens, including MCP results. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( guardrail_name="test_presidio", event_hook="post_mcp_call", output_parse_pii=True, @@ -2375,7 +2375,7 @@ async def test_apply_to_output_streaming_chat_chunks_are_masked_as_one_response( Structured chat completion chunks are buffered, assembled and masked as a whole, so a card number split across deltas cannot reach the caller. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, apply_to_output=True, mock_redacted_text={"text": "my card is "}, @@ -2406,7 +2406,7 @@ async def test_apply_to_output_streaming_bytes_after_chat_chunks_are_passed_thro Once structured chunks have been buffered, a trailing bytes frame belongs to the same stream and must be forwarded rather than treated as a new SSE stream. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, apply_to_output=True, mock_redacted_text={"text": "hello"}, @@ -2438,7 +2438,7 @@ async def test_apply_to_output_streaming_anthropic_sse_bytes_masks_text_split_ac bytes. Output masking must run over the whole content block so a card number split across text_delta events cannot reach the caller. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, apply_to_output=True, mock_redacted_text={"text": ""}, @@ -2488,7 +2488,7 @@ async def test_apply_to_output_streaming_anthropic_sse_bytes_masks_text_split_ac @pytest.mark.asyncio async def test_apply_to_output_streaming_anthropic_sse_bytes_without_pii_are_forwarded_unchanged(): - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, apply_to_output=True, mock_redacted_text={"text": "Hello world"}, @@ -2538,7 +2538,7 @@ def _gemini_sse(text: str) -> bytes: @pytest.mark.asyncio async def test_apply_to_output_streaming_gemini_sse_bytes_are_forwarded_incrementally_until_upstream_aborts(): - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, apply_to_output=True, mock_redacted_text={"text": ""}, @@ -2567,7 +2567,7 @@ async def test_apply_to_output_streaming_gemini_sse_bytes_are_forwarded_incremen @pytest.mark.asyncio async def test_apply_to_output_streaming_anthropic_first_frame_split_across_transport_chunks_is_still_masked(): - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, apply_to_output=True, mock_redacted_text={"text": ""}, @@ -2634,7 +2634,7 @@ def _anthropic_stream_head() -> list[bytes]: @pytest.mark.asyncio async def test_apply_to_output_streaming_anthropic_first_frame_split_inside_a_utf8_character_is_still_masked(): - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, apply_to_output=True, mock_redacted_text={"text": ""}, @@ -2677,7 +2677,7 @@ async def test_apply_to_output_streaming_anthropic_first_frame_split_inside_a_ut @pytest.mark.asyncio async def test_apply_to_output_streaming_anthropic_stream_led_by_sse_comment_keepalive_is_still_masked(): - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, apply_to_output=True, mock_redacted_text={"text": ""}, @@ -2714,7 +2714,7 @@ async def test_apply_to_output_streaming_anthropic_stream_led_by_sse_comment_kee @pytest.mark.asyncio async def test_apply_to_output_streaming_anthropic_stream_led_by_data_less_ping_event_is_still_masked(): - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, apply_to_output=True, mock_redacted_text={"text": ""}, @@ -2751,7 +2751,7 @@ async def test_apply_to_output_streaming_anthropic_stream_led_by_data_less_ping_ @pytest.mark.asyncio async def test_apply_to_output_streaming_leading_keepalive_is_forwarded_before_upstream_data_arrives(): - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, apply_to_output=True, mock_redacted_text={"text": ""}, @@ -2791,7 +2791,7 @@ async def test_apply_to_output_streaming_leading_keepalive_is_forwarded_before_u @pytest.mark.asyncio async def test_apply_to_output_streaming_leading_comments_over_the_frame_cap_are_still_masked(): - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, apply_to_output=True, mock_redacted_text={"text": ""}, @@ -2832,7 +2832,7 @@ async def test_apply_to_output_streaming_leading_comments_over_the_frame_cap_are @pytest.mark.asyncio async def test_apply_to_output_streaming_comment_only_stream_is_forwarded_unchanged(): - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, apply_to_output=True, mock_redacted_text={"text": ""}, @@ -2854,7 +2854,7 @@ async def test_apply_to_output_streaming_comment_only_stream_is_forwarded_unchan @pytest.mark.asyncio async def test_apply_to_output_streaming_gemini_first_frame_split_across_transport_chunks_streams_incrementally(): - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, apply_to_output=True, mock_redacted_text={"text": ""}, @@ -2885,7 +2885,7 @@ async def test_apply_to_output_streaming_gemini_first_frame_split_across_transpo @pytest.mark.asyncio async def test_apply_to_output_streaming_unterminated_first_frame_is_released_once_it_exceeds_the_cap(): - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, apply_to_output=True, mock_redacted_text={"text": ""}, @@ -2920,7 +2920,7 @@ async def test_apply_to_output_streaming_anthropic_sse_bytes_fail_closed_when_pr must surface as an error to the caller: replaying the unscanned frames would hand over whatever PII the model generated. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, apply_to_output=True, presidio_analyzer_api_base="http://127.0.0.1:9", @@ -2970,7 +2970,7 @@ async def test_apply_to_output_streaming_anthropic_sse_bytes_block_action_raises A BLOCK on generated PII must refuse the streaming /v1/messages response the same way it refuses the non streaming one, not replay the raw frames. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( apply_to_output=True, mock_testing=False, presidio_analyzer_api_base="http://test-analyzer/", @@ -3023,7 +3023,7 @@ async def test_apply_to_output_streaming_propagates_upstream_error_when_nothing_ An upstream guardrail that rejects the stream before the first chunk must surface as an error to the caller, not as an empty 200 stream. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, apply_to_output=True, mock_redacted_text={"text": ""}, @@ -3050,7 +3050,7 @@ async def test_output_parse_pii_streaming_responses_events_passthrough( Regression test: when output_parse_pii=True and pii_tokens exist, /v1/responses streaming events must pass through instead of being dropped. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, output_parse_pii=True, ) @@ -3099,7 +3099,7 @@ async def test_output_parse_pii_streaming_responses_completed_event_unmasked( ) from litellm.types.responses.main import GenericResponseOutputItem, OutputText - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, output_parse_pii=True, ) @@ -3159,7 +3159,7 @@ async def test_output_parse_pii_streaming_mixed_chunks_flushes_buffered( chunks must still be forwarded (in order) instead of being dropped at the saw_non_chat_chunk early return. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, output_parse_pii=True, ) @@ -3244,7 +3244,7 @@ async def test_anonymize_text_uses_correct_positions_no_parse_pii(): ], } - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( presidio_analyzer_api_base="http://test-analyzer/", presidio_anonymizer_api_base="http://test-anonymizer/", mock_testing=False, @@ -3318,7 +3318,7 @@ async def test_anonymize_text_uses_correct_positions_with_parse_pii(): ], } - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( presidio_analyzer_api_base="http://test-analyzer/", presidio_anonymizer_api_base="http://test-anonymizer/", mock_testing=False, @@ -3370,7 +3370,7 @@ def test_unmask_sse_bytes_chunk_replaces_text_delta(): } chunk = ("data: " + json.dumps(event) + "\n\n").encode("utf-8") - result = _OPTIONAL_PresidioPIIMasking._unmask_sse_bytes_chunk(chunk, pii_tokens) + result = OPTIONAL_PresidioPIIMasking._unmask_sse_bytes_chunk(chunk, pii_tokens) decoded = result.decode("utf-8") parsed = json.loads(decoded.split("data: ", 1)[1].strip()) @@ -3385,7 +3385,7 @@ def test_unmask_sse_bytes_chunk_ignores_non_text_delta(): # message_start event — no delta event = {"type": "message_start", "message": {"id": "msg_01", "role": "assistant"}} chunk = ("data: " + json.dumps(event) + "\n\n").encode("utf-8") - result = _OPTIONAL_PresidioPIIMasking._unmask_sse_bytes_chunk(chunk, pii_tokens) + result = OPTIONAL_PresidioPIIMasking._unmask_sse_bytes_chunk(chunk, pii_tokens) assert result == chunk # input_json_delta — should not be touched @@ -3395,19 +3395,19 @@ def test_unmask_sse_bytes_chunk_ignores_non_text_delta(): "delta": {"type": "input_json_delta", "partial_json": '{"name": ""}'}, } chunk2 = ("data: " + json.dumps(event2) + "\n\n").encode("utf-8") - result2 = _OPTIONAL_PresidioPIIMasking._unmask_sse_bytes_chunk(chunk2, pii_tokens) + result2 = OPTIONAL_PresidioPIIMasking._unmask_sse_bytes_chunk(chunk2, pii_tokens) assert result2 == chunk2 def test_unmask_sse_bytes_chunk_handles_malformed_json(): chunk = b"data: {not valid json}\n\n" - result = _OPTIONAL_PresidioPIIMasking._unmask_sse_bytes_chunk(chunk, {"": "Bobby"}) + result = OPTIONAL_PresidioPIIMasking._unmask_sse_bytes_chunk(chunk, {"": "Bobby"}) assert result == chunk def test_unmask_sse_bytes_chunk_handles_unicode_decode_error(): chunk = b"\xff\xfe invalid utf-8" - result = _OPTIONAL_PresidioPIIMasking._unmask_sse_bytes_chunk(chunk, {"": "Bobby"}) + result = OPTIONAL_PresidioPIIMasking._unmask_sse_bytes_chunk(chunk, {"": "Bobby"}) assert result == chunk @@ -3422,7 +3422,7 @@ def test_unmask_sse_bytes_chunk_non_ascii_pii_not_escaped(): } chunk = ("data: " + json.dumps(event) + "\n\n").encode("utf-8") - result = _OPTIONAL_PresidioPIIMasking._unmask_sse_bytes_chunk(chunk, pii_tokens) + result = OPTIONAL_PresidioPIIMasking._unmask_sse_bytes_chunk(chunk, pii_tokens) decoded = result.decode("utf-8") assert "Jos\\u" not in decoded @@ -3441,7 +3441,7 @@ def test_unmask_sse_bytes_chunk_handles_crlf_line_endings(): } crlf_chunk = ("data: " + json.dumps(event) + "\r\ndata: [DONE]\r\n").encode("utf-8") - result = _OPTIONAL_PresidioPIIMasking._unmask_sse_bytes_chunk(crlf_chunk, pii_tokens) + result = OPTIONAL_PresidioPIIMasking._unmask_sse_bytes_chunk(crlf_chunk, pii_tokens) decoded = result.decode("utf-8") parsed = json.loads(decoded.split("data: ", 1)[1].split("\n")[0].strip()) @@ -3453,7 +3453,7 @@ def test_unmask_sse_bytes_chunk_handles_crlf_line_endings(): async def test_stream_pii_unmasking_unmaskes_bytes_chunks(mock_user_api_key): import json - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, output_parse_pii=True, ) @@ -3489,7 +3489,7 @@ async def test_stream_pii_unmasking_unmaskes_bytes_chunks(mock_user_api_key): @pytest.mark.asyncio async def test_stream_pii_unmasking_passthrough_when_no_tokens(mock_user_api_key): - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, output_parse_pii=True, ) @@ -3514,7 +3514,7 @@ def test_new_entities_pass_through_analyze_payload(): """ import json - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, pii_entities_config={ PiiEntityType.DE_TAX_ID: PiiAction.MASK, @@ -3634,7 +3634,7 @@ def _make_marker_session_iterator( def _chunking_guardrail(chunk_size_bytes=100, **kwargs): - return _OPTIONAL_PresidioPIIMasking( + return OPTIONAL_PresidioPIIMasking( presidio_analyzer_api_base="http://test-analyzer/", presidio_anonymizer_api_base="http://test-anonymizer/", presidio_analyze_chunk_size_bytes=chunk_size_bytes, @@ -3651,7 +3651,7 @@ def _oversized_marker_text(): def test_split_text_for_analysis_offsets_and_byte_budget(): text = " ".join(f"word{i}" for i in range(200)) - chunks = _OPTIONAL_PresidioPIIMasking._split_text_for_analysis(text=text, chunk_size_bytes=100, overlap_chars=20) + chunks = OPTIONAL_PresidioPIIMasking._split_text_for_analysis(text=text, chunk_size_bytes=100, overlap_chars=20) assert len(chunks) > 1 for offset, chunk in chunks: assert len(chunk.encode("utf-8")) <= 100 @@ -3666,7 +3666,7 @@ def test_split_text_for_analysis_offsets_and_byte_budget(): def test_split_text_for_analysis_multibyte_characters(): text = "émoji🙂 çafé " * 120 - chunks = _OPTIONAL_PresidioPIIMasking._split_text_for_analysis(text=text, chunk_size_bytes=64, overlap_chars=8) + chunks = OPTIONAL_PresidioPIIMasking._split_text_for_analysis(text=text, chunk_size_bytes=64, overlap_chars=8) assert len(chunks) > 1 for offset, chunk in chunks: assert len(chunk.encode("utf-8")) <= 64 @@ -3676,7 +3676,7 @@ def test_split_text_for_analysis_multibyte_characters(): def test_split_text_for_analysis_under_budget_returns_single_chunk(): text = "short text" - chunks = _OPTIONAL_PresidioPIIMasking._split_text_for_analysis(text=text, chunk_size_bytes=100, overlap_chars=20) + chunks = OPTIONAL_PresidioPIIMasking._split_text_for_analysis(text=text, chunk_size_bytes=100, overlap_chars=20) assert chunks == [(0, text)] @@ -3805,18 +3805,18 @@ async def test_analyze_text_chunked_failure_stays_fail_closed(): def test_presidio_analyze_chunk_size_default_and_validation(): from litellm.constants import DEFAULT_PRESIDIO_ANALYZE_CHUNK_SIZE_BYTES - guardrail = _OPTIONAL_PresidioPIIMasking(mock_testing=True) + guardrail = OPTIONAL_PresidioPIIMasking(mock_testing=True) assert guardrail.presidio_analyze_chunk_size_bytes == DEFAULT_PRESIDIO_ANALYZE_CHUNK_SIZE_BYTES - nonpositive = _OPTIONAL_PresidioPIIMasking(mock_testing=True, presidio_analyze_chunk_size_bytes=-5) + nonpositive = OPTIONAL_PresidioPIIMasking(mock_testing=True, presidio_analyze_chunk_size_bytes=-5) assert nonpositive.presidio_analyze_chunk_size_bytes == DEFAULT_PRESIDIO_ANALYZE_CHUNK_SIZE_BYTES - custom = _OPTIONAL_PresidioPIIMasking(mock_testing=True, presidio_analyze_chunk_size_bytes=1234) + custom = OPTIONAL_PresidioPIIMasking(mock_testing=True, presidio_analyze_chunk_size_bytes=1234) assert custom.presidio_analyze_chunk_size_bytes == 1234 def test_update_in_memory_applies_analyze_chunk_size(): - guardrail = _OPTIONAL_PresidioPIIMasking(mock_testing=True) + guardrail = OPTIONAL_PresidioPIIMasking(mock_testing=True) params = LitellmParams( guardrail="presidio", mode="pre_call", @@ -3827,8 +3827,8 @@ def test_update_in_memory_applies_analyze_chunk_size(): def test_update_in_memory_keeps_output_masker_from_unmasking(): - masker = _OPTIONAL_PresidioPIIMasking(mock_testing=True, apply_to_output=True, output_parse_pii=False) - unmasker = _OPTIONAL_PresidioPIIMasking(mock_testing=True, output_parse_pii=True) + masker = OPTIONAL_PresidioPIIMasking(mock_testing=True, apply_to_output=True, output_parse_pii=False) + unmasker = OPTIONAL_PresidioPIIMasking(mock_testing=True, output_parse_pii=True) params = LitellmParams(guardrail="presidio", mode="pre_call", output_parse_pii=True) masker.update_in_memory_litellm_params(params) @@ -3844,7 +3844,7 @@ def test_merge_drops_truncated_same_type_fragment_from_overlap(): numbered-token rewriter and double-counts entities.""" truncated = {"entity_type": "IP_ADDRESS", "start": 10, "end": 21, "score": 0.6} full_local = {"entity_type": "IP_ADDRESS", "start": 5, "end": 18, "score": 0.95} - merged = _OPTIONAL_PresidioPIIMasking._merge_chunked_analyze_results( + merged = OPTIONAL_PresidioPIIMasking._merge_chunked_analyze_results( text_chunks=[(0, "x" * 21), (5, "x" * 25)], chunk_results=[[truncated], [full_local]], ) @@ -3856,7 +3856,7 @@ def test_merge_drops_truncated_same_type_fragment_from_overlap(): def test_merge_exact_duplicate_keeps_higher_score(): low = {"entity_type": "EMAIL_ADDRESS", "start": 3, "end": 9, "score": 0.4} high = {"entity_type": "EMAIL_ADDRESS", "start": 0, "end": 6, "score": 0.9} - merged = _OPTIONAL_PresidioPIIMasking._merge_chunked_analyze_results( + merged = OPTIONAL_PresidioPIIMasking._merge_chunked_analyze_results( text_chunks=[(0, "x" * 9), (3, "x" * 9)], chunk_results=[[low], [high]], ) @@ -3869,7 +3869,7 @@ def test_merge_preserves_cross_type_overlap(): (e.g. URL inside EMAIL_ADDRESS); the chunk merge must not drop those.""" email = {"entity_type": "EMAIL_ADDRESS", "start": 0, "end": 20, "score": 1.0} url = {"entity_type": "URL", "start": 5, "end": 20, "score": 0.5} - merged = _OPTIONAL_PresidioPIIMasking._merge_chunked_analyze_results( + merged = OPTIONAL_PresidioPIIMasking._merge_chunked_analyze_results( text_chunks=[(0, "x" * 25)], chunk_results=[[email, url]], ) @@ -3879,7 +3879,7 @@ def test_merge_preserves_cross_type_overlap(): def test_update_in_memory_coerces_invalid_chunk_size(): from litellm.constants import DEFAULT_PRESIDIO_ANALYZE_CHUNK_SIZE_BYTES - guardrail = _OPTIONAL_PresidioPIIMasking(mock_testing=True, presidio_analyze_chunk_size_bytes=99_000) + guardrail = OPTIONAL_PresidioPIIMasking(mock_testing=True, presidio_analyze_chunk_size_bytes=99_000) params = LitellmParams( guardrail="presidio", mode="pre_call", @@ -3890,7 +3890,7 @@ def test_update_in_memory_coerces_invalid_chunk_size(): def test_split_text_handles_chunk_size_below_char_width(): - chunks = _OPTIONAL_PresidioPIIMasking._split_text_for_analysis( + chunks = OPTIONAL_PresidioPIIMasking._split_text_for_analysis( text="\U0001f642\U0001f642", chunk_size_bytes=3, overlap_chars=8 ) assert all(chunk for _, chunk in chunks) @@ -3967,7 +3967,7 @@ def test_split_text_accounts_for_json_body_expansion(): text = "これは個人情報テストです。" * 200 # 3-byte UTF-8 chars, 6-byte escapes budget = 1000 - chunks = _OPTIONAL_PresidioPIIMasking._split_text_for_analysis(text=text, chunk_size_bytes=budget, overlap_chars=8) + chunks = OPTIONAL_PresidioPIIMasking._split_text_for_analysis(text=text, chunk_size_bytes=budget, overlap_chars=8) assert len(chunks) > 1 for offset, chunk in chunks: assert len(json_module.dumps(chunk).encode("utf-8")) - 2 <= budget @@ -4158,7 +4158,7 @@ async def test_pii_masking_replays_a_byte_identical_prefix_across_turns(mock_use the history lose their binding. The analyzer and anonymizer are an in-process fake handed to the guardrail through its api_base settings.""" async with TestServer(_fake_presidio_app()) as server: - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( presidio_analyzer_api_base=str(server.make_url("/")), presidio_anonymizer_api_base=str(server.make_url("/")), pii_entities_config={PiiEntityType.PERSON: PiiAction.MASK}, @@ -4279,7 +4279,7 @@ async def test_restoration_never_contacts_presidio(has_tokens: bool) -> None: async def test_standalone_restoration_preserves_post_call_selection(event_hook: str | list[str]) -> None: from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import UnifiedLLMGuardrails - callback: Final = _OPTIONAL_PresidioPIIMasking( + callback: Final = OPTIONAL_PresidioPIIMasking( event_hook=event_hook, default_on=True, output_parse_pii=True, @@ -4302,7 +4302,7 @@ async def test_standalone_restoration_preserves_post_call_selection(event_hook: ], ) def test_validate_environment_missing_http(base_url): - pii_masking = _OPTIONAL_PresidioPIIMasking(mock_testing=True) + pii_masking = OPTIONAL_PresidioPIIMasking(mock_testing=True) env_vars = { "PRESIDIO_ANALYZER_API_BASE": f"{base_url}/analyze", @@ -4333,7 +4333,7 @@ async def test_output_parsing(): """ litellm.set_verbose = True litellm.output_parse_pii = True - pii_masking = _OPTIONAL_PresidioPIIMasking(mock_testing=True) + pii_masking = OPTIONAL_PresidioPIIMasking(mock_testing=True) initial_message = [ { @@ -4413,7 +4413,7 @@ async def test_presidio_pii_masking_input_a(): """ Tests to see if correct parts of sentence anonymized """ - pii_masking = _OPTIONAL_PresidioPIIMasking( + pii_masking = OPTIONAL_PresidioPIIMasking( mock_testing=True, mock_redacted_text=input_a_anonymizer_results ) @@ -4444,7 +4444,7 @@ async def test_presidio_pii_masking_input_b(): """ Tests to see if correct parts of sentence anonymized """ - pii_masking = _OPTIONAL_PresidioPIIMasking( + pii_masking = OPTIONAL_PresidioPIIMasking( mock_testing=True, mock_redacted_text=input_b_anonymizer_results ) @@ -4474,7 +4474,7 @@ async def test_presidio_pii_masking_input_b(): async def test_presidio_pii_masking_logging_output_only_no_pre_api_hook(): from litellm.types.guardrails import GuardrailEventHooks - pii_masking = _OPTIONAL_PresidioPIIMasking( + pii_masking = OPTIONAL_PresidioPIIMasking( logging_only=True, mock_testing=True, mock_redacted_text=input_b_anonymizer_results, @@ -4505,7 +4505,7 @@ async def test_presidio_language_configuration(): """Test that presidio_language parameter is properly set and used in analyze requests""" litellm.turn_on_debug() - presidio_guardrail_de = _OPTIONAL_PresidioPIIMasking( + presidio_guardrail_de = OPTIONAL_PresidioPIIMasking( pii_entities_config={}, presidio_language="de", mock_testing=True, @@ -4520,7 +4520,7 @@ async def test_presidio_language_configuration(): assert analyze_request["language"] == "de" assert analyze_request["text"] == test_text - presidio_guardrail_es = _OPTIONAL_PresidioPIIMasking( + presidio_guardrail_es = OPTIONAL_PresidioPIIMasking( pii_entities_config={}, presidio_language="es", mock_testing=True ) @@ -4533,7 +4533,7 @@ async def test_presidio_language_configuration(): assert analyze_request_es["language"] == "es" assert analyze_request_es["text"] == test_text_es - presidio_guardrail_default = _OPTIONAL_PresidioPIIMasking( + presidio_guardrail_default = OPTIONAL_PresidioPIIMasking( pii_entities_config={}, mock_testing=True ) @@ -4554,7 +4554,7 @@ async def test_presidio_language_configuration_with_per_request_override(): """Test that per-request language configuration overrides the default configured language""" litellm.turn_on_debug() - presidio_guardrail = _OPTIONAL_PresidioPIIMasking( + presidio_guardrail = OPTIONAL_PresidioPIIMasking( pii_entities_config={}, presidio_language="de", mock_testing=True ) diff --git a/tests/unit/proxy/guardrails/test_deferred_guardrail_logging.py b/tests/unit/proxy/guardrails/test_deferred_guardrail_logging.py index 35137e4693e..17427e839b9 100644 --- a/tests/unit/proxy/guardrails/test_deferred_guardrail_logging.py +++ b/tests/unit/proxy/guardrails/test_deferred_guardrail_logging.py @@ -1114,7 +1114,7 @@ class TestDeferredStreamingClosure: with ( patch("litellm.callbacks", [guardrail]), patch( - "litellm.proxy.utils._check_and_merge_model_level_guardrails", + "litellm.proxy.utils.check_and_merge_model_level_guardrails", side_effect=mock_merge, ), ): @@ -1178,7 +1178,7 @@ class TestDeferredStreamingClosure: with ( patch("litellm.callbacks", [guardrail_a, guardrail_b]), patch( - "litellm.proxy.utils._check_and_merge_model_level_guardrails", + "litellm.proxy.utils.check_and_merge_model_level_guardrails", side_effect=mock_merge, ), ): @@ -1216,7 +1216,7 @@ class TestDeferredStreamingClosure: raise RuntimeError("Simulated init failure") with patch( - "litellm.proxy.utils._check_and_merge_model_level_guardrails", + "litellm.proxy.utils.check_and_merge_model_level_guardrails", side_effect=exploding_merge, ): await ProxyBaseLLMRequestProcessing._run_deferred_stream_guardrails( diff --git a/tests/unit/proxy/guardrails/test_guardrail_coverage.py b/tests/unit/proxy/guardrails/test_guardrail_coverage.py index 474ca9d8fa0..28e70ec84fa 100644 --- a/tests/unit/proxy/guardrails/test_guardrail_coverage.py +++ b/tests/unit/proxy/guardrails/test_guardrail_coverage.py @@ -605,9 +605,9 @@ async def test_azure_content_safety_pre_call_fires_on_runtime_call_types( ``aresponses`` for the Responses API. The hook must inspect text fragments under both, not only the literal ``"completion"`` string used by some SDK callers.""" - from litellm.proxy.hooks.azure_content_safety import _PROXY_AzureContentSafety + from litellm.proxy.hooks.azure_content_safety import PROXY_AzureContentSafety - guard = _PROXY_AzureContentSafety.__new__(_PROXY_AzureContentSafety) + guard = PROXY_AzureContentSafety.__new__(PROXY_AzureContentSafety) seen = [] async def fake_test_violation(content, source=None): @@ -628,9 +628,9 @@ async def test_azure_content_safety_post_call_checks_all_choices(user_api_key): """Krrish blocker: ``n>1`` responses must not bypass Azure Content Safety by placing the unsafe text in ``choices[1+]``.""" from fastapi import HTTPException - from litellm.proxy.hooks.azure_content_safety import _PROXY_AzureContentSafety + from litellm.proxy.hooks.azure_content_safety import PROXY_AzureContentSafety - guard = _PROXY_AzureContentSafety.__new__(_PROXY_AzureContentSafety) + guard = PROXY_AzureContentSafety.__new__(PROXY_AzureContentSafety) seen_outputs = [] async def fake_test_violation(content, source=None): diff --git a/tests/unit/proxy/guardrails/test_init_guardrails.py b/tests/unit/proxy/guardrails/test_init_guardrails.py index 2cac5d60068..c8eeaf11fde 100644 --- a/tests/unit/proxy/guardrails/test_init_guardrails.py +++ b/tests/unit/proxy/guardrails/test_init_guardrails.py @@ -209,7 +209,7 @@ def test_initialize_presidio_forwards_analyze_chunk_size_bytes(): """ import litellm from litellm.proxy.guardrails.guardrail_hooks.presidio import ( - _OPTIONAL_PresidioPIIMasking, + OPTIONAL_PresidioPIIMasking, ) test_guardrail = { @@ -229,7 +229,7 @@ def test_initialize_presidio_forwards_analyze_chunk_size_bytes(): initialized = [ callback for callback in litellm.callbacks - if isinstance(callback, _OPTIONAL_PresidioPIIMasking) and callback.guardrail_name == "test_presidio_chunk_size" + if isinstance(callback, OPTIONAL_PresidioPIIMasking) and callback.guardrail_name == "test_presidio_chunk_size" ] assert initialized, "presidio guardrail was not registered as a callback" assert initialized[-1].presidio_analyze_chunk_size_bytes == 250_000 diff --git a/tests/unit/proxy/health_endpoints/test_health_endpoints.py b/tests/unit/proxy/health_endpoints/test_health_endpoints.py index 95aad7b772f..266fd06333c 100644 --- a/tests/unit/proxy/health_endpoints/test_health_endpoints.py +++ b/tests/unit/proxy/health_endpoints/test_health_endpoints.py @@ -426,7 +426,7 @@ async def test_test_model_connection_loads_config_from_router(): mock_run_with_timeout, ), patch( - "litellm.proxy.health_endpoints._health_endpoints._update_litellm_params_for_health_check", + "litellm.proxy.health_endpoints._health_endpoints.update_litellm_params_for_health_check", mock_update_params, ), patch( @@ -575,7 +575,7 @@ async def test_test_model_connection_uses_model_info_id_to_disambiguate_duplicat mock_run_with_timeout, ), patch( - "litellm.proxy.health_endpoints._health_endpoints._update_litellm_params_for_health_check", + "litellm.proxy.health_endpoints._health_endpoints.update_litellm_params_for_health_check", mock_update_params, ), patch( @@ -677,7 +677,7 @@ async def test_test_model_connection_falls_back_to_deployments_zero_without_id() mock_run_with_timeout, ), patch( - "litellm.proxy.health_endpoints._health_endpoints._update_litellm_params_for_health_check", + "litellm.proxy.health_endpoints._health_endpoints.update_litellm_params_for_health_check", mock_update_params, ), patch( @@ -2566,13 +2566,13 @@ def test_no_federation_field_reaches_a_non_admin_health_entry(federation_field: deployment is healthy must learn neither. Both lists that enforce that are derived from the same key sets this runs over, so a field added to the funnel without joining either one shows up here as a value a non-admin could read.""" - from litellm.proxy.health_check import _clean_endpoint_data + from litellm.proxy.health_check import clean_endpoint_data from litellm.proxy.health_endpoints._health_endpoints import ( _strip_admin_only_fields_from_health_result, ) canary = f"CANARY-{federation_field}-VALUE" - cleaned = _clean_endpoint_data( + cleaned = clean_endpoint_data( {"model": "anthropic/claude-sonnet-5", federation_field: canary}, details=True, ) @@ -2591,10 +2591,10 @@ def test_no_federation_secret_reaches_even_an_admin_health_entry(secret_field: s token, key, or reference it federates with, so these fields drop at the health-check layer ahead of any per-caller stripping. Reading the same set the drop list is built from is what catches a new secret-bearing field that was only ever added to the admin-gated half.""" - from litellm.proxy.health_check import _clean_endpoint_data + from litellm.proxy.health_check import clean_endpoint_data canary = f"CANARY-{secret_field}-VALUE" - cleaned = _clean_endpoint_data( + cleaned = clean_endpoint_data( {"model": "anthropic/claude-sonnet-5", secret_field: canary}, details=True, ) @@ -3364,7 +3364,7 @@ def test_clean_endpoint_data_strips_credentials_keeps_routing_fields(): layer based on user role, not in the cleaning helper. This guarantees proxy admins continue to see those fields in the /health response. """ - from litellm.proxy.health_check import _clean_endpoint_data + from litellm.proxy.health_check import clean_endpoint_data raw = { "model": "openai/gpt-4o", @@ -3374,7 +3374,7 @@ def test_clean_endpoint_data_strips_credentials_keeps_routing_fields(): "aws_access_key_id": "AKIAEXAMPLE", } - cleaned = _clean_endpoint_data(raw, details=True) + cleaned = clean_endpoint_data(raw, details=True) assert "api_key" not in cleaned assert "aws_access_key_id" not in cleaned @@ -3388,7 +3388,7 @@ def test_clean_endpoint_data_strips_extra_headers_and_aws_session_token(): `extra_headers` / `headers` / `aws_session_token`. Before the fix these were returned in plaintext (api_key was stripped, but these were not). """ - from litellm.proxy.health_check import _clean_endpoint_data + from litellm.proxy.health_check import clean_endpoint_data raw = { "model": "openai/gpt-4o", @@ -3402,7 +3402,7 @@ def test_clean_endpoint_data_strips_extra_headers_and_aws_session_token(): "aws_session_token": "CANARY_AWS_SESSION_TOKEN_VALUE", } - cleaned = _clean_endpoint_data(raw, details=True) + cleaned = clean_endpoint_data(raw, details=True) assert "extra_headers" not in cleaned assert "headers" not in cleaned @@ -3439,10 +3439,10 @@ def test_clean_endpoint_data_never_displays_credential_fields(credential_field, LIT-6239 / gh-36898: /health entries, healthy and unhealthy alike, must never carry credential-bearing litellm_params, with or without details. """ - from litellm.proxy.health_check import _clean_endpoint_data + from litellm.proxy.health_check import clean_endpoint_data canary = f"CANARY-{credential_field}-VALUE" - cleaned = _clean_endpoint_data( + cleaned = clean_endpoint_data( { "model": "azure/gpt-5-mini", "api_base": "https://example.test/v1", @@ -4063,9 +4063,9 @@ def test_clean_endpoint_data_keeps_only_json_safe_diagnostics(): """ from fastapi.encoders import jsonable_encoder - from litellm.proxy.health_check import _clean_endpoint_data + from litellm.proxy.health_check import clean_endpoint_data - cleaned = _clean_endpoint_data( + cleaned = clean_endpoint_data( { "model": "bedrock/us.amazon.nova-2-lite-v1:0", "custom_llm_provider": "bedrock", diff --git a/tests/unit/proxy/hooks/test_batch_file_validation.py b/tests/unit/proxy/hooks/test_batch_file_validation.py index 38ee6997899..8eeee2e6468 100644 --- a/tests/unit/proxy/hooks/test_batch_file_validation.py +++ b/tests/unit/proxy/hooks/test_batch_file_validation.py @@ -151,9 +151,9 @@ async def test_pre_call_rejects_unauthorized_model_in_batch_file(): """Pre-fix the hook only validated the outer `model` parameter and forwarded the file as-is. With this fix, a model named inside the JSONL that the caller cannot use must trigger a 403.""" - from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + from litellm.proxy.hooks.batch_rate_limiter import PROXY_BatchRateLimiter - rate_limiter = _PROXY_BatchRateLimiter( + rate_limiter = PROXY_BatchRateLimiter( internal_usage_cache=MagicMock(), parallel_request_limiter=MagicMock(), ) @@ -199,9 +199,9 @@ async def test_pre_call_allows_all_team_models_key_when_model_in_team_allowlist( """Keys with ``all-team-models`` must inherit the team allowlist when validating models embedded in batch JSONL.""" from litellm.proxy._types import SpecialModelNames - from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + from litellm.proxy.hooks.batch_rate_limiter import PROXY_BatchRateLimiter - rate_limiter = _PROXY_BatchRateLimiter( + rate_limiter = PROXY_BatchRateLimiter( internal_usage_cache=MagicMock(), parallel_request_limiter=MagicMock(), ) @@ -233,9 +233,9 @@ async def test_pre_call_allows_all_team_models_key_when_model_in_team_allowlist( @pytest.mark.asyncio async def test_pre_call_uses_current_team_allowlist_for_all_team_models_key(): from litellm.proxy._types import LiteLLM_TeamTable, SpecialModelNames - from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + from litellm.proxy.hooks.batch_rate_limiter import PROXY_BatchRateLimiter - rate_limiter = _PROXY_BatchRateLimiter( + rate_limiter = PROXY_BatchRateLimiter( internal_usage_cache=MagicMock(), parallel_request_limiter=MagicMock(), ) @@ -287,9 +287,9 @@ async def test_pre_call_allows_all_team_models_key_via_current_team_object(): allowlist must be authorized through the freshly-fetched team object, not the cached-``team_models`` fallback.""" from litellm.proxy._types import LiteLLM_TeamTable, SpecialModelNames - from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + from litellm.proxy.hooks.batch_rate_limiter import PROXY_BatchRateLimiter - rate_limiter = _PROXY_BatchRateLimiter( + rate_limiter = PROXY_BatchRateLimiter( internal_usage_cache=MagicMock(), parallel_request_limiter=MagicMock(), ) @@ -352,9 +352,9 @@ async def test_pre_call_denies_all_team_models_key_via_member_scope(): LiteLLM_TeamTable, SpecialModelNames, ) - from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + from litellm.proxy.hooks.batch_rate_limiter import PROXY_BatchRateLimiter - rate_limiter = _PROXY_BatchRateLimiter( + rate_limiter = PROXY_BatchRateLimiter( internal_usage_cache=MagicMock(), parallel_request_limiter=MagicMock(), ) @@ -416,9 +416,9 @@ async def test_pre_call_fails_closed_when_current_team_fetch_fails_for_all_team_ team_fetch_error, expected_status ): from litellm.proxy._types import SpecialModelNames - from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + from litellm.proxy.hooks.batch_rate_limiter import PROXY_BatchRateLimiter - rate_limiter = _PROXY_BatchRateLimiter( + rate_limiter = PROXY_BatchRateLimiter( internal_usage_cache=MagicMock(), parallel_request_limiter=MagicMock(), ) @@ -470,9 +470,9 @@ async def test_pre_call_allows_teamless_all_team_models_key(): someone re-introduces a teamless denial in _resolve_key_models_for_auth_check or adds a team_id guard that blocks the batch path.""" from litellm.proxy._types import SpecialModelNames - from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + from litellm.proxy.hooks.batch_rate_limiter import PROXY_BatchRateLimiter - rate_limiter = _PROXY_BatchRateLimiter( + rate_limiter = PROXY_BatchRateLimiter( internal_usage_cache=MagicMock(), parallel_request_limiter=MagicMock(), ) @@ -503,9 +503,9 @@ async def test_pre_call_allows_teamless_all_team_models_key(): async def test_pre_call_allows_authorized_model_in_batch_file(): """If every model in the JSONL is on the caller's allowlist, the hook must not raise.""" - from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + from litellm.proxy.hooks.batch_rate_limiter import PROXY_BatchRateLimiter - rate_limiter = _PROXY_BatchRateLimiter( + rate_limiter = PROXY_BatchRateLimiter( internal_usage_cache=MagicMock(), parallel_request_limiter=MagicMock(), ) @@ -542,9 +542,9 @@ async def test_pre_call_allows_authorized_model_in_batch_file(): @pytest.mark.asyncio async def test_pre_call_skips_file_fetch_when_disabled_in_general_settings(): - from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + from litellm.proxy.hooks.batch_rate_limiter import PROXY_BatchRateLimiter - rate_limiter = _PROXY_BatchRateLimiter( + rate_limiter = PROXY_BatchRateLimiter( internal_usage_cache=MagicMock(), parallel_request_limiter=MagicMock(), ) @@ -562,14 +562,14 @@ async def test_pre_call_skips_file_fetch_when_disabled_in_general_settings(): ) assert result == {"input_file_id": "file-abc123"} - rate_limiter.parallel_request_limiter._create_rate_limit_descriptors.assert_not_called() + rate_limiter.parallel_request_limiter.create_rate_limit_descriptors.assert_not_called() @pytest.mark.asyncio async def test_pre_call_skips_file_fetch_for_configured_provider(): - from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + from litellm.proxy.hooks.batch_rate_limiter import PROXY_BatchRateLimiter - rate_limiter = _PROXY_BatchRateLimiter( + rate_limiter = PROXY_BatchRateLimiter( internal_usage_cache=MagicMock(), parallel_request_limiter=MagicMock(), ) @@ -600,7 +600,7 @@ async def test_pre_call_skips_file_fetch_for_configured_provider(): # work — assert the skip happened rather than the hook's error-recovery # path (which also returns data unchanged). mock_afile_content.assert_not_awaited() - rate_limiter.parallel_request_limiter._create_rate_limit_descriptors.assert_not_called() + rate_limiter.parallel_request_limiter.create_rate_limit_descriptors.assert_not_called() @pytest.mark.asyncio @@ -609,9 +609,9 @@ async def test_pre_call_does_not_skip_for_spoofed_provider(): user-supplied ``custom_llm_provider`` that is not backed by the routing deployment must not trigger a skip: the input file must still be fetched and the rate-limit counters incremented.""" - from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + from litellm.proxy.hooks.batch_rate_limiter import PROXY_BatchRateLimiter - rate_limiter = _PROXY_BatchRateLimiter( + rate_limiter = PROXY_BatchRateLimiter( internal_usage_cache=MagicMock(), parallel_request_limiter=MagicMock(), ) @@ -619,7 +619,7 @@ async def test_pre_call_does_not_skip_for_spoofed_provider(): # only thing that could prevent the fetch below is the provider skip. If the # spoofed ``custom_llm_provider`` were honored, afile_content would never be # awaited. - rate_limiter.parallel_request_limiter._create_rate_limit_descriptors.return_value = [ + rate_limiter.parallel_request_limiter.create_rate_limit_descriptors.return_value = [ {"rate_limit": {"requests_per_unit": 100}} ] rate_limiter.parallel_request_limiter.atomic_check_and_increment_by_n = AsyncMock( @@ -672,7 +672,7 @@ async def test_pre_call_does_not_skip_for_spoofed_provider(): async def test_count_input_file_usage_decodes_model_embedded_file_id(): import base64 - from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + from litellm.proxy.hooks.batch_rate_limiter import PROXY_BatchRateLimiter original_file_id = "file-provider-xyz" encoded_payload = ( @@ -684,7 +684,7 @@ async def test_count_input_file_usage_decodes_model_embedded_file_id(): ) encoded_file_id = f"file-{encoded_payload}" - rate_limiter = _PROXY_BatchRateLimiter( + rate_limiter = PROXY_BatchRateLimiter( internal_usage_cache=MagicMock(), parallel_request_limiter=MagicMock(), ) @@ -727,9 +727,9 @@ async def test_pre_call_allows_stripped_provider_model_when_key_has_proxy_alias( """After replace_model_in_jsonl, body.model is the provider id (e.g. gpt-5.5). Auth must check target_model_names from the unified file id, not reverse-map the stripped id.""" - from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + from litellm.proxy.hooks.batch_rate_limiter import PROXY_BatchRateLimiter - rate_limiter = _PROXY_BatchRateLimiter( + rate_limiter = PROXY_BatchRateLimiter( internal_usage_cache=MagicMock(), parallel_request_limiter=MagicMock(), ) @@ -791,9 +791,9 @@ async def test_pre_call_uses_target_model_names_not_stripped_reverse_lookup( ): """LIT-3593: three deployments strip to gpt-5.5; auth must use the upload target alias from target_model_names, not first-match reverse lookup.""" - from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + from litellm.proxy.hooks.batch_rate_limiter import PROXY_BatchRateLimiter - rate_limiter = _PROXY_BatchRateLimiter( + rate_limiter = PROXY_BatchRateLimiter( internal_usage_cache=MagicMock(), parallel_request_limiter=MagicMock(), ) @@ -862,9 +862,9 @@ async def test_pre_call_uses_target_model_names_not_stripped_reverse_lookup( async def test_pre_call_skips_check_when_no_models_present(): """Files without any `body.model` (corrupt or empty) must not 500; the rate limiter logs a warning elsewhere and proceeds.""" - from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + from litellm.proxy.hooks.batch_rate_limiter import PROXY_BatchRateLimiter - rate_limiter = _PROXY_BatchRateLimiter( + rate_limiter = PROXY_BatchRateLimiter( internal_usage_cache=MagicMock(), parallel_request_limiter=MagicMock(), ) @@ -889,9 +889,9 @@ async def test_pre_call_skips_check_when_no_models_present(): def _make_rate_limiter(): - from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + from litellm.proxy.hooks.batch_rate_limiter import PROXY_BatchRateLimiter - return _PROXY_BatchRateLimiter( + return PROXY_BatchRateLimiter( internal_usage_cache=MagicMock(), parallel_request_limiter=MagicMock(), ) @@ -954,7 +954,7 @@ def test_get_batch_routing_model_uses_unified_file_id_target(): return_value=None, ), patch( - "litellm.proxy.openai_files_endpoints.common_utils._is_base64_encoded_unified_file_id", + "litellm.proxy.openai_files_endpoints.common_utils.is_base64_encoded_unified_file_id", return_value="unified-id", ), patch( @@ -969,9 +969,9 @@ def test_get_batch_routing_model_uses_unified_file_id_target(): def test_key_requires_batch_model_access_check_branches(): - from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + from litellm.proxy.hooks.batch_rate_limiter import PROXY_BatchRateLimiter - check = _PROXY_BatchRateLimiter._key_requires_batch_model_access_check + check = PROXY_BatchRateLimiter._key_requires_batch_model_access_check assert check(UserAPIKeyAuth(api_key="sk", models=["*"])) is False assert check(UserAPIKeyAuth(api_key="sk", models=["all-proxy-models"])) is False assert ( @@ -1007,9 +1007,9 @@ def test_key_requires_batch_model_access_check_branches(): def test_has_applicable_batch_rate_limits(): - from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + from litellm.proxy.hooks.batch_rate_limiter import PROXY_BatchRateLimiter - has_limits = _PROXY_BatchRateLimiter._has_applicable_batch_rate_limits + has_limits = PROXY_BatchRateLimiter._has_applicable_batch_rate_limits assert has_limits([{"rate_limit": {"tokens_per_unit": 100}}]) is True assert has_limits([{"rate_limit": {"requests_per_unit": 5}}]) is True assert has_limits([{"rate_limit": {"max_parallel_requests": 2}}]) is True @@ -1032,7 +1032,7 @@ def test_should_skip_ignores_client_supplied_metadata_flag(): body. The skip decision is server-controlled only, so with applicable rate limits the JSONL is still processed despite the client flag.""" rate_limiter = _make_rate_limiter() - rate_limiter.parallel_request_limiter._create_rate_limit_descriptors.return_value = [ + rate_limiter.parallel_request_limiter.create_rate_limit_descriptors.return_value = [ {"rate_limit": {"requests_per_unit": 5}} ] user = UserAPIKeyAuth(api_key="sk", models=["*"]) @@ -1059,7 +1059,7 @@ def test_should_not_skip_for_forged_model_embedded_file_id(): import base64 rate_limiter = _make_rate_limiter() - rate_limiter.parallel_request_limiter._create_rate_limit_descriptors.return_value = [ + rate_limiter.parallel_request_limiter.create_rate_limit_descriptors.return_value = [ {"rate_limit": {"requests_per_unit": 5}} ] user = UserAPIKeyAuth(api_key="sk", models=["*"]) @@ -1088,7 +1088,7 @@ def test_should_not_skip_for_skip_listed_top_level_model(): ``body.model`` entries. No per-model skip exists, so a skip-listed model over a plain file still gets processed.""" rate_limiter = _make_rate_limiter() - rate_limiter.parallel_request_limiter._create_rate_limit_descriptors.return_value = [ + rate_limiter.parallel_request_limiter.create_rate_limit_descriptors.return_value = [ {"rate_limit": {"requests_per_unit": 5}} ] user = UserAPIKeyAuth(api_key="sk", models=["*"]) @@ -1114,7 +1114,7 @@ def test_should_not_skip_when_file_bound_provider_is_rate_limited(): import base64 rate_limiter = _make_rate_limiter() - rate_limiter.parallel_request_limiter._create_rate_limit_descriptors.return_value = [ + rate_limiter.parallel_request_limiter.create_rate_limit_descriptors.return_value = [ {"rate_limit": {"requests_per_unit": 5}} ] user = UserAPIKeyAuth(api_key="sk", models=["*"]) @@ -1156,7 +1156,7 @@ def test_should_skip_when_file_bound_provider_is_skip_listed(): import base64 rate_limiter = _make_rate_limiter() - rate_limiter.parallel_request_limiter._create_rate_limit_descriptors.return_value = [ + rate_limiter.parallel_request_limiter.create_rate_limit_descriptors.return_value = [ {"rate_limit": {"requests_per_unit": 5}} ] user = UserAPIKeyAuth(api_key="sk", models=["*"]) @@ -1194,7 +1194,7 @@ def test_warns_once_for_unsupported_model_skip_setting(): """Operators who set the no-op per-model skip key get a single warning so a misconfigured deployment does not silently leave batch limits unenforced.""" rate_limiter = _make_rate_limiter() - rate_limiter.parallel_request_limiter._create_rate_limit_descriptors.return_value = [ + rate_limiter.parallel_request_limiter.create_rate_limit_descriptors.return_value = [ {"rate_limit": {"requests_per_unit": 5}} ] user = UserAPIKeyAuth(api_key="sk", models=["*"]) @@ -1221,7 +1221,7 @@ def test_warns_once_for_unsupported_model_skip_setting(): def test_no_warning_when_model_skip_setting_absent(): rate_limiter = _make_rate_limiter() - rate_limiter.parallel_request_limiter._create_rate_limit_descriptors.return_value = [ + rate_limiter.parallel_request_limiter.create_rate_limit_descriptors.return_value = [ {"rate_limit": {"requests_per_unit": 5}} ] user = UserAPIKeyAuth(api_key="sk", models=["*"]) @@ -1243,7 +1243,7 @@ def test_no_warning_when_model_skip_setting_absent(): def test_should_skip_when_no_rate_limits_configured(): rate_limiter = _make_rate_limiter() - rate_limiter.parallel_request_limiter._create_rate_limit_descriptors.return_value = [ + rate_limiter.parallel_request_limiter.create_rate_limit_descriptors.return_value = [ {"rate_limit": {}} ] user = UserAPIKeyAuth(api_key="sk", models=["*"]) @@ -1261,7 +1261,7 @@ def test_should_skip_when_no_rate_limits_configured(): def test_should_not_skip_and_reuses_descriptors_when_limits_present(): rate_limiter = _make_rate_limiter() descriptors = [{"rate_limit": {"tokens_per_unit": 100}}] - rate_limiter.parallel_request_limiter._create_rate_limit_descriptors.return_value = ( + rate_limiter.parallel_request_limiter.create_rate_limit_descriptors.return_value = ( descriptors ) user = UserAPIKeyAuth(api_key="sk", models=["*"]) @@ -1357,17 +1357,17 @@ def test_resolve_fetch_params_model_embedded_fails_open_on_credential_error(): async def test_check_and_increment_computes_descriptors_when_not_passed(): from litellm.proxy.hooks.batch_rate_limiter import ( BatchFileUsage, - _PROXY_BatchRateLimiter, + PROXY_BatchRateLimiter, ) parallel_request_limiter = MagicMock() - parallel_request_limiter._create_rate_limit_descriptors.return_value = [ + parallel_request_limiter.create_rate_limit_descriptors.return_value = [ {"rate_limit": {"tokens_per_unit": 100}} ] parallel_request_limiter.atomic_check_and_increment_by_n = AsyncMock( return_value={"overall_code": "OK", "statuses": []} ) - rate_limiter = _PROXY_BatchRateLimiter( + rate_limiter = PROXY_BatchRateLimiter( internal_usage_cache=MagicMock(), parallel_request_limiter=parallel_request_limiter, ) @@ -1379,7 +1379,7 @@ async def test_check_and_increment_computes_descriptors_when_not_passed(): descriptors=None, ) - parallel_request_limiter._create_rate_limit_descriptors.assert_called_once() + parallel_request_limiter.create_rate_limit_descriptors.assert_called_once() @pytest.mark.asyncio @@ -1390,17 +1390,17 @@ async def test_pre_call_enforces_project_otpm_limit_for_batch(): quota. The project OTPM descriptor must now be present and charged with the batch's estimated *output* tokens, not its input tokens.""" from litellm import DualCache - from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + from litellm.proxy.hooks.batch_rate_limiter import PROXY_BatchRateLimiter from litellm.proxy.hooks.parallel_request_limiter_v3 import ( - _PROXY_MaxParallelRequestsHandler_v3, + PROXY_MaxParallelRequestsHandler_v3, ) from litellm.proxy.utils import InternalUsageCache local_cache = DualCache() - parallel_request_limiter = _PROXY_MaxParallelRequestsHandler_v3( + parallel_request_limiter = PROXY_MaxParallelRequestsHandler_v3( internal_usage_cache=InternalUsageCache(local_cache) ) - rate_limiter = _PROXY_BatchRateLimiter( + rate_limiter = PROXY_BatchRateLimiter( internal_usage_cache=InternalUsageCache(local_cache), parallel_request_limiter=parallel_request_limiter, ) @@ -1448,17 +1448,17 @@ async def test_pre_call_enforces_project_itpm_limit_for_batch(): """Companion to the OTPM regression above: a project's ITPM quota must also apply to batch submissions.""" from litellm import DualCache - from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + from litellm.proxy.hooks.batch_rate_limiter import PROXY_BatchRateLimiter from litellm.proxy.hooks.parallel_request_limiter_v3 import ( - _PROXY_MaxParallelRequestsHandler_v3, + PROXY_MaxParallelRequestsHandler_v3, ) from litellm.proxy.utils import InternalUsageCache local_cache = DualCache() - parallel_request_limiter = _PROXY_MaxParallelRequestsHandler_v3( + parallel_request_limiter = PROXY_MaxParallelRequestsHandler_v3( internal_usage_cache=InternalUsageCache(local_cache) ) - rate_limiter = _PROXY_BatchRateLimiter( + rate_limiter = PROXY_BatchRateLimiter( internal_usage_cache=InternalUsageCache(local_cache), parallel_request_limiter=parallel_request_limiter, ) @@ -1505,17 +1505,17 @@ async def test_pre_call_enforces_project_otpm_limit_for_non_routing_row_model(): different, quota-limited model. That row's tokens must still be charged against its own model's project OTPM quota.""" from litellm import DualCache - from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + from litellm.proxy.hooks.batch_rate_limiter import PROXY_BatchRateLimiter from litellm.proxy.hooks.parallel_request_limiter_v3 import ( - _PROXY_MaxParallelRequestsHandler_v3, + PROXY_MaxParallelRequestsHandler_v3, ) from litellm.proxy.utils import InternalUsageCache local_cache = DualCache() - parallel_request_limiter = _PROXY_MaxParallelRequestsHandler_v3( + parallel_request_limiter = PROXY_MaxParallelRequestsHandler_v3( internal_usage_cache=InternalUsageCache(local_cache) ) - rate_limiter = _PROXY_BatchRateLimiter( + rate_limiter = PROXY_BatchRateLimiter( internal_usage_cache=InternalUsageCache(local_cache), parallel_request_limiter=parallel_request_limiter, ) @@ -1566,18 +1566,18 @@ async def test_pre_call_charges_each_row_model_against_its_own_project_quota(): model's request must succeed even though the over-limit model's row would fail on its own.""" from litellm import DualCache - from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + from litellm.proxy.hooks.batch_rate_limiter import PROXY_BatchRateLimiter from litellm.proxy.hooks.parallel_request_limiter_v3 import ( PROJECT_OTPM_DESCRIPTOR_KEY, - _PROXY_MaxParallelRequestsHandler_v3, + PROXY_MaxParallelRequestsHandler_v3, ) from litellm.proxy.utils import InternalUsageCache local_cache = DualCache() - parallel_request_limiter = _PROXY_MaxParallelRequestsHandler_v3( + parallel_request_limiter = PROXY_MaxParallelRequestsHandler_v3( internal_usage_cache=InternalUsageCache(local_cache) ) - rate_limiter = _PROXY_BatchRateLimiter( + rate_limiter = PROXY_BatchRateLimiter( internal_usage_cache=InternalUsageCache(local_cache), parallel_request_limiter=parallel_request_limiter, ) @@ -1650,7 +1650,7 @@ def test_should_not_skip_when_project_has_io_limit_for_non_routing_model(): rate_limiter = _make_rate_limiter() # No key/team/model-level limits at all -- only a project OTPM limit for a # model unrelated to the routing model below. - rate_limiter.parallel_request_limiter._create_rate_limit_descriptors.return_value = [ + rate_limiter.parallel_request_limiter.create_rate_limit_descriptors.return_value = [ {"rate_limit": {}} ] user = UserAPIKeyAuth( @@ -1673,7 +1673,7 @@ def test_should_skip_when_project_has_no_io_limits_and_no_other_limits(): with no ITPM/OTPM configuration anywhere must still get the fast-path skip when no other rate limits apply, exactly as before this fix.""" rate_limiter = _make_rate_limiter() - rate_limiter.parallel_request_limiter._create_rate_limit_descriptors.return_value = [ + rate_limiter.parallel_request_limiter.create_rate_limit_descriptors.return_value = [ {"rate_limit": {}} ] user = UserAPIKeyAuth( @@ -1693,9 +1693,9 @@ def test_should_skip_when_project_has_no_io_limits_and_no_other_limits(): @pytest.mark.asyncio async def test_count_input_file_usage_raises_on_non_bytes_content(): - from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + from litellm.proxy.hooks.batch_rate_limiter import PROXY_BatchRateLimiter - rate_limiter = _PROXY_BatchRateLimiter( + rate_limiter = PROXY_BatchRateLimiter( internal_usage_cache=MagicMock(), parallel_request_limiter=MagicMock(), ) @@ -1800,9 +1800,9 @@ async def test_count_input_file_usage_streams_without_building_list(): """count_input_file_usage must count requests/tokens in one streaming pass. Mocks the download; asserts the count is correct and that the dict-list helper is never called (a revert to the list approach would call it).""" - from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + from litellm.proxy.hooks.batch_rate_limiter import PROXY_BatchRateLimiter - rate_limiter = _PROXY_BatchRateLimiter( + rate_limiter = PROXY_BatchRateLimiter( internal_usage_cache=MagicMock(), parallel_request_limiter=MagicMock(), ) @@ -1852,9 +1852,9 @@ async def test_count_input_file_usage_enforces_models_when_token_counting_fails( NOT skip the model allowlist check. async_pre_call_hook swallows non-HTTP exceptions and submits the batch, so a raised counting error would otherwise fail open. The access check must still run and deny the restricted model.""" - from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + from litellm.proxy.hooks.batch_rate_limiter import PROXY_BatchRateLimiter - rate_limiter = _PROXY_BatchRateLimiter( + rate_limiter = PROXY_BatchRateLimiter( internal_usage_cache=MagicMock(), parallel_request_limiter=MagicMock(), ) @@ -1897,9 +1897,9 @@ async def test_count_input_file_usage_estimates_tokens_when_counting_fails_for_a zero the token total, which would let a caller evade the TPM limit by sending rows the counter cannot measure. The row falls back to a conservative size-based estimate so the batch proceeds with a non-zero count.""" - from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + from litellm.proxy.hooks.batch_rate_limiter import PROXY_BatchRateLimiter - rate_limiter = _PROXY_BatchRateLimiter( + rate_limiter = PROXY_BatchRateLimiter( internal_usage_cache=MagicMock(), parallel_request_limiter=MagicMock(), ) @@ -1941,9 +1941,9 @@ async def test_count_input_file_usage_collects_models_after_malformed_line(): named on a row AFTER a malformed line must still be collected and denied by the allowlist check, otherwise a caller could hide a restricted model behind a bad row.""" - from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + from litellm.proxy.hooks.batch_rate_limiter import PROXY_BatchRateLimiter - rate_limiter = _PROXY_BatchRateLimiter( + rate_limiter = PROXY_BatchRateLimiter( internal_usage_cache=MagicMock(), parallel_request_limiter=MagicMock(), ) @@ -1991,15 +1991,15 @@ def _output_estimator(): """A `_PROXY_BatchRateLimiter` whose output-token floor is observable: the no-`max_tokens` floor mock returns a distinctive sentinel so tests can tell "floor was used" apart from "an explicit cap was read".""" - from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + from litellm.proxy.hooks.batch_rate_limiter import PROXY_BatchRateLimiter from litellm.proxy.hooks.parallel_request_limiter_v3 import ( - _PROXY_MaxParallelRequestsHandler_v3, + PROXY_MaxParallelRequestsHandler_v3, ) limiter = MagicMock() limiter.no_max_tokens_output_floor.return_value = 999 - limiter.get_output_candidate_count = _PROXY_MaxParallelRequestsHandler_v3.get_output_candidate_count - return _PROXY_BatchRateLimiter( + limiter.get_output_candidate_count = PROXY_MaxParallelRequestsHandler_v3.get_output_candidate_count + return PROXY_BatchRateLimiter( internal_usage_cache=MagicMock(), parallel_request_limiter=limiter, ) @@ -2103,16 +2103,16 @@ def test_estimate_entry_output_tokens_multiplies_candidate_count(body_extra, exp def _enqueued_rate_limiter(): from litellm import DualCache - from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + from litellm.proxy.hooks.batch_rate_limiter import PROXY_BatchRateLimiter from litellm.proxy.hooks.parallel_request_limiter_v3 import ( - _PROXY_MaxParallelRequestsHandler_v3, + PROXY_MaxParallelRequestsHandler_v3, ) from litellm.proxy.utils import InternalUsageCache local_cache = DualCache(default_in_memory_ttl=60) internal_usage_cache = InternalUsageCache(local_cache) - parallel_request_limiter = _PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=internal_usage_cache) - rate_limiter = _PROXY_BatchRateLimiter( + parallel_request_limiter = PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=internal_usage_cache) + rate_limiter = PROXY_BatchRateLimiter( internal_usage_cache=internal_usage_cache, parallel_request_limiter=parallel_request_limiter, ) diff --git a/tests/unit/proxy/hooks/test_batch_rate_limiter.py b/tests/unit/proxy/hooks/test_batch_rate_limiter.py index 930f62fcd10..a3e60a89c9f 100644 --- a/tests/unit/proxy/hooks/test_batch_rate_limiter.py +++ b/tests/unit/proxy/hooks/test_batch_rate_limiter.py @@ -19,7 +19,7 @@ from litellm.constants import BATCH_TPD_WINDOW_SECONDS from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.hooks.batch_rate_limiter import BatchFileUsage from litellm.proxy.hooks.parallel_request_limiter_v3 import ( - _PROXY_MaxParallelRequestsHandler_v3, + PROXY_MaxParallelRequestsHandler_v3, ) from litellm.proxy.utils import InternalUsageCache, hash_token @@ -34,7 +34,7 @@ class _Clock: def _make_limiters(clock: _Clock | None = None): internal_usage_cache = InternalUsageCache(dual_cache=DualCache()) - rate_limiter = _PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=internal_usage_cache, time_provider=clock) + rate_limiter = PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=internal_usage_cache, time_provider=clock) batch_limiter = rate_limiter._get_batch_rate_limiter() assert batch_limiter is not None return internal_usage_cache, rate_limiter, batch_limiter @@ -252,7 +252,7 @@ def test_tpd_only_key_is_not_skipped_as_having_no_limits(): def test_online_descriptors_ignore_tpd_limit(): _internal_usage_cache, rate_limiter, _batch_limiter = _make_limiters() api_key = hash_token("online-key") - descriptors = rate_limiter._create_rate_limit_descriptors( + descriptors = rate_limiter.create_rate_limit_descriptors( user_api_key_dict=UserAPIKeyAuth(api_key=api_key, rpm_limit=5, tpd_limit=100, team_id="t", team_tpd_limit=9), data={"model": "gpt-4o"}, rpm_limit_type=None, diff --git a/tests/unit/proxy/hooks/test_dynamic_rate_limiter.py b/tests/unit/proxy/hooks/test_dynamic_rate_limiter.py index 611924ed658..cbb3e5f981c 100644 --- a/tests/unit/proxy/hooks/test_dynamic_rate_limiter.py +++ b/tests/unit/proxy/hooks/test_dynamic_rate_limiter.py @@ -5,9 +5,9 @@ import asyncio, importlib, litellm, os, pytest from litellm.caching.caching import DualCache from litellm.proxy.hooks.dynamic_rate_limiter import( - _PROXY_DynamicRateLimitHandler as DynamicRateLimitHandler, + PROXY_DynamicRateLimitHandler as DynamicRateLimitHandler, DynamicRateLimiterCache, - _PROXY_DynamicRateLimitHandler, + PROXY_DynamicRateLimitHandler, ) from litellm.types.utils import HiddenParams, ModelResponse from litellm import DualCache as DualCache_dynamic_rate, Router @@ -46,7 +46,7 @@ async def test_minute_rollover_between_sadd_and_get_reads_empty_window(): @pytest.mark.asyncio async def test_handler_threads_time_fn_to_internal_cache(): - handler = _PROXY_DynamicRateLimitHandler( + handler = PROXY_DynamicRateLimitHandler( internal_usage_cache=DualCache(), time_fn=lambda: datetime(2024, 1, 1, 10, 30, 0, tzinfo=timezone.utc), ) @@ -66,7 +66,7 @@ async def test_success_hook_updates_existing_hidden_params_storage() -> None: } ] ) - handler: Final = _PROXY_DynamicRateLimitHandler(internal_usage_cache=DualCache()) + handler: Final = PROXY_DynamicRateLimitHandler(internal_usage_cache=DualCache()) handler.update_variables(llm_router=router) response: Final = ModelResponse() hidden_params: Final = HiddenParams(model_id=model_id) diff --git a/tests/unit/proxy/hooks/test_dynamic_rate_limiter_v3.py b/tests/unit/proxy/hooks/test_dynamic_rate_limiter_v3.py index 527449bbc48..969c1eed63c 100644 --- a/tests/unit/proxy/hooks/test_dynamic_rate_limiter_v3.py +++ b/tests/unit/proxy/hooks/test_dynamic_rate_limiter_v3.py @@ -17,7 +17,7 @@ import litellm from litellm import DualCache, Router from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.hooks.dynamic_rate_limiter_v3 import ( - _PROXY_DynamicRateLimitHandlerV3 as DynamicRateLimitHandler, + PROXY_DynamicRateLimitHandlerV3 as DynamicRateLimitHandler, ) diff --git a/tests/unit/proxy/hooks/test_max_budget_per_session_limiter.py b/tests/unit/proxy/hooks/test_max_budget_per_session_limiter.py index a1b3f313814..45cc7eb8922 100644 --- a/tests/unit/proxy/hooks/test_max_budget_per_session_limiter.py +++ b/tests/unit/proxy/hooks/test_max_budget_per_session_limiter.py @@ -18,7 +18,7 @@ from litellm.caching.caching import DualCache from litellm.caching.redis_cache import _redis_circuit_breaker_guard from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.hooks.max_budget_per_session_limiter import ( - _PROXY_MaxBudgetPerSessionHandler, + PROXY_MaxBudgetPerSessionHandler, ) from litellm.proxy.utils import InternalUsageCache from litellm.types.agents import AgentResponse @@ -39,7 +39,7 @@ async def test_budget_per_session_under_budget_passes(): Requests under budget should pass through without error. """ local_cache = DualCache() - handler = _PROXY_MaxBudgetPerSessionHandler( + handler = PROXY_MaxBudgetPerSessionHandler( internal_usage_cache=InternalUsageCache(local_cache), ) user_api_key_dict = UserAPIKeyAuth( @@ -70,7 +70,7 @@ async def test_budget_per_session_exceeds_budget(): pre-call check should raise 429. """ local_cache = DualCache() - handler = _PROXY_MaxBudgetPerSessionHandler( + handler = PROXY_MaxBudgetPerSessionHandler( internal_usage_cache=InternalUsageCache(local_cache), ) user_api_key_dict = UserAPIKeyAuth( @@ -107,7 +107,7 @@ async def test_budget_per_session_independent_sessions(): Exhausting session A does not affect session B. """ local_cache = DualCache() - handler = _PROXY_MaxBudgetPerSessionHandler( + handler = PROXY_MaxBudgetPerSessionHandler( internal_usage_cache=InternalUsageCache(local_cache), ) user_api_key_dict = UserAPIKeyAuth( @@ -151,7 +151,7 @@ async def test_no_agent_id_passes(): When no agent_id is set on the key, all requests pass through. """ local_cache = DualCache() - handler = _PROXY_MaxBudgetPerSessionHandler( + handler = PROXY_MaxBudgetPerSessionHandler( internal_usage_cache=InternalUsageCache(local_cache), ) user_api_key_dict = UserAPIKeyAuth( @@ -190,7 +190,7 @@ class _OpenBreakerRedis: @pytest.mark.asyncio async def test_an_open_circuit_breaker_reads_session_spend_locally_without_a_warning(caplog): cache = DualCache(redis_cache=_OpenBreakerRedis()) # pyright: ignore[reportArgumentType] # duck-typed Redis double - handler = _PROXY_MaxBudgetPerSessionHandler(internal_usage_cache=InternalUsageCache(cache)) + handler = PROXY_MaxBudgetPerSessionHandler(internal_usage_cache=InternalUsageCache(cache)) caplog.clear() with caplog.at_level(logging.DEBUG, logger="LiteLLM Proxy"): diff --git a/tests/unit/proxy/hooks/test_max_iterations_limiter.py b/tests/unit/proxy/hooks/test_max_iterations_limiter.py index 20928ef46d5..5eb499042b5 100644 --- a/tests/unit/proxy/hooks/test_max_iterations_limiter.py +++ b/tests/unit/proxy/hooks/test_max_iterations_limiter.py @@ -13,7 +13,7 @@ from fastapi import HTTPException from litellm.caching.caching import DualCache from litellm.proxy._types import UserAPIKeyAuth -from litellm.proxy.hooks.max_iterations_limiter import _PROXY_MaxIterationsHandler +from litellm.proxy.hooks.max_iterations_limiter import PROXY_MaxIterationsHandler from litellm.proxy.utils import InternalUsageCache from litellm.types.agents import AgentResponse @@ -36,7 +36,7 @@ async def test_max_iterations_basic_enforcement(): - 4th request should raise 429 """ local_cache = DualCache() - handler = _PROXY_MaxIterationsHandler( + handler = PROXY_MaxIterationsHandler( internal_usage_cache=InternalUsageCache(local_cache), ) user_api_key_dict = UserAPIKeyAuth( @@ -46,9 +46,7 @@ async def test_max_iterations_basic_enforcement(): mock_agent = _make_mock_agent(max_iterations=3) - with patch( - "litellm.proxy.agent_endpoints.agent_registry.global_agent_registry" - ) as mock_registry: + with patch("litellm.proxy.agent_endpoints.agent_registry.global_agent_registry") as mock_registry: mock_registry.get_agent_by_id.return_value = mock_agent # First 3 requests should succeed @@ -81,7 +79,7 @@ async def test_max_iterations_different_sessions_independent(): - Exhausting Session A does not affect Session B """ local_cache = DualCache() - handler = _PROXY_MaxIterationsHandler( + handler = PROXY_MaxIterationsHandler( internal_usage_cache=InternalUsageCache(local_cache), ) user_api_key_dict = UserAPIKeyAuth( @@ -91,9 +89,7 @@ async def test_max_iterations_different_sessions_independent(): mock_agent = _make_mock_agent(max_iterations=2) - with patch( - "litellm.proxy.agent_endpoints.agent_registry.global_agent_registry" - ) as mock_registry: + with patch("litellm.proxy.agent_endpoints.agent_registry.global_agent_registry") as mock_registry: mock_registry.get_agent_by_id.return_value = mock_agent # Session A: 2 calls succeed @@ -140,7 +136,7 @@ async def test_max_iterations_no_agent_id_passes(): When no agent_id is set on the key, all requests pass through. """ local_cache = DualCache() - handler = _PROXY_MaxIterationsHandler( + handler = PROXY_MaxIterationsHandler( internal_usage_cache=InternalUsageCache(local_cache), ) user_api_key_dict = UserAPIKeyAuth( diff --git a/tests/unit/proxy/hooks/test_model_max_budget_limiter.py b/tests/unit/proxy/hooks/test_model_max_budget_limiter.py index ffb60fb4651..8de34e6391e 100644 --- a/tests/unit/proxy/hooks/test_model_max_budget_limiter.py +++ b/tests/unit/proxy/hooks/test_model_max_budget_limiter.py @@ -8,7 +8,7 @@ import pytest from litellm.caching.caching import DualCache from litellm.proxy.hooks.model_max_budget_limiter import ( - _PROXY_VirtualKeyModelMaxBudgetLimiter, + PROXY_VirtualKeyModelMaxBudgetLimiter, ) from litellm.types.caching import RedisPipelineIncrementOperation from litellm.types.utils import LiteLLMBatch, Usage @@ -54,23 +54,23 @@ def _event(call_type: str, response_cost: float) -> dict[str, object]: } -async def _poll(limiter: _PROXY_VirtualKeyModelMaxBudgetLimiter, batch: LiteLLMBatch, response_cost: float) -> None: +async def _poll(limiter: PROXY_VirtualKeyModelMaxBudgetLimiter, batch: LiteLLMBatch, response_cost: float) -> None: await limiter.async_log_success_event( _event("aretrieve_batch", response_cost), response_obj=batch, start_time=None, end_time=None ) -async def _chat(limiter: _PROXY_VirtualKeyModelMaxBudgetLimiter) -> None: +async def _chat(limiter: PROXY_VirtualKeyModelMaxBudgetLimiter) -> None: await limiter.async_log_success_event( _event("acompletion", CHAT_COST), response_obj=None, start_time=None, end_time=None ) -async def _spend(limiter: _PROXY_VirtualKeyModelMaxBudgetLimiter, spend_key: str) -> float: +async def _spend(limiter: PROXY_VirtualKeyModelMaxBudgetLimiter, spend_key: str) -> float: return await limiter.dual_cache.async_get_cache(key=spend_key) or 0.0 -def _local_spend(limiter: _PROXY_VirtualKeyModelMaxBudgetLimiter, spend_key: str) -> float: +def _local_spend(limiter: PROXY_VirtualKeyModelMaxBudgetLimiter, spend_key: str) -> float: return limiter.dual_cache.in_memory_cache.get_cache(key=spend_key) or 0.0 @@ -133,8 +133,8 @@ class _SharedRedisDouble: return [await self.async_increment(op["key"], op["increment_value"], ttl=op["ttl"]) for op in increment_list] -def _worker(redis: _SharedRedisDouble) -> _PROXY_VirtualKeyModelMaxBudgetLimiter: - return _PROXY_VirtualKeyModelMaxBudgetLimiter( +def _worker(redis: _SharedRedisDouble) -> PROXY_VirtualKeyModelMaxBudgetLimiter: + return PROXY_VirtualKeyModelMaxBudgetLimiter( dual_cache=DualCache(redis_cache=redis) # pyright: ignore[reportArgumentType] # duck-typed Redis double ) @@ -145,7 +145,7 @@ async def _drain_redis_pushes() -> None: @pytest.mark.asyncio async def test_polls_of_a_finished_batch_charge_each_per_model_budget_once(): - limiter: Final = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=DualCache()) + limiter: Final = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=DualCache()) first: Final = _batch("batch_first", "completed") await _poll(limiter, _batch("batch_first", "in_progress"), response_cost=0) @@ -158,7 +158,7 @@ async def test_polls_of_a_finished_batch_charge_each_per_model_budget_once(): @pytest.mark.asyncio async def test_a_second_batch_and_chat_requests_still_charge_the_budget(): - limiter: Final = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=DualCache()) + limiter: Final = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=DualCache()) await _poll(limiter, _batch("batch_first", "completed"), response_cost=BATCH_COST) await _poll(limiter, _batch("batch_first", "completed"), response_cost=BATCH_COST) diff --git a/tests/unit/proxy/hooks/test_parallel_request_limiter.py b/tests/unit/proxy/hooks/test_parallel_request_limiter.py index c42d3b51799..1772f41a9d3 100644 --- a/tests/unit/proxy/hooks/test_parallel_request_limiter.py +++ b/tests/unit/proxy/hooks/test_parallel_request_limiter.py @@ -9,7 +9,7 @@ import pytest from litellm.caching.caching import DualCache from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.hooks.parallel_request_limiter import ( - _PROXY_MaxParallelRequestsHandler, + PROXY_MaxParallelRequestsHandler, ) from litellm.proxy.utils import InternalUsageCache, hash_token from litellm.types.utils import EmbeddingResponse, TextCompletionResponse, Usage @@ -17,7 +17,7 @@ from litellm.types.utils import EmbeddingResponse, TextCompletionResponse, Usage @pytest.mark.asyncio async def test_pre_call_hook_counts_a_cli_session_under_the_per_user_alias_not_the_login_token(): - handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(DualCache())) + handler = PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(DualCache())) session = UserAPIKeyAuth( api_key="cli-session-Qm7xJ2kP9sLw4vT1nR8yAa", user_id="alice", @@ -62,7 +62,7 @@ async def test_async_log_success_event_counts_non_chat_response_tokens(response_ team_id = "litellm-team" end_user_id = "customer-1" - parallel_request_handler = _PROXY_MaxParallelRequestsHandler( + parallel_request_handler = PROXY_MaxParallelRequestsHandler( internal_usage_cache=InternalUsageCache(DualCache()) ) diff --git a/tests/unit/proxy/hooks/test_parallel_request_limiter_v3.py b/tests/unit/proxy/hooks/test_parallel_request_limiter_v3.py index fecedd8c498..d9cf8536912 100644 --- a/tests/unit/proxy/hooks/test_parallel_request_limiter_v3.py +++ b/tests/unit/proxy/hooks/test_parallel_request_limiter_v3.py @@ -26,6 +26,7 @@ from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache from litellm.proxy.hooks.parallel_request_limiter_v3 import ( PARALLEL_REQUEST_SLOT_TTL_SECONDS, ParallelSlotAcquisition, + PROXY_MaxParallelRequestsHandler_v3, RateLimitDescriptor, RateLimitedModel, RateLimitResponse, @@ -36,7 +37,7 @@ from litellm.proxy.hooks.parallel_request_limiter_v3 import ( get_request_stash, ) from litellm.proxy.hooks.parallel_request_limiter_v3 import ( - _PROXY_MaxParallelRequestsHandler_v3 as _PROXY_MaxParallelRequestsHandler, + PROXY_MaxParallelRequestsHandler_v3 as _PROXY_MaxParallelRequestsHandler, ) from litellm.proxy.utils import InternalUsageCache, ProxyLogging, hash_token from litellm.types.caching import RedisPipelineIncrementOperation @@ -99,7 +100,7 @@ def test_api_key_descriptor_applies_budget_throttle( budget_throttle_pct=throttle_pct, ) - descriptors = handler._create_rate_limit_descriptors( + descriptors = handler.create_rate_limit_descriptors( user_api_key_dict=user_api_key_dict, data={}, rpm_limit_type=None, @@ -2914,7 +2915,7 @@ class TestGetTotalTokensFromUsageCacheExclusion: def handler(self): """Create a handler instance for testing.""" local_cache = DualCache() - return _PROXY_MaxParallelRequestsHandler( + return PROXY_MaxParallelRequestsHandler_v3( internal_usage_cache=InternalUsageCache(local_cache), ) @@ -3786,9 +3787,9 @@ async def test_failure_event_settles_project_itpm_otpm_at_recovered_partial_usag # ----------------------- Per-MCP-server rate limiting (v3) ----------------------- -def _make_mcp_handler() -> tuple[_PROXY_MaxParallelRequestsHandler, DualCache]: +def _make_mcp_handler() -> tuple[PROXY_MaxParallelRequestsHandler_v3, DualCache]: local_cache: Final = DualCache() - handler: Final = _PROXY_MaxParallelRequestsHandler( + handler: Final = PROXY_MaxParallelRequestsHandler_v3( internal_usage_cache=InternalUsageCache(local_cache) ) return handler, local_cache @@ -4878,7 +4879,7 @@ async def test_per_tag_rate_limit_independent_counters_v3(monkeypatch): @pytest.mark.asyncio async def test_per_tag_descriptor_creation_v3(): """ - _create_rate_limit_descriptors emits a tag_per_key descriptor carrying the + create_rate_limit_descriptors emits a tag_per_key descriptor carrying the configured RPM limit only for request tags present in the configured map. """ _api_key = hash_token("sk-per-tag-desc") @@ -4890,7 +4891,7 @@ async def test_per_tag_descriptor_creation_v3(): internal_usage_cache=InternalUsageCache(DualCache()) ) - descriptors = handler._create_rate_limit_descriptors( + descriptors = handler.create_rate_limit_descriptors( user_api_key_dict=user_api_key_dict, data={"model": "gpt-3.5-turbo", "metadata": {"tags": ["cell-1", "cell-2"]}}, rpm_limit_type=None, @@ -4916,7 +4917,7 @@ async def test_per_tag_descriptor_absent_without_config_v3(): internal_usage_cache=InternalUsageCache(DualCache()) ) - descriptors = handler._create_rate_limit_descriptors( + descriptors = handler.create_rate_limit_descriptors( user_api_key_dict=user_api_key_dict, data={"model": "gpt-3.5-turbo", "metadata": {"tags": ["cell-1"]}}, rpm_limit_type=None, @@ -6425,7 +6426,7 @@ async def test_conflicting_token_limits_cannot_bypass_tpm_reservation(): def _enqueued_test_handler() -> _PROXY_MaxParallelRequestsHandler: - return _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(DualCache(default_in_memory_ttl=60))) + return PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=InternalUsageCache(DualCache(default_in_memory_ttl=60))) def _batch_response(batch_id: str, status: str): @@ -6745,18 +6746,18 @@ def _handler_with_redis( ): internal_usage_cache = InternalUsageCache(DualCache(redis_cache=redis)) # pyright: ignore[reportArgumentType] # duck-typed Redis double if fail_closed is None and force_hash_tag_grouping is None: - return _PROXY_MaxParallelRequestsHandler(internal_usage_cache=internal_usage_cache) + return PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=internal_usage_cache) if fail_closed is None: - return _PROXY_MaxParallelRequestsHandler( + return PROXY_MaxParallelRequestsHandler_v3( internal_usage_cache=internal_usage_cache, force_hash_tag_grouping_resolver=lambda: force_hash_tag_grouping, ) if force_hash_tag_grouping is None: - return _PROXY_MaxParallelRequestsHandler( + return PROXY_MaxParallelRequestsHandler_v3( internal_usage_cache=internal_usage_cache, fail_closed_resolver=lambda: fail_closed, ) - return _PROXY_MaxParallelRequestsHandler( + return PROXY_MaxParallelRequestsHandler_v3( internal_usage_cache=internal_usage_cache, fail_closed_resolver=lambda: fail_closed, force_hash_tag_grouping_resolver=lambda: force_hash_tag_grouping, @@ -6853,7 +6854,7 @@ async def _admit(handler, auth, data=None): async def _read_only_check(handler, auth): - descriptors = handler._create_rate_limit_descriptors( + descriptors = handler.create_rate_limit_descriptors( user_api_key_dict=auth, data={"model": "test-model"}, rpm_limit_type=None, @@ -7675,7 +7676,7 @@ async def test_managed_invocations_enforce_actor_and_target_rate_policies( cache: Final = DualCache() handler: Final = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(cache)) monkeypatch.setattr(handler, "_get_agent_from_registry", lambda _: None) - descriptors: Final = handler._create_rate_limit_descriptors( + descriptors: Final = handler.create_rate_limit_descriptors( user_api_key_dict=auth, data={"model": "a2a/target", "litellm_session_id": "session"}, rpm_limit_type=None, diff --git a/tests/unit/proxy/hooks/test_prompt_injection_detection.py b/tests/unit/proxy/hooks/test_prompt_injection_detection.py index b82b1be3b2a..3c871ca7333 100644 --- a/tests/unit/proxy/hooks/test_prompt_injection_detection.py +++ b/tests/unit/proxy/hooks/test_prompt_injection_detection.py @@ -11,7 +11,7 @@ import litellm from litellm.caching.caching import DualCache from litellm.proxy._types import LiteLLMPromptInjectionParams, UserAPIKeyAuth from litellm.proxy.hooks.prompt_injection_detection import ( - _OPTIONAL_PromptInjectionDetection, + OPTIONAL_PromptInjectionDetection, ) from litellm.proxy.utils import ProxyLogging from litellm.router import Router @@ -21,8 +21,8 @@ from litellm.utils import _invalidate_model_cost_lowercase_map from tests._vcr_conftest_common import install_live_call_probe, record_vcr_outcome -def _moderation_detector(verdict: str) -> _OPTIONAL_PromptInjectionDetection: - detector = _OPTIONAL_PromptInjectionDetection( +def _moderation_detector(verdict: str) -> OPTIONAL_PromptInjectionDetection: + detector = OPTIONAL_PromptInjectionDetection( prompt_injection_params=LiteLLMPromptInjectionParams( heuristics_check=False, llm_api_check=True, @@ -48,7 +48,7 @@ LONG_SAFE_PROMPT = "Summarize the quarterly revenue report for the finance team. @pytest.mark.asyncio async def test_acompletion_call_type_rejects_prompt_injection(): - prompt_injection_detection = _OPTIONAL_PromptInjectionDetection() + prompt_injection_detection = OPTIONAL_PromptInjectionDetection() user_key = UserAPIKeyAuth(api_key="sk-test") cache = DualCache() data = { @@ -74,7 +74,7 @@ async def test_acompletion_call_type_rejects_prompt_injection(): @pytest.mark.asyncio async def test_acompletion_call_type_allows_safe_prompt(): - prompt_injection_detection = _OPTIONAL_PromptInjectionDetection() + prompt_injection_detection = OPTIONAL_PromptInjectionDetection() user_key = UserAPIKeyAuth(api_key="sk-test") cache = DualCache() data = { @@ -153,7 +153,7 @@ async def test_proxy_during_call_hook_runs_configured_llm_api_check(monkeypatch) @pytest.mark.asyncio async def test_heuristics_check_keeps_event_loop_responsive(): - detector = _OPTIONAL_PromptInjectionDetection( + detector = OPTIONAL_PromptInjectionDetection( prompt_injection_params=LiteLLMPromptInjectionParams(heuristics_check=True) ) data = {"model": "test-model", "messages": [{"role": "user", "content": LONG_SAFE_PROMPT}]} @@ -182,7 +182,7 @@ async def test_heuristics_check_keeps_event_loop_responsive(): @pytest.mark.asyncio async def test_heuristics_check_does_not_occupy_default_executor(): - detector = _OPTIONAL_PromptInjectionDetection( + detector = OPTIONAL_PromptInjectionDetection( prompt_injection_params=LiteLLMPromptInjectionParams(heuristics_check=True) ) data = {"model": "test-model", "messages": [{"role": "user", "content": LONG_SAFE_PROMPT}]} @@ -329,7 +329,7 @@ async def test_prompt_injection_attack_valid_attack(): """ Tests if prompt injection detection catches a valid attack """ - prompt_injection_detection = _OPTIONAL_PromptInjectionDetection() + prompt_injection_detection = OPTIONAL_PromptInjectionDetection() _api_key = "sk-98765" user_api_key_dict = UserAPIKeyAuth(api_key=_api_key) @@ -360,7 +360,7 @@ async def test_prompt_injection_attack_invalid_attack(): Tests if prompt injection detection passes an invalid attack, which contains just 1 word """ litellm.set_verbose = True - prompt_injection_detection = _OPTIONAL_PromptInjectionDetection() + prompt_injection_detection = OPTIONAL_PromptInjectionDetection() _api_key = "sk-98765" user_api_key_dict = UserAPIKeyAuth(api_key=_api_key) @@ -398,7 +398,7 @@ async def test_prompt_injection_llm_eval(): llm_api_system_prompt="Detect if a prompt is safe to run. Return 'UNSAFE' if not.", llm_api_fail_call_string="UNSAFE", ) - prompt_injection_detection = _OPTIONAL_PromptInjectionDetection( + prompt_injection_detection = OPTIONAL_PromptInjectionDetection( prompt_injection_params=_prompt_injection_params, ) diff --git a/tests/unit/proxy/hooks/test_proxy_hooks_init.py b/tests/unit/proxy/hooks/test_proxy_hooks_init.py index 7f07fa7966c..a6edd3db944 100644 --- a/tests/unit/proxy/hooks/test_proxy_hooks_init.py +++ b/tests/unit/proxy/hooks/test_proxy_hooks_init.py @@ -23,14 +23,14 @@ def test_managed_files_hook_registered(): pytest.importorskip("litellm_enterprise") assert "managed_files" in PROXY_HOOKS hook_cls = get_proxy_hook("managed_files") - assert hook_cls.__name__ == "PROXY_LiteLLMManagedFiles" + assert hook_cls.__name__ == "_PROXY_LiteLLMManagedFiles" def test_managed_vector_stores_hook_registered(): pytest.importorskip("litellm_enterprise") assert "managed_vector_stores" in PROXY_HOOKS hook_cls = get_proxy_hook("managed_vector_stores") - assert hook_cls.__name__ == "PROXY_LiteLLMManagedVectorStores" + assert hook_cls.__name__ == "_PROXY_LiteLLMManagedVectorStores" def test_isolation_module_does_not_pull_in_proxy_utils(): diff --git a/tests/unit/proxy/hooks/test_proxy_rate_limit_provider_field.py b/tests/unit/proxy/hooks/test_proxy_rate_limit_provider_field.py index 49bbd498cb9..16b5406bc21 100644 --- a/tests/unit/proxy/hooks/test_proxy_rate_limit_provider_field.py +++ b/tests/unit/proxy/hooks/test_proxy_rate_limit_provider_field.py @@ -44,21 +44,21 @@ from litellm.exceptions import RateLimitError from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.hooks.batch_rate_limiter import ( BatchFileUsage, - _PROXY_BatchRateLimiter, + PROXY_BatchRateLimiter, ) -from litellm.proxy.hooks.dynamic_rate_limiter import _PROXY_DynamicRateLimitHandler +from litellm.proxy.hooks.dynamic_rate_limiter import PROXY_DynamicRateLimitHandler from litellm.proxy.hooks.dynamic_rate_limiter_v3 import ( - _PROXY_DynamicRateLimitHandlerV3, + PROXY_DynamicRateLimitHandlerV3, ) from litellm.proxy.hooks.max_budget_per_session_limiter import ( - _PROXY_MaxBudgetPerSessionHandler, + PROXY_MaxBudgetPerSessionHandler, ) -from litellm.proxy.hooks.max_iterations_limiter import _PROXY_MaxIterationsHandler +from litellm.proxy.hooks.max_iterations_limiter import PROXY_MaxIterationsHandler from litellm.proxy.hooks.parallel_request_limiter import ( - _PROXY_MaxParallelRequestsHandler, + PROXY_MaxParallelRequestsHandler, ) from litellm.proxy.hooks.parallel_request_limiter_v3 import ( - _PROXY_MaxParallelRequestsHandler_v3, + PROXY_MaxParallelRequestsHandler_v3, ) from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError from litellm.proxy.hooks.rate_limiter_utils import ( @@ -166,9 +166,7 @@ class TestResolveLLMProviderForRateLimit: "litellm.proxy.proxy_server.llm_router", None, ): - resolved_model, provider = resolve_llm_provider_for_rate_limit( - "anything" - ) + resolved_model, provider = resolve_llm_provider_for_rate_limit("anything") assert provider == PROXY_LLM_PROVIDER_FALLBACK assert resolved_model == "anything" @@ -265,9 +263,7 @@ class TestResolveLLMProviderForRateLimit: "litellm.proxy.proxy_server.llm_router", _FakeRouter(), ): - resolved_model, provider = resolve_llm_provider_for_rate_limit( - "not-an-alias" - ) + resolved_model, provider = resolve_llm_provider_for_rate_limit("not-an-alias") assert provider == PROXY_LLM_PROVIDER_FALLBACK assert resolved_model == "not-an-alias" @@ -309,9 +305,7 @@ async def test_parallel_request_limiter_v1_populates_provider_when_at_rpm_limit( Trip the per-key RPM cap and assert the raised exception carries ``model`` / ``llm_provider`` resolved from ``data["model"]``. """ - handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(DualCache()) - ) + handler = PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(DualCache())) user_api_key_dict = UserAPIKeyAuth( api_key="sk-rl-test", max_parallel_requests=10, @@ -350,9 +344,7 @@ async def test_parallel_request_limiter_v1_zero_limit_path_populates_provider(): ``raise_rate_limit_error`` path. That path receives ``requested_model`` via the call-site change and must pass it through. """ - handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(DualCache()) - ) + handler = PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(DualCache())) user_api_key_dict = UserAPIKeyAuth( api_key="sk-rl-zero", max_parallel_requests=0, @@ -378,9 +370,7 @@ async def test_parallel_request_limiter_v1_zero_limit_path_populates_provider(): @pytest.mark.asyncio async def test_parallel_request_limiter_v1_global_limit_populates_provider(): """global_max_parallel_requests path also threads the model through.""" - handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(DualCache()) - ) + handler = PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(DualCache())) user_api_key_dict = UserAPIKeyAuth(api_key="sk-global") # Pre-fill the global counter so the next call exceeds it. @@ -414,9 +404,7 @@ async def test_parallel_request_limiter_v1_unknown_model_falls_back(): When ``data["model"]`` is unparseable, the resolver falls back to ``litellm_proxy`` — and crucially does *not* leak a secondary exception. """ - handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(DualCache()) - ) + handler = PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(DualCache())) user_api_key_dict = UserAPIKeyAuth( api_key="sk-rl-unknown", max_parallel_requests=10, @@ -450,9 +438,7 @@ async def test_parallel_request_limiter_v1_unknown_model_falls_back(): @pytest.mark.asyncio async def test_parallel_request_limiter_v1_missing_model_falls_back(): - handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(DualCache()) - ) + handler = PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(DualCache())) user_api_key_dict = UserAPIKeyAuth( api_key="sk-rl-no-model", max_parallel_requests=10, @@ -509,9 +495,7 @@ def _v3_over_limit_response(rate_limit_type: str = "requests") -> dict: ], ) async def test_parallel_request_limiter_v3_populates_provider(model, expected_provider): - handler = _PROXY_MaxParallelRequestsHandler_v3( - internal_usage_cache=InternalUsageCache(DualCache()) - ) + handler = PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=InternalUsageCache(DualCache())) descriptors = [{"key": "key", "value": "v", "rate_limit": {"requests_per_unit": 1}}] over = _v3_over_limit_response() @@ -535,9 +519,7 @@ async def test_parallel_request_limiter_v3_populates_provider(model, expected_pr @pytest.mark.asyncio async def test_parallel_request_limiter_v3_unknown_model_falls_back(): - handler = _PROXY_MaxParallelRequestsHandler_v3( - internal_usage_cache=InternalUsageCache(DualCache()) - ) + handler = PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=InternalUsageCache(DualCache())) descriptors = [{"key": "key", "value": "v", "rate_limit": {"requests_per_unit": 1}}] with pytest.raises(HTTPException) as exc_info: @@ -553,9 +535,7 @@ async def test_parallel_request_limiter_v3_unknown_model_falls_back(): @pytest.mark.asyncio async def test_parallel_request_limiter_v3_missing_model_falls_back(): - handler = _PROXY_MaxParallelRequestsHandler_v3( - internal_usage_cache=InternalUsageCache(DualCache()) - ) + handler = PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=InternalUsageCache(DualCache())) descriptors = [{"key": "key", "value": "v", "rate_limit": {"requests_per_unit": 1}}] with pytest.raises(HTTPException) as exc_info: @@ -576,7 +556,7 @@ async def test_parallel_request_limiter_v3_missing_model_falls_back(): @pytest.mark.asyncio async def test_dynamic_rate_limiter_v1_tpm_zero_populates_provider(): - handler = _PROXY_DynamicRateLimitHandler(internal_usage_cache=DualCache()) + handler = PROXY_DynamicRateLimitHandler(internal_usage_cache=DualCache()) handler.check_available_usage = AsyncMock(return_value=(0, 5, 100, 5, 1)) user_api_key_dict = UserAPIKeyAuth(api_key="sk-dyn") @@ -599,7 +579,7 @@ async def test_dynamic_rate_limiter_v1_tpm_zero_populates_provider(): @pytest.mark.asyncio async def test_dynamic_rate_limiter_v1_rpm_zero_populates_provider(): - handler = _PROXY_DynamicRateLimitHandler(internal_usage_cache=DualCache()) + handler = PROXY_DynamicRateLimitHandler(internal_usage_cache=DualCache()) handler.check_available_usage = AsyncMock(return_value=(5, 0, 5, 100, 1)) user_api_key_dict = UserAPIKeyAuth(api_key="sk-dyn") @@ -620,7 +600,7 @@ async def test_dynamic_rate_limiter_v1_rpm_zero_populates_provider(): @pytest.mark.asyncio async def test_dynamic_rate_limiter_v1_unknown_model_falls_back(): - handler = _PROXY_DynamicRateLimitHandler(internal_usage_cache=DualCache()) + handler = PROXY_DynamicRateLimitHandler(internal_usage_cache=DualCache()) handler.check_available_usage = AsyncMock(return_value=(0, 5, 100, 5, 1)) user_api_key_dict = UserAPIKeyAuth(api_key="sk-dyn") @@ -655,7 +635,7 @@ async def test_dynamic_rate_limiter_v3_model_capacity_path_populates_provider(): """ from litellm.types.router import ModelGroupInfo - handler = _PROXY_DynamicRateLimitHandlerV3(internal_usage_cache=DualCache()) + handler = PROXY_DynamicRateLimitHandlerV3(internal_usage_cache=DualCache()) handler.v3_limiter.atomic_check_and_increment_by_n = AsyncMock( return_value={ "overall_code": "OVER_LIMIT", @@ -704,7 +684,7 @@ async def test_dynamic_rate_limiter_v3_unknown_descriptor_path_populates_provide """Fail-closed unknown-descriptor branch must still attribute provider.""" from litellm.types.router import ModelGroupInfo - handler = _PROXY_DynamicRateLimitHandlerV3(internal_usage_cache=DualCache()) + handler = PROXY_DynamicRateLimitHandlerV3(internal_usage_cache=DualCache()) handler.v3_limiter.atomic_check_and_increment_by_n = AsyncMock( return_value={ "overall_code": "OVER_LIMIT", @@ -773,16 +753,12 @@ async def test_batch_rate_limiter_populates_provider(): """ parallel_limiter = MagicMock() parallel_limiter.window_size = 60 - parallel_limiter._create_rate_limit_descriptors = MagicMock( - return_value=[ - {"key": "key", "value": "v", "rate_limit": {"requests_per_unit": 10}} - ] - ) - parallel_limiter.atomic_check_and_increment_by_n = AsyncMock( - return_value=_batch_over_limit_response() + parallel_limiter.create_rate_limit_descriptors = MagicMock( + return_value=[{"key": "key", "value": "v", "rate_limit": {"requests_per_unit": 10}}] ) + parallel_limiter.atomic_check_and_increment_by_n = AsyncMock(return_value=_batch_over_limit_response()) - handler = _PROXY_BatchRateLimiter( + handler = PROXY_BatchRateLimiter( internal_usage_cache=InternalUsageCache(DualCache()), parallel_request_limiter=parallel_limiter, ) @@ -805,16 +781,12 @@ async def test_batch_rate_limiter_populates_provider(): async def test_batch_rate_limiter_unknown_model_falls_back(): parallel_limiter = MagicMock() parallel_limiter.window_size = 60 - parallel_limiter._create_rate_limit_descriptors = MagicMock( - return_value=[ - {"key": "key", "value": "v", "rate_limit": {"requests_per_unit": 10}} - ] - ) - parallel_limiter.atomic_check_and_increment_by_n = AsyncMock( - return_value=_batch_over_limit_response() + parallel_limiter.create_rate_limit_descriptors = MagicMock( + return_value=[{"key": "key", "value": "v", "rate_limit": {"requests_per_unit": 10}}] ) + parallel_limiter.atomic_check_and_increment_by_n = AsyncMock(return_value=_batch_over_limit_response()) - handler = _PROXY_BatchRateLimiter( + handler = PROXY_BatchRateLimiter( internal_usage_cache=InternalUsageCache(DualCache()), parallel_request_limiter=parallel_limiter, ) @@ -846,14 +818,10 @@ def _make_iter_agent(max_iterations: int) -> AgentResponse: @pytest.mark.asyncio async def test_max_iterations_limiter_populates_provider(): local_cache = DualCache() - handler = _PROXY_MaxIterationsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + handler = PROXY_MaxIterationsHandler(internal_usage_cache=InternalUsageCache(local_cache)) user_api_key_dict = UserAPIKeyAuth(api_key="sk-iter", agent_id="agent-iter") - with patch( - "litellm.proxy.agent_endpoints.agent_registry.global_agent_registry" - ) as mock_registry: + with patch("litellm.proxy.agent_endpoints.agent_registry.global_agent_registry") as mock_registry: mock_registry.get_agent_by_id.return_value = _make_iter_agent(max_iterations=1) await handler.async_pre_call_hook( @@ -887,14 +855,10 @@ async def test_max_iterations_limiter_populates_provider(): @pytest.mark.asyncio async def test_max_iterations_limiter_unknown_model_falls_back(): local_cache = DualCache() - handler = _PROXY_MaxIterationsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + handler = PROXY_MaxIterationsHandler(internal_usage_cache=InternalUsageCache(local_cache)) user_api_key_dict = UserAPIKeyAuth(api_key="sk-iter", agent_id="agent-iter") - with patch( - "litellm.proxy.agent_endpoints.agent_registry.global_agent_registry" - ) as mock_registry: + with patch("litellm.proxy.agent_endpoints.agent_registry.global_agent_registry") as mock_registry: mock_registry.get_agent_by_id.return_value = _make_iter_agent(max_iterations=1) await handler.async_pre_call_hook( @@ -937,22 +901,12 @@ def _make_session_budget_agent(max_budget: float) -> AgentResponse: @pytest.mark.asyncio async def test_max_budget_per_session_limiter_populates_provider(): - handler = _PROXY_MaxBudgetPerSessionHandler( - internal_usage_cache=InternalUsageCache(DualCache()) - ) - user_api_key_dict = UserAPIKeyAuth( - api_key="sk-session-budget", agent_id="agent-session-budget" - ) + handler = PROXY_MaxBudgetPerSessionHandler(internal_usage_cache=InternalUsageCache(DualCache())) + user_api_key_dict = UserAPIKeyAuth(api_key="sk-session-budget", agent_id="agent-session-budget") - with patch( - "litellm.proxy.agent_endpoints.agent_registry.global_agent_registry" - ) as mock_registry: - mock_registry.get_agent_by_id.return_value = _make_session_budget_agent( - max_budget=1.0 - ) - with patch.object( - handler, "_get_current_spend", new=AsyncMock(return_value=5.0) - ): + with patch("litellm.proxy.agent_endpoints.agent_registry.global_agent_registry") as mock_registry: + mock_registry.get_agent_by_id.return_value = _make_session_budget_agent(max_budget=1.0) + with patch.object(handler, "_get_current_spend", new=AsyncMock(return_value=5.0)): with pytest.raises(HTTPException) as exc_info: await handler.async_pre_call_hook( user_api_key_dict=user_api_key_dict, @@ -972,22 +926,12 @@ async def test_max_budget_per_session_limiter_populates_provider(): @pytest.mark.asyncio async def test_max_budget_per_session_limiter_unknown_model_falls_back(): - handler = _PROXY_MaxBudgetPerSessionHandler( - internal_usage_cache=InternalUsageCache(DualCache()) - ) - user_api_key_dict = UserAPIKeyAuth( - api_key="sk-session-budget", agent_id="agent-session-budget" - ) + handler = PROXY_MaxBudgetPerSessionHandler(internal_usage_cache=InternalUsageCache(DualCache())) + user_api_key_dict = UserAPIKeyAuth(api_key="sk-session-budget", agent_id="agent-session-budget") - with patch( - "litellm.proxy.agent_endpoints.agent_registry.global_agent_registry" - ) as mock_registry: - mock_registry.get_agent_by_id.return_value = _make_session_budget_agent( - max_budget=1.0 - ) - with patch.object( - handler, "_get_current_spend", new=AsyncMock(return_value=5.0) - ): + with patch("litellm.proxy.agent_endpoints.agent_registry.global_agent_registry") as mock_registry: + mock_registry.get_agent_by_id.return_value = _make_session_budget_agent(max_budget=1.0) + with patch.object(handler, "_get_current_spend", new=AsyncMock(return_value=5.0)): with pytest.raises(HTTPException) as exc_info: await handler.async_pre_call_hook( user_api_key_dict=user_api_key_dict, @@ -1056,10 +1000,7 @@ def test_prometheus_exception_class_name_back_compat_for_budget_exceeded_error() # Default (empty llm_provider) path — same literal label. err_no_provider = litellm.BudgetExceededError(current_cost=1.0, max_budget=0.5) - assert ( - PrometheusLogger._get_exception_class_name(err_no_provider) - == "BudgetExceededError" - ) + assert PrometheusLogger._get_exception_class_name(err_no_provider) == "BudgetExceededError" if __name__ == "__main__": diff --git a/tests/unit/proxy/hooks/test_proxy_track_cost_callback.py b/tests/unit/proxy/hooks/test_proxy_track_cost_callback.py index 7afa275c801..811176c4f4e 100644 --- a/tests/unit/proxy/hooks/test_proxy_track_cost_callback.py +++ b/tests/unit/proxy/hooks/test_proxy_track_cost_callback.py @@ -17,7 +17,7 @@ from litellm.proxy.db.db_spend_update_writer import DBSpendUpdateWriter from litellm.proxy.db.spend_log_tool_index import response_tool_call_names from litellm.proxy.hooks.proxy_track_cost_callback import ( _get_budget_reservation_from_metadata, - _ProxyDBLogger, + ProxyDBLogger, _should_track_cost_callback, _update_database_and_spend_counters, run_spend_event, @@ -33,7 +33,7 @@ from litellm.types.utils import CallTypes, LiteLLMBatch, ModelResponse, Usage @pytest.mark.asyncio async def test_async_post_call_failure_hook(): # Setup - logger = _ProxyDBLogger() + logger = ProxyDBLogger() # Mock user_api_key_dict user_api_key_dict = UserAPIKeyAuth( @@ -103,7 +103,7 @@ async def test_async_post_call_failure_hook_carries_guardrail_info_from_litellm_ consume provider usage units, so the info must be carried over or the failure row logs guardrail_information: null. """ - logger = _ProxyDBLogger() + logger = ProxyDBLogger() guardrail_info = [ { "guardrail_name": "bedrock-guard", @@ -136,7 +136,7 @@ async def test_async_post_call_failure_hook_carries_guardrail_info_from_litellm_ @pytest.mark.asyncio async def test_async_post_call_failure_hook_does_not_clobber_guardrail_info_in_metadata(): - logger = _ProxyDBLogger() + logger = ProxyDBLogger() metadata_bucket_info = [{"guardrail_name": "from-metadata-bucket"}] request_data = { "model": "gpt-4", @@ -173,7 +173,7 @@ async def test_async_post_call_failure_hook_carries_used_client_oauth_token_from and leave request_data["metadata"] to the caller's native metadata, so a failed request on those routes wrote a spend row whose used_client_oauth_token was null instead of the stamped value """ - logger = _ProxyDBLogger() + logger = ProxyDBLogger() request_data = { "model": "claude-sonnet-5", "custom_llm_provider": custom_llm_provider, @@ -217,7 +217,7 @@ async def test_async_post_call_failure_hook_never_lets_caller_metadata_set_used_ On /v1/messages and /v1/responses the request's own metadata field belongs to the caller, so a used_client_oauth_token they put there must never outrank the proxy's stamp or stand in for a missing one """ - logger = _ProxyDBLogger() + logger = ProxyDBLogger() request_data = { "model": "claude-sonnet-5", "custom_llm_provider": "anthropic", @@ -264,7 +264,7 @@ async def test_async_post_call_failure_hook_never_lets_caller_metadata_set_used_ async def test_async_post_call_failure_hook_reads_used_client_oauth_token_from_the_routes_stamped_bucket( request_route: str, metadata_buckets: dict, expected: bool | None ): - logger = _ProxyDBLogger() + logger = ProxyDBLogger() request_data = { "model": "claude-sonnet-5", "custom_llm_provider": "anthropic", @@ -297,7 +297,7 @@ async def test_async_post_call_failure_hook_bills_guardrail_cost_on_blocked_requ """LIT-5651: a request blocked by a guardrail never reaches the LLM, but the guardrail invocation itself is billed by the provider. The failure row must charge that cost against the key instead of recording zero spend.""" - logger = _ProxyDBLogger() + logger = ProxyDBLogger() request_data = { "model": "gpt-4", "messages": [{"role": "user", "content": "Hello"}], @@ -329,7 +329,7 @@ async def test_async_post_call_failure_hook_bills_guardrail_cost_on_blocked_requ @pytest.mark.asyncio async def test_async_post_call_failure_hook_adds_guardrail_cost_to_recovered_stream_cost(): - logger = _ProxyDBLogger() + logger = ProxyDBLogger() request_data = { "model": "gpt-4", "messages": [{"role": "user", "content": "Hello"}], @@ -359,7 +359,7 @@ async def test_async_post_call_failure_hook_adds_guardrail_cost_to_recovered_str @pytest.mark.asyncio async def test_async_post_call_failure_hook_non_llm_route(): # Setup - logger = _ProxyDBLogger() + logger = ProxyDBLogger() # Mock user_api_key_dict with a non-LLM route user_api_key_dict = UserAPIKeyAuth( @@ -403,7 +403,7 @@ async def test_async_post_call_failure_hook_non_llm_route(): @pytest.mark.asyncio async def test_async_post_call_failure_hook_releases_budget_reservation_before_route_skip(): - logger = _ProxyDBLogger() + logger = ProxyDBLogger() budget_reservation = {"reserved_cost": 0.5, "entries": []} user_api_key_dict = UserAPIKeyAuth( api_key="test_api_key", @@ -437,7 +437,7 @@ async def test_async_post_call_failure_hook_releases_budget_reservation_before_r @pytest.mark.asyncio async def test_should_continue_failure_tracking_when_budget_release_fails(): - logger = _ProxyDBLogger() + logger = ProxyDBLogger() budget_reservation = {"reserved_cost": 0.5, "entries": []} user_api_key_dict = UserAPIKeyAuth( api_key="test_api_key", @@ -491,7 +491,7 @@ async def test_should_continue_failure_tracking_when_budget_release_fails(): @pytest.mark.asyncio async def test_track_cost_callback_releases_budget_reservation_when_spend_tracking_skips(): - logger = _ProxyDBLogger() + logger = ProxyDBLogger() budget_reservation = {"reserved_cost": 0.5, "entries": []} user_api_key_auth = UserAPIKeyAuth(budget_reservation=budget_reservation) @@ -527,7 +527,7 @@ async def test_track_cost_callback_releases_budget_reservation_when_spend_tracki @pytest.mark.asyncio async def test_track_cost_callback_releases_budget_reservation_when_response_cost_missing(): - logger = _ProxyDBLogger() + logger = ProxyDBLogger() budget_reservation = {"reserved_cost": 0.5, "entries": []} user_api_key_auth = UserAPIKeyAuth(budget_reservation=budget_reservation) @@ -970,7 +970,7 @@ async def test_track_cost_callback_skips_when_no_standard_logging_object(): File operations have no model and no standard_logging_object. The callback should skip gracefully instead of raising. """ - logger = _ProxyDBLogger() + logger = ProxyDBLogger() kwargs = { "call_type": "afile_delete", @@ -1011,7 +1011,7 @@ async def test_track_cost_callback_defers_in_progress_background_interaction(): """ from litellm.types.interactions import InteractionsAPIResponse - logger = _ProxyDBLogger() + logger = ProxyDBLogger() kwargs = { "call_type": "acreate_interaction", @@ -1105,7 +1105,7 @@ async def test_track_cost_callback_charges_a_batch_once_and_only_when_final( # are gated, since creating a batch is its own billable request, and a retrieve that charges nothing hands its budget reservation back instead. """ - logger = _ProxyDBLogger() + logger = ProxyDBLogger() budget_reservation = None if charged else {"reserved_cost": 0.5, "entries": []} kwargs = _batch_retrieve_kwargs(call_type, reservation=budget_reservation) @@ -1170,7 +1170,7 @@ async def test_track_cost_callback_keeps_reservation_open_for_in_progress_backgr """ from litellm.types.interactions import InteractionsAPIResponse - logger = _ProxyDBLogger() + logger = ProxyDBLogger() reservation = {"reserved_cost": 0.05, "entries": [], "finalized": False} in_progress_response = InteractionsAPIResponse( id="interactions/bg-abc", @@ -1210,7 +1210,7 @@ async def test_track_cost_callback_releases_reservation_for_in_progress_interact monkeypatch.setattr(callback_module, "BACKGROUND_INTERACTION_COST_POLLING_ENABLED", False) - logger = _ProxyDBLogger() + logger = ProxyDBLogger() reservation = {"reserved_cost": 0.05, "entries": [], "finalized": False} in_progress_response = InteractionsAPIResponse( id="interactions/bg-abc", @@ -1256,7 +1256,7 @@ async def test_track_cost_callback_releases_reservation_for_unpollable_interacti """ from litellm.types.interactions import InteractionsAPIResponse - logger = _ProxyDBLogger() + logger = ProxyDBLogger() reservation = {"reserved_cost": 0.05, "entries": [], "finalized": False} terminal_response = InteractionsAPIResponse( id="interactions/bg-abc", @@ -1297,7 +1297,7 @@ async def test_track_cost_callback_alerts_when_an_interaction_that_produced_outp """ from litellm.types.interactions import InteractionsAPIResponse - logger = _ProxyDBLogger() + logger = ProxyDBLogger() reservation = {"reserved_cost": 0.05, "entries": [], "finalized": False} usageless_response = InteractionsAPIResponse( id="interactions/bg-abc", @@ -1333,7 +1333,7 @@ async def test_track_cost_callback_releases_reservation_for_interaction_without_ """ from litellm.types.interactions import InteractionsAPIResponse - logger = _ProxyDBLogger() + logger = ProxyDBLogger() reservation = {"reserved_cost": 0.05, "entries": [], "finalized": False} idless_response = InteractionsAPIResponse( id="", @@ -1379,7 +1379,7 @@ async def test_callback_handles_every_status_the_interactions_api_can_return(): released = set() for status in sorted(member.value for member in Status1): - logger = _ProxyDBLogger() + logger = ProxyDBLogger() reservation = {"reserved_cost": 0.05, "entries": [], "finalized": False} response = InteractionsAPIResponse( id="interactions/bg-abc", @@ -1425,7 +1425,7 @@ async def test_async_post_call_failure_hook_propagates_trace_id_from_logging_obj The failure hook should propagate this so the DB spend log's session_id matches the Langfuse trace_id. """ - logger = _ProxyDBLogger() + logger = ProxyDBLogger() user_api_key_dict = UserAPIKeyAuth( api_key="test_api_key", @@ -1494,7 +1494,7 @@ async def test_enrich_failure_metadata_with_team_alias(): "user_api_key_team_id": "test_team_id", "user_api_key_team_alias": None, } - result = await _ProxyDBLogger._enrich_failure_metadata_with_key_info(metadata) + result = await ProxyDBLogger._enrich_failure_metadata_with_key_info(metadata) assert result["user_api_key_team_alias"] == "my-team-alias" @@ -1536,7 +1536,7 @@ async def test_enrich_failure_metadata_with_full_key_lookup(): "user_api_key_org_id": None, "user_api_key_project_id": None, } - result = await _ProxyDBLogger._enrich_failure_metadata_with_key_info(metadata) + result = await ProxyDBLogger._enrich_failure_metadata_with_key_info(metadata) assert result["user_api_key_alias"] == "fetched-key-alias" assert result["user_api_key_user_id"] == "fetched-user-id" assert result["user_api_key_team_id"] == "fetched-team-id" @@ -1567,7 +1567,7 @@ async def test_enrich_failure_metadata_skips_when_team_alias_present(): "user_api_key_team_id": "test_team_id", "user_api_key_team_alias": "already-set", } - result = await _ProxyDBLogger._enrich_failure_metadata_with_key_info(metadata) + result = await ProxyDBLogger._enrich_failure_metadata_with_key_info(metadata) assert result["user_api_key_team_alias"] == "already-set" mock_get_key.assert_not_called() mock_get_team.assert_not_called() @@ -1589,7 +1589,7 @@ async def test_enrich_failure_metadata_skips_when_no_api_key(): "user_api_key_team_id": None, "user_api_key_team_alias": None, } - await _ProxyDBLogger._enrich_failure_metadata_with_key_info(metadata) + await ProxyDBLogger._enrich_failure_metadata_with_key_info(metadata) mock_get_key.assert_not_called() @@ -1629,7 +1629,7 @@ async def test_enrich_failure_metadata_keeps_captured_identity_when_not_resolvin "user_api_key_team_alias": None, "user_api_key_org_id": None, } - result = await _ProxyDBLogger._enrich_failure_metadata_with_key_info( + result = await ProxyDBLogger._enrich_failure_metadata_with_key_info( metadata, resolve_missing_key_identity=False ) @@ -1670,7 +1670,7 @@ async def test_enrich_failure_metadata_ignores_flag_when_alias_present(): "user_api_key_team_alias": None, "user_api_key_org_id": None, } - result = await _ProxyDBLogger._enrich_failure_metadata_with_key_info( + result = await ProxyDBLogger._enrich_failure_metadata_with_key_info( metadata, resolve_missing_key_identity=resolve ) mock_get_key.assert_not_called() @@ -1693,7 +1693,7 @@ async def test_track_cost_callback_reads_key_only_for_in_request_logs(call_type, identity persisted at create time. Every other call type still backfills from the key. """ - logger = _ProxyDBLogger() + logger = ProxyDBLogger() mock_key_obj = MagicMock() mock_key_obj.key_alias = "alias-assigned-later" @@ -1760,7 +1760,7 @@ async def test_async_post_call_failure_hook_enriches_auth_error_metadata(): UserAPIKeyAuth is created with only api_key set. The failure hook should look up the key and team from cache/DB to populate all missing fields. """ - logger = _ProxyDBLogger() + logger = ProxyDBLogger() # This is what auth_exception_handler creates for 401 errors user_api_key_dict = UserAPIKeyAuth( @@ -1817,7 +1817,7 @@ async def test_async_post_call_failure_hook_enriches_auth_error_metadata(): @pytest.mark.asyncio async def test_async_post_call_failure_hook_skips_the_key_lookup_when_the_failure_is_a_db_stall(): - logger = _ProxyDBLogger() + logger = ProxyDBLogger() user_api_key_dict = UserAPIKeyAuth(api_key="hashed_key") request_data = { "model": "gpt-5.6", @@ -1859,7 +1859,7 @@ async def test_async_post_call_failure_hook_skips_the_key_lookup_when_the_failur async def test_async_post_call_failure_hook_still_enriches_metadata_for_a_non_stall_failure(): """Only a DBLookupDeadlineExceeded skips the key lookup; a transport error from the provider call must still resolve the key's alias for the failure row.""" - logger = _ProxyDBLogger() + logger = ProxyDBLogger() user_api_key_dict = UserAPIKeyAuth(api_key="hashed_key") request_data = { "model": "gpt-5.6", @@ -1908,7 +1908,7 @@ async def test_async_post_call_failure_hook_enriches_missing_team_alias(): should look up the team from cache and populate user_api_key_team_alias in the spend log metadata written to the DB. """ - logger = _ProxyDBLogger() + logger = ProxyDBLogger() user_api_key_dict = UserAPIKeyAuth( api_key="test_api_key", @@ -1959,7 +1959,7 @@ async def test_track_cost_callback_skips_for_falsy_model_and_no_slo(model_value) Same bug as above but model can also be empty string (e.g. health check callbacks). The guard should catch all falsy model values when sl_object is missing. """ - logger = _ProxyDBLogger() + logger = ProxyDBLogger() kwargs = { "call_type": "acompletion", @@ -1996,7 +1996,7 @@ async def test_async_post_call_failure_hook_uses_actual_start_time(): """ from datetime import timedelta - logger = _ProxyDBLogger() + logger = ProxyDBLogger() user_api_key_dict = UserAPIKeyAuth( api_key="test_api_key", @@ -2055,7 +2055,7 @@ async def _invoke_failure_hook_with_raised_exception(): Returns the metadata dict that was forwarded to ``update_database`` so the caller can assert on its ``error_information`` payload. """ - logger = _ProxyDBLogger() + logger = ProxyDBLogger() user_api_key_dict = UserAPIKeyAuth( api_key="test_api_key", user_id="u", @@ -2135,7 +2135,7 @@ async def test_async_post_call_failure_hook_records_recovered_partial_spend(): """ from litellm.types.utils import Usage - logger = _ProxyDBLogger() + logger = ProxyDBLogger() user_api_key_dict = UserAPIKeyAuth(api_key="test_api_key", user_id="u", team_id="t") request_data = { @@ -2166,7 +2166,7 @@ async def test_track_cost_callback_enriches_user_id_for_mcp_style_metadata(): """MCP tool calls may only carry user_api_key; user/team rollups still need user_id.""" from litellm.proxy._types import UserAPIKeyAuth - logger = _ProxyDBLogger() + logger = ProxyDBLogger() key_obj = UserAPIKeyAuth( api_key="hashed-key", user_id="mcp-user@example.com", @@ -2235,7 +2235,7 @@ async def test_track_cost_callback_keeps_guardrail_cost_on_cache_hit(): guardrail's provider charge must still reach spend logs and budgets. The payload already prices the LLM share at 0 on a cache hit, so its response_cost is the guardrail cost alone and the callback must pass it through untouched.""" - logger = _ProxyDBLogger() + logger = ProxyDBLogger() kwargs = { "call_type": "acompletion", "model": "gpt-4o", @@ -2354,7 +2354,7 @@ async def test_track_cost_callback_logs_unauthenticated_pass_through_request(cal aretrieve_batch is included because CheckBatchCost's completed-batch cost event reaches this same callback with no attributable key/user/team when the batch was created with the master key or a team-less key.""" - logger = _ProxyDBLogger() + logger = ProxyDBLogger() kwargs = { "call_type": call_type, @@ -2425,7 +2425,7 @@ async def _groups_charged_by_the_callback(kwargs, deployments=None): The callback resolves ``proxy_logging_obj`` and the router by importing them off ``proxy_server`` inside its own body, so there is no seam to inject either through. """ - logger = _ProxyDBLogger() + logger = ProxyDBLogger() with ( patch( # test-quality-ok: callback imports proxy_logging_obj off proxy_server in its body, no seam "litellm.proxy.proxy_server.proxy_logging_obj" @@ -2606,7 +2606,7 @@ async def test_async_log_success_event_hands_the_sidecar_a_compact_event_and_ski producer = SpendEventProducer( address=address, on_unavailable="fallback", buffer_size=10, connect_timeout=1.0, fallback=_no_fallback ) - logger = _ProxyDBLogger(producer) + logger = ProxyDBLogger(producer) with ( patch( # test-quality-ok: the callback imports this from proxy_server inside its body, so there is no injection seam @@ -2642,7 +2642,7 @@ async def test_async_log_success_event_keeps_batch_retrieves_in_process(): connect_timeout=1.0, fallback=_no_fallback, ) - logger = _ProxyDBLogger(producer) + logger = ProxyDBLogger(producer) kwargs = {**_offload_kwargs(), "call_type": CallTypes.aretrieve_batch.value} completed_batch = LiteLLMBatch( id="batch_abc", @@ -2709,7 +2709,7 @@ async def test_sidecar_writes_the_same_spend_row_and_counters_as_the_in_process_ end_time = datetime(2026, 1, 1, 0, 0, 2) async def in_process() -> None: - await _ProxyDBLogger().async_log_success_event(_offload_kwargs(), _offload_response(), start_time, end_time) + await ProxyDBLogger().async_log_success_event(_offload_kwargs(), _offload_response(), start_time, end_time) async def via_sidecar() -> None: line = build_spend_event(_offload_kwargs(), _offload_response(), start_time, end_time, store_bodies=False) @@ -2753,7 +2753,7 @@ async def test_async_post_call_failure_hook_persists_no_raw_model_on_an_unknown_ raw_model: Final = "opus-4.6 Please summarize my medical records\nPatient has diabetes" writer: Final = MagicMock(spec=DBSpendUpdateWriter) writer.update_database = AsyncMock() - logger: Final = _ProxyDBLogger(spend_writer=lambda: writer) + logger: Final = ProxyDBLogger(spend_writer=lambda: writer) await logger.async_post_call_failure_hook( request_data={"model": raw_model, "messages": [{"role": "user", "content": "hi"}]}, @@ -2800,7 +2800,7 @@ def _spend_write_kwargs_with_metadata_value(metadata_value: object) -> dict: @pytest.mark.asyncio @pytest.mark.parametrize("log_level", [logging.WARNING, logging.DEBUG]) async def test_track_cost_callback_failure_alert_never_carries_request_metadata_values(log_level): - logger: Final = _ProxyDBLogger() + logger: Final = ProxyDBLogger() records: list[logging.LogRecord] = [] handler: Final = logging.Handler() handler.emit = records.append @@ -2862,7 +2862,7 @@ async def test_autonomous_llm_callback_persists_without_human_or_key(identity_fi new_callable=AsyncMock, return_value=False, ) as persist: - await _ProxyDBLogger()._PROXY_track_cost_callback( + await ProxyDBLogger()._PROXY_track_cost_callback( kwargs=kwargs, completion_response=ModelResponse(), start_time=datetime.now(), end_time=datetime.now() ) persist.assert_awaited_once() @@ -2889,7 +2889,7 @@ async def test_track_cost_callback_enqueue_emits_no_service_span(): # test-qual emitted; the flush that writes the queue emits its own table-named spans.""" from litellm.proxy.proxy_server import proxy_logging_obj - logger = _ProxyDBLogger() + logger = ProxyDBLogger() kwargs = { "model": "gpt-4", "call_type": "acompletion", diff --git a/tests/unit/proxy/hooks/test_rate_limiter_toctou.py b/tests/unit/proxy/hooks/test_rate_limiter_toctou.py index 23c717b0e3a..6a767a6486f 100644 --- a/tests/unit/proxy/hooks/test_rate_limiter_toctou.py +++ b/tests/unit/proxy/hooks/test_rate_limiter_toctou.py @@ -28,10 +28,10 @@ from litellm import DualCache, Router from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.hooks.batch_rate_limiter import BatchFileUsage from litellm.proxy.hooks.dynamic_rate_limiter_v3 import ( - _PROXY_DynamicRateLimitHandlerV3 as DynamicRateLimitHandler, + PROXY_DynamicRateLimitHandlerV3 as DynamicRateLimitHandler, ) from litellm.proxy.hooks.parallel_request_limiter_v3 import ( - _PROXY_MaxParallelRequestsHandler_v3, + PROXY_MaxParallelRequestsHandler_v3, ) from litellm.proxy.utils import InternalUsageCache, hash_token @@ -89,7 +89,7 @@ async def test_batch_limiter_concurrent_bypasses_tpm_via_toctou(): dual_cache = DualCache() internal_usage_cache = InternalUsageCache(dual_cache=dual_cache) - rate_limiter = _PROXY_MaxParallelRequestsHandler_v3( + rate_limiter = PROXY_MaxParallelRequestsHandler_v3( internal_usage_cache=internal_usage_cache ) batch_limiter = rate_limiter._get_batch_rate_limiter() @@ -142,7 +142,7 @@ async def test_batch_limiter_uses_atomic_check_and_increment(): """ dual_cache = DualCache() internal_usage_cache = InternalUsageCache(dual_cache=dual_cache) - rate_limiter = _PROXY_MaxParallelRequestsHandler_v3( + rate_limiter = PROXY_MaxParallelRequestsHandler_v3( internal_usage_cache=internal_usage_cache ) batch_limiter = rate_limiter._get_batch_rate_limiter() @@ -356,7 +356,7 @@ async def test_batch_zero_token_consumes_rpm_only(): """ dual_cache = DualCache() internal_usage_cache = InternalUsageCache(dual_cache=dual_cache) - rate_limiter = _PROXY_MaxParallelRequestsHandler_v3( + rate_limiter = PROXY_MaxParallelRequestsHandler_v3( internal_usage_cache=internal_usage_cache ) batch_limiter = rate_limiter._get_batch_rate_limiter() diff --git a/tests/unit/proxy/hooks/test_sensitive_data_routing.py b/tests/unit/proxy/hooks/test_sensitive_data_routing.py index 48a42e42c62..fbe97e2917a 100644 --- a/tests/unit/proxy/hooks/test_sensitive_data_routing.py +++ b/tests/unit/proxy/hooks/test_sensitive_data_routing.py @@ -24,7 +24,7 @@ from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.hooks.sensitive_data_routing import ( DEFAULT_SENSITIVE_ROUTING_TTL, SENSITIVE_ROUTING_CACHE_PREFIX, - _PROXY_SensitiveDataRoutingHandler, + PROXY_SensitiveDataRoutingHandler, ) from litellm.proxy.utils import InternalUsageCache @@ -48,7 +48,7 @@ class TestSensitiveDataRoutingHandler: @pytest.fixture def handler(self): cache = MockInternalUsageCache() - return _PROXY_SensitiveDataRoutingHandler(internal_usage_cache=cache) + return PROXY_SensitiveDataRoutingHandler(internal_usage_cache=cache) @pytest.fixture def user_api_key_dict(self): @@ -87,7 +87,7 @@ class TestSensitiveDataRoutingHandler: await super().async_set_cache(key, value, ttl=ttl, **kwargs) cache = TargetRecordingCache() - handler = _PROXY_SensitiveDataRoutingHandler(internal_usage_cache=cache) + handler = PROXY_SensitiveDataRoutingHandler(internal_usage_cache=cache) await handler.set_session_routing( session_id="s-1", model="on-premise-model", user_api_key_dict=user_api_key_dict, guardrail_name="g" ) @@ -320,7 +320,7 @@ class TestStickySessionRouting: @pytest.fixture def handler(self): cache = MockInternalUsageCache() - return _PROXY_SensitiveDataRoutingHandler(internal_usage_cache=cache) + return PROXY_SensitiveDataRoutingHandler(internal_usage_cache=cache) @pytest.fixture def user_api_key_dict(self): @@ -436,37 +436,37 @@ class TestCacheKeyAndTTL: def test_make_cache_key_format(self): cache = MockInternalUsageCache() - handler = _PROXY_SensitiveDataRoutingHandler(internal_usage_cache=cache) + handler = PROXY_SensitiveDataRoutingHandler(internal_usage_cache=cache) key = handler._make_cache_key("test-session-123", "hashed-key") assert key == "{sensitive_route:hashed-key:test-session-123}:model" def test_make_cache_key_is_tenant_scoped(self): cache = MockInternalUsageCache() - handler = _PROXY_SensitiveDataRoutingHandler(internal_usage_cache=cache) + handler = PROXY_SensitiveDataRoutingHandler(internal_usage_cache=cache) key_a = handler._make_cache_key("shared-session", "key-a") key_b = handler._make_cache_key("shared-session", "key-b") assert key_a != key_b def test_resolve_tenant_prefers_api_key(self): - tenant = _PROXY_SensitiveDataRoutingHandler._resolve_tenant( + tenant = PROXY_SensitiveDataRoutingHandler._resolve_tenant( UserAPIKeyAuth(api_key="hashed-key", user_id="alice") ) assert tenant == "hashed-key" def test_resolve_tenant_falls_back_to_jwt_principal(self): - tenant = _PROXY_SensitiveDataRoutingHandler._resolve_tenant( + tenant = PROXY_SensitiveDataRoutingHandler._resolve_tenant( UserAPIKeyAuth(api_key=None, user_id="alice", team_id="t1", org_id="o1") ) assert tenant == "user:alice|team:t1|org:o1" def test_resolve_tenant_distinguishes_keyless_principals(self): - tenant_a = _PROXY_SensitiveDataRoutingHandler._resolve_tenant(UserAPIKeyAuth(api_key=None, user_id="alice")) - tenant_b = _PROXY_SensitiveDataRoutingHandler._resolve_tenant(UserAPIKeyAuth(api_key=None, user_id="bob")) + tenant_a = PROXY_SensitiveDataRoutingHandler._resolve_tenant(UserAPIKeyAuth(api_key=None, user_id="alice")) + tenant_b = PROXY_SensitiveDataRoutingHandler._resolve_tenant(UserAPIKeyAuth(api_key=None, user_id="bob")) assert tenant_a != tenant_b def test_resolve_tenant_defaults_when_anonymous(self): - assert _PROXY_SensitiveDataRoutingHandler._resolve_tenant(None) == "default" - assert _PROXY_SensitiveDataRoutingHandler._resolve_tenant(UserAPIKeyAuth(api_key=None)) == "default" + assert PROXY_SensitiveDataRoutingHandler._resolve_tenant(None) == "default" + assert PROXY_SensitiveDataRoutingHandler._resolve_tenant(UserAPIKeyAuth(api_key=None)) == "default" class TestCustomGuardrailSessionIdExtraction: @@ -557,7 +557,7 @@ class TestRedisCache: cache = MockInternalUsageCache() mock_redis = AsyncMock() cache.dual_cache.redis_cache = mock_redis - return _PROXY_SensitiveDataRoutingHandler(internal_usage_cache=cache) + return PROXY_SensitiveDataRoutingHandler(internal_usage_cache=cache) @pytest.mark.asyncio async def test_get_routed_model_from_redis(self, handler_with_redis): @@ -652,7 +652,7 @@ class TestPreCallHookEdgeCases: @pytest.fixture def handler(self): cache = MockInternalUsageCache() - return _PROXY_SensitiveDataRoutingHandler(internal_usage_cache=cache) + return PROXY_SensitiveDataRoutingHandler(internal_usage_cache=cache) @pytest.fixture def user_api_key_dict(self): @@ -715,7 +715,7 @@ class TestProxyHandleSensitiveDataRouteException: @pytest.fixture def routing_hook(self): cache = MockInternalUsageCache() - return _PROXY_SensitiveDataRoutingHandler(internal_usage_cache=cache) + return PROXY_SensitiveDataRoutingHandler(internal_usage_cache=cache) @pytest.mark.asyncio async def test_sticky_routing_persists_override(self, proxy_logging, routing_hook): @@ -1022,7 +1022,7 @@ class _OpenBreakerRedis: @pytest.mark.asyncio async def test_an_open_circuit_breaker_keeps_session_routing_in_memory_without_a_warning(caplog): cache = DualCache(redis_cache=_OpenBreakerRedis()) # pyright: ignore[reportArgumentType] # duck-typed Redis double - handler = _PROXY_SensitiveDataRoutingHandler(internal_usage_cache=InternalUsageCache(cache)) + handler = PROXY_SensitiveDataRoutingHandler(internal_usage_cache=InternalUsageCache(cache)) caplog.clear() with caplog.at_level(logging.DEBUG, logger="LiteLLM Proxy"): diff --git a/tests/unit/proxy/hooks/test_tpm_concurrent.py b/tests/unit/proxy/hooks/test_tpm_concurrent.py index 42c1f489bdd..2d23b134433 100644 --- a/tests/unit/proxy/hooks/test_tpm_concurrent.py +++ b/tests/unit/proxy/hooks/test_tpm_concurrent.py @@ -27,7 +27,7 @@ from litellm.proxy.hooks.parallel_request_limiter_v3 import ( PROJECT_OTPM_DESCRIPTOR_KEY, RateLimitedModel, _AUDIO_BYTES_PER_TOKEN, - _PROXY_MaxParallelRequestsHandler_v3 as RateLimitHandler, + PROXY_MaxParallelRequestsHandler_v3 as RateLimitHandler, ) from litellm.proxy.hooks.parallel_request_limiter_v3 import ( _call_id_from_callback_kwargs, diff --git a/tests/unit/proxy/hooks/test_unit_test_max_model_budget_limiter.py b/tests/unit/proxy/hooks/test_unit_test_max_model_budget_limiter.py index f194e43c74a..dc9dac9a57b 100644 --- a/tests/unit/proxy/hooks/test_unit_test_max_model_budget_limiter.py +++ b/tests/unit/proxy/hooks/test_unit_test_max_model_budget_limiter.py @@ -14,7 +14,7 @@ from litellm.litellm_core_utils.duration_parser import duration_in_seconds from litellm.proxy._types import Litellm_EntityType from litellm.proxy.hooks.model_max_budget_limiter import ( _budget_model_candidates, - _PROXY_VirtualKeyModelMaxBudgetLimiter, + PROXY_VirtualKeyModelMaxBudgetLimiter, build_model_max_budget_usage, resolve_model_budget, ) @@ -27,7 +27,7 @@ from litellm.types.utils import BudgetConfig as GenericBudgetInfo @pytest.fixture def budget_limiter(): dual_cache = DualCache() - return _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) + return PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) # Test _budget_model_candidates @@ -462,7 +462,7 @@ async def test_async_log_success_event_pushes_redis_increments_when_redis_config """ dual_cache = DualCache() dual_cache.redis_cache = object() # truthy placeholder; push only checks is not None - limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) + limiter = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) model = "gpt-4" kwargs = { "standard_logging_object": { @@ -491,7 +491,7 @@ async def test_async_log_success_event_pushes_redis_increments_when_redis_config @pytest.mark.asyncio async def test_model_budget_limiter_initializes_redis_increment_queue_lock(): dual_cache = DualCache() - limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) + limiter = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) spend_key = "virtual_key_spend:test-key:gpt-4:1d" await limiter._increment_spend_in_current_window( @@ -653,7 +653,7 @@ async def test_logged_spend_is_visible_to_key_info_usage_and_enforcement(request actively blocked at 429 while reporting current_spend 0. """ dual_cache = DualCache() - limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) + limiter = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) key_hash = "vk-hash" model_max_budget = {"gpt-4": {"budget_limit": 1.0, "time_period": "1d"}} @@ -717,7 +717,7 @@ async def test_user_model_budget_is_tracked_and_enforced(): enforced, independently of any key-level budget. """ dual_cache = DualCache() - limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) + limiter = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) user_id = "user-1" user_model_max_budget = {"gpt-4": {"budget_limit": 1.0, "time_period": "1mo"}} @@ -765,7 +765,7 @@ async def test_user_model_budget_counter_is_separate_from_the_key_counter(): counters, so one request must charge each exactly once. """ dual_cache = DualCache() - limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) + limiter = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) model_max_budget = {"gpt-4": {"budget_limit": 10.0, "time_period": "1d"}} await limiter.async_log_success_event( @@ -793,7 +793,7 @@ async def test_two_models_on_one_key_do_not_share_a_budget_window(): per model: a shared start lets the shorter period restart the longer one. """ dual_cache = DualCache() - limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) + limiter = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) model_max_budget = { "gpt-4": {"budget_limit": 10.0, "time_period": "1d"}, "claude-3": {"budget_limit": 10.0, "time_period": "30d"}, @@ -824,7 +824,7 @@ async def test_two_models_on_one_key_do_not_share_a_budget_window(): @pytest.mark.asyncio async def test_no_increment_when_no_scope_budgets_the_model(): dual_cache = DualCache() - limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) + limiter = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) with patch.object(limiter, "_increment_spend_for_key", new_callable=AsyncMock) as mock_increment: await limiter.async_log_success_event( _success_kwargs( @@ -869,7 +869,7 @@ async def test_bedrock_traffic_charges_the_bare_family_name_budget(): the key went. """ dual_cache = DualCache() - limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) + limiter = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) key_hash = "vk-hash" model_max_budget = {"claude-opus-4-8": {"budget_limit": 1.0, "time_period": "18h"}} user_api_key = UserAPIKeyAuth(token=key_hash, model_max_budget=model_max_budget) @@ -916,7 +916,7 @@ async def test_user_model_budget_window_resets_when_the_period_elapses(): ) dual_cache = DualCache() - limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) + limiter = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) user_id = "user-1" user_model_max_budget = {"gpt-4": {"budget_limit": 1.0, "time_period": "1mo"}} spend_key = model_budget_spend_cache_key( @@ -976,7 +976,7 @@ async def test_a_zero_dollar_cap_blocks_the_model(): mean something. """ dual_cache = DualCache() - limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) + limiter = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) key = UserAPIKeyAuth( token="hash-zero", model_max_budget={"gpt-4": {"budget_limit": 0, "time_period": "1d"}}, @@ -1006,7 +1006,7 @@ async def test_spend_exactly_at_the_cap_is_refused(): (RouterBudgetLimiting, the key and team budget checks) uses `>=`. """ dual_cache = DualCache() - limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) + limiter = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) budget = {"gpt-4": {"budget_limit": 2.0, "time_period": "1d"}} key = UserAPIKeyAuth(token="hash-exact", model_max_budget=budget) @@ -1075,7 +1075,7 @@ async def test_one_malformed_scope_does_not_abort_the_other_scopes(): charged despite the user's entry being garbage. """ dual_cache = DualCache() - limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) + limiter = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) key_budget = {"gpt-4": {"budget_limit": 10.0, "time_period": "1d"}} await limiter.async_log_success_event( @@ -1107,7 +1107,7 @@ async def test_an_unusable_budget_entry_is_not_enforced_instead_of_raising(): cannot be keyed, so it cannot be enforced; the write path rejects these, so reaching here means config.yaml or a direct DB edit. """ - limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=DualCache()) + limiter = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=DualCache()) key = UserAPIKeyAuth( token="hash-malformed", model_max_budget={"gpt-4": {"budget_limit": "not-a-number", "time_period": "1d"}}, @@ -1151,7 +1151,7 @@ def test_a_malformed_specific_entry_does_not_hide_a_usable_family_budget(): async def test_a_malformed_specific_entry_still_enforces_the_family_budget(): """The fall-through has to reach enforcement, not just resolution.""" dual_cache = DualCache() - limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) + limiter = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) budget = { "openai/gpt-4": {"budget_limit": "not-a-number", "time_period": "1d"}, "gpt-4": {"budget_limit": 1.0, "time_period": "1d"}, @@ -1228,7 +1228,7 @@ async def test_a_pre_upgrade_counter_keyed_on_the_request_model_still_enforces(e configured-model key finds that counter empty and admits another full budget until the window expires. """ - limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=DualCache()) + limiter = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=DualCache()) model_max_budget = {"gpt-4": {"budget_limit": 10.0, "time_period": "1d"}} await limiter.dual_cache.async_set_cache(key=f"{prefix}:entity-1:openai/gpt-4:1d", value=25.0, ttl=86400) @@ -1264,7 +1264,7 @@ async def test_the_pre_upgrade_and_post_upgrade_counters_add_up_over_one_window( under-reports the window: 6 + 5 is over a cap of 10 that neither half reaches on its own. """ - limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=DualCache()) + limiter = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=DualCache()) model_max_budget = {"gpt-4": {"budget_limit": 10.0, "time_period": "1d"}} await limiter.dual_cache.async_set_cache( key="virtual_key_spend:entity-1:openai/gpt-4:1d", value=legacy_spend, ttl=86400 @@ -1293,7 +1293,7 @@ async def test_the_configured_model_counter_is_never_counted_twice(): added them without noticing would charge 12 against a cap of 10 and refuse a key that has spent 6. """ - limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=DualCache()) + limiter = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=DualCache()) await limiter.dual_cache.async_set_cache(key="virtual_key_spend:entity-1:gpt-4:1d", value=6.0, ttl=86400) assert ( @@ -1318,7 +1318,7 @@ async def test_the_pre_upgrade_counter_is_no_longer_read_a_window_after_start_up """ import litellm.proxy.hooks.model_max_budget_limiter as limiter_module - limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=DualCache()) + limiter = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=DualCache()) model_max_budget = {"gpt-4": {"budget_limit": 10.0, "time_period": "1d"}} await limiter.dual_cache.async_set_cache(key="virtual_key_spend:entity-1:openai/gpt-4:1d", value=25.0, ttl=86400) user_api_key = UserAPIKeyAuth(token="entity-1", model_max_budget=model_max_budget) @@ -1339,7 +1339,7 @@ async def test_the_user_scope_has_no_pre_upgrade_counter_to_carry(): Reading one would invent a counter no previous version ever wrote, which is the opposite of preserving one. """ - limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=DualCache()) + limiter = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=DualCache()) model_max_budget = {"gpt-4": {"budget_limit": 10.0, "time_period": "1d"}} await limiter.dual_cache.async_set_cache(key="user_model_spend:u1:openai/gpt-4:1d", value=25.0, ttl=86400) @@ -1414,8 +1414,8 @@ async def test_spend_logged_on_one_replica_is_enforced_and_reported_on_another() that local share, while the shared counter was already over the cap. """ shared_redis = _SharedFakeRedis() - replica_a = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=DualCache(redis_cache=shared_redis)) - replica_b = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=DualCache(redis_cache=shared_redis)) + replica_a = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=DualCache(redis_cache=shared_redis)) + replica_b = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=DualCache(redis_cache=shared_redis)) key_hash = "vk-shared" model_max_budget = {"gpt-4": {"budget_limit": 1.0, "time_period": "30d"}} user_api_key = UserAPIKeyAuth(token=key_hash, model_max_budget=model_max_budget) @@ -1436,7 +1436,7 @@ async def test_spend_logged_on_one_replica_is_enforced_and_reported_on_another() assert usage_on_b["gpt-4"]["current_spend"] == 1.25 # Control: a replica that never served this key reads the same total. - replica_c = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=DualCache(redis_cache=shared_redis)) + replica_c = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=DualCache(redis_cache=shared_redis)) with pytest.raises(litellm.BudgetExceededError): await replica_c.is_key_within_model_budget(user_api_key, "gpt-4") @@ -1459,7 +1459,7 @@ async def test_team_model_budget_is_shared_by_every_key_without_an_override(requ charge one team counter and are both refused once it is spent. """ dual_cache = DualCache() - limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) + limiter = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) team_model_max_budget = {"gpt-4": {"budget_limit": 1.0, "time_period": "1d"}} check = lambda: limiter.is_team_within_model_budget( team_id="team-1", @@ -1507,7 +1507,7 @@ async def test_key_override_replaces_the_team_cap_for_that_model(): the team counter. """ dual_cache = DualCache() - limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) + limiter = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) team_model_max_budget = {"gpt-4": {"budget_limit": 1.0, "time_period": "1d"}} key_model_max_budget = {"gpt-4": {"budget_limit": 5.0, "time_period": "1d"}} await dual_cache.async_set_cache(key="team_model_spend:team-1:gpt-4:1d", value=9.0) @@ -1540,7 +1540,7 @@ async def test_key_override_replaces_the_team_cap_for_that_model(): async def test_key_entry_for_another_model_does_not_lift_the_team_cap(): """A key override only covers the model it names; other models stay on the team counter.""" dual_cache = DualCache() - limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) + limiter = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) team_model_max_budget = {"gpt-4": {"budget_limit": 1.0, "time_period": "1d"}} key_model_max_budget = {"claude-3": {"budget_limit": 5.0, "time_period": "1d"}} @@ -1567,7 +1567,7 @@ async def test_key_entry_for_another_model_does_not_lift_the_team_cap(): @pytest.mark.asyncio async def test_team_budget_leaves_unconfigured_models_alone(): dual_cache = DualCache() - limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) + limiter = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) team_model_max_budget = {"gpt-4": {"budget_limit": 0.0, "time_period": "1d"}} assert ( @@ -1593,7 +1593,7 @@ async def test_team_budget_leaves_unconfigured_models_alone(): async def test_team_counters_are_isolated_by_team_model_and_window(): """Same model on two teams, and two models with different windows on one team, never share a counter.""" dual_cache = DualCache() - limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) + limiter = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) team_model_max_budget = { "gpt-4": {"budget_limit": 10.0, "time_period": "1d"}, "claude-3": {"budget_limit": 10.0, "time_period": "30d"}, @@ -1616,7 +1616,7 @@ async def test_team_counters_are_isolated_by_team_model_and_window(): @pytest.mark.asyncio async def test_malformed_team_entry_is_skipped_and_its_sibling_still_enforced(): - limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=DualCache()) + limiter = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=DualCache()) team_model_max_budget = { "gpt-4": {"budget_limit": "not-a-number", "time_period": "1d"}, "claude-3": {"budget_limit": 0.0, "time_period": "1d"}, @@ -1644,7 +1644,7 @@ async def test_malformed_team_entry_is_skipped_and_its_sibling_still_enforced(): async def test_malformed_key_entry_does_not_count_as_an_override(): """A key entry the limiter cannot enforce must not also switch the team cap off.""" dual_cache = DualCache() - limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) + limiter = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) team_model_max_budget = {"gpt-4": {"budget_limit": 1.0, "time_period": "1d"}} key_model_max_budget = {"gpt-4": {"budget_limit": "not-a-number", "time_period": "1d"}} @@ -1680,7 +1680,7 @@ async def test_malformed_key_entry_does_not_count_as_an_override(): async def test_key_entry_without_a_spend_cap_does_not_lift_the_team_cap(key_entry): """A key row that only rate-limits the model, or has no enforceable cap, leaves the team cap in force.""" dual_cache = DualCache() - limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) + limiter = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) team_model_max_budget = {"gpt-4": {"budget_limit": 1.0, "time_period": "1d"}} key_model_max_budget = {"gpt-4": key_entry} diff --git a/tests/unit/proxy/lens/test_endpoints.py b/tests/unit/proxy/lens/test_endpoints.py index 5b088462aed..fee53320802 100644 --- a/tests/unit/proxy/lens/test_endpoints.py +++ b/tests/unit/proxy/lens/test_endpoints.py @@ -1048,6 +1048,7 @@ async def test_service_status_uses_internal_auth_and_only_advertises_the_public_ result: Final = await service_connection(UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER)) assert result.url == "https://traces.example/lens-ingest" assert result.connected is connected + assert result.configured is True assert result.status.storage_ready is connected assert route.calls[0].request.headers["Authorization"] == "Bearer " + "x" * 32 assert "private storage details" not in result.model_dump_json() @@ -1057,6 +1058,23 @@ async def test_service_status_uses_internal_auth_and_only_advertises_the_public_ assert denied.value.status_code == 403 +@pytest.mark.asyncio +@pytest.mark.parametrize("url,configured", (("", False), ("http://lens", True))) +async def test_service_setup_distinguishes_missing_installation_from_incomplete_configuration( + monkeypatch: pytest.MonkeyPatch, url: str, configured: bool +) -> None: + from litellm.proxy.lens.endpoints import service_connection + + monkeypatch.setenv("LITELLM_LENS_URL", url) + monkeypatch.delenv("LITELLM_LENS_SERVICE_TOKEN", raising=False) + monkeypatch.setenv("LITELLM_RELEASE_TAG", "v1.2.3") + result: Final = await service_connection(UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER)) + assert result.configured is configured + assert result.connected is False + assert result.release == "v1.2.3" + assert result.status.storage_ready is False + + @pytest.mark.asyncio async def test_credential_snapshot_excludes_expired_keys_and_disables_caching(monkeypatch: pytest.MonkeyPatch) -> None: from unittest.mock import AsyncMock diff --git a/tests/unit/proxy/lens/test_feedback_endpoints.py b/tests/unit/proxy/lens/test_feedback_endpoints.py new file mode 100644 index 00000000000..3ee7c79f169 --- /dev/null +++ b/tests/unit/proxy/lens/test_feedback_endpoints.py @@ -0,0 +1,323 @@ +import hashlib +from collections.abc import Mapping, Sequence +from datetime import datetime, timedelta, timezone +from typing import Final + +import pytest +from fastapi import HTTPException +from pydantic import BaseModel, ValidationError + +from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth +from litellm.proxy.lens.feedback_endpoints import ( + FeedbackDeletion, + FeedbackSubmission, + FeedbackTarget, + delete_feedback, + feedback_summary, + read_feedback, + submit_feedback, +) +from litellm.proxy.lens.feedback_models import TraceFeedbackRequest +from litellm.proxy.lens.feedback_repository import FEEDBACK_TABLE, ClickHouseFeedbackStore, session_trace_id +from litellm.proxy.lens.models import TraceIdentity +from litellm.rust_bridge.trace.generated.models import ( + FeedbackRow, + FeedbackSummaryRow, + FeedbackTargetRow, + LensFeedbackParams, + LensFeedbackSummaryParams, + LensFeedbackTargetParams, +) +from litellm.rust_bridge.trace.queries import ReadQuery +from litellm.rust_bridge.trace.storage import ClickHouseStorage + +ADMIN: Final = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin") +OTHER_ADMIN: Final = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="other") +VIEWER: Final = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY, user_id="viewer") +INTERNAL: Final = UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER, user_id="dev") +TEAM_APP: Final = UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER, team_id="team-a", token="app-key") +SOLO_APP: Final = UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER, token="key-solo") +T0: Final = datetime(2026, 3, 1, 12, 0, tzinfo=timezone.utc) + + +def ref(team: str, key: str, trace: str) -> str: + return hashlib.sha256(f"{team}\0{key}\0{trace}".encode()).hexdigest().upper() + + +class FakeClickHouse(ClickHouseStorage): + """Mirrors lens_feedback: ReplacingMergeTree(UpdatedAt, IsDeleted) keyed by team, key, trace, author.""" + + def __init__(self, traces: Mapping[str, tuple[tuple[str, str], ...]]) -> None: + self.traces: Final = traces + self.rows: Final[list[Mapping[str, object]]] = [] # mutable-ok: stands in for the table + + async def insert_rows(self, table: str, rows: Sequence[Mapping[str, object]]) -> None: + assert table == FEEDBACK_TABLE + self.rows.extend(rows) + + def _visible( + self, team: str, key: str, access: LensFeedbackTargetParams | LensFeedbackParams | LensFeedbackSummaryParams + ) -> bool: + if access.all_teams: + return True + return team == access.team and (access.key_hash == "" or key == access.key_hash) + + def _latest(self) -> tuple[Mapping[str, object], ...]: + newest: Final = {} # mutable-ok: emulates FINAL collapse + for row in sorted(self.rows, key=lambda r: str(r["UpdatedAt"])): + newest[(row["TeamId"], row["ApiKeyHash"], row["TraceId"], row["Author"])] = row + return tuple(r for r in newest.values() if r["IsDeleted"] == 0) + + async def query(self, query: ReadQuery[BaseModel, object], parameters: BaseModel) -> tuple[object, ...]: # pyright: ignore[reportIncompatibleMethodOverride] # fake dispatches on the concrete params type + match parameters: + case LensFeedbackTargetParams(): + return tuple( + FeedbackTargetRow(team_id=team, key_hash=key, trace_ref=ref(team, key, parameters.trace_id)) + for team, key in self.traces.get(parameters.trace_id, ()) + if self._visible(team, key, parameters) + and parameters.trace_ref in ("", ref(team, key, parameters.trace_id)) + ) + case LensFeedbackParams(): + return tuple( + FeedbackRow( + trace_id=str(r["TraceId"]), + trace_ref=ref(str(r["TeamId"]), str(r["ApiKeyHash"]), str(r["TraceId"])), + author=str(r["Author"]), + score=int(str(r["Score"])), + comment=str(r["Comment"]), + created_at=str(r["CreatedAt"]), + updated_at=str(r["UpdatedAt"]), + ) + for r in self._latest() + if r["TraceId"] == parameters.trace_id + and ref(str(r["TeamId"]), str(r["ApiKeyHash"]), str(r["TraceId"])) == parameters.trace_ref + and self._visible(str(r["TeamId"]), str(r["ApiKeyHash"]), parameters) + ) + case LensFeedbackSummaryParams(): + live = tuple( + r + for r in self._latest() + if r["TraceId"] in parameters.trace_ids + and self._visible(str(r["TeamId"]), str(r["ApiKeyHash"]), parameters) + ) + keys = sorted({(str(r["TeamId"]), str(r["ApiKeyHash"]), str(r["TraceId"])) for r in live}) + return tuple( + FeedbackSummaryRow( + trace_id=trace, + trace_ref=ref(team, key, trace), + count=len(scores), + average=sum(scores) / len(scores), + lowest=min(scores), + ) + for team, key, trace in keys + for scores in [ + [ + int(str(r["Score"])) + for r in live + if (r["TeamId"], r["ApiKeyHash"], r["TraceId"]) == (team, key, trace) + ] + ] + ) + case _: + raise AssertionError(f"unexpected query {query.name}") + + +def store(**traces: tuple[tuple[str, str], ...]) -> ClickHouseFeedbackStore: + return ClickHouseFeedbackStore(FakeClickHouse(traces or {"t1": (("team-a", "key-a"),)})) + + +def submission(score: int, comment: str = "", **fields: str) -> FeedbackSubmission: + target: Final = {} if "trace_id" in fields or "session_id" in fields else {"trace_id": "t1"} + return FeedbackSubmission.model_validate({"score": score, "comment": comment, **target, **fields}) + + +@pytest.mark.asyncio +async def test_resubmitting_replaces_the_authors_feedback_and_keeps_other_authors() -> None: + feedback: Final = store() + first: Final = await submit_feedback(submission(3, "wrong file"), ADMIN, feedback, T0) + await submit_feedback(submission(9, "great"), OTHER_ADMIN, feedback, T0) + second: Final = await submit_feedback( + submission(7, "fixed after retry"), ADMIN, feedback, T0 + timedelta(minutes=5) + ) + + listed: Final = await read_feedback(FeedbackTarget(trace_id="t1"), VIEWER, feedback) + + assert (first.created_at, second.created_at, second.updated_at) == (T0, T0, T0 + timedelta(minutes=5)) + assert {f.author: f.created_at for f in listed.feedback}["admin"] == T0 + assert listed.trace_ref == ref("team-a", "key-a", "t1") + assert {(f.author, f.score, f.comment) for f in listed.feedback} == { + ("admin", 7, "fixed after retry"), + ("other", 9, "great"), + } + + +@pytest.mark.asyncio +async def test_session_id_resolves_to_the_trace_lens_derives_at_ingest() -> None: + session_trace: Final = session_trace_id("session-one") + feedback: Final = store(**{session_trace: (("team-a", "key-a"),)}) + + saved: Final = await submit_feedback(submission(4, session_id="session-one"), ADMIN, feedback, T0) + listed: Final = await read_feedback(FeedbackTarget(session_id="session-one"), ADMIN, feedback) + + assert saved.trace_id == session_trace + assert listed.trace_ref == ref("team-a", "key-a", session_trace) + assert [f.score for f in listed.feedback] == [4] + + +def test_session_trace_id_matches_the_rust_ingest_hash() -> None: + # Pinned in litellm-rust/crates/traces/tests/otlp.rs (session_capture_joins_native_logs_...). + assert session_trace_id("session-one") == "5fddf060372c8501dca4f331b9da882b" + + +@pytest.mark.asyncio +async def test_unknown_traces_are_not_found_and_write_nothing() -> None: + feedback: Final = store() + + with pytest.raises(HTTPException) as write: + await submit_feedback(submission(5, trace_id="missing"), ADMIN, feedback, T0) + with pytest.raises(HTTPException) as read: + await read_feedback(FeedbackTarget(trace_id="missing"), ADMIN, feedback) + + assert (write.value.status_code, read.value.status_code) == (404, 404) + assert isinstance(feedback.storage, FakeClickHouse) and feedback.storage.rows == [] + + +@pytest.mark.parametrize("score", (-1, 11)) +def test_scores_outside_zero_to_ten_are_rejected(score: int) -> None: + with pytest.raises(ValidationError): + submission(score) + + +@pytest.mark.parametrize("target", ({}, {"trace_id": "t1", "session_id": "s1"})) +def test_target_needs_exactly_one_of_trace_or_session(target: dict[str, str]) -> None: + with pytest.raises(ValidationError): + FeedbackTarget.model_validate(target) + + +@pytest.mark.asyncio +async def test_viewers_can_read_but_not_write_and_non_admins_cannot_read_in_lens() -> None: + feedback: Final = store() + await read_feedback(FeedbackTarget(trace_id="t1"), VIEWER, feedback) + + with pytest.raises(HTTPException) as write: + await submit_feedback(submission(5), VIEWER, feedback, T0) + with pytest.raises(HTTPException) as read: + await read_feedback(FeedbackTarget(trace_id="t1"), INTERNAL, feedback) + + assert (write.value.status_code, read.value.status_code) == (403, 403) + + +@pytest.mark.asyncio +async def test_tenant_comes_from_the_trace_and_author_defaults_to_the_caller() -> None: + feedback: Final = store() + saved: Final = await submit_feedback(submission(5), OTHER_ADMIN, feedback, T0) + + assert saved.author == "other" + assert isinstance(feedback.storage, FakeClickHouse) + assert {(r["TeamId"], r["ApiKeyHash"], r["Author"]) for r in feedback.storage.rows} == { + ("team-a", "key-a", "other") + } + with pytest.raises(ValidationError): + FeedbackSubmission.model_validate({"trace_id": "t1", "score": 5, "team_id": "someone-else"}) + + +@pytest.mark.asyncio +async def test_delete_hides_only_the_callers_feedback() -> None: + feedback: Final = store() + await submit_feedback(submission(2), ADMIN, feedback, T0) + await submit_feedback(submission(8), OTHER_ADMIN, feedback, T0) + + await delete_feedback(FeedbackDeletion(trace_id="t1"), ADMIN, feedback, T0 + timedelta(minutes=9)) + with pytest.raises(HTTPException) as missing: + await delete_feedback(FeedbackDeletion(trace_id="t1"), ADMIN, feedback, T0 + timedelta(minutes=9)) + + remaining: Final = await read_feedback(FeedbackTarget(trace_id="t1"), ADMIN, feedback) + assert missing.value.status_code == 404 + assert [f.author for f in remaining.feedback] == ["other"] + + +@pytest.mark.asyncio +async def test_summary_flags_rated_traces_and_leaves_unrated_ones_empty() -> None: + feedback: Final = store(t1=(("team-a", "key-a"),), t2=(("team-a", "key-a"),)) + await submit_feedback(submission(2), ADMIN, feedback, T0) + await submit_feedback(submission(8), OTHER_ADMIN, feedback, T0) + rated: Final = TraceIdentity(trace_id="t1", trace_ref=ref("team-a", "key-a", "t1")) + unrated: Final = TraceIdentity(trace_id="t2", trace_ref=ref("team-a", "key-a", "t2")) + + summaries: Final = await feedback_summary(TraceFeedbackRequest(traces=(rated, unrated)), VIEWER, feedback) + + assert {s.trace_id: (s.count, s.average, s.lowest) for s in summaries} == { + "t1": (2, 5.0, 2), + "t2": (0, None, None), + } + + +@pytest.mark.asyncio +async def test_a_trace_id_shared_by_two_keys_needs_its_trace_ref() -> None: + feedback: Final = store(t1=(("team-a", "key-a"), ("team-a", "key-b"))) + + with pytest.raises(HTTPException) as ambiguous: + await submit_feedback(submission(5), ADMIN, feedback, T0) + saved: Final = await submit_feedback( + submission(5, trace_id="t1", trace_ref=ref("team-a", "key-b", "t1")), ADMIN, feedback, T0 + ) + + assert ambiguous.value.status_code == 404 + assert saved.trace_ref == ref("team-a", "key-b", "t1") + + +@pytest.mark.asyncio +async def test_summary_without_trace_ref_reports_the_rated_trace_with_its_resolved_ref() -> None: + feedback: Final = store() + await submit_feedback(submission(6), ADMIN, feedback, T0) + + summaries: Final = await feedback_summary( + TraceFeedbackRequest(traces=(TraceIdentity(trace_id="t1"),)), VIEWER, feedback + ) + + assert [(s.trace_ref, s.count, s.lowest) for s in summaries] == [(ref("team-a", "key-a", "t1"), 1, 6)] + + +@pytest.mark.asyncio +async def test_an_app_key_records_its_end_users_feedback_on_its_own_teams_trace() -> None: + feedback: Final = store() + await submit_feedback(submission(2, "It ignored my file", user="customer-1"), TEAM_APP, feedback, T0) + await submit_feedback(submission(9, "Perfect", user="customer-2"), TEAM_APP, feedback, T0) + await submit_feedback( + submission(4, "Better after retry", user="customer-1"), TEAM_APP, feedback, T0 + timedelta(minutes=1) + ) + + listed: Final = await read_feedback(FeedbackTarget(trace_id="t1"), ADMIN, feedback) + + assert {(f.author, f.score, f.comment) for f in listed.feedback} == { + ("customer-1", 4, "Better after retry"), + ("customer-2", 9, "Perfect"), + } + + +@pytest.mark.asyncio +async def test_an_app_key_cannot_write_feedback_on_another_tenants_trace() -> None: + feedback: Final = store(t1=(("team-b", "key-b"),), solo=(("", "key-solo"),), other=(("", "key-other"),)) + + with pytest.raises(HTTPException) as other_team: + await submit_feedback(submission(5, user="customer-1"), TEAM_APP, feedback, T0) + with pytest.raises(HTTPException) as other_key: + await submit_feedback(submission(5, trace_id="other", user="customer-1"), SOLO_APP, feedback, T0) + saved: Final = await submit_feedback(submission(5, trace_id="solo", user="customer-1"), SOLO_APP, feedback, T0) + + assert (other_team.value.status_code, other_key.value.status_code) == (404, 404) + assert saved.author == "customer-1" + + +@pytest.mark.asyncio +async def test_an_app_can_remove_one_end_users_feedback() -> None: + feedback: Final = store() + await submit_feedback(submission(2, user="customer-1"), TEAM_APP, feedback, T0) + await submit_feedback(submission(9, user="customer-2"), TEAM_APP, feedback, T0) + + await delete_feedback( + FeedbackDeletion(trace_id="t1", user="customer-1"), TEAM_APP, feedback, T0 + timedelta(minutes=1) + ) + + listed: Final = await read_feedback(FeedbackTarget(trace_id="t1"), ADMIN, feedback) + assert [f.author for f in listed.feedback] == ["customer-2"] diff --git a/tests/unit/proxy/management_endpoints/scim/test_scim_key_deactivation.py b/tests/unit/proxy/management_endpoints/scim/test_scim_key_deactivation.py index d35a676a28b..0c3600cbcc7 100644 --- a/tests/unit/proxy/management_endpoints/scim/test_scim_key_deactivation.py +++ b/tests/unit/proxy/management_endpoints/scim/test_scim_key_deactivation.py @@ -73,7 +73,7 @@ async def test_set_user_keys_blocked_flips_state_and_invalidates_cache(): patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()), patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()), patch( - "litellm.proxy.management_endpoints.scim.scim_v2._delete_cache_key_object", + "litellm.proxy.management_endpoints.scim.scim_v2.delete_cache_key_object", AsyncMock(side_effect=fake_delete), ), ): @@ -104,7 +104,7 @@ async def test_set_user_keys_blocked_noop_when_no_matching_keys(): patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()), patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()), patch( - "litellm.proxy.management_endpoints.scim.scim_v2._delete_cache_key_object", + "litellm.proxy.management_endpoints.scim.scim_v2.delete_cache_key_object", AsyncMock(), ) as mocked_delete, ): @@ -139,7 +139,7 @@ async def test_set_user_keys_unblocked_skips_admin_blocked_keys(): patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()), patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()), patch( - "litellm.proxy.management_endpoints.scim.scim_v2._delete_cache_key_object", + "litellm.proxy.management_endpoints.scim.scim_v2.delete_cache_key_object", AsyncMock(side_effect=fake_delete), ), ): @@ -172,7 +172,7 @@ async def test_scim_delete_user_blocks_keys_before_deleting_user(): patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()), patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()), patch( - "litellm.proxy.management_endpoints.scim.scim_v2._delete_cache_key_object", + "litellm.proxy.management_endpoints.scim.scim_v2.delete_cache_key_object", AsyncMock(), ), ): @@ -220,7 +220,7 @@ async def test_scim_delete_user_clears_fk_referenced_rows_before_user_delete(): patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()), patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()), patch( - "litellm.proxy.management_endpoints.scim.scim_v2._delete_cache_key_object", + "litellm.proxy.management_endpoints.scim.scim_v2.delete_cache_key_object", AsyncMock(), ), ): @@ -294,7 +294,7 @@ async def test_scim_patch_user_active_false_blocks_keys(): AsyncMock(return_value=mock_scim_user), ), patch( - "litellm.proxy.management_endpoints.scim.scim_v2._delete_cache_key_object", + "litellm.proxy.management_endpoints.scim.scim_v2.delete_cache_key_object", AsyncMock(), ), ): @@ -354,7 +354,7 @@ async def test_scim_patch_user_active_true_unblocks_keys(): AsyncMock(return_value=mock_scim_user), ), patch( - "litellm.proxy.management_endpoints.scim.scim_v2._delete_cache_key_object", + "litellm.proxy.management_endpoints.scim.scim_v2.delete_cache_key_object", AsyncMock(), ), ): @@ -411,7 +411,7 @@ async def test_scim_patch_user_no_active_change_does_not_touch_keys(): AsyncMock(return_value=mock_scim_user), ), patch( - "litellm.proxy.management_endpoints.scim.scim_v2._delete_cache_key_object", + "litellm.proxy.management_endpoints.scim.scim_v2.delete_cache_key_object", AsyncMock(), ), ): @@ -474,7 +474,7 @@ async def test_scim_put_user_omitting_active_preserves_deactivated_state(): AsyncMock(return_value=mock_scim_user), ), patch( - "litellm.proxy.management_endpoints.scim.scim_v2._delete_cache_key_object", + "litellm.proxy.management_endpoints.scim.scim_v2.delete_cache_key_object", AsyncMock(), ), ): @@ -530,7 +530,7 @@ async def test_scim_put_user_explicit_active_false_blocks_keys(): AsyncMock(return_value=mock_scim_user), ), patch( - "litellm.proxy.management_endpoints.scim.scim_v2._delete_cache_key_object", + "litellm.proxy.management_endpoints.scim.scim_v2.delete_cache_key_object", AsyncMock(), ), ): diff --git a/tests/unit/proxy/management_endpoints/scim/test_scim_v2_endpoints.py b/tests/unit/proxy/management_endpoints/scim/test_scim_v2_endpoints.py index 62d77a00f25..88b28fffbd3 100644 --- a/tests/unit/proxy/management_endpoints/scim/test_scim_v2_endpoints.py +++ b/tests/unit/proxy/management_endpoints/scim/test_scim_v2_endpoints.py @@ -33,7 +33,7 @@ from litellm.proxy.management_endpoints.scim.scim_v2 import ( _handle_group_membership_changes, _handle_team_membership_changes, _parse_member_entries, - _premium_user_check, + premium_user_check, _process_group_patch_operations, _recompute_scim_member_roles, _resolve_group_member_ids, @@ -493,7 +493,7 @@ async def test_scim_create_user_respects_default_role_set_via_ui(mocker, monkeyp def scim_test_client(): """An in-process SCIM application with authorization dependencies bypassed.""" app = FastAPI() - app.dependency_overrides[_premium_user_check] = lambda: None + app.dependency_overrides[premium_user_check] = lambda: None app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN) app.include_router(scim_router) return AsyncClient(transport=ASGITransport(app=app), base_url="http://test") diff --git a/tests/unit/proxy/management_endpoints/test_auto_router_endpoints.py b/tests/unit/proxy/management_endpoints/test_auto_router_endpoints.py index 360fc5cc292..a6555faad98 100644 --- a/tests/unit/proxy/management_endpoints/test_auto_router_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_auto_router_endpoints.py @@ -3465,7 +3465,7 @@ async def test_member_billable_preview_checks_and_charges_destination_team( raise litellm.BudgetExceededError(current_cost=2, max_budget=1) checks: Final = AsyncMock(side_effect=check_and_tag) - monkeypatch.setattr(auth_module, "_run_centralized_common_checks", checks) + monkeypatch.setattr(auth_module, "run_centralized_common_checks", checks) http_request: Final = Request( { "type": "http", diff --git a/tests/unit/proxy/management_endpoints/test_cache_settings_endpoints.py b/tests/unit/proxy/management_endpoints/test_cache_settings_endpoints.py index 9a2dd914866..3721ab79bf1 100644 --- a/tests/unit/proxy/management_endpoints/test_cache_settings_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_cache_settings_endpoints.py @@ -201,11 +201,11 @@ async def test_update_cache_settings_persists_url_precedence(monkeypatch): mock_prisma.db.litellm_cacheconfig.upsert = AsyncMock() proxy_config = MagicMock() - proxy_config._encrypt_env_variables = MagicMock( + proxy_config.encrypt_env_variables = MagicMock( side_effect=lambda environment_variables: dict(environment_variables) ) - proxy_config._decrypt_db_variables = MagicMock(side_effect=lambda variables_dict: dict(variables_dict)) - proxy_config._init_cache = MagicMock() + proxy_config.decrypt_db_variables = MagicMock(side_effect=lambda variables_dict: dict(variables_dict)) + proxy_config.init_cache = MagicMock() proxy_config.switch_on_llm_response_caching = MagicMock() with ( @@ -229,7 +229,7 @@ async def test_update_cache_settings_persists_url_precedence(monkeypatch): litellm_changed_by=None, ) - persisted = proxy_config._encrypt_env_variables.call_args.kwargs["environment_variables"] + persisted = proxy_config.encrypt_env_variables.call_args.kwargs["environment_variables"] assert persisted["url"] == "redis://:pw@host:6379/1" assert persisted["namespace"] == "ns" assert "host" not in persisted @@ -237,7 +237,7 @@ async def test_update_cache_settings_persists_url_precedence(monkeypatch): assert "db" not in persisted assert "password" not in persisted - init_params = proxy_config._init_cache.call_args.kwargs["cache_params"] + init_params = proxy_config.init_cache.call_args.kwargs["cache_params"] assert "host" not in init_params assert init_params["url"] == "redis://:pw@host:6379/1" @@ -262,7 +262,7 @@ async def test_get_cache_settings_masks_password_bearing_url(): mock_prisma.db.litellm_cacheconfig.find_unique = AsyncMock(return_value=cache_row) proxy_config = MagicMock() - proxy_config._decrypt_db_variables = MagicMock(side_effect=lambda variables_dict: dict(variables_dict)) + proxy_config.decrypt_db_variables = MagicMock(side_effect=lambda variables_dict: dict(variables_dict)) with ( patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), @@ -505,11 +505,11 @@ async def test_update_cache_settings_emits_audit_log_when_enabled(monkeypatch): mock_prisma.db.litellm_cacheconfig.upsert = AsyncMock() proxy_config = MagicMock() - proxy_config._encrypt_env_variables = MagicMock( + proxy_config.encrypt_env_variables = MagicMock( side_effect=lambda environment_variables: dict(environment_variables) ) - proxy_config._decrypt_db_variables = MagicMock(side_effect=lambda variables_dict: dict(variables_dict)) - proxy_config._init_cache = MagicMock() + proxy_config.decrypt_db_variables = MagicMock(side_effect=lambda variables_dict: dict(variables_dict)) + proxy_config.init_cache = MagicMock() proxy_config.switch_on_llm_response_caching = MagicMock() audit_calls = [] @@ -575,11 +575,11 @@ async def test_update_cache_settings_no_audit_when_disabled(monkeypatch): mock_prisma.db.litellm_cacheconfig.upsert = AsyncMock() proxy_config = MagicMock() - proxy_config._encrypt_env_variables = MagicMock( + proxy_config.encrypt_env_variables = MagicMock( side_effect=lambda environment_variables: dict(environment_variables) ) - proxy_config._decrypt_db_variables = MagicMock(side_effect=lambda variables_dict: dict(variables_dict)) - proxy_config._init_cache = MagicMock() + proxy_config.decrypt_db_variables = MagicMock(side_effect=lambda variables_dict: dict(variables_dict)) + proxy_config.init_cache = MagicMock() proxy_config.switch_on_llm_response_caching = MagicMock() audit_calls = [] @@ -802,7 +802,7 @@ async def test_get_cache_settings_falls_back_to_redis_env(monkeypatch): mock_prisma.db.litellm_cacheconfig.find_unique = AsyncMock(return_value=None) proxy_config = MagicMock() - proxy_config._decrypt_db_variables = MagicMock(side_effect=lambda variables_dict: dict(variables_dict)) + proxy_config.decrypt_db_variables = MagicMock(side_effect=lambda variables_dict: dict(variables_dict)) with ( patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), @@ -829,7 +829,7 @@ async def test_get_cache_settings_redacts_password_with_marker(monkeypatch): mock_prisma = MagicMock() mock_prisma.db.litellm_cacheconfig.find_unique = AsyncMock(return_value=cache_row) proxy_config = MagicMock() - proxy_config._decrypt_db_variables = MagicMock(side_effect=lambda variables_dict: dict(variables_dict)) + proxy_config.decrypt_db_variables = MagicMock(side_effect=lambda variables_dict: dict(variables_dict)) with ( patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), @@ -858,7 +858,7 @@ async def test_get_cache_settings_url_mode_hides_env_discrete_fields(monkeypatch mock_prisma = MagicMock() mock_prisma.db.litellm_cacheconfig.find_unique = AsyncMock(return_value=cache_row) proxy_config = MagicMock() - proxy_config._decrypt_db_variables = MagicMock(side_effect=lambda variables_dict: dict(variables_dict)) + proxy_config.decrypt_db_variables = MagicMock(side_effect=lambda variables_dict: dict(variables_dict)) with ( patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), @@ -875,11 +875,11 @@ async def test_get_cache_settings_url_mode_hides_env_discrete_fields(monkeypatch def _mock_proxy_config_identity_crypto(): proxy_config = MagicMock() - proxy_config._encrypt_env_variables = MagicMock( + proxy_config.encrypt_env_variables = MagicMock( side_effect=lambda environment_variables: dict(environment_variables) ) - proxy_config._decrypt_db_variables = MagicMock(side_effect=lambda variables_dict: dict(variables_dict)) - proxy_config._init_cache = MagicMock() + proxy_config.decrypt_db_variables = MagicMock(side_effect=lambda variables_dict: dict(variables_dict)) + proxy_config.init_cache = MagicMock() proxy_config.switch_on_llm_response_caching = MagicMock() return proxy_config @@ -913,7 +913,7 @@ async def test_update_preserves_stored_password_on_redacted_resubmit(monkeypatch litellm_changed_by=None, ) - persisted = proxy_config._encrypt_env_variables.call_args.kwargs["environment_variables"] + persisted = proxy_config.encrypt_env_variables.call_args.kwargs["environment_variables"] assert persisted["host"] == "oldhost" assert persisted["namespace"] == "edited" assert persisted["password"] == "realpw" @@ -945,7 +945,7 @@ async def test_update_drops_env_sourced_redacted_secret(monkeypatch): litellm_changed_by=None, ) - persisted = proxy_config._encrypt_env_variables.call_args.kwargs["environment_variables"] + persisted = proxy_config.encrypt_env_variables.call_args.kwargs["environment_variables"] assert "password" not in persisted @@ -974,7 +974,7 @@ async def test_update_applies_new_password(monkeypatch): litellm_changed_by=None, ) - persisted = proxy_config._encrypt_env_variables.call_args.kwargs["environment_variables"] + persisted = proxy_config.encrypt_env_variables.call_args.kwargs["environment_variables"] assert persisted["password"] == "brandnewpw" @@ -992,7 +992,7 @@ async def test_test_cache_connection_survives_saved_lookup_failure(monkeypatch): # a client whose find_unique is not awaitable, so the saved read raises bad_prisma = MagicMock() proxy_config = MagicMock() - proxy_config._decrypt_db_variables = MagicMock(side_effect=lambda variables_dict: dict(variables_dict)) + proxy_config.decrypt_db_variables = MagicMock(side_effect=lambda variables_dict: dict(variables_dict)) cache_instance = MagicMock() cache_instance.cache = MagicMock() @@ -1029,7 +1029,7 @@ async def test_get_cache_settings_does_not_surface_non_display_env_credentials(m mock_prisma = MagicMock() mock_prisma.db.litellm_cacheconfig.find_unique = AsyncMock(return_value=None) proxy_config = MagicMock() - proxy_config._decrypt_db_variables = MagicMock(side_effect=lambda variables_dict: dict(variables_dict)) + proxy_config.decrypt_db_variables = MagicMock(side_effect=lambda variables_dict: dict(variables_dict)) with ( patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), @@ -1058,7 +1058,7 @@ async def test_test_cache_connection_does_not_log_plaintext_credentials(monkeypa mock_prisma = MagicMock() mock_prisma.db.litellm_cacheconfig.find_unique = AsyncMock(return_value=existing) proxy_config = MagicMock() - proxy_config._decrypt_db_variables = MagicMock(side_effect=lambda variables_dict: dict(variables_dict)) + proxy_config.decrypt_db_variables = MagicMock(side_effect=lambda variables_dict: dict(variables_dict)) cache_instance = MagicMock() cache_instance.cache = MagicMock() @@ -1096,7 +1096,7 @@ async def test_test_cache_connection_does_not_replay_saved_password_to_new_host( mock_prisma = MagicMock() mock_prisma.db.litellm_cacheconfig.find_unique = AsyncMock(return_value=existing) proxy_config = MagicMock() - proxy_config._decrypt_db_variables = MagicMock(side_effect=lambda variables_dict: dict(variables_dict)) + proxy_config.decrypt_db_variables = MagicMock(side_effect=lambda variables_dict: dict(variables_dict)) cache_instance = MagicMock() cache_instance.cache = MagicMock() diff --git a/tests/unit/proxy/management_endpoints/test_common_utils.py b/tests/unit/proxy/management_endpoints/test_common_utils.py index 67920c9c7fe..43c74286d75 100644 --- a/tests/unit/proxy/management_endpoints/test_common_utils.py +++ b/tests/unit/proxy/management_endpoints/test_common_utils.py @@ -28,12 +28,12 @@ from litellm.proxy._types import ( from litellm.proxy.management_endpoints.common_utils import ( _has_non_empty_value, _org_admin_can_invite_user, - _set_object_metadata_field, _team_admin_can_invite_user, - _update_metadata_fields, - _user_has_admin_privileges, - _user_has_admin_view, admin_can_invite_user, + set_object_metadata_field, + update_metadata_fields, + user_api_key_has_admin_view, + user_has_admin_privileges, ) from litellm.types.utils import BudgetConfig @@ -53,7 +53,7 @@ class TestUpdateMetadataFieldsEmptyCollections: guardrails by sending `guardrails: []`). """ - @patch("litellm.proxy.management_endpoints.common_utils._premium_user_check") + @patch("litellm.proxy.management_endpoints.common_utils.premium_user_check") def test_empty_list_does_not_trigger_premium_check(self, mock_premium_check): """Empty lists for premium fields must not trigger the premium check.""" updated_kv = { @@ -62,10 +62,10 @@ class TestUpdateMetadataFieldsEmptyCollections: "policies": [], "logging": [], } - _update_metadata_fields(updated_kv=updated_kv) + update_metadata_fields(updated_kv=updated_kv) mock_premium_check.assert_not_called() - @patch("litellm.proxy.management_endpoints.common_utils._premium_user_check") + @patch("litellm.proxy.management_endpoints.common_utils.premium_user_check") def test_empty_list_still_updates_metadata(self, mock_premium_check): """ Empty lists must still be moved into metadata so users can clear @@ -76,7 +76,7 @@ class TestUpdateMetadataFieldsEmptyCollections: "guardrails": [], "policies": [], } - _update_metadata_fields(updated_kv=updated_kv) + update_metadata_fields(updated_kv=updated_kv) # The fields should have been moved into metadata assert ( "guardrails" not in updated_kv @@ -85,17 +85,17 @@ class TestUpdateMetadataFieldsEmptyCollections: assert updated_kv["metadata"]["guardrails"] == [] assert updated_kv["metadata"]["policies"] == [] - @patch("litellm.proxy.management_endpoints.common_utils._premium_user_check") + @patch("litellm.proxy.management_endpoints.common_utils.premium_user_check") def test_empty_dict_does_not_trigger_premium_check(self, mock_premium_check): """Empty dicts for premium fields must not trigger the premium check.""" updated_kv = { "team_id": "test-team", "secret_manager_settings": {}, } - _update_metadata_fields(updated_kv=updated_kv) + update_metadata_fields(updated_kv=updated_kv) mock_premium_check.assert_not_called() - @patch("litellm.proxy.management_endpoints.common_utils._premium_user_check") + @patch("litellm.proxy.management_endpoints.common_utils.premium_user_check") def test_empty_dict_still_updates_metadata(self, mock_premium_check): """ Empty dicts must still be moved into metadata so users can clear @@ -105,13 +105,13 @@ class TestUpdateMetadataFieldsEmptyCollections: "team_id": "test-team", "secret_manager_settings": {}, } - _update_metadata_fields(updated_kv=updated_kv) + update_metadata_fields(updated_kv=updated_kv) assert ( "secret_manager_settings" not in updated_kv ), "secret_manager_settings should be popped from top-level" assert updated_kv["metadata"]["secret_manager_settings"] == {} - @patch("litellm.proxy.management_endpoints.common_utils._premium_user_check") + @patch("litellm.proxy.management_endpoints.common_utils.premium_user_check") def test_none_value_does_not_trigger_premium_check(self, mock_premium_check): """None values for premium fields should be silently ignored.""" updated_kv = { @@ -119,51 +119,51 @@ class TestUpdateMetadataFieldsEmptyCollections: "guardrails": None, "policies": None, } - _update_metadata_fields(updated_kv=updated_kv) + update_metadata_fields(updated_kv=updated_kv) mock_premium_check.assert_not_called() - @patch("litellm.proxy.management_endpoints.common_utils._premium_user_check") + @patch("litellm.proxy.management_endpoints.common_utils.premium_user_check") def test_absent_fields_do_not_trigger_premium_check(self, mock_premium_check): """Fields not present in the dict should not trigger premium check.""" updated_kv = { "team_id": "test-team", "team_alias": "example-team", } - _update_metadata_fields(updated_kv=updated_kv) + update_metadata_fields(updated_kv=updated_kv) mock_premium_check.assert_not_called() - @patch("litellm.proxy.management_endpoints.common_utils._premium_user_check") + @patch("litellm.proxy.management_endpoints.common_utils.premium_user_check") def test_non_empty_list_triggers_premium_check(self, mock_premium_check): """Non-empty lists for premium fields should trigger the premium check.""" updated_kv = { "team_id": "test-team", "guardrails": ["my-guardrail"], } - _update_metadata_fields(updated_kv=updated_kv) + update_metadata_fields(updated_kv=updated_kv) mock_premium_check.assert_called() - @patch("litellm.proxy.management_endpoints.common_utils._premium_user_check") + @patch("litellm.proxy.management_endpoints.common_utils.premium_user_check") def test_non_empty_value_triggers_premium_check(self, mock_premium_check): """Non-empty string values for premium fields should trigger the premium check.""" updated_kv = { "team_id": "test-team", "tags": ["production"], } - _update_metadata_fields(updated_kv=updated_kv) + update_metadata_fields(updated_kv=updated_kv) mock_premium_check.assert_called() - @patch("litellm.proxy.management_endpoints.common_utils._premium_user_check") + @patch("litellm.proxy.management_endpoints.common_utils.premium_user_check") def test_non_empty_list_updates_metadata(self, mock_premium_check): """Non-empty lists should be moved into metadata.""" updated_kv = { "team_id": "test-team", "guardrails": ["my-guardrail"], } - _update_metadata_fields(updated_kv=updated_kv) + update_metadata_fields(updated_kv=updated_kv) assert "guardrails" not in updated_kv assert updated_kv["metadata"]["guardrails"] == ["my-guardrail"] - @patch("litellm.proxy.management_endpoints.common_utils._premium_user_check") + @patch("litellm.proxy.management_endpoints.common_utils.premium_user_check") def test_false_boolean_does_not_trigger_premium_check(self, mock_premium_check): """ Regression #30285: /team/update sends disable_global_guardrails=False @@ -171,25 +171,25 @@ class TestUpdateMetadataFieldsEmptyCollections: premium check, so non-premium users are not wrongly 403'd. """ updated_kv = {"team_id": "test-team", "disable_global_guardrails": False} - _update_metadata_fields(updated_kv=updated_kv) + update_metadata_fields(updated_kv=updated_kv) mock_premium_check.assert_not_called() - @patch("litellm.proxy.management_endpoints.common_utils._premium_user_check") + @patch("litellm.proxy.management_endpoints.common_utils.premium_user_check") def test_false_boolean_still_updates_metadata(self, mock_premium_check): """A falsy boolean must still be moved into metadata so it persists.""" updated_kv = {"team_id": "test-team", "disable_global_guardrails": False} - _update_metadata_fields(updated_kv=updated_kv) + update_metadata_fields(updated_kv=updated_kv) assert "disable_global_guardrails" not in updated_kv assert updated_kv["metadata"]["disable_global_guardrails"] is False - @patch("litellm.proxy.management_endpoints.common_utils._premium_user_check") + @patch("litellm.proxy.management_endpoints.common_utils.premium_user_check") def test_true_boolean_triggers_premium_check(self, mock_premium_check): """Control: enabling the premium feature (True) still requires a license.""" updated_kv = {"team_id": "test-team", "disable_global_guardrails": True} - _update_metadata_fields(updated_kv=updated_kv) + update_metadata_fields(updated_kv=updated_kv) mock_premium_check.assert_called() - @patch("litellm.proxy.management_endpoints.common_utils._premium_user_check") + @patch("litellm.proxy.management_endpoints.common_utils.premium_user_check") def test_ui_typical_payload_does_not_trigger_premium_check( self, mock_premium_check ): @@ -208,7 +208,7 @@ class TestUpdateMetadataFieldsEmptyCollections: }, "policies": [], } - _update_metadata_fields(updated_kv=updated_kv) + update_metadata_fields(updated_kv=updated_kv) mock_premium_check.assert_not_called() @@ -228,7 +228,7 @@ class TestUserHasAdminView: """Parametrized test: admin roles return True, non-admin return False.""" mock_auth = MagicMock() mock_auth.user_role = user_role - assert _user_has_admin_view(mock_auth) == expected + assert user_api_key_has_admin_view(mock_auth) == expected def test_user_has_admin_view_with_user_api_key_auth(self): """Test with actual UserAPIKeyAuth object.""" @@ -242,8 +242,8 @@ class TestUserHasAdminView: api_key="sk-yyy", user_role=LitellmUserRoles.INTERNAL_USER, ) - assert _user_has_admin_view(auth_admin) is True - assert _user_has_admin_view(auth_user) is False + assert user_api_key_has_admin_view(auth_admin) is True + assert user_api_key_has_admin_view(auth_user) is False def test_published_enterprise_import_of_team_admin_check_still_answers(): @@ -384,7 +384,7 @@ class TestUserHasAdminPrivileges: api_key="sk-x", user_role=LitellmUserRoles.PROXY_ADMIN, ) - result = await _user_has_admin_privileges( + result = await user_has_admin_privileges( user_api_key_dict=auth, prisma_client=None, ) @@ -398,7 +398,7 @@ class TestUserHasAdminPrivileges: api_key="sk-x", user_role=LitellmUserRoles.INTERNAL_USER, ) - result = await _user_has_admin_privileges( + result = await user_has_admin_privileges( user_api_key_dict=auth, prisma_client=None, ) @@ -455,9 +455,9 @@ class TestSetObjectMetadataField: """Parametrized test: premium fields trigger _premium_user_check.""" team = LiteLLM_TeamTable(team_id="t1", metadata={}) with patch( - "litellm.proxy.management_endpoints.common_utils._premium_user_check" + "litellm.proxy.management_endpoints.common_utils.premium_user_check" ) as mock_premium: - _set_object_metadata_field(team, field_name, value) + set_object_metadata_field(team, field_name, value) if should_call_premium: mock_premium.assert_called_once() else: @@ -468,9 +468,9 @@ class TestSetObjectMetadataField: """Test initializes metadata dict when object has None.""" team = LiteLLM_TeamTable(team_id="t1", metadata=None) with patch( - "litellm.proxy.management_endpoints.common_utils._premium_user_check" + "litellm.proxy.management_endpoints.common_utils.premium_user_check" ): - _set_object_metadata_field(team, "model_rpm_limit", {"x": 1}) + set_object_metadata_field(team, "model_rpm_limit", {"x": 1}) assert team.metadata == {"model_rpm_limit": {"x": 1}} def test_mcp_rpm_limit_is_hoisted_into_metadata(self): @@ -492,11 +492,11 @@ class TestSetObjectMetadataField: data = SimpleNamespace(mcp_rpm_limit=mcp_rpm_limit) with patch( - "litellm.proxy.management_endpoints.common_utils._premium_user_check" + "litellm.proxy.management_endpoints.common_utils.premium_user_check" ): for field in LiteLLM_ManagementEndpoint_MetadataFields: if getattr(data, field, None) is not None: - _set_object_metadata_field(team, field, getattr(data, field)) + set_object_metadata_field(team, field, getattr(data, field)) assert team.metadata["mcp_rpm_limit"] == mcp_rpm_limit @@ -681,7 +681,7 @@ class TestCheckPassthroughRoutesCallerPermission: from pydantic import BaseModel from litellm.proxy.management_endpoints.common_utils import ( - _check_passthrough_routes_caller_permission, + check_passthrough_routes_caller_permission, ) class _RouteData(BaseModel): @@ -690,7 +690,7 @@ class TestCheckPassthroughRoutesCallerPermission: data = _RouteData(allowed_passthrough_routes=["/v1/foo"]) with pytest.raises(HTTPException) as exc_info: - _check_passthrough_routes_caller_permission(data, self._non_admin()) + check_passthrough_routes_caller_permission(data, self._non_admin()) assert exc_info.value.status_code == 403 assert exc_info.value.detail == { @@ -702,7 +702,7 @@ class TestCheckPassthroughRoutesCallerPermission: from pydantic import BaseModel from litellm.proxy.management_endpoints.common_utils import ( - _check_passthrough_routes_caller_permission, + check_passthrough_routes_caller_permission, ) class _RouteData(BaseModel): @@ -711,7 +711,7 @@ class TestCheckPassthroughRoutesCallerPermission: data = _RouteData(metadata={"allowed_passthrough_routes": ["/v1/foo"]}) with pytest.raises(HTTPException) as exc_info: - _check_passthrough_routes_caller_permission(data, self._non_admin()) + check_passthrough_routes_caller_permission(data, self._non_admin()) assert exc_info.value.detail == { "error": "Only proxy admins can set `metadata.allowed_passthrough_routes` on a key." @@ -721,13 +721,13 @@ class TestCheckPassthroughRoutesCallerPermission: from pydantic import BaseModel from litellm.proxy.management_endpoints.common_utils import ( - _check_passthrough_routes_caller_permission, + check_passthrough_routes_caller_permission, ) class _Bare(BaseModel): unrelated: str = "x" - assert _check_passthrough_routes_caller_permission(_Bare(), self._non_admin()) is None + assert check_passthrough_routes_caller_permission(_Bare(), self._non_admin()) is None @pytest.mark.parametrize( "kwargs, field", @@ -741,7 +741,7 @@ class TestCheckPassthroughRoutesCallerPermission: from pydantic import BaseModel from litellm.proxy.management_endpoints.common_utils import ( - _check_passthrough_routes_caller_permission, + check_passthrough_routes_caller_permission, ) class _RouteData(BaseModel): @@ -749,7 +749,7 @@ class TestCheckPassthroughRoutesCallerPermission: metadata: dict[str, object] | None = None with pytest.raises(HTTPException) as exc_info: - _check_passthrough_routes_caller_permission( + check_passthrough_routes_caller_permission( _RouteData.model_validate(kwargs), self._non_admin(), entity="team" ) @@ -777,10 +777,10 @@ class TestDeniedPassthroughRoutesCallerPermission: ids=["cleared", "replaced", "dropped-by-metadata-replace", "dropped-by-null-metadata"], ) def test_non_admin_cannot_change_an_existing_deny_list(self, kwargs: dict[str, object], field: str) -> None: - from litellm.proxy.management_endpoints.common_utils import _check_passthrough_routes_caller_permission + from litellm.proxy.management_endpoints.common_utils import check_passthrough_routes_caller_permission with pytest.raises(HTTPException) as exc_info: - _check_passthrough_routes_caller_permission( + check_passthrough_routes_caller_permission( _DenyRouteData.model_validate(kwargs), UserAPIKeyAuth(user_id="u1", user_role=LitellmUserRoles.INTERNAL_USER), existing_metadata=_EXISTING_DENY, @@ -799,28 +799,28 @@ class TestDeniedPassthroughRoutesCallerPermission: ids=["resent-top-level", "resent-in-metadata", "unrelated-field"], ) def test_non_admin_may_leave_an_existing_deny_list_unchanged(self, kwargs: dict[str, object]) -> None: - from litellm.proxy.management_endpoints.common_utils import _check_passthrough_routes_caller_permission + from litellm.proxy.management_endpoints.common_utils import check_passthrough_routes_caller_permission - _check_passthrough_routes_caller_permission( + check_passthrough_routes_caller_permission( _DenyRouteData.model_validate(kwargs), UserAPIKeyAuth(user_id="u1", user_role=LitellmUserRoles.INTERNAL_USER), existing_metadata=_EXISTING_DENY, ) def test_non_admin_may_send_null_metadata_when_no_deny_list_exists(self) -> None: - from litellm.proxy.management_endpoints.common_utils import _check_passthrough_routes_caller_permission + from litellm.proxy.management_endpoints.common_utils import check_passthrough_routes_caller_permission - _check_passthrough_routes_caller_permission( + check_passthrough_routes_caller_permission( _DenyRouteData(metadata=None), UserAPIKeyAuth(user_id="u1", user_role=LitellmUserRoles.INTERNAL_USER), existing_metadata={"team": "core"}, ) def test_malformed_metadata_deny_entries_are_rejected_even_for_proxy_admins(self) -> None: - from litellm.proxy.management_endpoints.common_utils import _check_passthrough_routes_caller_permission + from litellm.proxy.management_endpoints.common_utils import check_passthrough_routes_caller_permission with pytest.raises(HTTPException) as exc_info: - _check_passthrough_routes_caller_permission( + check_passthrough_routes_caller_permission( _DenyRouteData(metadata={"denied_passthrough_routes": [123, None]}), UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN), ) @@ -849,11 +849,11 @@ class TestCheckDisableGlobalGuardrailsCallerPermission: from fastapi import HTTPException from litellm.proxy.management_endpoints.common_utils import ( - _check_disable_global_guardrails_caller_permission, + check_disable_global_guardrails_caller_permission, ) with pytest.raises(HTTPException) as exc_info: - _check_disable_global_guardrails_caller_permission(True, None, self._non_admin()) + check_disable_global_guardrails_caller_permission(True, None, self._non_admin()) assert exc_info.value.status_code == 403 assert exc_info.value.detail == {"error": "Only proxy admins can set `disable_global_guardrails` on a key."} @@ -862,11 +862,11 @@ class TestCheckDisableGlobalGuardrailsCallerPermission: from fastapi import HTTPException from litellm.proxy.management_endpoints.common_utils import ( - _check_disable_global_guardrails_caller_permission, + check_disable_global_guardrails_caller_permission, ) with pytest.raises(HTTPException) as exc_info: - _check_disable_global_guardrails_caller_permission( + check_disable_global_guardrails_caller_permission( None, {"disable_global_guardrails": True}, self._non_admin() ) @@ -877,11 +877,11 @@ class TestCheckDisableGlobalGuardrailsCallerPermission: from fastapi import HTTPException from litellm.proxy.management_endpoints.common_utils import ( - _check_disable_global_guardrails_caller_permission, + check_disable_global_guardrails_caller_permission, ) with pytest.raises(HTTPException) as exc_info: - _check_disable_global_guardrails_caller_permission( + check_disable_global_guardrails_caller_permission( False, {"disable_global_guardrails": True}, self._non_admin() ) @@ -892,37 +892,37 @@ class TestCheckDisableGlobalGuardrailsCallerPermission: from fastapi import HTTPException from litellm.proxy.management_endpoints.common_utils import ( - _check_disable_global_guardrails_caller_permission, + check_disable_global_guardrails_caller_permission, ) with pytest.raises(HTTPException) as exc_info: - _check_disable_global_guardrails_caller_permission(True, None, self._non_admin(), entity="team") + check_disable_global_guardrails_caller_permission(True, None, self._non_admin(), entity="team") assert exc_info.value.detail == {"error": "Only proxy admins can set `disable_global_guardrails` on a team."} def test_false_and_absent_flag_do_not_raise(self): from litellm.proxy.management_endpoints.common_utils import ( - _check_disable_global_guardrails_caller_permission, + check_disable_global_guardrails_caller_permission, ) non_admin = self._non_admin() - assert _check_disable_global_guardrails_caller_permission(False, None, non_admin) is None - assert _check_disable_global_guardrails_caller_permission(None, None, non_admin) is None - assert _check_disable_global_guardrails_caller_permission(None, {}, non_admin) is None + assert check_disable_global_guardrails_caller_permission(False, None, non_admin) is None + assert check_disable_global_guardrails_caller_permission(None, None, non_admin) is None + assert check_disable_global_guardrails_caller_permission(None, {}, non_admin) is None assert ( - _check_disable_global_guardrails_caller_permission(None, {"disable_global_guardrails": False}, non_admin) + check_disable_global_guardrails_caller_permission(None, {"disable_global_guardrails": False}, non_admin) is None ) def test_unchanged_stored_flag_does_not_raise(self): """Re-sending a flag that is already stored is not an opt-out.""" from litellm.proxy.management_endpoints.common_utils import ( - _check_disable_global_guardrails_caller_permission, + check_disable_global_guardrails_caller_permission, ) non_admin = self._non_admin() assert ( - _check_disable_global_guardrails_caller_permission( + check_disable_global_guardrails_caller_permission( True, {"disable_global_guardrails": True}, non_admin, @@ -935,11 +935,11 @@ class TestCheckDisableGlobalGuardrailsCallerPermission: from fastapi import HTTPException from litellm.proxy.management_endpoints.common_utils import ( - _check_disable_global_guardrails_caller_permission, + check_disable_global_guardrails_caller_permission, ) with pytest.raises(HTTPException) as exc_info: - _check_disable_global_guardrails_caller_permission( + check_disable_global_guardrails_caller_permission( True, None, self._non_admin(), @@ -951,11 +951,11 @@ class TestCheckDisableGlobalGuardrailsCallerPermission: def test_proxy_admin_may_set_the_flag(self): from litellm.proxy.management_endpoints.common_utils import ( - _check_disable_global_guardrails_caller_permission, + check_disable_global_guardrails_caller_permission, ) assert ( - _check_disable_global_guardrails_caller_permission(True, {"disable_global_guardrails": True}, self._admin()) + check_disable_global_guardrails_caller_permission(True, {"disable_global_guardrails": True}, self._admin()) is None ) @@ -963,7 +963,7 @@ class TestCheckDisableGlobalGuardrailsCallerPermission: class TestTeamMemberHasPermission: def test_requires_caller_to_be_a_team_member(self): from litellm.proxy.management_endpoints.common_utils import ( - _team_member_has_permission, + team_member_has_permission, ) team = LiteLLM_TeamTable( @@ -974,7 +974,7 @@ class TestTeamMemberHasPermission: key = UserAPIKeyAuth( user_id="u1", api_key="sk-x", user_role=LitellmUserRoles.INTERNAL_USER ) - assert _team_member_has_permission(key, team, "/key/generate") is False + assert team_member_has_permission(key, team, "/key/generate") is False class TestUserHasAdminPrivilegesGuard: @@ -986,7 +986,7 @@ class TestUserHasAdminPrivilegesGuard: ) mock_get_user = AsyncMock(return_value=None) with patch("litellm.proxy.auth.auth_checks.get_user_object", mock_get_user): - result = await _user_has_admin_privileges( + result = await user_has_admin_privileges( user_api_key_dict=auth, prisma_client=None ) assert result is False @@ -1013,7 +1013,7 @@ class TestUserHasAdminPrivilegesGuard: ) mock_get_user = AsyncMock(return_value=user_obj) with patch("litellm.proxy.auth.auth_checks.get_user_object", mock_get_user): - result = await _user_has_admin_privileges( + result = await user_has_admin_privileges( user_api_key_dict=auth, prisma_client=MagicMock() ) assert result is True @@ -1104,9 +1104,9 @@ class TestSetObjectMetadataFieldPremiumArg: def test_premium_check_receives_the_field_name(self): team = LiteLLM_TeamTable(team_id="t1", metadata={}) with patch( - "litellm.proxy.management_endpoints.common_utils._premium_user_check" + "litellm.proxy.management_endpoints.common_utils.premium_user_check" ) as mock_premium: - _set_object_metadata_field(team, "guardrails", ["g1"]) + set_object_metadata_field(team, "guardrails", ["g1"]) mock_premium.assert_called_once_with("guardrails") @@ -1124,9 +1124,9 @@ class TestUpdateMetadataFieldMove: def test_set_premium_field_is_moved_into_metadata(self): updated_kv = {"guardrails": ["g1"]} with patch( - "litellm.proxy.management_endpoints.common_utils._premium_user_check" + "litellm.proxy.management_endpoints.common_utils.premium_user_check" ): - _update_metadata_fields(updated_kv) + update_metadata_fields(updated_kv) assert "guardrails" not in updated_kv assert updated_kv["metadata"]["guardrails"] == ["g1"] @@ -1171,7 +1171,7 @@ class TestUpdateMetadataFieldsPremiumCheck: """ @patch( - "litellm.proxy.management_endpoints.common_utils._premium_user_check", + "litellm.proxy.management_endpoints.common_utils.premium_user_check", side_effect=Exception("Should not be called"), ) def test_empty_policies_skips_premium_check(self, mock_check): @@ -1181,11 +1181,11 @@ class TestUpdateMetadataFieldsPremiumCheck: "team_alias": "my-team", "policies": [], } - _update_metadata_fields(updated_kv) + update_metadata_fields(updated_kv) mock_check.assert_not_called() @patch( - "litellm.proxy.management_endpoints.common_utils._premium_user_check", + "litellm.proxy.management_endpoints.common_utils.premium_user_check", side_effect=Exception("Should not be called"), ) def test_empty_guardrails_skips_premium_check(self, mock_check): @@ -1194,11 +1194,11 @@ class TestUpdateMetadataFieldsPremiumCheck: "team_id": "team-123", "guardrails": [], } - _update_metadata_fields(updated_kv) + update_metadata_fields(updated_kv) mock_check.assert_not_called() @patch( - "litellm.proxy.management_endpoints.common_utils._premium_user_check", + "litellm.proxy.management_endpoints.common_utils.premium_user_check", side_effect=Exception("Should not be called"), ) def test_empty_string_team_member_key_duration_skips_premium_check( @@ -1209,11 +1209,11 @@ class TestUpdateMetadataFieldsPremiumCheck: "team_id": "team-123", "team_member_key_duration": "", } - _update_metadata_fields(updated_kv) + update_metadata_fields(updated_kv) mock_check.assert_not_called() @patch( - "litellm.proxy.management_endpoints.common_utils._premium_user_check", + "litellm.proxy.management_endpoints.common_utils.premium_user_check", side_effect=Exception("Should not be called"), ) def test_full_ui_payload_with_empty_premium_fields_skips_premium_check( @@ -1231,11 +1231,11 @@ class TestUpdateMetadataFieldsPremiumCheck: "team_member_key_duration": "", "prompts": [], } - _update_metadata_fields(updated_kv) + update_metadata_fields(updated_kv) mock_check.assert_not_called() @patch( - "litellm.proxy.management_endpoints.common_utils._premium_user_check", + "litellm.proxy.management_endpoints.common_utils.premium_user_check", ) def test_non_empty_policies_triggers_premium_check(self, mock_check): """policies: ['real-policy'] SHOULD trigger premium user check.""" @@ -1243,11 +1243,11 @@ class TestUpdateMetadataFieldsPremiumCheck: "team_id": "team-123", "policies": ["real-policy"], } - _update_metadata_fields(updated_kv) + update_metadata_fields(updated_kv) mock_check.assert_called() @patch( - "litellm.proxy.management_endpoints.common_utils._premium_user_check", + "litellm.proxy.management_endpoints.common_utils.premium_user_check", ) def test_non_empty_guardrails_triggers_premium_check(self, mock_check): """guardrails: ['my-guardrail'] SHOULD trigger premium user check.""" @@ -1255,11 +1255,11 @@ class TestUpdateMetadataFieldsPremiumCheck: "team_id": "team-123", "guardrails": ["my-guardrail"], } - _update_metadata_fields(updated_kv) + update_metadata_fields(updated_kv) mock_check.assert_called() @patch( - "litellm.proxy.management_endpoints.common_utils._premium_user_check", + "litellm.proxy.management_endpoints.common_utils.premium_user_check", ) def test_non_empty_team_member_key_duration_triggers_premium_check( self, mock_check @@ -1269,7 +1269,7 @@ class TestUpdateMetadataFieldsPremiumCheck: "team_id": "team-123", "team_member_key_duration": "30d", } - _update_metadata_fields(updated_kv) + update_metadata_fields(updated_kv) mock_check.assert_called() diff --git a/tests/unit/proxy/management_endpoints/test_config_override_endpoints.py b/tests/unit/proxy/management_endpoints/test_config_override_endpoints.py index 49b0ed1b28a..8d736bfdbdd 100644 --- a/tests/unit/proxy/management_endpoints/test_config_override_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_config_override_endpoints.py @@ -14,7 +14,7 @@ from litellm.proxy.management_endpoints.config_override_endpoints import ( CYBERARK_ENV_VAR_MAPPING, HASHICORP_ENV_VAR_MAPPING, _build_field_schema, - _set_env_vars, + set_env_vars, ) from litellm.proxy.proxy_server import app from litellm.types.proxy.management_endpoints.config_overrides import ( @@ -54,6 +54,14 @@ def _make_mock_proxy_config(): k: v.replace("enc_", "") if isinstance(v, str) else v for k, v in d.items() } ) + cfg.encrypt_env_variables = MagicMock( + side_effect=lambda d: {k: f"enc_{v}" for k, v in d.items()} + ) + cfg.decrypt_db_variables = MagicMock( + side_effect=lambda d: { + k: v.replace("enc_", "") if isinstance(v, str) else v for k, v in d.items() + } + ) return cfg @@ -194,7 +202,7 @@ async def test_hashicorp_vault_crud_lifecycle(client, monkeypatch): # 10. _set_env_vars: empty string unsets monkeypatch.setenv("HCP_VAULT_TOKEN", "existing") - _set_env_vars({"vault_token": "", "vault_addr": "https://v.com"}) + set_env_vars({"vault_token": "", "vault_addr": "https://v.com"}) assert os.environ.get("HCP_VAULT_TOKEN") is None assert os.environ["HCP_VAULT_ADDR"] == "https://v.com" @@ -209,9 +217,9 @@ async def test_hashicorp_vault_crud_lifecycle(client, monkeypatch): monkeypatch.setenv("LITELLM_SALT_KEY", "sk-test-salt-key") pc = ProxyConfig() orig = {"vault_addr": "https://v.com", "vault_token": "secret"} - encrypted = pc._encrypt_env_variables(orig) + encrypted = pc.encrypt_env_variables(orig) assert all(encrypted[k] != orig[k] for k in orig) - decrypted = pc._decrypt_db_variables(encrypted) + decrypted = pc.decrypt_db_variables(encrypted) assert all(decrypted[k] == orig[k] for k in orig) finally: diff --git a/tests/unit/proxy/management_endpoints/test_coordination_redis_endpoints.py b/tests/unit/proxy/management_endpoints/test_coordination_redis_endpoints.py index 288d8847fbb..2ee7dd1364a 100644 --- a/tests/unit/proxy/management_endpoints/test_coordination_redis_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_coordination_redis_endpoints.py @@ -191,7 +191,7 @@ async def test_get_source_does_not_build_a_client(monkeypatch): with ( patch("litellm.proxy.proxy_server.prisma_client", _prisma_with_general_settings({})), patch("litellm.proxy.proxy_server.proxy_config", _proxy_config()), - patch("litellm.proxy.proxy_server._build_redis_usage_cache") as mock_build, + patch("litellm.proxy.proxy_server.build_redis_usage_cache") as mock_build, ): response = await get_coordination_redis_settings(user_api_key_dict=_admin_auth()) @@ -480,7 +480,7 @@ async def test_connection_test_returns_healthy_on_successful_ping(): with ( patch("litellm.proxy.proxy_server.prisma_client", _prisma_with_general_settings({})), patch("litellm.proxy.proxy_server.proxy_config", _proxy_config()), - patch("litellm.proxy.proxy_server._build_redis_usage_cache", return_value=mock_client) as mock_build, + patch("litellm.proxy.proxy_server.build_redis_usage_cache", return_value=mock_client) as mock_build, ): response = await check_coordination_redis_connection( request=CoordinationRedisSettingsRequest( @@ -509,7 +509,7 @@ async def test_connection_test_reports_unhealthy_without_leaking_the_password(): with ( patch("litellm.proxy.proxy_server.prisma_client", _prisma_with_general_settings({})), patch("litellm.proxy.proxy_server.proxy_config", _proxy_config()), - patch("litellm.proxy.proxy_server._build_redis_usage_cache", return_value=mock_client), + patch("litellm.proxy.proxy_server.build_redis_usage_cache", return_value=mock_client), ): response = await check_coordination_redis_connection( request=CoordinationRedisSettingsRequest( @@ -543,7 +543,7 @@ async def test_connection_test_uses_the_saved_password_for_a_redacted_field(): _prisma_with_general_settings({"coordination_redis": _SAVED_SETTINGS}), ), patch("litellm.proxy.proxy_server.proxy_config", _proxy_config()), - patch("litellm.proxy.proxy_server._build_redis_usage_cache", return_value=mock_client) as mock_build, + patch("litellm.proxy.proxy_server.build_redis_usage_cache", return_value=mock_client) as mock_build, ): response = await check_coordination_redis_connection( request=CoordinationRedisSettingsRequest( @@ -568,7 +568,7 @@ async def test_connection_test_times_out_instead_of_hanging(): with ( patch("litellm.proxy.proxy_server.prisma_client", _prisma_with_general_settings({})), patch("litellm.proxy.proxy_server.proxy_config", _proxy_config()), - patch("litellm.proxy.proxy_server._build_redis_usage_cache", return_value=mock_client), + patch("litellm.proxy.proxy_server.build_redis_usage_cache", return_value=mock_client), patch( "litellm.proxy.management_endpoints.coordination_redis_endpoints._PING_TIMEOUT_SECONDS", 0.01, diff --git a/tests/unit/proxy/management_endpoints/test_credential_migration.py b/tests/unit/proxy/management_endpoints/test_credential_migration.py index ac5a45499d3..ffb855173fa 100644 --- a/tests/unit/proxy/management_endpoints/test_credential_migration.py +++ b/tests/unit/proxy/management_endpoints/test_credential_migration.py @@ -17,7 +17,7 @@ import pytest from litellm._service_logger import ServiceTypes from litellm.proxy import proxy_server from litellm.proxy.common_utils.encrypt_decrypt_utils import ( - _V2_GCM_PREFIX, + V2_GCM_PREFIX, encrypt_value_helper, ) from litellm.proxy.management_endpoints import credential_migration as cm @@ -82,7 +82,7 @@ def test_reencrypt_value_legacy_to_v2(salt_key, monkeypatch): out = cm.reencrypt_value(legacy) assert out != legacy - assert out.startswith(_V2_GCM_PREFIX) + assert out.startswith(V2_GCM_PREFIX) def test_reencrypt_value_is_idempotent(salt_key, monkeypatch): @@ -113,7 +113,7 @@ def test_reencrypt_selective_dict(salt_key, monkeypatch): data = {"api_key": legacy_key, "base_url": "https://x", "integration_token": None} out = cm.reencrypt_selective_dict(data, ["api_key", "integration_token"]) - assert out["api_key"].startswith(_V2_GCM_PREFIX) + assert out["api_key"].startswith(V2_GCM_PREFIX) assert out["base_url"] == "https://x" # untouched non-sensitive assert out["integration_token"] is None # null skipped @@ -164,7 +164,7 @@ async def test_vantage_walker_migrates_legacy_field(salt_key, monkeypatch): written = json.loads( client.db.litellm_config.update.call_args.kwargs["data"]["param_value"] ) - assert written["api_key"].startswith(_V2_GCM_PREFIX) + assert written["api_key"].startswith(V2_GCM_PREFIX) assert written["base_url"] == "https://api.vantage.sh" # non-sensitive untouched @@ -584,11 +584,11 @@ async def test_migrate_covered_tables_reports_real_counts(salt_key, monkeypatch) client.db.litellm_config.find_unique = AsyncMock(return_value=None) async def fake_rotate(**kwargs): - # Stand in for _rotate_master_key: re-encrypt the model api_key in place. + # Stand in for rotate_master_key: re-encrypt the model api_key in place. row.litellm_params["api_key"] = v2 monkeypatch.setattr( - "litellm.proxy.management_endpoints.key_management_endpoints._rotate_master_key", + "litellm.proxy.management_endpoints.key_management_endpoints.rotate_master_key", fake_rotate, ) diff --git a/tests/unit/proxy/management_endpoints/test_delete_verification_tokens_failed.py b/tests/unit/proxy/management_endpoints/test_delete_verification_tokens_failed.py index e33945df7dc..bfd0d663eac 100644 --- a/tests/unit/proxy/management_endpoints/test_delete_verification_tokens_failed.py +++ b/tests/unit/proxy/management_endpoints/test_delete_verification_tokens_failed.py @@ -94,7 +94,7 @@ async def test_delete_all_tokens_admin_returns_empty_failed_tokens(monkeypatch): mock_cache.delete_cache = MagicMock() monkeypatch.setattr( - "litellm.proxy.management_endpoints.key_management_endpoints._hash_token_if_needed", + "litellm.proxy.management_endpoints.key_management_endpoints.hash_token_if_needed", lambda token: token, ) monkeypatch.setattr( @@ -132,7 +132,7 @@ async def test_delete_tokens_non_admin_all_succeed_returns_empty_failed_tokens( mock_cache.delete_cache = MagicMock() monkeypatch.setattr( - "litellm.proxy.management_endpoints.key_management_endpoints._hash_token_if_needed", + "litellm.proxy.management_endpoints.key_management_endpoints.hash_token_if_needed", lambda token: token, ) monkeypatch.setattr( @@ -183,7 +183,7 @@ async def test_delete_tokens_non_admin_token_not_in_db_returns_failed_tokens( mock_cache.delete_cache = MagicMock() monkeypatch.setattr( - "litellm.proxy.management_endpoints.key_management_endpoints._hash_token_if_needed", + "litellm.proxy.management_endpoints.key_management_endpoints.hash_token_if_needed", lambda token: token, ) monkeypatch.setattr( @@ -234,7 +234,7 @@ async def test_delete_tokens_admin_partial_db_failure_returns_failed_tokens( mock_cache.delete_cache = MagicMock() monkeypatch.setattr( - "litellm.proxy.management_endpoints.key_management_endpoints._hash_token_if_needed", + "litellm.proxy.management_endpoints.key_management_endpoints.hash_token_if_needed", lambda token: token, ) monkeypatch.setattr( diff --git a/tests/unit/proxy/management_endpoints/test_internal_user_endpoints.py b/tests/unit/proxy/management_endpoints/test_internal_user_endpoints.py index 712107e4a32..39265b37b95 100644 --- a/tests/unit/proxy/management_endpoints/test_internal_user_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_internal_user_endpoints.py @@ -4004,7 +4004,7 @@ async def test_admin_user_update_spend_invalidates_counter(mocker): mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) mocker.patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin") mock_invalidate = mocker.patch( - "litellm.proxy.proxy_server._invalidate_spend_counter", + "litellm.proxy.proxy_server.invalidate_spend_counter", new=mocker.AsyncMock(), ) @@ -4038,7 +4038,7 @@ async def test_user_update_rejects_non_finite_spend(mocker): mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) mocker.patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin") mock_invalidate = mocker.patch( - "litellm.proxy.proxy_server._invalidate_spend_counter", + "litellm.proxy.proxy_server.invalidate_spend_counter", new=mocker.AsyncMock(), ) @@ -4276,7 +4276,7 @@ def _object_permission_mocks(mocker, existing_object_permission_id=None): mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) mocker.patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin") mocker.patch( - "litellm.proxy.proxy_server._invalidate_spend_counter", + "litellm.proxy.proxy_server.invalidate_spend_counter", new=mocker.AsyncMock(), ) return mock_prisma_client diff --git a/tests/unit/proxy/management_endpoints/test_key_generate_prisma.py b/tests/unit/proxy/management_endpoints/test_key_generate_prisma.py index a4126c476bd..3d9667ace1b 100644 --- a/tests/unit/proxy/management_endpoints/test_key_generate_prisma.py +++ b/tests/unit/proxy/management_endpoints/test_key_generate_prisma.py @@ -514,9 +514,9 @@ def test_call_with_user_over_budget(prisma_client): # update spend using track_cost callback, make 2nd request, it should fail from litellm import Choices, Message, ModelResponse, Usage - from litellm.proxy.proxy_server import _ProxyDBLogger + from litellm.proxy.proxy_server import ProxyDBLogger - proxy_db_logger = _ProxyDBLogger() + proxy_db_logger = ProxyDBLogger() resp = ModelResponse( id="chatcmpl-e41836bb-bb8b-4df2-8e70-8f3e160155ac", @@ -612,9 +612,9 @@ def test_call_with_end_user_over_budget(prisma_client): # update spend using track_cost callback, make 2nd request, it should fail from litellm import Choices, Message, ModelResponse, Usage - from litellm.proxy.proxy_server import _ProxyDBLogger + from litellm.proxy.proxy_server import ProxyDBLogger - proxy_db_logger = _ProxyDBLogger() + proxy_db_logger = ProxyDBLogger() resp = ModelResponse( id="chatcmpl-e41836bb-bb8b-4df2-8e70-8f3e160155ac", @@ -722,9 +722,9 @@ def test_call_with_proxy_over_budget(prisma_client): # update spend using track_cost callback, make 2nd request, it should fail from litellm import Choices, Message, ModelResponse, Usage - from litellm.proxy.proxy_server import _ProxyDBLogger + from litellm.proxy.proxy_server import ProxyDBLogger - proxy_db_logger = _ProxyDBLogger() + proxy_db_logger = ProxyDBLogger() resp = ModelResponse( id="chatcmpl-e41836bb-bb8b-4df2-8e70-8f3e160155ac", @@ -814,9 +814,9 @@ def test_call_with_user_over_budget_stream(prisma_client): # update spend using track_cost callback, make 2nd request, it should fail from litellm import Choices, Message, ModelResponse, Usage - from litellm.proxy.proxy_server import _ProxyDBLogger + from litellm.proxy.proxy_server import ProxyDBLogger - proxy_db_logger = _ProxyDBLogger() + proxy_db_logger = ProxyDBLogger() resp = ModelResponse( id="chatcmpl-e41836bb-bb8b-4df2-8e70-8f3e160155ac", @@ -921,9 +921,9 @@ def test_call_with_proxy_over_budget_stream(prisma_client): # update spend using track_cost callback, make 2nd request, it should fail from litellm import Choices, Message, ModelResponse, Usage - from litellm.proxy.proxy_server import _ProxyDBLogger + from litellm.proxy.proxy_server import ProxyDBLogger - proxy_db_logger = _ProxyDBLogger() + proxy_db_logger = ProxyDBLogger() resp = ModelResponse( id="chatcmpl-e41836bb-bb8b-4df2-8e70-8f3e160155ac", @@ -1573,9 +1573,9 @@ def test_call_with_key_over_budget(prisma_client): # update spend using track_cost callback, make 2nd request, it should fail from litellm import Choices, Message, ModelResponse, Usage from litellm.caching.caching import Cache - from litellm.proxy.proxy_server import _ProxyDBLogger + from litellm.proxy.proxy_server import ProxyDBLogger - proxy_db_logger = _ProxyDBLogger() + proxy_db_logger = ProxyDBLogger() litellm.cache = Cache() import time @@ -1690,7 +1690,7 @@ def test_call_with_key_over_budget_no_cache(prisma_client): print("result from user auth with new key", result) # update spend using track_cost callback, make 2nd request, it should fail - from litellm.proxy.proxy_server import _ProxyDBLogger + from litellm.proxy.proxy_server import ProxyDBLogger from litellm.proxy.proxy_server import user_api_key_cache user_api_key_cache.in_memory_cache.cache_dict = {} @@ -1720,7 +1720,7 @@ def test_call_with_key_over_budget_no_cache(prisma_client): model="gpt-35-turbo", # azure always has model written like this usage=Usage(prompt_tokens=210, completion_tokens=200, total_tokens=410), ) - proxy_db_logger = _ProxyDBLogger() + proxy_db_logger = ProxyDBLogger() await proxy_db_logger._PROXY_track_cost_callback( kwargs={ "model": "chatgpt-v-3", @@ -1943,9 +1943,9 @@ async def test_call_with_key_never_over_budget(prisma_client): from litellm._uuid import uuid from litellm import Choices, Message, ModelResponse, Usage - from litellm.proxy.proxy_server import _ProxyDBLogger + from litellm.proxy.proxy_server import ProxyDBLogger - proxy_db_logger = _ProxyDBLogger() + proxy_db_logger = ProxyDBLogger() request_id = f"chatcmpl-{uuid.uuid4()}" @@ -2034,9 +2034,9 @@ async def test_call_with_key_over_budget_stream(prisma_client): from litellm._uuid import uuid from litellm import Choices, Message, ModelResponse, Usage - from litellm.proxy.proxy_server import _ProxyDBLogger + from litellm.proxy.proxy_server import ProxyDBLogger - proxy_db_logger = _ProxyDBLogger() + proxy_db_logger = ProxyDBLogger() request_id = f"chatcmpl-e41836bb-bb8b-4df2-8e70-8f3e160155ac{uuid.uuid4()}" resp = ModelResponse( @@ -2346,31 +2346,31 @@ async def test_upperbound_key_param_none_duration(prisma_client): def test_get_bearer_token(): - from litellm.proxy.auth.user_api_key_auth import _get_bearer_token + from litellm.proxy.auth.user_api_key_auth import get_bearer_token # Test valid Bearer token api_key = "Bearer valid_token" - result = _get_bearer_token(api_key) + result = get_bearer_token(api_key) assert result == "valid_token", f"Expected 'valid_token', got '{result}'" # Test empty API key api_key = "" - result = _get_bearer_token(api_key) + result = get_bearer_token(api_key) assert result == "", f"Expected '', got '{result}'" # Test API key without Bearer prefix api_key = "invalid_token" - result = _get_bearer_token(api_key) + result = get_bearer_token(api_key) assert result == "", f"Expected '', got '{result}'" # Test API key with Bearer prefix and extra spaces api_key = " Bearer valid_token " - result = _get_bearer_token(api_key) + result = get_bearer_token(api_key) assert result == "", f"Expected '', got '{result}'" # Test API key with Bearer prefix and no token api_key = "Bearer sk-9876" - result = _get_bearer_token(api_key) + result = get_bearer_token(api_key) assert result == "sk-9876", f"Expected 'sk-9876', got '{result}'" @@ -2507,7 +2507,7 @@ async def track_cost_callback_helper_fn(generated_key: str, user_id: str): from litellm._uuid import uuid from litellm import Choices, Message, ModelResponse, Usage - from litellm.proxy.proxy_server import _ProxyDBLogger + from litellm.proxy.proxy_server import ProxyDBLogger request_id = f"chatcmpl-e41836bb-bb8b-4df2-8e70-8f3e160155ac{uuid.uuid4()}" resp = ModelResponse( @@ -2525,7 +2525,7 @@ async def track_cost_callback_helper_fn(generated_key: str, user_id: str): model="gpt-35-turbo", # azure always has model written like this usage=Usage(prompt_tokens=210, completion_tokens=200, total_tokens=410), ) - proxy_db_logger = _ProxyDBLogger() + proxy_db_logger = ProxyDBLogger() await proxy_db_logger._PROXY_track_cost_callback( kwargs={ "call_type": "acompletion", diff --git a/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py b/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py index 20672e67358..0f21f3ee3e0 100644 --- a/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py @@ -38,7 +38,7 @@ from litellm.proxy._types import ( ) from litellm.models.object_permission import LiteLLM_ObjectPermissionTable from litellm.proxy.auth.auth_checks import ( - _delete_cache_key_object, + delete_cache_key_object, jwt_key_mapping_cache_key, ) from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth @@ -55,8 +55,8 @@ from litellm.proxy.management_endpoints.key_management_endpoints import ( _enforce_upperbound_key_params, _execute_virtual_key_regeneration, _get_and_validate_existing_key, - _list_key_helper, - _persist_deleted_verification_tokens, + list_key_helper, + persist_deleted_verification_tokens, _process_single_key_update, _requested_end_user_budget_id, _save_deleted_verification_token_records, @@ -108,7 +108,7 @@ async def test_list_keys(): "admin_team_ids": ["28bd3181-02c5-48f2-b408-ce790fb3d5ba"], } try: - result = await _list_key_helper(**args) + result = await list_key_helper(**args) except Exception as e: print(f"error: {e}") @@ -155,7 +155,7 @@ async def test_list_keys_include_created_by_keys(): } try: - result = await _list_key_helper(**args) + result = await list_key_helper(**args) except Exception as e: print(f"error: {e}") @@ -219,7 +219,7 @@ async def test_list_keys_include_created_by_keys(): args["include_created_by_keys"] = False try: - result = await _list_key_helper(**args) + result = await list_key_helper(**args) except Exception as e: print(f"error: {e}") @@ -246,7 +246,7 @@ async def test_list_keys_include_created_by_keys(): ) try: - result = await _list_key_helper(**args) + result = await list_key_helper(**args) except Exception as e: print(f"error: {e}") @@ -1468,7 +1468,7 @@ async def test_list_keys_full_object_returns_lifetime_total_spend(): ) mock_prisma_client.db.litellm_verificationtoken.count = AsyncMock(return_value=1) - result = await _list_key_helper( + result = await list_key_helper( prisma_client=mock_prisma_client, page=1, size=50, @@ -2708,7 +2708,7 @@ async def test_unblock_key_supports_both_sk_and_hashed_tokens(monkeypatch): pass monkeypatch.setattr( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", mock_delete_cache_key_object, ) @@ -3066,7 +3066,7 @@ async def test_update_key_by_alias_only(monkeypatch): request_data = UpdateKeyRequest(key_alias="prod-alias", max_budget=50.0) with patch( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object" + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object" ) as mock_delete_cache: mock_delete_cache.return_value = None result = await update_key_fn( @@ -3121,7 +3121,7 @@ async def test_update_key_changed_alias_must_match_key_alias_pattern(monkeypatch mock_prisma_client.update_data.assert_not_awaited() with patch( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", return_value=None, ): await update_key_fn( @@ -3237,7 +3237,7 @@ async def test_update_key_with_key_and_alias_selects_by_key(monkeypatch): ) with patch( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object" + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object" ) as mock_delete_cache: mock_delete_cache.return_value = None result = await update_key_fn( @@ -3307,7 +3307,7 @@ async def test_block_key_existing_key_succeeds(monkeypatch): pass monkeypatch.setattr( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", mock_delete_cache_key_object, ) @@ -5379,7 +5379,7 @@ async def test_persist_deleted_verification_tokens(): allowed_routes=[], ) - await _persist_deleted_verification_tokens( + await persist_deleted_verification_tokens( keys=[key], prisma_client=mock_prisma_client, user_api_key_dict=user_api_key_dict, @@ -5467,7 +5467,7 @@ async def test_delete_verification_tokens_persists_deleted_keys(monkeypatch): return token if not token.startswith("sk-") else f"hashed-{token}" monkeypatch.setattr( - "litellm.proxy.management_endpoints.key_management_endpoints._hash_token_if_needed", + "litellm.proxy.management_endpoints.key_management_endpoints.hash_token_if_needed", mock_hash_token, ) monkeypatch.setattr( @@ -5580,7 +5580,7 @@ async def test_delete_verification_tokens_evicts_jwt_key_mapping_cache(monkeypat recording_evict, ) monkeypatch.setattr( - "litellm.proxy.management_endpoints.key_management_endpoints._hash_token_if_needed", + "litellm.proxy.management_endpoints.key_management_endpoints.hash_token_if_needed", lambda token: token, ) monkeypatch.setattr( @@ -6150,7 +6150,7 @@ async def test_list_keys_with_expand_user(): "expand": ["user"], # Test the expand parameter } - result = await _list_key_helper(**args) + result = await list_key_helper(**args) # Verify that keys were fetched mock_find_many_keys.assert_called_once() @@ -6261,7 +6261,7 @@ async def test_list_keys_with_expand_user_includes_created_by_user(): "expand": ["user"], } - result = await _list_key_helper(**args) + result = await list_key_helper(**args) # Verify that the user lookup included both user_id and created_by call_args = mock_find_many_users.call_args @@ -6342,7 +6342,7 @@ async def test_list_keys_with_status_deleted(): "status": "deleted", # Test the status parameter } - result = await _list_key_helper(**args) + result = await list_key_helper(**args) # Verify that deleted table was queried mock_find_many_deleted.assert_called_once() @@ -6481,7 +6481,7 @@ async def test_list_key_helper_revoked_status_filters_live_table_on_blocked(): mock_prisma_client.db.litellm_verificationtoken.count = AsyncMock(return_value=0) mock_prisma_client.db.litellm_deletedverificationtoken.find_many = AsyncMock(return_value=[]) - await _list_key_helper( + await list_key_helper( prisma_client=mock_prisma_client, page=1, size=50, @@ -6676,7 +6676,7 @@ async def test_list_keys_non_admin_user_id_auto_set(): return_value=[], ): with patch( - "litellm.proxy.management_endpoints.key_management_endpoints._list_key_helper", + "litellm.proxy.management_endpoints.key_management_endpoints.list_key_helper", mock_list_key_helper, ): mock_request = Mock() @@ -6755,7 +6755,7 @@ async def _invoke_list_keys_and_capture_helper_kwargs( AsyncMock(return_value=team_objects), ), patch( - "litellm.proxy.management_endpoints.key_management_endpoints._list_key_helper", + "litellm.proxy.management_endpoints.key_management_endpoints.list_key_helper", mock_list_key_helper, ), ): @@ -7176,7 +7176,7 @@ async def test_list_key_helper_applies_search_to_prisma_where(): mock_prisma_client.db.litellm_verificationtoken.find_many = mock_find_many mock_prisma_client.db.litellm_verificationtoken.count = AsyncMock(return_value=0) - await _list_key_helper( + await list_key_helper( prisma_client=mock_prisma_client, page=1, size=50, @@ -7219,7 +7219,7 @@ async def _run_bulk_update_on_one_key( with ( patch( # test-quality-ok: the handler reads the cache and hook singletons from module globals, no injection seam - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", new_callable=AsyncMock, ), patch( # test-quality-ok: the permission check is a classmethod the handler calls directly, no injection seam @@ -7603,7 +7603,7 @@ async def test_get_and_validate_existing_key(): ) with patch( - "litellm.proxy.management_endpoints.key_management_endpoints._hash_token_if_needed", + "litellm.proxy.management_endpoints.key_management_endpoints.hash_token_if_needed", return_value="hashed-test-key-123", ): result = await _get_and_validate_existing_key( @@ -7622,7 +7622,7 @@ async def test_get_and_validate_existing_key(): ) with patch( - "litellm.proxy.management_endpoints.key_management_endpoints._hash_token_if_needed", + "litellm.proxy.management_endpoints.key_management_endpoints.hash_token_if_needed", return_value="hashed-non-existent-key", ): with pytest.raises(ProxyException) as exc_info: @@ -7704,7 +7704,7 @@ async def test_process_single_key_update(): # Mock _delete_cache_key_object with patch( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object" + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object" ) as mock_delete_cache: mock_delete_cache.return_value = None @@ -7714,7 +7714,7 @@ async def test_process_single_key_update(): # Mock _hash_token_if_needed with patch( - "litellm.proxy.management_endpoints.key_management_endpoints._hash_token_if_needed", + "litellm.proxy.management_endpoints.key_management_endpoints.hash_token_if_needed", return_value="hashed-test-key-123", ): # Mock KeyManagementEventHooks @@ -7853,7 +7853,7 @@ async def test_bulk_update_keys_success(monkeypatch): "litellm.proxy.management_endpoints.key_management_endpoints.TeamMemberPermissionChecks.can_team_member_execute_key_management_endpoint" ): with patch( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object" + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object" ): with patch("litellm.proxy._types.hash_token") as mock_hash: mock_hash.side_effect = ["hashed-key-1", "hashed-key-2"] @@ -7865,7 +7865,7 @@ async def test_bulk_update_keys_success(monkeypatch): }[token] with patch( - "litellm.proxy.management_endpoints.key_management_endpoints._hash_token_if_needed", + "litellm.proxy.management_endpoints.key_management_endpoints.hash_token_if_needed", side_effect=_hash_for_bulk_success, ): with patch( @@ -7981,7 +7981,7 @@ async def test_bulk_update_keys_partial_failures(monkeypatch): "litellm.proxy.management_endpoints.key_management_endpoints.TeamMemberPermissionChecks.can_team_member_execute_key_management_endpoint" ): with patch( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object" + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object" ): with patch("litellm.proxy._types.hash_token") as mock_hash: mock_hash.return_value = "hashed-key-1" @@ -7993,7 +7993,7 @@ async def test_bulk_update_keys_partial_failures(monkeypatch): }[token] with patch( - "litellm.proxy.management_endpoints.key_management_endpoints._hash_token_if_needed", + "litellm.proxy.management_endpoints.key_management_endpoints.hash_token_if_needed", side_effect=_hash_for_bulk_partial, ): with patch( @@ -8186,7 +8186,7 @@ async def test_reset_key_spend_success(monkeypatch): "litellm.proxy.management_endpoints.key_management_endpoints._check_proxy_or_team_admin_for_key" ) as mock_check_admin, patch( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object" + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object" ) as mock_delete_cache, ): mock_hash_token.return_value = hashed_key @@ -8306,7 +8306,7 @@ async def test_reset_key_spend_resets_budget_windows(monkeypatch): "litellm.proxy.management_endpoints.key_management_endpoints._check_proxy_or_team_admin_for_key" ) as mock_check_admin, patch( # test-quality-ok: no HTTP boundary; same pattern as test_reset_key_spend_success - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object" + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object" ) as mock_delete_cache, ): mock_hash_token.return_value = hashed_key @@ -8415,7 +8415,7 @@ async def test_reset_key_spend_no_budget_limits_skips_window_reset(monkeypatch): "litellm.proxy.management_endpoints.key_management_endpoints._check_proxy_or_team_admin_for_key" ) as mock_check_admin, patch( # test-quality-ok: no HTTP boundary; same pattern as test_reset_key_spend_success - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object" + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object" ) as mock_delete_cache, ): mock_hash_token.return_value = hashed_key @@ -8461,7 +8461,7 @@ async def test_delete_cache_key_object_broadcasts_invalidation(monkeypatch): "litellm.proxy.auth.auth_checks.publish_auth_cache_invalidation" ) as mock_publish: mock_publish.return_value = None - await _delete_cache_key_object( + await delete_cache_key_object( hashed_token="hashed-broadcast-key", user_api_key_cache=real_user_api_key_cache, proxy_logging_obj=mock_proxy_logging_obj, @@ -8519,7 +8519,7 @@ async def test_update_key_spend_updates_counter(monkeypatch): ) with patch( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object" + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object" ) as mock_delete_cache: mock_delete_cache.return_value = None @@ -8616,7 +8616,7 @@ async def test_reset_key_spend_success_team_admin(monkeypatch): with ( patch("litellm.proxy.proxy_server.hash_token") as mock_hash_token, patch( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object" + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object" ) as mock_delete_cache, ): mock_hash_token.return_value = hashed_key @@ -8830,7 +8830,7 @@ async def test_reset_key_spend_hashed_key(monkeypatch): "litellm.proxy.management_endpoints.key_management_endpoints._check_proxy_or_team_admin_for_key" ) as mock_check_admin, patch( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object" + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object" ) as mock_delete_cache, ): mock_check_admin.return_value = None @@ -9326,7 +9326,7 @@ async def test_rotate_master_key_reencrypts_model_params_in_place( from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth from litellm.proxy.management_endpoints.key_management_endpoints import ( - _rotate_master_key, + rotate_master_key, ) # Setup mock prisma client @@ -9394,7 +9394,7 @@ async def test_rotate_master_key_reencrypts_model_params_in_place( "litellm.proxy.proxy_server.proxy_config", mock_proxy_config, ): - await _rotate_master_key( + await rotate_master_key( prisma_client=mock_prisma_client, user_api_key_dict=user_api_key_dict, current_master_key="sk-old-master-key", @@ -11016,7 +11016,7 @@ def _setup_block_unblock_mocks(monkeypatch, mock_key_team_id=None): pass monkeypatch.setattr( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", mock_delete_cache_key_object, ) @@ -11270,7 +11270,7 @@ async def test_update_key_non_budget_fields_allowed_for_internal_user(monkeypatc pass monkeypatch.setattr( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", mock_delete_cache_key_object, ) @@ -11431,7 +11431,7 @@ async def test_update_key_throttle_unchanged_allows_non_budget_edit_for_internal pass monkeypatch.setattr( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", _noop, ) monkeypatch.setattr( @@ -11669,7 +11669,7 @@ async def test_update_key_team_member_with_permission_can_update_non_budget( mock_enforce_unique_key_alias, ) monkeypatch.setattr( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", mock_delete_cache_key_object, ) monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) @@ -12811,11 +12811,11 @@ async def test_execute_virtual_key_regeneration_rejects_over_limit_duration(monk new_callable=AsyncMock, ), patch( # test-quality-ok: archival path is outside upperbound rejection - "litellm.proxy.management_endpoints.key_management_endpoints._persist_deleted_verification_tokens", + "litellm.proxy.management_endpoints.key_management_endpoints.persist_deleted_verification_tokens", new_callable=AsyncMock, ) as persist_deleted_verification_tokens, patch( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", new_callable=AsyncMock, ), ): @@ -12861,11 +12861,11 @@ async def test_execute_virtual_key_regeneration_changed_alias_must_match_key_ali new_callable=AsyncMock, ), patch( - "litellm.proxy.management_endpoints.key_management_endpoints._persist_deleted_verification_tokens", + "litellm.proxy.management_endpoints.key_management_endpoints.persist_deleted_verification_tokens", new_callable=AsyncMock, ), patch( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", new_callable=AsyncMock, ), ): @@ -12921,7 +12921,7 @@ async def test_execute_virtual_key_regeneration_allows_within_limit_duration(mon new_callable=AsyncMock, ), patch( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", new_callable=AsyncMock, ), patch( @@ -12967,11 +12967,11 @@ async def test_execute_virtual_key_regeneration_rejects_when_custom_key_update_h new_callable=AsyncMock, ) as insert_deprecated_key, patch( # test-quality-ok: archival path is outside policy rejection - "litellm.proxy.management_endpoints.key_management_endpoints._persist_deleted_verification_tokens", + "litellm.proxy.management_endpoints.key_management_endpoints.persist_deleted_verification_tokens", new_callable=AsyncMock, ) as persist_deleted_verification_tokens, patch( # test-quality-ok: cache eviction is outside policy rejection - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", new_callable=AsyncMock, ), patch( # test-quality-ok: rotation callback is outside policy rejection @@ -13027,11 +13027,11 @@ async def test_execute_virtual_key_regeneration_allows_when_custom_key_update_ho new_callable=AsyncMock, ), patch( # test-quality-ok: verify archival follows policy approval - "litellm.proxy.management_endpoints.key_management_endpoints._persist_deleted_verification_tokens", + "litellm.proxy.management_endpoints.key_management_endpoints.persist_deleted_verification_tokens", new_callable=AsyncMock, ) as persist_deleted_verification_tokens, patch( # test-quality-ok: cache eviction is outside policy approval - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", new_callable=AsyncMock, ), patch( # test-quality-ok: rotation callback is outside policy approval @@ -13080,7 +13080,7 @@ async def test_execute_virtual_key_regeneration_skips_custom_key_update_hook_wit new_callable=AsyncMock, ), patch( # test-quality-ok: cache eviction is outside unchanged request - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", new_callable=AsyncMock, ), patch( # test-quality-ok: rotation callback is outside unchanged request @@ -13129,7 +13129,7 @@ async def test_execute_virtual_key_regeneration_hides_the_untouched_modal_expiry new_callable=AsyncMock, ), patch( # test-quality-ok: cache eviction is outside the hook input - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", new_callable=AsyncMock, ), patch( # test-quality-ok: rotation callback is outside the hook input @@ -13197,13 +13197,13 @@ def _regenerate_policy_mocks(policy, insert_deprecated_key: AsyncMock, persist: ) stack.enter_context( patch( # test-quality-ok: archival write must not run on a denied regenerate - "litellm.proxy.management_endpoints.key_management_endpoints._persist_deleted_verification_tokens", + "litellm.proxy.management_endpoints.key_management_endpoints.persist_deleted_verification_tokens", persist, ) ) stack.enter_context( patch( # test-quality-ok: cache eviction is outside the policy path - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", new_callable=AsyncMock, ) ) @@ -13340,7 +13340,7 @@ async def test_update_key_fn_runs_custom_key_policy_on_the_effective_row(monkeyp with ( patch( # test-quality-ok: cache eviction is outside the policy path - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", new_callable=AsyncMock, ), patch("litellm.proxy.proxy_server.user_custom_key_policy", policy), # test-quality-ok: inject policy hook @@ -13387,7 +13387,7 @@ async def test_update_key_fn_rejects_when_custom_key_policy_denies(monkeypatch): async def _process_single_key_update_under_policy(prisma_client: AsyncMock, data: UpdateKeyRequest, policy): with ( patch( # test-quality-ok: cache eviction is outside the policy path - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", new_callable=AsyncMock, ), patch( # test-quality-ok: update callback is outside the policy path @@ -13492,7 +13492,7 @@ def _setup_update_key_fn_object_permission_mocks(monkeypatch, allowed: bool) -> _record_object_permission_writes(mock_prisma_client, events) monkeypatch.setattr("litellm.proxy.proxy_server.user_custom_key_policy", _recording_policy(events, allowed)) monkeypatch.setattr( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", AsyncMock() + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", AsyncMock() ) return mock_prisma_client, events @@ -13628,7 +13628,7 @@ async def test_bulk_update_keys_runs_custom_key_policy_per_key(monkeypatch): with ( patch( # test-quality-ok: cache eviction is outside the policy path - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", new_callable=AsyncMock, ), patch( # test-quality-ok: update callback is outside the policy path @@ -13996,7 +13996,7 @@ async def test_regenerate_evicts_jwt_key_mapping_cache_so_next_jwt_call_gets_new new_callable=AsyncMock, ), patch( # test-quality-ok: key-object eviction is separate from the mapping eviction under test - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", new_callable=AsyncMock, ), patch( # test-quality-ok: background rotation hook is irrelevant to cache eviction @@ -14090,7 +14090,7 @@ async def test_execute_virtual_key_regeneration_rejects_over_limit_max_budget(mo new_callable=AsyncMock, ), patch( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", new_callable=AsyncMock, ), ): @@ -14146,7 +14146,7 @@ async def test_execute_virtual_key_regeneration_skips_none_values(monkeypatch): new_callable=AsyncMock, ), patch( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", new_callable=AsyncMock, ), patch( @@ -14193,7 +14193,7 @@ async def test_execute_virtual_key_regeneration_no_upperbound_config_is_noop(mon new_callable=AsyncMock, ), patch( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", new_callable=AsyncMock, ), patch( @@ -14685,7 +14685,7 @@ async def test_process_single_key_update_cache_invalidation_with_token_hash(): return_value=None, ), patch( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", new_callable=AsyncMock, ) as mock_delete_cache, patch( @@ -14875,7 +14875,7 @@ async def test_execute_virtual_key_regeneration_cache_invalidation_with_token_ha new_callable=AsyncMock, ), patch( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", new_callable=AsyncMock, ) as mock_delete_cache, patch( @@ -14968,7 +14968,7 @@ def _setup_team_keys_mocks( f"{_BULK_PKG}.prepare_key_update_data", AsyncMock(return_value={"max_budget": 50.0}), ) - monkeypatch.setattr(f"{_BULK_PKG}._delete_cache_key_object", AsyncMock()) + monkeypatch.setattr(f"{_BULK_PKG}.delete_cache_key_object", AsyncMock()) monkeypatch.setattr( f"{_BULK_PKG}.KeyManagementEventHooks.async_key_updated_hook", AsyncMock() ) @@ -14980,7 +14980,7 @@ def _setup_team_keys_mocks( ) if hash_identity: # Tests use already-hashed tokens; the raw-sk regression opts out. - monkeypatch.setattr(f"{_BULK_PKG}._hash_token_if_needed", lambda token: token) + monkeypatch.setattr(f"{_BULK_PKG}.hash_token_if_needed", lambda token: token) return mock_prisma @@ -15565,7 +15565,7 @@ def _patch_regenerate_side_effects(): new_callable=AsyncMock, ), patch( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", new_callable=AsyncMock, ), patch( @@ -15672,7 +15672,7 @@ async def test_regenerate_premium_gate_allows_actual_master_key_holder(): patch("litellm.proxy.proxy_server.master_key", master), patch("litellm.proxy.proxy_server.prisma_client", AsyncMock()), patch( - "litellm.proxy.management_endpoints.key_management_endpoints._rotate_master_key", + "litellm.proxy.management_endpoints.key_management_endpoints.rotate_master_key", new_callable=AsyncMock, ), ): @@ -17158,7 +17158,7 @@ async def _list_keys_capture_helper_kwargs(user_api_key_dict, **list_kwargs): return_value=mock_user_info, ): with patch( - "litellm.proxy.management_endpoints.key_management_endpoints._list_key_helper", + "litellm.proxy.management_endpoints.key_management_endpoints.list_key_helper", helper, ): await list_keys( @@ -18467,7 +18467,7 @@ async def test_list_keys_forwards_expires_filter(expires_value, expected_forward return_value=mock_user_info, ), patch( - "litellm.proxy.management_endpoints.key_management_endpoints._list_key_helper", + "litellm.proxy.management_endpoints.key_management_endpoints.list_key_helper", mock_helper, ), ): @@ -18506,7 +18506,7 @@ async def test_list_keys_without_expires_param_forwards_none(): return_value=mock_user_info, ), patch( - "litellm.proxy.management_endpoints.key_management_endpoints._list_key_helper", + "litellm.proxy.management_endpoints.key_management_endpoints.list_key_helper", mock_helper, ), ): @@ -18546,7 +18546,7 @@ async def test_rotate_master_key_rotates_sso_identity_assertions( from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth from litellm.proxy.management_endpoints.key_management_endpoints import ( - _rotate_master_key, + rotate_master_key, ) mock_prisma_client = AsyncMock() @@ -18580,7 +18580,7 @@ async def test_rotate_master_key_rotates_sso_identity_assertions( "litellm.proxy.proxy_server.proxy_config", mock_proxy_config, ): - await _rotate_master_key( + await rotate_master_key( prisma_client=mock_prisma_client, user_api_key_dict=user_api_key_dict, current_master_key="sk-old-master-key", @@ -18604,7 +18604,7 @@ async def test_rotate_master_key_rotates_search_tools(monkeypatch): encrypt_value_helper, ) from litellm.proxy.management_endpoints.key_management_endpoints import ( - _rotate_master_key, + rotate_master_key, ) monkeypatch.delenv("LITELLM_SALT_KEY", raising=False) @@ -18639,7 +18639,7 @@ async def test_rotate_master_key_rotates_search_tools(monkeypatch): user_id="test-user", ) - await _rotate_master_key( + await rotate_master_key( prisma_client=mock_prisma_client, user_api_key_dict=user_api_key_dict, current_master_key="sk-old-master-key", @@ -18901,7 +18901,7 @@ def _wire_update_key_fn(monkeypatch, existing_key): pass monkeypatch.setattr( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", _noop, ) monkeypatch.setattr( @@ -19570,7 +19570,7 @@ async def test_update_key_syncs_access_group_assigned_key_ids_in_both_directions with ( patch( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", new_callable=AsyncMock, ), patch( @@ -19657,7 +19657,7 @@ async def test_update_key_leaves_access_groups_alone_when_field_is_unset(monkeyp _setup_update_key_mocks(monkeypatch, mock_prisma_client) with patch( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", new_callable=AsyncMock, ): await update_key_fn( @@ -19705,7 +19705,7 @@ async def test_bulk_update_keys_syncs_access_group_assigned_key_ids(monkeypatch) with ( patch( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", new_callable=AsyncMock, ), patch( @@ -19882,7 +19882,7 @@ async def test_regenerate_key_repoints_access_group_assigned_key_ids(monkeypatch new_callable=AsyncMock, ), patch( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", new_callable=AsyncMock, ), patch( @@ -19961,7 +19961,7 @@ async def test_key_write_paths_revoke_the_key_cache_before_syncing_access_groups with ( patch( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", new_callable=AsyncMock, side_effect=lambda **kwargs: order.append("revoke_key_cache"), ), @@ -20035,7 +20035,7 @@ async def test_update_key_syncs_many_access_groups_in_one_statement_per_directio with ( patch( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", new_callable=AsyncMock, ), patch( @@ -20118,7 +20118,7 @@ async def test_regenerate_key_repoints_live_membership_not_the_key_row_it_read( new_callable=AsyncMock, ), patch( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", new_callable=AsyncMock, ), patch( @@ -20899,7 +20899,7 @@ async def test_key_regeneration_invalidates_cached_object_permission(monkeypatch new_callable=AsyncMock, ), patch( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", new_callable=AsyncMock, ), patch( @@ -20967,7 +20967,7 @@ async def test_key_update_evicts_object_permission_before_key_object(monkeypatch """ from litellm.proxy._types import LiteLLM_ObjectPermissionBase from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache, object_permission_cache_key - from litellm.proxy.utils import _hash_token_if_needed + from litellm.proxy.utils import hash_token_if_needed deleted: list[str] = [] @@ -21028,7 +21028,7 @@ async def test_key_update_evicts_object_permission_before_key_object(monkeypatch ) assert deleted.index(object_permission_cache_key(permission_id)) < deleted.index( - _hash_token_if_needed("sk-lit5479") + hash_token_if_needed("sk-lit5479") ), deleted @@ -21207,7 +21207,7 @@ async def test_rotate_master_key_reencrypts_guardrail_params(monkeypatch): ) from litellm.proxy.management_endpoints import key_management_endpoints from litellm.proxy.management_endpoints.key_management_endpoints import ( - _rotate_master_key, + rotate_master_key, ) for rotator in ( @@ -21232,7 +21232,7 @@ async def test_rotate_master_key_reencrypts_guardrail_params(monkeypatch): mock_prisma_client.db.litellm_guardrailstable.find_many = AsyncMock(return_value=[guardrail_row]) mock_prisma_client.db.litellm_guardrailstable.update_many = AsyncMock(return_value=1) - await _rotate_master_key( + await rotate_master_key( prisma_client=mock_prisma_client, user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="test-user"), current_master_key="sk-old-master-key", diff --git a/tests/unit/proxy/management_endpoints/test_mcp_management_endpoints.py b/tests/unit/proxy/management_endpoints/test_mcp_management_endpoints.py index 3c2e099e06c..21ff5534e03 100644 --- a/tests/unit/proxy/management_endpoints/test_mcp_management_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_mcp_management_endpoints.py @@ -422,7 +422,7 @@ class TestListMCPServers: mock_manager, ), patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints._user_has_admin_view", + "litellm.proxy.management_endpoints.mcp_management_endpoints.user_api_key_has_admin_view", return_value=True, ), patch( @@ -669,7 +669,7 @@ class TestListMCPServers: mock_manager, ), patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints._user_has_admin_view", + "litellm.proxy.management_endpoints.mcp_management_endpoints.user_api_key_has_admin_view", return_value=True, ), patch( @@ -788,7 +788,7 @@ class TestListMCPServers: mock_manager, ), patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints._user_has_admin_view", + "litellm.proxy.management_endpoints.mcp_management_endpoints.user_api_key_has_admin_view", return_value=False, ), patch( @@ -924,7 +924,7 @@ class TestListMCPServers: AsyncMock(return_value=mock_health_result), ), patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints._user_has_admin_view", + "litellm.proxy.management_endpoints.mcp_management_endpoints.user_api_key_has_admin_view", return_value=True, ), ): @@ -971,7 +971,7 @@ class TestListMCPServers: AsyncMock(return_value=mock_health_result), ), patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints._user_has_admin_view", + "litellm.proxy.management_endpoints.mcp_management_endpoints.user_api_key_has_admin_view", return_value=True, ), ): @@ -1028,7 +1028,7 @@ class TestListMCPServers: AsyncMock(return_value=mock_health_result), ), patch( # test-quality-ok: endpoint test must patch module globals - "litellm.proxy.management_endpoints.mcp_management_endpoints._user_has_admin_view", + "litellm.proxy.management_endpoints.mcp_management_endpoints.user_api_key_has_admin_view", return_value=True, ), ): @@ -1105,7 +1105,7 @@ class TestListMCPServers: AsyncMock(return_value=mock_health_result), ), patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints._user_has_admin_view", + "litellm.proxy.management_endpoints.mcp_management_endpoints.user_api_key_has_admin_view", return_value=True, ), ): @@ -1149,7 +1149,7 @@ class TestListMCPServers: AsyncMock(return_value=mock_health_result), ), patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints._user_has_admin_view", + "litellm.proxy.management_endpoints.mcp_management_endpoints.user_api_key_has_admin_view", return_value=True, ), ): @@ -1195,7 +1195,7 @@ class TestListMCPServers: AsyncMock(return_value=mock_health_result), ), patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints._user_has_admin_view", + "litellm.proxy.management_endpoints.mcp_management_endpoints.user_api_key_has_admin_view", return_value=True, ), ): @@ -1237,7 +1237,7 @@ class TestListMCPServers: side_effect=lambda sid: config_server if sid == "serper_custom_dev" else None ) mock_manager.get_mcp_server_by_name = MagicMock(return_value=None) - mock_manager._build_mcp_server_table = MagicMock( + mock_manager.build_mcp_server_table = MagicMock( return_value=generate_mock_mcp_server_db_record( server_id="serper_custom_dev", alias="Serper MCP", @@ -1264,7 +1264,7 @@ class TestListMCPServers: mock_manager, ), patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints._user_has_admin_view", + "litellm.proxy.management_endpoints.mcp_management_endpoints.user_api_key_has_admin_view", return_value=True, ), ): @@ -1281,7 +1281,7 @@ class TestListMCPServers: assert result.server_id == "serper_custom_dev" assert result.status == "healthy" mock_manager.get_mcp_server_by_id.assert_called_with("serper_custom_dev") - mock_manager._build_mcp_server_table.assert_called_once() + mock_manager.build_mcp_server_table.assert_called_once() @pytest.mark.asyncio async def test_fetch_single_mcp_server_from_registry_by_name_passes_client_ip(self): @@ -1299,7 +1299,7 @@ class TestListMCPServers: mock_manager = MagicMock() mock_manager.get_mcp_server_by_id = MagicMock(return_value=None) mock_manager.get_mcp_server_by_name = MagicMock(return_value=config_server) - mock_manager._build_mcp_server_table = MagicMock( + mock_manager.build_mcp_server_table = MagicMock( return_value=generate_mock_mcp_server_db_record( server_id="serper_custom_dev", alias="Serper MCP", @@ -1328,7 +1328,7 @@ class TestListMCPServers: mock_manager, ), patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints._user_has_admin_view", + "litellm.proxy.management_endpoints.mcp_management_endpoints.user_api_key_has_admin_view", return_value=True, ), ): @@ -1363,7 +1363,7 @@ class TestListMCPServers: side_effect=lambda sid: config_server if sid == "restricted_server" else None ) mock_manager.get_mcp_server_by_name = MagicMock(return_value=None) - mock_manager._build_mcp_server_table = MagicMock( + mock_manager.build_mcp_server_table = MagicMock( return_value=generate_mock_mcp_server_db_record( server_id="restricted_server", alias="Restricted MCP", @@ -1391,7 +1391,7 @@ class TestListMCPServers: mock_manager, ), patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints._user_has_admin_view", + "litellm.proxy.management_endpoints.mcp_management_endpoints.user_api_key_has_admin_view", return_value=False, ), ): @@ -1431,7 +1431,7 @@ class TestListMCPServers: side_effect=lambda sid: config_server if sid == "allowed_config_server" else None ) mock_manager.get_mcp_server_by_name = MagicMock(return_value=None) - mock_manager._build_mcp_server_table = MagicMock( + mock_manager.build_mcp_server_table = MagicMock( return_value=generate_mock_mcp_server_db_record( server_id="allowed_config_server", alias="Allowed MCP", @@ -1458,7 +1458,7 @@ class TestListMCPServers: mock_manager, ), patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints._user_has_admin_view", + "litellm.proxy.management_endpoints.mcp_management_endpoints.user_api_key_has_admin_view", return_value=False, ), ): @@ -1545,7 +1545,7 @@ class TestListMCPServers: AsyncMock(return_value=[generate_mock_mcp_server_db_record(server_id="env-server")]), ), patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints._user_has_admin_view", + "litellm.proxy.management_endpoints.mcp_management_endpoints.user_api_key_has_admin_view", return_value=False, ), ): @@ -1707,7 +1707,7 @@ class TestTeamScopedMCPServerAccess: with ( patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints._user_has_admin_view", + "litellm.proxy.management_endpoints.mcp_management_endpoints.user_api_key_has_admin_view", return_value=False, ), patch( @@ -1743,13 +1743,13 @@ class TestTeamScopedMCPServerAccess: mock_server = generate_mock_mcp_server_config_record(server_id="server-1", name="Team Server") mock_manager = MagicMock() mock_manager.get_mcp_server_by_id = MagicMock(return_value=mock_server) - mock_manager._build_mcp_server_table = MagicMock( + mock_manager.build_mcp_server_table = MagicMock( return_value=generate_mock_mcp_server_db_record(server_id="server-1") ) with ( patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints._user_has_admin_view", + "litellm.proxy.management_endpoints.mcp_management_endpoints.user_api_key_has_admin_view", return_value=False, ), patch( @@ -1783,7 +1783,7 @@ class TestTeamScopedMCPServerAccess: with ( patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints._user_has_admin_view", + "litellm.proxy.management_endpoints.mcp_management_endpoints.user_api_key_has_admin_view", return_value=True, ), patch( @@ -1836,7 +1836,7 @@ class TestFetchAllMCPServersOrdering: mock_manager, ), patch( # test-quality-ok: admin view is derived from module-global proxy settings - "litellm.proxy.management_endpoints.mcp_management_endpoints._user_has_admin_view", + "litellm.proxy.management_endpoints.mcp_management_endpoints.user_api_key_has_admin_view", return_value=True, ), patch( # test-quality-ok: auth contexts need a live prisma client @@ -1902,10 +1902,10 @@ class TestTemporaryMCPSessionEndpoints: mock_manager, ): from litellm.proxy.management_endpoints.mcp_management_endpoints import ( - _inherit_credentials_from_existing_server, + inherit_credentials_from_existing_server, ) - updated_payload = _inherit_credentials_from_existing_server(payload) + updated_payload = inherit_credentials_from_existing_server(payload) assert updated_payload.credentials == { "auth_value": "token-abc", @@ -1946,10 +1946,10 @@ class TestTemporaryMCPSessionEndpoints: mock_manager, ): from litellm.proxy.management_endpoints.mcp_management_endpoints import ( - _inherit_credentials_from_existing_server, + inherit_credentials_from_existing_server, ) - return _inherit_credentials_from_existing_server(payload) + return inherit_credentials_from_existing_server(payload) @pytest.mark.parametrize( "credentials", @@ -2509,7 +2509,7 @@ class TestTemporaryMCPSessionEndpoints: mock_manager = MagicMock() mock_manager.get_mcp_server_by_id.return_value = registry_server mock_manager.get_mcp_server_by_name.return_value = None - mock_manager._build_mcp_server_table.return_value = generate_mock_mcp_server_db_record(server_id="server-x") + mock_manager.build_mcp_server_table.return_value = generate_mock_mcp_server_db_record(server_id="server-x") mock_manager.get_allowed_mcp_servers = AsyncMock(return_value=["server-x"]) with ( @@ -2555,7 +2555,7 @@ class TestTemporaryMCPSessionEndpoints: mock_manager = MagicMock() mock_manager.get_mcp_server_by_id.return_value = registry_server mock_manager.get_mcp_server_by_name.return_value = None - mock_manager._build_mcp_server_table.return_value = generate_mock_mcp_server_db_record(server_id="server-x") + mock_manager.build_mcp_server_table.return_value = generate_mock_mcp_server_db_record(server_id="server-x") def allowed_for(auth): return ["server-x"] if auth.team_id == "team-with-mcp-grant" else [] @@ -2750,11 +2750,11 @@ class TestTemporaryMCPSessionEndpoints: with ( patch.dict(sys.modules, {"litellm.proxy.proxy_server": fake_proxy_server}), patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints._user_api_key_auth_builder", + "litellm.proxy.management_endpoints.mcp_management_endpoints.user_api_key_auth_builder", AsyncMock(return_value=expected_auth), ) as auth_builder_mock, patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints._read_request_body", + "litellm.proxy.management_endpoints.mcp_management_endpoints.read_request_body", AsyncMock(return_value={}), ), patch( @@ -2785,11 +2785,11 @@ class TestTemporaryMCPSessionEndpoints: with ( patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints._user_api_key_auth_builder", + "litellm.proxy.management_endpoints.mcp_management_endpoints.user_api_key_auth_builder", AsyncMock(return_value=expected_auth), ) as auth_builder_mock, patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints._read_request_body", + "litellm.proxy.management_endpoints.mcp_management_endpoints.read_request_body", AsyncMock(return_value={}), ), patch( @@ -2834,11 +2834,11 @@ class TestTemporaryMCPSessionEndpoints: mock_manager, ), patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints._user_api_key_auth_builder", + "litellm.proxy.management_endpoints.mcp_management_endpoints.user_api_key_auth_builder", AsyncMock(return_value=expected_auth), ) as auth_builder_mock, patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints._read_request_body", + "litellm.proxy.management_endpoints.mcp_management_endpoints.read_request_body", AsyncMock(return_value={}), ), patch( @@ -2888,11 +2888,11 @@ class TestTemporaryMCPSessionEndpoints: mock_manager, ), patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints._user_api_key_auth_builder", + "litellm.proxy.management_endpoints.mcp_management_endpoints.user_api_key_auth_builder", AsyncMock(return_value=expected_auth), ) as auth_builder_mock, patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints._read_request_body", + "litellm.proxy.management_endpoints.mcp_management_endpoints.read_request_body", AsyncMock(return_value={}), ), patch( @@ -3726,7 +3726,7 @@ class TestTemporaryMCPSessionEndpoints: return_value=nullcontext(server), ) as get_server, patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints._read_request_body", + "litellm.proxy.management_endpoints.mcp_management_endpoints.read_request_body", AsyncMock(return_value=request_body), ) as read_body, patch( @@ -3790,7 +3790,7 @@ class TestTemporaryMCPSessionEndpoints: return_value=nullcontext(server), ), patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints._read_request_body", + "litellm.proxy.management_endpoints.mcp_management_endpoints.read_request_body", AsyncMock(return_value=request_body), ), patch( @@ -3835,7 +3835,7 @@ class TestTemporaryMCPSessionEndpoints: return_value=nullcontext(server), ), patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints._read_request_body", + "litellm.proxy.management_endpoints.mcp_management_endpoints.read_request_body", AsyncMock(return_value=request_body), ), patch( @@ -5444,7 +5444,7 @@ async def test_store_mcp_oauth_user_credential_returns_status(): new=AsyncMock(return_value=generate_mock_mcp_server_db_record(server_id=server_id)), ), patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints._user_has_admin_view", + "litellm.proxy.management_endpoints.mcp_management_endpoints.user_api_key_has_admin_view", return_value=True, ), patch( @@ -5513,7 +5513,7 @@ async def test_store_mcp_oauth_user_credential_blocked_when_identity_binding_enf new=AsyncMock(return_value=generate_mock_mcp_server_db_record(server_id=server_id)), ), patch( # test-quality-ok: mirrors the existing store-credential tests in this file - "litellm.proxy.management_endpoints.mcp_management_endpoints._user_has_admin_view", + "litellm.proxy.management_endpoints.mcp_management_endpoints.user_api_key_has_admin_view", return_value=True, ), patch.object( # test-quality-ok: registry is a module-level singleton; injecting it would change the endpoint signature @@ -5607,7 +5607,7 @@ async def test_store_mcp_oauth_user_credential_invalidates_cached_token(): new=AsyncMock(return_value=generate_mock_mcp_server_db_record(server_id=server_id)), ), patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints._user_has_admin_view", + "litellm.proxy.management_endpoints.mcp_management_endpoints.user_api_key_has_admin_view", return_value=True, ), patch( @@ -7360,7 +7360,7 @@ class TestPerUserCredentialConfigServerResolution: manager.get_mcp_server_by_id = MagicMock( side_effect=lambda sid: config_server if sid == self.CONFIG_SERVER_ID else None ) - manager._build_mcp_server_table = MagicMock(return_value=record) + manager.build_mcp_server_table = MagicMock(return_value=record) manager.get_allowed_mcp_servers = AsyncMock(return_value=[]) return manager @@ -7488,7 +7488,7 @@ class TestPerUserCredentialConfigServerResolution: manager.get_mcp_server_by_id = MagicMock( return_value=generate_mock_mcp_server_config_record(server_id=self.CONFIG_SERVER_ID) ) - manager._build_mcp_server_table = MagicMock(return_value=env_var_server) + manager.build_mcp_server_table = MagicMock(return_value=env_var_server) manager.get_allowed_mcp_servers = AsyncMock(return_value=[self.CONFIG_SERVER_ID]) merge_mock = AsyncMock(return_value={"CORP_USERNAME": "alice"}) with ( @@ -8700,7 +8700,7 @@ class TestPinMCPServerTools: } ) ) - manager._get_tools_from_server = AsyncMock( + manager.get_tools_from_server = AsyncMock( return_value=[ MCPTool(name=name, description=description, inputSchema=schema) for name, description, schema in upstream_tools @@ -8747,7 +8747,7 @@ class TestPinMCPServerTools: "count_notes": PinnedMCPTool(description="", input_schema={}), } assert result == expected - listing = manager._get_tools_from_server.await_args.kwargs + listing = manager.get_tools_from_server.await_args.kwargs assert listing["server"].pinned_tools is None assert listing["server"].tool_name_to_description is None assert listing["proxy_logging_obj"] is None @@ -8778,7 +8778,7 @@ class TestPinMCPServerTools: assert result == {"server_id": "srv-1", "status": "unpinned"} assert store_mock.await_args.args[1:] == ("srv-1", None) assert store_mock.await_args.kwargs == {"touched_by": "admin"} - manager._get_tools_from_server.assert_not_awaited() + manager.get_tools_from_server.assert_not_awaited() manager.reload_servers_from_database.assert_awaited_once() @pytest.mark.asyncio @@ -8822,7 +8822,7 @@ class TestPinMCPServerTools: assert (pin_exc.value.status_code, unpin_exc.value.status_code) == (403, 403) store_mock.assert_not_awaited() - manager._get_tools_from_server.assert_not_awaited() + manager.get_tools_from_server.assert_not_awaited() @pytest.mark.asyncio async def test_pin_unknown_server_is_404(self): diff --git a/tests/unit/proxy/management_endpoints/test_model_management_endpoints.py b/tests/unit/proxy/management_endpoints/test_model_management_endpoints.py index 4439a3511df..f44b2b57ccd 100644 --- a/tests/unit/proxy/management_endpoints/test_model_management_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_model_management_endpoints.py @@ -176,7 +176,7 @@ class MockProxyConfig: self.success = success self.deployment_called = False - async def _add_deployment_locked(self, prisma_client, proxy_logging_obj): + async def add_deployment_locked(self, prisma_client, proxy_logging_obj): self.deployment_called = True if not self.success: raise Exception("Failed to add deployment") @@ -830,7 +830,7 @@ class TestClearCache: mock_router.model_list = ["openai/gpt-4o", "openai/gpt-4o-mini"] mock_config = MagicMock() - mock_config._add_deployment_locked = AsyncMock( + mock_config.add_deployment_locked = AsyncMock( return_value=ReconcileOutcome(still_desired=frozenset(), live_after=frozenset()) ) @@ -888,7 +888,7 @@ class TestClearCache: mock_router.complexity_routers = {"db-complexity-router": MagicMock(), "config-router": MagicMock()} mock_config = MagicMock() - mock_config._add_deployment_locked = AsyncMock( + mock_config.add_deployment_locked = AsyncMock( return_value=ReconcileOutcome(still_desired=frozenset(), live_after=frozenset()) ) @@ -922,7 +922,7 @@ class TestClearCache: assert "config-router" in mock_router.complexity_routers # Should have called the already-locked reload to restore DB models - mock_config._add_deployment_locked.assert_called_once_with( + mock_config.add_deployment_locked.assert_called_once_with( prisma_client=mock_prisma, proxy_logging_obj=mock_logging ) @@ -967,7 +967,7 @@ class TestClearCache: mock_router.quality_routers = {} mock_config = MagicMock() - mock_config._add_deployment_locked = AsyncMock( + mock_config.add_deployment_locked = AsyncMock( return_value=ReconcileOutcome(still_desired=frozenset(), live_after=frozenset()) ) @@ -1023,7 +1023,7 @@ class TestClearCachePreservesConfigRouters: } mock_config = MagicMock() - mock_config._add_deployment_locked = AsyncMock( + mock_config.add_deployment_locked = AsyncMock( return_value=ReconcileOutcome(still_desired=frozenset(), live_after=frozenset()) ) @@ -1065,7 +1065,7 @@ class TestClearCachePreservesConfigRouters: mock_router.complexity_routers = {"shared-name": MagicMock()} mock_config = MagicMock() - mock_config._add_deployment_locked = AsyncMock( + mock_config.add_deployment_locked = AsyncMock( return_value=ReconcileOutcome(still_desired=frozenset(), live_after=frozenset()) ) @@ -1109,7 +1109,7 @@ class TestClearCachePreservesConfigRouters: mock_router.adaptive_routers = {"a1": MagicMock()} mock_config = MagicMock() - mock_config._add_deployment_locked = AsyncMock( + mock_config.add_deployment_locked = AsyncMock( return_value=ReconcileOutcome(still_desired=frozenset(), live_after=frozenset()) ) @@ -1632,7 +1632,7 @@ class TestTeamModelSiblingRouting: team_model_add to register the public name on the team's models list. """ from litellm.proxy.management_endpoints.model_management_endpoints import ( - _add_team_model_to_db, + add_team_model_to_db, ) from litellm.types.router import ModelInfo @@ -1659,7 +1659,7 @@ class TestTeamModelSiblingRouting: ) with ( patch( - "litellm.proxy.management_endpoints.model_management_endpoints._add_model_to_db", + "litellm.proxy.management_endpoints.model_management_endpoints.add_model_to_db", side_effect=mock_add_model_to_db, ), patch( @@ -1667,7 +1667,7 @@ class TestTeamModelSiblingRouting: mock_team_model_add, ), ): - await _add_team_model_to_db( + await add_team_model_to_db( model_params=dep, user_api_key_dict=user, prisma_client=prisma_client, @@ -2521,7 +2521,7 @@ class TestAddAndDeleteModelLifecycle: mock_prisma.db.litellm_proxymodeltable.delete = AsyncMock(return_value=db_row) mock_proxy_config = MagicMock() - mock_proxy_config._add_deployment_locked = AsyncMock( + mock_proxy_config.add_deployment_locked = AsyncMock( return_value=ReconcileOutcome(still_desired=frozenset(), live_after=frozenset()) ) @@ -2644,7 +2644,7 @@ class TestDeleteTeamBYOKModelGhost: patch(f"{_PS}.llm_router", MagicMock()), patch(f"{_PS}.proxy_logging_obj", MagicMock()), patch(f"{_PS}.user_api_key_cache", MagicMock()), - patch(f"{_MOD}._refresh_cached_team", new=AsyncMock()) as mock_refresh, + patch(f"{_MOD}.refresh_cached_team", new=AsyncMock()) as mock_refresh, ): result = await delete_model_endpoint( model_info=ModelInfoDelete(id=model_id), @@ -2720,7 +2720,7 @@ class TestDeleteTeamBYOKModelGhost: patch(f"{_PS}.llm_router", MagicMock()), patch(f"{_PS}.proxy_logging_obj", MagicMock()), patch(f"{_PS}.user_api_key_cache", MagicMock()), - patch(f"{_MOD}._refresh_cached_team", new=AsyncMock()), + patch(f"{_MOD}.refresh_cached_team", new=AsyncMock()), ): result = await delete_model_endpoint( model_info=ModelInfoDelete(id=model_id), @@ -2793,7 +2793,7 @@ class TestDeleteTeamBYOKModelGhost: patch(f"{_PS}.llm_router", MagicMock()), patch(f"{_PS}.proxy_logging_obj", MagicMock()), patch(f"{_PS}.user_api_key_cache", MagicMock()), - patch(f"{_MOD}._refresh_cached_team", new=AsyncMock()) as mock_refresh, + patch(f"{_MOD}.refresh_cached_team", new=AsyncMock()) as mock_refresh, ): result = await delete_model_endpoint( model_info=ModelInfoDelete(id=deleted_id), @@ -2872,7 +2872,7 @@ class TestDeleteTeamBYOKModelGhost: patch(f"{_PS}.llm_router", mock_router), patch(f"{_PS}.proxy_logging_obj", MagicMock()), patch(f"{_PS}.user_api_key_cache", MagicMock()), - patch(f"{_MOD}._refresh_cached_team", new=AsyncMock()) as mock_refresh, + patch(f"{_MOD}.refresh_cached_team", new=AsyncMock()) as mock_refresh, ): result = await delete_model_endpoint( model_info=ModelInfoDelete(id=model_id), @@ -2948,7 +2948,7 @@ class TestDeleteTeamBYOKModelGhost: patch(f"{_PS}.llm_router", mock_router), patch(f"{_PS}.proxy_logging_obj", MagicMock()), patch(f"{_PS}.user_api_key_cache", MagicMock()), - patch(f"{_MOD}._refresh_cached_team", new=AsyncMock()), + patch(f"{_MOD}.refresh_cached_team", new=AsyncMock()), ): result = await delete_model_endpoint( model_info=ModelInfoDelete(id=model_id), @@ -3023,7 +3023,7 @@ class TestDeleteModelTeamAuth: patch(f"{_PS}.llm_router", MagicMock()), patch(f"{_PS}.proxy_logging_obj", MagicMock()), patch(f"{_PS}.user_api_key_cache", MagicMock()), - patch(f"{_MOD}._refresh_cached_team", new=AsyncMock()), + patch(f"{_MOD}.refresh_cached_team", new=AsyncMock()), ): result = await delete_model_endpoint( model_info=ModelInfoDelete(id=model_id), @@ -3059,7 +3059,7 @@ class TestDeleteModelTeamAuth: patch(f"{_PS}.llm_router", MagicMock()), patch(f"{_PS}.proxy_logging_obj", MagicMock()), patch(f"{_PS}.user_api_key_cache", MagicMock()), - patch(f"{_MOD}._refresh_cached_team", new=AsyncMock()), + patch(f"{_MOD}.refresh_cached_team", new=AsyncMock()), ): with pytest.raises(ProxyException) as exc_info: await delete_model_endpoint( @@ -3124,7 +3124,7 @@ class TestDeleteModelTeamAuth: patch(f"{_PS}.llm_router", MagicMock()), patch(f"{_PS}.proxy_logging_obj", MagicMock()), patch(f"{_PS}.user_api_key_cache", MagicMock()), - patch(f"{_MOD}._refresh_cached_team", new=AsyncMock()), + patch(f"{_MOD}.refresh_cached_team", new=AsyncMock()), ): with pytest.raises(ProxyException) as exc_info: await delete_model_endpoint( @@ -5201,7 +5201,7 @@ class TestConcurrentModelWritesDoNotEvictEachOther: depth -= 1 return ReconcileOutcome(still_desired=frozenset(), live_after=frozenset()) - monkeypatch.setattr(ProxyConfig, "_add_deployment_locked", fake_locked) + monkeypatch.setattr(ProxyConfig, "add_deployment_locked", fake_locked) config = ProxyConfig() await asyncio.gather( @@ -5244,7 +5244,7 @@ class TestConcurrentModelWritesDoNotEvictEachOther: async def fake_locked(self, **kwargs): return ReconcileOutcome(still_desired=frozenset({"m-db"}), live_after=frozenset({"m-db"})) - monkeypatch.setattr(ProxyConfig, "_add_deployment_locked", fake_locked) + monkeypatch.setattr(ProxyConfig, "add_deployment_locked", fake_locked) outcome = await asyncio.wait_for(clear_cache(), timeout=5) @@ -6408,7 +6408,7 @@ class TestStrategyRouterWriteValidation: lock holder waiting for a connection the waiters are occupying.""" from contextlib import asynccontextmanager - from litellm.proxy.management_endpoints.model_management_endpoints import _add_team_model_to_db + from litellm.proxy.management_endpoints.model_management_endpoints import add_team_model_to_db from litellm.types.router import ModelInfo events: list[str] = [] @@ -6438,7 +6438,7 @@ class TestStrategyRouterWriteValidation: side_effect=team_model_add, ), ): - result = await _add_team_model_to_db( + result = await add_team_model_to_db( model_params=deployment, user_api_key_dict=UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN), prisma_client=MagicMock(), @@ -8341,7 +8341,7 @@ class TestAddModelToDbBlocked: @pytest.mark.asyncio async def test_add_model_to_db_writes_blocked_true(self): from litellm.proxy.management_endpoints.model_management_endpoints import ( - _add_model_to_db, + add_model_to_db, ) mock_prisma = MagicMock() @@ -8351,7 +8351,7 @@ class TestAddModelToDbBlocked: with patch( # test-quality-ok: the proxy wiring under test is what this patches "litellm.proxy.proxy_server.master_key", "sk-test-master" ): # test-quality-ok: the proxy wiring under test is what this patches - await _add_model_to_db( + await add_model_to_db( model_params=self._deployment(True), user_api_key_dict=admin, prisma_client=mock_prisma ) @@ -8361,7 +8361,7 @@ class TestAddModelToDbBlocked: @pytest.mark.asyncio async def test_add_model_to_db_writes_blocked_false(self): from litellm.proxy.management_endpoints.model_management_endpoints import ( - _add_model_to_db, + add_model_to_db, ) mock_prisma = MagicMock() @@ -8371,7 +8371,7 @@ class TestAddModelToDbBlocked: with patch( # test-quality-ok: the proxy wiring under test is what this patches "litellm.proxy.proxy_server.master_key", "sk-test-master" ): # test-quality-ok: the proxy wiring under test is what this patches - await _add_model_to_db( + await add_model_to_db( model_params=self._deployment(False), user_api_key_dict=admin, prisma_client=mock_prisma ) @@ -8383,7 +8383,7 @@ class TestAddModelToDbBlocked: """None means "don't set it" -- the Prisma column defaults to False -- not "explicitly unblocked", so the key must be absent from the write entirely.""" from litellm.proxy.management_endpoints.model_management_endpoints import ( - _add_model_to_db, + add_model_to_db, ) mock_prisma = MagicMock() @@ -8393,7 +8393,7 @@ class TestAddModelToDbBlocked: with patch( # test-quality-ok: the proxy wiring under test is what this patches "litellm.proxy.proxy_server.master_key", "sk-test-master" ): # test-quality-ok: the proxy wiring under test is what this patches - await _add_model_to_db( + await add_model_to_db( model_params=self._deployment(None), user_api_key_dict=admin, prisma_client=mock_prisma ) @@ -8944,7 +8944,7 @@ class TestOneCredentialFeedsManyModelsNoWifCopy: async def test_two_discovered_models_share_the_credential_reference_only(self): from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper from litellm.proxy.management_endpoints.model_management_endpoints import ( - _add_model_to_db, + add_model_to_db, ) from litellm.types.router import ModelInfo @@ -8957,7 +8957,7 @@ class TestOneCredentialFeedsManyModelsNoWifCopy: "litellm.proxy.proxy_server.master_key", "sk-test-master" ), patch( # test-quality-ok: the proxy wiring under test is what this patches - "litellm.proxy.common_utils.encrypt_decrypt_utils._get_salt_key", return_value="sk-test-master" + "litellm.proxy.common_utils.encrypt_decrypt_utils.get_salt_key", return_value="sk-test-master" ), ): for i, discovered_id in enumerate(["claude-a", "claude-b"]): @@ -8969,7 +8969,7 @@ class TestOneCredentialFeedsManyModelsNoWifCopy: model_info=ModelInfo(id=f"dep-shared-{i}"), blocked=False, ) - await _add_model_to_db(model_params=model_params, user_api_key_dict=admin, prisma_client=mock_prisma) + await add_model_to_db(model_params=model_params, user_api_key_dict=admin, prisma_client=mock_prisma) assert mock_prisma.db.litellm_proxymodeltable.create.await_count == 2 for call in mock_prisma.db.litellm_proxymodeltable.create.await_args_list: @@ -9335,7 +9335,7 @@ class TestFederationGateScopesToWhatTheWriteTouches: patch(f"{_PS}.llm_router", MagicMock()), # test-quality-ok: proxy wiring under test patch(f"{_PS}.proxy_logging_obj", MagicMock()), # test-quality-ok: proxy wiring under test patch(f"{_PS}.user_api_key_cache", MagicMock()), # test-quality-ok: proxy wiring under test - patch(f"{_MOD}._refresh_cached_team", new=AsyncMock()), # test-quality-ok: proxy wiring under test + patch(f"{_MOD}.refresh_cached_team", new=AsyncMock()), # test-quality-ok: proxy wiring under test ): result = await delete_model_endpoint( model_info=ModelInfoDelete(id="m1"), diff --git a/tests/unit/proxy/management_endpoints/test_org_admin_team_access.py b/tests/unit/proxy/management_endpoints/test_org_admin_team_access.py index aab67dccf1d..4a4aaf13154 100644 --- a/tests/unit/proxy/management_endpoints/test_org_admin_team_access.py +++ b/tests/unit/proxy/management_endpoints/test_org_admin_team_access.py @@ -41,9 +41,7 @@ def _make_team(team_id="team-1", organization_id="org-1") -> LiteLLM_TeamTable: ) -def _make_user_key( - user_id="org-admin-user", role=LitellmUserRoles.INTERNAL_USER.value -) -> UserAPIKeyAuth: +def _make_user_key(user_id="org-admin-user", role=LitellmUserRoles.INTERNAL_USER.value) -> UserAPIKeyAuth: return UserAPIKeyAuth(user_id=user_id, user_role=role) @@ -57,9 +55,7 @@ def _make_membership(user_id, org_id, role="org_admin"): ) -def _make_caller_user( - user_id="org-admin-user", org_id="org-1", org_role="org_admin" -) -> LiteLLM_UserTable: +def _make_caller_user(user_id="org-admin-user", org_id="org-1", org_role="org_admin") -> LiteLLM_UserTable: return LiteLLM_UserTable( user_id=user_id, organization_memberships=[_make_membership(user_id, org_id, org_role)], @@ -76,9 +72,7 @@ def _patch_org_admin_deps(get_user_return): ), patch("litellm.proxy.proxy_server.prisma_client", MagicMock(), create=True), patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock(), create=True), - patch( - "litellm.proxy.proxy_server.user_api_key_cache", MagicMock(), create=True - ), + patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock(), create=True), ) @@ -133,9 +127,7 @@ class TestValidateMembership: team = _make_team(organization_id="org-1") key = _make_user_key(user_id="random-user") - caller = _make_caller_user( - user_id="random-user", org_id="org-2", org_role="user" - ) + caller = _make_caller_user(user_id="random-user", org_id="org-2", org_role="user") p1, p2, p3, p4 = _patch_org_admin_deps(caller) with p1, p2, p3, p4: @@ -150,9 +142,7 @@ class TestValidateMembership: ) team = _make_team(team_id="team-1") - key = UserAPIKeyAuth( - team_id="team-1", user_role=LitellmUserRoles.INTERNAL_USER.value - ) + key = UserAPIKeyAuth(team_id="team-1", user_role=LitellmUserRoles.INTERNAL_USER.value) await validate_membership(user_api_key_dict=key, team_table=team) @@ -168,55 +158,49 @@ class TestUserIsOrgAdminRouteCheck: """ def test_no_candidate_org_ids_returns_false(self): - from litellm.proxy.auth.auth_checks_organization import _user_is_org_admin + from litellm.proxy.auth.auth_checks_organization import user_is_org_admin user = LiteLLM_UserTable( user_id="org-admin-user", organization_memberships=[_make_membership("org-admin-user", "org-1")], ) - result = _user_is_org_admin(request_data={}, user_object=user) + result = user_is_org_admin(request_data={}, user_object=user) assert result is False, "Must NOT grant blanket access when no org in request" def test_matching_org_id_returns_true(self): - from litellm.proxy.auth.auth_checks_organization import _user_is_org_admin + from litellm.proxy.auth.auth_checks_organization import user_is_org_admin user = LiteLLM_UserTable( user_id="org-admin-user", organization_memberships=[_make_membership("org-admin-user", "org-1")], ) - result = _user_is_org_admin( - request_data={"organization_id": "org-1"}, user_object=user - ) + result = user_is_org_admin(request_data={"organization_id": "org-1"}, user_object=user) assert result is True def test_non_matching_org_id_returns_false(self): - from litellm.proxy.auth.auth_checks_organization import _user_is_org_admin + from litellm.proxy.auth.auth_checks_organization import user_is_org_admin user = LiteLLM_UserTable( user_id="org-admin-user", organization_memberships=[_make_membership("org-admin-user", "org-1")], ) - result = _user_is_org_admin( - request_data={"organization_id": "org-99"}, user_object=user - ) + result = user_is_org_admin(request_data={"organization_id": "org-99"}, user_object=user) assert result is False def test_organizations_list_field(self): - from litellm.proxy.auth.auth_checks_organization import _user_is_org_admin + from litellm.proxy.auth.auth_checks_organization import user_is_org_admin user = LiteLLM_UserTable( user_id="org-admin-user", organization_memberships=[_make_membership("org-admin-user", "org-1")], ) - result = _user_is_org_admin( - request_data={"organizations": ["org-1"]}, user_object=user - ) + result = user_is_org_admin(request_data={"organizations": ["org-1"]}, user_object=user) assert result is True def test_none_user_object_returns_false(self): - from litellm.proxy.auth.auth_checks_organization import _user_is_org_admin + from litellm.proxy.auth.auth_checks_organization import user_is_org_admin - result = _user_is_org_admin(request_data={}, user_object=None) + result = user_is_org_admin(request_data={}, user_object=None) assert result is False def test_user_list_in_self_managed_routes(self): diff --git a/tests/unit/proxy/management_endpoints/test_organization_endpoints.py b/tests/unit/proxy/management_endpoints/test_organization_endpoints.py index 7d3bd049a03..a98f1344279 100644 --- a/tests/unit/proxy/management_endpoints/test_organization_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_organization_endpoints.py @@ -108,7 +108,7 @@ async def test_get_organization_daily_activity_admin_param_passing(monkeypatch): # Admin view -> skip membership restriction monkeypatch.setattr( - "litellm.proxy.management_endpoints.organization_endpoints._user_has_admin_view", + "litellm.proxy.management_endpoints.organization_endpoints.user_api_key_has_admin_view", lambda _: True, ) @@ -175,7 +175,7 @@ async def test_get_organization_daily_activity_non_admin_defaults_to_admin_orgs( # Non-admin view monkeypatch.setattr( - "litellm.proxy.management_endpoints.organization_endpoints._user_has_admin_view", + "litellm.proxy.management_endpoints.organization_endpoints.user_api_key_has_admin_view", lambda _: False, ) @@ -227,7 +227,7 @@ async def test_get_organization_daily_activity_non_admin_unauthorized_org_raises # Non-admin view monkeypatch.setattr( - "litellm.proxy.management_endpoints.organization_endpoints._user_has_admin_view", + "litellm.proxy.management_endpoints.organization_endpoints.user_api_key_has_admin_view", lambda _: False, ) @@ -853,7 +853,7 @@ async def _run_update_organization_v2( mock_prisma_client.call_order = call_order monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) - monkeypatch.setattr(organization_endpoints, "_verify_org_access", AsyncMock()) + monkeypatch.setattr(organization_endpoints, "verify_org_access", AsyncMock()) auth = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin-1") await update_organization_v2( @@ -1019,7 +1019,7 @@ async def test_v2_rejects_caller_without_org_access(monkeypatch): mock_prisma_client = AsyncMock() monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) - monkeypatch.setattr(organization_endpoints, "_user_has_admin_view", lambda _: False) + monkeypatch.setattr(organization_endpoints, "user_api_key_has_admin_view", lambda _: False) caller = MagicMock() caller.organization_memberships = [] @@ -1110,7 +1110,7 @@ async def test_v2_rejects_empty_object_permission(monkeypatch): mock_prisma_client = AsyncMock() monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) - monkeypatch.setattr(organization_endpoints, "_verify_org_access", AsyncMock()) + monkeypatch.setattr(organization_endpoints, "verify_org_access", AsyncMock()) auth = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin-1") with pytest.raises(HTTPException) as exc: @@ -1178,7 +1178,7 @@ async def _run_legacy_update_organization( mock_prisma_client.db.litellm_budgettable.update = AsyncMock() monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) - monkeypatch.setattr(organization_endpoints, "_verify_org_access", AsyncMock()) + monkeypatch.setattr(organization_endpoints, "verify_org_access", AsyncMock()) request = MagicMock() request.json = AsyncMock(return_value=body) @@ -1299,7 +1299,7 @@ async def test_get_organization_daily_activity_non_admin_without_org_admin_role_ monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) monkeypatch.setattr( - "litellm.proxy.management_endpoints.organization_endpoints._user_has_admin_view", + "litellm.proxy.management_endpoints.organization_endpoints.user_api_key_has_admin_view", lambda _: False, ) diff --git a/tests/unit/proxy/management_endpoints/test_policy_endpoints.py b/tests/unit/proxy/management_endpoints/test_policy_endpoints.py index f5c4e9b69ff..e334305dd98 100644 --- a/tests/unit/proxy/management_endpoints/test_policy_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_policy_endpoints.py @@ -648,47 +648,47 @@ class TestApplyPoliciesDirectGuardrailNames: # Tests for competitor enrichment helper functions # --------------------------------------------------------------------------- from litellm.proxy.management_endpoints.policy_endpoints import ( - _build_all_names_per_competitor, - _build_comparison_blocked_words, - _build_competitor_guardrail_definitions, - _build_name_blocked_words, - _build_recommendation_blocked_words, - _build_refinement_prompt, - _clean_competitor_line, - _parse_variations_response, + build_all_names_per_competitor, + build_comparison_blocked_words, + build_competitor_guardrail_definitions, + build_name_blocked_words, + build_recommendation_blocked_words, + build_refinement_prompt, + clean_competitor_line, + parse_variations_response, ) class TestCleanCompetitorLine: - """Tests for _clean_competitor_line.""" + """Tests for clean_competitor_line.""" def test_strips_bullets_and_dashes(self): - assert _clean_competitor_line("- United Airlines") == "United Airlines" - assert _clean_competitor_line(" - JetBlue ") == "JetBlue" + assert clean_competitor_line("- United Airlines") == "United Airlines" + assert clean_competitor_line(" - JetBlue ") == "JetBlue" def test_strips_trailing_punctuation(self): - assert _clean_competitor_line("Delta Airlines.") == "Delta Airlines" - assert _clean_competitor_line("Southwest)") == "Southwest" + assert clean_competitor_line("Delta Airlines.") == "Delta Airlines" + assert clean_competitor_line("Southwest)") == "Southwest" def test_returns_none_for_empty(self): - assert _clean_competitor_line("") is None - assert _clean_competitor_line(" ") is None + assert clean_competitor_line("") is None + assert clean_competitor_line(" ") is None def test_returns_none_for_single_char(self): - assert _clean_competitor_line("A") is None - assert _clean_competitor_line(" - ") is None + assert clean_competitor_line("A") is None + assert clean_competitor_line(" - ") is None def test_plain_name(self): - assert _clean_competitor_line("Qatar Airways") == "Qatar Airways" + assert clean_competitor_line("Qatar Airways") == "Qatar Airways" class TestParseVariationsResponse: - """Tests for _parse_variations_response.""" + """Tests for parse_variations_response.""" def test_parses_standard_format(self): raw = "Delta Airlines: Delta Air Lines, DeltaAirlines, Delta\nUnited Airlines: United, UAL" competitors = ["Delta Airlines", "United Airlines"] - result = _parse_variations_response(raw, competitors) + result = parse_variations_response(raw, competitors) assert "Delta Airlines" in result assert "Delta Air Lines" in result["Delta Airlines"] assert "United" in result["United Airlines"] @@ -696,71 +696,71 @@ class TestParseVariationsResponse: def test_case_insensitive_matching(self): raw = "delta airlines: Delta Air Lines, DeltaAirlines" competitors = ["Delta Airlines"] - result = _parse_variations_response(raw, competitors) + result = parse_variations_response(raw, competitors) assert "Delta Airlines" in result assert len(result["Delta Airlines"]) == 2 def test_skips_lines_without_colon(self): raw = "This is a header\nDelta Airlines: Delta Air Lines" competitors = ["Delta Airlines"] - result = _parse_variations_response(raw, competitors) + result = parse_variations_response(raw, competitors) assert len(result) == 1 def test_skips_unknown_competitors(self): raw = "Unknown Corp: Foo, Bar\nDelta Airlines: Delta" competitors = ["Delta Airlines"] - result = _parse_variations_response(raw, competitors) + result = parse_variations_response(raw, competitors) assert "Unknown Corp" not in result assert "Delta Airlines" in result def test_filters_out_self_reference(self): raw = "Delta Airlines: Delta Airlines, Delta Air Lines" competitors = ["Delta Airlines"] - result = _parse_variations_response(raw, competitors) + result = parse_variations_response(raw, competitors) # "Delta Airlines" should be filtered out (same as canonical) assert "Delta Airlines" not in result["Delta Airlines"] assert "Delta Air Lines" in result["Delta Airlines"] def test_empty_input(self): - assert _parse_variations_response("", []) == {} + assert parse_variations_response("", []) == {} class TestBuildRefinementPrompt: - """Tests for _build_refinement_prompt.""" + """Tests for build_refinement_prompt.""" def test_includes_brand_name(self): - prompt = _build_refinement_prompt("add 10 more", ["Delta"], "Emirates") + prompt = build_refinement_prompt("add 10 more", ["Delta"], "Emirates") assert "Emirates" in prompt def test_includes_existing_competitors(self): - prompt = _build_refinement_prompt("add more", ["Delta", "United"], "Emirates") + prompt = build_refinement_prompt("add more", ["Delta", "United"], "Emirates") assert "Delta" in prompt assert "United" in prompt def test_includes_instruction(self): - prompt = _build_refinement_prompt("add 10 from Asia", ["Delta"], "Emirates") + prompt = build_refinement_prompt("add 10 from Asia", ["Delta"], "Emirates") assert "add 10 from Asia" in prompt def test_asks_for_new_names_only(self): - prompt = _build_refinement_prompt("add more", ["Delta"], "Emirates") + prompt = build_refinement_prompt("add more", ["Delta"], "Emirates") assert "NEW" in prompt class TestBuildAllNamesPerCompetitor: - """Tests for _build_all_names_per_competitor.""" + """Tests for build_all_names_per_competitor.""" def test_includes_canonical_and_variations(self): - result = _build_all_names_per_competitor( + result = build_all_names_per_competitor( ["Delta Airlines"], {"Delta Airlines": ["Delta", "DeltaAir"]} ) assert result["Delta Airlines"] == ["Delta Airlines", "Delta", "DeltaAir"] def test_no_variations(self): - result = _build_all_names_per_competitor(["Delta Airlines"], {}) + result = build_all_names_per_competitor(["Delta Airlines"], {}) assert result["Delta Airlines"] == ["Delta Airlines"] def test_multiple_competitors(self): - result = _build_all_names_per_competitor( + result = build_all_names_per_competitor( ["Delta", "United"], {"Delta": ["DL"], "United": ["UA"]}, ) @@ -770,11 +770,11 @@ class TestBuildAllNamesPerCompetitor: class TestBuildNameBlockedWords: - """Tests for _build_name_blocked_words.""" + """Tests for build_name_blocked_words.""" def test_basic_output(self): all_names = {"Delta": ["Delta", "DL"]} - result = _build_name_blocked_words(["Delta"], all_names) + result = build_name_blocked_words(["Delta"], all_names) keywords = [r["keyword"] for r in result] assert "Delta" in keywords assert "DL" in keywords @@ -782,18 +782,18 @@ class TestBuildNameBlockedWords: def test_descriptions_differ_for_variations(self): all_names = {"Delta": ["Delta", "DL"]} - result = _build_name_blocked_words(["Delta"], all_names) + result = build_name_blocked_words(["Delta"], all_names) descs = {r["keyword"]: r["description"] for r in result} assert "Competitor: Delta" == descs["Delta"] assert "variation" in descs["DL"].lower() class TestBuildRecommendationBlockedWords: - """Tests for _build_recommendation_blocked_words.""" + """Tests for build_recommendation_blocked_words.""" def test_generates_prefix_combinations(self): all_names = {"Delta": ["Delta"]} - result = _build_recommendation_blocked_words(["Delta"], all_names) + result = build_recommendation_blocked_words(["Delta"], all_names) keywords = [r["keyword"] for r in result] assert "try Delta" in keywords assert "use Delta" in keywords @@ -802,23 +802,23 @@ class TestBuildRecommendationBlockedWords: def test_includes_variations(self): all_names = {"Delta": ["Delta", "DL"]} - result = _build_recommendation_blocked_words(["Delta"], all_names) + result = build_recommendation_blocked_words(["Delta"], all_names) keywords = [r["keyword"] for r in result] assert "try DL" in keywords class TestBuildComparisonBlockedWords: - """Tests for _build_comparison_blocked_words.""" + """Tests for build_comparison_blocked_words.""" def test_generates_competitor_comparisons(self): all_names = {"Delta": ["Delta"]} - result = _build_comparison_blocked_words(["Delta"], all_names, "Emirates") + result = build_comparison_blocked_words(["Delta"], all_names, "Emirates") keywords = [r["keyword"] for r in result] assert "Delta is better" in keywords def test_generates_brand_comparisons_once(self): all_names = {"Delta": ["Delta"], "United": ["United"]} - result = _build_comparison_blocked_words( + result = build_comparison_blocked_words( ["Delta", "United"], all_names, "Emirates" ) keywords = [r["keyword"] for r in result] @@ -828,13 +828,13 @@ class TestBuildComparisonBlockedWords: def test_includes_variation_comparisons(self): all_names = {"Delta": ["Delta", "DL"]} - result = _build_comparison_blocked_words(["Delta"], all_names, "Emirates") + result = build_comparison_blocked_words(["Delta"], all_names, "Emirates") keywords = [r["keyword"] for r in result] assert "DL is better" in keywords class TestBuildCompetitorGuardrailDefinitions: - """Tests for _build_competitor_guardrail_definitions.""" + """Tests for build_competitor_guardrail_definitions.""" def test_populates_blocked_words_for_known_guardrail_names(self): definitions = [ @@ -847,7 +847,7 @@ class TestBuildCompetitorGuardrailDefinitions: "litellm_params": {"blocked_words": []}, }, ] - result = _build_competitor_guardrail_definitions( + result = build_competitor_guardrail_definitions( definitions, ["Delta"], "Emirates", {"Delta": ["DL"]} ) # Name blocker should have entries @@ -871,7 +871,7 @@ class TestBuildCompetitorGuardrailDefinitions: "litellm_params": {"blocked_words": ["original"]}, }, ] - result = _build_competitor_guardrail_definitions( + result = build_competitor_guardrail_definitions( definitions, ["Delta"], "Emirates" ) assert result[0]["litellm_params"]["blocked_words"] == ["original"] @@ -883,7 +883,7 @@ class TestBuildCompetitorGuardrailDefinitions: "litellm_params": {"blocked_words": []}, }, ] - _build_competitor_guardrail_definitions(definitions, ["Delta"], "Emirates") + build_competitor_guardrail_definitions(definitions, ["Delta"], "Emirates") # Original should be unchanged assert definitions[0]["litellm_params"]["blocked_words"] == [] @@ -898,7 +898,7 @@ class TestBuildCompetitorGuardrailDefinitions: "litellm_params": {"blocked_words": []}, }, ] - result = _build_competitor_guardrail_definitions( + result = build_competitor_guardrail_definitions( definitions, ["Delta"], "Emirates" ) for defn in result: diff --git a/tests/unit/proxy/management_endpoints/test_prompt_cache_prediction.py b/tests/unit/proxy/management_endpoints/test_prompt_cache_prediction.py index 987cacf7676..1ae978e9a2e 100644 --- a/tests/unit/proxy/management_endpoints/test_prompt_cache_prediction.py +++ b/tests/unit/proxy/management_endpoints/test_prompt_cache_prediction.py @@ -13,8 +13,8 @@ import litellm from litellm.caching.caching import DualCache from litellm.integrations.custom_logger import CustomLogger from litellm.proxy._types import ProxyException, UserAPIKeyAuth -from litellm.proxy.common_utils.http_parsing_utils import _read_request_body, _safe_set_request_parsed_body -from litellm.proxy.hooks.parallel_request_limiter_v3 import _PROXY_MaxParallelRequestsHandler_v3 +from litellm.proxy.common_utils.http_parsing_utils import read_request_body, safe_set_request_parsed_body +from litellm.proxy.hooks.parallel_request_limiter_v3 import PROXY_MaxParallelRequestsHandler_v3 from litellm.llms.anthropic.prompt_cache_prediction import PromptPrefix, cache_scope, parse_prompt from litellm.proxy.hooks.prompt_cache_prediction import ( CacheObservation, @@ -224,7 +224,7 @@ def _app( if caller is not None: usage_cache: Final = InternalUsageCache(cache) configured_limiter: Final = ( - _PROXY_MaxParallelRequestsHandler_v3(usage_cache) if isinstance(limiter, str) else limiter + PROXY_MaxParallelRequestsHandler_v3(usage_cache) if isinstance(limiter, str) else limiter ) monkeypatch.setattr(proxy_server, "proxy_logging_obj", _ProxyLogging(usage_cache, configured_limiter)) app.dependency_overrides[endpoint.user_api_key_auth] = lambda: caller @@ -384,7 +384,7 @@ async def test_missing_or_unsupported_limiter_returns_unknown_before_counting( @pytest.mark.asyncio async def test_occupied_parallel_capacity_rejects_before_provider_count(monkeypatch: pytest.MonkeyPatch) -> None: cache: Final = DualCache() - limiter: Final = _PROXY_MaxParallelRequestsHandler_v3(InternalUsageCache(cache)) + limiter: Final = PROXY_MaxParallelRequestsHandler_v3(InternalUsageCache(cache)) caller: Final = UserAPIKeyAuth(api_key=_CALLER, max_parallel_requests=1) app: Final = _app(monkeypatch, cache, caller=caller, counts=_unexpected_count, limiter=limiter) async with limiter.request_capacity(caller, "opus"): @@ -428,8 +428,8 @@ async def test_each_count_preserves_auth_cached_request_tag_limits( return await Counts()(model, api_key, body) async def authenticated_request(request: Request) -> UserAPIKeyAuth: - data: Final = await _read_request_body(request) - _safe_set_request_parsed_body(request, {**data, metadata_key: {"tags": ["cache-cost"]}}) + data: Final = await read_request_body(request) + safe_set_request_parsed_body(request, {**data, metadata_key: {"tags": ["cache-cost"]}}) return caller app: Final = _app(monkeypatch, DualCache(), caller=caller, counts=count) diff --git a/tests/unit/proxy/management_endpoints/test_ptu_model_settings.py b/tests/unit/proxy/management_endpoints/test_ptu_model_settings.py index d35b77f732c..5b4ce52baa5 100644 --- a/tests/unit/proxy/management_endpoints/test_ptu_model_settings.py +++ b/tests/unit/proxy/management_endpoints/test_ptu_model_settings.py @@ -16,7 +16,7 @@ from litellm.proxy._types import ( UserAPIKeyAuth, ) from litellm.litellm_core_utils.llm_cost_calc.utils import generic_cost_per_token -from litellm.proxy.auth.auth_checks import _is_model_cost_zero +from litellm.proxy.auth.auth_checks import is_model_cost_zero from litellm.llms.gemini.cost_calculator import cost_per_web_search_request from litellm.proxy.management_endpoints.model_management_endpoints import ( _PTU_ZEROED_PRICING_FIELDS, @@ -677,8 +677,8 @@ class TestAddNewModelPtuGate: f"{endpoints}.ModelManagementAuthChecks.can_user_make_model_call", AsyncMock(return_value=True), ), - patch(f"{endpoints}._add_model_to_db", add_model_to_db), - patch(f"{endpoints}._add_team_model_to_db", add_team_model_to_db), + patch(f"{endpoints}.add_model_to_db", add_model_to_db), + patch(f"{endpoints}.add_team_model_to_db", add_team_model_to_db), ] @staticmethod @@ -1090,8 +1090,8 @@ class TestPtuDeploymentsAreNotBilledPerToken: ) ) router = Router(model_list=[priced.to_json(exclude_none=True)]) - assert _is_model_cost_zero(model="model_name_team-1_dep-ptu", llm_router=router) is False - assert _is_model_cost_zero(model="ptu-model", llm_router=router) is False + assert is_model_cost_zero(model="model_name_team-1_dep-ptu", llm_router=router) is False + assert is_model_cost_zero(model="ptu-model", llm_router=router) is False def test_an_unrelated_patch_heals_a_deployment_stored_before_this_rule(self): """Both blobs, because litellm_params wins over model_info wherever the two are merged.""" diff --git a/tests/unit/proxy/management_endpoints/test_session_endpoints.py b/tests/unit/proxy/management_endpoints/test_session_endpoints.py index d5960a88937..afde88c1c66 100644 --- a/tests/unit/proxy/management_endpoints/test_session_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_session_endpoints.py @@ -65,7 +65,7 @@ async def test_session_logout_revokes_presented_session(): p2, p3, patch( - "litellm.proxy.management_endpoints.session_endpoints._persist_deleted_verification_tokens", + "litellm.proxy.management_endpoints.session_endpoints.persist_deleted_verification_tokens", persist_mock, ), patch( @@ -100,7 +100,7 @@ async def test_session_logout_clears_token_cookie(): p2, p3, patch( - "litellm.proxy.management_endpoints.session_endpoints._persist_deleted_verification_tokens", + "litellm.proxy.management_endpoints.session_endpoints.persist_deleted_verification_tokens", AsyncMock(), ), patch( @@ -194,7 +194,7 @@ async def test_revoke_ui_session_keys_revokes_all_and_broadcasts(): p2, p3, patch( - "litellm.proxy.management_endpoints.session_endpoints._persist_deleted_verification_tokens", + "litellm.proxy.management_endpoints.session_endpoints.persist_deleted_verification_tokens", persist_mock, ), patch( @@ -228,7 +228,7 @@ async def test_revoke_ui_session_keys_keeps_callers_session(): p2, p3, patch( - "litellm.proxy.management_endpoints.session_endpoints._persist_deleted_verification_tokens", + "litellm.proxy.management_endpoints.session_endpoints.persist_deleted_verification_tokens", AsyncMock(), ), patch( @@ -275,7 +275,7 @@ async def test_revoke_ui_session_keys_failure_is_swallowed(): p2, p3, patch( - "litellm.proxy.management_endpoints.session_endpoints._persist_deleted_verification_tokens", + "litellm.proxy.management_endpoints.session_endpoints.persist_deleted_verification_tokens", AsyncMock(), ), ): diff --git a/tests/unit/proxy/management_endpoints/test_team_callback_endpoints.py b/tests/unit/proxy/management_endpoints/test_team_callback_endpoints.py index 331e3a7983c..f576bac4742 100644 --- a/tests/unit/proxy/management_endpoints/test_team_callback_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_team_callback_endpoints.py @@ -89,7 +89,7 @@ def stub_team_cache_refresh(): test_disable_team_logging_refreshes_cached_team. """ with patch( - "litellm.proxy.management_endpoints.team_callback_endpoints._refresh_cached_team", + "litellm.proxy.management_endpoints.team_callback_endpoints.refresh_cached_team", new_callable=AsyncMock, ) as refresh: yield refresh @@ -614,7 +614,7 @@ async def test_get_team_callbacks_decrypts_vars_stored_under_non_sensitive_keys( encrypted at rest under a key that later stops being masked on read. Without the decrypt step that value comes back as an unusable litellm_enc:: blob. """ - from litellm.proxy.common_utils.callback_utils import _CALLBACK_VAR_ENCRYPTED_PREFIX, is_sensitive_callback_key + from litellm.proxy.common_utils.callback_utils import CALLBACK_VAR_ENCRYPTED_PREFIX, is_sensitive_callback_key from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_value_helper monkeypatch.setenv("LITELLM_SALT_KEY", "test-salt-32-bytes-aaaaaaaaaaaaaa") @@ -626,7 +626,7 @@ async def test_get_team_callbacks_decrypts_vars_stored_under_non_sensitive_keys( "callback_name": "langsmith", "callback_type": "success", "callback_vars": { - "langsmith_project": _CALLBACK_VAR_ENCRYPTED_PREFIX + encrypt_value_helper("tenant-project"), + "langsmith_project": CALLBACK_VAR_ENCRYPTED_PREFIX + encrypt_value_helper("tenant-project"), }, } ] @@ -655,11 +655,11 @@ async def test_get_team_callbacks_masks_values_that_fail_to_decrypt(monkeypatch) classified as sensitive it would otherwise reach the caller as an opaque blob that is indistinguishable from a real value. """ - from litellm.proxy.common_utils.callback_utils import _CALLBACK_VAR_ENCRYPTED_PREFIX + from litellm.proxy.common_utils.callback_utils import CALLBACK_VAR_ENCRYPTED_PREFIX from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_value_helper monkeypatch.setenv("LITELLM_SALT_KEY", "test-salt-32-bytes-aaaaaaaaaaaaaa") - stale = _CALLBACK_VAR_ENCRYPTED_PREFIX + encrypt_value_helper("tenant-project") + stale = CALLBACK_VAR_ENCRYPTED_PREFIX + encrypt_value_helper("tenant-project") monkeypatch.setenv("LITELLM_SALT_KEY", "test-salt-32-bytes-bbbbbbbbbbbbbb") metadata = { @@ -685,7 +685,7 @@ async def test_get_team_callbacks_masks_values_that_fail_to_decrypt(monkeypatch) assert response["data"]["success_callbacks"] == ["langsmith"] assert response["data"]["callback_vars"]["langsmith_project"] == "***REDACTED***" - assert _CALLBACK_VAR_ENCRYPTED_PREFIX not in json.dumps(response) + assert CALLBACK_VAR_ENCRYPTED_PREFIX not in json.dumps(response) @pytest.mark.asyncio @@ -752,7 +752,7 @@ async def test_disable_team_logging_stops_callbacks_registered_via_api(): the endpoint and then asks the real request-time resolver what the written row would do. """ - from litellm.proxy.litellm_pre_call_utils import _get_dynamic_logging_metadata + from litellm.proxy.litellm_pre_call_utils import get_dynamic_logging_metadata metadata = { "logging": [ @@ -780,7 +780,7 @@ async def test_disable_team_logging_stops_callbacks_registered_via_api(): written = json.loads(mock_prisma.db.litellm_teamtable.update.await_args.kwargs["data"]["metadata"]) assert written["logging"] == [] - resolved = _get_dynamic_logging_metadata( + resolved = get_dynamic_logging_metadata( UserAPIKeyAuth(api_key="hashed", team_id="team-1", team_metadata=written), proxy_config=MagicMock(**{"load_team_config.return_value": {}}), ) @@ -859,7 +859,7 @@ async def test_add_team_callbacks_refreshes_cached_team(stub_team_cache_refresh) @pytest.mark.asyncio async def test_disable_team_logging_clears_both_metadata_shapes(): """A team carrying both shapes ends up with neither active.""" - from litellm.proxy.litellm_pre_call_utils import _get_dynamic_logging_metadata + from litellm.proxy.litellm_pre_call_utils import get_dynamic_logging_metadata metadata = { "logging": [ @@ -893,7 +893,7 @@ async def test_disable_team_logging_clears_both_metadata_shapes(): assert written["callback_settings"]["success_callback"] == [] assert written["callback_settings"]["failure_callback"] == [] - resolved = _get_dynamic_logging_metadata( + resolved = get_dynamic_logging_metadata( UserAPIKeyAuth(api_key="hashed", team_id="team-1", team_metadata=written), proxy_config=MagicMock(**{"load_team_config.return_value": {}}), ) @@ -1023,7 +1023,7 @@ async def test_delete_team_callback_leaves_the_other_callback_firing(): Asks the real request-time resolver what the written row would do, the same way the disable_logging regression test does. """ - from litellm.proxy.litellm_pre_call_utils import _get_dynamic_logging_metadata + from litellm.proxy.litellm_pre_call_utils import get_dynamic_logging_metadata mock_prisma = _patch_prisma(_team_row(team_id="team-1", metadata=_two_callback_metadata())) @@ -1040,7 +1040,7 @@ async def test_delete_team_callback_leaves_the_other_callback_firing(): ) written = json.loads(mock_prisma.db.litellm_teamtable.update.await_args.kwargs["data"]["metadata"]) - resolved = _get_dynamic_logging_metadata( + resolved = get_dynamic_logging_metadata( UserAPIKeyAuth(api_key="hashed", team_id="team-1", team_metadata=written), proxy_config=MagicMock(**{"load_team_config.return_value": {}}), ) @@ -1219,7 +1219,7 @@ async def test_delete_team_callback_keeps_last_removal_from_reviving_legacy_shap dropping the key would fall through to a legacy callback_settings block and silently re-enable a destination the caller just removed. """ - from litellm.proxy.litellm_pre_call_utils import _get_dynamic_logging_metadata + from litellm.proxy.litellm_pre_call_utils import get_dynamic_logging_metadata metadata = { "logging": [ @@ -1253,7 +1253,7 @@ async def test_delete_team_callback_keeps_last_removal_from_reviving_legacy_shap assert written["logging"] == [] assert response.data.success_callbacks == () - resolved = _get_dynamic_logging_metadata( + resolved = get_dynamic_logging_metadata( UserAPIKeyAuth(api_key="hashed", team_id="team-1", team_metadata=written), proxy_config=MagicMock(**{"load_team_config.return_value": {}}), ) diff --git a/tests/unit/proxy/management_endpoints/test_team_endpoints.py b/tests/unit/proxy/management_endpoints/test_team_endpoints.py index 6c5447fce64..00f624ea518 100644 --- a/tests/unit/proxy/management_endpoints/test_team_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_team_endpoints.py @@ -2313,7 +2313,7 @@ async def test_team_model_add_delete_refresh_team_cache(endpoint_name): patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache, patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_logging, patch( - "litellm.proxy.management_endpoints.team_endpoints._cache_team_object", + "litellm.proxy.management_endpoints.team_endpoints.cache_team_object", new_callable=AsyncMock, ) as mock_cache_team, ): @@ -2492,7 +2492,7 @@ async def test_team_write_404s_when_row_vanishes_before_update(endpoint_name): patch("litellm.proxy.proxy_server.user_api_key_cache"), # test-quality-ok: proxy_server module global is the endpoint's only injection point patch("litellm.proxy.proxy_server.proxy_logging_obj"), # test-quality-ok: proxy_server module global is the endpoint's only injection point patch( # test-quality-ok: stubs the cache write so the test observes only the DB result handling - "litellm.proxy.management_endpoints.team_endpoints._cache_team_object", + "litellm.proxy.management_endpoints.team_endpoints.cache_team_object", new_callable=AsyncMock, ), patch( # test-quality-ok: stubs the collaborator so the test pins the endpoint's own error contract @@ -2545,7 +2545,8 @@ async def test_update_team_team_member_budget_not_passed_to_db( patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_logging, patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), patch( - "litellm.proxy.management_endpoints.team_endpoints._cache_team_object" + "litellm.proxy.management_endpoints.team_endpoints.cache_team_object", + new_callable=AsyncMock, ) as mock_cache_team, patch( "litellm.proxy.management_endpoints.team_endpoints.TeamMemberBudgetHandler.upsert_team_member_budget_table" @@ -3113,7 +3114,8 @@ async def test_update_team_with_team_member_budget_duration( patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_logging, patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), patch( - "litellm.proxy.management_endpoints.team_endpoints._cache_team_object" + "litellm.proxy.management_endpoints.team_endpoints.cache_team_object", + new_callable=AsyncMock, ) as mock_cache_team, patch( "litellm.proxy.management_endpoints.team_endpoints.TeamMemberBudgetHandler.upsert_team_member_budget_table" @@ -11537,25 +11539,25 @@ def _non_admin_auth(): def test_check_passthrough_routes_caller_permission_team(): from litellm.proxy._types import NewTeamRequest from litellm.proxy.management_endpoints.common_utils import ( - _check_passthrough_routes_caller_permission, + check_passthrough_routes_caller_permission, ) admin = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN) non_admin = _non_admin_auth() - _check_passthrough_routes_caller_permission( + check_passthrough_routes_caller_permission( NewTeamRequest(allowed_passthrough_routes=["/foo/*"]), admin, entity="team" ) - _check_passthrough_routes_caller_permission( + check_passthrough_routes_caller_permission( NewTeamRequest(), non_admin, entity="team" ) - _check_passthrough_routes_caller_permission( + check_passthrough_routes_caller_permission( NewTeamRequest(allowed_passthrough_routes=[]), non_admin, entity="team" ) with pytest.raises(HTTPException) as exc: - _check_passthrough_routes_caller_permission( + check_passthrough_routes_caller_permission( NewTeamRequest(allowed_passthrough_routes=["/admin/*"]), non_admin, entity="team", @@ -11565,7 +11567,7 @@ def test_check_passthrough_routes_caller_permission_team(): assert "team" in str(exc.value.detail) with pytest.raises(HTTPException) as exc: - _check_passthrough_routes_caller_permission( + check_passthrough_routes_caller_permission( NewTeamRequest(metadata={"allowed_passthrough_routes": ["/admin/*"]}), non_admin, entity="team", @@ -11628,24 +11630,24 @@ async def test_update_team_blocks_non_admin_passthrough_routes(mock_db_client): def test_check_disable_global_guardrails_caller_permission_team(): from litellm.proxy._types import NewTeamRequest from litellm.proxy.management_endpoints.common_utils import ( - _check_disable_global_guardrails_caller_permission, + check_disable_global_guardrails_caller_permission, ) admin = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN) non_admin = _non_admin_auth() - _check_disable_global_guardrails_caller_permission(True, {"disable_global_guardrails": True}, admin, entity="team") - _check_disable_global_guardrails_caller_permission(None, None, non_admin, entity="team") - _check_disable_global_guardrails_caller_permission(False, None, non_admin, entity="team") + check_disable_global_guardrails_caller_permission(True, {"disable_global_guardrails": True}, admin, entity="team") + check_disable_global_guardrails_caller_permission(None, None, non_admin, entity="team") + check_disable_global_guardrails_caller_permission(False, None, non_admin, entity="team") with pytest.raises(HTTPException) as exc: - _check_disable_global_guardrails_caller_permission(True, None, non_admin, entity="team") + check_disable_global_guardrails_caller_permission(True, None, non_admin, entity="team") assert exc.value.status_code == 403 assert "disable_global_guardrails" in str(exc.value.detail) assert "team" in str(exc.value.detail) with pytest.raises(HTTPException) as exc: - _check_disable_global_guardrails_caller_permission( + check_disable_global_guardrails_caller_permission( None, {"disable_global_guardrails": True}, non_admin, entity="team" ) assert exc.value.status_code == 403 @@ -12376,7 +12378,7 @@ async def _drive_team_write( _patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()), _patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), _patch( - "litellm.proxy.management_endpoints.team_endpoints._refresh_cached_team", + "litellm.proxy.management_endpoints.team_endpoints.refresh_cached_team", new=AsyncMock(), ), ): @@ -13860,7 +13862,7 @@ async def test_team_member_update_role_change_emits_a_roster_audit_event(monkeyp AsyncMock(side_effect=[_team_info_as_read_from_db("user"), _team_info_as_read_from_db("admin")]), ), patch( # test-quality-ok: no live DB here; matches this file's established convention for endpoint-logic unit tests - "litellm.proxy.management_endpoints.team_endpoints._upsert_budget_and_membership", + "litellm.proxy.management_endpoints.team_endpoints.upsert_budget_and_membership", AsyncMock(), ), ): @@ -13917,7 +13919,7 @@ def _member_update_patches(team_snapshot: LiteLLM_TeamTable): ), ), patch( # test-quality-ok: no live DB here; matches this file's established convention for endpoint-logic unit tests - "litellm.proxy.management_endpoints.team_endpoints._upsert_budget_and_membership", + "litellm.proxy.management_endpoints.team_endpoints.upsert_budget_and_membership", AsyncMock(), ), ) @@ -14439,7 +14441,12 @@ def _wire_update_team(stack, existing_metadata): stack.enter_context(patch("litellm.proxy.proxy_server.user_api_key_cache")) stack.enter_context(patch("litellm.proxy.proxy_server.proxy_logging_obj")) stack.enter_context(patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin")) - stack.enter_context(patch("litellm.proxy.management_endpoints.team_endpoints._cache_team_object")) + stack.enter_context( + patch( + "litellm.proxy.management_endpoints.team_endpoints.cache_team_object", + new_callable=AsyncMock, + ) + ) existing_team = MagicMock() existing_team.metadata = existing_metadata @@ -14843,7 +14850,10 @@ async def test_update_team_syncs_access_group_assigned_team_ids_in_both_directio patch("litellm.proxy.proxy_server.user_api_key_cache"), patch("litellm.proxy.proxy_server.proxy_logging_obj"), patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), - patch("litellm.proxy.management_endpoints.team_endpoints._refresh_cached_team"), + patch( + "litellm.proxy.management_endpoints.team_endpoints.refresh_cached_team", + new_callable=AsyncMock, + ), patch( "litellm.proxy.management_helpers.access_group_team_sync.invalidate_access_group_cache", new_callable=AsyncMock, @@ -15037,7 +15047,7 @@ async def test_invalidate_access_group_cache_deletes_the_cached_object(): patch("litellm.proxy.proxy_server.user_api_key_cache", cache), patch("litellm.proxy.proxy_server.proxy_logging_obj", logging_obj), patch( - "litellm.proxy.management_helpers.access_group_team_sync._delete_cache_access_object", + "litellm.proxy.management_helpers.access_group_team_sync.delete_cache_access_object", new_callable=AsyncMock, ) as delete_cached, ): @@ -15569,7 +15579,7 @@ async def test_team_member_update_invalidates_team_member_spend_state_when_budge AsyncMock(return_value=team_info_response), ), patch( # test-quality-ok: no live DB here; matches this file's established convention for endpoint-logic unit tests - "litellm.proxy.management_endpoints.team_endpoints._upsert_budget_and_membership", + "litellm.proxy.management_endpoints.team_endpoints.upsert_budget_and_membership", AsyncMock(), ), ): @@ -15622,7 +15632,7 @@ async def test_team_member_update_skips_invalidation_when_no_budget_fields_sent( AsyncMock(return_value=team_info_response), ), patch( # test-quality-ok: no live DB here; matches this file's established convention for endpoint-logic unit tests - "litellm.proxy.management_endpoints.team_endpoints._upsert_budget_and_membership", + "litellm.proxy.management_endpoints.team_endpoints.upsert_budget_and_membership", AsyncMock(), ), ): diff --git a/tests/unit/proxy/management_endpoints/test_team_model_alias_merge.py b/tests/unit/proxy/management_endpoints/test_team_model_alias_merge.py index 7cdf60f043e..22e03328fe6 100644 --- a/tests/unit/proxy/management_endpoints/test_team_model_alias_merge.py +++ b/tests/unit/proxy/management_endpoints/test_team_model_alias_merge.py @@ -29,9 +29,7 @@ class TestTeamModelAddAtomicAppend: from litellm.proxy.management_endpoints.team_endpoints import team_model_add mock_request = MagicMock() - mock_user = UserAPIKeyAuth( - user_role=LitellmUserRoles.PROXY_ADMIN, user_id="test_user" - ) + mock_user = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="test_user") existing_team = MagicMock() existing_team.model_dump.return_value = { @@ -49,19 +47,15 @@ class TestTeamModelAddAtomicAppend: with ( patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch( - "litellm.proxy.management_endpoints.team_endpoints._cache_team_object", + "litellm.proxy.management_endpoints.team_endpoints.cache_team_object", new_callable=AsyncMock, ), patch("litellm.proxy.proxy_server.user_api_key_cache"), patch("litellm.proxy.proxy_server.proxy_logging_obj"), ): - mock_prisma.db.litellm_teamtable.find_unique = AsyncMock( - return_value=existing_team - ) + mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=existing_team) mock_prisma.db.execute_raw = AsyncMock(return_value=None) - mock_prisma.db.litellm_teamtable.update = AsyncMock( - return_value=updated_team - ) + mock_prisma.db.litellm_teamtable.update = AsyncMock(return_value=updated_team) await team_model_add( data=TeamModelAddRequest(team_id="team-1", models=["new-model"]), diff --git a/tests/unit/proxy/management_endpoints/test_ui_sso.py b/tests/unit/proxy/management_endpoints/test_ui_sso.py index e053c4c94ea..4e128440703 100644 --- a/tests/unit/proxy/management_endpoints/test_ui_sso.py +++ b/tests/unit/proxy/management_endpoints/test_ui_sso.py @@ -2270,7 +2270,7 @@ class TestUISSO_FunctionsExistence: assert SSOAuthenticationHandler is not None # Check that the new _get_cli_state method exists - assert hasattr(SSOAuthenticationHandler, "_get_cli_state") + assert hasattr(SSOAuthenticationHandler, "get_cli_state") assert callable(SSOAuthenticationHandler._get_cli_state) @@ -3028,7 +3028,7 @@ class TestCLIKeyRegenerationFlow: return_value="https://proxy.example.com/sso/callback", ), patch( - "litellm.proxy.management_endpoints.ui_sso.SSOAuthenticationHandler._get_cli_state", + "litellm.proxy.management_endpoints.ui_sso.SSOAuthenticationHandler.get_cli_state", return_value=None, ) as mock_get_cli_state, ): @@ -8161,7 +8161,7 @@ class TestPKCEStateCookieBinding: ), patch.object( SSOAuthenticationHandler, - "_pkce_token_exchange", + "pkce_token_exchange", AsyncMock( return_value={ "access_token": "tok", @@ -8173,7 +8173,7 @@ class TestPKCEStateCookieBinding: ), patch.object( SSOAuthenticationHandler, - "_delete_pkce_verifier", + "delete_pkce_verifier", AsyncMock(), ), patch("fastapi_sso.sso.base.DiscoveryDocument"), @@ -8824,7 +8824,7 @@ async def test_pkce_arm_captures_sso_assertion(): ), patch.object( SSOAuthenticationHandler, - "_pkce_token_exchange", + "pkce_token_exchange", AsyncMock( return_value={ "access_token": "tok", @@ -8835,7 +8835,7 @@ async def test_pkce_arm_captures_sso_assertion(): } ), ), - patch.object(SSOAuthenticationHandler, "_delete_pkce_verifier", AsyncMock()), + patch.object(SSOAuthenticationHandler, "delete_pkce_verifier", AsyncMock()), patch("fastapi_sso.sso.base.DiscoveryDocument"), patch("fastapi_sso.sso.generic.create_provider", return_value=MagicMock()), patch.dict( diff --git a/tests/unit/proxy/management_helpers/test_management_helpers_utils.py b/tests/unit/proxy/management_helpers/test_management_helpers_utils.py index 82eafc70077..b450a4dfd02 100644 --- a/tests/unit/proxy/management_helpers/test_management_helpers_utils.py +++ b/tests/unit/proxy/management_helpers/test_management_helpers_utils.py @@ -899,7 +899,7 @@ class _FakeDb: async def test_team_update_reaches_inherited_members_but_not_overridden_ones(): from litellm.proxy._types import LitellmUserRoles from litellm.proxy.auth.auth_checks import _check_team_member_budget - from litellm.proxy.management_endpoints.common_utils import _upsert_budget_and_membership + from litellm.proxy.management_endpoints.common_utils import upsert_budget_and_membership from litellm.proxy.management_endpoints.team_endpoints import TeamMemberBudgetHandler from litellm.proxy.utils import ProxyLogging @@ -922,7 +922,7 @@ async def test_team_update_reaches_inherited_members_but_not_overridden_ones(): default_team_budget_id=default_budget.budget_id, ) - await _upsert_budget_and_membership( + await upsert_budget_and_membership( db, team_id=team_id, user_id="overridden", diff --git a/tests/unit/proxy/management_helpers/test_object_permission_utils.py b/tests/unit/proxy/management_helpers/test_object_permission_utils.py index 2fba39b30f6..33557713162 100644 --- a/tests/unit/proxy/management_helpers/test_object_permission_utils.py +++ b/tests/unit/proxy/management_helpers/test_object_permission_utils.py @@ -19,7 +19,7 @@ from litellm.proxy.management_helpers.object_permission_utils import ( _extract_requested_mcp_access_groups, _extract_requested_mcp_server_ids, _resolve_team_allowed_mcp_servers, - _set_object_permission, + set_object_permission, enforce_all_proxy_mcp_servers_grant_is_admin_only, prepare_object_permission_upsert, validate_key_mcp_servers_against_team, @@ -62,7 +62,7 @@ async def test_set_object_permission(): } # Call the function - result = await _set_object_permission( + result = await set_object_permission( data_json=data_json, prisma_client=mock_prisma_client ) @@ -116,7 +116,7 @@ async def test_set_object_permission_persists_mcp_tool_search_enabled(): }, } - await _set_object_permission(data_json=data_json, prisma_client=mock_prisma_client) + await set_object_permission(data_json=data_json, prisma_client=mock_prisma_client) created_data = ( mock_prisma_client.db.litellm_objectpermissiontable.create.call_args.kwargs[ @@ -139,7 +139,7 @@ async def test_set_object_permission_persists_skills(): "object_permission": LiteLLM_ObjectPermissionBase(skills=["private-skill"]).model_dump(), } - await _set_object_permission(data_json=data_json, prisma_client=mock_prisma_client) + await set_object_permission(data_json=data_json, prisma_client=mock_prisma_client) created_data = ( mock_prisma_client.db.litellm_objectpermissiontable.create.call_args.kwargs[ @@ -273,11 +273,11 @@ def _make_mock_mcp_manager(*existing_ids: str, servers=None): @pytest.mark.asyncio @patch( - "litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids", + "litellm.proxy.management_helpers.object_permission_utils.get_allow_all_keys_server_ids", return_value=set(), ) @patch( - "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups", + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ) @@ -295,11 +295,11 @@ async def test_validate_no_object_permission(mock_access_groups, mock_allow_all) new=_make_mock_mcp_manager("server-1", "server-2"), ) @patch( - "litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids", + "litellm.proxy.management_helpers.object_permission_utils.get_allow_all_keys_server_ids", return_value=set(), ) @patch( - "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups", + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ) @@ -320,11 +320,11 @@ async def test_validate_key_servers_within_team_scope( new=_make_mock_mcp_manager("server-1", "server-outside"), ) @patch( - "litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids", + "litellm.proxy.management_helpers.object_permission_utils.get_allow_all_keys_server_ids", return_value=set(), ) @patch( - "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups", + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ) @@ -348,11 +348,11 @@ async def test_validate_key_servers_outside_team_scope_raises( new=_make_mock_mcp_manager("server-1", "global-server"), ) @patch( - "litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids", + "litellm.proxy.management_helpers.object_permission_utils.get_allow_all_keys_server_ids", return_value={"global-server"}, ) @patch( - "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups", + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ) @@ -373,11 +373,11 @@ async def test_validate_allow_all_keys_servers_always_allowed( new=_make_mock_mcp_manager("global-server"), ) @patch( - "litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids", + "litellm.proxy.management_helpers.object_permission_utils.get_allow_all_keys_server_ids", return_value={"global-server"}, ) @patch( - "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups", + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ) @@ -395,11 +395,11 @@ async def test_validate_no_team_only_allow_all_keys(mock_access_groups, mock_all new=_make_mock_mcp_manager("private-server"), ) @patch( - "litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids", + "litellm.proxy.management_helpers.object_permission_utils.get_allow_all_keys_server_ids", return_value={"global-server"}, ) @patch( - "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups", + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ) @@ -422,11 +422,11 @@ async def test_validate_no_team_non_global_server_raises( new=_make_mock_mcp_manager("private-server"), ) @patch( - "litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids", + "litellm.proxy.management_helpers.object_permission_utils.get_allow_all_keys_server_ids", return_value=set(), ) @patch( - "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups", + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ) @@ -448,11 +448,11 @@ async def test_validate_no_team_proxy_admin_can_assign_private_server( new=_make_mock_mcp_manager("private-server"), ) @patch( - "litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids", + "litellm.proxy.management_helpers.object_permission_utils.get_allow_all_keys_server_ids", return_value=set(), ) @patch( - "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups", + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ) @@ -471,11 +471,11 @@ async def test_validate_no_team_non_admin_private_server_still_raises( @pytest.mark.asyncio @patch( - "litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids", + "litellm.proxy.management_helpers.object_permission_utils.get_allow_all_keys_server_ids", return_value=set(), ) @patch( - "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups", + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ) @@ -497,11 +497,11 @@ async def test_validate_no_team_proxy_admin_can_assign_access_group( new=_make_mock_mcp_manager("server-1", "server-outside"), ) @patch( - "litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids", + "litellm.proxy.management_helpers.object_permission_utils.get_allow_all_keys_server_ids", return_value=set(), ) @patch( - "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups", + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ) @@ -526,11 +526,11 @@ async def test_validate_proxy_admin_still_bounded_by_team_scope( new=_make_mock_mcp_manager("some-server"), ) @patch( - "litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids", + "litellm.proxy.management_helpers.object_permission_utils.get_allow_all_keys_server_ids", return_value=set(), ) @patch( - "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups", + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ) @@ -679,11 +679,11 @@ async def test_team_unified_access_group_without_servers_preserves_direct_grants new=_make_mock_mcp_manager("server-outside"), ) @patch( - "litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids", + "litellm.proxy.management_helpers.object_permission_utils.get_allow_all_keys_server_ids", return_value=set(), ) @patch( - "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups", + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ) @@ -707,11 +707,11 @@ async def test_validate_tool_permissions_validated_against_team( new=_make_mock_mcp_manager(), # empty registry — all IDs are stale ) @patch( - "litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids", + "litellm.proxy.management_helpers.object_permission_utils.get_allow_all_keys_server_ids", return_value=set(), ) @patch( - "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups", + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ) @@ -739,11 +739,11 @@ async def test_validate_stale_mcp_server_ids_are_silently_dropped( new=_make_mock_mcp_manager(), # empty registry — all IDs are stale ) @patch( - "litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids", + "litellm.proxy.management_helpers.object_permission_utils.get_allow_all_keys_server_ids", return_value=set(), ) @patch( - "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups", + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ) @@ -769,11 +769,11 @@ async def test_validate_stale_ids_in_mcp_tool_permissions_silently_dropped( new=_make_mock_mcp_manager(), # empty registry — all IDs are stale ) @patch( - "litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids", + "litellm.proxy.management_helpers.object_permission_utils.get_allow_all_keys_server_ids", return_value=set(), ) @patch( - "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups", + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ) @@ -804,11 +804,11 @@ async def test_validate_stale_mcp_server_ids_are_removed_from_object_permission( ), ) @patch( - "litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids", + "litellm.proxy.management_helpers.object_permission_utils.get_allow_all_keys_server_ids", return_value=set(), ) @patch( - "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups", + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ) @@ -839,11 +839,11 @@ async def test_validate_mcp_server_alias_outside_team_scope_raises( ), ) @patch( - "litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids", + "litellm.proxy.management_helpers.object_permission_utils.get_allow_all_keys_server_ids", return_value=set(), ) @patch( - "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups", + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ) @@ -891,11 +891,11 @@ def test_alias_grant_expands_on_other_region_after_save(): new=_make_mock_mcp_manager(), ) @patch( - "litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids", + "litellm.proxy.management_helpers.object_permission_utils.get_allow_all_keys_server_ids", return_value=set(), ) @patch( - "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups", + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ) @@ -925,11 +925,11 @@ async def test_validate_db_mcp_server_alias_outside_team_scope_raises_when_regis @pytest.mark.asyncio @patch( - "litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids", + "litellm.proxy.management_helpers.object_permission_utils.get_allow_all_keys_server_ids", return_value=set(), ) @patch( - "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups", + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ) @@ -946,11 +946,11 @@ async def test_validate_access_groups_within_team_scope( @pytest.mark.asyncio @patch( - "litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids", + "litellm.proxy.management_helpers.object_permission_utils.get_allow_all_keys_server_ids", return_value=set(), ) @patch( - "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups", + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ) @@ -970,11 +970,11 @@ async def test_validate_access_groups_outside_team_scope_raises( @pytest.mark.asyncio @patch( - "litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids", + "litellm.proxy.management_helpers.object_permission_utils.get_allow_all_keys_server_ids", return_value=set(), ) @patch( - "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups", + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ) @@ -993,11 +993,11 @@ async def test_validate_access_groups_no_team_raises( @pytest.mark.asyncio @patch( - "litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids", + "litellm.proxy.management_helpers.object_permission_utils.get_allow_all_keys_server_ids", return_value=set(), ) @patch( - "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups", + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=["server-from-group"], ) @@ -1018,7 +1018,7 @@ async def test_validate_team_access_groups_resolve_to_servers( @pytest.mark.asyncio @patch( - "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups", + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ) @@ -1037,7 +1037,7 @@ async def test_resolve_team_allowed_mcp_servers_string_tool_permissions( @pytest.mark.asyncio @patch( - "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups", + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ) @@ -1059,7 +1059,7 @@ async def test_resolve_team_allowed_mcp_servers_dict_tool_permissions( @pytest.mark.asyncio @patch( - "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups", + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ) @@ -1100,11 +1100,11 @@ async def test_resolve_team_all_proxy_sentinel_resolves_dynamically(mock_access_ new=_make_mock_mcp_manager("srv-x", "srv-y", "srv-z"), ) @patch( - "litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids", + "litellm.proxy.management_helpers.object_permission_utils.get_allow_all_keys_server_ids", return_value=set(), ) @patch( - "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups", + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ) @@ -1131,11 +1131,11 @@ async def test_validate_key_scoped_to_server_added_after_team_all_proxy( new=_make_mock_mcp_manager("srv-x", "srv-z"), ) @patch( - "litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids", + "litellm.proxy.management_helpers.object_permission_utils.get_allow_all_keys_server_ids", return_value=set(), ) @patch( - "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups", + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ) @@ -1274,11 +1274,11 @@ async def test_validate_search_tools_raises_when_not_subset(): @pytest.mark.asyncio @patch( - "litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids", + "litellm.proxy.management_helpers.object_permission_utils.get_allow_all_keys_server_ids", return_value=set(), ) @patch( - "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups", + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ) @@ -1297,11 +1297,11 @@ async def test_personal_non_admin_cannot_assign_mcp_toolsets( @pytest.mark.asyncio @patch( - "litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids", + "litellm.proxy.management_helpers.object_permission_utils.get_allow_all_keys_server_ids", return_value=set(), ) @patch( - "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups", + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ) @@ -1550,7 +1550,7 @@ async def test_set_object_permission_rejects_shared_alias_or_name_tool_permissio data_json = {"object_permission": {"mcp_tool_permissions": {identifier: ["read_wiki_structure"]}}} with pytest.raises(HTTPException) as exc_info: - await _set_object_permission(data_json=data_json, prisma_client=mock_prisma) + await set_object_permission(data_json=data_json, prisma_client=mock_prisma) assert exc_info.value.status_code == 400 assert all(server_id in str(exc_info.value.detail) for server_id in colliding_ids) diff --git a/tests/unit/proxy/management_helpers/test_team_metadata_validation.py b/tests/unit/proxy/management_helpers/test_team_metadata_validation.py index 26bcba775a5..6defb5e45b2 100644 --- a/tests/unit/proxy/management_helpers/test_team_metadata_validation.py +++ b/tests/unit/proxy/management_helpers/test_team_metadata_validation.py @@ -405,7 +405,7 @@ async def _drive_update(kind, existing_metadata, payload): patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()), patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), patch( - "litellm.proxy.management_endpoints.team_endpoints._refresh_cached_team", + "litellm.proxy.management_endpoints.team_endpoints.refresh_cached_team", new=AsyncMock(), ), ): diff --git a/tests/unit/proxy/openai_files_endpoint/test_batch_guardrails.py b/tests/unit/proxy/openai_files_endpoint/test_batch_guardrails.py index 8c8dc5d799f..2387caca33b 100644 --- a/tests/unit/proxy/openai_files_endpoint/test_batch_guardrails.py +++ b/tests/unit/proxy/openai_files_endpoint/test_batch_guardrails.py @@ -940,10 +940,10 @@ async def test_a_real_non_guardrail_enforcement_hook_drops_its_record(monkeypatc pins that, because the other tests raise their own exceptions. """ import litellm - from litellm.proxy.hooks.prompt_injection_detection import _OPTIONAL_PromptInjectionDetection + from litellm.proxy.hooks.prompt_injection_detection import OPTIONAL_PromptInjectionDetection from litellm.proxy._types import LiteLLMPromptInjectionParams - hook = _OPTIONAL_PromptInjectionDetection( + hook = OPTIONAL_PromptInjectionDetection( prompt_injection_params=LiteLLMPromptInjectionParams(heuristics_check=True) ) monkeypatch.setattr(litellm, "callbacks", [hook]) diff --git a/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_anthropic_passthrough_logging_handler.py b/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_anthropic_passthrough_logging_handler.py index 238970d5a8e..07e6c3ac89d 100644 --- a/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_anthropic_passthrough_logging_handler.py +++ b/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_anthropic_passthrough_logging_handler.py @@ -165,7 +165,7 @@ class TestAnthropicLoggingHandlerModelFallback: return mock_handler @patch.object( - AnthropicPassthroughLoggingHandler, "_build_complete_streaming_response" + AnthropicPassthroughLoggingHandler, "build_complete_streaming_response" ) @patch.object( AnthropicPassthroughLoggingHandler, "_create_anthropic_response_logging_payload" @@ -2310,7 +2310,7 @@ class TestAnthropicUsageOnlyFallback: @patch("litellm.completion_cost") @patch.object( - AnthropicPassthroughLoggingHandler, "_build_complete_streaming_response" + AnthropicPassthroughLoggingHandler, "build_complete_streaming_response" ) def test_handler_falls_back_when_assembly_returns_none( self, mock_assemble, mock_cost @@ -2336,7 +2336,7 @@ class TestAnthropicUsageOnlyFallback: @patch("litellm.completion_cost") @patch.object( - AnthropicPassthroughLoggingHandler, "_build_complete_streaming_response" + AnthropicPassthroughLoggingHandler, "build_complete_streaming_response" ) def test_handler_falls_back_when_assembly_raises(self, mock_assemble, mock_cost): import litellm @@ -2368,7 +2368,7 @@ class TestAnthropicUsageOnlyFallback: assert result["kwargs"]["response_cost"] == 0.0021 @patch.object( - AnthropicPassthroughLoggingHandler, "_build_complete_streaming_response" + AnthropicPassthroughLoggingHandler, "build_complete_streaming_response" ) def test_handler_returns_none_when_no_usage_recoverable(self, mock_assemble): # assembly fails AND the chunks carry no usage event, so there is nothing @@ -2392,10 +2392,10 @@ class TestAnthropicUsageOnlyFallback: assert result["kwargs"] == {} @patch.object( - AnthropicPassthroughLoggingHandler, "_build_usage_only_response_from_chunks" + AnthropicPassthroughLoggingHandler, "build_usage_only_response_from_chunks" ) @patch.object( - AnthropicPassthroughLoggingHandler, "_build_complete_streaming_response" + AnthropicPassthroughLoggingHandler, "build_complete_streaming_response" ) def test_handler_does_not_crash_when_usage_only_fallback_raises( self, mock_assemble, mock_fallback @@ -2839,11 +2839,7 @@ def test_handle_logging_anthropic_collected_chunks(all_chunks): "all_chunks": all_chunks, } - result = ( - AnthropicPassthroughLoggingHandler._handle_logging_anthropic_collected_chunks( - **sent_args - ) - ) + result = AnthropicPassthroughLoggingHandler.handle_logging_anthropic_collected_chunks(**sent_args) assert isinstance(result["result"], ModelResponse) print("result=", json.dumps(result, indent=4, default=str)) @@ -2857,7 +2853,7 @@ def test_build_complete_streaming_response(all_chunks): litellm_logging_obj = Mock() - result = AnthropicPassthroughLoggingHandler._build_complete_streaming_response( + result = AnthropicPassthroughLoggingHandler.build_complete_streaming_response( all_chunks=all_chunks, model="claude-sonnet-4-5-20250929", litellm_logging_obj=litellm_logging_obj, diff --git a/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_gemini_passthrough_logging_handler.py b/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_gemini_passthrough_logging_handler.py index 52c664a65a5..c30f6d9f15a 100644 --- a/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_gemini_passthrough_logging_handler.py +++ b/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_gemini_passthrough_logging_handler.py @@ -284,7 +284,7 @@ class TestGeminiPassthroughLoggingHandler: mock_logging_obj = self._create_mock_logging_obj() # Mock the _handle_logging method to capture the call - handler._handle_logging = AsyncMock() + handler.handle_logging = AsyncMock() # Mock httpx response mock_response = self._create_mock_httpx_response() @@ -316,8 +316,8 @@ class TestGeminiPassthroughLoggingHandler: assert mock_logging_obj.model_call_details["custom_llm_provider"] == "gemini" # Verify that _handle_logging was called with the correct kwargs - handler._handle_logging.assert_called_once() - call_kwargs = handler._handle_logging.call_args[1] + handler.handle_logging.assert_called_once() + call_kwargs = handler.handle_logging.call_args[1] assert call_kwargs["response_cost"] == 0.000050 assert call_kwargs["model"] == "gemini-2.0-flash" assert call_kwargs["custom_llm_provider"] == "gemini" diff --git a/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_openai_passthrough_logging_handler.py b/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_openai_passthrough_logging_handler.py index 1115b55e027..86a670bcb3e 100644 --- a/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_openai_passthrough_logging_handler.py +++ b/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_openai_passthrough_logging_handler.py @@ -1555,7 +1555,7 @@ class TestOpenAIPassthroughIntegration: ) # Mock the _handle_logging method to capture calls - self.handler._handle_logging = AsyncMock() + self.handler.handle_logging = AsyncMock() # Act result = await self.handler.pass_through_async_success_handler( @@ -1575,7 +1575,7 @@ class TestOpenAIPassthroughIntegration: ) # Assert - Should call the base handler, not our OpenAI handler - self.handler._handle_logging.assert_called_once() + self.handler.handle_logging.assert_called_once() @patch("litellm.cost_calculator.default_image_cost_calculator") def test_calculate_image_generation_cost(self, mock_image_cost_calculator): diff --git a/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_typesafe_passthrough_logging_handler.py b/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_typesafe_passthrough_logging_handler.py index 7961d2a911b..7a411c95afa 100644 --- a/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_typesafe_passthrough_logging_handler.py +++ b/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_typesafe_passthrough_logging_handler.py @@ -213,9 +213,9 @@ async def test_oss_gateway_accounts_for_checkpoint_usage_and_registered_cost( from litellm.caching.caching import DualCache from litellm.exceptions import BudgetExceededError - from litellm.proxy.hooks.model_max_budget_limiter import _PROXY_VirtualKeyModelMaxBudgetLimiter + from litellm.proxy.hooks.model_max_budget_limiter import PROXY_VirtualKeyModelMaxBudgetLimiter - budget_limiter: Final = _PROXY_VirtualKeyModelMaxBudgetLimiter(DualCache()) + budget_limiter: Final = PROXY_VirtualKeyModelMaxBudgetLimiter(DualCache()) assert await budget_limiter.is_key_within_model_budget(auth, f"{provider}/{requested}") await budget_limiter.async_log_success_event(logged, None, start, datetime.now()) with pytest.raises(BudgetExceededError): diff --git a/tests/unit/proxy/pass_through_endpoints/test_deepgram_ws_passthrough_routes.py b/tests/unit/proxy/pass_through_endpoints/test_deepgram_ws_passthrough_routes.py index 44533f35c72..404ed0f5800 100644 --- a/tests/unit/proxy/pass_through_endpoints/test_deepgram_ws_passthrough_routes.py +++ b/tests/unit/proxy/pass_through_endpoints/test_deepgram_ws_passthrough_routes.py @@ -17,7 +17,7 @@ import litellm from litellm.caching.dual_cache import DualCache from litellm.proxy._lazy_features import LAZY_FEATURES from litellm.proxy._types import LiteLLMRoutes, UserAPIKeyAuth -from litellm.proxy.auth.auth_checks import _cache_key_object +from litellm.proxy.auth.auth_checks import cache_key_object from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( _websocket_relay, deepgram_listen_websocket_route, @@ -365,7 +365,7 @@ def test_deepgram_listen_authenticates_the_litellm_key_and_relays_to_deepgram(mo async def _cache_restricted_key(virtual_key: str, models: list[str]) -> DualCache: cache = DualCache() - await _cache_key_object( + await cache_key_object( hashed_token=hash_token(virtual_key), user_api_key_obj=UserAPIKeyAuth(token=hash_token(virtual_key), models=models), user_api_key_cache=cache, diff --git a/tests/unit/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py b/tests/unit/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py index 14fd60c9793..1887021a53a 100644 --- a/tests/unit/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py +++ b/tests/unit/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py @@ -1393,7 +1393,7 @@ class TestBedrockLLMProxyRoute: with ( patch( - "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints._read_request_body", + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.read_request_body", return_value=mock_request_body, ), patch( @@ -1435,7 +1435,7 @@ class TestBedrockLLMProxyRoute: with ( patch( - "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints._read_request_body", + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.read_request_body", return_value=mock_request_body, ), patch( @@ -2493,7 +2493,7 @@ class TestForwardHeaders: with ( patch( - "litellm.proxy.pass_through_endpoints.pass_through_endpoints._read_request_body", + "litellm.proxy.pass_through_endpoints.pass_through_endpoints.read_request_body", return_value=mock_request_body, ), patch( @@ -2591,7 +2591,7 @@ class TestForwardHeaders: with ( patch( - "litellm.proxy.pass_through_endpoints.pass_through_endpoints._read_request_body", + "litellm.proxy.pass_through_endpoints.pass_through_endpoints.read_request_body", return_value=mock_request_body, ), patch( @@ -2678,7 +2678,7 @@ class TestForwardHeaders: "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.passthrough_endpoint_router.get_credentials" ) as mock_get_creds, patch( - "litellm.proxy.pass_through_endpoints.pass_through_endpoints._read_request_body", + "litellm.proxy.pass_through_endpoints.pass_through_endpoints.read_request_body", return_value={"messages": [{"role": "user", "content": "test"}]}, ), patch( @@ -2782,7 +2782,7 @@ class TestMilvusProxyRoute: "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.is_allowed_to_call_vector_store_endpoint" ) as mock_is_allowed, patch( - "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints._safe_set_request_parsed_body" + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.safe_set_request_parsed_body" ) as mock_safe_set, patch( "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.create_pass_through_route" @@ -2996,7 +2996,7 @@ class TestMilvusProxyRoute: patch( "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.is_allowed_to_call_vector_store_endpoint" ), - patch("litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints._safe_set_request_parsed_body"), + patch("litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.safe_set_request_parsed_body"), patch.object(litellm, "vector_store_index_registry") as mock_index_registry, patch.object(litellm, "vector_store_registry") as mock_vector_registry, ): @@ -3046,7 +3046,7 @@ class TestMilvusProxyRoute: patch( "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.is_allowed_to_call_vector_store_endpoint" ), - patch("litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints._safe_set_request_parsed_body"), + patch("litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.safe_set_request_parsed_body"), patch.object(litellm, "vector_store_index_registry") as mock_index_registry, patch.object(litellm, "vector_store_registry") as mock_vector_registry, ): @@ -3101,7 +3101,7 @@ class TestMilvusProxyRoute: patch( "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.is_allowed_to_call_vector_store_endpoint" ), - patch("litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints._safe_set_request_parsed_body"), + patch("litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.safe_set_request_parsed_body"), patch( "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.create_pass_through_route" ) as mock_create_route, @@ -4816,7 +4816,7 @@ class TestAzureProxyRouteCrossIndexAuthorization: new=AsyncMock(), ), patch( - "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.BaseOpenAIPassThroughHandler._base_openai_pass_through_handler", + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.BaseOpenAIPassThroughHandler.base_openai_pass_through_handler", new=AsyncMock(return_value=Response()), ), patch.object(litellm, "vector_store_index_registry") as mock_index_registry, @@ -4862,7 +4862,7 @@ class TestAzureProxyRouteCrossIndexAuthorization: return_value="azure-key", ), patch( - "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.BaseOpenAIPassThroughHandler._base_openai_pass_through_handler", + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.BaseOpenAIPassThroughHandler.base_openai_pass_through_handler", new=AsyncMock(return_value=Response()), ) as mock_handler, patch.object(litellm, "vector_store_index_registry") as mock_index_registry, @@ -4922,7 +4922,7 @@ class TestAzureProxyRouteServiceLevelIndexCreate: return_value="https://svc.search.windows.net", ), patch( - "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.BaseOpenAIPassThroughHandler._base_openai_pass_through_handler", + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.BaseOpenAIPassThroughHandler.base_openai_pass_through_handler", new=AsyncMock(return_value=Response()), ) as mock_handler, ): @@ -4954,7 +4954,7 @@ class TestAzureProxyRouteServiceLevelIndexCreate: return_value="azure-key", ), patch( - "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.BaseOpenAIPassThroughHandler._base_openai_pass_through_handler", + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.BaseOpenAIPassThroughHandler.base_openai_pass_through_handler", new=AsyncMock(return_value=Response()), ) as mock_handler, ): @@ -7542,16 +7542,16 @@ class TestTypeSafePassthroughRoute: ) -> None: from litellm.caching.caching import DualCache from litellm.proxy import proxy_server - from litellm.proxy.hooks.cache_control_check import _PROXY_CacheControlCheck + from litellm.proxy.hooks.cache_control_check import PROXY_CacheControlCheck from litellm.proxy.hooks.parallel_request_limiter_v3 import ( - _PROXY_MaxParallelRequestsHandler_v3, + PROXY_MaxParallelRequestsHandler_v3, get_request_stash, ) from litellm.proxy.utils import InternalUsageCache, ProxyLogging cache: Final = DualCache() - limiter: Final = _PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=InternalUsageCache(cache)) - monkeypatch.setattr(litellm, "callbacks", list((limiter, _PROXY_CacheControlCheck()))) + limiter: Final = PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=InternalUsageCache(cache)) + monkeypatch.setattr(litellm, "callbacks", list((limiter, PROXY_CacheControlCheck()))) monkeypatch.setattr(proxy_server, "proxy_logging_obj", ProxyLogging(user_api_key_cache=cache)) monkeypatch.setenv("OPENROUTER_API_KEY", "openrouter-test-key") monkeypatch.setenv("OPENROUTER_API_BASE", "https://typesafe.example/base") @@ -7719,12 +7719,12 @@ class TestOssDecisionPassthroughRoute: self, client: TestClient, monkeypatch: pytest.MonkeyPatch, metadata_slot: str, provider: str, checkpoint: str ) -> None: from litellm.integrations.custom_logger import CustomLogger - from litellm.proxy.hooks.parallel_request_limiter_v3 import _PROXY_MaxParallelRequestsHandler_v3 + from litellm.proxy.hooks.parallel_request_limiter_v3 import PROXY_MaxParallelRequestsHandler_v3 from litellm.proxy.utils import InternalUsageCache from litellm.proxy.proxy_server import app cache: Final = DualCache() - limiter: Final = _PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=InternalUsageCache(cache)) + limiter: Final = PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=InternalUsageCache(cache)) auth: Final = UserAPIKeyAuth( api_key="oss-native-rpm", metadata={"model_rpm_limit": {f"{provider}/{checkpoint}": 1}}, ) diff --git a/tests/unit/proxy/pass_through_endpoints/test_pass_through_endpoints.py b/tests/unit/proxy/pass_through_endpoints/test_pass_through_endpoints.py index 6c2c7d9c904..b604de5ea97 100644 --- a/tests/unit/proxy/pass_through_endpoints/test_pass_through_endpoints.py +++ b/tests/unit/proxy/pass_through_endpoints/test_pass_through_endpoints.py @@ -580,7 +580,7 @@ async def test_custom_passthrough_predict_path_logs_via_generic_handler(): ) handler = PassThroughEndpointLogging() - handler._handle_logging = AsyncMock() + handler.handle_logging = AsyncMock() mock_logging_obj = MagicMock(spec=LiteLLMLoggingObj) mock_logging_obj.model_call_details = {} @@ -612,8 +612,8 @@ async def test_custom_passthrough_predict_path_logs_via_generic_handler(): ) mock_vertex_handler.assert_not_called() - handler._handle_logging.assert_awaited_once() - logged_object = handler._handle_logging.call_args.kwargs["standard_logging_response_object"] + handler.handle_logging.assert_awaited_once() + logged_object = handler.handle_logging.call_args.kwargs["standard_logging_response_object"] assert logged_object == {"response": '{"forecast": [1, 2, 3]}'} @@ -1073,7 +1073,7 @@ async def test_pass_through_success_handler_with_cost_per_request(): mock_logging_obj.model_call_details = {} # Mock the _handle_logging method to capture the call - handler._handle_logging = AsyncMock() + handler.handle_logging = AsyncMock() # Mock httpx response mock_response = MagicMock(spec=httpx.Response) @@ -1108,8 +1108,8 @@ async def test_pass_through_success_handler_with_cost_per_request(): assert mock_logging_obj.model_call_details["response_cost"] == 1.25 # Verify that _handle_logging was called with the correct kwargs - handler._handle_logging.assert_called_once() - call_kwargs = handler._handle_logging.call_args[1] + handler.handle_logging.assert_called_once() + call_kwargs = handler.handle_logging.call_args[1] assert call_kwargs["response_cost"] == 1.25 @@ -3913,11 +3913,11 @@ async def _drive_pass_through_block(raised_exception): patch("litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging), patch(f"{_PT_MODULE}.verbose_proxy_logger", logger), patch( - f"{_PT_MODULE}._read_request_body", + f"{_PT_MODULE}.read_request_body", new_callable=AsyncMock, return_value={}, ), - patch(f"{_PT_MODULE}._safe_get_request_headers", return_value={}), + patch(f"{_PT_MODULE}.safe_get_request_headers", return_value={}), patch( "litellm.proxy.pass_through_endpoints.passthrough_guardrails." "PassthroughGuardrailHandler.collect_guardrails", @@ -6708,7 +6708,7 @@ def _passthrough_kwargs_for_reservation( async def _track_cost_for_passthrough_kwargs(kwargs: dict) -> AsyncMock: from datetime import datetime - from litellm.proxy.hooks.proxy_track_cost_callback import _ProxyDBLogger + from litellm.proxy.hooks.proxy_track_cost_callback import ProxyDBLogger callback_kwargs = { **kwargs, @@ -6731,7 +6731,7 @@ async def _track_cost_for_passthrough_kwargs(kwargs: dict) -> AsyncMock: mock_proxy_logging.db_spend_update_writer.update_database = AsyncMock() mock_proxy_logging.slack_alerting_instance.customer_spend_alert = AsyncMock() - await _ProxyDBLogger()._PROXY_track_cost_callback( + await ProxyDBLogger()._PROXY_track_cost_callback( kwargs=callback_kwargs, completion_response=None, start_time=datetime.now(), @@ -7033,10 +7033,10 @@ async def test_user_defined_passthrough_is_neither_tracked_nor_enforced(metadata from litellm.caching.caching import DualCache from litellm.proxy.auth.auth_utils import get_model_from_request - from litellm.proxy.hooks.model_max_budget_limiter import _PROXY_VirtualKeyModelMaxBudgetLimiter + from litellm.proxy.hooks.model_max_budget_limiter import PROXY_VirtualKeyModelMaxBudgetLimiter budget: Final = {"managed-model": {"budget_limit": 0.1, "time_period": "1d"}} - limiter: Final = _PROXY_VirtualKeyModelMaxBudgetLimiter(DualCache()) + limiter: Final = PROXY_VirtualKeyModelMaxBudgetLimiter(DualCache()) auth: Final = UserAPIKeyAuth( api_key="custom-key", token="custom-key", team_id="shared-team", team_model_max_budget=budget, ) diff --git a/tests/unit/proxy/pass_through_endpoints/test_passthrough_guardrail_block_otel_span.py b/tests/unit/proxy/pass_through_endpoints/test_passthrough_guardrail_block_otel_span.py index 09987b2781c..e0687636263 100644 --- a/tests/unit/proxy/pass_through_endpoints/test_passthrough_guardrail_block_otel_span.py +++ b/tests/unit/proxy/pass_through_endpoints/test_passthrough_guardrail_block_otel_span.py @@ -147,8 +147,8 @@ async def _drive(response_text: str): patch("litellm.proxy.proxy_server.llm_router", None), patch(f"{_PT_MOD}.pass_through_endpoint_logging", mock_pt_logging), patch(f"{_PT_MOD}.get_async_httpx_client", return_value=mock_async_client_obj), - patch(f"{_PT_MOD}._read_request_body", new_callable=AsyncMock, return_value={}), - patch(f"{_PT_MOD}._safe_get_request_headers", return_value={}), + patch(f"{_PT_MOD}.read_request_body", new_callable=AsyncMock, return_value={}), + patch(f"{_PT_MOD}.safe_get_request_headers", return_value={}), patch(_COLLECT, return_value=["block-demo"]), ] try: diff --git a/tests/unit/proxy/pass_through_endpoints/test_passthrough_post_call_guardrails.py b/tests/unit/proxy/pass_through_endpoints/test_passthrough_post_call_guardrails.py index c7696079adc..7b06905c8aa 100644 --- a/tests/unit/proxy/pass_through_endpoints/test_passthrough_post_call_guardrails.py +++ b/tests/unit/proxy/pass_through_endpoints/test_passthrough_post_call_guardrails.py @@ -89,8 +89,8 @@ def _common_patches(mock_proxy_logging, mock_response): patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging), patch(f"{_PT_MOD}.pass_through_endpoint_logging", mock_pt_logging), patch(f"{_PT_MOD}.get_async_httpx_client", return_value=mock_async_client_obj), - patch(f"{_PT_MOD}._read_request_body", new_callable=AsyncMock, return_value={}), - patch(f"{_PT_MOD}._safe_get_request_headers", return_value={}), + patch(f"{_PT_MOD}.read_request_body", new_callable=AsyncMock, return_value={}), + patch(f"{_PT_MOD}.safe_get_request_headers", return_value={}), ] stack = ExitStack() diff --git a/tests/unit/proxy/pass_through_endpoints/test_streaming_handler_interrupt.py b/tests/unit/proxy/pass_through_endpoints/test_streaming_handler_interrupt.py index e6b19f4eec1..b278d7dea01 100644 --- a/tests/unit/proxy/pass_through_endpoints/test_streaming_handler_interrupt.py +++ b/tests/unit/proxy/pass_through_endpoints/test_streaming_handler_interrupt.py @@ -55,7 +55,7 @@ async def test_chunk_processor_logs_on_normal_completion(): with patch.object( PassThroughStreamingHandler, - "_route_streaming_logging_to_handler", + "route_streaming_logging_to_handler", new=AsyncMock(), ) as mock_route: received = [] @@ -88,7 +88,7 @@ async def test_chunk_processor_logs_on_client_disconnect(): with patch.object( PassThroughStreamingHandler, - "_route_streaming_logging_to_handler", + "route_streaming_logging_to_handler", new=AsyncMock(), ) as mock_route: gen = PassThroughStreamingHandler.chunk_processor( @@ -126,7 +126,7 @@ async def test_chunk_processor_does_not_schedule_success_logging_for_upstream_er with patch.object( PassThroughStreamingHandler, - "_route_streaming_logging_to_handler", + "route_streaming_logging_to_handler", new=AsyncMock(), ) as mock_route: received = [] @@ -156,7 +156,7 @@ async def test_chunk_processor_does_not_schedule_logging_when_no_chunks(): with patch.object( PassThroughStreamingHandler, - "_route_streaming_logging_to_handler", + "route_streaming_logging_to_handler", new=AsyncMock(), ) as mock_route: received = [] @@ -193,7 +193,7 @@ async def test_chunk_processor_routes_logging_through_logging_worker(): with ( patch.object( PassThroughStreamingHandler, - "_route_streaming_logging_to_handler", + "route_streaming_logging_to_handler", new=AsyncMock(), ), patch.object( @@ -235,7 +235,7 @@ async def test_chunk_processor_routes_logging_through_logging_worker_on_disconne with ( patch.object( PassThroughStreamingHandler, - "_route_streaming_logging_to_handler", + "route_streaming_logging_to_handler", new=AsyncMock(), ), patch.object( @@ -287,7 +287,7 @@ async def test_chunk_processor_stamps_completion_start_time_on_first_chunk(): with patch.object( PassThroughStreamingHandler, - "_route_streaming_logging_to_handler", + "route_streaming_logging_to_handler", new=AsyncMock(), ): received = [] @@ -326,7 +326,7 @@ async def test_chunk_processor_does_not_reset_completion_start_time_on_later_chu with patch.object( PassThroughStreamingHandler, - "_route_streaming_logging_to_handler", + "route_streaming_logging_to_handler", new=AsyncMock(), ): async for _ in PassThroughStreamingHandler.chunk_processor( @@ -361,7 +361,7 @@ async def test_chunk_processor_stamps_completion_start_time_on_cost_injection_pa try: with patch.object( PassThroughStreamingHandler, - "_route_streaming_logging_to_handler", + "route_streaming_logging_to_handler", new=AsyncMock(), ): async for _ in PassThroughStreamingHandler.chunk_processor( diff --git a/tests/unit/proxy/proxy_server/test_background_health.py b/tests/unit/proxy/proxy_server/test_background_health.py index b15349d6705..bb2fac403a7 100644 --- a/tests/unit/proxy/proxy_server/test_background_health.py +++ b/tests/unit/proxy/proxy_server/test_background_health.py @@ -178,7 +178,7 @@ async def test_schedule_background_health_check_db_save_creates_task(monkeypatch import litellm.proxy.health_endpoints._health_endpoints as he - monkeypatch.setattr(he, "_save_background_health_checks_to_db", _fake_save) + monkeypatch.setattr(he, "save_background_health_checks_to_db", _fake_save) prisma_client = MagicMock() shared_manager = SimpleNamespace(pod_id="pod-xyz") @@ -227,7 +227,7 @@ async def test_schedule_background_health_check_db_save_invalid_no_event_loop_ra import litellm.proxy.health_endpoints._health_endpoints as he - monkeypatch.setattr(he, "_save_background_health_checks_to_db", _fake_save) + monkeypatch.setattr(he, "save_background_health_checks_to_db", _fake_save) def _broken_create_task(_coro): raise RuntimeError("no running event loop") @@ -261,7 +261,7 @@ def _capture_saves(monkeypatch, persisted=True): import litellm.proxy.health_endpoints._health_endpoints as he - monkeypatch.setattr(he, "_save_background_health_checks_to_db", _fake_save) + monkeypatch.setattr(he, "save_background_health_checks_to_db", _fake_save) return saves @@ -271,7 +271,7 @@ def _cancel_during_save(monkeypatch): import litellm.proxy.health_endpoints._health_endpoints as he - monkeypatch.setattr(he, "_save_background_health_checks_to_db", _fake_save) + monkeypatch.setattr(he, "save_background_health_checks_to_db", _fake_save) def _schedule_with(lock_manager): diff --git a/tests/unit/proxy/proxy_server/test_lifecycle.py b/tests/unit/proxy/proxy_server/test_lifecycle.py index ba5501315d9..69f9c72c9c8 100644 --- a/tests/unit/proxy/proxy_server/test_lifecycle.py +++ b/tests/unit/proxy/proxy_server/test_lifecycle.py @@ -33,7 +33,7 @@ from typing_extensions import TypedDict import litellm.proxy.proxy_server as ps from litellm.proxy.proxy_server import ( ProxyStartupEvent, - _initialize_shared_aiohttp_session, + initialize_shared_aiohttp_session, _resolve_pydantic_type, _resolve_typed_dict_type, cleanup_router_config_variables, @@ -349,7 +349,7 @@ async def test_flush_spend_counters_on_shutdown_logs_and_swallows_commit_errors( async def test_initialize_shared_aiohttp_session_returns_client_session(): from aiohttp import ClientSession - session = await _initialize_shared_aiohttp_session() + session = await initialize_shared_aiohttp_session() try: observed = { "is_client_session": isinstance(session, ClientSession), @@ -381,7 +381,7 @@ async def test_initialize_shared_aiohttp_session_aiohttp_missing_returns_none_on return real_import(name, *args, **kwargs) monkeypatch.setattr(builtins, "__import__", _raise_for_aiohttp) - result = await _initialize_shared_aiohttp_session() + result = await initialize_shared_aiohttp_session() assert result is None @@ -914,7 +914,7 @@ def test_otel_global_provider_published_after_callback_init(): """ wrapped = getattr(proxy_startup_event, "__wrapped__", proxy_startup_event) source = inspect.getsource(wrapped) - init_pos = source.find("_initialize_startup_logging(") + init_pos = source.find("ProxyStartupEvent.initialize_startup_logging(") publish_pos = source.find("publish_global_otel_v2_provider(") assert init_pos != -1, "callback init call not found in proxy_startup_event" assert publish_pos != -1, "OTEL global publish not found in proxy_startup_event" @@ -927,7 +927,7 @@ def test_otel_global_provider_published_after_callback_init(): def test_startup_warns_for_global_budget_without_database(caplog): with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"): - ProxyStartupEvent._warn_budget_without_db(max_budget=100.0, prisma_client=None) + ProxyStartupEvent.warn_budget_without_db(max_budget=100.0, prisma_client=None) assert "litellm.max_budget=100.0" in caplog.text assert "will NOT be enforced" in caplog.text @@ -936,7 +936,7 @@ def test_startup_warns_for_global_budget_without_database(caplog): def test_startup_does_not_warn_for_global_budget_with_database(caplog): with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"): - ProxyStartupEvent._warn_budget_without_db(max_budget=100.0, prisma_client=MagicMock()) + ProxyStartupEvent.warn_budget_without_db(max_budget=100.0, prisma_client=MagicMock()) assert "litellm.max_budget" not in caplog.text @@ -944,7 +944,7 @@ def test_startup_does_not_warn_for_global_budget_with_database(caplog): @pytest.mark.parametrize("max_budget", [0, None]) def test_startup_does_not_warn_without_global_budget(caplog, max_budget): with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"): - ProxyStartupEvent._warn_budget_without_db(max_budget=max_budget, prisma_client=None) + ProxyStartupEvent.warn_budget_without_db(max_budget=max_budget, prisma_client=None) assert "litellm.max_budget" not in caplog.text @@ -1001,7 +1001,7 @@ def test_proxy_startup_event_warns_for_global_budget_without_database(): wrapped = getattr(proxy_startup_event, "__wrapped__", proxy_startup_event) source = inspect.getsource(wrapped) budget_check_pos = source.find("if prisma_client is not None and litellm.max_budget > 0:") - warn_pos = source.find("_warn_budget_without_db(") + warn_pos = source.find("warn_budget_without_db(") next_startup_section_pos = source.find( "await ProxyStartupEvent.initialize_scheduled_background_jobs(", budget_check_pos, @@ -1089,7 +1089,7 @@ async def test_scorer_baseline_upgrade_preserves_existing_routers_and_is_not_ref @pytest.mark.asyncio async def test_tuning_baseline_waits_for_a_complete_db_model_census(monkeypatch): prisma_client = MagicMock() - monkeypatch.setattr(ps.proxy_config, "_get_models_from_db", AsyncMock(return_value=None)) + monkeypatch.setattr(ps.proxy_config, "get_models_from_db", AsyncMock(return_value=None)) result = await ProxyStartupEvent.enforce_heuristic_v1_tuning_baseline( prisma_client=prisma_client, diff --git a/tests/unit/proxy/proxy_server/test_proxy_config.py b/tests/unit/proxy/proxy_server/test_proxy_config.py index 903557194ec..d1b80627b3a 100644 --- a/tests/unit/proxy/proxy_server/test_proxy_config.py +++ b/tests/unit/proxy/proxy_server/test_proxy_config.py @@ -3883,7 +3883,7 @@ async def test_ProxyConfig_add_deployment_applies_db_router_settings(monkeypatch return {} monkeypatch.setattr(pc, "get_config", fake_get_config) - monkeypatch.setattr(pc, "_get_models_from_db", AsyncMock(return_value=[])) + monkeypatch.setattr(pc, "get_models_from_db", AsyncMock(return_value=[])) monkeypatch.setattr(pc, "_init_non_llm_objects_in_db", AsyncMock()) monkeypatch.setattr(proxy_server, "prefetch_config_params", AsyncMock()) monkeypatch.setattr(proxy_server, "get_config_param", AsyncMock(return_value=None)) @@ -3963,7 +3963,7 @@ async def test_ProxyConfig_add_deployment_loads_db_credentials_before_reconcilin async def install_models(new_models: object, proxy_logging_obj: object) -> None: installed(credential=CredentialAccessor.get_credential_values("openai-cred")) - monkeypatch.setattr(pc, "_get_models_from_db", read_models_while_a_credential_lands) + monkeypatch.setattr(pc, "get_models_from_db", read_models_while_a_credential_lands) monkeypatch.setattr(pc, "_update_llm_router", install_models) await pc.add_deployment(prisma_client=fake_prisma, proxy_logging_obj=MagicMock()) @@ -3987,7 +3987,7 @@ async def test_ProxyConfig_add_deployment_loads_db_credentials_even_when_models_ _stub_add_deployment_collaborators(monkeypatch, pc, fake_prisma) monkeypatch.setattr(proxy_server, "general_settings", {"supported_db_objects": ["mcp"]}) models_fetch = AsyncMock(return_value=[]) - monkeypatch.setattr(pc, "_get_models_from_db", models_fetch) + monkeypatch.setattr(pc, "get_models_from_db", models_fetch) await pc.add_deployment(prisma_client=fake_prisma, proxy_logging_obj=MagicMock()) diff --git a/tests/unit/proxy/proxy_server/test_routes_chat_completions.py b/tests/unit/proxy/proxy_server/test_routes_chat_completions.py index b186bb5ef5e..e124bde8120 100644 --- a/tests/unit/proxy/proxy_server/test_routes_chat_completions.py +++ b/tests/unit/proxy/proxy_server/test_routes_chat_completions.py @@ -37,9 +37,7 @@ HAPPY_RESPONSE = { def patched_chat(monkeypatch): """Stub chat-completions pipeline at ProxyBaseLLMRequestProcessing.""" monkeypatch.setattr(proxy_server, "llm_router", MagicMock()) - monkeypatch.setattr( - proxy_server, "proxy_logging_obj", MagicMock(post_call_failure_hook=AsyncMock()) - ) + monkeypatch.setattr(proxy_server, "proxy_logging_obj", MagicMock(post_call_failure_hook=AsyncMock())) async def _fake_process(self, *args, **kwargs): return dict(HAPPY_RESPONSE) @@ -56,9 +54,7 @@ def patched_chat(monkeypatch): def patched_chat_error(monkeypatch): """Variant that makes the pipeline raise -> 400 via _handle_llm_api_exception.""" monkeypatch.setattr(proxy_server, "llm_router", MagicMock()) - monkeypatch.setattr( - proxy_server, "proxy_logging_obj", MagicMock(post_call_failure_hook=AsyncMock()) - ) + monkeypatch.setattr(proxy_server, "proxy_logging_obj", MagicMock(post_call_failure_hook=AsyncMock())) from litellm.proxy._types import ProxyException @@ -66,9 +62,7 @@ def patched_chat_error(monkeypatch): raise ValueError("boom") async def _handler(self, *, e, user_api_key_dict, proxy_logging_obj): - return ProxyException( - message="boom", type="bad_request_error", param="model", code=400 - ) + return ProxyException(message="boom", type="bad_request_error", param="model", code=400) monkeypatch.setattr( common_request_processing.ProxyBaseLLMRequestProcessing, @@ -77,7 +71,7 @@ def patched_chat_error(monkeypatch): ) monkeypatch.setattr( common_request_processing.ProxyBaseLLMRequestProcessing, - "_handle_llm_api_exception", + "handle_llm_api_exception", _handler, ) yield diff --git a/tests/unit/proxy/proxy_server/test_routes_config.py b/tests/unit/proxy/proxy_server/test_routes_config.py index c470fa96d42..4c8265b82c3 100644 --- a/tests/unit/proxy/proxy_server/test_routes_config.py +++ b/tests/unit/proxy/proxy_server/test_routes_config.py @@ -1518,7 +1518,7 @@ def test_get_config_callbacks_excludes_internal_runtime_callbacks(client, auth_a from litellm.integrations.s3_v2 import S3Logger from litellm.integrations.sqs import SQSLogger from litellm.integrations.vector_store_integrations.vector_store_pre_call_hook import VectorStorePreCallHook - from litellm.proxy.hooks.cache_control_check import _PROXY_CacheControlCheck + from litellm.proxy.hooks.cache_control_check import PROXY_CacheControlCheck from litellm.router import Router class _InventoryTestGuardrail(CustomGuardrail): @@ -1546,7 +1546,7 @@ def test_get_config_callbacks_excludes_internal_runtime_callbacks(client, auth_a litellm, "callbacks", [ - _PROXY_CacheControlCheck(), + PROXY_CacheControlCheck(), PROXY_LiteLLMManagedFiles(internal_usage_cache=MagicMock(), prisma_client=MagicMock()), ServiceLogging(), VectorStorePreCallHook(), diff --git a/tests/unit/proxy/proxy_server/test_routes_embeddings.py b/tests/unit/proxy/proxy_server/test_routes_embeddings.py index 98249cb5ad5..f511938ee23 100644 --- a/tests/unit/proxy/proxy_server/test_routes_embeddings.py +++ b/tests/unit/proxy/proxy_server/test_routes_embeddings.py @@ -31,9 +31,7 @@ def patched_embedding(monkeypatch): router.model_names = ["text-embedding-ada-002"] router.get_deployment_by_model_group_name = MagicMock(return_value=None) monkeypatch.setattr(proxy_server, "llm_router", router) - monkeypatch.setattr( - proxy_server, "proxy_logging_obj", MagicMock(post_call_failure_hook=AsyncMock()) - ) + monkeypatch.setattr(proxy_server, "proxy_logging_obj", MagicMock(post_call_failure_hook=AsyncMock())) async def _fake_process(self, *args, **kwargs): return dict(HAPPY_RESPONSE) @@ -51,9 +49,7 @@ def embedding_pipeline_raises(monkeypatch): router = MagicMock() router.model_names = [] monkeypatch.setattr(proxy_server, "llm_router", router) - monkeypatch.setattr( - proxy_server, "proxy_logging_obj", MagicMock(post_call_failure_hook=AsyncMock()) - ) + monkeypatch.setattr(proxy_server, "proxy_logging_obj", MagicMock(post_call_failure_hook=AsyncMock())) from litellm.proxy._types import ProxyException @@ -61,9 +57,7 @@ def embedding_pipeline_raises(monkeypatch): raise ValueError("boom") async def _handler(self, *, e, user_api_key_dict, proxy_logging_obj, version=None): - return ProxyException( - message="boom", type="bad_request_error", param="model", code=400 - ) + return ProxyException(message="boom", type="bad_request_error", param="model", code=400) monkeypatch.setattr( common_request_processing.ProxyBaseLLMRequestProcessing, @@ -72,7 +66,7 @@ def embedding_pipeline_raises(monkeypatch): ) monkeypatch.setattr( common_request_processing.ProxyBaseLLMRequestProcessing, - "_handle_llm_api_exception", + "handle_llm_api_exception", _handler, ) yield diff --git a/tests/unit/proxy/proxy_server/test_routes_invitation.py b/tests/unit/proxy/proxy_server/test_routes_invitation.py index 5b54a63d8a2..35f2991996a 100644 --- a/tests/unit/proxy/proxy_server/test_routes_invitation.py +++ b/tests/unit/proxy/proxy_server/test_routes_invitation.py @@ -101,8 +101,8 @@ def test_invitation_new_non_admin_forbidden(client, auth_as, monkeypatch, mock_p return False # Patch at the proxy_server import site (used by the route). - monkeypatch.setattr(ps, "_user_has_admin_privileges", _no_privileges) - monkeypatch.setattr(common_utils, "_user_has_admin_privileges", _no_privileges) + monkeypatch.setattr(ps, "user_has_admin_privileges", _no_privileges) + monkeypatch.setattr(common_utils, "user_has_admin_privileges", _no_privileges) with auth_as(LitellmUserRoles.INTERNAL_USER): response = client.post("/invitation/new", json={"user_id": "user-target"}) @@ -186,7 +186,7 @@ def test_invitation_info_not_admin_forbidden(client, auth_as, monkeypatch, mock_ monkeypatch.setattr(ps, "prisma_client", mock_prisma) # _user_has_admin_view is referenced from proxy_server's import. - monkeypatch.setattr(ps, "_user_has_admin_view", lambda u: False) + monkeypatch.setattr(ps, "user_api_key_has_admin_view", lambda u: False) with auth_as(LitellmUserRoles.INTERNAL_USER): response = client.get("/invitation/info", params={"invitation_id": "inv-xyz"}) @@ -339,7 +339,7 @@ def test_invitation_delete_non_admin_forbidden( async def _no_privileges(**kwargs): return False - monkeypatch.setattr(ps, "_user_has_admin_privileges", _no_privileges) + monkeypatch.setattr(ps, "user_has_admin_privileges", _no_privileges) with auth_as(LitellmUserRoles.INTERNAL_USER): response = client.post( diff --git a/tests/unit/proxy/proxy_server/test_routes_model_info.py b/tests/unit/proxy/proxy_server/test_routes_model_info.py index 5175d92084c..fae4a6470cf 100644 --- a/tests/unit/proxy/proxy_server/test_routes_model_info.py +++ b/tests/unit/proxy/proxy_server/test_routes_model_info.py @@ -648,7 +648,7 @@ def model_group_info_router(monkeypatch): monkeypatch.setattr(proxy_server, "general_settings", {}) monkeypatch.setattr(proxy_server, "prisma_client", None) monkeypatch.setattr(proxy_server, "user_api_key_cache", None) - monkeypatch.setattr(proxy_server, "_get_model_group_info", model_group_info) + monkeypatch.setattr(proxy_server, "get_model_group_info", model_group_info) from litellm.proxy.agent_endpoints import model_list_helpers diff --git a/tests/unit/proxy/proxy_server/test_streaming_helpers.py b/tests/unit/proxy/proxy_server/test_streaming_helpers.py index 69fa195e9d6..09d71cb4ead 100644 --- a/tests/unit/proxy/proxy_server/test_streaming_helpers.py +++ b/tests/unit/proxy/proxy_server/test_streaming_helpers.py @@ -678,7 +678,7 @@ def _patch_logging_flags(monkeypatch, needs_wrap=False, needs_per_chunk=False): # touching real logging globals. monkeypatch.setattr( ps.ProxyLogging, - "_fire_deferred_stream_logging", + "fire_deferred_stream_logging", staticmethod(lambda request_data: None), ) diff --git a/tests/unit/proxy/public_endpoints/test_public_endpoints.py b/tests/unit/proxy/public_endpoints/test_public_endpoints.py index 99a76d59448..30538c1167f 100644 --- a/tests/unit/proxy/public_endpoints/test_public_endpoints.py +++ b/tests/unit/proxy/public_endpoints/test_public_endpoints.py @@ -548,11 +548,11 @@ def test_public_model_hub_with_healthy_model(): with ( patch("litellm.public_model_groups", ["gpt-3.5-turbo"]), - patch("litellm.proxy.proxy_server._get_model_group_info") as mock_get_info, + patch("litellm.proxy.proxy_server.get_model_group_info") as mock_get_info, patch("litellm.proxy.proxy_server.llm_router", mock_llm_router), patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), patch( - "litellm.proxy.health_endpoints._health_endpoints._convert_health_check_to_dict" + "litellm.proxy.health_endpoints._health_endpoints.convert_health_check_to_dict" ) as mock_convert, ): @@ -606,11 +606,11 @@ def test_public_model_hub_with_unhealthy_model(): with ( patch("litellm.public_model_groups", ["gpt-4"]), - patch("litellm.proxy.proxy_server._get_model_group_info") as mock_get_info, + patch("litellm.proxy.proxy_server.get_model_group_info") as mock_get_info, patch("litellm.proxy.proxy_server.llm_router", mock_llm_router), patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), patch( - "litellm.proxy.health_endpoints._health_endpoints._convert_health_check_to_dict" + "litellm.proxy.health_endpoints._health_endpoints.convert_health_check_to_dict" ) as mock_convert, ): @@ -655,7 +655,7 @@ def test_public_model_hub_without_health_check(): with ( patch("litellm.public_model_groups", ["claude-3"]), - patch("litellm.proxy.proxy_server._get_model_group_info") as mock_get_info, + patch("litellm.proxy.proxy_server.get_model_group_info") as mock_get_info, patch("litellm.proxy.proxy_server.llm_router", mock_llm_router), patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), ): @@ -737,11 +737,11 @@ def test_public_model_hub_mixed_health_statuses(): with ( patch("litellm.public_model_groups", ["gpt-3.5-turbo", "gpt-4", "claude-3"]), - patch("litellm.proxy.proxy_server._get_model_group_info") as mock_get_info, + patch("litellm.proxy.proxy_server.get_model_group_info") as mock_get_info, patch("litellm.proxy.proxy_server.llm_router", mock_llm_router), patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), patch( - "litellm.proxy.health_endpoints._health_endpoints._convert_health_check_to_dict" + "litellm.proxy.health_endpoints._health_endpoints.convert_health_check_to_dict" ) as mock_convert, ): diff --git a/tests/unit/proxy/response_api_endpoints/test_endpoints.py b/tests/unit/proxy/response_api_endpoints/test_endpoints.py index 456c3609419..5dac1d25572 100644 --- a/tests/unit/proxy/response_api_endpoints/test_endpoints.py +++ b/tests/unit/proxy/response_api_endpoints/test_endpoints.py @@ -137,7 +137,7 @@ async def test_responses_api_background_polling_rejects_missing_input(): async def return_exception(*, e: Exception, **kwargs: object) -> Exception: return e - processor._handle_llm_api_exception = AsyncMock(side_effect=return_exception) + processor.handle_llm_api_exception = AsyncMock(side_effect=return_exception) processor.common_processing_pre_call_logic = AsyncMock(return_value=({"model": "gpt-4o"}, MagicMock())) async def receive(): @@ -728,11 +728,17 @@ class TestResponsesWSFirstFrameModelAuth: async def fake_llm_call(): return None + authenticated_models: Final[list[str]] = [] + + async def record_model_auth(*, model: str, **_kwargs: object) -> None: + authenticated_models.append(model) + with ( patch( "litellm.proxy.response_api_endpoints.endpoints._enforce_responses_ws_first_frame_model_auth", new_callable=AsyncMock, - ) as mock_model_auth, + side_effect=record_model_auth, + ), patch( "litellm.proxy.response_api_endpoints.endpoints.ProxyBaseLLMRequestProcessing", return_value=processor, @@ -749,7 +755,7 @@ class TestResponsesWSFirstFrameModelAuth: user_api_key_dict=MagicMock(), ) - mock_model_auth.assert_awaited_once() + assert authenticated_models == ["gpt-4o-mini"] @pytest.mark.asyncio @pytest.mark.parametrize("nested", [False, True]) @@ -957,14 +963,15 @@ class TestResponsesWSFirstFrameModelAuth: request = Request({"type": "http", "method": "POST", "path": "/v1/responses", "headers": []}) user_api_key_dict = MagicMock() llm_router = MagicMock() + empty_settings: Final[dict[str, object]] = {} with ( patch( - "litellm.proxy.auth.user_api_key_auth._enforce_key_and_fallback_model_access", + "litellm.proxy.auth.user_api_key_auth.enforce_key_and_fallback_model_access", new_callable=AsyncMock, ) as mock_key_check, patch( - "litellm.proxy.auth.user_api_key_auth._run_centralized_common_checks", + "litellm.proxy.auth.user_api_key_auth.run_centralized_common_checks", new_callable=AsyncMock, ) as mock_common_checks, patch( @@ -973,7 +980,7 @@ class TestResponsesWSFirstFrameModelAuth: ), patch("litellm.proxy.proxy_server.master_key", "sk-test"), patch("litellm.proxy.proxy_server.user_custom_auth", None), - patch("litellm.proxy.proxy_server.general_settings", {}), + patch("litellm.proxy.proxy_server.general_settings", empty_settings), ): await _enforce_responses_ws_first_frame_model_auth( request=request, @@ -1000,7 +1007,7 @@ class TestResponsesWSFirstFrameModelAuth: class TestReadWSModelFromFirstFrameErrors: @pytest.mark.asyncio - async def test_timeout_closes_without_error_frame(self): + async def test_transport_error_first_frame_closes_with_internal_error(self): import asyncio from litellm.proxy.response_api_endpoints.endpoints import ( @@ -1016,7 +1023,7 @@ class TestReadWSModelFromFirstFrameErrors: assert result is None ws.send_text.assert_not_awaited() - ws.close.assert_awaited_once_with(code=1008, reason="Timed out waiting for first message") + ws.close.assert_awaited_once_with(code=1011, reason="Internal server error") @pytest.mark.asyncio async def test_invalid_json_sends_error_and_closes(self): @@ -1108,6 +1115,34 @@ class TestReadWSModelFromFirstFrameErrors: ws.close.assert_not_awaited() + +@pytest.mark.parametrize( + "configured,expected", + [ + (None, 3600.0), + (60, 60.0), + (1200, 1200.0), + (7200, 7200.0), + (59, 3600.0), + (0, 3600.0), + (9000, 3600.0), + ("not-a-number", 3600.0), + ], +) +def test_responses_ws_session_limit_resolution(monkeypatch, configured, expected): + from litellm.proxy.proxy_server import general_settings + from litellm.proxy.response_api_endpoints.endpoints import ( + _resolve_responses_ws_session_limit_seconds, + ) + + if configured is None: + monkeypatch.delitem(general_settings, "responses_websocket_session_limit_seconds", raising=False) + else: + monkeypatch.setitem(general_settings, "responses_websocket_session_limit_seconds", configured) + + assert _resolve_responses_ws_session_limit_seconds() == expected + + class TestManagedResponsesSameProvider: def _handler(self, model, custom_llm_provider=None): from litellm.responses.streaming_iterator import ( @@ -1260,7 +1295,7 @@ def test_cursor_chat_completions_input_body_uses_responses_pipeline_and_strips_s from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.common_utils.http_parsing_utils import ( - _read_request_body as real_read_request_body, + read_request_body as real_read_request_body, ) from litellm.types.llms.openai import ResponsesAPIResponse @@ -1297,7 +1332,7 @@ def test_cursor_chat_completions_input_body_uses_responses_pipeline_and_strips_s with ( patch.object(ps, "llm_router", mock_router), patch( - "litellm.proxy.response_api_endpoints.endpoints._read_request_body", + "litellm.proxy.response_api_endpoints.endpoints.read_request_body", side_effect=capturing_read_request_body, ), ): @@ -1407,9 +1442,9 @@ class TestCursorMessagesArmToolNormalization: seen = {} async def fake_chat_completion(request, fastapi_response, model, user_api_key_dict): - from litellm.proxy.common_utils.http_parsing_utils import _read_request_body + from litellm.proxy.common_utils.http_parsing_utils import read_request_body - seen["body"] = await _read_request_body(request=request) + seen["body"] = await read_request_body(request=request) return {"id": "chatcmpl-fake", "object": "chat.completion", "choices": []} app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(api_key=MASTER_KEY) @@ -1470,9 +1505,9 @@ class TestCursorMessagesArmToolNormalization: seen = {} async def fake_chat_completion(request, fastapi_response, model, user_api_key_dict): - from litellm.proxy.common_utils.http_parsing_utils import _read_request_body + from litellm.proxy.common_utils.http_parsing_utils import read_request_body - seen["body"] = await _read_request_body(request=request) + seen["body"] = await read_request_body(request=request) return {"id": "chatcmpl-fake", "object": "chat.completion", "choices": []} body = { @@ -1949,9 +1984,9 @@ class TestCursorModelSuffixResolutionEndToEnd: seen = {} async def fake_chat_completion(request, fastapi_response, model, user_api_key_dict): - from litellm.proxy.common_utils.http_parsing_utils import _read_request_body + from litellm.proxy.common_utils.http_parsing_utils import read_request_body - seen["body"] = await _read_request_body(request=request) + seen["body"] = await read_request_body(request=request) return {"id": "chatcmpl-fake", "object": "chat.completion", "choices": []} app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(api_key=MASTER_KEY) @@ -2031,7 +2066,7 @@ def _cursor_budget_auth_env(base_model: str, spend: float): from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.hooks.model_max_budget_limiter import ( VIRTUAL_KEY_SPEND_CACHE_KEY_PREFIX, - _PROXY_VirtualKeyModelMaxBudgetLimiter, + PROXY_VirtualKeyModelMaxBudgetLimiter, ) valid_token = UserAPIKeyAuth( @@ -2039,7 +2074,7 @@ def _cursor_budget_auth_env(base_model: str, spend: float): token="hashed-cursor-budget-token", model_max_budget={base_model: {"budget_limit": 0.00001, "time_period": "1d"}}, ) - limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=DualCache()) + limiter = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=DualCache()) limiter.dual_cache.in_memory_cache.set_cache( key=f"{VIRTUAL_KEY_SPEND_CACHE_KEY_PREFIX}:{valid_token.token}:{base_model}:1d", value=spend, @@ -2127,14 +2162,14 @@ class TestCursorVariantResolvedBeforeAuth: def _run_with_recording_auth(self, mock_router, request_model: str): from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.auth.user_api_key_auth import user_api_key_auth - from litellm.proxy.common_utils.http_parsing_utils import _read_request_body + from litellm.proxy.common_utils.http_parsing_utils import read_request_body from fastapi import Request bodies_seen_by_auth = [] async def recording_auth(request: Request) -> UserAPIKeyAuth: - bodies_seen_by_auth.append(await _read_request_body(request=request)) + bodies_seen_by_auth.append(await read_request_body(request=request)) return UserAPIKeyAuth(api_key="sk-test-cursor") async def fake_chat_completion(request, fastapi_response, model, user_api_key_dict): diff --git a/tests/unit/proxy/spend_tracking/test_search_api_logging.py b/tests/unit/proxy/spend_tracking/test_search_api_logging.py index 92108f19c8e..bcb079cf37b 100644 --- a/tests/unit/proxy/spend_tracking/test_search_api_logging.py +++ b/tests/unit/proxy/spend_tracking/test_search_api_logging.py @@ -18,7 +18,7 @@ import litellm from litellm import Router from litellm.caching import DualCache from litellm.proxy._types import UserAPIKeyAuth -from litellm.proxy.hooks.proxy_track_cost_callback import _ProxyDBLogger +from litellm.proxy.hooks.proxy_track_cost_callback import ProxyDBLogger from litellm.proxy.spend_tracking.spend_management_endpoints import view_spend_logs from litellm.proxy.utils import ProxyLogging, hash_token, update_spend from litellm.llms.base_llm.search.transformation import SearchResponse, SearchResult @@ -129,7 +129,7 @@ async def test_search_api_logging_and_cost_tracking(prisma_client): setattr(litellm.proxy.proxy_server, "proxy_logging_obj", proxy_logging_obj) # Call the track_cost_callback directly to simulate what happens after a search - proxy_db_logger = _ProxyDBLogger() + proxy_db_logger = ProxyDBLogger() # Simulate the kwargs that would be passed from the search endpoint request_id = "search_test_123" diff --git a/tests/unit/proxy/spend_tracking/test_spend_management_endpoints.py b/tests/unit/proxy/spend_tracking/test_spend_management_endpoints.py index 2e3b3c13bb3..8cc8dccd13e 100644 --- a/tests/unit/proxy/spend_tracking/test_spend_management_endpoints.py +++ b/tests/unit/proxy/spend_tracking/test_spend_management_endpoints.py @@ -273,7 +273,7 @@ from litellm.proxy._types import ( SpendLogsPayload, UserAPIKeyAuth, ) -from litellm.proxy.hooks.proxy_track_cost_callback import _ProxyDBLogger +from litellm.proxy.hooks.proxy_track_cost_callback import ProxyDBLogger from litellm.proxy.management.teams import authz as team_access from litellm.proxy.proxy_server import app from litellm.proxy.spend_tracking import spend_management_endpoints @@ -3684,7 +3684,7 @@ class TestSpendLogsPayload: @pytest.mark.asyncio async def test_spend_logs_payload_e2e(self): - litellm.callbacks = [_ProxyDBLogger(message_logging=False)] + litellm.callbacks = [ProxyDBLogger(message_logging=False)] # litellm.turn_on_debug() with ( diff --git a/tests/unit/proxy/spend_tracking/test_spend_tracking_utils.py b/tests/unit/proxy/spend_tracking/test_spend_tracking_utils.py index d4b397e3fa2..0df6eecd4b9 100644 --- a/tests/unit/proxy/spend_tracking/test_spend_tracking_utils.py +++ b/tests/unit/proxy/spend_tracking/test_spend_tracking_utils.py @@ -27,10 +27,10 @@ from litellm.proxy.spend_tracking.spend_tracking_utils import ( _get_session_id_for_spend_log, _get_spend_logs_metadata, _get_vector_store_request_for_spend_logs_payload, - _is_master_key, + is_master_key, _redact_logged_api_key, _redact_prompt_leaks_in_error_string, - _sanitize_error_information_for_spend_logs, + sanitize_error_information_for_spend_logs, _sanitize_guardrail_information_for_spend_logs, _sanitize_request_body_for_spend_logs_payload, _scrub_raw_model_from_error_information, @@ -1476,7 +1476,7 @@ def test_get_logging_payload_persists_no_raw_model_for_a_prompt_shaped_moderatio model=_RAW_MODEL_WITH_PROMPT, llm_provider="openai", ) - error_information: Final = _sanitize_error_information_for_spend_logs( + error_information: Final = sanitize_error_information_for_spend_logs( StandardLoggingPayloadSetup.get_error_information( original_exception=provider_rejection, traceback_str=( @@ -1643,7 +1643,7 @@ async def test_api_key_preserved_through_failure_hook_to_database(): If this test fails in CI/CD, the build MUST fail. """ from litellm.proxy._types import UserAPIKeyAuth - from litellm.proxy.hooks.proxy_track_cost_callback import _ProxyDBLogger + from litellm.proxy.hooks.proxy_track_cost_callback import ProxyDBLogger from litellm.proxy.utils import hash_token # Setup @@ -1727,7 +1727,7 @@ async def test_api_key_preserved_through_failure_hook_to_database(): exception = Exception("BadRequestError: Invalid parameter 'invalid_param'") # Execute the ACTUAL failure hook code path - logger = _ProxyDBLogger() + logger = ProxyDBLogger() with patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj): await logger.async_post_call_failure_hook( @@ -2793,19 +2793,19 @@ class TestIsMasterKey: def test_none_api_key_returns_false(self): """Regression: _is_master_key(None, 'sk-master') should return False, not raise TypeError.""" - assert _is_master_key(api_key=None, _master_key="sk-master-key") is False + assert is_master_key(api_key=None, _master_key="sk-master-key") is False def test_none_master_key_returns_false(self): - assert _is_master_key(api_key="sk-some-key", _master_key=None) is False + assert is_master_key(api_key="sk-some-key", _master_key=None) is False def test_both_none_returns_false(self): - assert _is_master_key(api_key=None, _master_key=None) is False + assert is_master_key(api_key=None, _master_key=None) is False def test_matching_key_returns_true(self): - assert _is_master_key(api_key="sk-master", _master_key="sk-master") is True + assert is_master_key(api_key="sk-master", _master_key="sk-master") is True def test_non_matching_key_returns_false(self): - assert _is_master_key(api_key="sk-other", _master_key="sk-master") is False + assert is_master_key(api_key="sk-other", _master_key="sk-master") is False def test_master_key_hash_is_rejected(self): """ @@ -2816,7 +2816,7 @@ class TestIsMasterKey: master = "sk-master-key-123" hashed = hash_token(master) - assert _is_master_key(api_key=hashed, _master_key=master) is False + assert is_master_key(api_key=hashed, _master_key=master) is False def test_sanitize_request_body_strips_secret_fields(): @@ -3199,7 +3199,7 @@ def test_sanitize_error_information_redacts_when_not_storing_prompts( ), } - sanitized = _sanitize_error_information_for_spend_logs(error_info) + sanitized = sanitize_error_information_for_spend_logs(error_info) assert sanitized is not None assert "leaked-prompt-content" not in sanitized["error_message"] @@ -3224,7 +3224,7 @@ def test_sanitize_error_information_skips_redaction_when_storing_prompts( "error_message": ('OpenAIException - {"error":{"input":[{"role":"user","content":"kept"}]}}'), } - sanitized = _sanitize_error_information_for_spend_logs(error_info) + sanitized = sanitize_error_information_for_spend_logs(error_info) assert sanitized is not None # User opted in via store_prompts_in_spend_logs — no key-level redaction. @@ -3251,7 +3251,7 @@ def test_sanitize_error_information_caps_size_regardless_of_prompt_flag( "error_message": huge_error, } - sanitized = _sanitize_error_information_for_spend_logs(error_info) + sanitized = sanitize_error_information_for_spend_logs(error_info) assert sanitized is not None assert len(sanitized["error_message"]) < len(huge_error) @@ -3260,7 +3260,7 @@ def test_sanitize_error_information_caps_size_regardless_of_prompt_flag( def test_sanitize_error_information_none_passthrough(): - assert _sanitize_error_information_for_spend_logs(None) is None + assert sanitize_error_information_for_spend_logs(None) is None @patch("litellm.proxy.spend_tracking.spend_tracking_utils.should_store_prompts_and_responses_in_spend_logs") @@ -3291,7 +3291,7 @@ def test_sanitize_error_information_reproduces_lit_2992(mock_should_store): "error_message": error_message, } - sanitized = _sanitize_error_information_for_spend_logs(error_info) + sanitized = sanitize_error_information_for_spend_logs(error_info) assert sanitized is not None assert huge_conversation_blob not in sanitized["error_message"] @@ -3370,7 +3370,7 @@ def test_sanitize_error_information_redacts_traceback_when_not_storing_prompts( "error_message": "invalid request", } - sanitized = _sanitize_error_information_for_spend_logs(error_info) + sanitized = sanitize_error_information_for_spend_logs(error_info) assert sanitized is not None assert "tb-leaked-prompt" not in sanitized["traceback"] @@ -3394,7 +3394,7 @@ def test_sanitize_error_information_skips_traceback_redaction_when_storing_promp "error_message": "invalid request", } - sanitized = _sanitize_error_information_for_spend_logs(error_info) + sanitized = sanitize_error_information_for_spend_logs(error_info) assert sanitized is not None assert "tb-kept" in sanitized["traceback"] @@ -3510,7 +3510,7 @@ def test_sanitize_error_information_redacts_pydantic_assignment_form( ), } - sanitized = _sanitize_error_information_for_spend_logs(error_info) + sanitized = sanitize_error_information_for_spend_logs(error_info) assert sanitized is not None assert "leaked-via-pydantic-msg" not in sanitized["error_message"] @@ -3535,7 +3535,7 @@ def test_sanitize_error_information_persists_no_raw_model_for_an_unknown_model_r ): error_information: Final = StandardLoggingPayloadSetup.get_error_information(original_exception=original_exception) - sanitized: Final = _sanitize_error_information_for_spend_logs( + sanitized: Final = sanitize_error_information_for_spend_logs( error_information, original_exception=original_exception ) diff --git a/tests/unit/proxy/test_aiohttp_session_recovery.py b/tests/unit/proxy/test_aiohttp_session_recovery.py index 29bd9a491b7..1ac67b76313 100644 --- a/tests/unit/proxy/test_aiohttp_session_recovery.py +++ b/tests/unit/proxy/test_aiohttp_session_recovery.py @@ -51,7 +51,7 @@ async def test_add_shared_session_recreates_closed_session(): ): with patch.object( proxy_server_module, - "_initialize_shared_aiohttp_session", + "initialize_shared_aiohttp_session", new_callable=AsyncMock, return_value=new_session, ) as mock_init: @@ -83,7 +83,7 @@ async def test_add_shared_session_handles_recreation_failure(): ): with patch.object( proxy_server_module, - "_initialize_shared_aiohttp_session", + "initialize_shared_aiohttp_session", new_callable=AsyncMock, return_value=None, ): @@ -112,7 +112,7 @@ async def test_add_shared_session_handles_recreation_exception(): ): with patch.object( proxy_server_module, - "_initialize_shared_aiohttp_session", + "initialize_shared_aiohttp_session", new_callable=AsyncMock, side_effect=RuntimeError("connection pool exhausted"), ): @@ -166,7 +166,7 @@ async def test_add_shared_session_concurrent_recreation_uses_lock(): ): with patch.object( proxy_server_module, - "_initialize_shared_aiohttp_session", + "initialize_shared_aiohttp_session", new_callable=AsyncMock, side_effect=mock_init, ): diff --git a/tests/unit/proxy/test_batch_x_litellm_model_encoding.py b/tests/unit/proxy/test_batch_x_litellm_model_encoding.py index 3161fe99e68..f8c734c4efd 100644 --- a/tests/unit/proxy/test_batch_x_litellm_model_encoding.py +++ b/tests/unit/proxy/test_batch_x_litellm_model_encoding.py @@ -91,7 +91,7 @@ async def test_create_batch_with_x_litellm_model_encodes_batch_id(): with ( patch( - "litellm.proxy.batches_endpoints.endpoints._read_request_body", + "litellm.proxy.batches_endpoints.endpoints.read_request_body", new=AsyncMock( return_value={ "input_file_id": "file-input456", @@ -206,7 +206,7 @@ async def test_create_batch_with_x_litellm_model_encodes_output_and_error_file_i with ( patch( - "litellm.proxy.batches_endpoints.endpoints._read_request_body", + "litellm.proxy.batches_endpoints.endpoints.read_request_body", new=AsyncMock( return_value={ "input_file_id": "file-input456", @@ -293,7 +293,7 @@ async def test_create_batch_without_x_litellm_model_returns_raw_ids(monkeypatch) with ( patch( - "litellm.proxy.batches_endpoints.endpoints._read_request_body", + "litellm.proxy.batches_endpoints.endpoints.read_request_body", new=AsyncMock( return_value={ "input_file_id": "file-input456", diff --git a/tests/unit/proxy/test_budget_reservation.py b/tests/unit/proxy/test_budget_reservation.py index c8e4df1030f..4635e145914 100644 --- a/tests/unit/proxy/test_budget_reservation.py +++ b/tests/unit/proxy/test_budget_reservation.py @@ -1937,7 +1937,7 @@ async def test_should_skip_reservation_when_counter_initialization_fails( return_value=0.5, ), patch( - "litellm.proxy.proxy_server._ensure_spend_counter_initialized", + "litellm.proxy.proxy_server.ensure_spend_counter_initialized", side_effect=RuntimeError("redis unavailable"), ), patch( @@ -1993,7 +1993,7 @@ async def test_should_release_tracked_entry_when_reservation_fails_after_increme side_effect=fail_after_increment, ), patch( - "litellm.proxy.proxy_server._invalidate_spend_counter", + "litellm.proxy.proxy_server.invalidate_spend_counter", side_effect=RuntimeError("invalidate unavailable"), ), ): @@ -2941,7 +2941,7 @@ async def _never_ending_stream(): def _drive_streaming_cancel(valid_token, iterator_hook): streaming_logging_obj = MagicMock() streaming_logging_obj.async_post_call_streaming_iterator_hook = iterator_hook - streaming_logging_obj._arelease_max_parallel_requests_on_disconnect = AsyncMock() + streaming_logging_obj.arelease_max_parallel_requests_on_disconnect = AsyncMock() generator = ProxyBaseLLMRequestProcessing.async_streaming_data_generator( response=MagicMock(), user_api_key_dict=valid_token, @@ -2985,7 +2985,7 @@ async def test_streaming_cancel_before_any_chunk_reconciles_to_input_cost( key="spend:key:key-cancel-no-chunk" ) == pytest.approx(0.5) assert reservation["finalized"] is True - streaming_logging_obj._arelease_max_parallel_requests_on_disconnect.assert_awaited_once() + streaming_logging_obj.arelease_max_parallel_requests_on_disconnect.assert_awaited_once() @pytest.mark.asyncio @@ -3019,7 +3019,7 @@ async def test_streaming_cancel_after_chunk_keeps_reservation( key="spend:key:key-cancel-after-chunk" ) == pytest.approx(2.0) assert reservation.get("finalized") is not True - streaming_logging_obj._arelease_max_parallel_requests_on_disconnect.assert_awaited_once() + streaming_logging_obj.arelease_max_parallel_requests_on_disconnect.assert_awaited_once() @pytest.mark.asyncio @@ -3098,7 +3098,7 @@ async def test_streaming_cancel_while_holding_back_provider_output_keeps_reserva streaming_logging_obj = MagicMock() streaming_logging_obj.async_post_call_streaming_iterator_hook = ping_then_cancel - streaming_logging_obj._arelease_max_parallel_requests_on_disconnect = AsyncMock() + streaming_logging_obj.arelease_max_parallel_requests_on_disconnect = AsyncMock() generator = ProxyBaseLLMRequestProcessing.async_streaming_data_generator( response=response, user_api_key_dict=valid_token, @@ -3156,7 +3156,7 @@ async def test_streaming_cancel_in_slow_path_before_yield_refunds(spend_counter_ streaming_logging_obj = MagicMock() streaming_logging_obj.async_post_call_streaming_iterator_hook = one_chunk - streaming_logging_obj._arelease_max_parallel_requests_on_disconnect = AsyncMock() + streaming_logging_obj.arelease_max_parallel_requests_on_disconnect = AsyncMock() # On the slow path the per-chunk hook is awaited before the chunk is yielded # to the client; cancel there. Nothing has reached the client yet. streaming_logging_obj.async_post_call_streaming_hook = AsyncMock( @@ -3189,7 +3189,7 @@ async def test_streaming_cancel_in_slow_path_before_yield_refunds(spend_counter_ key="spend:key:key-cancel-slowpath" ) == pytest.approx(0.5) assert reservation["finalized"] is True - streaming_logging_obj._arelease_max_parallel_requests_on_disconnect.assert_awaited_once() + streaming_logging_obj.arelease_max_parallel_requests_on_disconnect.assert_awaited_once() @pytest.mark.asyncio @@ -3219,7 +3219,7 @@ async def test_streaming_disconnect_after_consuming_chunk_keeps_reservation( key="spend:key:key-disconnect-after-chunk" ) == pytest.approx(2.0) assert reservation.get("finalized") is not True - streaming_logging_obj._arelease_max_parallel_requests_on_disconnect.assert_awaited_once() + streaming_logging_obj.arelease_max_parallel_requests_on_disconnect.assert_awaited_once() @pytest.mark.asyncio diff --git a/tests/unit/proxy/test_chat_completion_metadata.py b/tests/unit/proxy/test_chat_completion_metadata.py index 38dcdc13c50..7b84684d6f2 100644 --- a/tests/unit/proxy/test_chat_completion_metadata.py +++ b/tests/unit/proxy/test_chat_completion_metadata.py @@ -10,25 +10,17 @@ async def test_chat_completion_metadata_population(): # Setup request = MagicMock(spec=Request) # Mock _read_request_body to return a dict - with patch( - "litellm.proxy.proxy_server._read_request_body", new_callable=AsyncMock - ) as mock_read_body: + with patch("litellm.proxy.proxy_server.read_request_body", new_callable=AsyncMock) as mock_read_body: mock_read_body.return_value = {"model": "gpt-3.5-turbo", "messages": []} - user_api_key_dict = UserAPIKeyAuth( - user_id="test_user_id", team_id="test_team_id", org_id="test_org_id" - ) + user_api_key_dict = UserAPIKeyAuth(user_id="test_user_id", team_id="test_team_id", org_id="test_org_id") fastapi_response = MagicMock(spec=Response) # Mock ProxyBaseLLMRequestProcessing - with patch( - "litellm.proxy.proxy_server.ProxyBaseLLMRequestProcessing" - ) as MockProcessor: + with patch("litellm.proxy.proxy_server.ProxyBaseLLMRequestProcessing") as MockProcessor: mock_instance = MockProcessor.return_value - mock_instance.base_process_llm_request = AsyncMock( - return_value={"choices": []} - ) + mock_instance.base_process_llm_request = AsyncMock(return_value={"choices": []}) # Execute await chat_completion( @@ -57,9 +49,7 @@ async def test_embedding_metadata_population(): from UserAPIKeyAuth. """ # Setup - with patch( - "litellm.proxy.proxy_server.ProxyBaseLLMRequestProcessing.base_process_llm_request" - ): + with patch("litellm.proxy.proxy_server.ProxyBaseLLMRequestProcessing.base_process_llm_request"): with patch( "litellm.proxy.proxy_server.ProxyBaseLLMRequestProcessing.__init__", return_value=None, @@ -72,15 +62,11 @@ async def test_embedding_metadata_population(): # Create a mock Request object mock_request = MagicMock(spec=Request) - mock_request.json = AsyncMock( - return_value={"model": "gpt-3.5-turbo", "input": "hello"} - ) + mock_request.json = AsyncMock(return_value={"model": "gpt-3.5-turbo", "input": "hello"}) # Mock _read_request_body to return our data with patch( - "litellm.proxy.proxy_server._read_request_body", - new=AsyncMock( - return_value={"model": "gpt-3.5-turbo", "input": "hello"} - ), + "litellm.proxy.proxy_server.read_request_body", + new=AsyncMock(return_value={"model": "gpt-3.5-turbo", "input": "hello"}), ): # Call the endpoint function directly await embeddings( @@ -98,12 +84,8 @@ async def test_embedding_metadata_population(): else: data_arg = call_args.args[0] - assert ( - data_arg["metadata"]["user_api_key_user_id"] == "test_user_id_emb" - ) - assert ( - data_arg["metadata"]["user_api_key_team_id"] == "test_team_id_emb" - ) + assert data_arg["metadata"]["user_api_key_user_id"] == "test_user_id_emb" + assert data_arg["metadata"]["user_api_key_team_id"] == "test_team_id_emb" assert data_arg["metadata"]["user_api_key_org_id"] == "test_org_id_emb" @@ -112,28 +94,20 @@ async def test_completion_metadata_population(): # Setup request = MagicMock(spec=Request) # Mock _read_request_body to return a dict - with patch( - "litellm.proxy.proxy_server._read_request_body", new_callable=AsyncMock - ) as mock_read_body: + with patch("litellm.proxy.proxy_server.read_request_body", new_callable=AsyncMock) as mock_read_body: mock_read_body.return_value = { "model": "gpt-3.5-turbo-instruct", "prompt": "test", } - user_api_key_dict = UserAPIKeyAuth( - user_id="test_user_id_2", team_id="test_team_id_2", org_id="test_org_id_2" - ) + user_api_key_dict = UserAPIKeyAuth(user_id="test_user_id_2", team_id="test_team_id_2", org_id="test_org_id_2") fastapi_response = MagicMock(spec=Response) # Mock ProxyBaseLLMRequestProcessing - with patch( - "litellm.proxy.proxy_server.ProxyBaseLLMRequestProcessing" - ) as MockProcessor: + with patch("litellm.proxy.proxy_server.ProxyBaseLLMRequestProcessing") as MockProcessor: mock_instance = MockProcessor.return_value - mock_instance.base_process_llm_request = AsyncMock( - return_value={"choices": []} - ) + mock_instance.base_process_llm_request = AsyncMock(return_value={"choices": []}) # Execute await completion( diff --git a/tests/unit/proxy/test_common_request_processing.py b/tests/unit/proxy/test_common_request_processing.py index 368f55b1eb1..1aed60ee5e2 100644 --- a/tests/unit/proxy/test_common_request_processing.py +++ b/tests/unit/proxy/test_common_request_processing.py @@ -44,14 +44,14 @@ from litellm.proxy.common_request_processing import ( CostBreakdownHeaderValues, _has_attribute_error_in_chain, include_guardrail_response_requested, - _is_azure_model_router_request, + is_azure_model_router_request, open_sse_before_first_byte, resolve_litellm_call_id, ttft_keepalive_interval, _override_openai_response_model, _parse_event_data_for_error, _resolve_per_request_model_group_alias, - _should_return_raw_model_name, + should_return_raw_model_name, _sse_error_frames, _UpstreamClosingStreamingResponse, create_response, @@ -3020,7 +3020,7 @@ class TestOverrideOpenAIResponseModel: ], ) def test_raw_model_name_toggle_metadata(self, request_data, expected): - assert _should_return_raw_model_name(request_data) is expected + assert should_return_raw_model_name(request_data) is expected def test_override_model_preserves_fallback_model_when_fallback_occurred_object( self, @@ -3471,17 +3471,17 @@ class TestIsAzureModelRouterRequest: """Tests for _is_azure_model_router_request helper""" def test_detects_model_router_with_underscore(self): - assert _is_azure_model_router_request("azure_ai/model_router") is True - assert _is_azure_model_router_request("azure_ai/model_router/my-deployment") is True + assert is_azure_model_router_request("azure_ai/model_router") is True + assert is_azure_model_router_request("azure_ai/model_router/my-deployment") is True def test_detects_model_router_with_hyphen(self): - assert _is_azure_model_router_request("azure_ai/model-router") is True - assert _is_azure_model_router_request("model-router") is True + assert is_azure_model_router_request("azure_ai/model-router") is True + assert is_azure_model_router_request("model-router") is True def test_rejects_regular_models(self): - assert _is_azure_model_router_request("azure_ai/gpt-4") is False - assert _is_azure_model_router_request("gpt-4") is False - assert _is_azure_model_router_request("openai/gpt-3.5-turbo") is False + assert is_azure_model_router_request("azure_ai/gpt-4") is False + assert is_azure_model_router_request("gpt-4") is False + assert is_azure_model_router_request("openai/gpt-3.5-turbo") is False class TestStreamingOverheadHeader: @@ -5263,7 +5263,7 @@ class TestStreamingClientDisconnectLogging: fire_spy = MagicMock() monkeypatch.setattr( - "litellm.proxy.utils.ProxyLogging._fire_deferred_stream_logging", + "litellm.proxy.utils.ProxyLogging.fire_deferred_stream_logging", fire_spy, ) @@ -5297,7 +5297,7 @@ class TestStreamingClientDisconnectLogging: fire_spy = MagicMock() monkeypatch.setattr( - "litellm.proxy.utils.ProxyLogging._fire_deferred_stream_logging", + "litellm.proxy.utils.ProxyLogging.fire_deferred_stream_logging", fire_spy, ) @@ -5328,7 +5328,7 @@ class TestStreamingClientDisconnectLogging: ) monkeypatch.setattr( - "litellm.proxy.utils.ProxyLogging._fire_deferred_stream_logging", + "litellm.proxy.utils.ProxyLogging.fire_deferred_stream_logging", MagicMock(), ) @@ -6944,7 +6944,7 @@ class TestPreCallWithFallbacksOnLocalRateLimit: ProxyRateLimitError, ) from litellm.proxy.hooks.parallel_request_limiter import ( - _PROXY_MaxParallelRequestsHandler, + PROXY_MaxParallelRequestsHandler, ) from litellm.proxy.utils import InternalUsageCache @@ -6962,7 +6962,7 @@ class TestPreCallWithFallbacksOnLocalRateLimit: # Real per-key per-model TPM limiter + a key carrying the customer's # `model_tpm_limit` metadata (only the primary is capped). - limiter = _PROXY_MaxParallelRequestsHandler( + limiter = PROXY_MaxParallelRequestsHandler( internal_usage_cache=InternalUsageCache(DualCache()) ) user_api_key_dict = UserAPIKeyAuth( @@ -7064,11 +7064,11 @@ class TestPreCallWithFallbacksOnLocalRateLimit: ``add_litellm_data_to_request`` with a live OTel span, ``function_setup``, then the limiter.""" from litellm.caching.caching import DualCache from litellm.proxy import proxy_server - from litellm.proxy.hooks.parallel_request_limiter_v3 import _PROXY_MaxParallelRequestsHandler_v3 + from litellm.proxy.hooks.parallel_request_limiter_v3 import PROXY_MaxParallelRequestsHandler_v3 from litellm.proxy.utils import InternalUsageCache monkeypatch.setattr(proxy_server, "prisma_client", None) - limiter = _PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=InternalUsageCache(DualCache())) + limiter = PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=InternalUsageCache(DualCache())) limiter_models: list[str] = [] async def run_limiter( @@ -7561,7 +7561,7 @@ class TestStreamingClientDisconnectBilling: try: response = await self._start_partial_stream() proxy_logging_obj = types.SimpleNamespace( - _arelease_max_parallel_requests_on_disconnect=AsyncMock(), + arelease_max_parallel_requests_on_disconnect=AsyncMock(), ) billed = await _bill_partial_streamed_spend_on_disconnect( @@ -7581,7 +7581,7 @@ class TestStreamingClientDisconnectBilling: finally: litellm.callbacks = original_callbacks - proxy_logging_obj._arelease_max_parallel_requests_on_disconnect.assert_not_called() + proxy_logging_obj.arelease_max_parallel_requests_on_disconnect.assert_not_called() @pytest.mark.asyncio async def test_disconnect_without_billable_chunks_releases_slot(self): @@ -7596,7 +7596,7 @@ class TestStreamingClientDisconnectBilling: # No chunks to assemble -> billing dispatches no success event. empty_response = types.SimpleNamespace(chunks=[], messages=None) proxy_logging_obj = types.SimpleNamespace( - _arelease_max_parallel_requests_on_disconnect=AsyncMock(), + arelease_max_parallel_requests_on_disconnect=AsyncMock(), ) await ProxyBaseLLMRequestProcessing._finalize_streaming_generator_cleanup( @@ -7609,7 +7609,7 @@ class TestStreamingClientDisconnectBilling: proxy_logging_obj=proxy_logging_obj, ) - proxy_logging_obj._arelease_max_parallel_requests_on_disconnect.assert_awaited_once() + proxy_logging_obj.arelease_max_parallel_requests_on_disconnect.assert_awaited_once() async def _bill_and_collect_success_event(self, prepare=None, request_data=None): recorder = _RecordingSuccessLogger() @@ -8287,7 +8287,7 @@ class TestPerRequestModelGroupAlias: monkeypatch.setattr( litellm.proxy.common_request_processing, - "_check_and_merge_model_level_guardrails", + "check_and_merge_model_level_guardrails", recording_merge, ) diff --git a/tests/unit/proxy/test_dynamic_mcp_route.py b/tests/unit/proxy/test_dynamic_mcp_route.py index 83963fbd962..4a8dddb1bf4 100644 --- a/tests/unit/proxy/test_dynamic_mcp_route.py +++ b/tests/unit/proxy/test_dynamic_mcp_route.py @@ -33,7 +33,7 @@ _IS_ACCESS_GROUP = "litellm.proxy.proxy_server._is_mcp_access_group_cached" _USER_API_KEY_CACHE = "litellm.proxy.proxy_server.user_api_key_cache" _GET_ACCESS_GROUP_SERVERS = ( "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp." - "MCPRequestHandler._get_mcp_servers_from_access_groups" + "MCPRequestHandler.get_mcp_servers_from_access_groups" ) _FORWARD = "litellm.proxy.proxy_server._mcp_forward_as_path" _RESOLVE_CSV = "litellm.proxy.proxy_server._resolve_mcp_csv_tokens" @@ -286,10 +286,10 @@ async def test_dynamic_mcp_route_resolves_toolset(): async def fake_stream(fn, scope, receive): nonlocal captured_toolset_id from litellm.proxy._experimental.mcp_server.server import ( - _mcp_active_toolset_id, + mcp_active_toolset_id, ) - captured_toolset_id = _mcp_active_toolset_id.get() + captured_toolset_id = mcp_active_toolset_id.get() captured_scope.update(scope) with ( diff --git a/tests/unit/proxy/test_health_check_functions.py b/tests/unit/proxy/test_health_check_functions.py index 1b2fc73fca7..2c0198cd538 100644 --- a/tests/unit/proxy/test_health_check_functions.py +++ b/tests/unit/proxy/test_health_check_functions.py @@ -12,7 +12,7 @@ from litellm.proxy.health_endpoints._health_endpoints import ( _aggregate_health_check_results, _build_model_param_to_info_mapping, _perform_health_check_and_save, - _save_background_health_checks_to_db, + save_background_health_checks_to_db, _save_health_check_results_if_changed, _save_health_check_to_db, latest_health_checks_endpoint, @@ -427,7 +427,7 @@ async def test_save_background_health_checks_to_db(): start_time = 1234567890.0 - persisted = await _save_background_health_checks_to_db( + persisted = await save_background_health_checks_to_db( mock_prisma, model_list, healthy_endpoints, @@ -535,7 +535,7 @@ async def test_save_background_health_checks_to_db_returns_false_when_a_write_fa mock_prisma.save_health_check_result = AsyncMock(return_value=None) model_list, healthy_endpoints, unhealthy_endpoints = _one_model_setup() - persisted = await _save_background_health_checks_to_db( + persisted = await save_background_health_checks_to_db( mock_prisma, model_list, healthy_endpoints, unhealthy_endpoints, 1234567890.0, "background_health_check" ) @@ -552,7 +552,7 @@ async def test_save_background_health_checks_to_db_writes_nothing_when_the_lates mock_prisma.save_health_check_result = AsyncMock(return_value={"id": "row"}) model_list, healthy_endpoints, unhealthy_endpoints = _one_model_setup() - persisted = await _save_background_health_checks_to_db( + persisted = await save_background_health_checks_to_db( mock_prisma, model_list, healthy_endpoints, unhealthy_endpoints, 1234567890.0, "background_health_check" ) @@ -562,7 +562,7 @@ async def test_save_background_health_checks_to_db_writes_nothing_when_the_lates @pytest.mark.asyncio async def test_save_background_health_checks_to_db_no_prisma(): """Test graceful handling when no prisma client""" - result = await _save_background_health_checks_to_db(None, [], [], [], 0.0, "background_health_check") + result = await save_background_health_checks_to_db(None, [], [], [], 0.0, "background_health_check") assert result is False @@ -582,7 +582,7 @@ async def test_save_background_health_checks_to_db_exception_handling(): # Must not raise (the health check loop has to survive a DB outage) but must report # the failure, so the window lock can be released for another pod to retry - persisted = await _save_background_health_checks_to_db( + persisted = await save_background_health_checks_to_db( mock_prisma, model_list, [], [], 0.0, "background_health_check" ) @@ -653,7 +653,7 @@ async def test_save_background_health_checks_compares_raw_checked_at_against_utc {"model_name": "fresh-model", "model_info": {"id": "fresh-id"}, "litellm_params": {"model": "openai/fresh"}}, ] - await _save_background_health_checks_to_db( + await save_background_health_checks_to_db( mock_prisma, model_list, [{"model": "openai/stale"}, {"model": "openai/fresh"}], diff --git a/tests/unit/proxy/test_health_check_max_tokens.py b/tests/unit/proxy/test_health_check_max_tokens.py index e3641ac2c81..091de1e24b3 100644 --- a/tests/unit/proxy/test_health_check_max_tokens.py +++ b/tests/unit/proxy/test_health_check_max_tokens.py @@ -12,7 +12,7 @@ from litellm.proxy.health_check import ( _is_strategy_router_deployment, _resolve_health_check_max_tokens, resolve_health_check_mode, - _update_litellm_params_for_health_check, + update_litellm_params_for_health_check, ) @@ -26,7 +26,7 @@ async def test_update_litellm_params_max_tokens_default(monkeypatch): model_info = {} litellm_params = {"model": "gpt-4"} - updated_params = _update_litellm_params_for_health_check(model_info, litellm_params) + updated_params = update_litellm_params_for_health_check(model_info, litellm_params) assert updated_params["max_tokens"] == 16 @@ -39,7 +39,7 @@ async def test_update_litellm_params_max_tokens_custom(): model_info = {"health_check_max_tokens": 5} litellm_params = {"model": "gpt-4"} - updated_params = _update_litellm_params_for_health_check(model_info, litellm_params) + updated_params = update_litellm_params_for_health_check(model_info, litellm_params) assert updated_params["max_tokens"] == 5 @@ -52,7 +52,7 @@ async def test_update_litellm_params_max_tokens_wildcard(): model_info = {} litellm_params = {"model": "openai/*"} - updated_params = _update_litellm_params_for_health_check(model_info, litellm_params) + updated_params = update_litellm_params_for_health_check(model_info, litellm_params) assert "max_tokens" not in updated_params @@ -102,7 +102,7 @@ async def test_background_health_check_max_tokens_env_var(monkeypatch): model_info = {} litellm_params = {"model": "azure/gpt-4"} - updated_params = _update_litellm_params_for_health_check(model_info, litellm_params) + updated_params = update_litellm_params_for_health_check(model_info, litellm_params) assert updated_params["max_tokens"] == 10 @@ -118,7 +118,7 @@ async def test_per_model_overrides_global_env_var(monkeypatch): model_info = {"health_check_max_tokens": 5} litellm_params = {"model": "azure/gpt-4"} - updated_params = _update_litellm_params_for_health_check(model_info, litellm_params) + updated_params = update_litellm_params_for_health_check(model_info, litellm_params) assert updated_params["max_tokens"] == 5 @@ -133,7 +133,7 @@ async def test_global_env_var_applies_to_wildcard_models(monkeypatch): model_info = {} litellm_params = {"model": "openai/*"} - updated_params = _update_litellm_params_for_health_check(model_info, litellm_params) + updated_params = update_litellm_params_for_health_check(model_info, litellm_params) assert updated_params["max_tokens"] == 15 @@ -184,12 +184,12 @@ async def test_background_split_env_reasoning_vs_non_reasoning(monkeypatch): litellm_params = {"model": "azure/gpt-4"} with patch.object(hc_module.litellm, "supports_reasoning", return_value=False): - updated = _update_litellm_params_for_health_check(model_info, litellm_params) + updated = update_litellm_params_for_health_check(model_info, litellm_params) assert updated["max_tokens"] == 16 litellm_params2 = {"model": "openai/o1"} with patch.object(hc_module.litellm, "supports_reasoning", return_value=True): - updated2 = _update_litellm_params_for_health_check(model_info, litellm_params2) + updated2 = update_litellm_params_for_health_check(model_info, litellm_params2) assert updated2["max_tokens"] == 50 @@ -202,7 +202,7 @@ async def test_reasoning_env_precedence_over_global(monkeypatch): litellm_params = {"model": "openai/gpt-5.4"} with patch.object(hc_module.litellm, "supports_reasoning", return_value=True): - updated = _update_litellm_params_for_health_check(model_info, litellm_params) + updated = update_litellm_params_for_health_check(model_info, litellm_params) assert updated["max_tokens"] == 20 @@ -215,7 +215,7 @@ async def test_non_reasoning_uses_global_when_reasoning_env_set(monkeypatch): litellm_params = {"model": "azure/gpt-4"} with patch.object(hc_module.litellm, "supports_reasoning", return_value=False): - updated = _update_litellm_params_for_health_check(model_info, litellm_params) + updated = update_litellm_params_for_health_check(model_info, litellm_params) assert updated["max_tokens"] == 10 @@ -250,7 +250,7 @@ def test_image_generation_mode_skips_max_tokens(): model_info = {"mode": "image_generation"} litellm_params = {"model": "openai/dall-e-3", "api_key": "sk-test"} - updated = _update_litellm_params_for_health_check(model_info, litellm_params) + updated = update_litellm_params_for_health_check(model_info, litellm_params) assert "max_tokens" not in updated # connection-level params must still pass through unchanged @@ -267,7 +267,7 @@ def test_health_check_max_tokens_value_is_ignored_for_non_chat_modes(): model_info = {"mode": "image_generation", "health_check_max_tokens": 50} litellm_params = {"model": "openai/dall-e-3"} - updated = _update_litellm_params_for_health_check(model_info, litellm_params) + updated = update_litellm_params_for_health_check(model_info, litellm_params) assert "max_tokens" not in updated @@ -277,7 +277,7 @@ def test_chat_mode_still_injects_max_tokens(): model_info = {"mode": "chat"} litellm_params = {"model": "gpt-4"} - updated = _update_litellm_params_for_health_check(model_info, litellm_params) + updated = update_litellm_params_for_health_check(model_info, litellm_params) assert updated["max_tokens"] == 16 @@ -287,7 +287,7 @@ def test_no_mode_still_injects_max_tokens(): model_info: dict = {} litellm_params = {"model": "gpt-4"} - updated = _update_litellm_params_for_health_check(model_info, litellm_params) + updated = update_litellm_params_for_health_check(model_info, litellm_params) assert updated["max_tokens"] == 16 @@ -305,7 +305,7 @@ def test_no_mode_still_injects_max_tokens(): @pytest.mark.parametrize("mode", ["chat", "completion", "responses"]) def test_chat_style_modes_inject_max_tokens(mode): - updated = _update_litellm_params_for_health_check({"mode": mode}, {"model": f"openai/dummy-{mode}"}) + updated = update_litellm_params_for_health_check({"mode": mode}, {"model": f"openai/dummy-{mode}"}) assert updated["max_tokens"] == 16 @@ -326,7 +326,7 @@ def test_chat_style_modes_inject_max_tokens(mode): ], ) def test_non_chat_modes_skip_max_tokens(mode): - updated = _update_litellm_params_for_health_check({"mode": mode}, {"model": f"openai/dummy-{mode}"}) + updated = update_litellm_params_for_health_check({"mode": mode}, {"model": f"openai/dummy-{mode}"}) assert "max_tokens" not in updated @@ -339,7 +339,7 @@ def test_explicit_override_true_forces_injection_outside_allowlist(): } litellm_params = {"model": "openai/some-future-image-model"} - updated = _update_litellm_params_for_health_check(model_info, litellm_params) + updated = update_litellm_params_for_health_check(model_info, litellm_params) assert updated["max_tokens"] == 16 @@ -349,7 +349,7 @@ def test_explicit_override_false_suppresses_injection_inside_allowlist(): model_info = {"mode": "chat", "health_check_supports_max_tokens": False} litellm_params = {"model": "openai/strict-schema-chat"} - updated = _update_litellm_params_for_health_check(model_info, litellm_params) + updated = update_litellm_params_for_health_check(model_info, litellm_params) assert "max_tokens" not in updated @@ -358,29 +358,29 @@ def test_update_litellm_params_health_check_reasoning_effort(): """model_info.health_check_reasoning_effort sets reasoning_effort for chat-style health checks.""" model_info = {"health_check_reasoning_effort": "low"} litellm_params = {"model": "openai/gpt-5", "api_key": "x"} - out = _update_litellm_params_for_health_check(model_info, dict(litellm_params)) + out = update_litellm_params_for_health_check(model_info, dict(litellm_params)) assert out.get("reasoning_effort") == "low" model_info = {"mode": "chat", "health_check_reasoning_effort": "none"} - out = _update_litellm_params_for_health_check(model_info, {"model": "openai/gpt-5", "api_key": "x"}) + out = update_litellm_params_for_health_check(model_info, {"model": "openai/gpt-5", "api_key": "x"}) assert out.get("reasoning_effort") == "none" model_info = {"mode": "completion", "health_check_reasoning_effort": "low"} - out = _update_litellm_params_for_health_check(model_info, {"model": "openai/gpt-5", "api_key": "x"}) + out = update_litellm_params_for_health_check(model_info, {"model": "openai/gpt-5", "api_key": "x"}) assert out.get("reasoning_effort") == "low" model_info = { "health_check_reasoning_effort": {"effort": "none", "summary": "auto"}, } - out = _update_litellm_params_for_health_check(model_info, {"model": "openai/gpt-5.1", "api_key": "x"}) + out = update_litellm_params_for_health_check(model_info, {"model": "openai/gpt-5.1", "api_key": "x"}) assert out.get("reasoning_effort") == {"effort": "none", "summary": "auto"} model_info = {"mode": "embedding", "health_check_reasoning_effort": "low"} - out = _update_litellm_params_for_health_check(model_info, {"model": "text-embedding-3-small", "api_key": "x"}) + out = update_litellm_params_for_health_check(model_info, {"model": "text-embedding-3-small", "api_key": "x"}) assert "reasoning_effort" not in out model_info = {} - out = _update_litellm_params_for_health_check(model_info, {"model": "openai/gpt-4o", "api_key": "x"}) + out = update_litellm_params_for_health_check(model_info, {"model": "openai/gpt-4o", "api_key": "x"}) assert "reasoning_effort" not in out @@ -408,7 +408,7 @@ def test_bedrock_embedding_without_explicit_mode_skips_max_tokens(deployment_mod """Embedding mode auto-detected from model cost map -> no max_tokens, provider pinned.""" assert resolve_health_check_mode({}, {"model": deployment_model}) == "embedding" - updated = _update_litellm_params_for_health_check({}, {"model": deployment_model}) + updated = update_litellm_params_for_health_check({}, {"model": deployment_model}) assert "max_tokens" not in updated assert updated["custom_llm_provider"] == "bedrock" @@ -427,7 +427,7 @@ def test_resolve_health_check_mode_unknown_model_returns_none(): def test_bedrock_chat_without_mode_still_injects_max_tokens_and_pins_provider(): """Regression guard: chat-style Bedrock deployments keep max_tokens and get the provider pin.""" - updated = _update_litellm_params_for_health_check( + updated = update_litellm_params_for_health_check( {}, {"model": "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0"} ) @@ -443,7 +443,7 @@ def test_bedrock_prefix_strip_preserves_explicit_custom_llm_provider(): not clobber a more specific one, otherwise a converse deployment would be probed against the Invoke endpoint and report a spurious failure. """ - updated = _update_litellm_params_for_health_check( + updated = update_litellm_params_for_health_check( {}, { "model": "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0", @@ -509,7 +509,7 @@ def test_mantle_claude_without_mode_resolves_to_anthropic_messages(deployment_mo """Mantle only serves Claude over /anthropic/v1/messages, so that is the probe surface by default.""" assert resolve_health_check_mode({}, {"model": deployment_model}) == "anthropic_messages" - updated = _update_litellm_params_for_health_check({}, {"model": deployment_model}) + updated = update_litellm_params_for_health_check({}, {"model": deployment_model}) assert updated["max_tokens"] == 16 assert [message["role"] for message in updated["messages"]] == ["user"] @@ -597,7 +597,7 @@ def test_autodetected_embedding_skips_reasoning_effort(): Bedrock embedding probe, which embeddings reject as an unknown field. The mode is now resolved from the cost map, so embeddings are excluded. """ - updated = _update_litellm_params_for_health_check( + updated = update_litellm_params_for_health_check( {"health_check_reasoning_effort": "low"}, {"model": "bedrock/amazon.titan-embed-text-v2:0"}, ) @@ -649,7 +649,7 @@ def test_health_check_params_merge_into_probe_params(): """health_check_params reach the probe request for the deployment that declares them.""" media_source = {"s3Location": {"uri": "s3://my-bucket/clip.mp4"}} - updated = _update_litellm_params_for_health_check( + updated = update_litellm_params_for_health_check( {"mode": "chat", "health_check_params": {"mediaSource": media_source}}, {"model": "bedrock/us.twelvelabs.pegasus-1-2-v1:0"}, ) @@ -674,7 +674,7 @@ def test_health_check_params_lose_to_dedicated_health_check_knobs(): "health_check_reasoning_effort": "none", } - updated = _update_litellm_params_for_health_check(model_info, {"model": "openai/dummy"}) + updated = update_litellm_params_for_health_check(model_info, {"model": "openai/dummy"}) assert updated["max_tokens"] == 5 assert updated["model"] == "openai/cheap-model" @@ -684,7 +684,7 @@ def test_health_check_params_lose_to_dedicated_health_check_knobs(): def test_health_check_params_lose_to_the_audio_speech_voice_knob(): """health_check_voice still wins for audio_speech deployments.""" - updated = _update_litellm_params_for_health_check( + updated = update_litellm_params_for_health_check( { "mode": "audio_speech", "health_check_params": {"voice": "sage", "response_format": "wav"}, @@ -704,7 +704,7 @@ def test_health_check_params_lose_to_the_audio_speech_voice_knob(): def test_health_check_params_ignored_when_not_a_dict(bad_value, caplog): """A misconfigured health_check_params is skipped with a warning instead of breaking the probe.""" with caplog.at_level(logging.WARNING, logger="litellm.proxy.health_check"): - updated = _update_litellm_params_for_health_check( + updated = update_litellm_params_for_health_check( {"mode": "chat", "health_check_params": bad_value}, {"model": "openai/dummy"}, ) @@ -716,7 +716,7 @@ def test_health_check_params_ignored_when_not_a_dict(bad_value, caplog): def test_health_check_params_apply_to_non_chat_modes(): """Non-chat probes get health_check_params too, and still no max_tokens.""" - updated = _update_litellm_params_for_health_check( + updated = update_litellm_params_for_health_check( {"mode": "embedding", "health_check_params": {"dimensions": 8}}, {"model": "bedrock/amazon.titan-embed-text-v2:0"}, ) @@ -731,7 +731,7 @@ async def _pegasus_health_check_request_body( monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) litellm.in_memory_llm_clients_cache.flush_cache() - litellm_params = _update_litellm_params_for_health_check( + litellm_params = update_litellm_params_for_health_check( model_info, { "model": "bedrock/us.twelvelabs.pegasus-1-2-v1:0", diff --git a/tests/unit/proxy/test_litellm_pre_call_utils.py b/tests/unit/proxy/test_litellm_pre_call_utils.py index d4624719565..c6c473ae71e 100644 --- a/tests/unit/proxy/test_litellm_pre_call_utils.py +++ b/tests/unit/proxy/test_litellm_pre_call_utils.py @@ -22,9 +22,9 @@ from litellm.proxy.litellm_pre_call_utils import ( LiteLLMProxyRequestSetup, _apply_credential_overrides_from_model_config, _extract_credential_from_entry, - _get_dynamic_logging_metadata, + get_dynamic_logging_metadata, _get_enforced_params, - _get_metadata_variable_name, + get_metadata_variable_name, _match_and_track_policies, _promoted_trace_control_fields, _resolve_credential_from_model_config, @@ -84,45 +84,45 @@ class TestGetMetadataVariableName: def test_returns_litellm_metadata_for_thread_routes(self): request = self._make_request("/v1/threads/thread_123/messages") - assert _get_metadata_variable_name(request) == "litellm_metadata" + assert get_metadata_variable_name(request) == "litellm_metadata" def test_returns_litellm_metadata_for_assistant_routes(self): request = self._make_request("/v1/assistants/asst_123") - assert _get_metadata_variable_name(request) == "litellm_metadata" + assert get_metadata_variable_name(request) == "litellm_metadata" def test_returns_litellm_metadata_for_batches_route(self): request = self._make_request("/v1/batches") - assert _get_metadata_variable_name(request) == "litellm_metadata" + assert get_metadata_variable_name(request) == "litellm_metadata" def test_returns_litellm_metadata_for_messages_route(self): request = self._make_request("/v1/messages") - assert _get_metadata_variable_name(request) == "litellm_metadata" + assert get_metadata_variable_name(request) == "litellm_metadata" def test_returns_litellm_metadata_for_files_route(self): request = self._make_request("/v1/files") - assert _get_metadata_variable_name(request) == "litellm_metadata" + assert get_metadata_variable_name(request) == "litellm_metadata" def test_returns_metadata_for_chat_completions(self): request = self._make_request("/chat/completions") - assert _get_metadata_variable_name(request) == "metadata" + assert get_metadata_variable_name(request) == "metadata" def test_returns_metadata_for_completions(self): request = self._make_request("/v1/completions") - assert _get_metadata_variable_name(request) == "metadata" + assert get_metadata_variable_name(request) == "metadata" def test_returns_metadata_for_embeddings(self): request = self._make_request("/v1/embeddings") - assert _get_metadata_variable_name(request) == "metadata" + assert get_metadata_variable_name(request) == "metadata" def test_returns_litellm_metadata_for_bedrock_invoke(self): # GH#30629: bedrock passthrough must use litellm_metadata # to prevent key-level tags from leaking into provider body request = self._make_request("/bedrock/model/us.anthropic.claude-sonnet-4-6/invoke") - assert _get_metadata_variable_name(request) == "litellm_metadata" + assert get_metadata_variable_name(request) == "litellm_metadata" def test_returns_litellm_metadata_for_bedrock_converse(self): request = self._make_request("/bedrock/model/us.anthropic.claude-sonnet-4-6/converse") - assert _get_metadata_variable_name(request) == "litellm_metadata" + assert get_metadata_variable_name(request) == "litellm_metadata" def test_get_enforced_params_for_service_account_settings(): @@ -902,7 +902,7 @@ async def test_post_guardrail_snapshot_preserves_logging_only_masking_in_spend_l monkeypatch: pytest.MonkeyPatch, pre_call_ran: bool ) -> None: from litellm.litellm_core_utils.litellm_logging import Logging - from litellm.proxy.guardrails.guardrail_hooks.presidio import _OPTIONAL_PresidioPIIMasking + from litellm.proxy.guardrails.guardrail_hooks.presidio import OPTIONAL_PresidioPIIMasking from litellm.proxy.litellm_pre_call_utils import refresh_proxy_server_request_body_snapshot from litellm.proxy.spend_tracking.spend_tracking_utils import _get_proxy_server_request_for_spend_logs_payload @@ -921,7 +921,7 @@ async def test_post_guardrail_snapshot_preserves_logging_only_masking_in_spend_l logging_obj.update_messages(messages) snapshot: Final = logging_obj.shadow_eval_request_snapshot assert (snapshot is not None) is pre_call_ran - guardrail: Final = _OPTIONAL_PresidioPIIMasking( + guardrail: Final = OPTIONAL_PresidioPIIMasking( mock_testing=True, logging_only=True, mock_redacted_text={"text": "email [EMAIL]", "items": []} ) @@ -2419,7 +2419,7 @@ def test_get_dynamic_logging_metadata_with_arize_team_logging(): mock_proxy_config = MagicMock() # Call the function - result = _get_dynamic_logging_metadata(user_api_key_dict=user_api_key_dict, proxy_config=mock_proxy_config) + result = get_dynamic_logging_metadata(user_api_key_dict=user_api_key_dict, proxy_config=mock_proxy_config) # Verify the result assert result is not None @@ -2466,7 +2466,7 @@ def test_get_dynamic_logging_metadata_ignores_env_reference_from_key_metadata( team_metadata={}, ) - result = _get_dynamic_logging_metadata(user_api_key_dict=user_api_key_dict, proxy_config=MagicMock()) + result = get_dynamic_logging_metadata(user_api_key_dict=user_api_key_dict, proxy_config=MagicMock()) assert result is None @@ -3674,6 +3674,106 @@ def test_add_litellm_metadata_from_request_headers_explicit_trace_id_beats_trace assert data["litellm_session_id"] == "explicit-trace-id-value" +def test_add_litellm_metadata_from_request_headers_body_trace_id_beats_traceparent(): + headers = {"traceparent": "00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01"} + data = {"metadata": {"trace_id": "caller-chosen-trace-id"}} + LiteLLMProxyRequestSetup.add_litellm_metadata_from_request_headers( + headers=headers, data=data, _metadata_variable_name="metadata" + ) + assert data["metadata"]["trace_id"] == "caller-chosen-trace-id" + assert "litellm_trace_id" not in data + + +def test_add_litellm_metadata_from_request_headers_body_session_id_beats_baggage(): + headers = {"baggage": "session.id=baggage-session-42"} + data = {"metadata": {"session_id": "caller-chosen-session-id"}} + LiteLLMProxyRequestSetup.add_litellm_metadata_from_request_headers( + headers=headers, data=data, _metadata_variable_name="metadata" + ) + assert data["metadata"]["session_id"] == "caller-chosen-session-id" + assert "litellm_session_id" not in data + + +def test_add_litellm_metadata_from_request_headers_body_steering_is_per_field(): + headers = { + "traceparent": "00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01", + "baggage": "session.id=baggage-session-42", + } + data = {"metadata": {"trace_id": "caller-chosen-trace-id"}} + LiteLLMProxyRequestSetup.add_litellm_metadata_from_request_headers( + headers=headers, data=data, _metadata_variable_name="metadata" + ) + assert data["metadata"]["trace_id"] == "caller-chosen-trace-id" + assert data["litellm_session_id"] == "baggage-session-42" + + +def test_add_litellm_metadata_from_request_headers_litellm_metadata_steering_honoured(): + headers = {"traceparent": "00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01"} + data = {"litellm_metadata": {"trace_id": "caller-chosen-trace-id"}} + LiteLLMProxyRequestSetup.add_litellm_metadata_from_request_headers( + headers=headers, data=data, _metadata_variable_name="litellm_metadata" + ) + assert data["litellm_metadata"]["trace_id"] == "caller-chosen-trace-id" + assert "litellm_trace_id" not in data + + +@pytest.mark.parametrize("empty_session_id", ["", None]) +def test_add_litellm_metadata_from_request_headers_empty_body_session_id_falls_back_to_baggage( + empty_session_id: str | None, +): + headers = {"baggage": "session.id=baggage-session-42"} + data = {"metadata": {"session_id": empty_session_id}} + LiteLLMProxyRequestSetup.add_litellm_metadata_from_request_headers( + headers=headers, data=data, _metadata_variable_name="metadata" + ) + assert data["litellm_session_id"] == "baggage-session-42" + assert data["metadata"]["session_id"] == "baggage-session-42" + + +@pytest.mark.parametrize("field", ["trace_id", "session_id"]) +def test_add_litellm_metadata_from_request_headers_promoted_metadata_beats_headers(field: str): + headers = { + "traceparent": "00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01", + "baggage": "session.id=baggage-session-42", + } + data = {"metadata": {field: "caller-chosen"}, "litellm_metadata": {}} + LiteLLMProxyRequestSetup.add_litellm_metadata_from_request_headers( + headers=headers, data=data, _metadata_variable_name="litellm_metadata" + ) + assert field not in data["litellm_metadata"] + assert f"litellm_{field}" not in data + + +def test_add_litellm_metadata_from_request_headers_equal_ids_still_stamp_root_fields(): + headers = { + "traceparent": "00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01", + "baggage": "session.id=matching-session-42", + } + data = { + "metadata": {"trace_id": "4bf92f3577b34da6a3ce929d0e0e4736", "session_id": "matching-session-42"}, + } + LiteLLMProxyRequestSetup.add_litellm_metadata_from_request_headers( + headers=headers, data=data, _metadata_variable_name="metadata" + ) + assert data["litellm_trace_id"] == "4bf92f3577b34da6a3ce929d0e0e4736" + assert data["litellm_session_id"] == "matching-session-42" + assert data["metadata"]["trace_id"] == "4bf92f3577b34da6a3ce929d0e0e4736" + assert data["metadata"]["session_id"] == "matching-session-42" + + +@pytest.mark.parametrize("non_string_session_id", [4815162342, True, {"session": "nested"}]) +def test_add_litellm_metadata_from_request_headers_non_string_body_session_id_falls_back_to_baggage( + non_string_session_id: object, +): + headers = {"baggage": "session.id=header-session-42"} + data = {"metadata": {"session_id": non_string_session_id}} + LiteLLMProxyRequestSetup.add_litellm_metadata_from_request_headers( + headers=headers, data=data, _metadata_variable_name="metadata" + ) + assert data["litellm_session_id"] == "header-session-42" + assert data["metadata"]["session_id"] == "header-session-42" + + def _otel_span_with_trace_id(trace_id: int) -> NonRecordingSpan: return NonRecordingSpan(SpanContext(trace_id=trace_id, span_id=0x00F067AA0BA902B7, is_remote=False)) @@ -4008,7 +4108,7 @@ async def test_team_guardrails_append_to_key_guardrails(): team_metadata={"guardrails": ["team-guardrail-1", "key-guardrail-1"]}, ) - with patch("litellm.proxy.utils._premium_user_check"): + with patch("litellm.proxy.utils.premium_user_check"): updated_data = await add_litellm_data_to_request( data=data, request=request_mock, @@ -4057,7 +4157,7 @@ async def test_request_guardrails_do_not_override_key_guardrails(): "guardrails": [], } - with patch("litellm.proxy.utils._premium_user_check"): + with patch("litellm.proxy.utils.premium_user_check"): updated_data_empty = await add_litellm_data_to_request( data=data_with_empty, request=request_mock, @@ -4103,7 +4203,7 @@ async def test_project_guardrails_merge_with_key_and_team(): project_metadata={"guardrails": ["project-guardrail-1", "team-guardrail-1"]}, ) - with patch("litellm.proxy.utils._premium_user_check"): + with patch("litellm.proxy.utils.premium_user_check"): updated_data = await add_litellm_data_to_request( data=data, request=request_mock, @@ -4152,7 +4252,7 @@ async def test_project_guardrails_only(): project_metadata={"guardrails": ["project-guardrail-1", "project-guardrail-2"]}, ) - with patch("litellm.proxy.utils._premium_user_check"): + with patch("litellm.proxy.utils.premium_user_check"): updated_data = await add_litellm_data_to_request( data=data, request=request_mock, @@ -5459,7 +5559,7 @@ async def test_team_guardrail_merges_with_global_policy(): attachment_registry._initialized = True try: - with patch("litellm.proxy.utils._premium_user_check"): + with patch("litellm.proxy.utils.premium_user_check"): await move_guardrails_to_metadata( data=data, _metadata_variable_name="metadata", @@ -7127,7 +7227,7 @@ class TestPromotedTraceControlFields: ) def test_returns_litellm_metadata_for_responses_route(self): - assert _get_metadata_variable_name(self._make_request("/v1/responses")) == "litellm_metadata" + assert get_metadata_variable_name(self._make_request("/v1/responses")) == "litellm_metadata" def test_promotes_trace_prefixed_and_allow_listed_fields(self): requester_metadata = { @@ -8222,6 +8322,85 @@ async def test_missing_session_id_omit_keeps_client_supplied_session_id(): assert _spend_log_session_id(updated) == "client-session-1" +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("path", "client_body"), + [ + ("/v1/chat/completions", {"model": "gpt-4o", "messages": [], "metadata": {"session_id": ""}}), + ("/v1/responses", {"model": "gpt-4o", "input": "hi"}), + ], +) +async def test_missing_session_id_reject_accepts_baggage_session_id(path: str, client_body: dict[str, object]): + request = _request_for(path) + request.headers = {"baggage": "session.id=baggage-session-42"} + + updated = await add_litellm_data_to_request( + data=client_body, + request=request, + user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"), + proxy_config=MagicMock(), + general_settings={"missing_session_id": "reject"}, + ) + + assert updated["litellm_session_id"] == "baggage-session-42" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("path", ["/v1/responses", "/v1/messages"]) +@pytest.mark.parametrize("policy", [None, "reject", "generate"]) +async def test_promoted_caller_trace_ids_beat_traceparent_and_baggage(path: str, policy: str | None): + request = _request_for(path) + request.headers = { + "traceparent": "00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01", + "baggage": "session.id=baggage-session-42", + } + + updated = await add_litellm_data_to_request( + data={"model": "gpt-4o", "metadata": {"trace_id": "caller-trace", "session_id": "caller-session"}}, + request=request, + user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"), + proxy_config=MagicMock(), + general_settings={"missing_session_id": policy} if policy else {}, + ) + + assert updated["litellm_metadata"]["trace_id"] == "caller-trace" + assert updated["litellm_metadata"]["session_id"] == "caller-session" + + +@pytest.mark.asyncio +async def test_missing_session_id_reject_ignores_requester_session_id_shadowed_by_empty_litellm_metadata(): + with pytest.raises(ProxyException) as exc_info: + await add_litellm_data_to_request( + data={ + "model": "gpt-4o", + "input": "hi", + "metadata": {"session_id": "caller-session"}, + "litellm_metadata": {"session_id": ""}, + }, + request=_request_for("/v1/responses"), + user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"), + proxy_config=MagicMock(), + general_settings={"missing_session_id": "reject"}, + ) + assert exc_info.value.code == "400" + + +def test_add_litellm_metadata_from_request_headers_empty_litellm_metadata_field_falls_back_to_headers(): + headers = { + "traceparent": "00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01", + "baggage": "session.id=baggage-session-42", + } + data = { + "metadata": {"trace_id": "caller-trace", "session_id": "caller-session"}, + "litellm_metadata": {"trace_id": "", "session_id": ""}, + } + LiteLLMProxyRequestSetup.add_litellm_metadata_from_request_headers( + headers=headers, data=data, _metadata_variable_name="litellm_metadata" + ) + assert data["litellm_trace_id"] == "4bf92f3577b34da6a3ce929d0e0e4736" + assert data["litellm_session_id"] == "baggage-session-42" + + @pytest.mark.asyncio @pytest.mark.parametrize( "client_body", @@ -8341,6 +8520,91 @@ async def test_missing_session_id_generate_reuses_traceparent_trace_id(): assert _spend_log_session_id(updated) == "4bf92f3577b34da6a3ce929d0e0e4736" +@pytest.mark.asyncio +@pytest.mark.parametrize("path", ["/v1/responses", "/v1/messages"]) +@pytest.mark.parametrize( + "headers", + [ + {"traceparent": "00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01"}, + {}, + ], + ids=["with_traceparent", "no_traceparent"], +) +async def test_missing_session_id_generate_reuses_promoted_caller_trace_id(path: str, headers: dict[str, str]): + request = _request_for(path) + request.headers = headers + + updated = await add_litellm_data_to_request( + data={"model": "gpt-4o", "input": "hi", "metadata": {"trace_id": "caller-trace"}}, + request=request, + user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"), + proxy_config=MagicMock(), + general_settings={"missing_session_id": "generate"}, + ) + + caller_trace_msg: Final = "generate must derive the session from the caller's un-promoted metadata.trace_id" + assert updated["litellm_session_id"] == "caller-trace", caller_trace_msg + assert updated["litellm_metadata"]["session_id"] == "caller-trace", caller_trace_msg + assert updated["litellm_metadata"][SESSION_ID_GENERATED_METADATA_KEY] is True, ( + "the derived session id must still be marked generated" + ) + assert _spend_log_session_id(updated, "litellm_metadata") == "caller-trace", ( + "spend log and callback session ids must agree on the caller trace id" + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("policy", ["generate", "reject"]) +async def test_missing_session_id_policy_promotes_caller_session_to_root_field(policy: str): + """A caller-supplied usable session id satisfies the missing_session_id policies on + litellm_metadata routes (where it is not yet the managed metadata field) and must also + land on the root ``litellm_session_id`` field: consumers that read the root field + (router fallbacks, spend logs, sandbox reuse) otherwise mint a fresh uuid4 per request.""" + request = _request_for("/v1/responses") + + updated = await add_litellm_data_to_request( + data={ + "model": "gpt-4o", + "input": "hi", + "litellm_trace_id": "root-trace-42", + "metadata": {"session_id": "caller-session-42"}, + }, + request=request, + user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"), + proxy_config=MagicMock(), + general_settings={"missing_session_id": policy}, + ) + + assert updated["litellm_trace_id"] == "root-trace-42" + assert updated["litellm_session_id"] == "caller-session-42" + assert updated["litellm_metadata"]["session_id"] == "caller-session-42" + assert SESSION_ID_GENERATED_METADATA_KEY not in updated["litellm_metadata"] + + +@pytest.mark.asyncio +async def test_missing_session_id_generate_ignores_non_string_caller_session_id(): + """A non-string session id is not a usable session: the generate policy must fall through + to generation instead of letting an unusable value strand the root session field.""" + request = _request_for("/v1/responses") + + updated = await add_litellm_data_to_request( + data={ + "model": "gpt-4o", + "input": "hi", + "litellm_trace_id": "root-trace-42", + "metadata": {"session_id": 4815162342}, + }, + request=request, + user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"), + proxy_config=MagicMock(), + general_settings={"missing_session_id": "generate"}, + ) + + assert updated["litellm_session_id"] == "root-trace-42" + assert updated["litellm_metadata"]["session_id"] == "root-trace-42" + assert updated["litellm_metadata"][SESSION_ID_GENERATED_METADATA_KEY] is True + + @pytest.mark.asyncio @pytest.mark.parametrize("policy", ["generate", "reject"]) async def test_missing_session_id_policy_keeps_client_supplied_session_id(policy: str): diff --git a/tests/unit/proxy/test_model_level_guardrails.py b/tests/unit/proxy/test_model_level_guardrails.py index 9eaae49c46a..0bc9c348127 100644 --- a/tests/unit/proxy/test_model_level_guardrails.py +++ b/tests/unit/proxy/test_model_level_guardrails.py @@ -15,7 +15,7 @@ from unittest.mock import AsyncMock, MagicMock, patch sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "../../.."))) from litellm.proxy.utils import ( - _check_and_merge_model_level_guardrails, + check_and_merge_model_level_guardrails, _merge_guardrails_with_existing, ) @@ -38,7 +38,7 @@ class TestCheckAndMergeModelLevelGuardrails: mock_deployment.litellm_params.get.return_value = ["openai-moderation"] mock_router.get_deployment.return_value = mock_deployment - result = _check_and_merge_model_level_guardrails(data=data, llm_router=mock_router) + result = check_and_merge_model_level_guardrails(data=data, llm_router=mock_router) assert "openai-moderation" in result["metadata"]["guardrails"] mock_router.get_deployment.assert_called_once_with(model_id="model-uuid-123") @@ -57,7 +57,7 @@ class TestCheckAndMergeModelLevelGuardrails: mock_deployment.litellm_params.get.return_value = ["model-guardrail"] mock_router.get_deployment.return_value = mock_deployment - result = _check_and_merge_model_level_guardrails(data=data, llm_router=mock_router) + result = check_and_merge_model_level_guardrails(data=data, llm_router=mock_router) assert "existing-guardrail" in result["metadata"]["guardrails"] assert "model-guardrail" in result["metadata"]["guardrails"] @@ -76,14 +76,14 @@ class TestCheckAndMergeModelLevelGuardrails: mock_deployment.litellm_params.get.return_value = ["openai-moderation"] mock_router.get_deployment.return_value = mock_deployment - result = _check_and_merge_model_level_guardrails(data=data, llm_router=mock_router) + result = check_and_merge_model_level_guardrails(data=data, llm_router=mock_router) assert result["metadata"]["guardrails"].count("openai-moderation") == 1 def test_returns_data_unchanged_when_no_router(self): """Returns data unchanged when llm_router is None.""" data = {"model": "gpt-4", "metadata": {}} - result = _check_and_merge_model_level_guardrails(data=data, llm_router=None) + result = check_and_merge_model_level_guardrails(data=data, llm_router=None) assert result is data def test_returns_data_unchanged_when_no_model_info(self): @@ -95,7 +95,7 @@ class TestCheckAndMergeModelLevelGuardrails: # finds a deployment. mock_router.get_deployment.return_value = None mock_router.get_deployment_by_model_group_name.return_value = None - result = _check_and_merge_model_level_guardrails(data=data, llm_router=mock_router) + result = check_and_merge_model_level_guardrails(data=data, llm_router=mock_router) assert result is data def test_returns_data_unchanged_when_deployment_has_no_guardrails(self): @@ -109,7 +109,7 @@ class TestCheckAndMergeModelLevelGuardrails: mock_deployment.litellm_params.get.return_value = None mock_router.get_deployment.return_value = mock_deployment - result = _check_and_merge_model_level_guardrails(data=data, llm_router=mock_router) + result = check_and_merge_model_level_guardrails(data=data, llm_router=mock_router) assert result is data @@ -122,7 +122,7 @@ class TestCheckAndMergeModelLevelGuardrails: mock_router = MagicMock() mock_router.get_deployment.return_value = None - result = _check_and_merge_model_level_guardrails(data=data, llm_router=mock_router) + result = check_and_merge_model_level_guardrails(data=data, llm_router=mock_router) assert result is data @@ -140,7 +140,7 @@ class TestCheckAndMergeModelLevelGuardrails: mock_deployment.litellm_params.get.return_value = ["new-guardrail"] mock_router.get_deployment.return_value = mock_deployment - result = _check_and_merge_model_level_guardrails(data=data, llm_router=mock_router) + result = check_and_merge_model_level_guardrails(data=data, llm_router=mock_router) # Result is a different top-level dict assert result is not data diff --git a/tests/unit/proxy/test_model_list_callback_filter.py b/tests/unit/proxy/test_model_list_callback_filter.py index 00fbfee24ed..d397ee676f2 100644 --- a/tests/unit/proxy/test_model_list_callback_filter.py +++ b/tests/unit/proxy/test_model_list_callback_filter.py @@ -107,7 +107,7 @@ def team_admin_privileges(monkeypatch) -> None: async def _is_team_admin(**kwargs) -> bool: return True - monkeypatch.setattr(common_utils, "_user_has_admin_privileges", _is_team_admin) + monkeypatch.setattr(common_utils, "user_has_admin_privileges", _is_team_admin) def _non_admin(**kwargs) -> UserAPIKeyAuth: diff --git a/tests/unit/proxy/test_model_list_discoverable.py b/tests/unit/proxy/test_model_list_discoverable.py index bcd52479f2c..0be5342c45c 100644 --- a/tests/unit/proxy/test_model_list_discoverable.py +++ b/tests/unit/proxy/test_model_list_discoverable.py @@ -71,7 +71,7 @@ def team_admin_privileges(monkeypatch) -> None: async def _is_team_admin(**kwargs) -> bool: return True - monkeypatch.setattr(common_utils, "_user_has_admin_privileges", _is_team_admin) + monkeypatch.setattr(common_utils, "user_has_admin_privileges", _is_team_admin) def _non_admin() -> UserAPIKeyAuth: diff --git a/tests/unit/proxy/test_model_list_healthy_only.py b/tests/unit/proxy/test_model_list_healthy_only.py index 718c7e41da8..2652fbb30ef 100644 --- a/tests/unit/proxy/test_model_list_healthy_only.py +++ b/tests/unit/proxy/test_model_list_healthy_only.py @@ -120,7 +120,7 @@ async def test_model_list_healthy_only_applies_to_scope_expand( async def _fake_admin(**kwargs): return True - monkeypatch.setattr(common_utils, "_user_has_admin_privileges", _fake_admin) + monkeypatch.setattr(common_utils, "user_has_admin_privileges", _fake_admin) monkeypatch.setattr( model_checks, "get_complete_model_list", @@ -158,7 +158,7 @@ async def test_model_list_general_setting_applies_to_scope_expand(patched_model_ async def _fake_admin(**kwargs): return True - monkeypatch.setattr(common_utils, "_user_has_admin_privileges", _fake_admin) + monkeypatch.setattr(common_utils, "user_has_admin_privileges", _fake_admin) monkeypatch.setattr( model_checks, "get_complete_model_list", diff --git a/tests/unit/proxy/test_modify_response_streaming_passthrough.py b/tests/unit/proxy/test_modify_response_streaming_passthrough.py index da57d9c616e..f4d3b659ee1 100644 --- a/tests/unit/proxy/test_modify_response_streaming_passthrough.py +++ b/tests/unit/proxy/test_modify_response_streaming_passthrough.py @@ -40,7 +40,7 @@ async def _run_streaming_block_and_get_wrapper(exception): outer_body = {"model": "gpt-4o", "messages": [], "stream": True} with patch( - "litellm.proxy.proxy_server._read_request_body", + "litellm.proxy.proxy_server.read_request_body", new_callable=AsyncMock, return_value=outer_body, ), patch( diff --git a/tests/unit/proxy/test_moyai_endpoints.py b/tests/unit/proxy/test_moyai_endpoints.py new file mode 100644 index 00000000000..92a1261d952 --- /dev/null +++ b/tests/unit/proxy/test_moyai_endpoints.py @@ -0,0 +1,324 @@ +import json +import time +from types import SimpleNamespace +from typing import Final +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from litellm._internal_context import current_service_target +from litellm.caching.dual_cache import DualCache +from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth + + +def _admin() -> UserAPIKeyAuth: + return UserAPIKeyAuth(user_id="admin-user", user_role=LitellmUserRoles.PROXY_ADMIN) + + +def _request() -> MagicMock: + request = MagicMock() + request.base_url = "http://localhost:4000/" + return request + + +@pytest.mark.asyncio +async def test_start_rejects_non_admin(monkeypatch: pytest.MonkeyPatch) -> None: + from fastapi import HTTPException + from litellm.proxy import proxy_server + from litellm.proxy.moyai_endpoints import MoyaiConnectStartRequest, moyai_connect_start + + monkeypatch.setattr(proxy_server, "master_key", "sk-master") + actor = UserAPIKeyAuth(user_id="member", user_role="internal_user") + + with pytest.raises(HTTPException) as exc: + await moyai_connect_start( + _request(), + MoyaiConnectStartRequest(moyai_url="https://moyai.example.com", return_to="http://localhost:3000/ui/moyai"), + actor, + ) + assert exc.value.status_code == 403 + + +@pytest.mark.asyncio +async def test_start_requires_master_key(monkeypatch: pytest.MonkeyPatch) -> None: + from fastapi import HTTPException + from litellm.proxy import proxy_server + from litellm.proxy.moyai_endpoints import MoyaiConnectStartRequest, moyai_connect_start + + monkeypatch.setattr(proxy_server, "master_key", None) + + with pytest.raises(HTTPException) as exc: + await moyai_connect_start( + _request(), + MoyaiConnectStartRequest(moyai_url="https://moyai.example.com", return_to="http://localhost:3000/ui/moyai"), + _admin(), + ) + assert exc.value.status_code == 400 + assert "LITELLM_MASTER_KEY" in exc.value.detail + + +@pytest.mark.asyncio +async def test_start_returns_connect_url_with_all_params(monkeypatch: pytest.MonkeyPatch) -> None: + from litellm.proxy import proxy_server + from litellm.proxy.moyai_endpoints import MoyaiConnectStartRequest, moyai_connect_start + from urllib.parse import parse_qs, urlparse + + monkeypatch.setattr(proxy_server, "master_key", "sk-master") + + response = await moyai_connect_start( + _request(), + MoyaiConnectStartRequest(moyai_url="https://moyai.example.com/", return_to="http://localhost:3000/ui/moyai"), + _admin(), + ) + + parsed = urlparse(response.connect_url) + assert f"{parsed.scheme}://{parsed.netloc}{parsed.path}" == "https://moyai.example.com/connect/litellm" + params = parse_qs(parsed.query) + assert params["gateway_url"] == ["http://localhost:4000"] + assert params["return_to"] == ["http://localhost:3000/ui/moyai"] + assert params["code"] and "." in params["code"][0] + + +async def _exchange_env(monkeypatch: pytest.MonkeyPatch): + from litellm.proxy import proxy_server + + cache: dict = {} + + async def _get(key: str): + return cache.get(key) + + async def _set(key: str, value, ttl=None): + cache[key] = value + + monkeypatch.setattr(proxy_server, "master_key", "sk-master") + monkeypatch.setattr(proxy_server, "llm_router", None) + user_api_key_cache = SimpleNamespace( + async_get_cache=AsyncMock(side_effect=_get), async_set_cache=AsyncMock(side_effect=_set) + ) + monkeypatch.setattr(proxy_server, "user_api_key_cache", user_api_key_cache) + + prisma = MagicMock() + prisma.db.litellm_uisettings.find_unique = AsyncMock(return_value=SimpleNamespace(ui_settings={})) + prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) + + claimed_nonces: set = set() + + async def _config_create(*, data): + from prisma.errors import UniqueViolationError + + param_name = data["param_name"] + if param_name in claimed_nonces: + raise UniqueViolationError({}, message="Unique constraint failed on the fields: (`param_name`)") + claimed_nonces.add(param_name) + + nonce_create = AsyncMock(side_effect=_config_create) + config_table = SimpleNamespace(create=nonce_create) + prisma.db.litellm_config = config_table + prisma.writer_db.litellm_config = config_table + + persisted: dict = {} + + async def _upsert(where, data): + persisted.update(json.loads(data["update"]["ui_settings"])) + + prisma.db.litellm_uisettings.upsert = AsyncMock(side_effect=_upsert) + monkeypatch.setattr(proxy_server, "prisma_client", prisma) + + mint_calls: list = [] + + async def _mint(request_type, **kwargs): + mint_calls.append(kwargs) + return {"token": "sk-new-virtual-key"} + + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints.generate_key_helper_fn", + AsyncMock(side_effect=_mint), + ) + return persisted, mint_calls, nonce_create + + +async def _exchange_call(code: str, moyai_url: str): + from litellm.proxy.moyai_endpoints import MoyaiConnectExchangeRequest, moyai_connect_exchange + + return await moyai_connect_exchange(_request(), MoyaiConnectExchangeRequest(code=code, moyai_url=moyai_url)) + + +@pytest.mark.asyncio +async def test_exchange_happy_path_mints_key_and_saves_setting(monkeypatch: pytest.MonkeyPatch) -> None: + from litellm.proxy.moyai_endpoints import _sign_connect_code + + persisted, mint_calls, _ = await _exchange_env(monkeypatch) + code = _sign_connect_code("sk-master", "https://moyai.example.com", "admin-user") + + response = await _exchange_call(code, "https://moyai.example.com") + + assert response.api_key == "sk-new-virtual-key" + assert response.key_alias == "moyai-moyai.example.com" + assert response.api_base == "http://localhost:4000" + assert persisted["moyai_url"] == "https://moyai.example.com" + mint = mint_calls[0] + assert "user_id" not in mint + assert mint["allowed_routes"] == ["openai_routes", "anthropic_routes", "/model/info"] + assert mint["metadata"] == { + "created_via": "moyai_quick_connect", + "moyai_url": "https://moyai.example.com", + "connected_by": "admin-user", + } + + +@pytest.mark.asyncio +async def test_exchange_without_database_fails_before_nonce(monkeypatch: pytest.MonkeyPatch) -> None: + from fastapi import HTTPException + from litellm.proxy import proxy_server + from litellm.proxy.moyai_endpoints import _sign_connect_code + + _, _, nonce_create = await _exchange_env(monkeypatch) + monkeypatch.setattr(proxy_server, "prisma_client", None) + code = _sign_connect_code("sk-master", "https://moyai.example.com", "admin-user") + + with pytest.raises(HTTPException) as exc: + await _exchange_call(code, "https://moyai.example.com") + assert exc.value.status_code == 400 + assert "database" in exc.value.detail + nonce_create.assert_not_called() + + +@pytest.mark.asyncio +async def test_exchange_rejects_tampered_signature(monkeypatch: pytest.MonkeyPatch) -> None: + from fastapi import HTTPException + from litellm.proxy.moyai_endpoints import _sign_connect_code + + await _exchange_env(monkeypatch) + code = _sign_connect_code("sk-master", "https://moyai.example.com", "admin-user") + tampered = code[:-2] + ("aa" if not code.endswith("aa") else "bb") + + with pytest.raises(HTTPException) as exc: + await _exchange_call(tampered, "https://moyai.example.com") + assert exc.value.status_code == 400 + + +@pytest.mark.asyncio +async def test_exchange_rejects_expired_code(monkeypatch: pytest.MonkeyPatch) -> None: + import base64 + import hashlib + import hmac as hmac_mod + from fastapi import HTTPException + from litellm.proxy.moyai_endpoints import _b64url, _master_key_hmac_key + + await _exchange_env(monkeypatch) + payload = json.dumps( + { + "moyai_origin": "https://moyai.example.com", + "user_id": "admin-user", + "exp": int(time.time()) - 10, + "nonce": "n", + }, + separators=(",", ":"), + sort_keys=True, + ).encode() + sig = hmac_mod.new(_master_key_hmac_key("sk-master"), payload, hashlib.sha256).digest() + code = f"{_b64url(payload)}.{_b64url(sig)}" + + with pytest.raises(HTTPException) as exc: + await _exchange_call(code, "https://moyai.example.com") + assert exc.value.status_code == 400 + + +@pytest.mark.asyncio +async def test_exchange_rejects_origin_mismatch(monkeypatch: pytest.MonkeyPatch) -> None: + from fastapi import HTTPException + from litellm.proxy.moyai_endpoints import _sign_connect_code + + await _exchange_env(monkeypatch) + code = _sign_connect_code("sk-master", "https://moyai.example.com", "admin-user") + + with pytest.raises(HTTPException) as exc: + await _exchange_call(code, "https://evil.example.com") + assert exc.value.status_code == 400 + + +@pytest.mark.asyncio +async def test_exchange_rejects_replayed_nonce(monkeypatch: pytest.MonkeyPatch) -> None: + from fastapi import HTTPException + from litellm.proxy.moyai_endpoints import _sign_connect_code + + _, mint_calls, _ = await _exchange_env(monkeypatch) + code = _sign_connect_code("sk-master", "https://moyai.example.com", "admin-user") + + await _exchange_call(code, "https://moyai.example.com") + with pytest.raises(HTTPException) as exc: + await _exchange_call(code, "https://moyai.example.com") + assert exc.value.status_code == 400 + assert len(mint_calls) == 1 + + +@pytest.mark.asyncio +async def test_exchange_concurrent_replay_claims_nonce_once(monkeypatch: pytest.MonkeyPatch) -> None: + import asyncio + + from fastapi import HTTPException + from litellm.proxy.moyai_endpoints import _sign_connect_code + + _, mint_calls, _ = await _exchange_env(monkeypatch) + code = _sign_connect_code("sk-master", "https://moyai.example.com", "admin-user") + + results = await asyncio.gather( + _exchange_call(code, "https://moyai.example.com"), + _exchange_call(code, "https://moyai.example.com"), + return_exceptions=True, + ) + + successes = [r for r in results if not isinstance(r, BaseException)] + rejections = [r for r in results if isinstance(r, HTTPException) and r.status_code == 400] + assert len(successes) == 1 + assert len(rejections) == 1 + assert len(mint_calls) == 1 + + +def _target_recording_cache(cache: DualCache) -> SimpleNamespace: + async def _set(key: str, value: object, **kwargs: object) -> None: + await cache.async_set_cache(key=f"{key}:service_target", value=current_service_target()) + await cache.async_set_cache(key=key, value=value, **kwargs) + + return SimpleNamespace(async_set_cache=_set) + + +@pytest.mark.asyncio +async def test_persist_moyai_url_writes_ui_settings_cache_under_config_params_target( + monkeypatch: pytest.MonkeyPatch, +) -> None: + from litellm.proxy import proxy_server + from litellm.proxy.moyai_endpoints import _persist_moyai_url + from litellm.proxy.ui_crud_endpoints.proxy_setting_endpoints import UI_SETTINGS_CACHE_KEY + from litellm.proxy.utils import CONFIG_PARAMS_TARGET + + cache: Final = DualCache() + monkeypatch.setattr(proxy_server, "user_api_key_cache", _target_recording_cache(cache)) + + prisma: Final = MagicMock() + prisma.db.litellm_uisettings.find_unique = AsyncMock(return_value=None) + prisma.db.litellm_uisettings.upsert = AsyncMock() + + await _persist_moyai_url(prisma, "https://moyai.example.com") + + assert await cache.async_get_cache(key=f"{UI_SETTINGS_CACHE_KEY}:service_target") == CONFIG_PARAMS_TARGET + assert await cache.async_get_cache(key=UI_SETTINGS_CACHE_KEY) == {"moyai_url": "https://moyai.example.com"} + assert current_service_target() is None + + +@pytest.mark.asyncio +async def test_exchange_replay_survives_fresh_worker_cache(monkeypatch: pytest.MonkeyPatch) -> None: + from fastapi import HTTPException + from litellm.caching.caching import DualCache + from litellm.proxy import proxy_server + from litellm.proxy.moyai_endpoints import _sign_connect_code + + _, mint_calls, _ = await _exchange_env(monkeypatch) + code = _sign_connect_code("sk-master", "https://moyai.example.com", "admin-user") + + await _exchange_call(code, "https://moyai.example.com") + monkeypatch.setattr(proxy_server, "user_api_key_cache", DualCache()) + with pytest.raises(HTTPException) as exc: + await _exchange_call(code, "https://moyai.example.com") + assert exc.value.status_code == 400 + assert len(mint_calls) == 1 diff --git a/tests/unit/proxy/test_proxy_cli.py b/tests/unit/proxy/test_proxy_cli.py index 8fc807cd8a0..c73b60f3c26 100644 --- a/tests/unit/proxy/test_proxy_cli.py +++ b/tests/unit/proxy/test_proxy_cli.py @@ -607,7 +607,7 @@ class TestProxyInitializationHelpers: }, ), patch( - "litellm.proxy.proxy_cli.ProxyInitializationHelpers._get_default_unvicorn_init_args" + "litellm.proxy.proxy_cli.ProxyInitializationHelpers.get_default_unvicorn_init_args" ) as mock_get_args, ): mock_get_args.return_value = { @@ -713,7 +713,7 @@ class TestProxyInitializationHelpers: }, ), patch( - "litellm.proxy.proxy_cli.ProxyInitializationHelpers._get_default_unvicorn_init_args" + "litellm.proxy.proxy_cli.ProxyInitializationHelpers.get_default_unvicorn_init_args" ) as mock_get_args, ): mock_get_args.side_effect = lambda *a, **k: { @@ -806,7 +806,7 @@ class TestProxyInitializationHelpers: }, ), patch( # test-quality-ok: same isolation as the sibling CLI tests above - "litellm.proxy.proxy_cli.ProxyInitializationHelpers._get_default_unvicorn_init_args" + "litellm.proxy.proxy_cli.ProxyInitializationHelpers.get_default_unvicorn_init_args" ) as mock_get_args, ): mock_get_args.side_effect = lambda *a, **k: { @@ -935,7 +935,7 @@ class TestProxyInitializationHelpers: }, ), patch( - "litellm.proxy.proxy_cli.ProxyInitializationHelpers._get_default_unvicorn_init_args" + "litellm.proxy.proxy_cli.ProxyInitializationHelpers.get_default_unvicorn_init_args" ) as mock_get_args, patch( "litellm.proxy.proxy_cli.append_query_params", @@ -1065,7 +1065,7 @@ class TestProxyInitializationHelpers: }, ), patch( - "litellm.proxy.proxy_cli.ProxyInitializationHelpers._get_default_unvicorn_init_args" + "litellm.proxy.proxy_cli.ProxyInitializationHelpers.get_default_unvicorn_init_args" ) as mock_get_args, patch( "litellm.proxy.proxy_cli.append_query_params", @@ -1186,7 +1186,7 @@ class TestProxyInitializationHelpers: }, ), patch( - "litellm.proxy.proxy_cli.ProxyInitializationHelpers._get_default_unvicorn_init_args" + "litellm.proxy.proxy_cli.ProxyInitializationHelpers.get_default_unvicorn_init_args" ) as mock_get_args, patch( "litellm.proxy.proxy_cli.append_query_params", @@ -1279,7 +1279,7 @@ class TestProxyInitializationHelpers: }, ), patch( - "litellm.proxy.proxy_cli.ProxyInitializationHelpers._get_default_unvicorn_init_args" + "litellm.proxy.proxy_cli.ProxyInitializationHelpers.get_default_unvicorn_init_args" ) as mock_get_args, patch( "litellm.proxy.proxy_cli.append_query_params", @@ -1358,7 +1358,7 @@ class TestProxyInitializationHelpers: }, ), patch( - "litellm.proxy.proxy_cli.ProxyInitializationHelpers._get_default_unvicorn_init_args" + "litellm.proxy.proxy_cli.ProxyInitializationHelpers.get_default_unvicorn_init_args" ) as mock_get_args, patch( "litellm.proxy.proxy_cli.append_query_params", @@ -1417,7 +1417,7 @@ class TestProxyInitializationHelpers: }, ), patch( - "litellm.proxy.proxy_cli.ProxyInitializationHelpers._get_default_unvicorn_init_args" + "litellm.proxy.proxy_cli.ProxyInitializationHelpers.get_default_unvicorn_init_args" ) as mock_get_args, ): mock_get_args.return_value = { @@ -1478,10 +1478,10 @@ class TestProxyInitializationHelpers: }, ), patch( - "litellm.proxy.proxy_cli.ProxyInitializationHelpers._get_default_unvicorn_init_args" + "litellm.proxy.proxy_cli.ProxyInitializationHelpers.get_default_unvicorn_init_args" ) as mock_get_args, patch( - "litellm.proxy.proxy_cli.ProxyInitializationHelpers._is_port_in_use", + "litellm.proxy.proxy_cli.ProxyInitializationHelpers.is_port_in_use", return_value=False, ), ): @@ -1553,10 +1553,10 @@ class TestProxyInitializationHelpers: }, ), patch( - "litellm.proxy.proxy_cli.ProxyInitializationHelpers._get_default_unvicorn_init_args" + "litellm.proxy.proxy_cli.ProxyInitializationHelpers.get_default_unvicorn_init_args" ) as mock_get_args, patch( - "litellm.proxy.proxy_cli.ProxyInitializationHelpers._is_port_in_use", + "litellm.proxy.proxy_cli.ProxyInitializationHelpers.is_port_in_use", return_value=False, ), ): @@ -1620,7 +1620,7 @@ class TestProxyInitializationHelpers: }, ), patch( - "litellm.proxy.proxy_cli.ProxyInitializationHelpers._get_default_unvicorn_init_args" + "litellm.proxy.proxy_cli.ProxyInitializationHelpers.get_default_unvicorn_init_args" ) as mock_get_args, ): mock_get_args.return_value = { @@ -1678,7 +1678,7 @@ class TestProxyInitializationHelpers: }, ), patch( - "litellm.proxy.proxy_cli.ProxyInitializationHelpers._get_default_unvicorn_init_args" + "litellm.proxy.proxy_cli.ProxyInitializationHelpers.get_default_unvicorn_init_args" ) as mock_get_args, ): mock_get_args.return_value = { @@ -1706,7 +1706,7 @@ class TestProxyInitializationHelpers: assert call_args[1]["limit_max_requests"] == 1000 assert call_args[1]["limit_max_requests_jitter"] == 50 - @patch("litellm.proxy.proxy_cli.ProxyInitializationHelpers._run_gunicorn_server") + @patch("litellm.proxy.proxy_cli.ProxyInitializationHelpers.run_gunicorn_server") @patch("uvicorn.run") @patch("builtins.print") @patch("litellm.proxy.db.prisma_client.PrismaManager.setup_database") @@ -1738,7 +1738,7 @@ class TestProxyInitializationHelpers: }, ), patch( - "litellm.proxy.proxy_cli.ProxyInitializationHelpers._get_default_unvicorn_init_args" + "litellm.proxy.proxy_cli.ProxyInitializationHelpers.get_default_unvicorn_init_args" ) as mock_get_args, ): mock_get_args.return_value = { @@ -2015,7 +2015,7 @@ class TestProxyInitializationHelpers: {"litellm.proxy.proxy_server": mock_proxy_server_module}, ), patch( - "litellm.proxy.proxy_cli.ProxyInitializationHelpers._get_default_unvicorn_init_args" + "litellm.proxy.proxy_cli.ProxyInitializationHelpers.get_default_unvicorn_init_args" ) as mock_get_args, ): mock_get_args.return_value = { @@ -2080,7 +2080,7 @@ class TestQueryEngineReaperWiring: "litellm.proxy.proxy_cli.start_query_engine_reaper" ) as mock_start_reaper, patch( - "litellm.proxy.proxy_cli.ProxyInitializationHelpers._get_default_unvicorn_init_args" + "litellm.proxy.proxy_cli.ProxyInitializationHelpers.get_default_unvicorn_init_args" ) as mock_get_args, ): mock_get_args.return_value = { @@ -2180,7 +2180,7 @@ class TestRunServerDbSetup: }, ), patch( - "litellm.proxy.proxy_cli.ProxyInitializationHelpers._get_default_unvicorn_init_args" + "litellm.proxy.proxy_cli.ProxyInitializationHelpers.get_default_unvicorn_init_args" ) as mock_get_args, ): mock_get_args.return_value = { @@ -2257,7 +2257,7 @@ class TestRunServerDbSetup: }, ), patch( # test-quality-ok: same isolation as the sibling CLI tests above - "litellm.proxy.proxy_cli.ProxyInitializationHelpers._get_default_unvicorn_init_args" + "litellm.proxy.proxy_cli.ProxyInitializationHelpers.get_default_unvicorn_init_args" ) as mock_get_args, ): mock_get_args.return_value = { @@ -2374,7 +2374,7 @@ class TestRunServerDbSetup: }, ), patch( - "litellm.proxy.proxy_cli.ProxyInitializationHelpers._get_default_unvicorn_init_args" + "litellm.proxy.proxy_cli.ProxyInitializationHelpers.get_default_unvicorn_init_args" ) as mock_get_args, ): mock_get_args.return_value = { @@ -2743,7 +2743,7 @@ class TestRunServerDbSetup: {"proxy_server": mock_proxy_module, "litellm.proxy.proxy_server": mock_proxy_module}, ), patch( - "litellm.proxy.proxy_cli.ProxyInitializationHelpers._get_default_unvicorn_init_args" + "litellm.proxy.proxy_cli.ProxyInitializationHelpers.get_default_unvicorn_init_args" ) as mock_get_args, outcome as exc_info, ): @@ -3435,7 +3435,7 @@ class TestTokenAuthCliFlags: patch("litellm.proxy.db.prisma_client.PrismaManager.setup_database"), patch("uvicorn.run"), patch( - "litellm.proxy.proxy_cli.ProxyInitializationHelpers._get_default_unvicorn_init_args" + "litellm.proxy.proxy_cli.ProxyInitializationHelpers.get_default_unvicorn_init_args" ) as mock_get_args, ): mock_get_args.return_value = { diff --git a/tests/unit/proxy/test_proxy_server.py b/tests/unit/proxy/test_proxy_server.py index 9772988eac8..79f4edd4d62 100644 --- a/tests/unit/proxy/test_proxy_server.py +++ b/tests/unit/proxy/test_proxy_server.py @@ -1238,7 +1238,7 @@ async def test_team_update_redis(): """ from litellm.caching.caching import DualCache, RedisCache from litellm.proxy._types import LiteLLM_TeamTableCachedObj - from litellm.proxy.auth.auth_checks import _cache_team_object + from litellm.proxy.auth.auth_checks import cache_team_object proxy_logging_obj: ProxyLogging = getattr( litellm.proxy.proxy_server, "proxy_logging_obj" @@ -1251,7 +1251,7 @@ async def test_team_update_redis(): "async_set_cache", new=AsyncMock(), ) as mock_client: - await _cache_team_object( + await cache_team_object( team_id="1234", team_table=LiteLLM_TeamTableCachedObj(team_id="1234"), user_api_key_cache=DualCache(redis_cache=redis_cache), diff --git a/tests/unit/proxy/test_proxy_server_endpoints_and_startup.py b/tests/unit/proxy/test_proxy_server_endpoints_and_startup.py index ec91f1d8edc..920f9d90bb8 100644 --- a/tests/unit/proxy/test_proxy_server_endpoints_and_startup.py +++ b/tests/unit/proxy/test_proxy_server_endpoints_and_startup.py @@ -1652,7 +1652,7 @@ mock_prisma = MockPrisma() @patch( - "litellm.proxy.proxy_server.ProxyStartupEvent._setup_prisma_client", + "litellm.proxy.proxy_server.ProxyStartupEvent.setup_prisma_client", return_value=mock_prisma, ) @pytest.mark.asyncio @@ -3623,12 +3623,12 @@ async def test_startup_initializes_string_callbacks_after_all_litellm_settings_l def test_startup_hands_router_to_every_registered_prompt_injection_detector(monkeypatch): from litellm.proxy._types import LiteLLMPromptInjectionParams - from litellm.proxy.hooks.prompt_injection_detection import _OPTIONAL_PromptInjectionDetection + from litellm.proxy.hooks.prompt_injection_detection import OPTIONAL_PromptInjectionDetection from litellm.proxy.proxy_server import ProxyStartupEvent from litellm.router import Router monkeypatch.setattr(litellm, "callbacks", []) - detector = _OPTIONAL_PromptInjectionDetection( + detector = OPTIONAL_PromptInjectionDetection( prompt_injection_params=LiteLLMPromptInjectionParams( heuristics_check=False, llm_api_check=True, @@ -4545,7 +4545,7 @@ async def test_chat_completion_result_no_nested_none_values(): with ( patch( - "litellm.proxy.proxy_server._read_request_body", + "litellm.proxy.proxy_server.read_request_body", return_value={"model": "gpt-3.5-turbo", "messages": []}, ), patch( @@ -6554,7 +6554,7 @@ async def test_init_sso_settings_in_db(): mock_prisma_client.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=mock_sso_config) # Mock _decrypt_and_set_db_env_variables - with patch.object(proxy_config, "_decrypt_and_set_db_env_variables") as mock_decrypt_and_set: + with patch.object(proxy_config, "decrypt_and_set_db_env_variables") as mock_decrypt_and_set: await proxy_config._init_sso_settings_in_db(prisma_client=mock_prisma_client) # Verify find_unique was called with correct parameters @@ -6598,7 +6598,7 @@ async def test_init_sso_settings_in_db_no_settings(): mock_prisma_client.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=None) # Mock _decrypt_and_set_db_env_variables - with patch.object(proxy_config, "_decrypt_and_set_db_env_variables") as mock_decrypt_and_set: + with patch.object(proxy_config, "decrypt_and_set_db_env_variables") as mock_decrypt_and_set: await proxy_config._init_sso_settings_in_db(prisma_client=mock_prisma_client) # Verify find_unique was called @@ -6652,7 +6652,7 @@ async def test_init_sso_settings_in_db_empty_settings(): mock_prisma_client.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=mock_sso_config) # Mock _decrypt_and_set_db_env_variables - with patch.object(proxy_config, "_decrypt_and_set_db_env_variables") as mock_decrypt_and_set: + with patch.object(proxy_config, "decrypt_and_set_db_env_variables") as mock_decrypt_and_set: await proxy_config._init_sso_settings_in_db(prisma_client=mock_prisma_client) # Verify find_unique was called @@ -6694,7 +6694,7 @@ async def test_init_sso_settings_in_db_retries_on_transport_error(): mock_prisma_client._db_auth_reconnect_timeout_seconds = 2.0 mock_prisma_client._db_auth_reconnect_lock_timeout_seconds = 0.1 - with patch.object(proxy_config, "_decrypt_and_set_db_env_variables") as mock_decrypt: + with patch.object(proxy_config, "decrypt_and_set_db_env_variables") as mock_decrypt: await proxy_config._init_sso_settings_in_db(prisma_client=mock_prisma_client) assert len(invocations) == 2 @@ -6872,7 +6872,7 @@ def test_root_redirect_when_docs_url_not_root_and_redirect_url_set(monkeypatch): from fastapi.responses import RedirectResponse from litellm.proxy.proxy_server import cleanup_router_config_variables - from litellm.proxy.utils import _get_docs_url + from litellm.proxy.utils import get_docs_url cleanup_router_config_variables() filepath = os.path.dirname(os.path.abspath(__file__)) @@ -6885,7 +6885,7 @@ def test_root_redirect_when_docs_url_not_root_and_redirect_url_set(monkeypatch): asyncio.run(initialize(config=config_fp, debug=True)) - docs_url = _get_docs_url() + docs_url = get_docs_url() root_redirect_url = os.getenv("ROOT_REDIRECT_URL") # Remove any existing "/" route that might interfere @@ -7368,7 +7368,7 @@ class TestInvitationEndpoints: mock_prisma.db.litellm_invitationlink = MagicMock() # Avoid triggering async DB calls in _user_has_admin_privileges with patch( - "litellm.proxy.proxy_server._user_has_admin_privileges", + "litellm.proxy.proxy_server.user_has_admin_privileges", new_callable=AsyncMock, return_value=False, ): @@ -7537,7 +7537,7 @@ async def test_async_data_generator_uses_direct_stream_fast_path_without_callbac mock_proxy_logging_obj.post_call_failure_hook = AsyncMock() with patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj): - with patch.object(ProxyLogging, "_fire_deferred_stream_logging") as mock_deferred_logging: + with patch.object(ProxyLogging, "fire_deferred_stream_logging") as mock_deferred_logging: yielded_data = [] async for data in async_data_generator(mock_response, mock_user_api_key_dict, mock_request_data): yielded_data.append(data) @@ -7593,7 +7593,7 @@ async def test_async_data_generator_preserves_non_raw_sse_like_bytes(): mock_proxy_logging_obj.post_call_failure_hook = AsyncMock() with patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj): - with patch.object(ProxyLogging, "_fire_deferred_stream_logging"): + with patch.object(ProxyLogging, "fire_deferred_stream_logging"): yielded_data = [] async for data in async_data_generator(mock_response, mock_user_api_key_dict, mock_request_data): yielded_data.append(data) @@ -7650,7 +7650,7 @@ async def test_async_data_generator_buffers_split_google_native_sse_json_frame() mock_proxy_logging_obj.post_call_failure_hook = AsyncMock() with patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj): - with patch.object(ProxyLogging, "_fire_deferred_stream_logging"): + with patch.object(ProxyLogging, "fire_deferred_stream_logging"): yielded_data = [] async for data in async_data_generator(mock_response, mock_user_api_key_dict, mock_request_data): yielded_data.append(data) @@ -7698,7 +7698,7 @@ async def test_async_data_generator_flushes_raw_sse_stream_without_trailing_deli with ( patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj), - patch.object(ProxyLogging, "_fire_deferred_stream_logging"), + patch.object(ProxyLogging, "fire_deferred_stream_logging"), ): yielded_data = [] async for data in async_data_generator(mock_response, mock_user_api_key_dict, mock_request_data): @@ -7748,7 +7748,7 @@ async def test_async_data_generator_errors_when_raw_sse_frame_exceeds_buffer_lim with ( patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj), patch("litellm.proxy.proxy_server._MAX_RAW_SSE_BUFFER_CHARS", 8), - patch.object(ProxyLogging, "_fire_deferred_stream_logging"), + patch.object(ProxyLogging, "fire_deferred_stream_logging"), ): yielded_data = [] async for data in async_data_generator(mock_response, mock_user_api_key_dict, mock_request_data): @@ -7804,7 +7804,7 @@ async def test_async_data_generator_checks_raw_sse_buffer_limit_after_complete_f with ( patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj), patch("litellm.proxy.proxy_server._MAX_RAW_SSE_BUFFER_CHARS", 8), - patch.object(ProxyLogging, "_fire_deferred_stream_logging"), + patch.object(ProxyLogging, "fire_deferred_stream_logging"), ): yielded_data = [] async for data in async_data_generator(mock_response, mock_user_api_key_dict, mock_request_data): @@ -7854,7 +7854,7 @@ async def test_async_data_generator_google_genai_stream_omits_openai_done(): mock_proxy_logging_obj.post_call_failure_hook = AsyncMock() with patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj): - with patch.object(ProxyLogging, "_fire_deferred_stream_logging"): + with patch.object(ProxyLogging, "fire_deferred_stream_logging"): yielded_data = [] async for data in async_data_generator(mock_response, mock_user_api_key_dict, mock_request_data): yielded_data.append(data) @@ -7896,7 +7896,7 @@ async def test_async_data_generator_does_not_mark_completed_stream_as_disconnect mock_proxy_logging_obj.post_call_failure_hook = AsyncMock() with patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj): - with patch.object(ProxyLogging, "_fire_deferred_stream_logging"): + with patch.object(ProxyLogging, "fire_deferred_stream_logging"): yielded_data = [] async for data in async_data_generator( mock_response, @@ -7946,7 +7946,7 @@ async def test_async_data_generator_google_genai_stream_forwards_error_without_d mock_proxy_logging_obj.post_call_failure_hook = AsyncMock() with patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj): - with patch.object(ProxyLogging, "_fire_deferred_stream_logging"): + with patch.object(ProxyLogging, "fire_deferred_stream_logging"): yielded_data = [] async for data in async_data_generator(mock_response, mock_user_api_key_dict, mock_request_data): yielded_data.append(data) @@ -9384,7 +9384,7 @@ async def test_window_spend_counter_skips_invalid_window_start(): @pytest.mark.asyncio async def test_window_spend_counter_does_not_seed_zero_when_db_unavailable(): from litellm.caching.dual_cache import DualCache - from litellm.proxy.proxy_server import _ensure_window_spend_counter_initialized + from litellm.proxy.proxy_server import ensure_window_spend_counter_initialized counter_cache = DualCache() counter_key = "spend:key:key-window-db-unavailable:window:1h" @@ -9395,7 +9395,7 @@ async def test_window_spend_counter_does_not_seed_zero_when_db_unavailable(): ps.spend_counter_cache = counter_cache ps.prisma_client = None try: - initialized = await _ensure_window_spend_counter_initialized( + initialized = await ensure_window_spend_counter_initialized( counter_key=counter_key, entity_type="Key", entity_id="key-window-db-unavailable", @@ -9563,7 +9563,7 @@ async def test_increment_spend_counters_reseeds_from_db_on_bad_reserved_counter( @pytest.mark.asyncio async def test_increment_spend_counter_invalidates_stale_cache_on_redis_failure(): from litellm.caching.dual_cache import DualCache - from litellm.proxy.proxy_server import _increment_spend_counter_cache + from litellm.proxy.proxy_server import increment_spend_counter_cache counter_cache = DualCache() counter_cache.in_memory_cache.set_cache(key="spend:team:redis-fail", value=4.0) @@ -9578,7 +9578,7 @@ async def test_increment_spend_counter_invalidates_stale_cache_on_redis_failure( ps.spend_counter_cache = counter_cache try: with pytest.raises(RuntimeError): - await _increment_spend_counter_cache( + await increment_spend_counter_cache( counter_key="spend:team:redis-fail", increment=0.5, ) @@ -11161,14 +11161,14 @@ async def _lit6463_drive_realtime_session_holding_a_max_parallel_slot( endpoint did to the slot.""" from litellm.proxy import proxy_server as ps from litellm.proxy.hooks.parallel_request_limiter_v3 import ( - _PROXY_MaxParallelRequestsHandler_v3, + PROXY_MaxParallelRequestsHandler_v3, _request_stash, ) from litellm.proxy.utils import InternalUsageCache dual_cache: Final = DualCache() await dual_cache.async_set_cache(key=_LIT6463_COUNTER_KEY, value={"slot-1": 1.0, "slot-2": 2.0}, local_only=True) - limiter: Final = _PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=InternalUsageCache(dual_cache)) + limiter: Final = PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=InternalUsageCache(dual_cache)) stash: Final = RequestRateLimiterStash(parallel_slot={"slot_id": "slot-1", "counter_keys": [_LIT6463_COUNTER_KEY]}) reservation: Final = {"reserved_cost": 0.55, "input_cost": 0.0, "finalized": False, "entries": []} @@ -11271,7 +11271,7 @@ async def test_release_or_invalidate_falls_back_to_invalidating_the_counters(): br, "release_budget_reservation", new=AsyncMock(side_effect=RuntimeError("counter store down")) ) # test-quality-ok: forces the failure branch; assertion observes which counter key got invalidated sink = patch.object( - ps, "_invalidate_spend_counter", new=_record + ps, "invalidate_spend_counter", new=_record ) # test-quality-ok: fakes the counter-store sink so the invalidated key is observable with failing_release, sink: await br.release_or_invalidate_budget_reservation(budget_reservation=reservation) @@ -11518,7 +11518,7 @@ class TestDeleteDeploymentSync: mock_prisma = MagicMock() mock_prisma.db.litellm_proxymodeltable.find_many = AsyncMock(side_effect=Exception("DB connection lost")) - result = await proxy_config._get_models_from_db(prisma_client=mock_prisma) + result = await proxy_config.get_models_from_db(prisma_client=mock_prisma) assert result is None, f"Expected None on DB failure to signal fetch error, got {result!r}" @@ -11548,7 +11548,7 @@ class TestDeleteDeploymentSync: reader=PrismaWrapper(original_prisma=reader_inner, iam_token_db_auth=False), ) - result = await ProxyConfig()._get_models_from_db(prisma_client=mock_prisma) + result = await ProxyConfig().get_models_from_db(prisma_client=mock_prisma) assert result == [committed_row], f"Expected the writer's just-committed row, got {result!r}" reader_inner.litellm_proxymodeltable.find_many.assert_not_awaited() @@ -11587,7 +11587,7 @@ class TestDeleteDeploymentSync: ) mock_prisma.db._writer_unavailable = True - result = await ProxyConfig()._get_models_from_db(prisma_client=mock_prisma) + result = await ProxyConfig().get_models_from_db(prisma_client=mock_prisma) assert result == [replica_row], f"Expected the replica's rows in degraded mode, got {result!r}" writer_inner.litellm_proxymodeltable.find_many.assert_not_awaited() @@ -13755,7 +13755,7 @@ async def _collect_async_data_generator_frames(request_data: dict) -> list: mock_proxy_logging_obj.post_call_failure_hook = AsyncMock() with patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj): - with patch.object(proxy_server_module.ProxyLogging, "_fire_deferred_stream_logging"): + with patch.object(proxy_server_module.ProxyLogging, "fire_deferred_stream_logging"): return [ frame.decode("utf-8") if isinstance(frame, bytes) else frame async for frame in async_data_generator(MockStream(), MagicMock(spec=UserAPIKeyAuth), request_data) @@ -13842,7 +13842,7 @@ def test_startup_warns_when_mock_testing_params_enabled(caplog): ) with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"): - ProxyStartupEvent._warn_if_mock_testing_params_enabled(general_settings={MOCK_TESTING_CONFIG_KEY: True}) + ProxyStartupEvent.warn_if_mock_testing_params_enabled(general_settings={MOCK_TESTING_CONFIG_KEY: True}) assert MOCK_TESTING_CONFIG_KEY in caplog.text for param_name in GATED_MOCK_PARAM_NAMES: @@ -13857,7 +13857,7 @@ def test_startup_is_silent_when_mock_testing_params_disabled(caplog): from litellm.proxy.route_llm_request import MOCK_TESTING_CONFIG_KEY with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"): - ProxyStartupEvent._warn_if_mock_testing_params_enabled(general_settings={}) + ProxyStartupEvent.warn_if_mock_testing_params_enabled(general_settings={}) assert MOCK_TESTING_CONFIG_KEY not in caplog.text @@ -14932,7 +14932,7 @@ class TestEmbeddingsFailureHookRequestData: with ( patch.object( proxy_server_module, - "_read_request_body", + "read_request_body", new=AsyncMock(return_value={"model": "my-embed", "input": "hello"}), ), patch.object( diff --git a/tests/unit/proxy/test_proxy_setting_guardrails.py b/tests/unit/proxy/test_proxy_setting_guardrails.py index c1d2c640b93..a734dbdbd7d 100644 --- a/tests/unit/proxy/test_proxy_setting_guardrails.py +++ b/tests/unit/proxy/test_proxy_setting_guardrails.py @@ -46,7 +46,7 @@ def test_active_callbacks(client): expected_callback_names = [ "lakeraAI_Moderation", - "_OPTIONAL_PromptInjectionDetectio", + "OPTIONAL_PromptInjectionDetection", "_ENTERPRISE_SecretDetection", ] diff --git a/tests/unit/proxy/test_proxy_token_counter.py b/tests/unit/proxy/test_proxy_token_counter.py index fcd72e3777c..e7a32816464 100644 --- a/tests/unit/proxy/test_proxy_token_counter.py +++ b/tests/unit/proxy/test_proxy_token_counter.py @@ -167,8 +167,8 @@ async def test_anthropic_messages_count_tokens_endpoint(): # Patch the _read_request_body function import litellm.proxy.anthropic_endpoints.endpoints as anthropic_endpoints - original_read_request_body = anthropic_endpoints._read_request_body - anthropic_endpoints._read_request_body = mock_read_request_body + original_read_request_body = anthropic_endpoints.read_request_body + anthropic_endpoints.read_request_body = mock_read_request_body # Mock the internal token_counter function to return a controlled response async def mock_token_counter(request, call_endpoint=False): @@ -207,7 +207,7 @@ async def test_anthropic_messages_count_tokens_endpoint(): finally: # Restore original functions - anthropic_endpoints._read_request_body = original_read_request_body + anthropic_endpoints.read_request_body = original_read_request_body proxy_server.token_counter = original_token_counter @@ -241,8 +241,8 @@ async def test_anthropic_messages_count_tokens_with_non_anthropic_model(): # Patch the _read_request_body function import litellm.proxy.anthropic_endpoints.endpoints as anthropic_endpoints - original_read_request_body = anthropic_endpoints._read_request_body - anthropic_endpoints._read_request_body = mock_read_request_body + original_read_request_body = anthropic_endpoints.read_request_body + anthropic_endpoints.read_request_body = mock_read_request_body # Mock the internal token_counter function to return a controlled response async def mock_token_counter(request, call_endpoint=True): @@ -281,7 +281,7 @@ async def test_anthropic_messages_count_tokens_with_non_anthropic_model(): finally: # Restore original functions - anthropic_endpoints._read_request_body = original_read_request_body + anthropic_endpoints.read_request_body = original_read_request_body proxy_server.token_counter = original_token_counter @@ -381,8 +381,8 @@ async def test_anthropic_endpoint_error_handling(): import litellm.proxy.anthropic_endpoints.endpoints as anthropic_endpoints - original_read_request_body = anthropic_endpoints._read_request_body - anthropic_endpoints._read_request_body = mock_read_request_body + original_read_request_body = anthropic_endpoints.read_request_body + anthropic_endpoints.read_request_body = mock_read_request_body try: # Should raise HTTPException for missing model @@ -395,7 +395,7 @@ async def test_anthropic_endpoint_error_handling(): print("✅ Error handling test passed!") finally: - anthropic_endpoints._read_request_body = original_read_request_body + anthropic_endpoints.read_request_body = original_read_request_body @pytest.mark.asyncio @@ -1111,8 +1111,8 @@ async def test_anthropic_endpoint_returns_anthropic_error_format(): mock_user_api_key_dict = MagicMock() - original_read_request_body = anthropic_endpoints._read_request_body - anthropic_endpoints._read_request_body = mock_read_request_body + original_read_request_body = anthropic_endpoints.read_request_body + anthropic_endpoints.read_request_body = mock_read_request_body original_token_counter = proxy_server.token_counter @@ -1140,7 +1140,7 @@ async def test_anthropic_endpoint_returns_anthropic_error_format(): assert detail["error"]["type"] == "invalid_request_error" assert detail["error"]["message"] == "Input is too long for requested model." finally: - anthropic_endpoints._read_request_body = original_read_request_body + anthropic_endpoints.read_request_body = original_read_request_body proxy_server.token_counter = original_token_counter @@ -1163,8 +1163,8 @@ async def test_anthropic_endpoint_403_permission_error_format(): mock_user_api_key_dict = MagicMock() - original_read_request_body = anthropic_endpoints._read_request_body - anthropic_endpoints._read_request_body = mock_read_request_body + original_read_request_body = anthropic_endpoints.read_request_body + anthropic_endpoints.read_request_body = mock_read_request_body original_token_counter = proxy_server.token_counter @@ -1190,7 +1190,7 @@ async def test_anthropic_endpoint_403_permission_error_format(): assert detail["error"]["type"] == "permission_error" assert detail["error"]["message"] == "Bearer Token has expired" finally: - anthropic_endpoints._read_request_body = original_read_request_body + anthropic_endpoints.read_request_body = original_read_request_body proxy_server.token_counter = original_token_counter @@ -1213,8 +1213,8 @@ async def test_anthropic_endpoint_429_rate_limit_error_format(): mock_user_api_key_dict = MagicMock() - original_read_request_body = anthropic_endpoints._read_request_body - anthropic_endpoints._read_request_body = mock_read_request_body + original_read_request_body = anthropic_endpoints.read_request_body + anthropic_endpoints.read_request_body = mock_read_request_body original_token_counter = proxy_server.token_counter @@ -1240,5 +1240,5 @@ async def test_anthropic_endpoint_429_rate_limit_error_format(): assert detail["error"]["type"] == "rate_limit_error" assert detail["error"]["message"] == "Rate limit exceeded" finally: - anthropic_endpoints._read_request_body = original_read_request_body + anthropic_endpoints.read_request_body = original_read_request_body proxy_server.token_counter = original_token_counter diff --git a/tests/unit/proxy/test_proxy_utils.py b/tests/unit/proxy/test_proxy_utils.py index f52a5bd431a..ba61e393281 100644 --- a/tests/unit/proxy/test_proxy_utils.py +++ b/tests/unit/proxy/test_proxy_utils.py @@ -10,7 +10,7 @@ from fastapi import HTTPException, Request from starlette.datastructures import State from litellm.integrations.custom_guardrail import CustomGuardrail -from litellm.proxy.utils import _get_docs_url, _get_openapi_url, _get_redoc_url +from litellm.proxy.utils import get_docs_url, get_openapi_url, get_redoc_url from litellm.types.guardrails import GuardrailEventHooks from unittest.mock import AsyncMock, MagicMock, patch @@ -22,7 +22,7 @@ from litellm.proxy.auth.auth_utils import ( is_request_body_safe, ) from litellm.proxy.litellm_pre_call_utils import ( - _get_dynamic_logging_metadata, + get_dynamic_logging_metadata, add_litellm_data_to_request, ) from pydantic import ValidationError @@ -294,7 +294,7 @@ def test_dynamic_logging_metadata_key_and_team_metadata(callback_vars): rpm_limit_per_model=None, tpm_limit_per_model=None, ) - callbacks = _get_dynamic_logging_metadata( + callbacks = get_dynamic_logging_metadata( user_api_key_dict=user_api_key_dict, proxy_config=proxy_config ) @@ -332,7 +332,7 @@ def test_dynamic_logging_metadata_ignores_env_references_from_key_metadata( team_metadata={}, ) - callbacks = _get_dynamic_logging_metadata( + callbacks = get_dynamic_logging_metadata( user_api_key_dict=user_api_key_dict, proxy_config=proxy_config ) @@ -410,7 +410,7 @@ def test_dynamic_turn_off_message_logging(callback_vars): rpm_limit_per_model=None, tpm_limit_per_model=None, ) - callbacks = _get_dynamic_logging_metadata( + callbacks = get_dynamic_logging_metadata( user_api_key_dict=user_api_key_dict, proxy_config=proxy_config ) @@ -774,7 +774,7 @@ def test_get_redoc_url(env_vars, expected_url): for key, value in env_vars.items(): os.environ[key] = value - result = _get_redoc_url() + result = get_redoc_url() assert result == expected_url @@ -799,7 +799,7 @@ def test_get_docs_url(env_vars, expected_url): for key, value in env_vars.items(): os.environ[key] = value - result = _get_docs_url() + result = get_docs_url() assert result == expected_url @@ -824,7 +824,7 @@ def test_get_openapi_url(env_vars, expected_url): for key, value in env_vars.items(): os.environ[key] = value - result = _get_openapi_url() + result = get_openapi_url() assert result == expected_url @@ -1569,7 +1569,7 @@ def test_is_allowed_to_make_key_request(): def test_get_model_group_info(): from litellm import Router - from litellm.proxy.proxy_server import _get_model_group_info + from litellm.proxy.proxy_server import get_model_group_info router = Router( model_list=[ @@ -1589,7 +1589,7 @@ def test_get_model_group_info(): }, ] ) - model_list = _get_model_group_info( + model_list = get_model_group_info( llm_router=router, all_models_str=["openai/tts-1", "openai/gpt-3.5-turbo"], model_group="openai/tts-1", diff --git a/tests/unit/proxy/test_proxy_utils_model_creation_and_error_logging.py b/tests/unit/proxy/test_proxy_utils_model_creation_and_error_logging.py index 34260908a5c..ec7104fa459 100644 --- a/tests/unit/proxy/test_proxy_utils_model_creation_and_error_logging.py +++ b/tests/unit/proxy/test_proxy_utils_model_creation_and_error_logging.py @@ -195,7 +195,7 @@ async def test_proxy_only_error_log_keeps_the_request_litellm_call_id(monkeypatc def test_get_model_group_info_order(): from litellm import Router - from litellm.proxy.proxy_server import _get_model_group_info + from litellm.proxy.proxy_server import get_model_group_info router = Router( model_list=[ @@ -215,7 +215,7 @@ def test_get_model_group_info_order(): }, ] ) - model_list = _get_model_group_info( + model_list = get_model_group_info( llm_router=router, all_models_str=["openai/tts-1", "openai/gpt-3.5-turbo"], model_group=None, @@ -277,10 +277,10 @@ def _patch_today(monkeypatch, year, month, day): def test_get_projected_spend_over_limit_day_one(monkeypatch): - from litellm.proxy.utils import _get_projected_spend_over_limit + from litellm.proxy.utils import get_projected_spend_over_limit _patch_today(monkeypatch, 2026, 1, 1) - result = _get_projected_spend_over_limit(100.0, 1.0) + result = get_projected_spend_over_limit(100.0, 1.0) assert result is not None projected_spend, projected_exceeded_date = result @@ -289,10 +289,10 @@ def test_get_projected_spend_over_limit_day_one(monkeypatch): def test_get_projected_spend_over_limit_december(monkeypatch): - from litellm.proxy.utils import _get_projected_spend_over_limit + from litellm.proxy.utils import get_projected_spend_over_limit _patch_today(monkeypatch, 2026, 12, 15) - result = _get_projected_spend_over_limit(100.0, 1.0) + result = get_projected_spend_over_limit(100.0, 1.0) assert result is not None projected_spend, projected_exceeded_date = result @@ -301,10 +301,10 @@ def test_get_projected_spend_over_limit_december(monkeypatch): def test_get_projected_spend_over_limit_includes_current_spend(monkeypatch): - from litellm.proxy.utils import _get_projected_spend_over_limit + from litellm.proxy.utils import get_projected_spend_over_limit _patch_today(monkeypatch, 2026, 4, 11) - result = _get_projected_spend_over_limit(100.0, 200.0) + result = get_projected_spend_over_limit(100.0, 200.0) assert result is not None projected_spend, projected_exceeded_date = result @@ -629,7 +629,7 @@ class TestPostCallFailureHookLiftsCallTypeAndStartTime: from unittest.mock import AsyncMock from litellm.litellm_core_utils.litellm_logging import Logging - from litellm.proxy.hooks.proxy_track_cost_callback import _ProxyDBLogger + from litellm.proxy.hooks.proxy_track_cost_callback import ProxyDBLogger from litellm.proxy.spend_tracking.spend_tracking_utils import get_logging_payload request_start = real_datetime.datetime.now() - real_datetime.timedelta(seconds=2) @@ -659,7 +659,7 @@ class TestPostCallFailureHookLiftsCallTypeAndStartTime: proxy_logging_obj.alert_types = [] spend_writer = SimpleNamespace(update_database=AsyncMock()) original_callbacks = list(litellm.callbacks) - litellm.callbacks = [_ProxyDBLogger(spend_writer=lambda: spend_writer)] + litellm.callbacks = [ProxyDBLogger(spend_writer=lambda: spend_writer)] try: with patch.object(proxy_logging_obj, "update_request_status", new=AsyncMock()): await proxy_logging_obj.post_call_failure_hook( diff --git a/tests/unit/proxy/test_response_polling_pre_call_checks.py b/tests/unit/proxy/test_response_polling_pre_call_checks.py index 1dca00fd5fb..c6e4301a501 100644 --- a/tests/unit/proxy/test_response_polling_pre_call_checks.py +++ b/tests/unit/proxy/test_response_polling_pre_call_checks.py @@ -35,9 +35,7 @@ class TestSkipPreCallLogic: mock_proxy_logging.during_call_hook = AsyncMock() with ( - patch.object( - processor, "common_processing_pre_call_logic", new_callable=AsyncMock - ) as mock_pre_call, + patch.object(processor, "common_processing_pre_call_logic", new_callable=AsyncMock) as mock_pre_call, patch( "litellm.proxy.common_request_processing.route_request", new_callable=AsyncMock, @@ -119,7 +117,7 @@ class TestPollingEndpointPreCallGuard: generate_polling_id_mock = MagicMock(return_value="litellm_poll_test") proxy_server_patches = { - "litellm.proxy.proxy_server._read_request_body": AsyncMock( + "litellm.proxy.proxy_server.read_request_body": AsyncMock( return_value={"model": "gpt-4", "background": True} ), "litellm.proxy.proxy_server.general_settings": {}, @@ -155,15 +153,11 @@ class TestPollingEndpointPreCallGuard: ), patch.object( ProxyBaseLLMRequestProcessing, - "_handle_llm_api_exception", + "handle_llm_api_exception", new_callable=AsyncMock, - return_value=HTTPException( - status_code=429, detail="Rate limit exceeded" - ), - ), - patch.object( - ResponsePollingHandler, "generate_polling_id", generate_polling_id_mock + return_value=HTTPException(status_code=429, detail="Rate limit exceeded"), ), + patch.object(ResponsePollingHandler, "generate_polling_id", generate_polling_id_mock), # Prevent background task from running (avoids noise from incomplete mocks) patch("asyncio.create_task"), patch.object( diff --git a/tests/unit/proxy/test_route_llm_request.py b/tests/unit/proxy/test_route_llm_request.py index 3517f1412d7..a891e1079d8 100644 --- a/tests/unit/proxy/test_route_llm_request.py +++ b/tests/unit/proxy/test_route_llm_request.py @@ -1,6 +1,7 @@ import pytest +from types import SimpleNamespace from typing import Final from unittest.mock import MagicMock @@ -1200,6 +1201,175 @@ def test_required_present_body_param_without_router_default_still_raises() -> No assert exc_info.value.param == "max_tokens" +@pytest.mark.parametrize( + "route_type, data, param, default", + [ + ("anthropic_messages", {"model": "claude-router-default", "messages": []}, "max_tokens", 16), + ("arerank", {"model": "rerank-router-default", "query": "hi"}, "documents", ["router default document"]), + ], +) +def test_required_present_body_param_uses_router_wide_default( + route_type: str, data: dict[str, object], param: str, default: object +) -> None: + import litellm + from litellm.proxy.route_llm_request import raise_if_required_body_param_missing + + router = litellm.Router( + model_list=[ + { + "model_name": str(data["model"]), + "litellm_params": { + "model": "openai/gpt-4o-mini", + "api_key": "test-key", + }, + } + ], + default_litellm_params={param: default}, + ) + + raise_if_required_body_param_missing(route_type=route_type, data=data, llm_router=router) + + +@pytest.mark.parametrize( + "route_type, data, param", + [ + ("anthropic_messages", {"model": "claude", "messages": []}, "max_tokens"), + ("arerank", {"model": "rerank-model", "query": "hi"}, "documents"), + ], +) +def test_required_present_body_param_with_none_router_wide_default_still_raises( + route_type: str, data: dict[str, object], param: str +) -> None: + import litellm + from litellm.proxy.route_llm_request import ( + ProxyMissingRequiredParamError, + raise_if_required_body_param_missing, + ) + + router = litellm.Router( + model_list=[ + { + "model_name": str(data["model"]), + "litellm_params": { + "model": "openai/gpt-4o-mini", + "api_key": "test-key", + }, + } + ], + default_litellm_params={param: None}, + ) + + with pytest.raises(ProxyMissingRequiredParamError) as exc_info: + raise_if_required_body_param_missing(route_type=route_type, data=data, llm_router=router) + + assert exc_info.value.param == param + + +@pytest.mark.parametrize( + "route_type, data, param", + [ + ("asearch", {"model": "search-model"}, "query"), + ("acreate_eval", {"model": "eval-model"}, "data_source_config"), + ("acreate_run", {"model": "eval-model"}, "data_source"), + ("acreate_agent", {"model": "agent-model"}, "name"), + ("avector_store_search", {}, "query"), + ], +) +def test_required_present_body_param_ignores_router_wide_default_when_dispatch_skips_router( + route_type: str, data: dict[str, object], param: str +) -> None: + import litellm + from litellm.proxy.route_llm_request import ( + ProxyMissingRequiredParamError, + raise_if_required_body_param_missing, + ) + + router = litellm.Router( + model_list=[ + { + "model_name": str(data.get("model") or "any-model"), + "litellm_params": { + "model": "openai/gpt-4o-mini", + "api_key": "test-key", + }, + } + ], + default_litellm_params={param: "supplied-by-router-default"}, + ) + + with pytest.raises(ProxyMissingRequiredParamError) as exc_info: + raise_if_required_body_param_missing(route_type=route_type, data=data, llm_router=router) + + assert exc_info.value.param == param + + +def test_required_present_body_param_uses_user_config_router_defaults() -> None: + from litellm.proxy.route_llm_request import raise_if_required_body_param_missing + + raise_if_required_body_param_missing( + route_type="anthropic_messages", + data={ + "model": "claude-user-config", + "messages": [], + "user_config": {"default_litellm_params": {"max_tokens": 16}}, + }, + llm_router=None, + ) + + +def test_required_present_body_param_ignores_global_router_defaults_for_user_config_requests() -> None: + import litellm + from litellm.proxy.route_llm_request import ( + ProxyMissingRequiredParamError, + raise_if_required_body_param_missing, + ) + + router = litellm.Router( + model_list=[ + { + "model_name": "claude-user-config", + "litellm_params": { + "model": "openai/gpt-4o-mini", + "api_key": "test-key", + }, + } + ], + default_litellm_params={"max_tokens": 16}, + ) + + with pytest.raises(ProxyMissingRequiredParamError) as exc_info: + raise_if_required_body_param_missing( + route_type="anthropic_messages", + data={"model": "claude-user-config", "messages": [], "user_config": {}}, + llm_router=router, + ) + + assert exc_info.value.param == "max_tokens" + + +@pytest.mark.parametrize( + "llm_router", + [ + pytest.param(MagicMock(), id="magicmock-router"), + pytest.param(SimpleNamespace(get_model_list=lambda **_kwargs: ()), id="router-without-defaults-attr"), + ], +) +def test_required_present_body_param_ignores_non_mapping_router_defaults(llm_router: object) -> None: + from litellm.proxy.route_llm_request import ( + ProxyMissingRequiredParamError, + raise_if_required_body_param_missing, + ) + + with pytest.raises(ProxyMissingRequiredParamError) as exc_info: + raise_if_required_body_param_missing( + route_type="arerank", + data={"model": "rerank-model", "query": "hi"}, + llm_router=llm_router, # pyright: ignore[reportArgumentType] # deliberately duck-typed routers + ) + + assert exc_info.value.param == "documents" + + @pytest.mark.parametrize( "route_type, data, param", [ diff --git a/tests/unit/proxy/test_team_member_update.py b/tests/unit/proxy/test_team_member_update.py index ace4c4e65af..d77904fc05d 100644 --- a/tests/unit/proxy/test_team_member_update.py +++ b/tests/unit/proxy/test_team_member_update.py @@ -91,21 +91,17 @@ def happy_path_upsert(monkeypatch): AsyncMock( return_value={ "team_info": team_row, - "team_memberships": [ - types.SimpleNamespace(user_id="user-1", budget_id="bud-1") - ], + "team_memberships": [types.SimpleNamespace(user_id="user-1", budget_id="bud-1")], } ), ) upsert_mock = AsyncMock() - monkeypatch.setattr(team_endpoints, "_upsert_budget_and_membership", upsert_mock) + monkeypatch.setattr(team_endpoints, "upsert_budget_and_membership", upsert_mock) return upsert_mock def _member_update_request(**overrides): - data = TeamMemberUpdateRequest( - team_id="team-1234", user_id="user-1", role="user", **overrides - ) + data = TeamMemberUpdateRequest(team_id="team-1234", user_id="user-1", role="user", **overrides) request = Request({"type": "http", "method": "POST", "path": "/team/member_update"}) auth = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN.value, user_id="admin") return data, request, auth @@ -115,9 +111,7 @@ def _member_update_request(**overrides): async def test_team_member_update_sends_provided_fields_as_patch(happy_path_upsert): """Fields the request sets must reach _upsert_budget_and_membership as a budget patch, otherwise the member budget is never written/reset.""" - data, request, auth = _member_update_request( - max_budget_in_team=10.0, budget_duration="30d" - ) + data, request, auth = _member_update_request(max_budget_in_team=10.0, budget_duration="30d") response = await team_member_update(data, request, auth) @@ -137,9 +131,7 @@ async def test_team_member_update_explicit_null_clears_field(happy_path_upsert): await team_member_update(data, request, auth) - assert happy_path_upsert.await_args.kwargs["budget_patch"] == { - "budget_duration": None - } + assert happy_path_upsert.await_args.kwargs["budget_patch"] == {"budget_duration": None} @pytest.mark.asyncio @@ -163,15 +155,13 @@ async def test_team_member_update_omits_unset_fields_from_patch(happy_path_upser ], ) @pytest.mark.asyncio -async def test_team_member_update_rejects_invalid_budget_duration( - monkeypatch, bad_duration -): +async def test_team_member_update_rejects_invalid_budget_duration(monkeypatch, bad_duration): """An invalid budget_duration must be rejected with a 400 before any DB write, so it can never be persisted and later break the budget reset job.""" monkeypatch.setattr(proxy_server, "prisma_client", object()) monkeypatch.setattr(proxy_server, "premium_user", False) upsert_mock = AsyncMock() - monkeypatch.setattr(team_endpoints, "_upsert_budget_and_membership", upsert_mock) + monkeypatch.setattr(team_endpoints, "upsert_budget_and_membership", upsert_mock) data = TeamMemberUpdateRequest( team_id="team-1234", diff --git a/tests/unit/proxy/test_unit_test_proxy_hooks.py b/tests/unit/proxy/test_unit_test_proxy_hooks.py index e6ffea35e52..6ce7154f2d1 100644 --- a/tests/unit/proxy/test_unit_test_proxy_hooks.py +++ b/tests/unit/proxy/test_unit_test_proxy_hooks.py @@ -2,7 +2,7 @@ import asyncio from unittest.mock import Mock, patch, AsyncMock import pytest from fastapi import Request -from litellm.proxy.utils import _get_redoc_url, _get_docs_url +from litellm.proxy.utils import get_redoc_url, get_docs_url from datetime import datetime import litellm diff --git a/tests/unit/proxy/test_update_spend.py b/tests/unit/proxy/test_update_spend.py index 6b92320762b..62adb7303e2 100644 --- a/tests/unit/proxy/test_update_spend.py +++ b/tests/unit/proxy/test_update_spend.py @@ -1,6 +1,6 @@ import asyncio from unittest.mock import Mock -from litellm.proxy.utils import _get_redoc_url, _get_docs_url +from litellm.proxy.utils import get_redoc_url, get_docs_url import pytest from fastapi import Request diff --git a/tests/unit/proxy/test_zero_cost_model_budget_bypass.py b/tests/unit/proxy/test_zero_cost_model_budget_bypass.py index 56133f2d35b..c2931470482 100644 --- a/tests/unit/proxy/test_zero_cost_model_budget_bypass.py +++ b/tests/unit/proxy/test_zero_cost_model_budget_bypass.py @@ -23,8 +23,8 @@ from litellm.proxy._types import ( ) from litellm.proxy.auth.auth_checks import ( _check_team_member_budget, - _is_model_cost_zero, - _team_max_budget_check, + is_model_cost_zero, + team_max_budget_check, common_checks, ) from litellm.proxy.utils import ProxyLogging @@ -103,7 +103,7 @@ class TestIsModelCostZero: def test_zero_cost_model_in_router(self, mock_router_with_zero_cost_model): """Test that a zero-cost model in router is correctly identified.""" - result = _is_model_cost_zero( + result = is_model_cost_zero( model="on-prem-model", llm_router=mock_router_with_zero_cost_model ) assert result is True @@ -116,26 +116,26 @@ class TestIsModelCostZero: "input_cost_per_token": 0.0000015, "output_cost_per_token": 0.000002, } - result = _is_model_cost_zero( + result = is_model_cost_zero( model="cloud-model", llm_router=mock_router_with_zero_cost_model ) assert result is False def test_none_model(self, mock_router_with_zero_cost_model): """Test that None model returns False.""" - result = _is_model_cost_zero( + result = is_model_cost_zero( model=None, llm_router=mock_router_with_zero_cost_model ) assert result is False def test_none_router(self): """Test that None router returns False.""" - result = _is_model_cost_zero(model="some-model", llm_router=None) + result = is_model_cost_zero(model="some-model", llm_router=None) assert result is False def test_list_of_zero_cost_models(self, mock_router_with_zero_cost_model): """Test that a list of zero-cost models returns True.""" - result = _is_model_cost_zero( + result = is_model_cost_zero( model=["on-prem-model"], llm_router=mock_router_with_zero_cost_model ) assert result is True @@ -147,7 +147,7 @@ class TestIsModelCostZero: "input_cost_per_token": 0.0000015, "output_cost_per_token": 0.000002, } - result = _is_model_cost_zero( + result = is_model_cost_zero( model=["on-prem-model", "cloud-model"], llm_router=mock_router_with_zero_cost_model, ) @@ -514,7 +514,7 @@ class TestEdgeCases: with patch("litellm.get_model_info") as mock_get_model_info: # Simulate model not found mock_get_model_info.side_effect = Exception("Model not found") - result = _is_model_cost_zero( + result = is_model_cost_zero( model="nonexistent-model", llm_router=mock_router_with_zero_cost_model ) # Should return False (conservative approach) diff --git a/tests/unit/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py b/tests/unit/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py index bea305417f3..0cdf264d8fc 100644 --- a/tests/unit/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py +++ b/tests/unit/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py @@ -445,7 +445,7 @@ class TestProxySettingEndpoints: monkeypatch.setattr( proxy_config, - "_decrypt_and_set_db_env_variables", + "decrypt_and_set_db_env_variables", lambda environment_variables: environment_variables, ) @@ -655,7 +655,7 @@ class TestProxySettingEndpoints: monkeypatch.setattr( proxy_config, - "_encrypt_env_variables", + "encrypt_env_variables", lambda environment_variables: environment_variables, ) @@ -739,7 +739,7 @@ class TestProxySettingEndpoints: monkeypatch.setattr( proxy_config, - "_encrypt_env_variables", + "encrypt_env_variables", lambda environment_variables: environment_variables, ) @@ -806,7 +806,7 @@ class TestProxySettingEndpoints: monkeypatch.setattr( proxy_config, - "_encrypt_env_variables", + "encrypt_env_variables", lambda environment_variables: environment_variables, ) @@ -874,7 +874,7 @@ class TestProxySettingEndpoints: monkeypatch.setattr( proxy_config, - "_encrypt_env_variables", + "encrypt_env_variables", lambda environment_variables: environment_variables, ) @@ -953,7 +953,7 @@ class TestProxySettingEndpoints: monkeypatch.setattr( proxy_config, - "_encrypt_env_variables", + "encrypt_env_variables", lambda environment_variables: environment_variables, ) @@ -1024,7 +1024,7 @@ class TestProxySettingEndpoints: monkeypatch.setattr( proxy_config, - "_encrypt_env_variables", + "encrypt_env_variables", lambda environment_variables: environment_variables, ) @@ -1092,7 +1092,7 @@ class TestProxySettingEndpoints: monkeypatch.setattr( proxy_config, - "_encrypt_env_variables", + "encrypt_env_variables", lambda environment_variables: environment_variables, ) @@ -1964,7 +1964,7 @@ class TestProxySettingEndpoints: from litellm.proxy.proxy_server import proxy_config - monkeypatch.setattr(proxy_config, "_encrypt_env_variables", mock_encrypt) + monkeypatch.setattr(proxy_config, "encrypt_env_variables", mock_encrypt) # New SSO settings to save new_sso_settings = { @@ -2043,7 +2043,7 @@ class TestProxySettingEndpoints: monkeypatch.setattr( proxy_config, - "_encrypt_env_variables", + "encrypt_env_variables", lambda environment_variables: environment_variables, ) @@ -2092,7 +2092,7 @@ class TestProxySettingEndpoints: monkeypatch.setattr( proxy_config, - "_encrypt_env_variables", + "encrypt_env_variables", lambda environment_variables: environment_variables, ) @@ -2127,7 +2127,7 @@ class TestProxySettingEndpoints: from litellm.proxy.proxy_server import proxy_config monkeypatch.setattr( - proxy_config, "_decrypt_and_set_db_env_variables", mock_decrypt_and_set + proxy_config, "decrypt_and_set_db_env_variables", mock_decrypt_and_set ) response = client.get("/get/sso_settings") @@ -2212,7 +2212,7 @@ class TestProxySettingEndpoints: return environment_variables monkeypatch.setattr( - proxy_config, "_decrypt_and_set_db_env_variables", mock_decrypt + proxy_config, "decrypt_and_set_db_env_variables", mock_decrypt ) response = client.get("/get/sso_settings") @@ -2257,7 +2257,7 @@ class TestProxySettingEndpoints: monkeypatch.setattr( proxy_config, - "_encrypt_env_variables", + "encrypt_env_variables", lambda environment_variables: environment_variables, ) @@ -2309,7 +2309,7 @@ class TestProxySettingEndpoints: ) monkeypatch.setattr( proxy_config, - "_decrypt_and_set_db_env_variables", + "decrypt_and_set_db_env_variables", lambda environment_variables: environment_variables, ) @@ -2431,7 +2431,7 @@ class TestProxySettingEndpoints: monkeypatch.setattr( proxy_config, - "_decrypt_and_set_db_env_variables", + "decrypt_and_set_db_env_variables", lambda environment_variables: environment_variables, ) @@ -2581,7 +2581,7 @@ def test_update_sso_settings_writes_redacted_audit_log(mock_proxy_config, monkey monkeypatch.setattr(litellm, "store_audit_logs", True) monkeypatch.setattr( proxy_server_module.proxy_config, - "_encrypt_env_variables", + "encrypt_env_variables", lambda environment_variables: environment_variables, ) @@ -2653,14 +2653,14 @@ def test_update_sso_settings_audit_captures_redacted_before_snapshot( monkeypatch.setattr(litellm, "store_audit_logs", True) monkeypatch.setattr( proxy_server_module.proxy_config, - "_encrypt_env_variables", + "encrypt_env_variables", lambda environment_variables: environment_variables, ) # Pretend the stored value is already plaintext for the test (production # decrypts via Fernet); the audit helper still has to redact it. monkeypatch.setattr( proxy_server_module.proxy_config, - "_decrypt_db_variables", + "decrypt_db_variables", lambda variables_dict: dict(variables_dict), ) @@ -2934,7 +2934,7 @@ def test_update_ui_theme_settings_writes_audit_log(mock_proxy_config, monkeypatc monkeypatch.setattr(litellm, "store_audit_logs", True) monkeypatch.setattr( proxy_server_module.proxy_config, - "_encrypt_env_variables", + "encrypt_env_variables", lambda environment_variables: environment_variables, ) @@ -4165,3 +4165,71 @@ class TestSyncUiSettingsToGeneralSettings: assert general_settings["forward_client_headers_to_llm_api"] is False assert general_settings.source("forward_client_headers_to_llm_api") == "config" + + +class TestMoyaiUrlSetting: + @pytest.mark.parametrize( + "value,expected", + [ + (None, None), + ("", None), + ("https://moyai.example.com", "https://moyai.example.com"), + ("https://moyai.example.com/", "https://moyai.example.com"), + ("http://localhost:8787/", "http://localhost:8787"), + ], + ) + def test_moyai_url_validator_accepts_and_normalizes(self, value, expected): + from litellm.proxy.ui_crud_endpoints.proxy_setting_endpoints import UISettings + + assert UISettings(moyai_url=value).moyai_url == expected + + @pytest.mark.parametrize( + "value", + [ + "javascript:alert(1)", + "ftp://moyai.example.com", + "https://user:pass@moyai.example.com", + "https://user@moyai.example.com", + "not-a-url", + "https://", + ], + ) + def test_moyai_url_validator_rejects(self, value): + from pydantic import ValidationError + from litellm.proxy.ui_crud_endpoints.proxy_setting_endpoints import UISettings + + with pytest.raises(ValidationError): + UISettings(moyai_url=value) + + def test_moyai_url_is_in_allowed_ui_settings_fields(self): + from litellm.proxy.ui_crud_endpoints.proxy_setting_endpoints import ALLOWED_UI_SETTINGS_FIELDS + + assert "moyai_url" in ALLOWED_UI_SETTINGS_FIELDS + + @pytest.mark.asyncio + async def test_moyai_url_patch_sets_and_clears(self, monkeypatch): + from types import SimpleNamespace + from unittest.mock import AsyncMock, MagicMock + from litellm.proxy import proxy_server + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.ui_crud_endpoints.proxy_setting_endpoints import update_ui_settings + + prisma = MagicMock() + prisma.db.litellm_uisettings.find_unique = AsyncMock( + return_value=SimpleNamespace(ui_settings={"moyai_url": "https://old.example.com"}) + ) + persisted: dict = {} + + async def _upsert(where, data): + persisted.update(json.loads(data["update"]["ui_settings"])) + + prisma.db.litellm_uisettings.upsert = AsyncMock(side_effect=_upsert) + monkeypatch.setattr(proxy_server, "prisma_client", prisma) + monkeypatch.setattr(proxy_server, "store_model_in_db", True) + actor = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN) + + await update_ui_settings({"moyai_url": "https://new.example.com/"}, actor) + assert persisted["moyai_url"] == "https://new.example.com" + + await update_ui_settings({"moyai_url": None}, actor) + assert persisted["moyai_url"] is None diff --git a/tests/unit/proxy/utils/helpers/test_guardrail_merge.py b/tests/unit/proxy/utils/helpers/test_guardrail_merge.py index be00e36acae..e505517eeff 100644 --- a/tests/unit/proxy/utils/helpers/test_guardrail_merge.py +++ b/tests/unit/proxy/utils/helpers/test_guardrail_merge.py @@ -4,7 +4,7 @@ from unittest.mock import MagicMock import pytest from litellm.proxy.utils import ( - _check_and_merge_model_level_guardrails, + check_and_merge_model_level_guardrails, _merge_guardrails_with_existing, ) @@ -55,7 +55,7 @@ def test_check_and_merge_model_level_guardrails_happy_path_merges_lists(): "guardrails": ["user-policy"], }, } - result = _check_and_merge_model_level_guardrails(data, router) + result = check_and_merge_model_level_guardrails(data, router) snapshot = { "model": result["model"], "model_info_id": result["metadata"]["model_info"]["id"], @@ -70,7 +70,7 @@ def test_check_and_merge_model_level_guardrails_happy_path_merges_lists(): def test_check_and_merge_model_level_guardrails_returns_data_when_router_none(): data = {"metadata": {"model_info": {"id": "x"}}, "model": "m", "other": 1} - result = _check_and_merge_model_level_guardrails(data, None) + result = check_and_merge_model_level_guardrails(data, None) assert result is data assert normalize(result) == { "metadata": {"model_info": {"id": "x"}}, @@ -84,7 +84,7 @@ def test_check_and_merge_model_level_guardrails_returns_data_when_model_id_missi deployment (router returns None for both lookups), data is unchanged.""" router = _router_with_deployment(["pii"]) # by_alias=False by default data = {"metadata": {"model_info": {}}, "model": "m", "extra": "v"} - result = _check_and_merge_model_level_guardrails(data, router) + result = check_and_merge_model_level_guardrails(data, router) snapshot = { "is_same_object": result is data, "metadata": result["metadata"], @@ -108,7 +108,7 @@ def test_check_and_merge_model_level_guardrails_falls_back_to_model_alias_when_m model alias (#29652) so DB/UI-assigned guardrails still fire.""" router = _router_with_deployment(["pii"], by_alias=True) data = {"metadata": {"model_info": {}}, "model": "m", "extra": "v"} - result = _check_and_merge_model_level_guardrails(data, router) + result = check_and_merge_model_level_guardrails(data, router) # Merge happened via the alias fallback. assert "pii" in result["metadata"]["guardrails"] router.get_model_list.assert_called_once() @@ -121,7 +121,7 @@ def test_check_and_merge_model_level_guardrails_unions_guardrails_across_group_d The fix is to union the guardrails from all deployments in the group.""" router = _router_with_deployments([["pii"], ["secret-scan"], None]) data = {"metadata": {"model_info": {}}, "model": "m"} - result = _check_and_merge_model_level_guardrails(data, router) + result = check_and_merge_model_level_guardrails(data, router) assert sorted(result["metadata"]["guardrails"]) == ["pii", "secret-scan"] @@ -130,7 +130,7 @@ def test_check_and_merge_model_level_guardrails_dedups_guardrails_across_group_d entries in the merged guardrails list.""" router = _router_with_deployments([["pii"], ["pii", "secret-scan"]]) data = {"metadata": {"model_info": {}}, "model": "m"} - result = _check_and_merge_model_level_guardrails(data, router) + result = check_and_merge_model_level_guardrails(data, router) assert sorted(result["metadata"]["guardrails"]) == ["pii", "secret-scan"] @@ -139,7 +139,7 @@ def test_check_and_merge_model_level_guardrails_group_with_no_guardrails_returns the helper returns the data unchanged.""" router = _router_with_deployments([None, None, []]) data = {"metadata": {"model_info": {}}, "model": "m"} - result = _check_and_merge_model_level_guardrails(data, router) + result = check_and_merge_model_level_guardrails(data, router) assert result is data assert "guardrails" not in result["metadata"] @@ -163,7 +163,7 @@ def test_check_and_merge_model_level_guardrails_ignores_client_model_info_id_whe "model": "guarded-alias", "metadata": {"model_info": {"id": "spoofed-unguarded-deployment"}}, } - result = _check_and_merge_model_level_guardrails(data, router, trust_client_model_info=False) + result = check_and_merge_model_level_guardrails(data, router, trust_client_model_info=False) assert "alias-secret-scan" in result["metadata"]["guardrails"] # The model_id lookup must NOT have been used. router.get_deployment.assert_not_called() @@ -179,7 +179,7 @@ def test_check_and_merge_model_level_guardrails_trusts_client_model_info_id_by_d "model": "any", "metadata": {"model_info": {"id": "deployment-123"}}, } - result = _check_and_merge_model_level_guardrails(data, router) + result = check_and_merge_model_level_guardrails(data, router) assert "post-call-guardrail" in result["metadata"]["guardrails"] router.get_deployment.assert_called_once_with(model_id="deployment-123") @@ -193,7 +193,7 @@ def test_check_and_merge_model_level_guardrails_post_call_accepts_bare_string_gu router = MagicMock() router.get_deployment.return_value = deployment data = {"model": "any", "metadata": {"model_info": {"id": "deployment-x"}}} - result = _check_and_merge_model_level_guardrails(data, router) + result = check_and_merge_model_level_guardrails(data, router) assert "scalar-guardrail" in result["metadata"]["guardrails"] @@ -203,7 +203,7 @@ def test_check_and_merge_model_level_guardrails_alias_union_accepts_bare_string_ router.get_deployment.return_value = None router.get_model_list.return_value = [{"litellm_params": {"guardrails": "scalar-alias-guardrail"}}] data = {"model": "alias-m", "metadata": {"model_info": {}}} - result = _check_and_merge_model_level_guardrails(data, router) + result = check_and_merge_model_level_guardrails(data, router) assert "scalar-alias-guardrail" in result["metadata"]["guardrails"] @@ -222,7 +222,7 @@ def test_check_and_merge_model_level_guardrails_alias_fallback_passes_team_id(): "user_api_key_team_id": "team-abc", }, } - result = _check_and_merge_model_level_guardrails(data, router, trust_client_model_info=False) + result = check_and_merge_model_level_guardrails(data, router, trust_client_model_info=False) assert "team-guardrail" in result["metadata"]["guardrails"] router.get_model_list.assert_called_once_with(model_name="team-scoped-alias", team_id="team-abc") @@ -238,21 +238,21 @@ def test_check_and_merge_model_level_guardrails_alias_fallback_reads_team_id_fro "metadata": {"model_info": {}}, "litellm_metadata": {"user_api_key_team_id": "team-xyz"}, } - _check_and_merge_model_level_guardrails(data, router, trust_client_model_info=False) + check_and_merge_model_level_guardrails(data, router, trust_client_model_info=False) router.get_model_list.assert_called_once_with(model_name="alias-m", team_id="team-xyz") def test_check_and_merge_model_level_guardrails_returns_data_when_deployment_none(): router = _router_without_deployment() data = {"metadata": {"model_info": {"id": "x"}}, "model": "m"} - result = _check_and_merge_model_level_guardrails(data, router) + result = check_and_merge_model_level_guardrails(data, router) assert result is data def test_check_and_merge_model_level_guardrails_returns_data_when_guardrails_none(): router = _router_with_deployment(None) data = {"metadata": {"model_info": {"id": "x"}}, "model": "m"} - result = _check_and_merge_model_level_guardrails(data, router) + result = check_and_merge_model_level_guardrails(data, router) assert result is data @@ -260,7 +260,7 @@ def test_check_and_merge_model_level_guardrails_handles_missing_metadata(): """No metadata at all + alias unknown to the router → data unchanged.""" router = _router_with_deployment(["pii"]) # by_alias=False data = {"model": "m"} - result = _check_and_merge_model_level_guardrails(data, router) + result = check_and_merge_model_level_guardrails(data, router) snapshot = { "is_same_object": result is data, "model": result["model"], @@ -277,7 +277,7 @@ def test_check_and_merge_model_level_guardrails_raises_when_metadata_is_not_dict router = _router_with_deployment(["pii"]) data = {"metadata": "not-a-dict", "model": "m"} with pytest.raises(AttributeError): - _check_and_merge_model_level_guardrails(data, router) + check_and_merge_model_level_guardrails(data, router) def test_merge_guardrails_with_existing_happy_path_combines_lists(): diff --git a/tests/unit/proxy/utils/helpers/test_month_end_projection.py b/tests/unit/proxy/utils/helpers/test_month_end_projection.py index 5afe1f4faf8..8c1a4ae746d 100644 --- a/tests/unit/proxy/utils/helpers/test_month_end_projection.py +++ b/tests/unit/proxy/utils/helpers/test_month_end_projection.py @@ -4,8 +4,8 @@ import pytest from litellm.proxy.utils import ( _get_month_end_date, - _get_projected_spend_over_limit, - _is_projected_spend_over_limit, + get_projected_spend_over_limit, + is_projected_spend_over_limit, ) @@ -59,9 +59,7 @@ def test_get_month_end_date_raises_on_non_date_input(): def test_is_projected_spend_over_limit_happy_path_under_budget(monkeypatch): _freeze_today(monkeypatch, date(2024, 1, 11)) summary = { - "result": _is_projected_spend_over_limit( - current_spend=10.0, soft_budget_limit=1_000_000.0 - ), + "result": is_projected_spend_over_limit(current_spend=10.0, soft_budget_limit=1_000_000.0), "current_spend": 10.0, "soft_budget_limit": 1_000_000.0, } @@ -75,9 +73,7 @@ def test_is_projected_spend_over_limit_happy_path_under_budget(monkeypatch): def test_is_projected_spend_over_limit_happy_path_over_budget(monkeypatch): _freeze_today(monkeypatch, date(2024, 1, 11)) summary = { - "result": _is_projected_spend_over_limit( - current_spend=100.0, soft_budget_limit=50.0 - ), + "result": is_projected_spend_over_limit(current_spend=100.0, soft_budget_limit=50.0), "current_spend": 100.0, "soft_budget_limit": 50.0, } @@ -91,9 +87,7 @@ def test_is_projected_spend_over_limit_happy_path_over_budget(monkeypatch): def test_is_projected_spend_over_limit_first_of_month_no_division_by_zero(monkeypatch): _freeze_today(monkeypatch, date(2024, 1, 1)) summary = { - "result": _is_projected_spend_over_limit( - current_spend=5.0, soft_budget_limit=10.0 - ), + "result": is_projected_spend_over_limit(current_spend=5.0, soft_budget_limit=10.0), "current_spend": 5.0, "soft_budget_limit": 10.0, } @@ -105,10 +99,7 @@ def test_is_projected_spend_over_limit_first_of_month_no_division_by_zero(monkey def test_is_projected_spend_over_limit_none_limit_returns_false(): - assert ( - _is_projected_spend_over_limit(current_spend=10_000.0, soft_budget_limit=None) - is False - ) + assert is_projected_spend_over_limit(current_spend=10_000.0, soft_budget_limit=None) is False def test_is_projected_spend_over_limit_raises_when_today_missing(monkeypatch): @@ -119,14 +110,12 @@ def test_is_projected_spend_over_limit_raises_when_today_missing(monkeypatch): monkeypatch.setattr("litellm.proxy.utils.date", _Broken) with pytest.raises(RuntimeError): - _is_projected_spend_over_limit(current_spend=1.0, soft_budget_limit=1.0) + is_projected_spend_over_limit(current_spend=1.0, soft_budget_limit=1.0) def test_get_projected_spend_over_limit_happy_path_over_budget(monkeypatch): _freeze_today(monkeypatch, date(2024, 1, 11)) - result = _get_projected_spend_over_limit( - current_spend=100.0, soft_budget_limit=50.0 - ) + result = get_projected_spend_over_limit(current_spend=100.0, soft_budget_limit=50.0) assert result is not None projected, exceed_date = result summary = { @@ -147,7 +136,7 @@ def test_get_projected_spend_over_limit_first_of_month_uses_current_as_daily( monkeypatch, ): _freeze_today(monkeypatch, date(2024, 1, 1)) - result = _get_projected_spend_over_limit(current_spend=5.0, soft_budget_limit=10.0) + result = get_projected_spend_over_limit(current_spend=5.0, soft_budget_limit=10.0) assert result is not None projected, exceed_date = result expected_exceed = date(2024, 1, 1) + timedelta(days=1.0) @@ -167,7 +156,7 @@ def test_get_projected_spend_over_limit_first_of_month_uses_current_as_daily( def test_get_projected_spend_over_limit_zero_daily_spend_exceed_today(monkeypatch): _freeze_today(monkeypatch, date(2024, 1, 11)) - result = _get_projected_spend_over_limit(current_spend=0.0, soft_budget_limit=-1.0) + result = get_projected_spend_over_limit(current_spend=0.0, soft_budget_limit=-1.0) assert result is not None projected, exceed_date = result summary = { @@ -184,17 +173,12 @@ def test_get_projected_spend_over_limit_zero_daily_spend_exceed_today(monkeypatc def test_get_projected_spend_over_limit_under_budget_returns_none(monkeypatch): _freeze_today(monkeypatch, date(2024, 1, 11)) - assert ( - _get_projected_spend_over_limit( - current_spend=1.0, soft_budget_limit=1_000_000.0 - ) - is None - ) + assert get_projected_spend_over_limit(current_spend=1.0, soft_budget_limit=1_000_000.0) is None def test_get_projected_spend_over_limit_exceed_date_uses_remaining_budget(monkeypatch): _freeze_today(monkeypatch, date(2024, 1, 11)) - result = _get_projected_spend_over_limit(current_spend=20.0, soft_budget_limit=30.0) + result = get_projected_spend_over_limit(current_spend=20.0, soft_budget_limit=30.0) assert result is not None projected, exceed_date = result daily = 20.0 / 10 @@ -215,10 +199,7 @@ def test_get_projected_spend_over_limit_exceed_date_uses_remaining_budget(monkey def test_get_projected_spend_over_limit_none_limit_returns_none(): - assert ( - _get_projected_spend_over_limit(current_spend=1.0, soft_budget_limit=None) - is None - ) + assert get_projected_spend_over_limit(current_spend=1.0, soft_budget_limit=None) is None def test_get_projected_spend_over_limit_raises_when_today_missing(monkeypatch): @@ -229,4 +210,4 @@ def test_get_projected_spend_over_limit_raises_when_today_missing(monkeypatch): monkeypatch.setattr("litellm.proxy.utils.date", _Broken) with pytest.raises(RuntimeError): - _get_projected_spend_over_limit(current_spend=1.0, soft_budget_limit=1.0) + get_projected_spend_over_limit(current_spend=1.0, soft_budget_limit=1.0) diff --git a/tests/unit/proxy/utils/helpers/test_premium_user_check.py b/tests/unit/proxy/utils/helpers/test_premium_user_check.py index 0a9539c6dc3..f78d634157e 100644 --- a/tests/unit/proxy/utils/helpers/test_premium_user_check.py +++ b/tests/unit/proxy/utils/helpers/test_premium_user_check.py @@ -1,7 +1,7 @@ import pytest from fastapi import HTTPException -from litellm.proxy.utils import _premium_user_check +from litellm.proxy.utils import premium_user_check def normalize(value): @@ -13,7 +13,7 @@ def test_premium_user_check_happy_path_no_raise_when_premium(monkeypatch): monkeypatch.setattr(ps, "premium_user", True, raising=False) summary = { - "result": _premium_user_check(), + "result": premium_user_check(), "premium_user": True, "raised": False, } @@ -29,7 +29,7 @@ def test_premium_user_check_happy_path_with_feature_no_raise(monkeypatch): monkeypatch.setattr(ps, "premium_user", True, raising=False) summary = { - "result": _premium_user_check(feature="model-routing"), + "result": premium_user_check(feature="model-routing"), "premium_user": True, "feature": "model-routing", } @@ -45,7 +45,7 @@ def test_premium_user_check_raises_when_not_premium(monkeypatch): monkeypatch.setattr(ps, "premium_user", False, raising=False) with pytest.raises(HTTPException) as exc_info: - _premium_user_check() + premium_user_check() snapshot = { "status_code": exc_info.value.status_code, "is_dict_detail": isinstance(exc_info.value.detail, dict), @@ -63,7 +63,7 @@ def test_premium_user_check_raises_with_feature_message(monkeypatch): monkeypatch.setattr(ps, "premium_user", False, raising=False) with pytest.raises(HTTPException) as exc_info: - _premium_user_check(feature="custom-callbacks") + premium_user_check(feature="custom-callbacks") error_msg = exc_info.value.detail["error"] snapshot = { "status_code": exc_info.value.status_code, diff --git a/tests/unit/proxy/utils/helpers/test_team_configs.py b/tests/unit/proxy/utils/helpers/test_team_configs.py index 185d4d26ff4..2a9b36b92b8 100644 --- a/tests/unit/proxy/utils/helpers/test_team_configs.py +++ b/tests/unit/proxy/utils/helpers/test_team_configs.py @@ -1,6 +1,6 @@ import pytest -from litellm.proxy.utils import _is_valid_team_configs +from litellm.proxy.utils import is_valid_team_configs def normalize(value): @@ -11,7 +11,7 @@ def test_is_valid_team_configs_happy_path_allowed_model_mutates_config(): team_config = {"models": ["gpt-4o", "gpt-4o-mini"], "max_budget": 100.0} request_data = {"model": "gpt-4o"} snapshot = { - "result": _is_valid_team_configs( + "result": is_valid_team_configs( team_id="team-1", team_config=team_config, request_data=request_data, @@ -30,7 +30,7 @@ def test_is_valid_team_configs_no_models_key_is_noop(): team_config = {"max_budget": 100.0, "tpm_limit": 1000} request_data = {"model": "anything"} snapshot = { - "result": _is_valid_team_configs( + "result": is_valid_team_configs( team_id="team-1", team_config=team_config, request_data=request_data, @@ -48,7 +48,7 @@ def test_is_valid_team_configs_no_models_key_is_noop(): def test_is_valid_team_configs_short_circuits_when_team_id_none(): team_config = {"models": ["only-this"]} snapshot = { - "result": _is_valid_team_configs( + "result": is_valid_team_configs( team_id=None, team_config=team_config, request_data={"model": "anything-else"}, @@ -67,7 +67,7 @@ def test_is_valid_team_configs_raises_on_model_not_in_team_models(): team_config = {"models": ["gpt-4o"]} request_data = {"model": "claude-haiku"} with pytest.raises(Exception, match='claude-haiku\\. Valid models for team are') as exc_info: - _is_valid_team_configs( + is_valid_team_configs( team_id="team-1", team_config=team_config, request_data=request_data, diff --git a/tests/unit/proxy/utils/helpers/test_url_helpers.py b/tests/unit/proxy/utils/helpers/test_url_helpers.py index 5f23c5fc20e..ca7520b87dc 100644 --- a/tests/unit/proxy/utils/helpers/test_url_helpers.py +++ b/tests/unit/proxy/utils/helpers/test_url_helpers.py @@ -1,9 +1,9 @@ import pytest from litellm.proxy.utils import ( - _get_docs_url, - _get_openapi_url, - _get_redoc_url, + get_docs_url, + get_openapi_url, + get_redoc_url, get_custom_url, get_proxy_base_url, get_server_root_path, @@ -33,7 +33,7 @@ def _clear_url_env(monkeypatch): def test_get_redoc_url_default(monkeypatch): _clear_url_env(monkeypatch) summary = { - "result": _get_redoc_url(), + "result": get_redoc_url(), "redoc_url_env": None, "no_redoc_env": None, } @@ -48,7 +48,7 @@ def test_get_redoc_url_custom_env(monkeypatch): _clear_url_env(monkeypatch) monkeypatch.setenv("REDOC_URL", "/custom-redoc") summary = { - "result": _get_redoc_url(), + "result": get_redoc_url(), "redoc_url_env": "/custom-redoc", "default_overridden": True, } @@ -62,13 +62,13 @@ def test_get_redoc_url_custom_env(monkeypatch): def test_get_redoc_url_disabled_returns_none_error_path(monkeypatch): _clear_url_env(monkeypatch) monkeypatch.setenv("NO_REDOC", "True") - assert _get_redoc_url() is None + assert get_redoc_url() is None def test_get_docs_url_default(monkeypatch): _clear_url_env(monkeypatch) summary = { - "result": _get_docs_url(), + "result": get_docs_url(), "no_docs": None, "docs_url": None, } @@ -83,7 +83,7 @@ def test_get_docs_url_custom_env(monkeypatch): _clear_url_env(monkeypatch) monkeypatch.setenv("DOCS_URL", "/api-docs") summary = { - "result": _get_docs_url(), + "result": get_docs_url(), "env": "/api-docs", "default_overridden": True, } @@ -97,13 +97,13 @@ def test_get_docs_url_custom_env(monkeypatch): def test_get_docs_url_disabled_returns_none_error_path(monkeypatch): _clear_url_env(monkeypatch) monkeypatch.setenv("NO_DOCS", "True") - assert _get_docs_url() is None + assert get_docs_url() is None def test_get_openapi_url_default(monkeypatch): _clear_url_env(monkeypatch) summary = { - "result": _get_openapi_url(), + "result": get_openapi_url(), "no_openapi": None, "openapi_url": None, } @@ -118,7 +118,7 @@ def test_get_openapi_url_custom_env(monkeypatch): _clear_url_env(monkeypatch) monkeypatch.setenv("OPENAPI_URL", "/api-schema") summary = { - "result": _get_openapi_url(), + "result": get_openapi_url(), "env": "/api-schema", "default_overridden": True, } @@ -132,7 +132,7 @@ def test_get_openapi_url_custom_env(monkeypatch): def test_get_openapi_url_disabled_returns_none_error_path(monkeypatch): _clear_url_env(monkeypatch) monkeypatch.setenv("NO_OPENAPI", "True") - assert _get_openapi_url() is None + assert get_openapi_url() is None @pytest.mark.parametrize( diff --git a/tests/unit/proxy/utils/prisma_and_spend/conftest.py b/tests/unit/proxy/utils/prisma_and_spend/conftest.py index e37a82a023b..aefb95516af 100644 --- a/tests/unit/proxy/utils/prisma_and_spend/conftest.py +++ b/tests/unit/proxy/utils/prisma_and_spend/conftest.py @@ -359,7 +359,7 @@ def proxy_logging_with_redis(fake_redis: FakeRedisList) -> MagicMock: proxy_logging.db_spend_update_writer = MagicMock() proxy_logging.db_spend_update_writer.db_update_spend_transaction_handler = AsyncMock() buffer = RedisUpdateBuffer(redis_cache=fake_redis) - buffer._should_commit_spend_updates_to_redis = MagicMock(return_value=True) + buffer.should_commit_spend_updates_to_redis = MagicMock(return_value=True) proxy_logging.db_spend_update_writer.redis_update_buffer = buffer return proxy_logging diff --git a/tests/unit/proxy/utils/prisma_and_spend/test_cache_user_row.py b/tests/unit/proxy/utils/prisma_and_spend/test_cache_user_row.py index d1270b60b19..a40e4422d49 100644 --- a/tests/unit/proxy/utils/prisma_and_spend/test_cache_user_row.py +++ b/tests/unit/proxy/utils/prisma_and_spend/test_cache_user_row.py @@ -12,7 +12,7 @@ from unittest.mock import AsyncMock, MagicMock import pytest -from litellm.proxy.utils import _cache_user_row +from litellm.proxy.utils import cache_user_row @pytest.mark.asyncio @@ -28,7 +28,7 @@ async def test_cache_user_row_caches_on_miss( db = MagicMock() db.get_data = AsyncMock(return_value=user_row) - result = await _cache_user_row("u1", mock_dual_cache, db) + result = await cache_user_row("u1", mock_dual_cache, db) cache_key = "u1_user_api_key_user_id" pinned = { "result": result, @@ -54,7 +54,7 @@ async def test_cache_user_row_skips_db_on_cache_hit( mock_dual_cache._store[cache_key] = "cached-blob" db = MagicMock() db.get_data = AsyncMock(return_value=None) - result = await _cache_user_row("u-hit", mock_dual_cache, db) + result = await cache_user_row("u-hit", mock_dual_cache, db) assert result is None assert db.get_data.await_count == 0 @@ -66,7 +66,7 @@ async def test_cache_user_row_skips_set_when_user_row_lacks_model_dump_json( user_row = SimpleNamespace(user_id="u2", spend=1.0) db = MagicMock() db.get_data = AsyncMock(return_value=user_row) - await _cache_user_row("u2", mock_dual_cache, db) + await cache_user_row("u2", mock_dual_cache, db) assert mock_dual_cache._store == {} assert mock_dual_cache.set_cache.call_count == 0 @@ -78,4 +78,4 @@ async def test_cache_user_row_propagates_db_error( db = MagicMock() db.get_data = AsyncMock(side_effect=RuntimeError("db down")) with pytest.raises(RuntimeError, match="db down"): - await _cache_user_row("u3", mock_dual_cache, db) + await cache_user_row("u3", mock_dual_cache, db) diff --git a/tests/unit/proxy/utils/prisma_and_spend/test_password_helpers.py b/tests/unit/proxy/utils/prisma_and_spend/test_password_helpers.py index 3c028473479..f241cd00cf5 100644 --- a/tests/unit/proxy/utils/prisma_and_spend/test_password_helpers.py +++ b/tests/unit/proxy/utils/prisma_and_spend/test_password_helpers.py @@ -21,7 +21,7 @@ from unittest.mock import AsyncMock, MagicMock import pytest from litellm.proxy.utils import ( - _hash_token_if_needed, + hash_token_if_needed, hash_password, hash_token, migrate_passwords_to_scrypt_async, @@ -116,9 +116,9 @@ def test_hash_token_if_needed_handles_sk_prefix() -> None: already_hashed = hashlib.sha256(plain.encode()).hexdigest() not_a_secret = "token-without-sk-prefix" actual = { - "sk_input_is_hashed": _hash_token_if_needed(plain) == already_hashed, - "non_sk_passthrough": _hash_token_if_needed(not_a_secret) == not_a_secret, - "double_hash_stable": _hash_token_if_needed(already_hashed) == already_hashed, + "sk_input_is_hashed": hash_token_if_needed(plain) == already_hashed, + "non_sk_passthrough": hash_token_if_needed(not_a_secret) == not_a_secret, + "double_hash_stable": hash_token_if_needed(already_hashed) == already_hashed, } assert actual == { "sk_input_is_hashed": True, @@ -129,7 +129,7 @@ def test_hash_token_if_needed_handles_sk_prefix() -> None: def test_hash_token_if_needed_error_on_non_string() -> None: with pytest.raises(AttributeError): - _hash_token_if_needed(None) # type: ignore[arg-type] + hash_token_if_needed(None) # pyright: ignore[reportArgumentType] # intentional invalid input checks the error # --------------------------------------------------------------------------- @@ -191,12 +191,10 @@ async def test_migrate_passwords_upgrades_only_plaintext_rows() -> None: result = await migrate_passwords_to_scrypt_async(pc) updated_user_ids = sorted( - call.kwargs["where"]["user_id"] - for call in pc.db.litellm_usertable.update.await_args_list + call.kwargs["where"]["user_id"] for call in pc.db.litellm_usertable.update.await_args_list ) new_password_prefixes = sorted( - call.kwargs["data"]["password"][:7] - for call in pc.db.litellm_usertable.update.await_args_list + call.kwargs["data"]["password"][:7] for call in pc.db.litellm_usertable.update.await_args_list ) outcome = { "message": result, @@ -216,8 +214,6 @@ async def test_migrate_passwords_upgrades_only_plaintext_rows() -> None: async def test_migrate_passwords_raises_on_db_failure() -> None: pc = MagicMock() pc.db = MagicMock() - pc.db.litellm_usertable.find_many = AsyncMock( - side_effect=RuntimeError("db unavailable") - ) + pc.db.litellm_usertable.find_many = AsyncMock(side_effect=RuntimeError("db unavailable")) with pytest.raises(RuntimeError, match="db unavailable"): await migrate_passwords_to_scrypt_async(pc) diff --git a/tests/unit/proxy/utils/prisma_and_spend/test_spend_functions.py b/tests/unit/proxy/utils/prisma_and_spend/test_spend_functions.py index 7c4f4e582be..fc7ae531c34 100644 --- a/tests/unit/proxy/utils/prisma_and_spend/test_spend_functions.py +++ b/tests/unit/proxy/utils/prisma_and_spend/test_spend_functions.py @@ -23,8 +23,8 @@ import pytest from litellm.constants import REDIS_SPEND_LOGS_BUFFER_KEY from litellm.proxy.utils import ( MAX_SPEND_LOG_DRAIN_ITERATIONS, - _monitor_spend_logs_queue, - _raise_failed_update_spend_exception, + monitor_spend_logs_queue, + raise_failed_update_spend_exception, drain_spend_logs_queue, recover_parked_spend_logs, update_daily_tag_spend, @@ -111,15 +111,15 @@ async def test_update_daily_tag_spend_redis_path_when_buffered( writer = MagicMock() proxy_logging.db_spend_update_writer = writer writer.redis_update_buffer = MagicMock() - writer.redis_update_buffer._should_commit_spend_updates_to_redis = MagicMock(return_value=True) - writer._commit_daily_tag_spend_to_db_with_redis = AsyncMock() - writer._commit_daily_tag_spend_to_db = AsyncMock() + writer.redis_update_buffer.should_commit_spend_updates_to_redis = MagicMock(return_value=True) + writer.commit_daily_tag_spend_to_db_with_redis = AsyncMock() + writer.commit_daily_tag_spend_to_db = AsyncMock() await update_daily_tag_spend(prisma_client=mock_prisma_client, proxy_logging_obj=proxy_logging) - redis_kwargs = writer._commit_daily_tag_spend_to_db_with_redis.await_args.kwargs + redis_kwargs = writer.commit_daily_tag_spend_to_db_with_redis.await_args.kwargs pinned = { - "redis_calls": writer._commit_daily_tag_spend_to_db_with_redis.await_count, - "direct_calls": writer._commit_daily_tag_spend_to_db.await_count, + "redis_calls": writer.commit_daily_tag_spend_to_db_with_redis.await_count, + "direct_calls": writer.commit_daily_tag_spend_to_db.await_count, "redis_kwargs_keys": sorted(redis_kwargs.keys()), "redis_n_retries": redis_kwargs["n_retry_times"], } @@ -139,13 +139,13 @@ async def test_update_daily_tag_spend_direct_path_when_no_redis( writer = MagicMock() proxy_logging.db_spend_update_writer = writer writer.redis_update_buffer = MagicMock() - writer.redis_update_buffer._should_commit_spend_updates_to_redis = MagicMock(return_value=False) - writer._commit_daily_tag_spend_to_db_with_redis = AsyncMock() - writer._commit_daily_tag_spend_to_db = AsyncMock() + writer.redis_update_buffer.should_commit_spend_updates_to_redis = MagicMock(return_value=False) + writer.commit_daily_tag_spend_to_db_with_redis = AsyncMock() + writer.commit_daily_tag_spend_to_db = AsyncMock() await update_daily_tag_spend(prisma_client=mock_prisma_client, proxy_logging_obj=proxy_logging) - assert writer._commit_daily_tag_spend_to_db.await_count == 1 - assert writer._commit_daily_tag_spend_to_db_with_redis.await_count == 0 + assert writer.commit_daily_tag_spend_to_db.await_count == 1 + assert writer.commit_daily_tag_spend_to_db_with_redis.await_count == 0 @pytest.mark.asyncio @@ -159,10 +159,10 @@ async def test_update_daily_tag_spend_logs_and_swallows_errors( proxy_logging = MagicMock() proxy_logging.db_spend_update_writer = MagicMock() proxy_logging.db_spend_update_writer.redis_update_buffer = MagicMock() - proxy_logging.db_spend_update_writer.redis_update_buffer._should_commit_spend_updates_to_redis = MagicMock( + proxy_logging.db_spend_update_writer.redis_update_buffer.should_commit_spend_updates_to_redis = MagicMock( return_value=False ) - proxy_logging.db_spend_update_writer._commit_daily_tag_spend_to_db = AsyncMock( + proxy_logging.db_spend_update_writer.commit_daily_tag_spend_to_db = AsyncMock( side_effect=RuntimeError("commit boom") ) await update_daily_tag_spend(prisma_client=mock_prisma_client, proxy_logging_obj=proxy_logging) @@ -473,7 +473,7 @@ async def test_monitor_spend_logs_queue_invokes_job_when_queue_nonempty( monkeypatch.setattr(utils_mod, "update_spend_logs_job", _fake_job) with pytest.raises(asyncio.CancelledError): - await _monitor_spend_logs_queue( + await monitor_spend_logs_queue( prisma_client=mock_prisma_client, db_writer_client=None, proxy_logging_obj=proxy_logging, @@ -510,7 +510,7 @@ async def test_monitor_spend_logs_queue_swallows_errors_and_backs_off( mock_prisma_client._spend_log_transactions_lock = bad_lock with pytest.raises(asyncio.CancelledError): - await _monitor_spend_logs_queue( + await monitor_spend_logs_queue( prisma_client=mock_prisma_client, db_writer_client=None, proxy_logging_obj=proxy_logging, @@ -543,7 +543,7 @@ async def test_monitor_spend_logs_queue_flushes_as_soon_as_one_is_requested( monkeypatch.setattr(utils_mod, "update_spend_logs_job", _fake_job) monitor: Final = asyncio.create_task( - _monitor_spend_logs_queue( + monitor_spend_logs_queue( prisma_client=mock_prisma_client, db_writer_client=None, proxy_logging_obj=MagicMock(), @@ -589,7 +589,7 @@ def test_monitor_spend_logs_queue_flush_survives_an_earlier_event_loop( mock_prisma_client.spend_log_transactions = [] monitor: Final = asyncio.create_task( - _monitor_spend_logs_queue( + monitor_spend_logs_queue( prisma_client=mock_prisma_client, db_writer_client=None, proxy_logging_obj=MagicMock(), @@ -641,7 +641,7 @@ async def test_flush_requested_before_the_monitor_starts_costs_the_row_nothing( assert mock_prisma_client.spend_log_flush_requested is None monitor: Final = asyncio.create_task( - _monitor_spend_logs_queue( + monitor_spend_logs_queue( prisma_client=mock_prisma_client, db_writer_client=None, proxy_logging_obj=MagicMock(), @@ -661,7 +661,7 @@ def test_raise_failed_update_spend_exception_emits_failure_handler() -> None: async def _runner() -> Any: try: - _raise_failed_update_spend_exception( + raise_failed_update_spend_exception( e=RuntimeError("boom"), start_time=0.0, proxy_logging_obj=proxy_logging, @@ -701,7 +701,7 @@ def test_raise_failed_update_spend_exception_raises_original_error() -> None: proxy_logging.failure_handler = AsyncMock() async def _runner() -> None: - _raise_failed_update_spend_exception( + raise_failed_update_spend_exception( e=ValueError("specific"), start_time=0.0, proxy_logging_obj=proxy_logging, @@ -921,7 +921,7 @@ async def test_monitor_spend_logs_queue_pulls_parked_rows_before_each_flush( monkeypatch.setattr(utils_mod, "_wait_for_spend_log_flush_request", _poll) with pytest.raises(asyncio.CancelledError): - await _monitor_spend_logs_queue( + await monitor_spend_logs_queue( prisma_client=mock_prisma_client, db_writer_client=None, proxy_logging_obj=proxy_logging_with_redis, diff --git a/tests/unit/proxy/utils/proxy_logging/test_pre_call_hook.py b/tests/unit/proxy/utils/proxy_logging/test_pre_call_hook.py index 2bf6a8d7e6c..b7fc8c6dbeb 100644 --- a/tests/unit/proxy/utils/proxy_logging/test_pre_call_hook.py +++ b/tests/unit/proxy/utils/proxy_logging/test_pre_call_hook.py @@ -454,8 +454,8 @@ def test_every_pre_call_customlogger_is_deliberately_classified(): "_PROXY_SensitiveDataRoutingHandler", "ResponsesIDSecurity", "SkillsInjectionHook", - "PROXY_LiteLLMManagedFiles", - "PROXY_LiteLLMManagedVectorStores", + "_PROXY_LiteLLMManagedFiles", + "_PROXY_LiteLLMManagedVectorStores", } from litellm.proxy.hooks import PROXY_HOOKS @@ -464,8 +464,8 @@ def test_every_pre_call_customlogger_is_deliberately_classified(): for name, cls in ( ("banned_keywords", _load("enterprise.enterprise_hooks.banned_keywords", "ENTERPRISE_BannedKeywords")), ("blocked_user_check", _load("enterprise.enterprise_hooks.blocked_user_list", "ENTERPRISE_BlockedUserList")), - ("detect_prompt_injection", _load("litellm.proxy.hooks.prompt_injection_detection", "_OPTIONAL_PromptInjectionDetection")), - ("azure_content_safety", _load("litellm.proxy.hooks.azure_content_safety", "_PROXY_AzureContentSafety")), + ("detect_prompt_injection", _load("litellm.proxy.hooks.prompt_injection_detection", "OPTIONAL_PromptInjectionDetection")), + ("azure_content_safety", _load("litellm.proxy.hooks.azure_content_safety", "PROXY_AzureContentSafety")), ): if cls is not None: registered[name] = cls diff --git a/tests/unit/proxy/vector_store_endpoints/test_vector_store_tenant_guard.py b/tests/unit/proxy/vector_store_endpoints/test_vector_store_tenant_guard.py index 268000517d3..981fbfcf3e5 100644 --- a/tests/unit/proxy/vector_store_endpoints/test_vector_store_tenant_guard.py +++ b/tests/unit/proxy/vector_store_endpoints/test_vector_store_tenant_guard.py @@ -37,7 +37,7 @@ async def test_vector_store_search_forces_path_id_over_body_id(): request = _mock_request() with ( patch( - "litellm.proxy.proxy_server._read_request_body", + "litellm.proxy.proxy_server.read_request_body", new=AsyncMock( return_value={ "vector_store_id": "vs_body_victim", @@ -85,7 +85,7 @@ async def test_vector_store_file_create_forces_path_id_over_body_id(): request = _mock_request() with ( patch( - "litellm.proxy.proxy_server._read_request_body", + "litellm.proxy.proxy_server.read_request_body", new=AsyncMock( return_value={ "vector_store_id": "vs_body_victim", @@ -122,18 +122,14 @@ async def test_vector_store_file_list_resolves_managed_ids_and_cursors(): captured_data = {} provider_file_id: Final = "file-list-owned" - managed_file_data: Final = ( - SpecialEnums.LITELLM_MANAGED_FILE_COMPLETE_STR.value.format( - "application/json", - "unified-file", - "managed-deployment", - provider_file_id, - "managed-deployment-id", - ) - ) - managed_file_id: Final = ( - base64.urlsafe_b64encode(managed_file_data.encode()).decode().rstrip("=") + managed_file_data: Final = SpecialEnums.LITELLM_MANAGED_FILE_COMPLETE_STR.value.format( + "application/json", + "unified-file", + "managed-deployment", + provider_file_id, + "managed-deployment-id", ) + managed_file_id: Final = base64.urlsafe_b64encode(managed_file_data.encode()).decode().rstrip("=") user_api_key_dict: Final = UserAPIKeyAuth(team_models=["team-openai"]) managed_file: Final[VectorStoreFileObject] = { "id": provider_file_id, @@ -171,9 +167,7 @@ async def test_vector_store_file_list_resolves_managed_ids_and_cursors(): "provider_resource_id,vs_provider_native;" "model_id,managed-deployment" ) - vector_store_id = ( - base64.urlsafe_b64encode(raw_vector_store_id.encode()).decode().rstrip("=") - ) + vector_store_id = base64.urlsafe_b64encode(raw_vector_store_id.encode()).decode().rstrip("=") request = _mock_request() request.method = "GET" @@ -221,9 +215,7 @@ async def test_vector_store_file_list_resolves_managed_ids_and_cursors(): assert captured_data["vector_store_id"] == "vs_provider_native" assert captured_data["api_key"] == "sk-managed-deployment" assert captured_data["model"] == "openai/managed-deployment" - llm_router.get_deployment_credentials_with_provider.assert_called_once_with( - model_id="managed-deployment" - ) + llm_router.get_deployment_credentials_with_provider.assert_called_once_with(model_id="managed-deployment") proxy_logging_obj.get_proxy_hook.assert_called_once_with("managed_files") resolver.assert_awaited_once_with( provider_file_ids=(provider_file_id,), @@ -247,7 +239,7 @@ async def test_vector_store_file_create_denies_other_team_path_store(): request = _mock_request() with ( patch( - "litellm.proxy.proxy_server._read_request_body", + "litellm.proxy.proxy_server.read_request_body", new=AsyncMock(return_value={"file_id": "file_123"}), ), patch.object(litellm, "vector_store_registry", mock_registry), @@ -282,7 +274,7 @@ async def test_rag_query_denies_nested_other_team_vector_store(): request = _mock_request() with ( patch( - "litellm.proxy.rag_endpoints.endpoints._read_request_body", + "litellm.proxy.rag_endpoints.endpoints.read_request_body", new=AsyncMock( return_value={ "model": "gpt-4o-mini", @@ -501,9 +493,7 @@ async def test_get_managed_vector_store_uses_shared_cache_helper_for_db_fallback new=cache_helper, ), ): - vector_store = await get_litellm_managed_vector_store( - vector_store_id="vs_cached" - ) + vector_store = await get_litellm_managed_vector_store(vector_store_id="vs_cached") assert vector_store is not None assert vector_store["vector_store_id"] == "vs_cached" @@ -518,9 +508,7 @@ async def test_get_managed_vector_store_fails_closed_on_lookup_error(): ) mock_registry = MagicMock() - mock_registry.get_litellm_managed_vector_store_from_registry.side_effect = ( - RuntimeError("registry unavailable") - ) + mock_registry.get_litellm_managed_vector_store_from_registry.side_effect = RuntimeError("registry unavailable") with patch.object(litellm, "vector_store_registry", mock_registry): with pytest.raises(HTTPException) as exc_info: @@ -624,8 +612,8 @@ async def test_azure_passthrough_denies_other_team_vector_store_index(): index_object.litellm_params.vector_store_name = "tenant-b-store" mock_index_registry = MagicMock() - mock_index_registry.is_vector_store_index.side_effect = ( - lambda vector_store_index_name: vector_store_index_name == "managed_index" + mock_index_registry.is_vector_store_index.side_effect = lambda vector_store_index_name: ( + vector_store_index_name == "managed_index" ) mock_index_registry.get_vector_store_index_by_name.return_value = index_object diff --git a/tests/unit/proxy/video_endpoints/test_endpoints.py b/tests/unit/proxy/video_endpoints/test_endpoints.py index c5996f95f54..5d84068fede 100644 --- a/tests/unit/proxy/video_endpoints/test_endpoints.py +++ b/tests/unit/proxy/video_endpoints/test_endpoints.py @@ -138,11 +138,11 @@ def harness(): stack.enter_context( patch.object( ProxyBaseLLMRequestProcessing, - "_handle_llm_api_exception", + "handle_llm_api_exception", handle_exc, ) ) - stack.enter_context(patch.object(endpoints, "_read_request_body", read_body)) + stack.enter_context(patch.object(endpoints, "read_request_body", read_body)) stack.enter_context( patch.object(endpoints, "batch_to_bytesio", batch_to_bytesio) ) diff --git a/tests/unit/realtime_api/test_main.py b/tests/unit/realtime_api/test_main.py index 25f116f5991..a692f38ec04 100644 --- a/tests/unit/realtime_api/test_main.py +++ b/tests/unit/realtime_api/test_main.py @@ -658,6 +658,47 @@ async def test_arealtime_drops_model_from_the_upstream_url_only_for_transcriptio assert connect.url == expected_backend_url +async def _vertex_health_check_connect_for(monkeypatch: pytest.MonkeyPatch, api_base: str | None) -> _CapturingConnect: + async def fake_token_resolver( + credentials: object, project_id: str | None, custom_llm_provider: str + ) -> tuple[str, str]: + return "access-token", project_id or "" + + monkeypatch.setattr(realtime_main, "vertex_access_token_resolver", fake_token_resolver) + connect: Final = _CapturingConnect() + with patch("websockets.connect", connect): + assert await realtime_main._realtime_health_check( + model="gemini-3.8-live", + custom_llm_provider="vertex_ai", + api_key=None, + api_base=api_base, + model_params={ + "vertex_project": "proj-1", + "vertex_credentials": "fake-credentials", + "vertex_location": "us", + }, + ) + return connect + + +@pytest.mark.asyncio +async def test_vertex_health_check_sends_no_tls_argument_to_a_plain_ws_api_base( + monkeypatch: pytest.MonkeyPatch, +) -> None: + connect: Final = await _vertex_health_check_connect_for(monkeypatch, "http://127.0.0.1:8080") + assert connect.url == "ws://127.0.0.1:8080/ws/google.cloud.aiplatform.v1.LlmBidiService/BidiGenerateContent" + assert connect.kwargs["ssl"] is None + + +@pytest.mark.asyncio +async def test_vertex_health_check_keeps_tls_for_the_multi_region_host(monkeypatch: pytest.MonkeyPatch) -> None: + connect: Final = await _vertex_health_check_connect_for(monkeypatch, None) + assert connect.url == ( + "wss://aiplatform.us.rep.googleapis.com/ws/google.cloud.aiplatform.v1.LlmBidiService/BidiGenerateContent" + ) + assert connect.kwargs["ssl"] is not None + + BLOCKED_PHRASE = "XSECRETBLOCKTESTPHRASEX" diff --git a/tests/unit/responses/litellm_completion_transformation/test_litellm_completion_responses.py b/tests/unit/responses/litellm_completion_transformation/test_litellm_completion_responses.py index 2090ff3c9fa..ac02354e4d9 100644 --- a/tests/unit/responses/litellm_completion_transformation/test_litellm_completion_responses.py +++ b/tests/unit/responses/litellm_completion_transformation/test_litellm_completion_responses.py @@ -3954,7 +3954,9 @@ class TestEnsureOutputItemContentPartAdded: iterator._tool_item_id_by_call_id = {} iterator._tool_call_id_by_index = {} iterator._ambiguous_tool_call_indexes = set() - iterator._next_tool_output_index = 1 + iterator._next_output_index = 0 + iterator._message_output_index = None + iterator._reasoning_output_index = None iterator._final_tool_events_queued = False iterator._custom_tool_names = set() iterator.responses_api_request = {} @@ -4456,7 +4458,7 @@ class TestEnsureOutputItemContentPartAdded: Choices( finish_reason=finish_reason, index=0, - message=Message(content="", role="assistant"), + message=Message(content="Partial answer", role="assistant"), ) ], usage=Usage(prompt_tokens=10, completion_tokens=1, total_tokens=11), diff --git a/tests/unit/responses/litellm_completion_transformation/test_streaming_iterator_transformation.py b/tests/unit/responses/litellm_completion_transformation/test_streaming_iterator_transformation.py index a42eb74b0fa..cb764f0c17b 100644 --- a/tests/unit/responses/litellm_completion_transformation/test_streaming_iterator_transformation.py +++ b/tests/unit/responses/litellm_completion_transformation/test_streaming_iterator_transformation.py @@ -132,14 +132,14 @@ def test_tool_call_delta_is_emitted_as_responses_events(): evt1 = iterator._transform_chat_completion_chunk_to_response_api_chunk(chunk) assert evt1 is not None assert evt1.type == ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED - assert evt1.output_index == 1 + assert evt1.output_index == 0 # The arguments are now chunked, so we get the first delta chunk evt2 = iterator._transform_chat_completion_chunk_to_response_api_chunk(chunk) assert evt2 is not None assert evt2.type == ResponsesAPIStreamEvents.FUNCTION_CALL_ARGUMENTS_DELTA assert evt2.item_id == "fc_call_1" - assert evt2.output_index == 1 + assert evt2.output_index == 0 # The delta will be a chunk of the arguments, not the full arguments assert len(evt2.delta) <= 10 # Chunks are max 10 characters @@ -391,7 +391,7 @@ def test_tool_calls_present_only_in_final_response_are_emitted_before_completed( # First common_done_event_logic call should yield tool events, not response.completed. evt1 = iterator.common_done_event_logic(sync_mode=True) assert evt1.type == ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED - assert evt1.output_index == 1 + assert evt1.output_index == 0 # Now delta events are emitted (arguments split into chunks) # Collect all delta events @@ -412,12 +412,12 @@ def test_tool_calls_present_only_in_final_response_are_emitted_before_completed( # The last event should be FUNCTION_CALL_ARGUMENTS_DONE assert evt.type == ResponsesAPIStreamEvents.FUNCTION_CALL_ARGUMENTS_DONE assert evt.item_id == "fc_call_2" - assert evt.output_index == 1 + assert evt.output_index == 0 assert evt.arguments == '{"y":2}' evt_final = iterator.common_done_event_logic(sync_mode=True) assert evt_final.type == ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE - assert evt_final.output_index == 1 + assert evt_final.output_index == 0 def test_tool_call_arguments_are_chunked_to_match_openai_behavior(): @@ -472,7 +472,7 @@ def test_tool_call_arguments_are_chunked_to_match_openai_behavior(): # First event should be OUTPUT_ITEM_ADDED assert evt is not None assert evt.type == ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED - assert evt.output_index == 1 + assert evt.output_index == 0 assert hasattr(evt, "__dict__") and "sequence_number" in evt.__dict__ # Collect all remaining delta events from the pending queue by creating empty chunks @@ -506,7 +506,7 @@ def test_tool_call_arguments_are_chunked_to_match_openai_behavior(): for evt in delta_events: assert len(evt.delta) <= 10 assert evt.item_id == "fc_call_test" - assert evt.output_index == 1 + assert evt.output_index == 0 assert hasattr(evt, "__dict__") and "sequence_number" in evt.__dict__ # Verify all deltas concatenated equal the original arguments @@ -978,6 +978,23 @@ def _reasoning_chunk(reasoning: str, finish_reason: str | None = None) -> ModelR ) +def _annotation_only_chunk() -> ModelResponseStream: + citation: Final = {"start_index": 0, "end_index": 2, "url": "https://example.com", "title": "Example"} + return ModelResponseStream( + id=CHAT_COMPLETION_ID, + created=1748575031, + model="claude-haiku-4-5", + object="chat.completion.chunk", + choices=[ + StreamingChoices( + index=0, + delta=Delta(role="assistant", annotations=[{"type": "url_citation", "url_citation": citation}]), + finish_reason=None, + ) + ], + ) + + def _signature_only_thinking_chunk(signature: str) -> ModelResponseStream: return ModelResponseStream( id=CHAT_COMPLETION_ID, @@ -1034,6 +1051,68 @@ async def test_tool_only_stream_emits_no_message_item_events(sync_mode: bool): assert any(getattr(event, "type", None) == ResponsesAPIStreamEvents.RESPONSE_COMPLETED for event in events) +@pytest.mark.parametrize("sync_mode", [True, False]) +@pytest.mark.parametrize( + "leading_chunk", + [pytest.param(_chunk(""), id="empty-text-delta"), pytest.param(_reasoning_chunk(""), id="empty-reasoning-delta")], +) +@pytest.mark.asyncio +async def test_empty_leading_delta_does_not_open_a_message_item_ahead_of_a_tool_call( + sync_mode: bool, leading_chunk: ModelResponseStream +): + iterator: Final = _build_iterator([leading_chunk, _tool_call_chunk(), _chunk("", finish_reason="tool_calls")]) + + events: Final = await _collect_events(iterator, sync_mode) + + item_events: Final = [ + (event.type, event.item.type) + for event in events + if getattr(event, "type", None) + in (ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED, ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE) + ] + assert item_events == [ + (ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED, "function_call"), + (ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE, "function_call"), + ] + completed: Final = [ + event for event in events if getattr(event, "type", None) == ResponsesAPIStreamEvents.RESPONSE_COMPLETED + ] + assert [item.type for item in completed[0].response.output] == ["function_call"] + + +@pytest.mark.parametrize("sync_mode", [True, False]) +@pytest.mark.asyncio +async def test_empty_leading_delta_still_opens_the_message_item_for_the_first_text_delta(sync_mode: bool): + iterator: Final = _build_iterator([_chunk(""), _chunk("Hi"), _chunk("", finish_reason="stop")]) + + events: Final = await _collect_events(iterator, sync_mode) + + added_types: Final = [ + event.item.type + for event in events + if getattr(event, "type", None) == ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED + ] + assert added_types == ["message"] + completed: Final = [ + event for event in events if getattr(event, "type", None) == ResponsesAPIStreamEvents.RESPONSE_COMPLETED + ] + assert completed[0].response.output_text == "Hi" + + +@pytest.mark.parametrize("sync_mode", [True, False]) +@pytest.mark.asyncio +async def test_annotation_only_leading_delta_opens_the_message_item_before_its_annotation(sync_mode: bool): + iterator: Final = _build_iterator([_annotation_only_chunk(), _chunk("Hi"), _chunk("", finish_reason="stop")]) + + events: Final = await _collect_events(iterator, sync_mode) + + event_types: Final = [getattr(event, "type", None) for event in events] + assert ResponsesAPIStreamEvents.OUTPUT_TEXT_ANNOTATION_ADDED in event_types + assert event_types.index(ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED) < event_types.index( + ResponsesAPIStreamEvents.OUTPUT_TEXT_ANNOTATION_ADDED + ) + + @pytest.mark.parametrize("sync_mode", [True, False]) @pytest.mark.asyncio async def test_signature_only_thinking_streams_a_replayable_reasoning_item(sync_mode: bool): @@ -1152,6 +1231,82 @@ async def test_tool_then_reasoning_then_text_gives_message_its_own_output_index( assert len(output_indexes_by_item_id) == len(set(output_indexes_by_item_id.values())) +@pytest.mark.parametrize("sync_mode", [True, False]) +@pytest.mark.asyncio +async def test_tool_only_stream_puts_the_function_call_at_output_index_zero(sync_mode: bool): + iterator: Final = _build_iterator([_tool_call_chunk(), _chunk("", finish_reason="tool_calls")]) + + events: Final = await _collect_events(iterator, sync_mode) + + indexed_events: Final = [event for event in events if hasattr(event, "output_index")] + assert indexed_events + assert {event.output_index for event in indexed_events} == {0} + completed: Final = next( + event for event in events if getattr(event, "type", None) == ResponsesAPIStreamEvents.RESPONSE_COMPLETED + ) + assert [item.type for item in completed.response.output] == ["function_call"] + + +@pytest.mark.parametrize("sync_mode", [True, False]) +@pytest.mark.asyncio +async def test_completed_output_follows_the_streamed_output_indexes(sync_mode: bool): + iterator: Final = _build_iterator( + [ + _reasoning_chunk("thinking"), + _tool_call_chunk(), + _chunk("Hello"), + _chunk("!", finish_reason="stop"), + ] + ) + + events: Final = await _collect_events(iterator, sync_mode) + + added: Final = [ + event for event in events if getattr(event, "type", None) == ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED + ] + assert [event.output_index for event in added] == list(range(len(added))) + completed: Final = next( + event for event in events if getattr(event, "type", None) == ResponsesAPIStreamEvents.RESPONSE_COMPLETED + ) + assert [item.type for item in completed.response.output] == [event.item.type for event in added] + assert [item.id for item in completed.response.output] == [event.item.id for event in added] + + +@pytest.mark.parametrize("sync_mode", [True, False]) +@pytest.mark.asyncio +async def test_reasoning_only_stream_lists_no_message_item_it_never_announced(sync_mode: bool): + iterator: Final = _build_iterator([_reasoning_chunk("thinking"), _chunk("", finish_reason="stop")]) + + events: Final = await _collect_events(iterator, sync_mode) + + added: Final = [ + event for event in events if getattr(event, "type", None) == ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED + ] + assert [event.item.type for event in added] == ["reasoning"] + completed: Final = next( + event for event in events if getattr(event, "type", None) == ResponsesAPIStreamEvents.RESPONSE_COMPLETED + ) + assert [item.type for item in completed.response.output] == ["reasoning"] + + +@pytest.mark.parametrize("sync_mode", [True, False]) +@pytest.mark.asyncio +async def test_empty_answer_stream_keeps_the_message_item_it_announced(sync_mode: bool): + iterator: Final = _build_iterator([_chunk("", finish_reason="stop")]) + + events: Final = await _collect_events(iterator, sync_mode) + + added: Final = [ + event for event in events if getattr(event, "type", None) == ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED + ] + assert [event.item.type for event in added] == ["message"] + completed: Final = next( + event for event in events if getattr(event, "type", None) == ResponsesAPIStreamEvents.RESPONSE_COMPLETED + ) + assert [item.type for item in completed.response.output] == ["message"] + assert [item.id for item in completed.response.output] == [added[0].item.id] + + @pytest.mark.parametrize("sync_mode", [True, False]) @pytest.mark.asyncio async def test_plain_text_stream_announces_exactly_one_message_item(sync_mode: bool): diff --git a/tests/unit/responses/mcp/test_litellm_proxy_mcp_handler.py b/tests/unit/responses/mcp/test_litellm_proxy_mcp_handler.py index d215ce292aa..f43e15cb342 100644 --- a/tests/unit/responses/mcp/test_litellm_proxy_mcp_handler.py +++ b/tests/unit/responses/mcp/test_litellm_proxy_mcp_handler.py @@ -51,7 +51,7 @@ def _setup_mcp_call_environment(monkeypatch: pytest.MonkeyPatch) -> AsyncMock: get_registry=MagicMock(return_value={}), call_tool=AsyncMock(return_value=_DummyMCPResult()), # Newer logging path calls this to enrich spend logs metadata - _get_mcp_server_from_tool_name=MagicMock(return_value=None), + get_mcp_server_from_tool_name=MagicMock(return_value=None), get_mcp_server_by_name=MagicMock(return_value=None), ) monkeypatch.setattr( @@ -315,7 +315,7 @@ async def test_execute_tool_calls_strips_prefix_when_alias_differs_from_server_n ) from litellm.proxy._experimental.mcp_server import mcp_server_manager as _msm - _msm.global_mcp_server_manager._get_mcp_server_from_tool_name = MagicMock(return_value=fake_server) + _msm.global_mcp_server_manager.get_mcp_server_from_tool_name = MagicMock(return_value=fake_server) tool_name = "my_deepwiki-read_wiki_structure" tool_calls = [ @@ -357,7 +357,7 @@ async def test_execute_tool_calls_reverse_maps_display_name(monkeypatch): ) from litellm.proxy._experimental.mcp_server import mcp_server_manager as _msm - _msm.global_mcp_server_manager._get_mcp_server_from_tool_name = MagicMock(return_value=colliding_server) + _msm.global_mcp_server_manager.get_mcp_server_from_tool_name = MagicMock(return_value=colliding_server) _msm.global_mcp_server_manager.get_mcp_server_by_name = MagicMock(return_value=fake_server) tool_name = "browse_repo_docs" @@ -510,7 +510,7 @@ async def test_execute_tool_calls_applies_post_call_hook_content(monkeypatch): catalog=types.SimpleNamespace(operation=nullcontext), get_registry=MagicMock(return_value={}), call_tool=AsyncMock(return_value=result), - _get_mcp_server_from_tool_name=MagicMock(return_value=None), + get_mcp_server_from_tool_name=MagicMock(return_value=None), get_mcp_server_by_name=MagicMock(return_value=None), ) monkeypatch.setattr( @@ -552,7 +552,7 @@ async def test_execute_tool_calls_returns_proxy_result_without_logging(monkeypat catalog=types.SimpleNamespace(operation=nullcontext), get_registry=MagicMock(return_value={}), call_tool=AsyncMock(return_value=result), - _get_mcp_server_from_tool_name=MagicMock(return_value=None), + get_mcp_server_from_tool_name=MagicMock(return_value=None), get_mcp_server_by_name=MagicMock(return_value=None), ) monkeypatch.setattr( @@ -586,7 +586,7 @@ async def test_execute_tool_calls_passes_logging_details_to_proxy_hook(monkeypat catalog=types.SimpleNamespace(operation=nullcontext), get_registry=MagicMock(return_value={}), call_tool=AsyncMock(return_value=result), - _get_mcp_server_from_tool_name=MagicMock(return_value=None), + get_mcp_server_from_tool_name=MagicMock(return_value=None), get_mcp_server_by_name=MagicMock(return_value=None), ) monkeypatch.setattr( @@ -622,7 +622,7 @@ async def test_execute_tool_calls_continues_when_post_call_logging_fails(monkeyp catalog=types.SimpleNamespace(operation=nullcontext), get_registry=MagicMock(return_value={}), call_tool=AsyncMock(return_value=result), - _get_mcp_server_from_tool_name=MagicMock(return_value=None), + get_mcp_server_from_tool_name=MagicMock(return_value=None), get_mcp_server_by_name=MagicMock(return_value=None), ) monkeypatch.setattr( @@ -1370,7 +1370,7 @@ async def test_bridge_listing_leaves_the_callers_catalog_unchanged( with ( patch.dict(manager.tool_name_to_mcp_server_name_mapping), patch.object(mcp_operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[server])), - patch.object(manager, "_create_mcp_client", AsyncMock(return_value=object())), + patch.object(manager, "create_mcp_client", AsyncMock(return_value=object())), patch.object(manager, "_fetch_tools_with_timeout", AsyncMock(return_value=upstream)), ): try: @@ -1444,7 +1444,7 @@ async def test_concurrent_bridge_calls_use_their_own_served_metadata(monkeypatch ] client: Final = AsyncMock() client.call_tool.return_value = CallToolResult(content=[TextContent(type="text", text="ok")]) - manager._create_mcp_client = AsyncMock(return_value=client) + manager.create_mcp_client = AsyncMock(return_value=client) manager._fetch_tools_with_timeout = AsyncMock(return_value=upstream) guardrail: Final = _BridgeMetadataGuardrail() logger: Final = ProxyLogging(user_api_key_cache=DualCache()) diff --git a/tests/unit/responses/mcp/test_mcp_streaming_iterator.py b/tests/unit/responses/mcp/test_mcp_streaming_iterator.py index 3982081706f..52e2adbe34b 100644 --- a/tests/unit/responses/mcp/test_mcp_streaming_iterator.py +++ b/tests/unit/responses/mcp/test_mcp_streaming_iterator.py @@ -81,7 +81,7 @@ def _mock_mcp_environment(monkeypatch: pytest.MonkeyPatch) -> AsyncMock: catalog=types.SimpleNamespace(operation=nullcontext), get_registry=MagicMock(return_value={}), call_tool=call_tool, - _get_mcp_server_from_tool_name=MagicMock(return_value=None), + get_mcp_server_from_tool_name=MagicMock(return_value=None), get_mcp_server_by_name=MagicMock(return_value=None), ) monkeypatch.setattr( diff --git a/tests/unit/test_private_usage_aliases.py b/tests/unit/test_private_usage_aliases.py index 412cae4eb0c..a5d74092d16 100644 --- a/tests/unit/test_private_usage_aliases.py +++ b/tests/unit/test_private_usage_aliases.py @@ -6,6 +6,1138 @@ from typing import Final, cast import pytest +PROXY_CLASS_NAME_ALIAS_CASES: Final = ( + ( + "litellm.proxy.hooks.dynamic_rate_limiter", + "_PROXY_DynamicRateLimitHandler", + "PROXY_DynamicRateLimitHandler", + ), + ( + "litellm.proxy.hooks.dynamic_rate_limiter_v3", + "_PROXY_DynamicRateLimitHandlerV3", + "PROXY_DynamicRateLimitHandlerV3", + ), + ( + "litellm.proxy.common_utils.config_sync_pubsub", + "_ConfigSyncPubSub", + "ConfigSyncPubSub", + ), + ( + "litellm.proxy.guardrails.guardrail_hooks.presidio", + "_OPTIONAL_PresidioPIIMasking", + "OPTIONAL_PresidioPIIMasking", + ), + ( + "litellm.proxy.hooks.prompt_injection_detection", + "_OPTIONAL_PromptInjectionDetection", + "OPTIONAL_PromptInjectionDetection", + ), + ( + "litellm.proxy.hooks.batch_redis_get", + "_PROXY_BatchRedisRequests", + "PROXY_BatchRedisRequests", + ), + ( + "litellm.proxy.hooks.azure_content_safety", + "_PROXY_AzureContentSafety", + "PROXY_AzureContentSafety", + ), + ( + "litellm.proxy.guardrails.guardrail_hooks.cisco_ai_defense.cisco_ai_defense_mcp", + "_CiscoAIDefenseMcpMixin", + "CiscoAIDefenseMcpMixin", + ), + ( + "litellm.proxy.hooks.cache_control_check", + "_PROXY_CacheControlCheck", + "PROXY_CacheControlCheck", + ), + ( + "litellm.proxy.hooks.max_budget_per_session_limiter", + "_PROXY_MaxBudgetPerSessionHandler", + "PROXY_MaxBudgetPerSessionHandler", + ), + ( + "litellm.proxy.hooks.max_iterations_limiter", + "_PROXY_MaxIterationsHandler", + "PROXY_MaxIterationsHandler", + ), + ( + "litellm.proxy.hooks.parallel_request_limiter", + "_PROXY_MaxParallelRequestsHandler", + "PROXY_MaxParallelRequestsHandler", + ), + ( + "litellm.proxy.hooks.parallel_request_limiter_v3", + "_PROXY_MaxParallelRequestsHandler_v3", + "PROXY_MaxParallelRequestsHandler_v3", + ), + ( + "litellm.proxy.hooks.sensitive_data_routing", + "_PROXY_SensitiveDataRoutingHandler", + "PROXY_SensitiveDataRoutingHandler", + ), + ( + "litellm.proxy.hooks.batch_rate_limiter", + "_PROXY_BatchRateLimiter", + "PROXY_BatchRateLimiter", + ), + ( + "litellm.proxy.hooks.model_max_budget_limiter", + "_PROXY_VirtualKeyModelMaxBudgetLimiter", + "PROXY_VirtualKeyModelMaxBudgetLimiter", + ), + ( + "litellm.proxy.hooks.proxy_track_cost_callback", + "_ProxyDBLogger", + "ProxyDBLogger", + ), + ( + "enterprise.litellm_enterprise.proxy.hooks.managed_files", + "_PROXY_LiteLLMManagedFiles", + "PROXY_LiteLLMManagedFiles", + ), + ( + "enterprise.litellm_enterprise.proxy.hooks.managed_vector_stores", + "_PROXY_LiteLLMManagedVectorStores", + "PROXY_LiteLLMManagedVectorStores", + ), +) + +PACKAGE_EXPORT_ALIAS_CASES: Final = ( + ("litellm.proxy.hooks", "_PROXY_CacheControlCheck", "PROXY_CacheControlCheck"), + ( + "litellm.proxy.hooks", + "_PROXY_MaxBudgetPerSessionHandler", + "PROXY_MaxBudgetPerSessionHandler", + ), + ("litellm.proxy.hooks", "_PROXY_MaxIterationsHandler", "PROXY_MaxIterationsHandler"), + ("litellm.proxy.hooks", "_PROXY_MaxParallelRequestsHandler", "PROXY_MaxParallelRequestsHandler"), + ( + "litellm.proxy.hooks", + "_PROXY_MaxParallelRequestsHandler_v3", + "PROXY_MaxParallelRequestsHandler_v3", + ), + ( + "litellm.proxy.hooks", + "_PROXY_SensitiveDataRoutingHandler", + "PROXY_SensitiveDataRoutingHandler", + ), + ( + "litellm.proxy.management_endpoints.policy_endpoints", + "_build_all_names_per_competitor", + "build_all_names_per_competitor", + ), + ( + "litellm.proxy.management_endpoints.policy_endpoints", + "_build_comparison_blocked_words", + "build_comparison_blocked_words", + ), + ( + "litellm.proxy.management_endpoints.policy_endpoints", + "_build_competitor_guardrail_definitions", + "build_competitor_guardrail_definitions", + ), + ( + "litellm.proxy.management_endpoints.policy_endpoints", + "_build_name_blocked_words", + "build_name_blocked_words", + ), + ( + "litellm.proxy.management_endpoints.policy_endpoints", + "_build_recommendation_blocked_words", + "build_recommendation_blocked_words", + ), + ( + "litellm.proxy.management_endpoints.policy_endpoints", + "_build_refinement_prompt", + "build_refinement_prompt", + ), + ( + "litellm.proxy.management_endpoints.policy_endpoints", + "_clean_competitor_line", + "clean_competitor_line", + ), + ( + "litellm.proxy.management_endpoints.policy_endpoints", + "_parse_variations_response", + "parse_variations_response", + ), +) + +MODULE_IMPORT_ALIAS_CASES: Final = ( + ( + "litellm.litellm_core_utils.custom_logger_registry", + "_PROXY_DynamicRateLimitHandler", + "PROXY_DynamicRateLimitHandler", + ), + ( + "litellm.litellm_core_utils.custom_logger_registry", + "_PROXY_DynamicRateLimitHandlerV3", + "PROXY_DynamicRateLimitHandlerV3", + ), + ( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp", + "_read_request_body", + "read_request_body", + ), + ( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp", + "_run_centralized_common_checks", + "run_centralized_common_checks", + ), + ( + "litellm.proxy._experimental.mcp_server.bridge_token_flow", + "_V2_GCM_PREFIX", + "V2_GCM_PREFIX", + ), + ( + "litellm.proxy._experimental.mcp_server.db", + "_get_salt_key", + "get_salt_key", + ), + ( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints", + "_bridge_mint_error_response", + "bridge_mint_error_response", + ), + ( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints", + "_extract_user_id_from_request", + "extract_user_id_from_request", + ), + ( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints", + "_finish_bridge_mint", + "finish_bridge_mint", + ), + ( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints", + "_prepare_bridge_mint", + "prepare_bridge_mint", + ), + ( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints", + "_prepare_bridge_refresh", + "prepare_bridge_refresh", + ), + ( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints", + "_read_request_body", + "read_request_body", + ), + ( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints", + "_reload_active_user_by_id", + "reload_active_user_by_id", + ), + ( + "litellm.proxy._experimental.mcp_server.mcp_server_manager", + "_is_mcp_admitted_user_subject", + "is_mcp_admitted_user_subject", + ), + ( + "litellm.proxy._experimental.mcp_server.mcp_server_manager", + "_redact_mcp_resource_url", + "redact_mcp_resource_url", + ), + ( + "litellm.proxy._experimental.mcp_server.oauth2_flow_backfill", + "_decode_oauth_payload", + "decode_oauth_payload", + ), + ( + "litellm.proxy._experimental.mcp_server.operations", + "_caller_authorization_fans_out", + "caller_authorization_fans_out", + ), + ( + "litellm.proxy._experimental.mcp_server.operations", + "_client_forwarded_authorization_headers", + "client_forwarded_authorization_headers", + ), + ( + "litellm.proxy._experimental.mcp_server.operations", + "_redact_mcp_resource_url", + "redact_mcp_resource_url", + ), + ( + "litellm.proxy._experimental.mcp_server.operations", + "_request_auth_header", + "request_auth_header", + ), + ( + "litellm.proxy._experimental.mcp_server.operations", + "_request_extra_headers", + "request_extra_headers", + ), + ( + "litellm.proxy._experimental.mcp_server.operations", + "_request_resolved_auth_headers", + "request_resolved_auth_headers", + ), + ( + "litellm.proxy._experimental.mcp_server.operations", + "_resolve_openapi_tool_auth", + "resolve_openapi_tool_auth", + ), + ( + "litellm.proxy._experimental.mcp_server.operations", + "_should_strip_caller_authorization", + "should_strip_caller_authorization", + ), + ( + "litellm.proxy._experimental.mcp_server.rest_endpoints", + "_UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES", + "UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES", + ), + ( + "litellm.proxy._experimental.mcp_server.rest_endpoints", + "_apply_toolset_scope", + "apply_toolset_scope", + ), + ( + "litellm.proxy._experimental.mcp_server.rest_endpoints", + "_inherit_credentials_from_existing_server", + "inherit_credentials_from_existing_server", + ), + ( + "litellm.proxy._experimental.mcp_server.rest_endpoints", + "_redact_mcp_resource_url", + "redact_mcp_resource_url", + ), + ( + "litellm.proxy._experimental.mcp_server.rest_endpoints", + "_safe_get_request_headers", + "safe_get_request_headers", + ), + ( + "litellm.proxy._experimental.mcp_server.server", + "_is_mcp_admitted_user_subject", + "is_mcp_admitted_user_subject", + ), + ( + "litellm.proxy._experimental.mcp_server.server", + "_mcp_active_toolset_id", + "mcp_active_toolset_id", + ), + ( + "litellm.proxy._experimental.mcp_server.server", + "_mcp_gateway_initialize_instructions", + "mcp_gateway_initialize_instructions", + ), + ( + "litellm.proxy._experimental.mcp_server.server", + "_mcp_gateway_server_name", + "mcp_gateway_server_name", + ), + ( + "litellm.proxy.anthropic_endpoints.endpoints", + "_read_request_body", + "read_request_body", + ), + ( + "litellm.proxy.anthropic_endpoints.gateway_endpoints", + "_safe_set_request_parsed_body", + "safe_set_request_parsed_body", + ), + ( + "litellm.proxy.auth.auth_checks", + "_safe_get_request_headers", + "safe_get_request_headers", + ), + ( + "litellm.proxy.auth.auth_checks", + "_safe_get_request_query_params", + "safe_get_request_query_params", + ), + ( + "litellm.proxy.auth.auth_exception_handler", + "_get_request_ip_address", + "get_request_ip_address", + ), + ( + "litellm.proxy.auth.fallback_budget", + "_is_model_cost_zero", + "is_model_cost_zero", + ), + ( + "litellm.proxy.auth.ip_address_utils", + "_get_request_ip_address", + "get_request_ip_address", + ), + ( + "litellm.proxy.auth.resolvers.store", + "_cache_key_object", + "cache_key_object", + ), + ( + "litellm.proxy.auth.resolvers.store", + "_copy_user_api_key_auth_for_cache", + "copy_user_api_key_auth_for_cache", + ), + ( + "litellm.proxy.auth.resolvers.store", + "_fetch_key_object_from_db_with_reconnect", + "fetch_key_object_from_db_with_reconnect", + ), + ( + "litellm.proxy.auth.route_checks", + "_user_is_org_admin", + "user_is_org_admin", + ), + ( + "litellm.proxy.auth.user_api_key_auth", + "_cache_key_object", + "cache_key_object", + ), + ( + "litellm.proxy.auth.user_api_key_auth", + "_can_object_call_model", + "can_object_call_model", + ), + ( + "litellm.proxy.auth.user_api_key_auth", + "_check_end_user_budget", + "check_end_user_budget", + ), + ( + "litellm.proxy.auth.user_api_key_auth", + "_delete_cache_key_object", + "delete_cache_key_object", + ), + ( + "litellm.proxy.auth.user_api_key_auth", + "_get_user_role", + "get_user_role", + ), + ( + "litellm.proxy.auth.user_api_key_auth", + "_is_model_cost_zero", + "is_model_cost_zero", + ), + ( + "litellm.proxy.auth.user_api_key_auth", + "_is_user_proxy_admin", + "is_user_proxy_admin", + ), + ( + "litellm.proxy.auth.user_api_key_auth", + "_read_request_body", + "read_request_body", + ), + ( + "litellm.proxy.auth.user_api_key_auth", + "_realtime_request_body", + "realtime_request_body", + ), + ( + "litellm.proxy.auth.user_api_key_auth", + "_safe_get_request_headers", + "safe_get_request_headers", + ), + ( + "litellm.proxy.auth.user_api_key_auth", + "_safe_get_request_query_params", + "safe_get_request_query_params", + ), + ( + "litellm.proxy.auth.user_api_key_auth", + "_safe_set_request_parsed_body", + "safe_set_request_parsed_body", + ), + ( + "litellm.proxy.auth.user_api_key_auth", + "_team_member_max_budget_alert_check", + "team_member_max_budget_alert_check", + ), + ( + "litellm.proxy.auth.user_api_key_auth", + "_virtual_key_max_budget_alert_check", + "virtual_key_max_budget_alert_check", + ), + ( + "litellm.proxy.auth.user_api_key_auth", + "_virtual_key_max_budget_check", + "virtual_key_max_budget_check", + ), + ( + "litellm.proxy.auth.user_api_key_auth", + "_virtual_key_soft_budget_check", + "virtual_key_soft_budget_check", + ), + ( + "litellm.proxy.batches_endpoints.endpoints", + "_is_base64_encoded_unified_file_id", + "is_base64_encoded_unified_file_id", + ), + ( + "litellm.proxy.batches_endpoints.endpoints", + "_read_request_body", + "read_request_body", + ), + ( + "litellm.proxy.common_request_processing", + "_check_and_merge_model_level_guardrails", + "check_and_merge_model_level_guardrails", + ), + ( + "litellm.proxy.common_utils.auth_cache_invalidation_pubsub", + "_ConfigSyncPubSub", + "ConfigSyncPubSub", + ), + ( + "litellm.proxy.common_utils.auth_cache_invalidation_pubsub", + "_pubsub_capable_client", + "pubsub_capable_client", + ), + ( + "litellm.proxy.common_utils.key_rotation_manager", + "_calculate_key_rotation_time", + "calculate_key_rotation_time", + ), + ( + "litellm.proxy.common_utils.openai_endpoint_utils", + "_read_request_body", + "read_request_body", + ), + ( + "litellm.proxy.container_endpoints.endpoints", + "_read_request_body", + "read_request_body", + ), + ( + "litellm.proxy.custom_hooks.custom_ui_sso_hook", + "_safe_get_request_headers", + "safe_get_request_headers", + ), + ( + "litellm.proxy.db.db_span", + "_is_exception_related_to_db", + "is_exception_related_to_db", + ), + ( + "litellm.proxy.fine_tuning_endpoints.endpoints", + "_is_base64_encoded_unified_file_id", + "is_base64_encoded_unified_file_id", + ), + ( + "litellm.proxy.google_endpoints.agents_endpoints", + "_read_request_body", + "read_request_body", + ), + ( + "litellm.proxy.google_endpoints.agents_endpoints", + "_safe_get_request_query_params", + "safe_get_request_query_params", + ), + ( + "litellm.proxy.google_endpoints.endpoints", + "_read_request_body", + "read_request_body", + ), + ( + "litellm.proxy.guardrails.guardrail_endpoints", + "_user_has_admin_view", + "user_api_key_has_admin_view", + ), + ( + "litellm.proxy.guardrails.guardrail_hooks.azure.prompt_shield", + "_RESPONSES_API_CALL_TYPES", + "RESPONSES_API_CALL_TYPES", + ), + ( + "litellm.proxy.guardrails.guardrail_hooks.azure.text_moderation", + "_RESPONSES_API_CALL_TYPES", + "RESPONSES_API_CALL_TYPES", + ), + ( + "litellm.proxy.guardrails.guardrail_hooks.cisco_ai_defense.cisco_ai_defense", + "_CiscoAIDefenseMcpMixin", + "CiscoAIDefenseMcpMixin", + ), + ( + "litellm.proxy.guardrails.guardrail_hooks.litellm_content_filter.competitor_intent.airline", + "_compile_marker", + "compile_marker", + ), + ( + "litellm.proxy.guardrails.guardrail_hooks.litellm_content_filter.competitor_intent.airline", + "_count_signals", + "count_signals", + ), + ( + "litellm.proxy.guardrails.guardrail_hooks.litellm_content_filter.competitor_intent.airline", + "_word_boundary_match", + "word_boundary_match", + ), + ( + "litellm.proxy.guardrails.guardrail_registry", + "_OPTIONAL_PresidioPIIMasking", + "OPTIONAL_PresidioPIIMasking", + ), + ( + "litellm.proxy.health_endpoints._health_endpoints", + "_clean_endpoint_data", + "clean_endpoint_data", + ), + ( + "litellm.proxy.health_endpoints._health_endpoints", + "_update_litellm_params_for_health_check", + "update_litellm_params_for_health_check", + ), + ( + "litellm.proxy.hooks.dynamic_rate_limiter_v3", + "_PROXY_MaxParallelRequestsHandler_v3", + "PROXY_MaxParallelRequestsHandler_v3", + ), + ( + "litellm.proxy.hooks.key_management_event_hooks", + "_hash_token_if_needed", + "hash_token_if_needed", + ), + ( + "litellm.proxy.hooks.proxy_track_cost_callback", + "_sanitize_error_information_for_spend_logs", + "sanitize_error_information_for_spend_logs", + ), + ( + "litellm.proxy.litellm_pre_call_utils", + "_safe_get_request_headers", + "safe_get_request_headers", + ), + ( + "litellm.proxy.management_endpoints.access_group_endpoints", + "_cache_access_object", + "cache_access_object", + ), + ( + "litellm.proxy.management_endpoints.access_group_endpoints", + "_cache_key_object", + "cache_key_object", + ), + ( + "litellm.proxy.management_endpoints.access_group_endpoints", + "_cache_team_object", + "cache_team_object", + ), + ( + "litellm.proxy.management_endpoints.access_group_endpoints", + "_get_team_object_from_cache", + "get_team_object_from_cache", + ), + ( + "litellm.proxy.management_endpoints.auto_router_endpoints", + "_virtual_key_max_budget_check", + "virtual_key_max_budget_check", + ), + ( + "litellm.proxy.management_endpoints.budget_management_endpoints", + "_user_has_admin_view", + "user_api_key_has_admin_view", + ), + ( + "litellm.proxy.management_endpoints.common_utils", + "_premium_user_check", + "premium_user_check", + ), + ( + "litellm.proxy.management_endpoints.credential_migration", + "_ALGO_AES_GCM", + "ALGO_AES_GCM", + ), + ( + "litellm.proxy.management_endpoints.credential_migration", + "_ENCRYPTION_ALGORITHM_SETTING", + "ENCRYPTION_ALGORITHM_SETTING", + ), + ( + "litellm.proxy.management_endpoints.credential_migration", + "_V2_GCM_PREFIX", + "V2_GCM_PREFIX", + ), + ( + "litellm.proxy.management_endpoints.credential_migration", + "_get_salt_key", + "get_salt_key", + ), + ( + "litellm.proxy.management_endpoints.customer_endpoints", + "_set_object_permission", + "set_object_permission", + ), + ( + "litellm.proxy.management_endpoints.internal_user_endpoints", + "_check_permissions_caller_permission", + "check_permissions_caller_permission", + ), + ( + "litellm.proxy.management_endpoints.internal_user_endpoints", + "_set_object_permission", + "set_object_permission", + ), + ( + "litellm.proxy.management_endpoints.internal_user_endpoints", + "_user_has_admin_view", + "user_api_key_has_admin_view", + ), + ( + "litellm.proxy.management_endpoints.jwt_key_mapping_endpoints", + "_user_has_admin_view", + "user_api_key_has_admin_view", + ), + ( + "litellm.proxy.management_endpoints.key_management_endpoints", + "_add_model_to_db", + "add_model_to_db", + ), + ( + "litellm.proxy.management_endpoints.key_management_endpoints", + "_check_disable_global_guardrails_caller_permission", + "check_disable_global_guardrails_caller_permission", + ), + ( + "litellm.proxy.management_endpoints.key_management_endpoints", + "_check_passthrough_routes_caller_permission", + "check_passthrough_routes_caller_permission", + ), + ( + "litellm.proxy.management_endpoints.key_management_endpoints", + "_delete_cache_key_object", + "delete_cache_key_object", + ), + ( + "litellm.proxy.management_endpoints.key_management_endpoints", + "_hash_token_if_needed", + "hash_token_if_needed", + ), + ( + "litellm.proxy.management_endpoints.key_management_endpoints", + "_is_master_key", + "is_master_key", + ), + ( + "litellm.proxy.management_endpoints.key_management_endpoints", + "_set_object_metadata_field", + "set_object_metadata_field", + ), + ( + "litellm.proxy.management_endpoints.key_management_endpoints", + "_set_object_permission", + "set_object_permission", + ), + ( + "litellm.proxy.management_endpoints.key_management_endpoints", + "_team_member_has_permission", + "team_member_has_permission", + ), + ( + "litellm.proxy.management_endpoints.key_management_endpoints", + "_user_has_admin_view", + "user_api_key_has_admin_view", + ), + ( + "litellm.proxy.management_endpoints.mcp_management_endpoints", + "_raise_if_not_oauth2", + "raise_if_not_oauth2", + ), + ( + "litellm.proxy.management_endpoints.mcp_management_endpoints", + "_read_request_body", + "read_request_body", + ), + ( + "litellm.proxy.management_endpoints.mcp_management_endpoints", + "_user_api_key_auth_builder", + "user_api_key_auth_builder", + ), + ( + "litellm.proxy.management_endpoints.mcp_management_endpoints", + "_user_has_admin_view", + "user_api_key_has_admin_view", + ), + ( + "litellm.proxy.management_endpoints.model_management_endpoints", + "_refresh_cached_team", + "refresh_cached_team", + ), + ( + "litellm.proxy.management_endpoints.organization_endpoints", + "_set_object_metadata_field", + "set_object_metadata_field", + ), + ( + "litellm.proxy.management_endpoints.organization_endpoints", + "_user_has_admin_view", + "user_api_key_has_admin_view", + ), + ( + "litellm.proxy.management_endpoints.prompt_cache_prediction", + "_PROXY_MaxParallelRequestsHandler_v3", + "PROXY_MaxParallelRequestsHandler_v3", + ), + ( + "litellm.proxy.management_endpoints.prompt_cache_prediction", + "_read_request_body", + "read_request_body", + ), + ( + "litellm.proxy.management_endpoints.scim.scim_v2", + "_delete_cache_key_object", + "delete_cache_key_object", + ), + ( + "litellm.proxy.management_endpoints.scim.scim_v2", + "_premium_user_check", + "premium_user_check", + ), + ( + "litellm.proxy.management_endpoints.scim.scim_v2", + "_safe_get_request_headers", + "safe_get_request_headers", + ), + ( + "litellm.proxy.management_endpoints.session_endpoints", + "_persist_deleted_verification_tokens", + "persist_deleted_verification_tokens", + ), + ( + "litellm.proxy.management_endpoints.team_callback_endpoints", + "_CALLBACK_VAR_ENCRYPTED_PREFIX", + "CALLBACK_VAR_ENCRYPTED_PREFIX", + ), + ( + "litellm.proxy.management_endpoints.team_callback_endpoints", + "_get_validated_callback_metadata", + "get_validated_callback_metadata", + ), + ( + "litellm.proxy.management_endpoints.team_callback_endpoints", + "_refresh_cached_team", + "refresh_cached_team", + ), + ( + "litellm.proxy.management_endpoints.team_endpoints", + "_cache_team_object", + "cache_team_object", + ), + ( + "litellm.proxy.management_endpoints.team_endpoints", + "_check_disable_global_guardrails_caller_permission", + "check_disable_global_guardrails_caller_permission", + ), + ( + "litellm.proxy.management_endpoints.team_endpoints", + "_check_passthrough_routes_caller_permission", + "check_passthrough_routes_caller_permission", + ), + ( + "litellm.proxy.management_endpoints.team_endpoints", + "_set_object_metadata_field", + "set_object_metadata_field", + ), + ( + "litellm.proxy.management_endpoints.team_endpoints", + "_set_object_permission", + "set_object_permission", + ), + ( + "litellm.proxy.management_endpoints.team_endpoints", + "_team_member_has_permission", + "team_member_has_permission", + ), + ( + "litellm.proxy.management_endpoints.team_endpoints", + "_update_metadata_fields", + "update_metadata_fields", + ), + ( + "litellm.proxy.management_endpoints.team_endpoints", + "_upsert_budget_and_membership", + "upsert_budget_and_membership", + ), + ( + "litellm.proxy.management_endpoints.team_endpoints", + "_user_has_admin_view", + "user_api_key_has_admin_view", + ), + ( + "litellm.proxy.management_endpoints.ui_sso", + "_get_request_ip_address", + "get_request_ip_address", + ), + ( + "litellm.proxy.management_helpers.access_group_key_sync", + "_delete_cache_access_object", + "delete_cache_access_object", + ), + ( + "litellm.proxy.management_helpers.access_group_team_sync", + "_delete_cache_access_object", + "delete_cache_access_object", + ), + ( + "litellm.proxy.management_helpers.auto_router_permissions", + "_check_team_member_model_access", + "check_team_member_model_access", + ), + ( + "litellm.proxy.management_helpers.bulk_team_member_budgets", + "_upsert_budget_and_membership", + "upsert_budget_and_membership", + ), + ( + "litellm.proxy.management_helpers.bulk_user_creation", + "_check_permissions_caller_permission", + "check_permissions_caller_permission", + ), + ( + "litellm.proxy.management_helpers.bulk_user_creation", + "_set_object_permission", + "set_object_permission", + ), + ( + "litellm.proxy.management_helpers.bulk_user_deletion", + "_persist_deleted_verification_tokens", + "persist_deleted_verification_tokens", + ), + ( + "litellm.proxy.management_helpers.utils", + "_read_request_body", + "read_request_body", + ), + ( + "litellm.proxy.openai_files_endpoints.files_endpoints", + "_is_base64_encoded_unified_file_id", + "is_base64_encoded_unified_file_id", + ), + ( + "litellm.proxy.openai_files_endpoints.files_endpoints", + "_read_request_body", + "read_request_body", + ), + ( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints", + "_get_bearer_token", + "get_bearer_token", + ), + ( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints", + "_read_request_body", + "read_request_body", + ), + ( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints", + "_safe_get_request_headers", + "safe_get_request_headers", + ), + ( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints", + "_safe_set_request_parsed_body", + "safe_set_request_parsed_body", + ), + ( + "litellm.proxy.pass_through_endpoints.pass_through_endpoints", + "_get_dynamic_logging_metadata", + "get_dynamic_logging_metadata", + ), + ( + "litellm.proxy.pass_through_endpoints.pass_through_endpoints", + "_read_request_body", + "read_request_body", + ), + ( + "litellm.proxy.pass_through_endpoints.pass_through_endpoints", + "_safe_get_request_headers", + "safe_get_request_headers", + ), + ( + "litellm.proxy.proxy_server", + "_OPTIONAL_PromptInjectionDetection", + "OPTIONAL_PromptInjectionDetection", + ), + ( + "litellm.proxy.proxy_server", + "_PROXY_VirtualKeyModelMaxBudgetLimiter", + "PROXY_VirtualKeyModelMaxBudgetLimiter", + ), + ( + "litellm.proxy.proxy_server", + "_ProxyDBLogger", + "ProxyDBLogger", + ), + ( + "litellm.proxy.proxy_server", + "_add_model_to_db", + "add_model_to_db", + ), + ( + "litellm.proxy.proxy_server", + "_add_team_model_to_db", + "add_team_model_to_db", + ), + ( + "litellm.proxy.proxy_server", + "_cache_user_row", + "cache_user_row", + ), + ( + "litellm.proxy.proxy_server", + "_deduplicate_litellm_router_models", + "deduplicate_litellm_router_models", + ), + ( + "litellm.proxy.proxy_server", + "_fetch_global_spend_with_event_coordination", + "fetch_global_spend_with_event_coordination", + ), + ( + "litellm.proxy.proxy_server", + "_get_docs_url", + "get_docs_url", + ), + ( + "litellm.proxy.proxy_server", + "_get_openapi_url", + "get_openapi_url", + ), + ( + "litellm.proxy.proxy_server", + "_get_projected_spend_over_limit", + "get_projected_spend_over_limit", + ), + ( + "litellm.proxy.proxy_server", + "_get_redoc_url", + "get_redoc_url", + ), + ( + "litellm.proxy.proxy_server", + "_is_azure_model_router_request", + "is_azure_model_router_request", + ), + ( + "litellm.proxy.proxy_server", + "_is_projected_spend_over_limit", + "is_projected_spend_over_limit", + ), + ( + "litellm.proxy.proxy_server", + "_is_valid_team_configs", + "is_valid_team_configs", + ), + ( + "litellm.proxy.proxy_server", + "_read_request_body", + "read_request_body", + ), + ( + "litellm.proxy.proxy_server", + "_realtime_request_body", + "realtime_request_body", + ), + ( + "litellm.proxy.proxy_server", + "_safe_get_request_headers", + "safe_get_request_headers", + ), + ( + "litellm.proxy.proxy_server", + "_should_return_raw_model_name", + "should_return_raw_model_name", + ), + ( + "litellm.proxy.proxy_server", + "_user_has_admin_privileges", + "user_has_admin_privileges", + ), + ( + "litellm.proxy.proxy_server", + "_user_has_admin_view", + "user_api_key_has_admin_view", + ), + ( + "litellm.proxy.rag_endpoints.endpoints", + "_read_request_body", + "read_request_body", + ), + ( + "litellm.proxy.rag_endpoints.endpoints", + "_safe_get_request_headers", + "safe_get_request_headers", + ), + ( + "litellm.proxy.realtime_endpoints.endpoints", + "_read_request_body", + "read_request_body", + ), + ( + "litellm.proxy.response_api_endpoints.endpoints", + "_read_request_body", + "read_request_body", + ), + ( + "litellm.proxy.response_api_endpoints.endpoints", + "_safe_set_request_parsed_body", + "safe_set_request_parsed_body", + ), + ( + "litellm.proxy.search_endpoints.search_tool_registry", + "_get_salt_key", + "get_salt_key", + ), + ( + "litellm.proxy.spend_tracking.cloudzero_endpoints", + "_user_has_admin_view", + "user_api_key_has_admin_view", + ), + ( + "litellm.proxy.spend_tracking.vantage_endpoints", + "_user_has_admin_view", + "user_api_key_has_admin_view", + ), + ( + "litellm.proxy.utils", + "_PROXY_CacheControlCheck", + "PROXY_CacheControlCheck", + ), + ( + "litellm.proxy.utils", + "_PROXY_MaxParallelRequestsHandler", + "PROXY_MaxParallelRequestsHandler", + ), + ( + "litellm.proxy.utils", + "_PROXY_MaxParallelRequestsHandler_v3", + "PROXY_MaxParallelRequestsHandler_v3", + ), + ( + "litellm.proxy.utils", + "_PROXY_SensitiveDataRoutingHandler", + "PROXY_SensitiveDataRoutingHandler", + ), + ( + "litellm.proxy.utils", + "_is_exception_related_to_db", + "is_exception_related_to_db", + ), + ( + "litellm.proxy.vertex_ai_endpoints.langfuse_endpoints", + "_get_dynamic_logging_metadata", + "get_dynamic_logging_metadata", + ), + ( + "litellm.proxy.vertex_ai_endpoints.langfuse_endpoints", + "_safe_get_request_headers", + "safe_get_request_headers", + ), + ( + "litellm.proxy.video_endpoints.endpoints", + "_read_request_body", + "read_request_body", + ), +) + ALIAS_CASES: Final = ( ( "enterprise.enterprise_hooks.banned_keywords", @@ -1176,6 +2308,1843 @@ LLMS_ALIAS_CASES: Final = ( ("litellm.llms.watsonx.common_utils", "", "_generate_watsonx_token", "generate_watsonx_token", False), ("litellm.llms.watsonx.common_utils", "", "_get_api_params", "get_api_params", False), ) +PROXY_ALIAS_CASES: Final = ( + ( + 'litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp', + '', + '_is_mcp_admitted_user_subject', + 'is_mcp_admitted_user_subject', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp', + 'MCPRequestHandler', + '_get_mcp_auth_header_from_headers', + 'get_mcp_auth_header_from_headers', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp', + 'MCPRequestHandler', + '_get_mcp_server_auth_headers_from_headers', + 'get_mcp_server_auth_headers_from_headers', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp', + 'MCPRequestHandler', + '_get_mcp_servers_from_access_groups', + 'get_mcp_servers_from_access_groups', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp', + 'MCPRequestHandler', + '_get_oauth2_headers_from_headers', + 'get_oauth2_headers_from_headers', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp', + 'MCPRequestHandler', + '_safe_get_headers_from_scope', + 'safe_get_headers_from_scope', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.bridge_token_flow', + '', + '_bridge_mint_error_response', + 'bridge_mint_error_response', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.bridge_token_flow', + '', + '_extract_user_id_from_request', + 'extract_user_id_from_request', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.bridge_token_flow', + '', + '_finish_bridge_mint', + 'finish_bridge_mint', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.bridge_token_flow', + '', + '_prepare_bridge_mint', + 'prepare_bridge_mint', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.bridge_token_flow', + '', + '_prepare_bridge_refresh', + 'prepare_bridge_refresh', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.bridge_token_flow', + '', + '_reload_active_user_by_id', + 'reload_active_user_by_id', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.byok_oauth_endpoints', + '', + '_user_id_from_session_cookie', + 'user_id_from_session_cookie', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.db', + '', + '_decode_oauth_payload', + 'decode_oauth_payload', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.discoverable_endpoints', + '', + '_raise_if_not_oauth2', + 'raise_if_not_oauth2', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.mcp_context', + '', + '_mcp_active_toolset_id', + 'mcp_active_toolset_id', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.mcp_context', + '', + '_mcp_gateway_initialize_instructions', + 'mcp_gateway_initialize_instructions', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.mcp_context', + '', + '_mcp_gateway_server_name', + 'mcp_gateway_server_name', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.mcp_server_manager', + '', + '_UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES', + 'UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.mcp_server_manager', + '', + '_caller_authorization_fans_out', + 'caller_authorization_fans_out', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.mcp_server_manager', + '', + '_client_forwarded_authorization_headers', + 'client_forwarded_authorization_headers', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.mcp_server_manager', + '', + '_resolve_openapi_tool_auth', + 'resolve_openapi_tool_auth', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.mcp_server_manager', + '', + '_should_strip_caller_authorization', + 'should_strip_caller_authorization', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.mcp_server_manager', + 'MCPServerManager', + '_build_mcp_server_table', + 'build_mcp_server_table', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.mcp_server_manager', + 'MCPServerManager', + '_build_stdio_env', + 'build_stdio_env', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.mcp_server_manager', + 'MCPServerManager', + '_create_mcp_client', + 'create_mcp_client', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.mcp_server_manager', + 'MCPServerManager', + '_ensure_upstream_initialize_instructions_cached', + 'ensure_upstream_initialize_instructions_cached', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.mcp_server_manager', + 'MCPServerManager', + '_extract_subject_token', + 'extract_subject_token', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.mcp_server_manager', + 'MCPServerManager', + '_get_mcp_server_from_tool_name', + 'get_mcp_server_from_tool_name', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.mcp_server_manager', + 'MCPServerManager', + '_get_tools_from_server', + 'get_tools_from_server', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.mcp_server_manager', + 'MCPServerManager', + '_is_server_accessible_from_ip', + 'is_server_accessible_from_ip', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.oauth2_token_cache', + '', + '_compute_per_user_token_ttl', + 'compute_per_user_token_ttl', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.oauth_utils', + '', + '_redact_mcp_resource_url', + 'redact_mcp_resource_url', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.server', + '', + '_redact_mcp_resource_url', + 'redact_mcp_resource_url', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator', + '', + '_OPENAPI_TOOL_NAME_MAX_LEN', + 'OPENAPI_TOOL_NAME_MAX_LEN', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator', + '', + '_request_auth_header', + 'request_auth_header', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator', + '', + '_request_extra_headers', + 'request_extra_headers', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator', + '', + '_request_resolved_auth_headers', + 'request_resolved_auth_headers', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.semantic_tool_filter', + 'SemanticMCPToolFilter', + '_extract_tool_info', + 'extract_tool_info', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.server', + '', + '_apply_toolset_scope', + 'apply_toolset_scope', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.server_resolution', + 'MCPServerRegistry', + '_build_mcp_server_table', + 'build_mcp_server_table', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.server_resolution', + 'MCPServerRegistry', + '_is_server_accessible_from_ip', + 'is_server_accessible_from_ip', + False, + ), + ( + 'litellm.proxy.anthropic_endpoints.claude_code_endpoints.claude_code_marketplace', + '', + '_get_prisma_client', + 'get_prisma_client', + False, + ), + ( + 'litellm.proxy.auth.auth_checks', + '', + '_cache_access_object', + 'cache_access_object', + False, + ), + ( + 'litellm.proxy.auth.auth_checks', + '', + '_cache_key_object', + 'cache_key_object', + False, + ), + ( + 'litellm.proxy.auth.auth_checks', + '', + '_cache_team_object', + 'cache_team_object', + False, + ), + ( + 'litellm.proxy.auth.auth_checks', + '', + '_can_object_call_model', + 'can_object_call_model', + False, + ), + ( + 'litellm.proxy.auth.auth_checks', + '', + '_check_end_user_budget', + 'check_end_user_budget', + False, + ), + ( + 'litellm.proxy.auth.auth_checks', + '', + '_check_model_access_helper', + 'check_model_access_helper', + False, + ), + ( + 'litellm.proxy.auth.auth_checks', + '', + '_check_team_member_model_access', + 'check_team_member_model_access', + False, + ), + ( + 'litellm.proxy.auth.auth_checks', + '', + '_copy_user_api_key_auth_for_cache', + 'copy_user_api_key_auth_for_cache', + False, + ), + ( + 'litellm.proxy.auth.auth_checks', + '', + '_delete_cache_access_object', + 'delete_cache_access_object', + False, + ), + ( + 'litellm.proxy.auth.auth_checks', + '', + '_delete_cache_key_object', + 'delete_cache_key_object', + False, + ), + ( + 'litellm.proxy.auth.auth_checks', + '', + '_fetch_key_object_from_db_with_reconnect', + 'fetch_key_object_from_db_with_reconnect', + False, + ), + ( + 'litellm.proxy.auth.auth_checks', + '', + '_get_agent_ids_from_access_groups', + 'get_agent_ids_from_access_groups', + False, + ), + ( + 'litellm.proxy.auth.auth_checks', + '', + '_get_mcp_server_ids_from_access_groups', + 'get_mcp_server_ids_from_access_groups', + False, + ), + ( + 'litellm.proxy.auth.auth_checks', + '', + '_get_models_from_access_groups', + 'get_models_from_access_groups', + False, + ), + ( + 'litellm.proxy.auth.auth_checks', + '', + '_get_team_object_from_cache', + 'get_team_object_from_cache', + False, + ), + ( + 'litellm.proxy.auth.auth_checks', + '', + '_get_user_role', + 'get_user_role', + False, + ), + ( + 'litellm.proxy.auth.auth_checks', + '', + '_is_model_cost_zero', + 'is_model_cost_zero', + False, + ), + ( + 'litellm.proxy.auth.auth_checks', + '', + '_is_user_proxy_admin', + 'is_user_proxy_admin', + False, + ), + ( + 'litellm.proxy.auth.auth_checks', + '', + '_key_access_group_grants_model', + 'key_access_group_grants_model', + False, + ), + ( + 'litellm.proxy.auth.auth_checks', + '', + '_organization_max_budget_check', + 'organization_max_budget_check', + False, + ), + ( + 'litellm.proxy.auth.auth_checks', + '', + '_team_max_budget_check', + 'team_max_budget_check', + False, + ), + ( + 'litellm.proxy.auth.auth_checks', + '', + '_team_member_max_budget_alert_check', + 'team_member_max_budget_alert_check', + False, + ), + ( + 'litellm.proxy.auth.auth_checks', + '', + '_virtual_key_max_budget_alert_check', + 'virtual_key_max_budget_alert_check', + False, + ), + ( + 'litellm.proxy.auth.auth_checks', + '', + '_virtual_key_max_budget_check', + 'virtual_key_max_budget_check', + False, + ), + ( + 'litellm.proxy.auth.auth_checks', + '', + '_virtual_key_soft_budget_check', + 'virtual_key_soft_budget_check', + False, + ), + ( + 'litellm.proxy.auth.auth_checks_organization', + '', + '_user_is_org_admin', + 'user_is_org_admin', + False, + ), + ( + 'litellm.proxy.auth.auth_exception_handler', + 'UserAPIKeyAuthExceptionHandler', + '_handle_authentication_error', + 'handle_authentication_error', + False, + ), + ( + 'litellm.proxy.auth.auth_utils', + '', + '_get_request_ip_address', + 'get_request_ip_address', + False, + ), + ( + 'litellm.proxy.auth.resolvers.store', + 'IdentityStore', + '_principal_from_key', + 'principal_from_key', + False, + ), + ( + 'litellm.proxy.auth.route_checks', + 'RouteChecks', + '_get_request_method', + 'get_request_method', + False, + ), + ( + 'litellm.proxy.auth.route_checks', + 'RouteChecks', + '_is_assistants_api_request', + 'is_assistants_api_request', + False, + ), + ( + 'litellm.proxy.auth.route_checks', + 'RouteChecks', + '_is_wildcard_pattern', + 'is_wildcard_pattern', + False, + ), + ( + 'litellm.proxy.auth.user_api_key_auth', + '', + '_enforce_key_and_fallback_model_access', + 'enforce_key_and_fallback_model_access', + False, + ), + ( + 'litellm.proxy.auth.user_api_key_auth', + '', + '_fetch_global_spend_with_event_coordination', + 'fetch_global_spend_with_event_coordination', + False, + ), + ( + 'litellm.proxy.auth.user_api_key_auth', + '', + '_get_bearer_token', + 'get_bearer_token', + False, + ), + ( + 'litellm.proxy.auth.user_api_key_auth', + '', + '_run_centralized_common_checks', + 'run_centralized_common_checks', + False, + ), + ( + 'litellm.proxy.auth.user_api_key_auth', + '', + '_user_api_key_auth_builder', + 'user_api_key_auth_builder', + False, + ), + ( + 'litellm.proxy.common_request_processing', + '', + '_is_azure_model_router_request', + 'is_azure_model_router_request', + False, + ), + ( + 'litellm.proxy.common_request_processing', + '', + '_should_return_raw_model_name', + 'should_return_raw_model_name', + False, + ), + ( + 'litellm.proxy.common_request_processing', + 'ProxyBaseLLMRequestProcessing', + '_finalize_streaming_generator_cleanup', + 'finalize_streaming_generator_cleanup', + False, + ), + ( + 'litellm.proxy.common_request_processing', + 'ProxyBaseLLMRequestProcessing', + '_handle_llm_api_exception', + 'handle_llm_api_exception', + False, + ), + ( + 'litellm.proxy.common_request_processing', + 'ProxyBaseLLMRequestProcessing', + '_process_chunk_with_cost_injection', + 'process_chunk_with_cost_injection', + False, + ), + ( + 'litellm.proxy.common_utils.callback_utils', + '', + '_CALLBACK_VAR_ENCRYPTED_PREFIX', + 'CALLBACK_VAR_ENCRYPTED_PREFIX', + False, + ), + ( + 'litellm.proxy.common_utils.config_sync_pubsub', + '', + '_ConfigSyncPubSub', + 'ConfigSyncPubSub', + False, + ), + ( + 'litellm.proxy.common_utils.config_sync_pubsub', + '', + '_pubsub_capable_client', + 'pubsub_capable_client', + False, + ), + ( + 'litellm.proxy.common_utils.encrypt_decrypt_utils', + '', + '_ALGO_AES_GCM', + 'ALGO_AES_GCM', + False, + ), + ( + 'litellm.proxy.common_utils.encrypt_decrypt_utils', + '', + '_ENCRYPTION_ALGORITHM_SETTING', + 'ENCRYPTION_ALGORITHM_SETTING', + False, + ), + ( + 'litellm.proxy.common_utils.encrypt_decrypt_utils', + '', + '_V2_GCM_PREFIX', + 'V2_GCM_PREFIX', + False, + ), + ( + 'litellm.proxy.common_utils.encrypt_decrypt_utils', + '', + '_get_salt_key', + 'get_salt_key', + False, + ), + ( + 'litellm.proxy.common_utils.http_parsing_utils', + '', + '_read_request_body', + 'read_request_body', + False, + ), + ( + 'litellm.proxy.common_utils.http_parsing_utils', + '', + '_safe_get_request_headers', + 'safe_get_request_headers', + False, + ), + ( + 'litellm.proxy.common_utils.http_parsing_utils', + '', + '_safe_get_request_query_params', + 'safe_get_request_query_params', + False, + ), + ( + 'litellm.proxy.common_utils.http_parsing_utils', + '', + '_safe_set_request_parsed_body', + 'safe_set_request_parsed_body', + False, + ), + ( + 'litellm.proxy.common_utils.realtime_utils', + '', + '_realtime_request_body', + 'realtime_request_body', + False, + ), + ( + 'litellm.proxy.db.db_spend_update_writer', + 'DBSpendUpdateWriter', + '_commit_daily_tag_spend_to_db', + 'commit_daily_tag_spend_to_db', + False, + ), + ( + 'litellm.proxy.db.db_spend_update_writer', + 'DBSpendUpdateWriter', + '_commit_daily_tag_spend_to_db_with_redis', + 'commit_daily_tag_spend_to_db_with_redis', + False, + ), + ( + 'litellm.proxy.db.db_spend_update_writer', + 'DBSpendUpdateWriter', + '_handle_spend_update_failure', + 'handle_spend_update_failure', + False, + ), + ( + 'litellm.proxy.db.db_transaction_queue.redis_update_buffer', + 'RedisUpdateBuffer', + '_should_commit_spend_updates_to_redis', + 'should_commit_spend_updates_to_redis', + False, + ), + ( + 'litellm.proxy.db.log_db_metrics', + '', + '_is_exception_related_to_db', + 'is_exception_related_to_db', + False, + ), + ( + 'litellm.proxy.guardrails.guardrail_hooks.azure.base', + '', + '_RESPONSES_API_CALL_TYPES', + 'RESPONSES_API_CALL_TYPES', + False, + ), + ( + 'litellm.proxy.guardrails.guardrail_hooks.cisco_ai_defense.cisco_ai_defense_mcp', + '', + '_CiscoAIDefenseMcpMixin', + 'CiscoAIDefenseMcpMixin', + False, + ), + ( + 'litellm.proxy.guardrails.guardrail_hooks.litellm_content_filter.competitor_intent.base', + '', + '_compile_marker', + 'compile_marker', + False, + ), + ( + 'litellm.proxy.guardrails.guardrail_hooks.litellm_content_filter.competitor_intent.base', + '', + '_count_signals', + 'count_signals', + False, + ), + ( + 'litellm.proxy.guardrails.guardrail_hooks.litellm_content_filter.competitor_intent.base', + '', + '_word_boundary_match', + 'word_boundary_match', + False, + ), + ( + 'litellm.proxy.guardrails.guardrail_hooks.presidio', + '', + '_OPTIONAL_PresidioPIIMasking', + 'OPTIONAL_PresidioPIIMasking', + False, + ), + ( + 'litellm.proxy.health_check', + '', + '_clean_endpoint_data', + 'clean_endpoint_data', + False, + ), + ( + 'litellm.proxy.health_check', + '', + '_update_litellm_params_for_health_check', + 'update_litellm_params_for_health_check', + False, + ), + ( + 'litellm.proxy.health_endpoints._health_endpoints', + '', + '_convert_health_check_to_dict', + 'convert_health_check_to_dict', + False, + ), + ( + 'litellm.proxy.health_endpoints._health_endpoints', + '', + '_save_background_health_checks_to_db', + 'save_background_health_checks_to_db', + False, + ), + ( + 'litellm.proxy.hooks.azure_content_safety', + '', + '_PROXY_AzureContentSafety', + 'PROXY_AzureContentSafety', + False, + ), + ( + 'litellm.proxy.hooks.batch_rate_limiter', + '', + '_PROXY_BatchRateLimiter', + 'PROXY_BatchRateLimiter', + False, + ), + ( + 'litellm.proxy.hooks.batch_redis_get', + '', + '_PROXY_BatchRedisRequests', + 'PROXY_BatchRedisRequests', + False, + ), + ( + 'litellm.proxy.hooks.cache_control_check', + '', + '_PROXY_CacheControlCheck', + 'PROXY_CacheControlCheck', + False, + ), + ( + 'litellm.proxy.hooks.dynamic_rate_limiter', + '', + '_PROXY_DynamicRateLimitHandler', + 'PROXY_DynamicRateLimitHandler', + False, + ), + ( + 'litellm.proxy.hooks.dynamic_rate_limiter_v3', + '', + '_PROXY_DynamicRateLimitHandlerV3', + 'PROXY_DynamicRateLimitHandlerV3', + False, + ), + ( + 'litellm.proxy.hooks.max_budget_per_session_limiter', + '', + '_PROXY_MaxBudgetPerSessionHandler', + 'PROXY_MaxBudgetPerSessionHandler', + False, + ), + ( + 'litellm.proxy.hooks.max_iterations_limiter', + '', + '_PROXY_MaxIterationsHandler', + 'PROXY_MaxIterationsHandler', + False, + ), + ( + 'litellm.proxy.hooks.model_max_budget_limiter', + '', + '_PROXY_VirtualKeyModelMaxBudgetLimiter', + 'PROXY_VirtualKeyModelMaxBudgetLimiter', + False, + ), + ( + 'litellm.proxy.hooks.parallel_request_limiter', + '', + '_PROXY_MaxParallelRequestsHandler', + 'PROXY_MaxParallelRequestsHandler', + False, + ), + ( + 'litellm.proxy.hooks.parallel_request_limiter_v3', + '', + '_PROXY_MaxParallelRequestsHandler_v3', + 'PROXY_MaxParallelRequestsHandler_v3', + False, + ), + ( + 'litellm.proxy.hooks.parallel_request_limiter_v3', + '_PROXY_MaxParallelRequestsHandler_v3', + '_create_rate_limit_descriptors', + 'create_rate_limit_descriptors', + False, + ), + ( + 'litellm.proxy.hooks.prompt_injection_detection', + '', + '_OPTIONAL_PromptInjectionDetection', + 'OPTIONAL_PromptInjectionDetection', + False, + ), + ( + 'litellm.proxy.hooks.proxy_track_cost_callback', + '', + '_ProxyDBLogger', + 'ProxyDBLogger', + False, + ), + ( + 'litellm.proxy.hooks.sensitive_data_routing', + '', + '_PROXY_SensitiveDataRoutingHandler', + 'PROXY_SensitiveDataRoutingHandler', + False, + ), + ( + 'litellm.proxy.litellm_pre_call_utils', + '', + '_add_guardrails_from_key_or_team_metadata', + 'add_guardrails_from_key_or_team_metadata', + False, + ), + ( + 'litellm.proxy.litellm_pre_call_utils', + '', + '_get_dynamic_logging_metadata', + 'get_dynamic_logging_metadata', + False, + ), + ( + 'litellm.proxy.litellm_pre_call_utils', + '', + '_get_metadata_variable_name', + 'get_metadata_variable_name', + False, + ), + ( + 'litellm.proxy.litellm_pre_call_utils', + '', + '_get_validated_callback_metadata', + 'get_validated_callback_metadata', + False, + ), + ( + 'litellm.proxy.litellm_pre_call_utils', + 'LiteLLMProxyRequestSetup', + '_merge_tags', + 'merge_tags', + False, + ), + ( + 'litellm.proxy.management_endpoints.common_utils', + '', + '_check_disable_global_guardrails_caller_permission', + 'check_disable_global_guardrails_caller_permission', + False, + ), + ( + 'litellm.proxy.management_endpoints.common_utils', + '', + '_check_passthrough_routes_caller_permission', + 'check_passthrough_routes_caller_permission', + False, + ), + ( + 'litellm.proxy.management_endpoints.common_utils', + '', + '_set_object_metadata_field', + 'set_object_metadata_field', + False, + ), + ( + 'litellm.proxy.management_endpoints.common_utils', + '', + '_team_member_has_permission', + 'team_member_has_permission', + False, + ), + ( + 'litellm.proxy.management_endpoints.common_utils', + '', + '_update_metadata_fields', + 'update_metadata_fields', + False, + ), + ( + 'litellm.proxy.management_endpoints.common_utils', + '', + '_upsert_budget_and_membership', + 'upsert_budget_and_membership', + False, + ), + ( + 'litellm.proxy.management_endpoints.common_utils', + '', + '_user_has_admin_privileges', + 'user_has_admin_privileges', + False, + ), + ( + 'litellm.proxy.management_endpoints.common_utils', + '', + '_user_has_admin_view', + 'user_api_key_has_admin_view', + False, + ), + ( + 'litellm.proxy.management_endpoints.config_override_endpoints', + '', + '_clear_hashicorp_vault_state', + 'clear_hashicorp_vault_state', + False, + ), + ( + 'litellm.proxy.management_endpoints.config_override_endpoints', + '', + '_get_current_env_values', + 'get_current_env_values', + False, + ), + ( + 'litellm.proxy.management_endpoints.config_override_endpoints', + '', + '_parse_config_value', + 'parse_config_value', + False, + ), + ( + 'litellm.proxy.management_endpoints.config_override_endpoints', + '', + '_set_env_vars', + 'set_env_vars', + False, + ), + ( + 'litellm.proxy.management_endpoints.key_management_endpoints', + '', + '_calculate_key_rotation_time', + 'calculate_key_rotation_time', + False, + ), + ( + 'litellm.proxy.management_endpoints.key_management_endpoints', + '', + '_check_permissions_caller_permission', + 'check_permissions_caller_permission', + False, + ), + ( + 'litellm.proxy.management_endpoints.key_management_endpoints', + '', + '_get_caller_team_role', + 'get_caller_team_role', + False, + ), + ( + 'litellm.proxy.management_endpoints.key_management_endpoints', + '', + '_list_key_helper', + 'list_key_helper', + False, + ), + ( + 'litellm.proxy.management_endpoints.key_management_endpoints', + '', + '_persist_deleted_verification_tokens', + 'persist_deleted_verification_tokens', + False, + ), + ( + 'litellm.proxy.management_endpoints.key_management_endpoints', + '', + '_rotate_master_key', + 'rotate_master_key', + False, + ), + ( + 'litellm.proxy.management_endpoints.mcp_management_endpoints', + '', + '_inherit_credentials_from_existing_server', + 'inherit_credentials_from_existing_server', + False, + ), + ( + 'litellm.proxy.management_endpoints.model_management_endpoints', + '', + '_add_model_to_db', + 'add_model_to_db', + False, + ), + ( + 'litellm.proxy.management_endpoints.model_management_endpoints', + '', + '_add_team_model_to_db', + 'add_team_model_to_db', + False, + ), + ( + 'litellm.proxy.management_endpoints.model_management_endpoints', + '', + '_deduplicate_litellm_router_models', + 'deduplicate_litellm_router_models', + False, + ), + ( + 'litellm.proxy.management_endpoints.organization_endpoints', + '', + '_verify_org_access', + 'verify_org_access', + False, + ), + ( + 'litellm.proxy.management_endpoints.policy_endpoints.endpoints', + '', + '_build_all_names_per_competitor', + 'build_all_names_per_competitor', + False, + ), + ( + 'litellm.proxy.management_endpoints.policy_endpoints.endpoints', + '', + '_build_comparison_blocked_words', + 'build_comparison_blocked_words', + False, + ), + ( + 'litellm.proxy.management_endpoints.policy_endpoints.endpoints', + '', + '_build_competitor_guardrail_definitions', + 'build_competitor_guardrail_definitions', + False, + ), + ( + 'litellm.proxy.management_endpoints.policy_endpoints.endpoints', + '', + '_build_name_blocked_words', + 'build_name_blocked_words', + False, + ), + ( + 'litellm.proxy.management_endpoints.policy_endpoints.endpoints', + '', + '_build_recommendation_blocked_words', + 'build_recommendation_blocked_words', + False, + ), + ( + 'litellm.proxy.management_endpoints.policy_endpoints.endpoints', + '', + '_build_refinement_prompt', + 'build_refinement_prompt', + False, + ), + ( + 'litellm.proxy.management_endpoints.policy_endpoints.endpoints', + '', + '_clean_competitor_line', + 'clean_competitor_line', + False, + ), + ( + 'litellm.proxy.management_endpoints.policy_endpoints.endpoints', + '', + '_parse_variations_response', + 'parse_variations_response', + False, + ), + ( + 'litellm.proxy.management_endpoints.team_endpoints', + '', + '_cleanup_members_with_roles', + 'cleanup_members_with_roles', + False, + ), + ( + 'litellm.proxy.management_endpoints.team_endpoints', + '', + '_refresh_cached_team', + 'refresh_cached_team', + False, + ), + ( + 'litellm.proxy.management_endpoints.team_endpoints', + 'TeamMemberBudgetHandler', + '_clean_team_member_fields', + 'clean_team_member_fields', + False, + ), + ( + 'litellm.proxy.management_endpoints.ui_sso', + '', + '_sso_return_to_redirect', + 'sso_return_to_redirect', + False, + ), + ( + 'litellm.proxy.management_endpoints.ui_sso', + 'SSOAuthenticationHandler', + '_delete_pkce_verifier', + 'delete_pkce_verifier', + False, + ), + ( + 'litellm.proxy.management_endpoints.ui_sso', + 'SSOAuthenticationHandler', + '_get_cli_state', + 'get_cli_state', + False, + ), + ( + 'litellm.proxy.management_endpoints.ui_sso', + 'SSOAuthenticationHandler', + '_get_user_email_and_id_from_result', + 'get_user_email_and_id_from_result', + False, + ), + ( + 'litellm.proxy.management_endpoints.ui_sso', + 'SSOAuthenticationHandler', + '_pkce_token_exchange', + 'pkce_token_exchange', + False, + ), + ( + 'litellm.proxy.management_endpoints.ui_sso', + 'SSOAuthenticationHandler', + '_validate_return_to', + 'validate_return_to', + False, + ), + ( + 'litellm.proxy.management_helpers.object_permission_utils', + '', + '_get_allow_all_keys_server_ids', + 'get_allow_all_keys_server_ids', + False, + ), + ( + 'litellm.proxy.management_helpers.object_permission_utils', + '', + '_get_team_allowed_mcp_servers', + 'get_team_allowed_mcp_servers', + False, + ), + ( + 'litellm.proxy.management_helpers.object_permission_utils', + '', + '_set_object_permission', + 'set_object_permission', + False, + ), + ( + 'litellm.proxy.openai_files_endpoints.common_utils', + '', + '_is_base64_encoded_unified_file_id', + 'is_base64_encoded_unified_file_id', + False, + ), + ( + 'litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints', + '', + '_extract_model_from_bedrock_endpoint', + 'extract_model_from_bedrock_endpoint', + False, + ), + ( + 'litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints', + 'BaseOpenAIPassThroughHandler', + '_base_openai_pass_through_handler', + 'base_openai_pass_through_handler', + False, + ), + ( + 'litellm.proxy.pass_through_endpoints.llm_provider_handlers.anthropic_passthrough_logging_handler', + 'AnthropicPassthroughLoggingHandler', + '_build_complete_streaming_response', + 'build_complete_streaming_response', + False, + ), + ( + 'litellm.proxy.pass_through_endpoints.llm_provider_handlers.anthropic_passthrough_logging_handler', + 'AnthropicPassthroughLoggingHandler', + '_build_usage_only_response_from_chunks', + 'build_usage_only_response_from_chunks', + False, + ), + ( + 'litellm.proxy.pass_through_endpoints.llm_provider_handlers.anthropic_passthrough_logging_handler', + 'AnthropicPassthroughLoggingHandler', + '_handle_logging_anthropic_collected_chunks', + 'handle_logging_anthropic_collected_chunks', + False, + ), + ( + 'litellm.proxy.pass_through_endpoints.llm_provider_handlers.assembly_passthrough_logging_handler', + 'AssemblyAIPassthroughLoggingHandler', + '_get_assembly_base_url_from_region', + 'get_assembly_base_url_from_region', + False, + ), + ( + 'litellm.proxy.pass_through_endpoints.llm_provider_handlers.assembly_passthrough_logging_handler', + 'AssemblyAIPassthroughLoggingHandler', + '_get_assembly_region_from_url', + 'get_assembly_region_from_url', + False, + ), + ( + 'litellm.proxy.pass_through_endpoints.llm_provider_handlers.assembly_passthrough_logging_handler', + 'AssemblyAIPassthroughLoggingHandler', + '_should_log_request', + 'should_log_request', + False, + ), + ( + 'litellm.proxy.pass_through_endpoints.llm_provider_handlers.openai_passthrough_logging_handler', + '', + '_is_openai_compatible_url', + 'is_openai_compatible_url', + False, + ), + ( + 'litellm.proxy.pass_through_endpoints.llm_provider_handlers.openai_passthrough_logging_handler', + 'OpenAIPassthroughLoggingHandler', + '_handle_logging_openai_collected_chunks', + 'handle_logging_openai_collected_chunks', + False, + ), + ( + 'litellm.proxy.pass_through_endpoints.llm_provider_handlers.vertex_passthrough_logging_handler', + 'VertexPassthroughLoggingHandler', + '_handle_logging_vertex_collected_chunks', + 'handle_logging_vertex_collected_chunks', + False, + ), + ( + 'litellm.proxy.pass_through_endpoints.pass_through_endpoints', + 'HttpPassThroughEndpointHelpers', + '_init_kwargs_for_pass_through_endpoint', + 'init_kwargs_for_pass_through_endpoint', + False, + ), + ( + 'litellm.proxy.pass_through_endpoints.pass_through_endpoints', + 'HttpPassThroughEndpointHelpers', + '_update_stream_param_based_on_request_body', + 'update_stream_param_based_on_request_body', + False, + ), + ( + 'litellm.proxy.pass_through_endpoints.streaming_handler', + 'PassThroughStreamingHandler', + '_route_streaming_logging_to_handler', + 'route_streaming_logging_to_handler', + False, + ), + ( + 'litellm.proxy.pass_through_endpoints.success_handler', + 'PassThroughEndpointLogging', + '_handle_logging', + 'handle_logging', + False, + ), + ( + 'litellm.proxy.policy_engine.policy_registry', + 'PolicyRegistry', + '_parse_policy', + 'parse_policy', + False, + ), + ( + 'litellm.proxy.proxy_cli', + 'ProxyInitializationHelpers', + '_apply_uvicorn_max_requests_jitter', + 'apply_uvicorn_max_requests_jitter', + False, + ), + ( + 'litellm.proxy.proxy_cli', + 'ProxyInitializationHelpers', + '_configure_dev_reload', + 'configure_dev_reload', + False, + ), + ( + 'litellm.proxy.proxy_cli', + 'ProxyInitializationHelpers', + '_echo_litellm_version', + 'echo_litellm_version', + False, + ), + ( + 'litellm.proxy.proxy_cli', + 'ProxyInitializationHelpers', + '_get_default_unvicorn_init_args', + 'get_default_unvicorn_init_args', + False, + ), + ( + 'litellm.proxy.proxy_cli', + 'ProxyInitializationHelpers', + '_get_loop_type', + 'get_loop_type', + False, + ), + ( + 'litellm.proxy.proxy_cli', + 'ProxyInitializationHelpers', + '_init_granian_server', + 'init_granian_server', + False, + ), + ( + 'litellm.proxy.proxy_cli', + 'ProxyInitializationHelpers', + '_init_hypercorn_server', + 'init_hypercorn_server', + False, + ), + ( + 'litellm.proxy.proxy_cli', + 'ProxyInitializationHelpers', + '_is_port_in_use', + 'is_port_in_use', + False, + ), + ( + 'litellm.proxy.proxy_cli', + 'ProxyInitializationHelpers', + '_maybe_setup_prometheus_multiproc_dir', + 'maybe_setup_prometheus_multiproc_dir', + False, + ), + ( + 'litellm.proxy.proxy_cli', + 'ProxyInitializationHelpers', + '_run_config_validation', + 'run_config_validation', + False, + ), + ( + 'litellm.proxy.proxy_cli', + 'ProxyInitializationHelpers', + '_run_gunicorn_server', + 'run_gunicorn_server', + False, + ), + ( + 'litellm.proxy.proxy_cli', + 'ProxyInitializationHelpers', + '_run_health_check', + 'run_health_check', + False, + ), + ( + 'litellm.proxy.proxy_cli', + 'ProxyInitializationHelpers', + '_run_ollama_serve', + 'run_ollama_serve', + False, + ), + ( + 'litellm.proxy.proxy_cli', + 'ProxyInitializationHelpers', + '_run_test_chat_completion', + 'run_test_chat_completion', + False, + ), + ( + 'litellm.proxy.proxy_server', + '', + '_build_redis_usage_cache', + 'build_redis_usage_cache', + False, + ), + ( + 'litellm.proxy.proxy_server', + '', + '_ensure_spend_counter_initialized', + 'ensure_spend_counter_initialized', + False, + ), + ( + 'litellm.proxy.proxy_server', + '', + '_ensure_window_spend_counter_initialized', + 'ensure_window_spend_counter_initialized', + False, + ), + ( + 'litellm.proxy.proxy_server', + '', + '_environment_has_redis_connection_target', + 'environment_has_redis_connection_target', + False, + ), + ( + 'litellm.proxy.proxy_server', + '', + '_get_model_group_info', + 'get_model_group_info', + False, + ), + ( + 'litellm.proxy.proxy_server', + '', + '_increment_spend_counter_cache', + 'increment_spend_counter_cache', + False, + ), + ( + 'litellm.proxy.proxy_server', + '', + '_initialize_shared_aiohttp_session', + 'initialize_shared_aiohttp_session', + False, + ), + ( + 'litellm.proxy.proxy_server', + '', + '_invalidate_spend_counter', + 'invalidate_spend_counter', + False, + ), + ( + 'litellm.proxy.proxy_server', + '', + '_title', + 'title', + False, + ), + ( + 'litellm.proxy.proxy_server', + 'ProxyConfig', + '_add_deployment_locked', + 'add_deployment_locked', + False, + ), + ( + 'litellm.proxy.proxy_server', + 'ProxyConfig', + '_decrypt_and_set_db_env_variables', + 'decrypt_and_set_db_env_variables', + False, + ), + ( + 'litellm.proxy.proxy_server', + 'ProxyConfig', + '_decrypt_db_variables', + 'decrypt_db_variables', + False, + ), + ( + 'litellm.proxy.proxy_server', + 'ProxyConfig', + '_encrypt_env_variables', + 'encrypt_env_variables', + False, + ), + ( + 'litellm.proxy.proxy_server', + 'ProxyConfig', + '_encrypt_env_variables_for_db', + 'encrypt_env_variables_for_db', + False, + ), + ( + 'litellm.proxy.proxy_server', + 'ProxyConfig', + '_get_models_from_db', + 'get_models_from_db', + False, + ), + ( + 'litellm.proxy.proxy_server', + 'ProxyConfig', + '_init_cache', + 'init_cache', + False, + ), + ( + 'litellm.proxy.proxy_server', + 'ProxyConfig', + '_init_semantic_filter_settings_in_db', + 'init_semantic_filter_settings_in_db', + False, + ), + ( + 'litellm.proxy.proxy_server', + 'ProxyConfig', + '_serve_pass_through_endpoints', + 'serve_pass_through_endpoints', + False, + ), + ( + 'litellm.proxy.proxy_server', + 'ProxyStartupEvent', + '_add_proxy_budget_to_db', + 'add_proxy_budget_to_db', + False, + ), + ( + 'litellm.proxy.proxy_server', + 'ProxyStartupEvent', + '_attach_router_to_prompt_injection_detectors', + 'attach_router_to_prompt_injection_detectors', + False, + ), + ( + 'litellm.proxy.proxy_server', + 'ProxyStartupEvent', + '_get_transaction_buffer_redis_cache', + 'get_transaction_buffer_redis_cache', + False, + ), + ( + 'litellm.proxy.proxy_server', + 'ProxyStartupEvent', + '_init_coordination_redis_from_db', + 'init_coordination_redis_from_db', + False, + ), + ( + 'litellm.proxy.proxy_server', + 'ProxyStartupEvent', + '_init_dd_tracer', + 'init_dd_tracer', + False, + ), + ( + 'litellm.proxy.proxy_server', + 'ProxyStartupEvent', + '_init_pyroscope', + 'init_pyroscope', + False, + ), + ( + 'litellm.proxy.proxy_server', + 'ProxyStartupEvent', + '_initialize_jwt_auth', + 'initialize_jwt_auth', + False, + ), + ( + 'litellm.proxy.proxy_server', + 'ProxyStartupEvent', + '_initialize_semantic_tool_filter', + 'initialize_semantic_tool_filter', + False, + ), + ( + 'litellm.proxy.proxy_server', + 'ProxyStartupEvent', + '_initialize_startup_logging', + 'initialize_startup_logging', + False, + ), + ( + 'litellm.proxy.proxy_server', + 'ProxyStartupEvent', + '_setup_prisma_client', + 'setup_prisma_client', + False, + ), + ( + 'litellm.proxy.proxy_server', + 'ProxyStartupEvent', + '_sync_ui_settings_to_general_settings', + 'sync_ui_settings_to_general_settings', + False, + ), + ( + 'litellm.proxy.proxy_server', + 'ProxyStartupEvent', + '_update_default_team_member_budget', + 'update_default_team_member_budget', + False, + ), + ( + 'litellm.proxy.proxy_server', + 'ProxyStartupEvent', + '_validate_redis_transaction_buffer_config', + 'validate_redis_transaction_buffer_config', + False, + ), + ( + 'litellm.proxy.proxy_server', + 'ProxyStartupEvent', + '_warm_global_spend_cache', + 'warm_global_spend_cache', + False, + ), + ( + 'litellm.proxy.proxy_server', + 'ProxyStartupEvent', + '_warn_budget_without_db', + 'warn_budget_without_db', + False, + ), + ( + 'litellm.proxy.proxy_server', + 'ProxyStartupEvent', + '_warn_fail_closed_rate_limits_without_redis', + 'warn_fail_closed_rate_limits_without_redis', + False, + ), + ( + 'litellm.proxy.proxy_server', + 'ProxyStartupEvent', + '_warn_if_mock_testing_params_enabled', + 'warn_if_mock_testing_params_enabled', + False, + ), + ( + 'litellm.proxy.spend_tracking.spend_management_endpoints', + '', + '_get_spend_report_for_time_range', + 'get_spend_report_for_time_range', + False, + ), + ( + 'litellm.proxy.spend_tracking.spend_management_endpoints', + '', + '_is_admin_view_safe', + 'is_admin_view_safe', + False, + ), + ( + 'litellm.proxy.spend_tracking.spend_tracking_utils', + '', + '_is_master_key', + 'is_master_key', + False, + ), + ( + 'litellm.proxy.spend_tracking.spend_tracking_utils', + '', + '_sanitize_error_information_for_spend_logs', + 'sanitize_error_information_for_spend_logs', + False, + ), + ( + 'litellm.proxy.utils', + '', + '_cache_user_row', + 'cache_user_row', + False, + ), + ( + 'litellm.proxy.utils', + '', + '_check_and_merge_model_level_guardrails', + 'check_and_merge_model_level_guardrails', + False, + ), + ( + 'litellm.proxy.utils', + '', + '_get_docs_url', + 'get_docs_url', + False, + ), + ( + 'litellm.proxy.utils', + '', + '_get_openapi_url', + 'get_openapi_url', + False, + ), + ( + 'litellm.proxy.utils', + '', + '_get_projected_spend_over_limit', + 'get_projected_spend_over_limit', + False, + ), + ( + 'litellm.proxy.utils', + '', + '_get_redoc_url', + 'get_redoc_url', + False, + ), + ( + 'litellm.proxy.utils', + '', + '_hash_token_if_needed', + 'hash_token_if_needed', + False, + ), + ( + 'litellm.proxy.utils', + '', + '_is_projected_spend_over_limit', + 'is_projected_spend_over_limit', + False, + ), + ( + 'litellm.proxy.utils', + '', + '_is_valid_team_configs', + 'is_valid_team_configs', + False, + ), + ( + 'litellm.proxy.utils', + '', + '_monitor_spend_logs_queue', + 'monitor_spend_logs_queue', + False, + ), + ( + 'litellm.proxy.utils', + '', + '_premium_user_check', + 'premium_user_check', + False, + ), + ( + 'litellm.proxy.utils', + '', + '_raise_failed_update_spend_exception', + 'raise_failed_update_spend_exception', + False, + ), + ( + 'litellm.proxy.utils', + 'ProxyLogging', + '_arelease_max_parallel_requests_on_disconnect', + 'arelease_max_parallel_requests_on_disconnect', + False, + ), + ( + 'litellm.proxy.utils', + 'ProxyLogging', + '_callback_capabilities', + 'callback_capabilities', + False, + ), + ( + 'litellm.proxy.utils', + 'ProxyLogging', + '_convert_mcp_hook_response_to_kwargs', + 'convert_mcp_hook_response_to_kwargs', + False, + ), + ( + 'litellm.proxy.utils', + 'ProxyLogging', + '_create_mcp_request_object_from_kwargs', + 'create_mcp_request_object_from_kwargs', + False, + ), + ( + 'litellm.proxy.utils', + 'ProxyLogging', + '_fire_deferred_stream_logging', + 'fire_deferred_stream_logging', + False, + ), +) + LLMS_FORWARDER_CASES: Final = ( ( "litellm.llms.bedrock.chat.invoke_handler", @@ -2218,6 +5187,17 @@ LLMS_FORWARDER_CASES: Final = ( False, ), ) +PROXY_FORWARDER_CASES: Final = ( + ( + 'litellm.proxy.utils', + 'ProxyLogging', + '_convert_mcp_to_llm_format', + 'convert_mcp_to_llm_format', + 'instance', + False, + ), +) + def _get_owner(module_path: str, owner_name: str) -> object: @@ -2242,7 +5222,7 @@ def _get_instance(owner: object, public_name: str) -> object: @pytest.mark.parametrize( ("module_path", "owner_name", "old_name", "new_name", "use_instance"), - (*ALIAS_CASES, *LLMS_ALIAS_CASES), + (*ALIAS_CASES, *LLMS_ALIAS_CASES, *PROXY_ALIAS_CASES), ) def test_public_aliases( module_path: str, @@ -2260,6 +5240,51 @@ def test_public_aliases( assert old_value is new_value +@pytest.mark.parametrize( + ("package_name", "private_name", "public_name"), + PACKAGE_EXPORT_ALIAS_CASES, +) +def test_private_package_exports_are_available_and_match_public_alias( + package_name: str, + private_name: str, + public_name: str, +) -> None: + package: Final = import_module(package_name) + + assert getattr(package, private_name) is getattr(package, public_name) + + +@pytest.mark.parametrize( + ("module_path", "private_name", "public_name"), + MODULE_IMPORT_ALIAS_CASES, +) +def test_module_level_private_imports_remain_compatible( + module_path: str, + private_name: str, + public_name: str, +) -> None: + module: Final = import_module(module_path) + + assert getattr(module, private_name) is getattr(module, public_name) + + +@pytest.mark.parametrize( + ("module_path", "private_name", "public_name"), + PROXY_CLASS_NAME_ALIAS_CASES, +) +def test_proxy_class_aliases_keep_the_private_name( + module_path: str, + private_name: str, + public_name: str, +) -> None: + module: Final = import_module(module_path) + private_class: Final = cast(type[object], getattr(module, private_name)) + public_class: Final = cast(type[object], getattr(module, public_name)) + + assert public_class is private_class + assert public_class.__name__ == private_name + + def _make_private_override( descriptor: str, is_async: bool, @@ -2325,7 +5350,7 @@ def _forwarder_arguments( "descriptor", "is_async", ), - LLMS_FORWARDER_CASES, + (*LLMS_FORWARDER_CASES, *PROXY_FORWARDER_CASES), ) async def test_public_forwarders_dispatch_to_private_subclass_override( module_path: str, diff --git a/tests/unit/test_rate_limit_error_unification.py b/tests/unit/test_rate_limit_error_unification.py index e5acba938c7..ac2af7dba99 100644 --- a/tests/unit/test_rate_limit_error_unification.py +++ b/tests/unit/test_rate_limit_error_unification.py @@ -336,10 +336,10 @@ class TestProxyHooksActuallyRaiseProxyRateLimitError: from unittest.mock import MagicMock from litellm.proxy.hooks.parallel_request_limiter import ( - _PROXY_MaxParallelRequestsHandler, + PROXY_MaxParallelRequestsHandler, ) - handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=MagicMock()) + handler = PROXY_MaxParallelRequestsHandler(internal_usage_cache=MagicMock()) with pytest.raises(ProxyRateLimitError) as exc_info: handler.raise_rate_limit_error(additional_details="key-over-rpm") e = exc_info.value @@ -366,10 +366,10 @@ class TestProxyHooksActuallyRaiseProxyRateLimitError: from unittest.mock import MagicMock from litellm.proxy.hooks.parallel_request_limiter import ( - _PROXY_MaxParallelRequestsHandler, + PROXY_MaxParallelRequestsHandler, ) - handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=MagicMock()) + handler = PROXY_MaxParallelRequestsHandler(internal_usage_cache=MagicMock()) with pytest.raises(ProxyRateLimitError) as exc_info: handler.raise_rate_limit_error() # no additional_details detail_str = str(exc_info.value.detail) @@ -424,10 +424,10 @@ class TestProxyHooksActuallyRaiseProxyRateLimitError: from unittest.mock import MagicMock from litellm.proxy.hooks.parallel_request_limiter_v3 import ( - _PROXY_MaxParallelRequestsHandler_v3, + PROXY_MaxParallelRequestsHandler_v3, ) - handler = _PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=MagicMock()) + handler = PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=MagicMock()) # Minimal fabricated OVER_LIMIT response. The helper only reads a # handful of fields off `status` and ignores everything else. response = { @@ -475,13 +475,13 @@ class TestProxyHooksActuallyRaiseProxyRateLimitError: from litellm.caching.caching import DualCache from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.hooks.max_iterations_limiter import ( - _PROXY_MaxIterationsHandler, + PROXY_MaxIterationsHandler, ) from litellm.proxy.utils import InternalUsageCache from litellm.types.agents import AgentResponse cache = DualCache() - handler = _PROXY_MaxIterationsHandler( + handler = PROXY_MaxIterationsHandler( internal_usage_cache=InternalUsageCache(cache), ) user_api_key_dict = UserAPIKeyAuth( @@ -531,10 +531,10 @@ class TestProxyHooksActuallyRaiseProxyRateLimitError: from litellm.caching.caching import DualCache from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.hooks.dynamic_rate_limiter import ( - _PROXY_DynamicRateLimitHandler, + PROXY_DynamicRateLimitHandler, ) - handler = _PROXY_DynamicRateLimitHandler(internal_usage_cache=MagicMock()) + handler = PROXY_DynamicRateLimitHandler(internal_usage_cache=MagicMock()) # check_available_usage returns (available_tpm, available_rpm, # model_tpm, model_rpm, active_projects). Setting available_tpm == 0 # forces the TPM-exceeded raise. @@ -570,12 +570,12 @@ class TestProxyHooksActuallyRaiseProxyRateLimitError: from litellm.caching.caching import DualCache from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.hooks.parallel_request_limiter import ( - _PROXY_MaxParallelRequestsHandler, + PROXY_MaxParallelRequestsHandler, ) cache = MagicMock() cache.async_batch_set_cache = AsyncMock(return_value=None) - handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=cache) + handler = PROXY_MaxParallelRequestsHandler(internal_usage_cache=cache) with pytest.raises(ProxyRateLimitError) as exc_info: await handler.check_key_in_limits( user_api_key_dict=UserAPIKeyAuth(api_key="sk-key"), @@ -634,12 +634,12 @@ class TestProxyHooksActuallyRaiseProxyRateLimitError: from litellm.caching.caching import DualCache from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.hooks.parallel_request_limiter import ( - _PROXY_MaxParallelRequestsHandler, + PROXY_MaxParallelRequestsHandler, ) cache = MagicMock() cache.async_batch_set_cache = AsyncMock(return_value=None) - handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=cache) + handler = PROXY_MaxParallelRequestsHandler(internal_usage_cache=cache) with pytest.raises(ProxyRateLimitError) as exc_info: await handler.check_key_in_limits( user_api_key_dict=UserAPIKeyAuth(api_key="sk-key"), @@ -692,12 +692,12 @@ class TestProxyHooksActuallyRaiseProxyRateLimitError: from litellm.caching.caching import DualCache from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.hooks.parallel_request_limiter import ( - _PROXY_MaxParallelRequestsHandler, + PROXY_MaxParallelRequestsHandler, ) cache = MagicMock() cache.async_batch_set_cache = AsyncMock(return_value=None) - handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=cache) + handler = PROXY_MaxParallelRequestsHandler(internal_usage_cache=cache) with pytest.raises(ProxyRateLimitError) as exc_info: await handler.check_key_in_limits( user_api_key_dict=UserAPIKeyAuth(api_key="sk-key"), @@ -723,10 +723,10 @@ class TestProxyHooksActuallyRaiseProxyRateLimitError: from litellm.caching.caching import DualCache from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.hooks.dynamic_rate_limiter import ( - _PROXY_DynamicRateLimitHandler, + PROXY_DynamicRateLimitHandler, ) - handler = _PROXY_DynamicRateLimitHandler(internal_usage_cache=MagicMock()) + handler = PROXY_DynamicRateLimitHandler(internal_usage_cache=MagicMock()) # available_tpm > 0, available_rpm == 0 → RPM raise branch. handler.check_available_usage = AsyncMock( # type: ignore[method-assign] return_value=(100, 0, 1000, 100, 1) @@ -769,13 +769,13 @@ class TestProxyHooksActuallyRaiseProxyRateLimitError: from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.hooks.dynamic_rate_limiter_v3 import ( - _PROXY_DynamicRateLimitHandlerV3, + PROXY_DynamicRateLimitHandlerV3, ) # Bypass __init__ — we want to inject a stub v3_limiter without # paying for the full handler setup. - handler = _PROXY_DynamicRateLimitHandlerV3.__new__( - _PROXY_DynamicRateLimitHandlerV3 + handler = PROXY_DynamicRateLimitHandlerV3.__new__( + PROXY_DynamicRateLimitHandlerV3 ) v3_limiter = MagicMock() v3_limiter.window_size = 60 @@ -836,12 +836,12 @@ class TestProxyHooksActuallyRaiseProxyRateLimitError: from litellm.caching.caching import DualCache from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.hooks.max_budget_per_session_limiter import ( - _PROXY_MaxBudgetPerSessionHandler, + PROXY_MaxBudgetPerSessionHandler, ) internal_cache = MagicMock() internal_cache.async_get_cache = AsyncMock(return_value=10.0) - handler = _PROXY_MaxBudgetPerSessionHandler( + handler = PROXY_MaxBudgetPerSessionHandler( internal_usage_cache=internal_cache, ) user_api_key_dict = UserAPIKeyAuth( @@ -876,14 +876,14 @@ class TestProxyHooksActuallyRaiseProxyRateLimitError: from litellm.proxy.hooks.batch_rate_limiter import ( BatchFileUsage, - _PROXY_BatchRateLimiter, + PROXY_BatchRateLimiter, ) # Inject a parallel_request_limiter mock with a usable window_size so # the helper's str(window_size) call doesn't NameError. parallel_limiter = MagicMock() parallel_limiter.window_size = 60 - handler = _PROXY_BatchRateLimiter( + handler = PROXY_BatchRateLimiter( internal_usage_cache=MagicMock(), parallel_request_limiter=parallel_limiter, ) @@ -1117,10 +1117,10 @@ class TestProxyHooksWireTypeCorrectly: from unittest.mock import MagicMock from litellm.proxy.hooks.parallel_request_limiter import ( - _PROXY_MaxParallelRequestsHandler, + PROXY_MaxParallelRequestsHandler, ) - handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=MagicMock()) + handler = PROXY_MaxParallelRequestsHandler(internal_usage_cache=MagicMock()) with pytest.raises(ProxyRateLimitError) as exc_info: handler.raise_rate_limit_error() assert exc_info.value.rate_limit_type == "concurrent_requests" @@ -1129,10 +1129,10 @@ class TestProxyHooksWireTypeCorrectly: from unittest.mock import MagicMock from litellm.proxy.hooks.parallel_request_limiter import ( - _PROXY_MaxParallelRequestsHandler, + PROXY_MaxParallelRequestsHandler, ) - handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=MagicMock()) + handler = PROXY_MaxParallelRequestsHandler(internal_usage_cache=MagicMock()) with pytest.raises(ProxyRateLimitError) as exc_info: handler.raise_rate_limit_error( additional_details="tpm-zero", @@ -1174,10 +1174,10 @@ class TestProxyHooksWireTypeCorrectly: from unittest.mock import MagicMock from litellm.proxy.hooks.parallel_request_limiter_v3 import ( - _PROXY_MaxParallelRequestsHandler_v3, + PROXY_MaxParallelRequestsHandler_v3, ) - handler = _PROXY_MaxParallelRequestsHandler_v3( + handler = PROXY_MaxParallelRequestsHandler_v3( internal_usage_cache=MagicMock(), ) # Minimal RateLimitResponse + descriptors shape that the handler @@ -1225,10 +1225,10 @@ class TestProxyHooksWireTypeCorrectly: from unittest.mock import MagicMock from litellm.proxy.hooks.parallel_request_limiter_v3 import ( - _PROXY_MaxParallelRequestsHandler_v3, + PROXY_MaxParallelRequestsHandler_v3, ) - handler = _PROXY_MaxParallelRequestsHandler_v3( + handler = PROXY_MaxParallelRequestsHandler_v3( internal_usage_cache=MagicMock(), ) response = { @@ -1269,12 +1269,12 @@ class TestProxyHooksWireTypeCorrectly: from litellm.proxy.hooks.batch_rate_limiter import ( BatchFileUsage, - _PROXY_BatchRateLimiter, + PROXY_BatchRateLimiter, ) prl = MagicMock() prl.window_size = 60 - handler = _PROXY_BatchRateLimiter( + handler = PROXY_BatchRateLimiter( internal_usage_cache=MagicMock(), parallel_request_limiter=prl, ) @@ -1312,12 +1312,12 @@ class TestProxyHooksWireTypeCorrectly: from litellm.proxy.hooks.batch_rate_limiter import ( BatchFileUsage, - _PROXY_BatchRateLimiter, + PROXY_BatchRateLimiter, ) prl = MagicMock() prl.window_size = 60 - handler = _PROXY_BatchRateLimiter( + handler = PROXY_BatchRateLimiter( internal_usage_cache=MagicMock(), parallel_request_limiter=prl, ) @@ -1558,7 +1558,7 @@ class TestBudgetExceededErrorLlmProviderEnrichment: {"use_x_forwarded_for": False}, ), patch( - "litellm.proxy.auth.auth_exception_handler._get_request_ip_address", + "litellm.proxy.auth.auth_exception_handler.get_request_ip_address", return_value="127.0.0.1", ), ): diff --git a/tests/unit/test_router/test_router.py b/tests/unit/test_router/test_router.py index 8f3133dcaa4..ec9e4861fdd 100644 --- a/tests/unit/test_router/test_router.py +++ b/tests/unit/test_router/test_router.py @@ -65,6 +65,7 @@ from litellm.router_utils.cooldown_handlers import ( ) from litellm.router_utils.fallback_event_handlers import DISABLE_FALLBACKS_METADATA_KEY from litellm.router_utils.router_callbacks.track_deployment_metrics import get_deployment_successes_for_current_minute +from litellm.scheduler import FlowItem from litellm.types.llms.openai import ChatCompletionRequest from litellm.types.router import ( CustomRoutingStrategyBase, @@ -19949,6 +19950,64 @@ def test_access_windows_filter_reserved_deployments_method(): ] == ["reserved-deployment", "open-deployment"] +def _scheduled_router(timeout: float) -> Router: + return Router( + model_list=[ + { + "model_name": "sched-model", + "litellm_params": {"model": "openai/sched-model", "api_key": "sk-fake", "mock_response": "hi"}, + "model_info": {"id": "sched-deployment"}, + } + ], + timeout=timeout, + ) + + +async def _send_scheduled_chat(router: Router, priority: int) -> object: + return await router.acompletion( + model="sched-model", messages=[{"role": "user", "content": "hi"}], priority=priority + ) + + +async def _send_scheduled_text(router: Router, priority: int) -> object: + return await router.atext_completion(model="sched-model", prompt="hi", priority=priority) + + +@pytest.mark.parametrize( + "send", [_send_scheduled_chat, _send_scheduled_text], ids=["schedule_acompletion", "schedule_factory"] +) +@pytest.mark.asyncio +async def test_admitted_prioritized_request_does_not_block_later_request_during_cooldown( + send: Callable[[Router, int], Awaitable[object]], +): + from litellm.types.router import RouterRateLimitError + + router: Final = _scheduled_router(timeout=1) + await send(router, 1) + _cool_down(router, "sched-deployment") + + with pytest.raises(RouterRateLimitError, match="cooldown"): + await send(router, 2) + + +@pytest.mark.parametrize("stop_waiting", ["cancel", "timeout"]) +@pytest.mark.asyncio +async def test_prioritized_request_leaves_queue_when_it_stops_waiting(stop_waiting: Literal["cancel", "timeout"]): + router: Final = _scheduled_router(timeout=0.5) + _cool_down(router, "sched-deployment") + await router.scheduler.add_request(FlowItem(priority=0, request_id="head-of-queue", model_name="sched-model")) + waiting: Final = asyncio.create_task(_send_scheduled_chat(router, 5)) + await asyncio.sleep(0.05) + assert len(await router.scheduler.get_queue("sched-model")) == 2 + + if stop_waiting == "cancel": + waiting.cancel() + with pytest.raises(asyncio.CancelledError if stop_waiting == "cancel" else litellm.Timeout): + await waiting + + assert await router.scheduler.get_queue("sched-model") == [(0, "head-of-queue")] + + @pytest.mark.asyncio async def test_bare_model_group_served_by_wildcard_deployment_uses_provider_prefixed_fallback_key() -> None: """Claude Code sends the bare "claude-sonnet-4-6" to /v1/messages; routing serves it through the diff --git a/tests/unit/test_scheduler.py b/tests/unit/test_scheduler.py index 553fb44cb27..e0968e9b2cf 100644 --- a/tests/unit/test_scheduler.py +++ b/tests/unit/test_scheduler.py @@ -3,7 +3,10 @@ import asyncio import importlib +import json import os +from collections.abc import Sequence +from typing import Final import pytest @@ -26,8 +29,8 @@ async def test_scheduler_diff_model_names(): await scheduler.add_request(item1) await scheduler.add_request(item2) - assert await scheduler.poll(id="10", model_name="gpt-3.5-turbo", health_deployments=[{"key": "value"}]) == True - assert await scheduler.poll(id="11", model_name="gpt-4", health_deployments=[{"key": "value"}]) == True + assert await scheduler.poll(request=item1, health_deployments=[{"key": "value"}]) == True + assert await scheduler.poll(request=item2, health_deployments=[{"key": "value"}]) == True @pytest.mark.asyncio @@ -50,7 +53,7 @@ async def test_scheduler_poll_persists_queue_to_cache(): await scheduler.add_request(item1) await scheduler.add_request(item2) - await scheduler.poll(id="10", model_name="gpt-3.5-turbo", health_deployments=[]) + await scheduler.poll(request=item1, health_deployments=[]) queue_key = f"{SchedulerCacheKeys.queue.value}:{item1.model_name}" updated_queue = redis_cache.store[queue_key] @@ -145,6 +148,190 @@ async def test_scheduler_queue_cleanup_on_timeout(): assert queue_after[0][1] == "req-0", "Expected req-0 (priority 0) to be at front" +@pytest.mark.asyncio +async def test_poll_admits_request_missing_from_queue_while_a_deployment_is_healthy(): + scheduler: Final = Scheduler() + + assert await scheduler.poll( + request=FlowItem(priority=1, request_id="erased-by-concurrent-write", model_name="sched-model"), + health_deployments=[{"model_info": {"id": "a"}}], + ) + + +@pytest.mark.asyncio +async def test_poll_during_cooldown_admits_only_the_head_of_the_queue(): + scheduler: Final = Scheduler() + later: Final = FlowItem(priority=2, request_id="later", model_name="sched-model") + head: Final = FlowItem(priority=1, request_id="head", model_name="sched-model") + await scheduler.add_request(later) + await scheduler.add_request(head) + + assert not await scheduler.poll(request=later, health_deployments=[]) + assert await scheduler.poll(request=head, health_deployments=[]) + assert await scheduler.get_queue("sched-model") == [(2, "later")] + + +@pytest.mark.asyncio +async def test_poll_during_cooldown_re_enqueues_a_request_a_concurrent_writer_erased_behind_the_head(): + scheduler: Final = Scheduler() + still_queued: Final = FlowItem(priority=0, request_id="still-queued", model_name="sched-model") + erased: Final = FlowItem(priority=1, request_id="erased-by-concurrent-write", model_name="sched-model") + await scheduler.add_request(still_queued) + + assert not await scheduler.poll(request=erased, health_deployments=[]) + assert await scheduler.get_queue("sched-model") == [(0, "still-queued"), (1, "erased-by-concurrent-write")] + assert await scheduler.poll(request=still_queued, health_deployments=[]) + assert await scheduler.poll(request=erased, health_deployments=[]) + assert await scheduler.get_queue("sched-model") == [] + + +@pytest.mark.asyncio +async def test_poll_during_cooldown_admits_an_erased_request_that_outranks_the_queue(): + scheduler: Final = Scheduler() + await scheduler.add_request(FlowItem(priority=2, request_id="still-queued", model_name="sched-model")) + + assert await scheduler.poll( + request=FlowItem(priority=0, request_id="erased-urgent", model_name="sched-model"), health_deployments=[] + ) + assert await scheduler.get_queue("sched-model") == [(2, "still-queued")] + + +class _ExpiringCache: + def __init__(self) -> None: + self.store: dict[str, object] = {} + + async def async_get_cache(self, key: str, **kwargs: object) -> object: + return self.store.get(key) + + async def async_set_cache(self, key: str, value: object, **kwargs: object) -> None: + self.store[key] = value + + +@pytest.mark.asyncio +async def test_wait_for_turn_during_cooldown_survives_the_queue_key_expiring(): + cache: Final = _ExpiringCache() + scheduler: Final = Scheduler(redis_cache=cache) + scheduler.cache.in_memory_cache.cache_dict.clear() + + async def no_healthy_deployments_after_the_key_expired() -> Sequence[object]: + cache.store.clear() + scheduler.cache.in_memory_cache.cache_dict.clear() + return () + + await scheduler.wait_for_turn( + request=FlowItem(priority=1, request_id="sole-waiter", model_name="sched-model"), + timeout=5, + get_healthy_deployments=no_healthy_deployments_after_the_key_expired, + ) + + assert await scheduler.get_queue("sched-model") == [] + + + +class _JsonRoundTripRedisCache: + def __init__(self) -> None: + self.store: dict[str, str] = {} + + async def async_get_cache(self, key: str, **kwargs: object) -> object: + raw: Final = self.store.get(key) + return None if raw is None else json.loads(raw) + + async def async_set_cache(self, key: str, value: object, **kwargs: object) -> None: + self.store[key] = json.dumps(value) + + +@pytest.mark.asyncio +async def test_second_replica_enqueues_behind_a_queue_decoded_from_redis(): + redis_cache: Final = _JsonRoundTripRedisCache() + replica_a: Final = Scheduler(redis_cache=redis_cache) + replica_b: Final = Scheduler(redis_cache=redis_cache) + waiting_on_a: Final = FlowItem(priority=1, request_id="waiting-on-a", model_name="sched-model") + urgent_on_b: Final = FlowItem(priority=0, request_id="urgent-on-b", model_name="sched-model") + await replica_a.add_request(waiting_on_a) + + await replica_b.add_request(urgent_on_b) + + assert await replica_b.get_queue("sched-model") == [(0, "urgent-on-b"), (1, "waiting-on-a")] + assert await replica_b.poll(request=urgent_on_b, health_deployments=[]) + assert await replica_b.poll(request=waiting_on_a, health_deployments=[]) + await replica_b.remove_request(request_id="waiting-on-a", model_name="sched-model") + assert await replica_b.get_queue("sched-model") == [] + +class _PausesAfterEnqueueScheduler(Scheduler): + def __init__(self) -> None: + super().__init__() + self.enqueued: Final = asyncio.Event() + + async def add_request(self, request: FlowItem) -> None: + await super().add_request(request) + self.enqueued.set() + await asyncio.Event().wait() + + +async def _no_healthy_deployments() -> Sequence[object]: + return () + + +@pytest.mark.asyncio +async def test_wait_for_turn_removes_entry_when_cancelled_mid_enqueue(): + scheduler: Final = _PausesAfterEnqueueScheduler() + waiting: Final = asyncio.create_task( + scheduler.wait_for_turn( + request=FlowItem(priority=1, request_id="cancelled", model_name="sched-model"), + timeout=5, + get_healthy_deployments=_no_healthy_deployments, + ) + ) + await scheduler.enqueued.wait() + assert await scheduler.get_queue("sched-model") == [(1, "cancelled")] + + waiting.cancel() + with pytest.raises(asyncio.CancelledError): + await waiting + + assert await scheduler.get_queue("sched-model") == [] + + +class _HeldRemovalScheduler(Scheduler): + def __init__(self) -> None: + super().__init__() + self.removing: Final = asyncio.Event() + self.finish_removal: Final = asyncio.Event() + + async def remove_request(self, request_id: str, model_name: str) -> None: + self.removing.set() + await self.finish_removal.wait() + await super().remove_request(request_id=request_id, model_name=model_name) + + +@pytest.mark.asyncio +async def test_wait_for_turn_finishes_removal_when_cancelled_again_during_cleanup(): + scheduler: Final = _HeldRemovalScheduler() + await scheduler.add_request(FlowItem(priority=0, request_id="head", model_name="sched-model")) + polling: Final = asyncio.Event() + + async def no_healthy_deployments() -> Sequence[object]: + polling.set() + return () + + waiting: Final = asyncio.create_task( + scheduler.wait_for_turn( + request=FlowItem(priority=1, request_id="cancelled", model_name="sched-model"), + timeout=5, + get_healthy_deployments=no_healthy_deployments, + ) + ) + await polling.wait() + waiting.cancel() + await scheduler.removing.wait() + waiting.cancel() + scheduler.finish_removal.set() + with pytest.raises(asyncio.CancelledError): + await waiting + + assert await scheduler.get_queue("sched-model") == [(0, "head")] + + @pytest.fixture(autouse=True) def _vcr_outcome_gate(request, vcr): install_live_call_probe(request, vcr) diff --git a/tests/unit/test_utils.py b/tests/unit/test_utils.py index ac5dc6e7c45..7551a310977 100644 --- a/tests/unit/test_utils.py +++ b/tests/unit/test_utils.py @@ -1064,6 +1064,7 @@ def test_aaamodel_prices_and_context_window_json_is_valid(): "/vertex_ai/live", "/v1/listen", "/v1/systemone", + "/v1/decisions", "/v1beta/interactions", ], }, diff --git a/tests/unit/test_utils_get_optional_params.py b/tests/unit/test_utils_get_optional_params.py index 92ce66c6ebe..d9d049191a1 100644 --- a/tests/unit/test_utils_get_optional_params.py +++ b/tests/unit/test_utils_get_optional_params.py @@ -312,8 +312,8 @@ def test_azure_ai_mistral_optional_params(): assert "user" not in optional_params -def test_vertex_ai_llama_3_optional_params(): - litellm.vertex_llama3_models = ["meta/llama3-405b-instruct-maas"] +def test_vertex_ai_llama_3_optional_params(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(litellm, "vertex_llama3_models", {"meta/llama3-405b-instruct-maas"}) litellm.drop_params = True optional_params = get_optional_params( model="meta/llama3-405b-instruct-maas", @@ -325,8 +325,8 @@ def test_vertex_ai_llama_3_optional_params(): assert "user" not in optional_params -def test_vertex_ai_mistral_optional_params(): - litellm.vertex_mistral_models = ["mistral-large@2407"] +def test_vertex_ai_mistral_optional_params(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(litellm, "vertex_mistral_models", {"mistral-large@2407"}) litellm.drop_params = True optional_params = get_optional_params( model="mistral-large@2407", diff --git a/tests/unit/test_video_generation.py b/tests/unit/test_video_generation.py index dc15e82597d..5c4bc8ea652 100644 --- a/tests/unit/test_video_generation.py +++ b/tests/unit/test_video_generation.py @@ -2398,7 +2398,7 @@ async def test_edit_and_extension_read_cached_body_after_auth_consumes_stream( import litellm.proxy.video_endpoints.endpoints as endpoints from litellm.proxy._types import ProxyException, UserAPIKeyAuth - from litellm.proxy.common_utils.http_parsing_utils import _read_request_body + from litellm.proxy.common_utils.http_parsing_utils import read_request_body body = urlencode(form).encode() stream = {"sent": False} @@ -2423,7 +2423,7 @@ async def test_edit_and_extension_read_cached_body_after_auth_consumes_stream( receive, ) - await _read_request_body(request=request) + await read_request_body(request=request) handler = getattr(endpoints, handler_name) with pytest.raises(ProxyException) as exc_info: diff --git a/tests/unit/tracing/test_exporter.py b/tests/unit/tracing/test_exporter.py index de92a49e361..ebbe0ad97e5 100644 --- a/tests/unit/tracing/test_exporter.py +++ b/tests/unit/tracing/test_exporter.py @@ -9,6 +9,57 @@ import pytest from litellm.tracing.exporter import MAX_BUFFER_EVENTS, MAX_EVENT_BYTES, ExportFailure, LensExporter, encode_record +@pytest.mark.asyncio +@pytest.mark.parametrize("failed", (False, True), ids=("success-event", "failure-event")) +@pytest.mark.parametrize( + ("tags", "expected"), + ( + pytest.param((7,), ("7",), id="numeric"), + pytest.param(("env:prod", 0, -7, 2.5, True, None), ("env:prod", "0", "-7", "2.5", "True", "None"), id="mixed"), + pytest.param(("env:prod", "agent:research"), ("env:prod", "agent:research"), id="strings"), + ), +) +async def test_request_tags_are_normalized_without_dropping_the_record( + failed: bool, tags: tuple[object, ...], expected: tuple[str, ...] +) -> None: + requests: Final = asyncio.Queue[httpx.Request]() + + def accept(request: httpx.Request) -> httpx.Response: + requests.put_nowait(request) + return httpx.Response(204) + + async with httpx.AsyncClient(base_url="http://lens", transport=httpx.MockTransport(accept)) as client: + exporter: Final = LensExporter(client) + exporter.start() + callback: Final = exporter.async_log_failure_event if failed else exporter.async_log_success_event + await callback( + { + "response_cost": 0.12, + "standard_logging_object": { + "id": "tagged-request", + "status": "failure" if failed else "success", + "response_cost": 0.12, + "request_tags": list(tags), + }, + }, + None, + None, + None, + ) + await exporter.aclose() + assert exporter.rows_written == 1 + assert exporter.rows_dropped == 0 + assert requests.qsize() == 1 + request: Final = requests.get_nowait() + rows: Final = json.loads(request.content) + assert request.url.path == "/internal/spend" + assert len(rows) == 1 + assert rows[0]["request_id"] == "tagged-request" + assert rows[0]["status"] == ("failure" if failed else "success") + assert rows[0]["spend"] == 0.12 + assert rows[0]["request_tags"] == list(expected) + + @pytest.mark.asyncio async def test_request_export_ignores_unrelated_model_metadata_and_preserves_billing() -> None: received: Final = asyncio.Future[httpx.Request]() diff --git a/tests/unit/tracing/test_remote.py b/tests/unit/tracing/test_remote.py index 2572c2a2044..92922dbd3ca 100644 --- a/tests/unit/tracing/test_remote.py +++ b/tests/unit/tracing/test_remote.py @@ -168,7 +168,7 @@ async def test_gateway_cannot_relay_otlp_or_write_arbitrary_tables() -> None: await store.ensure_schema() with pytest.raises(RuntimeError, match="directly"): await store.ingest(b"{}", "application/json", {}) - with pytest.raises(ValueError, match="request records"): + with pytest.raises(ValueError, match="request records and feedback"): await store.insert_rows("otel_traces", ()) @@ -185,3 +185,18 @@ async def test_request_records_use_the_internal_service_endpoint() -> None: request: Final = requests.get_nowait() assert request.url.path == "/internal/spend" assert json.loads(request.content) == [{"request_id": "r"}] + + +@pytest.mark.asyncio +async def test_feedback_rows_use_the_internal_feedback_endpoint() -> None: + requests: Final = asyncio.Queue[httpx.Request]() + + def accept(request: httpx.Request) -> httpx.Response: + requests.put_nowait(request) + return httpx.Response(204) + + async with httpx.AsyncClient(base_url="http://lens", transport=httpx.MockTransport(accept)) as client: + await RemoteTraceStore(client).insert_rows("lens_feedback", ({"TraceId": "t", "Score": 2},)) + request: Final = requests.get_nowait() + assert request.url.path == "/internal/feedback" + assert json.loads(request.content) == [{"TraceId": "t", "Score": 2}] diff --git a/tests/unit/types/test_openai_decisions.py b/tests/unit/types/test_openai_decisions.py new file mode 100644 index 00000000000..d9387d644d5 --- /dev/null +++ b/tests/unit/types/test_openai_decisions.py @@ -0,0 +1,150 @@ +from __future__ import annotations + +from collections.abc import Mapping +from typing import Final + +import pytest +from pydantic import TypeAdapter, ValidationError + +from litellm.types.openai_decisions import ( + ChoiceAnswer, + DecisionsRequestBody, + DecisionsResponse, + RefusalAnswer, + ScoreQuestion, +) + +_REQUEST: Final[Mapping[str, object]] = { + "input": [ + { + "role": "user", + "content": [ + {"type": "input_text", "text": "Is this receipt a valid business expense?"}, + {"type": "input_image", "image_url": "https://example.com/receipt.png", "detail": "high"}, + ], + } + ], + "questions": [ + {"type": "predicate", "name": "is_expense", "instructions": "Is this a business expense?"}, + { + "type": "choice", + "name": "approve", + "instructions": "Should this be approved?", + "choices": [{"value": True, "description": "approve"}, {"value": False, "description": "reject"}], + }, + { + "type": "score", + "name": "risk", + "instructions": "How risky is this expense?", + "levels": [{"label": "low"}, {"label": "high", "description": "needs a manager"}], + }, + ], + "safety_identifier": "user-123", +} +_RESPONSE: Final[Mapping[str, object]] = { + "model": "gpt-6-luna", + "answers": [ + {"type": "predicate", "name": "is_expense", "probability": 0.92}, + { + "type": "choice", + "name": "approve", + "choice": True, + "probabilities": [{"value": True, "probability": 0.7}, {"value": False, "probability": 0.3}], + "confidence": 0.7, + }, + {"type": "refusal", "name": "risk"}, + ], + "usage": { + "input_tokens": 120, + "input_tokens_details": {"cached_tokens": 100, "cache_write_tokens": 0}, + "output_tokens": 12, + "output_tokens_details": {"reasoning_tokens": 4}, + "total_tokens": 132, + }, +} +_REQUEST_ADAPTER: Final[TypeAdapter[DecisionsRequestBody]] = TypeAdapter(DecisionsRequestBody) +_RESPONSE_ADAPTER: Final[TypeAdapter[DecisionsResponse]] = TypeAdapter(DecisionsResponse) + + +def test_the_documented_request_round_trips_with_its_boolean_choices_and_image_part() -> None: + request: Final = _REQUEST_ADAPTER.validate_python(_REQUEST) + + assert request.model_dump(mode="json", exclude_none=True) == _REQUEST + assert isinstance(request.questions[2], ScoreQuestion) + + +def test_the_documented_response_keeps_answer_order_refusals_and_token_details() -> None: + response: Final = _RESPONSE_ADAPTER.validate_python(_RESPONSE) + + assert response.model_dump(mode="json") == _RESPONSE + assert isinstance(response.answers[1], ChoiceAnswer) + assert isinstance(response.answers[2], RefusalAnswer) + + +_OFF_SPEC: Final[tuple[tuple[str, object], ...]] = ( + ("questions", [{"type": "noul", "name": "q", "instructions": "x"}]), + ("questions", [{"type": "choice", "name": "q", "instructions": "x", "choices": [{"value": 1}]}]), + ("questions", [{"type": "score", "name": "q", "levels": [{"label": "low"}]}]), + ("input", {"state": "not an OpenAI input"}), +) + + +@pytest.mark.parametrize(("field", "value"), _OFF_SPEC) +def test_requests_off_the_spec_are_rejected(field: str, value: object) -> None: + with pytest.raises(ValidationError): + _REQUEST_ADAPTER.validate_python({**_REQUEST, field: value}) + + +_EMPTY_COLLECTIONS: Final[tuple[tuple[str, list[object]], ...]] = ( + ("questions", []), + ("questions", [{"type": "choice", "name": "q", "instructions": "x", "choices": []}]), + ("questions", [{"type": "score", "name": "q", "instructions": "x", "levels": []}]), +) + + +@pytest.mark.parametrize(("field", "value"), _EMPTY_COLLECTIONS) +def test_empty_collections_are_left_for_the_provider_to_judge(field: str, value: list[object]) -> None: + request: Final = _REQUEST_ADAPTER.validate_python({**_REQUEST, field: value}) + + assert request.model_dump(mode="json", exclude_none=True)[field] == value + + +def test_choice_values_keep_their_type_so_a_string_true_and_a_boolean_true_stay_distinct() -> None: + answer: Final = { + "type": "choice", + "name": "approve", + "choice": "true", + "probabilities": [{"value": "true", "probability": 0.6}, {"value": True, "probability": 0.4}], + "confidence": 0.6, + } + + response: Final = _RESPONSE_ADAPTER.validate_python({**_RESPONSE, "answers": [answer]}) + + assert response.model_dump(mode="json")["answers"] == [answer] + with pytest.raises(ValidationError): + _RESPONSE_ADAPTER.validate_python({**_RESPONSE, "answers": [{**answer, "choice": 1}]}) + + +_OFF_SPEC_RESPONSE: Final[tuple[str, ...]] = ("model", "usage") + + +@pytest.mark.parametrize("field", _OFF_SPEC_RESPONSE) +def test_responses_missing_a_required_field_are_rejected(field: str) -> None: + with pytest.raises(ValidationError): + _RESPONSE_ADAPTER.validate_python({k: v for k, v in _RESPONSE.items() if k != field}) + + +def test_usage_without_token_details_is_rejected() -> None: + usage: Final = {"input_tokens": 120, "output_tokens": 12, "total_tokens": 132} + + with pytest.raises(ValidationError): + _RESPONSE_ADAPTER.validate_python({**_RESPONSE, "usage": usage}) + + +def test_hidden_params_live_outside_the_wire_body() -> None: + response: Final = _RESPONSE_ADAPTER.validate_python(_RESPONSE) + + response.set_hidden_params({"custom_llm_provider": "openai"}) + + assert response.hidden_params == {"custom_llm_provider": "openai"} + assert "_hidden_params" not in response.model_dump(mode="json") diff --git a/ui/litellm-dashboard/eslint-suppressions.json b/ui/litellm-dashboard/eslint-suppressions.json index 280d3cc4115..1e8e6ee0e39 100644 --- a/ui/litellm-dashboard/eslint-suppressions.json +++ b/ui/litellm-dashboard/eslint-suppressions.json @@ -1116,11 +1116,6 @@ "count": 1 } }, - "src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.tsx": { - "local/no-complex-jsx-arrow": { - "count": 2 - } - }, "src/app/(dashboard)/usage/_components/components/UsageAIChatPanel.tsx": { "no-nested-ternary": { "count": 1 @@ -1130,17 +1125,11 @@ } }, "src/app/(dashboard)/usage/_components/components/UsagePageView.tsx": { - "local/no-complex-jsx-arrow": { - "count": 2 - }, - "max-lines": { - "count": 1 - }, "react-hooks/purity": { "count": 1 }, "react-hooks/set-state-in-effect": { - "count": 2 + "count": 1 } }, "src/app/(dashboard)/users/_components/BulkEditUsers.tsx": { diff --git a/ui/litellm-dashboard/public/assets/moyai/logos/anthropic.svg b/ui/litellm-dashboard/public/assets/moyai/logos/anthropic.svg new file mode 100644 index 00000000000..a37f591fb76 --- /dev/null +++ b/ui/litellm-dashboard/public/assets/moyai/logos/anthropic.svg @@ -0,0 +1,5 @@ + + + + + diff --git a/ui/litellm-dashboard/public/assets/moyai/logos/bedrock.svg b/ui/litellm-dashboard/public/assets/moyai/logos/bedrock.svg new file mode 100644 index 00000000000..e0f929a7a97 --- /dev/null +++ b/ui/litellm-dashboard/public/assets/moyai/logos/bedrock.svg @@ -0,0 +1 @@ +Bedrock \ No newline at end of file diff --git a/ui/litellm-dashboard/public/assets/moyai/logos/deepseek.svg b/ui/litellm-dashboard/public/assets/moyai/logos/deepseek.svg new file mode 100644 index 00000000000..c4754047da2 --- /dev/null +++ b/ui/litellm-dashboard/public/assets/moyai/logos/deepseek.svg @@ -0,0 +1,25 @@ + + + + + + diff --git a/ui/litellm-dashboard/public/assets/moyai/logos/fireworks.svg b/ui/litellm-dashboard/public/assets/moyai/logos/fireworks.svg new file mode 100644 index 00000000000..a23445cf94b --- /dev/null +++ b/ui/litellm-dashboard/public/assets/moyai/logos/fireworks.svg @@ -0,0 +1 @@ +Fireworks \ No newline at end of file diff --git a/ui/litellm-dashboard/public/assets/moyai/logos/google.svg b/ui/litellm-dashboard/public/assets/moyai/logos/google.svg new file mode 100644 index 00000000000..7bc4a38ce7a --- /dev/null +++ b/ui/litellm-dashboard/public/assets/moyai/logos/google.svg @@ -0,0 +1,2 @@ + + \ No newline at end of file diff --git a/ui/litellm-dashboard/public/assets/moyai/logos/hermes.png b/ui/litellm-dashboard/public/assets/moyai/logos/hermes.png new file mode 100644 index 00000000000..de47b728d12 Binary files /dev/null and b/ui/litellm-dashboard/public/assets/moyai/logos/hermes.png differ diff --git a/ui/litellm-dashboard/public/assets/moyai/logos/langchain.svg b/ui/litellm-dashboard/public/assets/moyai/logos/langchain.svg new file mode 100644 index 00000000000..939b79989a7 --- /dev/null +++ b/ui/litellm-dashboard/public/assets/moyai/logos/langchain.svg @@ -0,0 +1 @@ +LangChain \ No newline at end of file diff --git a/ui/litellm-dashboard/public/assets/moyai/logos/mistral.svg b/ui/litellm-dashboard/public/assets/moyai/logos/mistral.svg new file mode 100644 index 00000000000..8e03e244bf1 --- /dev/null +++ b/ui/litellm-dashboard/public/assets/moyai/logos/mistral.svg @@ -0,0 +1 @@ +Mistral \ No newline at end of file diff --git a/ui/litellm-dashboard/public/assets/moyai/logos/openai.svg b/ui/litellm-dashboard/public/assets/moyai/logos/openai.svg new file mode 100644 index 00000000000..52dad8269ec --- /dev/null +++ b/ui/litellm-dashboard/public/assets/moyai/logos/openai.svg @@ -0,0 +1,5 @@ + + + + + \ No newline at end of file diff --git a/ui/litellm-dashboard/public/assets/moyai/logos/opencode.svg b/ui/litellm-dashboard/public/assets/moyai/logos/opencode.svg new file mode 100644 index 00000000000..7ed0af003bb --- /dev/null +++ b/ui/litellm-dashboard/public/assets/moyai/logos/opencode.svg @@ -0,0 +1,16 @@ + + + + + + + + + + + + + + + + diff --git a/ui/litellm-dashboard/public/assets/moyai/logos/xai.svg b/ui/litellm-dashboard/public/assets/moyai/logos/xai.svg new file mode 100644 index 00000000000..9491b192fd5 --- /dev/null +++ b/ui/litellm-dashboard/public/assets/moyai/logos/xai.svg @@ -0,0 +1,28 @@ + + + + + + + + + + + + + diff --git a/ui/litellm-dashboard/public/assets/moyai/moyai-head.svg b/ui/litellm-dashboard/public/assets/moyai/moyai-head.svg new file mode 100644 index 00000000000..6d69078c4e3 --- /dev/null +++ b/ui/litellm-dashboard/public/assets/moyai/moyai-head.svg @@ -0,0 +1,19 @@ + + + + + + + + + + + + + + + + + + + diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/tags/useTagSummary.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/tags/useTagSummary.ts new file mode 100644 index 00000000000..9ee14469374 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/tags/useTagSummary.ts @@ -0,0 +1,29 @@ +import { $api } from "@/lib/http/api"; +import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; + +const isoDay = (date: Date) => { + const month = String(date.getMonth() + 1).padStart(2, "0"); + const day = String(date.getDate()).padStart(2, "0"); + return `${date.getFullYear()}-${month}-${day}`; +}; + +/** Per-tag spend, tokens and requests for a date range, including the `User-Agent:` tags agents are read from. */ +export const useTagSummary = (startTime: Date | null, endTime: Date | null, enabled = true) => { + const { accessToken } = useAuthorized(); + return $api.useQuery( + "get", + "/tag/summary", + { + params: { + query: { + start_date: startTime ? isoDay(startTime) : "", + end_date: endTime ? isoDay(endTime) : "", + }, + }, + }, + { + enabled: enabled && Boolean(accessToken && startTime && endTime), + select: (data) => data.results, + }, + ); +}; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/useDisableBlogPosts.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/useDisableBlogPosts.ts index a7b37b78d42..8bffbe9f741 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/useDisableBlogPosts.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/useDisableBlogPosts.ts @@ -29,5 +29,5 @@ function getSnapshot() { } export function useDisableBlogPosts() { - return useSyncExternalStore(subscribe, getSnapshot); + return useSyncExternalStore(subscribe, getSnapshot, () => false); } diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/useDisableBouncingIcon.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/useDisableBouncingIcon.ts index f5d8087ebe7..0ec94f9915f 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/useDisableBouncingIcon.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/useDisableBouncingIcon.ts @@ -29,5 +29,5 @@ function getSnapshot() { } export function useDisableBouncingIcon() { - return useSyncExternalStore(subscribe, getSnapshot); + return useSyncExternalStore(subscribe, getSnapshot, () => false); } diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/useDisableShowNewBadge.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/useDisableShowNewBadge.ts index d0a618e27ba..0feba3a27fd 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/useDisableShowNewBadge.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/useDisableShowNewBadge.ts @@ -31,5 +31,5 @@ function getSnapshot() { } export function useDisableShowNewBadge() { - return useSyncExternalStore(subscribe, getSnapshot); + return useSyncExternalStore(subscribe, getSnapshot, () => false); } diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/useDisableShowPrompts.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/useDisableShowPrompts.ts index 801fbdbb99d..872da097868 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/useDisableShowPrompts.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/useDisableShowPrompts.ts @@ -31,5 +31,5 @@ function getSnapshot() { } export function useDisableShowPrompts() { - return useSyncExternalStore(subscribe, getSnapshot); + return useSyncExternalStore(subscribe, getSnapshot, () => false); } diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/useHideAutoRouterAnnouncement.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/useHideAutoRouterAnnouncement.ts index c213c4834bc..61b9b3eea7a 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/useHideAutoRouterAnnouncement.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/useHideAutoRouterAnnouncement.ts @@ -31,5 +31,5 @@ function getSnapshot() { } export function useHideAutoRouterAnnouncement() { - return useSyncExternalStore(subscribe, getSnapshot); + return useSyncExternalStore(subscribe, getSnapshot, () => false); } diff --git a/ui/litellm-dashboard/src/app/(dashboard)/model-insights/_components/ModelInsightsView.tsx b/ui/litellm-dashboard/src/app/(dashboard)/model-insights/_components/ModelInsightsView.tsx index 21a707b8c42..ec3b5066d2f 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/model-insights/_components/ModelInsightsView.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/model-insights/_components/ModelInsightsView.tsx @@ -2,8 +2,8 @@ import { Page } from "@/components/shared/Page"; import React from "react"; -import { Bar, BarChart, CartesianGrid, XAxis, YAxis } from "recharts"; import { ArrowDownRight, ArrowUpRight, BarChart3, Minus } from "lucide-react"; +import { StackedUsageChart } from "@/components/shared/charts"; import { apiClient } from "@/components/networking"; import { extractErrorMessage } from "@/utils/errorUtils"; @@ -11,7 +11,6 @@ import { ProviderLogo } from "@/components/molecules/models/ProviderLogo"; import { PageHeader, PageHeaderDescription, PageHeaderTitle } from "@/components/shared/PageHeader"; import { Alert, AlertDescription, AlertTitle } from "@/components/ui/alert"; import { Card, CardContent, CardDescription, CardHeader, CardTitle } from "@/components/ui/card"; -import { ChartConfig, ChartContainer, ChartTooltip, ChartTooltipContent } from "@/components/ui/chart"; import { Skeleton } from "@/components/ui/skeleton"; import { Tabs, TabsList, TabsTrigger } from "@/components/ui/tabs"; import { @@ -26,18 +25,6 @@ import { RankedModel, } from "./modelInsightsData"; -const PALETTE = [ - "#ec4899", - "#a855f7", - "#f59e0b", - "#3b82f6", - "#10b981", - "#ef4444", - "#14b8a6", - "#84cc16", - "#6366f1", - "#f97316", -]; const SCALES = ["linear", "log"] as const; const GRANULARITIES = ["day", "week"] as const; const GRANULARITY_LABELS: Record = { day: "Daily", week: "Weekly" }; @@ -143,10 +130,6 @@ export default function ModelInsightsView({ accessToken }: { accessToken: string ); } - const chartConfig = Object.fromEntries( - models.map((model, index) => [model, { label: model, color: PALETTE[index % PALETTE.length] }]), - ) satisfies ChartConfig; - return ( @@ -198,38 +181,15 @@ export default function ModelInsightsView({ accessToken }: { accessToken: string - - - - - formatMetric(Number(value), shown)} - /> - - `${label} · Gateway total ${formatMetric(bucketTotals.get(String(label)) ?? 0, shown)}` - } - /> - } - /> - {models.map((model, index) => ( - - ))} - - + formatMetric(value, shown)} + totalLabel="Gateway total" + totalFor={(label) => bucketTotals.get(label) ?? 0} + /> diff --git a/ui/litellm-dashboard/src/app/(dashboard)/playground/components/systemOneUI/SystemOneUI.integration.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/playground/components/systemOneUI/SystemOneUI.integration.test.tsx index 2cc7e704499..0812e40d3e4 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/playground/components/systemOneUI/SystemOneUI.integration.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/playground/components/systemOneUI/SystemOneUI.integration.test.tsx @@ -75,7 +75,7 @@ describe("SystemOneUI integration", () => { render(); screen.getByRole("combobox", { name: "Decision endpoint" }).focus(); await user.keyboard("{ArrowDown}"); - await user.click(await screen.findByRole("option", { name: "Decisions · /v1/decisions" })); + await user.click(await screen.findByRole("option", { name: "System One · /v1/systemone" })); expect(screen.getByRole("note", { name: "Decision endpoint notice" })).toHaveTextContent( "omit model to use the proxy's configured default.", @@ -162,7 +162,7 @@ describe("SystemOneUI integration", () => { render(); screen.getByRole("combobox", { name: "Decision endpoint" }).focus(); await user.keyboard("{ArrowDown}"); - await user.click(await screen.findByRole("option", { name: "Decisions · /v1/decisions" })); + await user.click(await screen.findByRole("option", { name: "System One · /v1/systemone" })); const editor = screen.getByRole("textbox", { name: "System One JSON payload" }); const draft = JSON.stringify({ model: "my-decider", @@ -185,12 +185,12 @@ describe("SystemOneUI integration", () => { screen.getByRole("combobox", { name: "Decision endpoint" }).focus(); await user.keyboard("{ArrowDown}"); - await user.click(await screen.findByRole("option", { name: "Decisions · /v1/decisions" })); + await user.click(await screen.findByRole("option", { name: "System One · /v1/systemone" })); expect(editor).toHaveValue(draft); expect(screen.queryByText("Selected choice")).not.toBeInTheDocument(); await user.click(screen.getByRole("button", { name: "Send" })); expect(await screen.findByText("Selected choice")).toBeInTheDocument(); - expect(mockFetch.mock.calls[1]?.[0]).toMatch(/\/v1\/decisions$/); + expect(mockFetch.mock.calls[1]?.[0]).toMatch(/\/v1\/systemone$/); expect(JSON.parse(mockFetch.mock.calls[1]?.[1]?.body as string)).toEqual(JSON.parse(draft)); }); @@ -199,7 +199,7 @@ describe("SystemOneUI integration", () => { render(); screen.getByRole("combobox", { name: "Decision endpoint" }).focus(); await user.keyboard("{ArrowDown}"); - await user.click(await screen.findByRole("option", { name: "Decisions · /v1/decisions" })); + await user.click(await screen.findByRole("option", { name: "System One · /v1/systemone" })); const payload = { state: "An outage", questions: { urgent: { type: "noul", instructions: "Is this urgent?" } } }; fireEvent.change(screen.getByRole("textbox", { name: "System One JSON payload" }), { target: { value: JSON.stringify(payload) }, @@ -207,7 +207,7 @@ describe("SystemOneUI integration", () => { expect(screen.getByRole("button", { name: "Send" })).toBeEnabled(); await user.click(screen.getByRole("button", { name: "Send" })); expect(await screen.findByText("jev-1.13.0")).toBeInTheDocument(); - expect(mockFetch.mock.calls[0]?.[0]).toMatch(/\/v1\/decisions$/); + expect(mockFetch.mock.calls[0]?.[0]).toMatch(/\/v1\/systemone$/); const body = JSON.parse(mockFetch.mock.calls[0]?.[1]?.body as string); expect(body).toEqual(payload); expect(body).not.toHaveProperty("model"); @@ -227,7 +227,7 @@ describe("SystemOneUI integration", () => { await screen.findByRole("button", { name: "Cancel request" }); screen.getByRole("combobox", { name: "Decision endpoint" }).focus(); await user.keyboard("{ArrowDown}"); - await user.click(await screen.findByRole("option", { name: "Decisions · /v1/decisions" })); + await user.click(await screen.findByRole("option", { name: "System One · /v1/systemone" })); expect(mockFetch.mock.calls[0]?.[1]?.signal?.aborted).toBe(true); expect(screen.getByRole("button", { name: "Send" })).toBeEnabled(); @@ -249,7 +249,7 @@ describe("SystemOneUI integration", () => { await user.click(screen.getByRole("button", { name: "Send" })); expect(await screen.findByText("jev-1.13.0")).toBeInTheDocument(); expect(screen.getByText("Selected choice")).toBeInTheDocument(); - expect(mockFetch.mock.calls[1]?.[0]).toMatch(/\/v1\/decisions$/); + expect(mockFetch.mock.calls[1]?.[0]).toMatch(/\/v1\/systemone$/); expect(mockFetch).toHaveBeenCalledTimes(2); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/playground/components/systemOneUI/SystemOneUI.tsx b/ui/litellm-dashboard/src/app/(dashboard)/playground/components/systemOneUI/SystemOneUI.tsx index 7b312eab081..e092efce161 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/playground/components/systemOneUI/SystemOneUI.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/playground/components/systemOneUI/SystemOneUI.tsx @@ -42,11 +42,11 @@ export default function SystemOneUI({ accessToken, disabledPersonalKeyCreation = const [customApiKey, setCustomApiKey] = useState(""); const [endpoint, setEndpoint] = useState("/typesafe/v1/systemone"); const [payloads, setPayloads] = useState>({ - "/v1/decisions": DECISIONS_EXAMPLE_PAYLOAD, + "/v1/systemone": DECISIONS_EXAMPLE_PAYLOAD, "/typesafe/v1/systemone": EXAMPLE_PAYLOAD, }); const rawPayload = payloads[endpoint]; - const examplePayload = endpoint === "/v1/decisions" ? DECISIONS_EXAMPLE_PAYLOAD : EXAMPLE_PAYLOAD; + const examplePayload = endpoint === "/v1/systemone" ? DECISIONS_EXAMPLE_PAYLOAD : EXAMPLE_PAYLOAD; const activeController = useRef(null); const validation = useMemo(() => validateSystemOnePayload(rawPayload, endpoint), [rawPayload, endpoint]); const effectiveApiKey = apiKeySource === "session" ? accessToken || "" : customApiKey.trim(); @@ -105,7 +105,7 @@ export default function SystemOneUI({ accessToken, disabledPersonalKeyCreation = @@ -182,11 +182,11 @@ export default function SystemOneUI({ accessToken, disabledPersonalKeyCreation = - {endpoint === "/v1/decisions" ? "Decision models · Jev format" : "TypeSafe Jev · System One"} + {endpoint === "/v1/systemone" ? "Decision models · System One" : "TypeSafe Jev · System One"} - {endpoint === "/v1/decisions" - ? "Sends choice, noul, and score questions through /v1/decisions. Replace the example model with a decision model configured on your proxy, or omit model to use the proxy's configured default." + {endpoint === "/v1/systemone" + ? "Sends choice, noul, and score questions through /v1/systemone. Replace the example model with a decision model configured on your proxy, or omit model to use the proxy's configured default." : "Sends requests through /typesafe/v1/systemone and requires TYPESAFE_API_KEY on the proxy."}{" "} Give us feedback on what you want for decision models diff --git a/ui/litellm-dashboard/src/app/(dashboard)/playground/components/systemOneUI/lib/decisions.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/playground/components/systemOneUI/lib/decisions.test.ts index 1ae9b2a130f..90f1b63b09f 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/playground/components/systemOneUI/lib/decisions.test.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/playground/components/systemOneUI/lib/decisions.test.ts @@ -12,7 +12,7 @@ const request = { }, provider_option: { enabled: true }, }; -const validate = (value: unknown) => validateSystemOnePayload(JSON.stringify(value), "/v1/decisions"); +const validate = (value: unknown) => validateSystemOnePayload(JSON.stringify(value), "/v1/systemone"); describe("native decisions validation", () => { it("accepts structured Jev criteria, optional instructions, and provider extensions without dropping fields", () => { diff --git a/ui/litellm-dashboard/src/app/(dashboard)/playground/components/systemOneUI/lib/schemas.ts b/ui/litellm-dashboard/src/app/(dashboard)/playground/components/systemOneUI/lib/schemas.ts index 585a2a36921..9da040256bc 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/playground/components/systemOneUI/lib/schemas.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/playground/components/systemOneUI/lib/schemas.ts @@ -1,6 +1,6 @@ import { z } from "zod"; -export type DecisionEndpoint = "/v1/decisions" | "/typesafe/v1/systemone"; +export type DecisionEndpoint = "/v1/systemone" | "/typesafe/v1/systemone"; const decisionsJson = z.union([z.string(), z.record(z.string(), z.unknown()), z.array(z.unknown())]); const decisionInstructions = decisionsJson.nullish(); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/playground/components/systemOneUI/lib/validatePayload.ts b/ui/litellm-dashboard/src/app/(dashboard)/playground/components/systemOneUI/lib/validatePayload.ts index eac96e67045..ac53f871333 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/playground/components/systemOneUI/lib/validatePayload.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/playground/components/systemOneUI/lib/validatePayload.ts @@ -54,7 +54,7 @@ export function validateSystemOnePayload( return invalid("syntax", `Invalid JSON syntax: ${json.message}`); } - const schema = endpoint === "/v1/decisions" ? decisionsRequestSchema : systemOneRequestSchema; + const schema = endpoint === "/v1/systemone" ? decisionsRequestSchema : systemOneRequestSchema; const result = schema.safeParse(json.value); if (!result.success) { return { diff --git a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/EndpointUsage.tsx b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/EndpointUsage.tsx index 51e451ca770..dcf248cee65 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/EndpointUsage.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/EndpointUsage.tsx @@ -56,10 +56,12 @@ const EndpointUsage: React.FC = ({ userSpendData }) => { }, [userSpendData]); return ( -
+
+
+ + +
- -
); }; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/components/EndpointUsageBarChart.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/components/EndpointUsageBarChart.test.tsx index a9e65b21f4b..0c9407d6ea1 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/components/EndpointUsageBarChart.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/components/EndpointUsageBarChart.test.tsx @@ -30,18 +30,20 @@ describe("EndpointUsageBarChart", () => { renderWithProviders(); expect(screen.getByText("Success vs Failed Requests by Endpoint")).toBeInTheDocument(); - expect(screen.getByText("Successful Requests")).toBeInTheDocument(); - expect(screen.getByText("Failed Requests")).toBeInTheDocument(); + expect(screen.getByText("Successful")).toBeInTheDocument(); + expect(screen.getByText("Failed")).toBeInTheDocument(); + expect(screen.getByText("160")).toBeInTheDocument(); + expect(screen.getByText("7")).toBeInTheDocument(); }); - it("renders stacked green and red bars per endpoint", () => { + it("renders stacked brand-blue and red bars per endpoint", () => { const { container } = renderWithProviders(); expect(container.querySelectorAll(".recharts-bar")).toHaveLength(2); const rectangles = Array.from(container.querySelectorAll("path.recharts-rectangle")); expect(rectangles).toHaveLength(4); const fills = new Set(rectangles.map((rect) => rect.getAttribute("fill"))); - expect(fills).toEqual(new Set(["var(--color-green-500, #22c55e)", "var(--color-red-500, #ef4444)"])); + expect(fills).toEqual(new Set(["#2b3fd6", "#ef4444"])); const xPositions = rectangles.map((rect) => rect.getAttribute("d")?.split(",")[0]); expect(new Set(xPositions).size).toBe(2); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/components/EndpointUsageBarChart.tsx b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/components/EndpointUsageBarChart.tsx index bf9868d77cf..b55ea0513c7 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/components/EndpointUsageBarChart.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/components/EndpointUsageBarChart.tsx @@ -1,53 +1,69 @@ import React from "react"; -import { BarChart, CustomLegend, CustomTooltip } from "@/components/shared/charts"; -import { Card, CardContent, CardHeader, CardTitle } from "@/components/ui/card"; +import { ChartColumnStacked } from "lucide-react"; +import { BarChart, CustomTooltip } from "@/components/shared/charts"; import { MetricWithMetadata } from "@/components/UsagePage/types"; +import { Panel } from "../../overview/Primitives"; interface EndpointUsageBarChartProps { endpointData?: Record; } +const SUCCESS_COLOR = "#2b3fd6"; +const FAILED_COLOR = "#ef4444"; +const CATEGORIES = ["Successful", "Failed"] as const; +const COLORS = [SUCCESS_COLOR, FAILED_COLOR] as const; + +const valueFormatter = (value: number) => value.toLocaleString(); + const EndpointUsageBarChart: React.FC = ({ endpointData }) => { // Transform endpoint data into chart format const chartData = React.useMemo(() => { return Object.entries(endpointData || {}).map(([endpoint, data]) => ({ endpoint, - "metrics.successful_requests": data.metrics.successful_requests, - "metrics.failed_requests": data.metrics.failed_requests, - metrics: { - successful_requests: data.metrics.successful_requests, - failed_requests: data.metrics.failed_requests, - }, + Successful: data.metrics.successful_requests, + Failed: data.metrics.failed_requests, })); }, [endpointData]); - const valueFormatter = (value: number) => value.toLocaleString(); + const totals = React.useMemo( + () => + chartData.reduce( + (acc, row) => ({ Successful: acc.Successful + row.Successful, Failed: acc.Failed + row.Failed }), + { Successful: 0, Failed: 0 }, + ), + [chartData], + ); return ( - - -
- Success vs Failed Requests by Endpoint - + + {CATEGORIES.map((category, i) => ( + + + ))}
-
- - - -
+ } + > + + ); }; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/components/EndpointUsageLineChart.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/components/EndpointUsageLineChart.test.tsx index 81ca90d8cb9..8e26e63d9d9 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/components/EndpointUsageLineChart.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/components/EndpointUsageLineChart.test.tsx @@ -2,6 +2,7 @@ import { screen } from "@testing-library/react"; import { describe, expect, it } from "vitest"; import { renderWithProviders } from "@/../tests/test-utils"; import { DailyData, MetricWithMetadata, SpendMetrics } from "@/components/UsagePage/types"; +import { STACKED_USAGE_PALETTE } from "@/components/shared/charts"; import EndpointUsageLineChart from "./EndpointUsageLineChart"; const spendMetrics = (apiRequests: number): SpendMetrics => ({ @@ -53,13 +54,13 @@ describe("EndpointUsageLineChart", () => { expect(screen.getByText("Endpoint Usage Trends")).toBeInTheDocument(); }); - it("renders one line per endpoint with the tremor palette strokes", () => { + it("renders one line per endpoint with the stacked usage palette strokes", () => { const { container } = renderWithProviders(); const curves = Array.from(container.querySelectorAll("path.recharts-line-curve")); expect(curves).toHaveLength(2); expect(new Set(curves.map((curve) => curve.getAttribute("stroke")))).toEqual( - new Set(["var(--color-blue-500, #3b82f6)", "var(--color-cyan-500, #06b6d4)"]), + new Set([STACKED_USAGE_PALETTE[0], STACKED_USAGE_PALETTE[1]]), ); }); @@ -87,7 +88,7 @@ describe("EndpointUsageLineChart", () => { expect(screen.getAllByText(/^\d,\d{3}$/).length).toBeGreaterThan(0); }); - it("draws smooth natural curves", () => { + it("draws smooth curves that never overshoot below zero (monotone)", () => { const { container } = renderWithProviders(); const path = container.querySelector("path.recharts-line-curve")?.getAttribute("d") ?? ""; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/components/EndpointUsageLineChart.tsx b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/components/EndpointUsageLineChart.tsx index 483a30b1639..cd35034fe1f 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/components/EndpointUsageLineChart.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/components/EndpointUsageLineChart.tsx @@ -1,7 +1,8 @@ import { useMemo } from "react"; -import { LineChart, type ChartColor } from "@/components/shared/charts"; -import { Card, CardContent, CardHeader, CardTitle } from "@/components/ui/card"; +import { ChartLine } from "lucide-react"; +import { LineChart, stackedUsageColor } from "@/components/shared/charts"; import { DailyData } from "@/components/UsagePage/types"; +import { Panel } from "../../overview/Primitives"; interface EndpointUsageLineChartProps { dailyData?: { results: DailyData[] }; @@ -58,41 +59,24 @@ export function EndpointUsageLineChart({ dailyData }: EndpointUsageLineChartProp return keys; }, [chartData]); - // Tremor color palette for multiple lines - const colors: readonly ChartColor[] = [ - "blue", - "cyan", - "indigo", - "violet", - "purple", - "fuchsia", - "pink", - "rose", - "red", - "orange", - ]; + const colors = useMemo(() => categories.map((_, i) => stackedUsageColor(i)), [categories]); return ( - - - Endpoint Usage Trends - - - value.toLocaleString()} - showLegend={true} - showGridLines={true} - yAxisWidth={60} - connectNulls={true} - curveType="natural" - /> - - + + value.toLocaleString()} + showLegend={true} + showGridLines={true} + yAxisWidth={56} + connectNulls={true} + curveType="monotone" + /> + ); } diff --git a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/components/EndpointUsageTable.tsx b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/components/EndpointUsageTable.tsx index 3d19d825c0c..2cccdfb47db 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/components/EndpointUsageTable.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/components/EndpointUsageTable.tsx @@ -1,5 +1,8 @@ import React from "react"; import type { ColumnDef } from "@tanstack/react-table"; +import { Route } from "lucide-react"; +import { cn } from "@/lib/cva.config"; +import { Panel } from "../../overview/Primitives"; import { Meter, MeterIndicator, MeterTrack } from "@/components/shared/Meter"; import { DataTable } from "@/components/shared/DataTable"; import { MoneyCell } from "@/components/shared/table_cells"; @@ -41,7 +44,7 @@ const EndpointUsageTable: React.FC = ({ endpointData }) { header: "Endpoint", accessorKey: "endpoint", - cell: ({ row }) => {row.original.endpoint}, + cell: ({ row }) => {row.original.endpoint}, }, { header: "Successful / Failed", @@ -54,18 +57,20 @@ const EndpointUsageTable: React.FC = ({ endpointData }) const totalPercentage = successPercentage + failurePercentage; return ( -
-
+
+
- 0 ? "bg-destructive" : undefined}> - + 0 ? "h-1 bg-destructive/70" : "h-1"}> +
-
- {record.successful_requests.toLocaleString()} +
+ {record.successful_requests.toLocaleString()} / - {record.failed_requests.toLocaleString()} + 0 ? "text-destructive" : "text-muted-foreground"}> + {record.failed_requests.toLocaleString()} +
); @@ -75,7 +80,7 @@ const EndpointUsageTable: React.FC = ({ endpointData }) header: "Total Request", accessorKey: "api_requests", meta: { numeric: true }, - cell: ({ row }) => row.original.api_requests.toLocaleString(), + cell: ({ row }) => {row.original.api_requests.toLocaleString()}, }, { header: "Success Rate", @@ -86,13 +91,10 @@ const EndpointUsageTable: React.FC = ({ endpointData }) const successRateStr = value.toFixed(2); return ( = 95 - ? "text-success font-medium" - : value >= 80 - ? "text-warning font-medium" - : "text-destructive font-medium" - } + className={cn( + "tabular-nums", + value >= 95 ? "text-foreground" : value >= 80 ? "text-warning" : "text-destructive", + )} > {successRateStr}% @@ -103,7 +105,7 @@ const EndpointUsageTable: React.FC = ({ endpointData }) header: "Total Tokens", accessorKey: "total_tokens", meta: { numeric: true }, - cell: ({ row }) => row.original.total_tokens.toLocaleString(), + cell: ({ row }) => {row.original.total_tokens.toLocaleString()}, }, { header: "Spend", @@ -114,13 +116,23 @@ const EndpointUsageTable: React.FC = ({ endpointData }) ]; return ( - row.key} - noDataMessage="No endpoint usage data" - size="compact" - /> + + {dataSource.length.toLocaleString()} {dataSource.length === 1 ? "endpoint" : "endpoints"} + + } + > + row.key} + noDataMessage="No endpoint usage data" + size="compact" + /> + ); }; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.test.tsx index 23bd772729b..b546b2c0558 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.test.tsx @@ -563,7 +563,8 @@ describe("EntityUsage", () => { expect(spendElements.length).toBeGreaterThan(0); }); - expect(screen.getByText("1,000")).toBeInTheDocument(); // Total Requests + // Scoped to the active Cost tab: the keep-mounted Key Activity tab shows the same totals. + expect(within(screen.getByRole("tabpanel")).getByText("1,000")).toBeInTheDocument(); // Total Requests }); it("should render with team entity type and call team API", async () => { @@ -683,7 +684,7 @@ describe("EntityUsage", () => { fireEvent.click(keyActivityTab); }); - expect(screen.getAllByText("Activity Metrics")[1]).toBeInTheDocument(); + expect(within(screen.getByRole("tabpanel")).getByRole("heading", { name: "Overall Usage" })).toBeInTheDocument(); }); it("loads key pages separately from the aggregate using the current entity scope", async () => { @@ -806,9 +807,10 @@ describe("EntityUsage", () => { }); expect(await screen.findByText("Tag Spend Overview")).toBeInTheDocument(); - expect(await screen.findByText("$0.00")).toBeInTheDocument(); - expect(screen.getByText("Total Spend")).toBeInTheDocument(); - expect(screen.getAllByText("0")[0]).toBeInTheDocument(); + const costTab = screen.getByRole("tabpanel"); + expect(await within(costTab).findByText("$0.00")).toBeInTheDocument(); + expect(within(costTab).getByText("Total Spend")).toBeInTheDocument(); + expect(within(costTab).getAllByText("0")[0]).toBeInTheDocument(); }); it("should display Model Activity tab for non-agent entity types", async () => { @@ -1093,31 +1095,26 @@ describe("EntityUsage", () => { }); }); - it("renders daily spend bars, per-entity bars, and the provider donut with cyan fills and a $ center total", async () => { + it("renders the stacked daily spend chart, the per-entity table, and the provider share bar with a $ total", async () => { const { container } = render(); await waitFor(() => { expect(mockTagDailyActivityCall).toHaveBeenCalled(); }); + // The fixture carries no model breakdown, so the day's spend stacks as a single "Other" segment. await waitFor(() => { - expect(container.querySelectorAll("path.recharts-rectangle")).toHaveLength(2); + expect(container.querySelectorAll("path.recharts-rectangle")).toHaveLength(1); }); + expect(container.querySelector("path.recharts-rectangle")).toHaveAttribute("fill", "#94a3b8"); - const barFills = new Set( - Array.from(container.querySelectorAll("path.recharts-rectangle")).map((rect) => rect.getAttribute("fill")), - ); - expect(barFills).toEqual(new Set(["var(--color-cyan-500, #06b6d4)"])); + expect(screen.getAllByText("Jan 1").length).toBeGreaterThan(0); + expect(screen.getAllByText("Tag 1").length).toBeGreaterThan(0); - expect(screen.getAllByText("2025-01-01").length).toBeGreaterThan(0); - expect(screen.getAllByText("Tag 1").length).toBeGreaterThan(1); - - const sectors = container.querySelectorAll(".recharts-pie-sector path"); - expect(sectors).toHaveLength(1); - expect(sectors[0]).toHaveAttribute("fill", "var(--color-cyan-500, #06b6d4)"); - - const centerLabels = Array.from(container.querySelectorAll("text.fill-foreground")).map((text) => text.textContent); - expect(centerLabels).toContain("$100.50"); + const segments = screen.getAllByTestId("provider-share-segment"); + expect(segments).toHaveLength(1); + expect(segments[0]).toHaveStyle({ backgroundColor: "rgb(236, 72, 153)" }); + expect(screen.getByTestId("provider-spend-total")).toHaveTextContent("$100.50"); }); it("should label the chart with user_email metadata instead of the raw UUID (LIT-3889)", async () => { diff --git a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.tsx b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.tsx index 3f239250967..ab1c678869a 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.tsx @@ -1,32 +1,24 @@ import useTeams from "@/app/(dashboard)/hooks/useTeams"; -import { BarChart, DonutChart } from "@/components/shared/charts"; +import { StackedUsageChart, type StackedUsageScale } from "@/components/shared/charts"; import { DataTable } from "@/components/shared/DataTable"; -import { - getProviderSpend, - getTopAgents, - getTopAPIKeys, - getTopModels, - type ProviderSpendRow, -} from "./entityUsageAggregations"; +import { getProviderSpend, getTopAgents, getTopAPIKeys, getTopModels } from "./entityUsageAggregations"; import { buildCostBreakdownTiles, buildSummaryTiles, hasFlatCost, type SummaryTile } from "./entityUsageSummary"; import { MoneyCell } from "@/components/shared/table_cells"; -import { Card as ShadcnCard, CardContent, CardHeader, CardTitle } from "@/components/ui/card"; import { hasCapability, type Capability } from "@/utils/capabilities"; -import { formatNumberWithCommas } from "@/utils/dataUtils"; import type { DateRangePickerValue } from "@/components/shared/date_picker_types"; -import { ChevronDown, ChevronRight, Info } from "lucide-react"; +import { Bot, Boxes, ChevronDown, ChevronRight, ExternalLink, Info, KeyRound, Layers, Server } from "lucide-react"; import type { ColumnDef } from "@tanstack/react-table"; import { Alert, AlertDescription } from "@/components/shared/Alert"; import { ChartLoader } from "@/components/shared/chart_loader"; import { Tabs, TabsContent, TabsList, TabsTrigger } from "@/components/ui/tabs"; import { Tooltip, TooltipContent, TooltipTrigger } from "@/components/ui/tooltip"; +import { cn } from "@/lib/cva.config"; import React, { type ReactNode, useCallback, useMemo, useState } from "react"; import TeamMultiSelect from "@/components/common_components/team_multi_select"; import UserDropdown from "@/components/common_components/UserDropdown"; import { ActivityMetrics, processActivityData } from "@/components/activity_metrics"; import { UsageExportHeader } from "@/components/EntityUsageExport"; import type { EntityType } from "@/components/EntityUsageExport/types"; -import { Logo } from "@/components/molecules/logo/Logo"; import { useAggregatedDailyActivity } from "../../hooks/useAggregatedDailyActivity"; import { ENTITY_API } from "./entityFetchFns"; import { @@ -35,28 +27,78 @@ import { type DailyActivityRequest, } from "@/components/UsagePage/dailyActivityApi"; import { keyDetailFromResponse, overallUsageMetrics } from "@/components/UsagePage/keyActivityData"; -import { EntityMetricWithMetadata } from "@/components/UsagePage/types"; -import { valueFormatterSpend } from "@/components/UsagePage/utils/value_formatters"; +import type { DailyData, EntityMetricWithMetadata } from "@/components/UsagePage/types"; import EndpointUsage from "../EndpointUsage/EndpointUsage"; import ModelViewToggle, { ModelViewType } from "../ModelViewToggle"; import TopKeyView from "@/components/UsagePage/components/EntityUsage/TopKeyView"; import KeyActivityPanel from "@/components/UsagePage/components/KeyActivityPanel"; +import { BreakdownControls, Leaderboard, useBreakdown, type BreakdownState } from "../overview/BreakdownChart"; +import { + bucketSeries, + bucketTotals, + labelForDate, + dailyTotals, + formatMetricValue, + type Granularity, + type Series, + type UsageMetric, +} from "../overview/overviewData"; +import { Panel, Segmented, Sparkline, Stat } from "../overview/Primitives"; import TopModelView from "./TopModelView"; import TeamUserSpendCard from "./TeamUserSpendCard"; +import { ProviderSpendBreakdown } from "./SpendByProvider"; -interface EntityMetrics { - metrics: { - spend: number; - prompt_tokens: number; - completion_tokens: number; - cache_read_input_tokens: number; - cache_creation_input_tokens: number; - total_tokens: number; - successful_requests: number; - failed_requests: number; - api_requests: number; +const BRAND = "#2b3fd6"; +const FLAT_COST_SERIES = "Flat cost"; +const FLAT_COST_COLOR = "#8b5cf6"; +const METRIC_TITLE: Record = { spend: "Spend", tokens: "Tokens", requests: "Requests" }; +/** Summary tiles that carry a sparkline, keyed by tile title, valued by the dailyTotals field. */ +const TILE_TRENDS: Readonly> = { + "Total Spend": "spend", + "Total Cost": "spend", + "Total Requests": "requests", + "Total Tokens": "tokens", +}; +const GRANULARITY_OPTIONS = [ + { value: "day", label: "Daily" }, + { value: "week", label: "Weekly" }, +] as const satisfies readonly { value: Granularity; label: string }[]; +const SCALE_OPTIONS = [ + { value: "linear", label: "Linear" }, + { value: "log", label: "Log" }, +] as const satisfies readonly { value: StackedUsageScale; label: string }[]; +const QUIET_HEADER = { headerClassName: "font-normal" }; +const FLAT_COST_KEY = "flat_cost"; + +/** Stacks reserved-capacity flat cost on top of the per-model spend, so each bar is the day's full cost. */ +const withFlatCost = (series: Series, results: readonly DailyData[]): Series => { + const flatByDate = new Map(results.map((day) => [day.date, day.metrics.flat_cost ?? 0])); + return { + data: series.data.map((day) => ({ ...day, [FLAT_COST_KEY]: flatByDate.get(day.date) ?? 0 })), + keys: [...series.keys, FLAT_COST_KEY], + labels: [...series.labels, FLAT_COST_SERIES], + colors: [...series.colors, FLAT_COST_COLOR], }; - metadata: Record; +}; + +function ShareBar({ value, max }: { value: number; max: number }) { + return ( +